434 lines
17 KiB
Go
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)
|
|
}
|