diff --git a/backend/internal/repository/migrations_runner.go b/backend/internal/repository/migrations_runner.go index 169a01675..0b786616a 100644 --- a/backend/internal/repository/migrations_runner.go +++ b/backend/internal/repository/migrations_runner.go @@ -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 } diff --git a/backend/internal/repository/migrations_runner_notx_test.go b/backend/internal/repository/migrations_runner_notx_test.go index 55588bb04..8f9882e6d 100644 --- a/backend/internal/repository/migrations_runner_notx_test.go +++ b/backend/internal/repository/migrations_runner_notx_test.go @@ -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) diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index a2b82dfa6..c69f374d4 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -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) diff --git a/backend/internal/repository/usage_log_repo_stats.go b/backend/internal/repository/usage_log_repo_stats.go index 53d60b7f4..9385900a7 100644 --- a/backend/internal/repository/usage_log_repo_stats.go +++ b/backend/internal/repository/usage_log_repo_stats.go @@ -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, "") diff --git a/backend/internal/repository/usage_log_repo_stats_integration_test.go b/backend/internal/repository/usage_log_repo_stats_integration_test.go index e9025b1bd..754ea32a9 100644 --- a/backend/internal/repository/usage_log_repo_stats_integration_test.go +++ b/backend/internal/repository/usage_log_repo_stats_integration_test.go @@ -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) { diff --git a/backend/migrations/226_add_usage_log_effective_model_indexes_notx.sql b/backend/migrations/226_add_usage_log_effective_model_indexes_notx.sql new file mode 100644 index 000000000..7bb6250fb --- /dev/null +++ b/backend/migrations/226_add_usage_log_effective_model_indexes_notx.sql @@ -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 + ); diff --git a/backend/migrations/usage_log_effective_model_index_test.go b/backend/migrations/usage_log_effective_model_index_test.go new file mode 100644 index 000000000..9505889dc --- /dev/null +++ b/backend/migrations/usage_log_effective_model_index_test.go @@ -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") +}