Files
cpa-plugin/internal/repository/sqlite_query.go
T
2026-08-15 22:31:12 +08:00

281 lines
10 KiB
Go

package repository
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"math"
"strings"
"time"
"billing/internal/collection"
)
type usageCursor struct {
Version int `json:"v"`
RequestedAt string `json:"t"`
ID int64 `json:"id"`
Direction string `json:"d"`
Page int `json:"p"`
Signature string `json:"s"`
}
type detailKey struct {
ID int64
RequestedAt string
}
func normalizeUsageQuery(query collection.UsageQuery) (collection.UsageQuery, error) {
if query.Page < 1 {
query.Page = 1
}
if query.PageSize < 1 {
query.PageSize = 100
}
if query.PageSize > 100 {
return query, errors.New("page_size 不能超过 100")
}
query.KeyID = strings.TrimSpace(query.KeyID)
query.Model = strings.TrimSpace(query.Model)
query.Result = strings.ToLower(strings.TrimSpace(query.Result))
query.AuthID = strings.TrimSpace(query.AuthID)
query.Endpoint = endpointKind(query.Endpoint)
query.RequestID = strings.TrimSpace(query.RequestID)
if query.From != nil && query.To != nil && !query.From.Before(*query.To) {
return query, errors.New("from 必须早于 to")
}
if query.Result != "" && query.Result != collection.UsageResultSucceeded && query.Result != collection.UsageResultFailed && query.Result != collection.UsageResultRejected && query.Result != collection.UsageResultCanceled {
return query, errors.New("result 必须是 succeeded、failed、rejected 或 canceled")
}
return query, nil
}
func usageQuerySignature(query collection.UsageQuery) string {
payload := struct {
PageSize int
From string
To string
KeyID string
Model string
Result string
AuthID string
Endpoint string
RequestID string
}{PageSize: query.PageSize, KeyID: query.KeyID, Model: strings.ToLower(query.Model), Result: query.Result, AuthID: query.AuthID, Endpoint: query.Endpoint, RequestID: query.RequestID}
if query.From != nil {
payload.From = formatTime(*query.From)
}
if query.To != nil {
payload.To = formatTime(*query.To)
}
raw, _ := json.Marshal(payload)
sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:])
}
func encodeUsageCursor(cursor usageCursor) string {
raw, _ := json.Marshal(cursor)
return base64.RawURLEncoding.EncodeToString(raw)
}
func decodeUsageCursor(raw, signature string) (usageCursor, error) {
decoded, err := base64.RawURLEncoding.DecodeString(strings.TrimSpace(raw))
if err != nil {
return usageCursor{}, errors.New("cursor 无效")
}
var cursor usageCursor
if err := json.Unmarshal(decoded, &cursor); err != nil || cursor.Version != 1 || cursor.ID < 1 || cursor.RequestedAt == "" || cursor.Page < 1 || (cursor.Direction != "next" && cursor.Direction != "previous") {
return usageCursor{}, errors.New("cursor 无效")
}
if cursor.Signature != signature {
return usageCursor{}, errors.New("cursor 与当前筛选条件不匹配")
}
return cursor, nil
}
func usageWhere(query collection.UsageQuery) (string, []any) {
conditions := []string{"1=1"}
args := make([]any, 0, 8)
if query.From != nil {
conditions = append(conditions, "d.requested_at>=?")
args = append(args, formatTime(*query.From))
}
if query.To != nil {
conditions = append(conditions, "d.requested_at<?")
args = append(args, formatTime(*query.To))
}
if query.KeyID != "" {
conditions = append(conditions, "d.managed_key_id=?")
args = append(args, query.KeyID)
}
if query.Model != "" {
conditions = append(conditions, "d.model=? COLLATE NOCASE")
args = append(args, query.Model)
}
if query.Result != "" {
conditions = append(conditions, "d.result=?")
args = append(args, query.Result)
}
if query.AuthID != "" {
conditions = append(conditions, "d.auth_id=?")
args = append(args, query.AuthID)
}
if query.Endpoint != "" {
conditions = append(conditions, "d.endpoint_kind=?")
args = append(args, query.Endpoint)
}
if query.RequestID != "" {
conditions = append(conditions, "d.request_id=?")
args = append(args, query.RequestID)
}
return strings.Join(conditions, " AND "), args
}
func (r *SQLiteUsageRepository) QueryUsage(ctx context.Context, input collection.UsageQuery) (collection.UsagePage, error) {
query, err := normalizeUsageQuery(input)
if err != nil {
return collection.UsagePage{}, err
}
where, args := usageWhere(query)
var total int64
if err := r.readDB.QueryRowContext(ctx, `SELECT COUNT(*) FROM request_detail_index d WHERE `+where, args...).Scan(&total); err != nil {
return collection.UsagePage{}, fmt.Errorf("统计请求明细: %w", err)
}
totalPages := 0
if total > 0 {
totalPages = int(math.Ceil(float64(total) / float64(query.PageSize)))
}
signature := usageQuerySignature(query)
page := query.Page
order := "d.requested_at DESC, d.id DESC"
limitArgs := append([]any{}, args...)
if query.Cursor != "" {
cursor, cursorErr := decodeUsageCursor(query.Cursor, signature)
if cursorErr != nil {
return collection.UsagePage{}, cursorErr
}
page = cursor.Page
if cursor.Direction == "next" {
where += " AND (d.requested_at<? OR (d.requested_at=? AND d.id<?))"
limitArgs = append(limitArgs, cursor.RequestedAt, cursor.RequestedAt, cursor.ID)
} else {
where += " AND (d.requested_at>? OR (d.requested_at=? AND d.id>?))"
limitArgs = append(limitArgs, cursor.RequestedAt, cursor.RequestedAt, cursor.ID)
order = "d.requested_at ASC, d.id ASC"
}
} else {
limitArgs = append(limitArgs, query.PageSize, (page-1)*query.PageSize)
}
statement := usageDetailSelectSQL + " WHERE " + where + " ORDER BY " + order
if query.Cursor != "" {
statement += " LIMIT ?"
limitArgs = append(limitArgs, query.PageSize)
} else {
statement += " LIMIT ? OFFSET ?"
}
rows, err := r.readDB.QueryContext(ctx, statement, limitArgs...)
if err != nil {
return collection.UsagePage{}, fmt.Errorf("查询请求明细: %w", err)
}
records, keys, err := scanUsageDetailRows(rows)
if err != nil {
return collection.UsagePage{}, err
}
if order == "d.requested_at ASC, d.id ASC" {
reverseRecords(records)
reverseKeys(keys)
}
result := collection.UsagePage{Records: records, Page: page, PageSize: query.PageSize, Total: total, TotalPages: totalPages}
if len(keys) > 0 {
if page > 1 {
result.PreviousCursor = encodeUsageCursor(usageCursor{Version: 1, RequestedAt: keys[0].RequestedAt, ID: keys[0].ID, Direction: "previous", Page: page - 1, Signature: signature})
}
if page < totalPages {
last := keys[len(keys)-1]
result.NextCursor = encodeUsageCursor(usageCursor{Version: 1, RequestedAt: last.RequestedAt, ID: last.ID, Direction: "next", Page: page + 1, Signature: signature})
}
}
return result, nil
}
const usageDetailSelectSQL = `
SELECT d.id, d.requested_at,
d.managed_key_id, d.request_id, COALESCE(u.execution_id, ''),
COALESCE(NULLIF(u.trace_id, ''), r.trace_id, ''), COALESCE(u.api_key, ''),
COALESCE(k.name, u.key_alias, ''), COALESCE(u.auth_id, ''), COALESCE(u.auth_index, ''), COALESCE(u.auth_type, ''),
COALESCE(NULLIF(u.model, ''), NULLIF(r.model, ''), r.requested_model, d.model, ''),
COALESCE(u.reasoning_effort, ''), COALESCE(u.service_tier, ''), COALESCE(u.speed, ''),
CASE WHEN d.result='succeeded' THEN 0 ELSE 1 END,
COALESCE(u.executor_type, ''),
CASE WHEN r.request_id IS NOT NULL THEN CASE WHEN r.stream THEN 'SSE' ELSE 'JSON' END ELSE COALESCE(u.request_type, '') END,
COALESCE(NULLIF(u.endpoint, ''), r.endpoint, ''),
COALESCE(u.input_tokens, 0), COALESCE(u.output_tokens, 0), COALESCE(u.reasoning_tokens, 0), COALESCE(u.cached_tokens, 0),
COALESCE(u.cache_read_tokens, 0), COALESCE(u.cache_write_tokens, 0), COALESCE(u.total_tokens, 0),
COALESCE(u.ttft_ns, 0), COALESCE(u.latency_ns, r.latency_ns, 0), u.cost_micros,
COALESCE(u.price_tier, ''), COALESCE(u.fast_requested, 0), COALESCE(u.fast_pricing_applied, 0),
COALESCE(u.price_multiplier_numerator, 1), COALESCE(u.price_multiplier_denominator, 1), COALESCE(u.client_ip, ''),
COALESCE(r.outcome, ''), COALESCE(r.status_code, 0), COALESCE(r.error, '')
FROM request_detail_index d
LEFT JOIN usage_records u ON u.id=d.usage_id
LEFT JOIN request_records r ON r.request_id=d.lifecycle_request_id
LEFT JOIN managed_keys k ON k.id=d.managed_key_id`
func scanUsageDetailRows(rows *sql.Rows) ([]collection.Record, []detailKey, error) {
defer rows.Close()
var records []collection.Record
var keys []detailKey
for rows.Next() {
var record collection.Record
var key detailKey
var failed bool
var ttftNS, latencyNS int64
var cost sql.NullInt64
if err := rows.Scan(&key.ID, &key.RequestedAt, &record.ManagedKeyID, &record.RequestID, &record.ExecutionID,
&record.TraceID, &record.APIKey, &record.KeyAlias, &record.AuthID, &record.AuthIndex, &record.AuthType,
&record.Model, &record.ReasoningEffort, &record.ServiceTier, &record.Speed, &failed, &record.ExecutorType,
&record.RequestType, &record.Endpoint, &record.InputTokens, &record.OutputTokens, &record.ReasoningTokens,
&record.CachedTokens, &record.CacheReadTokens, &record.CacheWriteTokens, &record.TotalTokens, &ttftNS,
&latencyNS, &cost, &record.PriceTier, &record.FastRequested, &record.FastPricingApplied,
&record.PriceMultiplierNumerator, &record.PriceMultiplierDenominator, &record.ClientIP,
&record.Outcome, &record.StatusCode, &record.Error); err != nil {
return nil, nil, fmt.Errorf("读取请求明细: %w", err)
}
requestedAt, err := time.Parse(time.RFC3339Nano, key.RequestedAt)
if err != nil {
return nil, nil, fmt.Errorf("解析请求明细时间: %w", err)
}
record.RequestedAt = requestedAt
record.Failed = failed
record.TTFT = time.Duration(ttftNS)
record.Latency = time.Duration(latencyNS)
if cost.Valid {
value := cost.Int64
record.CostMicros = &value
}
records = append(records, record)
keys = append(keys, key)
}
if err := rows.Err(); err != nil {
return nil, nil, fmt.Errorf("遍历请求明细: %w", err)
}
return records, keys, nil
}
func reverseRecords(values []collection.Record) {
for left, right := 0, len(values)-1; left < right; left, right = left+1, right-1 {
values[left], values[right] = values[right], values[left]
}
}
func reverseKeys(values []detailKey) {
for left, right := 0, len(values)-1; left < right; left, right = left+1, right-1 {
values[left], values[right] = values[right], values[left]
}
}