281 lines
10 KiB
Go
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]
|
|
}
|
|
}
|