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:
Wesley Liddick
2026-08-19 20:27:42 +08:00
committed by GitHub
7 changed files with 273 additions and 223 deletions
@@ -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")
}