Files
cpa-plugin/internal/repository/sqlite_access.go
T

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