421 lines
16 KiB
Go
421 lines
16 KiB
Go
package repository
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
managedaccess "cpa-ext/internal/access"
|
|
"cpa-ext/internal/collection"
|
|
)
|
|
|
|
var ErrManagedKeyNotFound = errors.New("managed key not found")
|
|
|
|
func (r *SQLiteUsageRepository) BootstrapManagedKey(ctx context.Context, name, secret string) (managedaccess.ManagedKey, error) {
|
|
var count int
|
|
if err := r.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM managed_keys`).Scan(&count); err != nil {
|
|
return managedaccess.ManagedKey{}, fmt.Errorf("查询 Key 数量: %w", err)
|
|
}
|
|
if count > 0 {
|
|
return managedaccess.ManagedKey{}, nil
|
|
}
|
|
now := time.Now().UTC()
|
|
key := managedaccess.ManagedKey{
|
|
ID: "key_default", Name: strings.TrimSpace(name), Secret: secret,
|
|
Status: managedaccess.StatusActive, RouteMode: managedaccess.RouteAuto,
|
|
AllModels: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
if err := r.CreateManagedKey(ctx, key); err != nil {
|
|
return managedaccess.ManagedKey{}, err
|
|
}
|
|
if _, err := r.db.ExecContext(ctx, `
|
|
UPDATE usage_records SET managed_key_id = ?, key_alias = ?
|
|
WHERE managed_key_id = '' AND api_key = ?`, key.ID, key.Name, key.Secret); err != nil {
|
|
return managedaccess.ManagedKey{}, fmt.Errorf("归档 default 历史用量: %w", err)
|
|
}
|
|
return key, nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) CreateManagedKey(ctx context.Context, key managedaccess.ManagedKey) error {
|
|
if key.CreatedAt.IsZero() {
|
|
key.CreatedAt = time.Now().UTC()
|
|
}
|
|
if key.UpdatedAt.IsZero() {
|
|
key.UpdatedAt = key.CreatedAt
|
|
}
|
|
tx, err := r.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return fmt.Errorf("开始创建 Key: %w", err)
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
if _, err = tx.ExecContext(ctx, `
|
|
INSERT INTO managed_keys (id, name, secret, status, route_mode, upstream_account_id, all_models, created_at, updated_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, key.ID, key.Name, key.Secret, key.Status, key.RouteMode,
|
|
key.UpstreamAccountID, key.AllModels, formatTime(key.CreatedAt), formatTime(key.UpdatedAt)); err != nil {
|
|
return fmt.Errorf("创建 Key: %w", err)
|
|
}
|
|
if err := replaceManagedKeyModels(ctx, tx, key.ID, key.Models); err != nil {
|
|
return err
|
|
}
|
|
if _, err = tx.ExecContext(ctx, `
|
|
UPDATE usage_records SET managed_key_id = ?, key_alias = ?
|
|
WHERE managed_key_id = '' AND api_key = ?`, key.ID, key.Name, key.Secret); err != nil {
|
|
return fmt.Errorf("关联历史用量: %w", err)
|
|
}
|
|
if err = tx.Commit(); err != nil {
|
|
return fmt.Errorf("提交创建 Key: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) UpdateManagedKey(ctx context.Context, key managedaccess.ManagedKey) error {
|
|
current, err := r.ManagedKeyByID(ctx, key.ID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if current.Status == managedaccess.StatusArchived {
|
|
return errors.New("已归档 Key 不可修改")
|
|
}
|
|
key.Secret = current.Secret
|
|
key.CreatedAt = current.CreatedAt
|
|
key.UpdatedAt = time.Now().UTC()
|
|
tx, err := r.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return fmt.Errorf("开始更新 Key: %w", err)
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
result, err := tx.ExecContext(ctx, `
|
|
UPDATE managed_keys SET name=?, status=?, route_mode=?, upstream_account_id=?, all_models=?, updated_at=?
|
|
WHERE id=? AND status <> 'archived'`, key.Name, key.Status, key.RouteMode, key.UpstreamAccountID,
|
|
key.AllModels, formatTime(key.UpdatedAt), key.ID)
|
|
if err != nil {
|
|
return fmt.Errorf("更新 Key: %w", err)
|
|
}
|
|
if affected, _ := result.RowsAffected(); affected == 0 {
|
|
return ErrManagedKeyNotFound
|
|
}
|
|
if err := replaceManagedKeyModels(ctx, tx, key.ID, key.Models); err != nil {
|
|
return err
|
|
}
|
|
if err = tx.Commit(); err != nil {
|
|
return fmt.Errorf("提交更新 Key: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ArchiveManagedKey(ctx context.Context, id string) error {
|
|
result, err := r.db.ExecContext(ctx, `
|
|
UPDATE managed_keys SET status='archived', updated_at=? WHERE id=? AND status <> 'archived'`, formatTime(time.Now().UTC()), id)
|
|
if err != nil {
|
|
return fmt.Errorf("归档 Key: %w", err)
|
|
}
|
|
if affected, _ := result.RowsAffected(); affected == 0 {
|
|
return ErrManagedKeyNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func replaceManagedKeyModels(ctx context.Context, tx *sql.Tx, keyID string, models []string) error {
|
|
if _, err := tx.ExecContext(ctx, `DELETE FROM managed_key_models WHERE managed_key_id=?`, keyID); err != nil {
|
|
return fmt.Errorf("清理 Key 模型: %w", err)
|
|
}
|
|
seen := make(map[string]struct{})
|
|
for _, raw := range models {
|
|
model := strings.ToLower(strings.TrimSpace(raw))
|
|
if model == "" {
|
|
continue
|
|
}
|
|
if _, exists := seen[model]; exists {
|
|
continue
|
|
}
|
|
seen[model] = struct{}{}
|
|
if _, err := tx.ExecContext(ctx, `INSERT INTO managed_key_models (managed_key_id, model) VALUES (?, ?)`, keyID, model); err != nil {
|
|
return fmt.Errorf("保存 Key 模型: %w", err)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ManagedKeyByCredential(ctx context.Context, credential string) (managedaccess.ManagedKey, error) {
|
|
return r.scanManagedKey(ctx, `WHERE id=? OR secret=? LIMIT 1`, credential, credential)
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ManagedKeyByID(ctx context.Context, id string) (managedaccess.ManagedKey, error) {
|
|
return r.scanManagedKey(ctx, `WHERE id=? LIMIT 1`, id)
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ManagedKeyByReference(ctx context.Context, reference string) (managedaccess.ManagedKey, error) {
|
|
if key, err := r.ManagedKeyByID(ctx, reference); err == nil {
|
|
return key, nil
|
|
}
|
|
keys, err := r.ListManagedKeys(ctx, true)
|
|
if err != nil {
|
|
return managedaccess.ManagedKey{}, err
|
|
}
|
|
for _, key := range keys {
|
|
if managedaccess.CallerScope(key.ID) == reference {
|
|
return key, nil
|
|
}
|
|
}
|
|
return managedaccess.ManagedKey{}, ErrManagedKeyNotFound
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ManagedKeyByName(ctx context.Context, name string) (managedaccess.ManagedKey, error) {
|
|
return r.scanManagedKey(ctx, `WHERE name=? COLLATE NOCASE LIMIT 1`, name)
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) scanManagedKey(ctx context.Context, where string, args ...any) (managedaccess.ManagedKey, error) {
|
|
var key managedaccess.ManagedKey
|
|
var createdAt, updatedAt string
|
|
err := r.db.QueryRowContext(ctx, `
|
|
SELECT id, name, secret, status, route_mode, upstream_account_id, all_models, created_at, updated_at
|
|
FROM managed_keys `+where, args...).Scan(&key.ID, &key.Name, &key.Secret, &key.Status, &key.RouteMode,
|
|
&key.UpstreamAccountID, &key.AllModels, &createdAt, &updatedAt)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return managedaccess.ManagedKey{}, ErrManagedKeyNotFound
|
|
}
|
|
if err != nil {
|
|
return managedaccess.ManagedKey{}, fmt.Errorf("查询 Key: %w", err)
|
|
}
|
|
key.CreatedAt, _ = time.Parse(time.RFC3339Nano, createdAt)
|
|
key.UpdatedAt, _ = time.Parse(time.RFC3339Nano, updatedAt)
|
|
key.Models, err = r.managedKeyModels(ctx, key.ID)
|
|
return key, err
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ListManagedKeys(ctx context.Context, includeArchived bool) ([]managedaccess.ManagedKey, error) {
|
|
where := `WHERE status <> 'archived'`
|
|
if includeArchived {
|
|
where = ""
|
|
}
|
|
rows, err := r.db.QueryContext(ctx, `
|
|
SELECT id, name, secret, status, route_mode, upstream_account_id, all_models, created_at, updated_at
|
|
FROM managed_keys `+where+` ORDER BY name COLLATE NOCASE`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询 Key 列表: %w", err)
|
|
}
|
|
var keys []managedaccess.ManagedKey
|
|
for rows.Next() {
|
|
var key managedaccess.ManagedKey
|
|
var createdAt, updatedAt string
|
|
if err := rows.Scan(&key.ID, &key.Name, &key.Secret, &key.Status, &key.RouteMode,
|
|
&key.UpstreamAccountID, &key.AllModels, &createdAt, &updatedAt); err != nil {
|
|
return nil, fmt.Errorf("读取 Key: %w", err)
|
|
}
|
|
key.CreatedAt, _ = time.Parse(time.RFC3339Nano, createdAt)
|
|
key.UpdatedAt, _ = time.Parse(time.RFC3339Nano, updatedAt)
|
|
keys = append(keys, key)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
_ = rows.Close()
|
|
return nil, fmt.Errorf("遍历 Key: %w", err)
|
|
}
|
|
if err := rows.Close(); err != nil {
|
|
return nil, fmt.Errorf("关闭 Key 查询: %w", err)
|
|
}
|
|
for index := range keys {
|
|
keys[index].Models, err = r.managedKeyModels(ctx, keys[index].ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
return keys, nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) managedKeyModels(ctx context.Context, keyID string) ([]string, error) {
|
|
rows, err := r.db.QueryContext(ctx, `SELECT model FROM managed_key_models WHERE managed_key_id=? ORDER BY model`, keyID)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询 Key 模型: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
models := make([]string, 0)
|
|
for rows.Next() {
|
|
var model string
|
|
if err := rows.Scan(&model); err != nil {
|
|
return nil, fmt.Errorf("读取 Key 模型: %w", err)
|
|
}
|
|
models = append(models, model)
|
|
}
|
|
return models, rows.Err()
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) SyncUpstreamAccounts(ctx context.Context, accounts []managedaccess.UpstreamAccount) error {
|
|
tx, err := r.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return fmt.Errorf("开始同步上游账号: %w", err)
|
|
}
|
|
defer func() { _ = tx.Rollback() }()
|
|
for _, account := range accounts {
|
|
var existingID string
|
|
err := tx.QueryRowContext(ctx, `
|
|
SELECT id FROM upstream_accounts WHERE cpa_auth_id=? OR (cpa_auth_index <> '' AND cpa_auth_index=?) LIMIT 1`,
|
|
account.CPAAuthID, account.CPAAuthIndex).Scan(&existingID)
|
|
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
|
return fmt.Errorf("匹配上游账号: %w", err)
|
|
}
|
|
if existingID == "" {
|
|
sum := sha256.Sum256([]byte(account.Provider + "\x00" + account.CPAAuthIndex + "\x00" + account.CPAAuthID))
|
|
existingID = fmt.Sprintf("up_%x", sum[:10])
|
|
}
|
|
_, err = tx.ExecContext(ctx, `
|
|
INSERT INTO upstream_accounts (id, cpa_auth_id, cpa_auth_index, provider, display_name, status, status_message, disabled, unavailable, priority, last_seen_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(id) DO UPDATE SET cpa_auth_id=excluded.cpa_auth_id, provider=excluded.provider,
|
|
cpa_auth_index=CASE WHEN excluded.cpa_auth_index <> '' THEN excluded.cpa_auth_index ELSE upstream_accounts.cpa_auth_index END,
|
|
display_name=CASE WHEN excluded.display_name = excluded.cpa_auth_id AND upstream_accounts.display_name <> '' THEN upstream_accounts.display_name ELSE excluded.display_name END,
|
|
status=excluded.status,
|
|
status_message=excluded.status_message, disabled=excluded.disabled, unavailable=excluded.unavailable,
|
|
priority=excluded.priority, last_seen_at=excluded.last_seen_at`, existingID, account.CPAAuthID,
|
|
account.CPAAuthIndex, account.Provider, account.DisplayName, account.Status, account.StatusMessage,
|
|
account.Disabled, account.Unavailable, account.Priority, formatTime(account.LastSeenAt))
|
|
if err != nil {
|
|
return fmt.Errorf("保存上游账号: %w", err)
|
|
}
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
return fmt.Errorf("提交上游账号同步: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ListUpstreamAccounts(ctx context.Context) ([]managedaccess.UpstreamAccount, error) {
|
|
rows, err := r.db.QueryContext(ctx, `
|
|
SELECT id, cpa_auth_id, cpa_auth_index, provider, display_name, status, status_message,
|
|
disabled, unavailable, priority, last_seen_at
|
|
FROM upstream_accounts ORDER BY provider, display_name COLLATE NOCASE`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询上游账号: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
var accounts []managedaccess.UpstreamAccount
|
|
for rows.Next() {
|
|
var account managedaccess.UpstreamAccount
|
|
var lastSeen string
|
|
if err := rows.Scan(&account.ID, &account.CPAAuthID, &account.CPAAuthIndex, &account.Provider,
|
|
&account.DisplayName, &account.Status, &account.StatusMessage, &account.Disabled,
|
|
&account.Unavailable, &account.Priority, &lastSeen); err != nil {
|
|
return nil, fmt.Errorf("读取上游账号: %w", err)
|
|
}
|
|
account.LastSeenAt, _ = time.Parse(time.RFC3339Nano, lastSeen)
|
|
accounts = append(accounts, account)
|
|
}
|
|
return accounts, rows.Err()
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) UpstreamAccountByID(ctx context.Context, id string) (managedaccess.UpstreamAccount, error) {
|
|
var account managedaccess.UpstreamAccount
|
|
var lastSeen string
|
|
err := r.db.QueryRowContext(ctx, `
|
|
SELECT id, cpa_auth_id, cpa_auth_index, provider, display_name, status, status_message,
|
|
disabled, unavailable, priority, last_seen_at FROM upstream_accounts WHERE id=?`, id).Scan(
|
|
&account.ID, &account.CPAAuthID, &account.CPAAuthIndex, &account.Provider, &account.DisplayName,
|
|
&account.Status, &account.StatusMessage, &account.Disabled, &account.Unavailable, &account.Priority, &lastSeen)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return managedaccess.UpstreamAccount{}, errors.New("upstream account not found")
|
|
}
|
|
if err != nil {
|
|
return managedaccess.UpstreamAccount{}, fmt.Errorf("查询上游账号: %w", err)
|
|
}
|
|
account.LastSeenAt, _ = time.Parse(time.RFC3339Nano, lastSeen)
|
|
return account, nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) KeyStats(ctx context.Context, keyID string, today time.Time) (managedaccess.KeyStats, error) {
|
|
stats := managedaccess.KeyStats{KeyID: keyID}
|
|
facts, err := r.keyUsageFacts(ctx, keyID)
|
|
if err != nil {
|
|
return stats, err
|
|
}
|
|
for _, fact := range mergeOrphanLifecycleRecords(facts) {
|
|
addUsageSummary(&stats.Total, fact)
|
|
if !fact.RequestedAt.Before(today) {
|
|
addUsageSummary(&stats.Today, fact)
|
|
}
|
|
if stats.LastUsedAt == nil || fact.RequestedAt.After(*stats.LastUsedAt) {
|
|
value := fact.RequestedAt
|
|
stats.LastUsedAt = &value
|
|
}
|
|
}
|
|
return stats, nil
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) keyUsageFacts(ctx context.Context, keyID string) ([]collection.Record, error) {
|
|
rows, err := r.db.QueryContext(ctx, `
|
|
SELECT request_id, requested_at, model, api_key, executor_type,
|
|
input_tokens, output_tokens, total_tokens, cost_micros, ''
|
|
FROM usage_records WHERE managed_key_id=?
|
|
UNION ALL
|
|
SELECT request_id, requested_at, COALESCE(NULLIF(model, ''), requested_model), '', '',
|
|
0, 0, 0, NULL, outcome
|
|
FROM request_records r
|
|
WHERE managed_key_id=? AND NOT EXISTS (SELECT 1 FROM usage_records u WHERE u.request_id=r.request_id)
|
|
ORDER BY requested_at DESC`, keyID, keyID)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询 Key 用量事实: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
facts := make([]collection.Record, 0)
|
|
for rows.Next() {
|
|
var fact collection.Record
|
|
var requestedAt string
|
|
var cost sql.NullInt64
|
|
if err := rows.Scan(&fact.RequestID, &requestedAt, &fact.Model, &fact.APIKey, &fact.ExecutorType,
|
|
&fact.InputTokens, &fact.OutputTokens, &fact.TotalTokens, &cost, &fact.Outcome); err != nil {
|
|
return nil, fmt.Errorf("读取 Key 用量事实: %w", err)
|
|
}
|
|
fact.ManagedKeyID = keyID
|
|
fact.RequestedAt, err = time.Parse(time.RFC3339Nano, requestedAt)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("解析 Key 用量时间: %w", err)
|
|
}
|
|
if cost.Valid {
|
|
value := cost.Int64
|
|
fact.CostMicros = &value
|
|
}
|
|
facts = append(facts, fact)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("遍历 Key 用量事实: %w", err)
|
|
}
|
|
return facts, nil
|
|
}
|
|
|
|
func addUsageSummary(summary *managedaccess.UsageSummary, fact collection.Record) {
|
|
summary.Requests++
|
|
summary.InputTokens += fact.InputTokens
|
|
summary.OutputTokens += fact.OutputTokens
|
|
summary.TotalTokens += fact.TotalTokens
|
|
if fact.CostMicros != nil {
|
|
summary.CostMicros += *fact.CostMicros
|
|
}
|
|
}
|
|
|
|
func (r *SQLiteUsageRepository) ModelSuggestions(ctx context.Context) ([]string, error) {
|
|
rows, err := r.db.QueryContext(ctx, `
|
|
SELECT model FROM model_prices
|
|
UNION SELECT model FROM usage_records WHERE model <> ''
|
|
UNION SELECT requested_model FROM request_records WHERE requested_model <> ''
|
|
ORDER BY 1 COLLATE NOCASE`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询模型建议: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
var models []string
|
|
for rows.Next() {
|
|
var model string
|
|
if err := rows.Scan(&model); err != nil {
|
|
return nil, err
|
|
}
|
|
models = append(models, model)
|
|
}
|
|
return models, rows.Err()
|
|
}
|
|
|
|
func formatTime(value time.Time) string {
|
|
return value.UTC().Format(time.RFC3339Nano)
|
|
}
|