perf(usage): aggregate stats in one scan

Reduce filtered admin usage statistics from four scans to one GROUPING SETS query so every breakdown shares the exact same filters. Add concurrent expression indexes for requested and upstream model filters on large usage_logs tables.
This commit is contained in:
jaxxjj
2026-08-18 14:26:19 +08:00
parent 938f1868ae
commit a9514a68d2
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")
}