Merge pull request #5762 from jaxxjj/codex/perf-usage-stats-grouping-sets
perf(usage): aggregate admin stats in one scan
This commit is contained in:
@@ -59,6 +59,9 @@ const latestAPIKeyIPIndexMigration = "174_add_usage_logs_api_key_latest_ip_index
|
||||
const latestAPIKeyIPIndex = "idx_usage_logs_api_key_latest_ip"
|
||||
const usageLogsUpstreamModelMismatchIndexMigration = "195_add_usage_log_upstream_model_mismatch_index_notx.sql"
|
||||
const usageLogsUpstreamModelMismatchIndex = "idx_usage_logs_upstream_model_mismatch_created_at"
|
||||
const usageLogsEffectiveModelIndexesMigration = "226_add_usage_log_effective_model_indexes_notx.sql"
|
||||
const usageLogsEffectiveRequestedModelIndex = "idx_usage_logs_effective_requested_model_created"
|
||||
const usageLogsEffectiveUpstreamModelIndex = "idx_usage_logs_effective_upstream_model_created"
|
||||
|
||||
type migrationChecksumCompatibilityRule struct {
|
||||
fileChecksum string
|
||||
@@ -295,6 +298,13 @@ func prepareNonTransactionalMigration(ctx context.Context, db migrationConnectio
|
||||
return dropInvalidIndexIfPresent(ctx, db, latestAPIKeyIPIndex)
|
||||
case usageLogsUpstreamModelMismatchIndexMigration:
|
||||
return dropInvalidIndexIfPresent(ctx, db, usageLogsUpstreamModelMismatchIndex)
|
||||
case usageLogsEffectiveModelIndexesMigration:
|
||||
for _, indexName := range []string{usageLogsEffectiveRequestedModelIndex, usageLogsEffectiveUpstreamModelIndex} {
|
||||
if err := dropInvalidIndexIfPresent(ctx, db, indexName); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -191,6 +191,47 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_upstream_model_mismatch_c
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestApplyMigrationsFS_NonTransactionalMigration_EffectiveModelIndexesDropInvalidIndexesBeforeRetry(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
prepareMigrationsBootstrapExpectations(mock)
|
||||
mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1").
|
||||
WithArgs(usageLogsEffectiveModelIndexesMigration).
|
||||
WillReturnError(sql.ErrNoRows)
|
||||
for _, indexName := range []string{usageLogsEffectiveRequestedModelIndex, usageLogsEffectiveUpstreamModelIndex} {
|
||||
mock.ExpectQuery("SELECT EXISTS \\(").
|
||||
WithArgs(indexName).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true))
|
||||
mock.ExpectExec("DROP INDEX CONCURRENTLY IF EXISTS " + indexName).
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
}
|
||||
mock.ExpectExec("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_requested_model_created").
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectExec("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_upstream_model_created").
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectExec("INSERT INTO schema_migrations \\(filename, checksum\\) VALUES \\(\\$1, \\$2\\)").
|
||||
WithArgs(usageLogsEffectiveModelIndexesMigration, sqlmock.AnyArg()).
|
||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
||||
mock.ExpectExec("SELECT pg_advisory_unlock\\(\\$1\\)").
|
||||
WithArgs(migrationsAdvisoryLockID).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
fsys := fstest.MapFS{
|
||||
usageLogsEffectiveModelIndexesMigration: &fstest.MapFile{Data: []byte(`
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_requested_model_created
|
||||
ON usage_logs ((COALESCE(NULLIF(BTRIM(requested_model), ''), model)), created_at DESC, id DESC);
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_upstream_model_created
|
||||
ON usage_logs ((COALESCE(NULLIF(BTRIM(upstream_model), ''), model)), created_at DESC, id DESC);
|
||||
`)},
|
||||
}
|
||||
|
||||
err = applyMigrationsFS(context.Background(), db, fsys)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestApplyMigrationsFS_PaymentOrdersOutTradeNoUniqueMigration_FailsFastOnDuplicatePrecheck(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -505,33 +505,33 @@ func TestUsageLogRepositoryGetStatsWithFiltersRequestedModelSource(t *testing.T)
|
||||
ModelFilterSource: usagestats.ModelSourceRequested,
|
||||
}
|
||||
|
||||
mock.ExpectQuery("FROM usage_logs\\s+WHERE COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$1").
|
||||
mock.ExpectQuery("(?s)FROM usage_logs\\s+WHERE COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$1.*GROUP BY GROUPING SETS").
|
||||
WithArgs("gpt-5").
|
||||
WillReturnRows(sqlmock.NewRows([]string{
|
||||
"total_requests",
|
||||
"total_input_tokens",
|
||||
"total_output_tokens",
|
||||
"total_cache_tokens",
|
||||
"total_cache_creation_tokens",
|
||||
"total_cache_read_tokens",
|
||||
"total_cost",
|
||||
"total_actual_cost",
|
||||
"total_account_cost",
|
||||
"inbound_grouped",
|
||||
"upstream_grouped",
|
||||
"inbound_endpoint",
|
||||
"upstream_endpoint",
|
||||
"requests",
|
||||
"input_tokens",
|
||||
"output_tokens",
|
||||
"cache_creation_tokens",
|
||||
"cache_read_tokens",
|
||||
"cost",
|
||||
"actual_cost",
|
||||
"account_cost",
|
||||
"avg_duration_ms",
|
||||
}).AddRow(int64(1), int64(2), int64(3), int64(4), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(inbound_endpoint\\), ''\\), 'unknown'\\) AS endpoint").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(upstream_endpoint\\), ''\\), 'unknown'\\) AS endpoint").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT CONCAT\\(").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
}).
|
||||
AddRow(1, 1, nil, nil, int64(1), int64(2), int64(3), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0).
|
||||
AddRow(0, 1, "/v1/responses", nil, int64(1), int64(2), int64(3), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0).
|
||||
AddRow(1, 0, nil, "/v1/responses", int64(1), int64(2), int64(3), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0).
|
||||
AddRow(0, 0, "/v1/responses", "/v1/responses", int64(1), int64(2), int64(3), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0))
|
||||
|
||||
stats, err := repo.GetStatsWithFilters(context.Background(), filters)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), stats.TotalRequests)
|
||||
require.Equal(t, "/v1/responses", stats.Endpoints[0].Endpoint)
|
||||
require.Equal(t, "/v1/responses -> /v1/responses", stats.EndpointPaths[0].Endpoint)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
@@ -546,29 +546,23 @@ func TestUsageLogRepositoryGetStatsWithFiltersRequestTypePriority(t *testing.T)
|
||||
Stream: &stream,
|
||||
}
|
||||
|
||||
mock.ExpectQuery("FROM usage_logs\\s+WHERE \\(request_type = \\$1 OR \\(request_type = 0 AND stream = FALSE AND openai_ws_mode = FALSE\\)\\)").
|
||||
mock.ExpectQuery("(?s)FROM usage_logs\\s+WHERE \\(request_type = \\$1 OR \\(request_type = 0 AND stream = FALSE AND openai_ws_mode = FALSE\\)\\).*GROUP BY GROUPING SETS").
|
||||
WithArgs(requestType).
|
||||
WillReturnRows(sqlmock.NewRows([]string{
|
||||
"total_requests",
|
||||
"total_input_tokens",
|
||||
"total_output_tokens",
|
||||
"total_cache_tokens",
|
||||
"total_cache_creation_tokens",
|
||||
"total_cache_read_tokens",
|
||||
"total_cost",
|
||||
"total_actual_cost",
|
||||
"total_account_cost",
|
||||
"inbound_grouped",
|
||||
"upstream_grouped",
|
||||
"inbound_endpoint",
|
||||
"upstream_endpoint",
|
||||
"requests",
|
||||
"input_tokens",
|
||||
"output_tokens",
|
||||
"cache_creation_tokens",
|
||||
"cache_read_tokens",
|
||||
"cost",
|
||||
"actual_cost",
|
||||
"account_cost",
|
||||
"avg_duration_ms",
|
||||
}).AddRow(int64(1), int64(2), int64(3), int64(4), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(inbound_endpoint\\), ''\\), 'unknown'\\) AS endpoint").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(upstream_endpoint\\), ''\\), 'unknown'\\) AS endpoint").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT CONCAT\\(").
|
||||
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
}).AddRow(1, 1, nil, nil, int64(1), int64(2), int64(3), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0))
|
||||
|
||||
stats, err := repo.GetStatsWithFilters(context.Background(), filters)
|
||||
require.NoError(t, err)
|
||||
@@ -689,19 +683,12 @@ func TestUsageLogRepositoryGetStatsWithFiltersAlwaysReturnsAccountCost(t *testin
|
||||
// No AccountID filter set - TotalAccountCost should still be returned
|
||||
filters := usagestats.UsageLogFilters{}
|
||||
|
||||
mock.ExpectQuery("FROM usage_logs").
|
||||
mock.ExpectQuery("(?s)FROM usage_logs.*GROUP BY GROUPING SETS").
|
||||
WillReturnRows(sqlmock.NewRows([]string{
|
||||
"total_requests", "total_input_tokens", "total_output_tokens",
|
||||
"total_cache_tokens", "total_cache_creation_tokens", "total_cache_read_tokens",
|
||||
"total_cost", "total_actual_cost",
|
||||
"total_account_cost", "avg_duration_ms",
|
||||
}).AddRow(int64(50), int64(1000), int64(2000), int64(100), int64(60), int64(40), 15.0, 12.5, 11.0, 100.0))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(inbound_endpoint\\)").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(upstream_endpoint\\)").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
mock.ExpectQuery("SELECT CONCAT\\(").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"}))
|
||||
"inbound_grouped", "upstream_grouped", "inbound_endpoint", "upstream_endpoint",
|
||||
"requests", "input_tokens", "output_tokens", "cache_creation_tokens", "cache_read_tokens",
|
||||
"cost", "actual_cost", "account_cost", "avg_duration_ms",
|
||||
}).AddRow(1, 1, nil, nil, int64(50), int64(1000), int64(2000), int64(60), int64(40), 15.0, 12.5, 11.0, 100.0))
|
||||
|
||||
stats, err := repo.GetStatsWithFilters(context.Background(), filters)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -3,9 +3,9 @@ package repository
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/lib/pq"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
// GetUserStatsAggregated returns aggregated usage statistics for a user using database-level aggregation
|
||||
@@ -696,108 +695,131 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(`
|
||||
WITH scoped AS (
|
||||
SELECT
|
||||
COALESCE(NULLIF(TRIM(inbound_endpoint), ''), 'unknown') AS inbound_endpoint,
|
||||
COALESCE(NULLIF(TRIM(upstream_endpoint), ''), 'unknown') AS upstream_endpoint,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_creation_tokens,
|
||||
cache_read_tokens,
|
||||
total_cost,
|
||||
actual_cost,
|
||||
COALESCE(account_stats_cost, total_cost) * COALESCE(account_rate_multiplier, 1) AS account_cost,
|
||||
duration_ms
|
||||
FROM usage_logs
|
||||
%s
|
||||
)
|
||||
SELECT
|
||||
COUNT(*) as total_requests,
|
||||
COALESCE(SUM(input_tokens), 0) as total_input_tokens,
|
||||
COALESCE(SUM(output_tokens), 0) as total_output_tokens,
|
||||
COALESCE(SUM(cache_creation_tokens + cache_read_tokens), 0) as total_cache_tokens,
|
||||
COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens,
|
||||
COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens,
|
||||
COALESCE(SUM(total_cost), 0) as total_cost,
|
||||
COALESCE(SUM(actual_cost), 0) as total_actual_cost,
|
||||
COALESCE(SUM(COALESCE(account_stats_cost, total_cost) * COALESCE(account_rate_multiplier, 1)), 0) as total_account_cost,
|
||||
COALESCE(AVG(duration_ms), 0) as avg_duration_ms
|
||||
FROM usage_logs
|
||||
%s
|
||||
GROUPING(inbound_endpoint) AS inbound_grouped,
|
||||
GROUPING(upstream_endpoint) AS upstream_grouped,
|
||||
inbound_endpoint,
|
||||
upstream_endpoint,
|
||||
COUNT(*) AS requests,
|
||||
COALESCE(SUM(input_tokens), 0) AS input_tokens,
|
||||
COALESCE(SUM(output_tokens), 0) AS output_tokens,
|
||||
COALESCE(SUM(cache_creation_tokens), 0) AS cache_creation_tokens,
|
||||
COALESCE(SUM(cache_read_tokens), 0) AS cache_read_tokens,
|
||||
COALESCE(SUM(total_cost), 0) AS cost,
|
||||
COALESCE(SUM(actual_cost), 0) AS actual_cost,
|
||||
COALESCE(SUM(account_cost), 0) AS account_cost,
|
||||
COALESCE(AVG(duration_ms), 0) AS avg_duration_ms
|
||||
FROM scoped
|
||||
GROUP BY GROUPING SETS (
|
||||
(),
|
||||
(inbound_endpoint),
|
||||
(upstream_endpoint),
|
||||
(inbound_endpoint, upstream_endpoint)
|
||||
)
|
||||
`, buildWhere(conditions))
|
||||
|
||||
stats := &UsageStats{}
|
||||
var totalAccountCost float64
|
||||
|
||||
start := time.Unix(0, 0).UTC()
|
||||
if filters.StartTime != nil {
|
||||
start = *filters.StartTime
|
||||
}
|
||||
end := time.Now().UTC()
|
||||
if filters.EndTime != nil {
|
||||
end = *filters.EndTime
|
||||
useAccountCostForEndpoint := filters.AccountID > 0 && filters.UserID == 0 && filters.APIKeyID == 0
|
||||
rows, err := r.sql.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var endpoints, upstreamEndpoints, endpointPaths []EndpointStat
|
||||
|
||||
// 汇总查询:失败即致命。
|
||||
runSummary := func(c context.Context) error {
|
||||
return scanSingleRow(
|
||||
c, r.sql, query, args,
|
||||
&stats.TotalRequests,
|
||||
&stats.TotalInputTokens,
|
||||
&stats.TotalOutputTokens,
|
||||
&stats.TotalCacheTokens,
|
||||
&stats.TotalCacheCreationTokens,
|
||||
&stats.TotalCacheReadTokens,
|
||||
&stats.TotalCost,
|
||||
&stats.TotalActualCost,
|
||||
&totalAccountCost,
|
||||
&stats.AverageDurationMs,
|
||||
for rows.Next() {
|
||||
var (
|
||||
inboundGrouped, upstreamGrouped int
|
||||
inboundEndpoint, upstreamEndpoint sql.NullString
|
||||
requests, inputTokens, outputTokens, cacheCreationTokens, cacheReads int64
|
||||
cost, actualCost, accountCost, averageDurationMs float64
|
||||
)
|
||||
}
|
||||
// endpoint 明细:best-effort(失败 log + 返空),不致命。
|
||||
runEndpoints := func(c context.Context) {
|
||||
res, err := r.getEndpointStatsByColumnWithFilters(c, "inbound_endpoint", start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode)
|
||||
if err != nil {
|
||||
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
|
||||
logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
|
||||
}
|
||||
res = []EndpointStat{}
|
||||
if err := rows.Scan(
|
||||
&inboundGrouped,
|
||||
&upstreamGrouped,
|
||||
&inboundEndpoint,
|
||||
&upstreamEndpoint,
|
||||
&requests,
|
||||
&inputTokens,
|
||||
&outputTokens,
|
||||
&cacheCreationTokens,
|
||||
&cacheReads,
|
||||
&cost,
|
||||
&actualCost,
|
||||
&accountCost,
|
||||
&averageDurationMs,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
endpoints = res
|
||||
}
|
||||
runUpstream := func(c context.Context) {
|
||||
res, err := r.getEndpointStatsByColumnWithFilters(c, "upstream_endpoint", start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode)
|
||||
if err != nil {
|
||||
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
|
||||
logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
|
||||
}
|
||||
res = []EndpointStat{}
|
||||
|
||||
totalTokens := inputTokens + outputTokens + cacheCreationTokens + cacheReads
|
||||
endpointActualCost := actualCost
|
||||
if useAccountCostForEndpoint {
|
||||
endpointActualCost = accountCost
|
||||
}
|
||||
upstreamEndpoints = res
|
||||
}
|
||||
runPaths := func(c context.Context) {
|
||||
res, err := r.getEndpointPathStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode)
|
||||
if err != nil {
|
||||
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
|
||||
logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", err)
|
||||
}
|
||||
res = []EndpointStat{}
|
||||
|
||||
switch {
|
||||
case inboundGrouped == 1 && upstreamGrouped == 1:
|
||||
stats.TotalRequests = requests
|
||||
stats.TotalInputTokens = inputTokens
|
||||
stats.TotalOutputTokens = outputTokens
|
||||
stats.TotalCacheCreationTokens = cacheCreationTokens
|
||||
stats.TotalCacheReadTokens = cacheReads
|
||||
stats.TotalCacheTokens = cacheCreationTokens + cacheReads
|
||||
stats.TotalCost = cost
|
||||
stats.TotalActualCost = actualCost
|
||||
totalAccountCost = accountCost
|
||||
stats.AverageDurationMs = averageDurationMs
|
||||
case inboundGrouped == 0 && upstreamGrouped == 1:
|
||||
stats.Endpoints = append(stats.Endpoints, EndpointStat{
|
||||
Endpoint: inboundEndpoint.String, Requests: requests, TotalTokens: totalTokens,
|
||||
Cost: cost, ActualCost: endpointActualCost,
|
||||
})
|
||||
case inboundGrouped == 1 && upstreamGrouped == 0:
|
||||
stats.UpstreamEndpoints = append(stats.UpstreamEndpoints, EndpointStat{
|
||||
Endpoint: upstreamEndpoint.String, Requests: requests, TotalTokens: totalTokens,
|
||||
Cost: cost, ActualCost: endpointActualCost,
|
||||
})
|
||||
case inboundGrouped == 0 && upstreamGrouped == 0:
|
||||
stats.EndpointPaths = append(stats.EndpointPaths, EndpointStat{
|
||||
Endpoint: inboundEndpoint.String + " -> " + upstreamEndpoint.String,
|
||||
Requests: requests, TotalTokens: totalTokens, Cost: cost, ActualCost: endpointActualCost,
|
||||
})
|
||||
}
|
||||
endpointPaths = res
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if r.db != nil {
|
||||
// 生产路径:r.sql 是 *sql.DB 连接池,可并发。4 条查询并行,延迟取最大值。
|
||||
g, gctx := errgroup.WithContext(ctx)
|
||||
g.Go(func() error { return runSummary(gctx) })
|
||||
g.Go(func() error { runEndpoints(gctx); return nil })
|
||||
g.Go(func() error { runUpstream(gctx); return nil })
|
||||
g.Go(func() error { runPaths(gctx); return nil })
|
||||
if err := g.Wait(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
} else {
|
||||
// 事务路径(ent.Tx 不能并发查询):顺序执行,行为与重构前一致。
|
||||
if err := runSummary(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
runEndpoints(ctx)
|
||||
runUpstream(ctx)
|
||||
runPaths(ctx)
|
||||
sortEndpointStats := func(values []EndpointStat) {
|
||||
sort.Slice(values, func(i, j int) bool {
|
||||
if values[i].Requests != values[j].Requests {
|
||||
return values[i].Requests > values[j].Requests
|
||||
}
|
||||
return values[i].Endpoint < values[j].Endpoint
|
||||
})
|
||||
}
|
||||
sortEndpointStats(stats.Endpoints)
|
||||
sortEndpointStats(stats.UpstreamEndpoints)
|
||||
sortEndpointStats(stats.EndpointPaths)
|
||||
|
||||
stats.TotalAccountCost = &totalAccountCost
|
||||
stats.TotalTokens = stats.TotalInputTokens + stats.TotalOutputTokens + stats.TotalCacheTokens
|
||||
stats.Endpoints = endpoints
|
||||
stats.UpstreamEndpoints = upstreamEndpoints
|
||||
stats.EndpointPaths = endpointPaths
|
||||
|
||||
return stats, nil
|
||||
}
|
||||
@@ -882,78 +904,6 @@ func (r *usageLogRepository) getEndpointStatsByColumnWithFilters(ctx context.Con
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (r *usageLogRepository) getEndpointPathStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []EndpointStat, err error) {
|
||||
actualCostExpr := "COALESCE(SUM(actual_cost), 0) as actual_cost"
|
||||
if accountID > 0 && userID == 0 && apiKeyID == 0 {
|
||||
actualCostExpr = "COALESCE(SUM(COALESCE(account_stats_cost, total_cost) * COALESCE(account_rate_multiplier, 1)), 0) as actual_cost"
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(`
|
||||
SELECT
|
||||
CONCAT(
|
||||
COALESCE(NULLIF(TRIM(inbound_endpoint), ''), 'unknown'),
|
||||
' -> ',
|
||||
COALESCE(NULLIF(TRIM(upstream_endpoint), ''), 'unknown')
|
||||
) AS endpoint,
|
||||
COUNT(*) AS requests,
|
||||
COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS total_tokens,
|
||||
COALESCE(SUM(total_cost), 0) as cost,
|
||||
%s
|
||||
FROM usage_logs
|
||||
WHERE created_at >= $1 AND created_at < $2
|
||||
`, actualCostExpr)
|
||||
|
||||
args := []any{startTime, endTime}
|
||||
if userID > 0 {
|
||||
query += fmt.Sprintf(" AND user_id = $%d", len(args)+1)
|
||||
args = append(args, userID)
|
||||
}
|
||||
if apiKeyID > 0 {
|
||||
query += fmt.Sprintf(" AND api_key_id = $%d", len(args)+1)
|
||||
args = append(args, apiKeyID)
|
||||
}
|
||||
if accountID > 0 {
|
||||
query += fmt.Sprintf(" AND account_id = $%d", len(args)+1)
|
||||
args = append(args, accountID)
|
||||
}
|
||||
if groupID > 0 {
|
||||
query += fmt.Sprintf(" AND group_id = $%d", len(args)+1)
|
||||
args = append(args, groupID)
|
||||
}
|
||||
query, args = appendUsageLogModelQueryFilter(query, args, model, modelSource)
|
||||
query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream)
|
||||
if billingType != nil {
|
||||
query += fmt.Sprintf(" AND billing_type = $%d", len(args)+1)
|
||||
args = append(args, int16(*billingType))
|
||||
}
|
||||
query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "")
|
||||
query += " GROUP BY endpoint ORDER BY requests DESC"
|
||||
|
||||
rows, err := r.sql.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() {
|
||||
if closeErr := rows.Close(); closeErr != nil && err == nil {
|
||||
err = closeErr
|
||||
results = nil
|
||||
}
|
||||
}()
|
||||
|
||||
results = make([]EndpointStat, 0)
|
||||
for rows.Next() {
|
||||
var row EndpointStat
|
||||
if err := rows.Scan(&row.Endpoint, &row.Requests, &row.TotalTokens, &row.Cost, &row.ActualCost); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, row)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// GetEndpointStatsWithFilters returns inbound endpoint statistics with optional filters.
|
||||
func (r *usageLogRepository) GetEndpointStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) ([]EndpointStat, error) {
|
||||
return r.getEndpointStatsByColumnWithFilters(ctx, "inbound_endpoint", startTime, endTime, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "")
|
||||
|
||||
@@ -43,6 +43,15 @@ func TestUsageLog_UpstreamModelMismatchFilterAndPartialIndex(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), stats.TotalRequests)
|
||||
require.Equal(t, []usagestats.EndpointStat{{
|
||||
Endpoint: "unknown", Requests: 1, TotalTokens: 2,
|
||||
}}, stats.Endpoints)
|
||||
require.Equal(t, []usagestats.EndpointStat{{
|
||||
Endpoint: "unknown", Requests: 1, TotalTokens: 2,
|
||||
}}, stats.UpstreamEndpoints)
|
||||
require.Equal(t, []usagestats.EndpointStat{{
|
||||
Endpoint: "unknown -> unknown", Requests: 1, TotalTokens: 2,
|
||||
}}, stats.EndpointPaths)
|
||||
|
||||
trend, err := repo.GetUsageTrendWithUsageFilters(ctx, start, end, "hour", usagestats.UsageLogFilters{
|
||||
UserID: user.ID, UpstreamModelMismatch: &trueValue,
|
||||
@@ -53,24 +62,45 @@ func TestUsageLog_UpstreamModelMismatchFilterAndPartialIndex(t *testing.T) {
|
||||
|
||||
_, err = tx.ExecContext(ctx, "SET LOCAL enable_seqscan = off")
|
||||
require.NoError(t, err)
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
assertPlanUsesIndex := func(query, indexName string, args ...any) {
|
||||
rows, queryErr := tx.QueryContext(ctx, query, args...)
|
||||
require.NoError(t, queryErr)
|
||||
var planLines []string
|
||||
for rows.Next() {
|
||||
var line string
|
||||
require.NoError(t, rows.Scan(&line))
|
||||
planLines = append(planLines, line)
|
||||
}
|
||||
require.NoError(t, rows.Err())
|
||||
require.NoError(t, rows.Close())
|
||||
require.Contains(t, strings.Join(planLines, "\n"), indexName)
|
||||
}
|
||||
assertPlanUsesIndex(`
|
||||
EXPLAIN (COSTS OFF)
|
||||
SELECT id
|
||||
FROM usage_logs
|
||||
WHERE upstream_model_mismatch IS TRUE
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT 100
|
||||
`)
|
||||
require.NoError(t, err)
|
||||
defer func() { require.NoError(t, rows.Close()) }()
|
||||
var planLines []string
|
||||
for rows.Next() {
|
||||
var line string
|
||||
require.NoError(t, rows.Scan(&line))
|
||||
planLines = append(planLines, line)
|
||||
}
|
||||
require.NoError(t, rows.Err())
|
||||
require.Contains(t, strings.Join(planLines, "\n"), usageLogsUpstreamModelMismatchIndex)
|
||||
`, usageLogsUpstreamModelMismatchIndex)
|
||||
assertPlanUsesIndex(`
|
||||
EXPLAIN (COSTS OFF)
|
||||
SELECT id
|
||||
FROM usage_logs
|
||||
WHERE COALESCE(NULLIF(TRIM(requested_model), ''), model) = $1
|
||||
AND created_at >= $2 AND created_at < $3
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT 100
|
||||
`, usageLogsEffectiveRequestedModelIndex, "gpt-5.5", start, end)
|
||||
assertPlanUsesIndex(`
|
||||
EXPLAIN (COSTS OFF)
|
||||
SELECT id
|
||||
FROM usage_logs
|
||||
WHERE COALESCE(NULLIF(TRIM(upstream_model), ''), model) = $1
|
||||
AND created_at >= $2 AND created_at < $3
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT 100
|
||||
`, usageLogsEffectiveUpstreamModelIndex, "gpt-5.5", start, end)
|
||||
}
|
||||
|
||||
func TestUsageLog_GetStatsWithFilters_AggregatesAndEndpoints(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_requested_model_created
|
||||
ON usage_logs (
|
||||
(COALESCE(NULLIF(BTRIM(requested_model), ''), model)),
|
||||
created_at DESC,
|
||||
id DESC
|
||||
);
|
||||
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_upstream_model_created
|
||||
ON usage_logs (
|
||||
(COALESCE(NULLIF(BTRIM(upstream_model), ''), model)),
|
||||
created_at DESC,
|
||||
id DESC
|
||||
);
|
||||
@@ -0,0 +1,19 @@
|
||||
package migrations
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestUsageLogEffectiveModelIndexesMigration(t *testing.T) {
|
||||
content, err := FS.ReadFile("226_add_usage_log_effective_model_indexes_notx.sql")
|
||||
require.NoError(t, err)
|
||||
|
||||
sql := strings.Join(strings.Fields(string(content)), " ")
|
||||
require.Contains(t, sql, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_requested_model_created")
|
||||
require.Contains(t, sql, "(COALESCE(NULLIF(BTRIM(requested_model), ''), model)), created_at DESC, id DESC")
|
||||
require.Contains(t, sql, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_effective_upstream_model_created")
|
||||
require.Contains(t, sql, "(COALESCE(NULLIF(BTRIM(upstream_model), ''), model)), created_at DESC, id DESC")
|
||||
}
|
||||
Reference in New Issue
Block a user