Files
cpa-plugin/internal/repository/sqlite_usage_test.go
T
chuan 6a461bf1f3 feat: 添加持久化用量计费与管理面板
- 使用 SQLite 保存关联后的请求生命周期和用量记录
- 支持长上下文阶梯价格和可配置的 Fast 计费倍率
- 添加管理接口和可配置列的用量面板
- 保留失败、取消、重试和 compact 请求,便于计费核对
2026-08-14 21:08:02 +08:00

95 lines
2.9 KiB
Go

package repository_test
import (
"context"
"fmt"
"path/filepath"
"testing"
"time"
"cpa-ext/internal/collection"
"cpa-ext/internal/repository"
)
func TestSQLiteUsagePersistsAllRecordsAndLimitsQueries(t *testing.T) {
databasePath := filepath.Join(t.TempDir(), "usage.db")
store, err := repository.OpenSQLiteUsage(databasePath)
if err != nil {
t.Fatal(err)
}
requestedAt := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC)
for index := 1; index <= 1002; index++ {
record := collection.Record{
RequestedAt: requestedAt.Add(time.Duration(index) * time.Second),
Model: fmt.Sprintf("model-%d", index),
TotalTokens: int64(index),
}
if err := store.Insert(context.Background(), record); err != nil {
t.Fatalf("insert record %d: %v", index, err)
}
}
recent, err := store.ListRecent(context.Background(), 1000)
if err != nil {
t.Fatal(err)
}
if len(recent) != 1000 || recent[0].Model != "model-1002" || recent[999].Model != "model-3" {
t.Fatalf("unexpected recent records: len=%d first=%q last=%q", len(recent), recent[0].Model, recent[len(recent)-1].Model)
}
if err := store.Close(); err != nil {
t.Fatal(err)
}
reopened, err := repository.OpenSQLiteUsage(databasePath)
if err != nil {
t.Fatal(err)
}
defer reopened.Close()
all, err := reopened.ListRecent(context.Background(), 2000)
if err != nil {
t.Fatal(err)
}
if len(all) != 1002 {
t.Fatalf("persisted records = %d, want 1002", len(all))
}
}
func TestSQLiteUsageRoundTripsNullableCostAndDurations(t *testing.T) {
store, err := repository.OpenSQLiteUsage(filepath.Join(t.TempDir(), "usage.db"))
if err != nil {
t.Fatal(err)
}
defer store.Close()
cost := int64(125_000)
want := collection.Record{
RequestedAt: time.Date(2026, 8, 14, 12, 0, 0, 123, time.FixedZone("CST", 8*60*60)),
APIKey: "key",
Model: "model",
Failed: true,
RequestType: "SSE",
Endpoint: "POST /v1/responses",
InputTokens: 10,
CacheReadTokens: 8,
TTFT: 250 * time.Millisecond,
Latency: 2 * time.Second,
CostMicros: &cost,
PriceTier: "base",
PriceMultiplierNumerator: 1,
PriceMultiplierDenominator: 1,
ClientIP: "192.0.2.10",
}
if err := store.Insert(context.Background(), want); err != nil {
t.Fatal(err)
}
records, err := store.ListRecent(context.Background(), 1)
if err != nil {
t.Fatal(err)
}
got := records[0]
if !got.RequestedAt.Equal(want.RequestedAt) || got.APIKey != want.APIKey || got.RequestType != want.RequestType || got.Endpoint != want.Endpoint || got.ClientIP != want.ClientIP || got.TTFT != want.TTFT || got.Latency != want.Latency || got.CostMicros == nil || *got.CostMicros != cost || got.PriceTier != "base" {
t.Fatalf("round trip mismatch: %+v", got)
}
}