Files

434 lines
17 KiB
Go

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)
}