package repository import ( "context" "crypto/sha256" "database/sql" "errors" "fmt" "strings" "time" managedaccess "billing/internal/access" ) 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.readDB.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, ShowInStats: 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, show_in_stats, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, key.ID, key.Name, key.Secret, key.Status, key.RouteMode, key.UpstreamAccountID, key.AllModels, key.ShowInStats, 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, ` INSERT INTO billing_accounts (managed_key_id, quota_micros, reset_period, max_concurrency, current_cycle_sequence, updated_at) VALUES (?, 0, 'none', 4, 1, ?)`, key.ID, formatTime(key.CreatedAt)); err != nil { return fmt.Errorf("创建 Key 额度账户: %w", err) } if _, err = tx.ExecContext(ctx, ` INSERT INTO billing_cycles (managed_key_id, sequence, started_at, quota_micros, spent_micros) VALUES (?, 1, ?, 0, 0)`, key.ID, formatTime(key.CreatedAt)); err != nil { return fmt.Errorf("创建 Key 额度周期: %w", 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) } event := "管理员创建用户 " + quotedName(key.Name) + " 的 Key" if key.ID == "key_default" { event = "系统初始化用户 " + quotedName(key.Name) + " 的 Key" } if err := insertBusinessEvent(ctx, tx, event, BusinessEventSucceeded, key.CreatedAt); err != nil { return 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=?, show_in_stats=?, updated_at=? WHERE id=? AND status <> 'archived'`, key.Name, key.Status, key.RouteMode, key.UpstreamAccountID, key.AllModels, key.ShowInStats, 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 } event := "管理员修改用户 " + quotedName(key.Name) + " 的 Key" if current.Name != key.Name { event = "管理员修改用户名称:" + quotedName(current.Name) + " → " + quotedName(key.Name) } if err := insertBusinessEvent(ctx, tx, event, BusinessEventSucceeded, key.UpdatedAt); 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 { now := time.Now().UTC() tx, err := r.db.BeginTx(ctx, nil) if err != nil { return fmt.Errorf("开始归档 Key: %w", err) } defer func() { _ = tx.Rollback() }() var name string if err := tx.QueryRowContext(ctx, `SELECT name FROM managed_keys WHERE id=? AND status <> 'archived'`, id).Scan(&name); err != nil { if errors.Is(err, sql.ErrNoRows) { return ErrManagedKeyNotFound } return fmt.Errorf("查询归档 Key: %w", err) } result, err := tx.ExecContext(ctx, ` UPDATE managed_keys SET status='archived', updated_at=? WHERE id=? AND status <> 'archived'`, formatTime(now), id) if err != nil { return fmt.Errorf("归档 Key: %w", err) } if affected, _ := result.RowsAffected(); affected == 0 { return ErrManagedKeyNotFound } if err := insertBusinessEvent(ctx, tx, "管理员归档用户 "+quotedName(name)+" 的 Key", BusinessEventSucceeded, now); err != nil { return err } if err := tx.Commit(); err != nil { return fmt.Errorf("提交归档 Key: %w", err) } 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 secret=? LIMIT 1`, 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) scanManagedKey(ctx context.Context, where string, args ...any) (managedaccess.ManagedKey, error) { var key managedaccess.ManagedKey var createdAt, updatedAt string err := r.readDB.QueryRowContext(ctx, ` SELECT id, name, secret, status, route_mode, upstream_account_id, all_models, show_in_stats, created_at, updated_at FROM managed_keys `+where, args...).Scan(&key.ID, &key.Name, &key.Secret, &key.Status, &key.RouteMode, &key.UpstreamAccountID, &key.AllModels, &key.ShowInStats, &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.readDB.QueryContext(ctx, ` SELECT id, name, secret, status, route_mode, upstream_account_id, all_models, show_in_stats, 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, &key.ShowInStats, &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.readDB.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 { return r.syncUpstreamAccounts(ctx, accounts, false) } // ReconcileUpstreamAccounts treats accounts as the complete current CPA auth snapshot. // Historical rows remain available for existing strict bindings, but are excluded from // new route selection once CPA no longer reports them. func (r *SQLiteUsageRepository) ReconcileUpstreamAccounts(ctx context.Context, accounts []managedaccess.UpstreamAccount) error { return r.syncUpstreamAccounts(ctx, accounts, true) } func (r *SQLiteUsageRepository) syncUpstreamAccounts(ctx context.Context, accounts []managedaccess.UpstreamAccount, snapshot bool) error { tx, err := r.db.BeginTx(ctx, nil) if err != nil { return fmt.Errorf("开始同步上游账号: %w", err) } defer func() { _ = tx.Rollback() }() if snapshot { if _, err := tx.ExecContext(ctx, `UPDATE upstream_accounts SET current=0`); err != nil { return fmt.Errorf("重置上游账号快照: %w", err) } } 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, current, priority, last_seen_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 1, ?, ?) 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, current=1, 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.readDB.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 WHERE current=1 ORDER BY provider, display_name COLLATE NOCASE, cpa_auth_index`) 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.readDB.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} var last sql.NullString err := r.readDB.QueryRowContext(ctx, ` SELECT COUNT(*), COALESCE(SUM(u.input_tokens),0), COALESCE(SUM(u.output_tokens),0), COALESCE(SUM(u.total_tokens),0), COALESCE(SUM(u.cost_micros),0), COALESCE(SUM(CASE WHEN d.requested_at>=? THEN 1 ELSE 0 END),0), COALESCE(SUM(CASE WHEN d.requested_at>=? THEN u.input_tokens ELSE 0 END),0), COALESCE(SUM(CASE WHEN d.requested_at>=? THEN u.output_tokens ELSE 0 END),0), COALESCE(SUM(CASE WHEN d.requested_at>=? THEN u.total_tokens ELSE 0 END),0), COALESCE(SUM(CASE WHEN d.requested_at>=? THEN u.cost_micros ELSE 0 END),0), MAX(d.requested_at) FROM request_detail_index d LEFT JOIN usage_records u ON u.id=d.usage_id WHERE d.managed_key_id=?`, formatTime(today), formatTime(today), formatTime(today), formatTime(today), formatTime(today), keyID).Scan( &stats.Total.Requests, &stats.Total.InputTokens, &stats.Total.OutputTokens, &stats.Total.TotalTokens, &stats.Total.CostMicros, &stats.Today.Requests, &stats.Today.InputTokens, &stats.Today.OutputTokens, &stats.Today.TotalTokens, &stats.Today.CostMicros, &last) if err != nil { return stats, fmt.Errorf("汇总 Key 用量: %w", err) } if last.Valid { if parsed, parseErr := time.Parse(time.RFC3339Nano, last.String); parseErr == nil { stats.LastUsedAt = &parsed } } return stats, nil } func (r *SQLiteUsageRepository) ModelSuggestions(ctx context.Context) ([]string, error) { rows, err := r.readDB.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) }