From c637e6cf318590f4284df3be63dbb6c554c37ce6 Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Sun, 15 Mar 2026 22:13:12 +0800 Subject: [PATCH 01/15] fix: use half-open date ranges for DST-safe usage queries Replace t.Add(24*time.Hour - time.Nanosecond) with t.AddDate(0, 0, 1) and use SQL < instead of <= for end-of-day boundaries. This avoids edge-case misses around DST transitions. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- backend/internal/handler/admin/usage_handler.go | 7 ++++--- backend/internal/handler/usage_handler.go | 8 ++++---- backend/internal/repository/usage_log_repo.go | 2 +- 3 files changed, 9 insertions(+), 8 deletions(-) diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go index 05fd00f1f5..7a3135b88a 100644 --- a/backend/internal/handler/admin/usage_handler.go +++ b/backend/internal/handler/admin/usage_handler.go @@ -159,8 +159,8 @@ func (h *UsageHandler) List(c *gin.Context) { response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") return } - // Set end time to end of day - t = t.Add(24*time.Hour - time.Nanosecond) + // Use half-open range [start, end), move to next calendar day start (DST-safe). + t = t.AddDate(0, 0, 1) endTime = &t } @@ -285,7 +285,8 @@ func (h *UsageHandler) Stats(c *gin.Context) { response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") return } - endTime = endTime.Add(24*time.Hour - time.Nanosecond) + // 与 SQL 条件 created_at < end 对齐,使用次日 00:00 作为上边界(DST-safe)。 + endTime = endTime.AddDate(0, 0, 1) } else { period := c.DefaultQuery("period", "today") switch period { diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go index 2bd0e0d7b5..483f51059b 100644 --- a/backend/internal/handler/usage_handler.go +++ b/backend/internal/handler/usage_handler.go @@ -114,8 +114,8 @@ func (h *UsageHandler) List(c *gin.Context) { response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") return } - // Set end time to end of day - t = t.Add(24*time.Hour - time.Nanosecond) + // Use half-open range [start, end), move to next calendar day start (DST-safe). + t = t.AddDate(0, 0, 1) endTime = &t } @@ -227,8 +227,8 @@ func (h *UsageHandler) Stats(c *gin.Context) { response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") return } - // 设置结束时间为当天结束 - endTime = endTime.Add(24*time.Hour - time.Nanosecond) + // 与 SQL 条件 created_at < end 对齐,使用次日 00:00 作为上边界(DST-safe)。 + endTime = endTime.AddDate(0, 0, 1) } else { // 使用 period 参数 period := c.DefaultQuery("period", "today") diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index cc949db2da..a1fab45be6 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -3004,7 +3004,7 @@ func (r *usageLogRepository) GetGlobalStats(ctx context.Context, startTime, endT COALESCE(SUM(actual_cost), 0) as total_actual_cost, COALESCE(AVG(duration_ms), 0) as avg_duration_ms FROM usage_logs - WHERE created_at >= $1 AND created_at <= $2 + WHERE created_at >= $1 AND created_at < $2 ` stats := &UsageStats{} From 1b79b0f3ffdb1776459b320eb19849fd33b28ac7 Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Sun, 15 Mar 2026 22:13:22 +0800 Subject: [PATCH 02/15] feat: add InboundEndpoint/UpstreamEndpoint fields to non-OpenAI usage records Extend RecordUsageInput and RecordUsageLongContextInput structs with InboundEndpoint and UpstreamEndpoint so that Claude, Gemini, and Sora handlers can record endpoint info alongside OpenAI handlers. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- backend/internal/service/gateway_service.go | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index cff9e9bb73..0b50162a14 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -7130,6 +7130,8 @@ type RecordUsageInput struct { User *User Account *Account Subscription *UserSubscription // 可选:订阅信息 + InboundEndpoint string // 入站端点(客户端请求路径) + UpstreamEndpoint string // 上游端点(标准化后的上游路径) UserAgent string // 请求的 User-Agent IPAddress string // 请求的客户端 IP 地址 RequestPayloadHash string // 请求体语义哈希,用于降低 request_id 误复用时的静默误去重风险 @@ -7528,6 +7530,8 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu RequestID: requestID, Model: result.Model, ReasoningEffort: result.ReasoningEffort, + InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), + UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), InputTokens: result.Usage.InputTokens, OutputTokens: result.Usage.OutputTokens, CacheCreationTokens: result.Usage.CacheCreationInputTokens, @@ -7608,6 +7612,8 @@ type RecordUsageLongContextInput struct { User *User Account *Account Subscription *UserSubscription // 可选:订阅信息 + InboundEndpoint string // 入站端点(客户端请求路径) + UpstreamEndpoint string // 上游端点(标准化后的上游路径) UserAgent string // 请求的 User-Agent IPAddress string // 请求的客户端 IP 地址 RequestPayloadHash string // 请求体语义哈希,用于降低 request_id 误复用时的静默误去重风险 @@ -7705,6 +7711,8 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input * RequestID: requestID, Model: result.Model, ReasoningEffort: result.ReasoningEffort, + InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), + UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), InputTokens: result.Usage.InputTokens, OutputTokens: result.Usage.OutputTokens, CacheCreationTokens: result.Usage.CacheCreationInputTokens, From 2c9dcfe27b8c187f44316393ada984ffc71d92dd Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Sun, 15 Mar 2026 22:13:31 +0800 Subject: [PATCH 03/15] refactor: add unified endpoint normalization infrastructure Introduce endpoint.go with shared constants, NormalizeInboundEndpoint, DeriveUpstreamEndpoint, InboundEndpointMiddleware, and context helpers. This replaces the two separate normalization implementations (OpenAI and Gateway) with a single source of truth. Includes comprehensive test coverage. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- backend/internal/handler/endpoint.go | 174 ++++++++++++++++++++++ backend/internal/handler/endpoint_test.go | 159 ++++++++++++++++++++ 2 files changed, 333 insertions(+) create mode 100644 backend/internal/handler/endpoint.go create mode 100644 backend/internal/handler/endpoint_test.go diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go new file mode 100644 index 0000000000..b120098875 --- /dev/null +++ b/backend/internal/handler/endpoint.go @@ -0,0 +1,174 @@ +package handler + +import ( + "strings" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +// ────────────────────────────────────────────────────────── +// Canonical inbound / upstream endpoint paths. +// All normalization and derivation reference this single set +// of constants — add new paths HERE when a new API surface +// is introduced. +// ────────────────────────────────────────────────────────── + +const ( + EndpointMessages = "/v1/messages" + EndpointChatCompletions = "/v1/chat/completions" + EndpointResponses = "/v1/responses" + EndpointGeminiModels = "/v1beta/models" +) + +// gin.Context keys used by the middleware and helpers below. +const ( + ctxKeyInboundEndpoint = "_gateway_inbound_endpoint" +) + +// ────────────────────────────────────────────────────────── +// Normalization functions +// ────────────────────────────────────────────────────────── + +// NormalizeInboundEndpoint maps a raw request path (which may carry +// prefixes like /antigravity, /openai, /sora) to its canonical form. +// +// "/antigravity/v1/messages" → "/v1/messages" +// "/v1/chat/completions" → "/v1/chat/completions" +// "/openai/v1/responses/foo" → "/v1/responses" +// "/v1beta/models/gemini:gen" → "/v1beta/models" +func NormalizeInboundEndpoint(path string) string { + path = strings.TrimSpace(path) + switch { + case strings.Contains(path, EndpointChatCompletions): + return EndpointChatCompletions + case strings.Contains(path, EndpointMessages): + return EndpointMessages + case strings.Contains(path, EndpointResponses): + return EndpointResponses + case strings.Contains(path, EndpointGeminiModels): + return EndpointGeminiModels + default: + return path + } +} + +// DeriveUpstreamEndpoint determines the upstream endpoint from the +// account platform and the normalized inbound endpoint. +// +// Platform-specific rules: +// - OpenAI always forwards to /v1/responses (with optional subpath +// such as /v1/responses/compact preserved from the raw URL). +// - Anthropic → /v1/messages +// - Gemini → /v1beta/models +// - Sora → /v1/chat/completions +// - Antigravity routes may target either Claude or Gemini, so the +// inbound endpoint is used to distinguish. +func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { + inbound = strings.TrimSpace(inbound) + + switch platform { + case service.PlatformOpenAI: + // OpenAI forwards everything to the Responses API. + // Preserve subresource suffix (e.g. /v1/responses/compact). + if suffix := responsesSubpathSuffix(rawRequestPath); suffix != "" { + return EndpointResponses + suffix + } + return EndpointResponses + + case service.PlatformAnthropic: + return EndpointMessages + + case service.PlatformGemini: + return EndpointGeminiModels + + case service.PlatformSora: + return EndpointChatCompletions + + case service.PlatformAntigravity: + // Antigravity accounts serve both Claude and Gemini. + if inbound == EndpointGeminiModels { + return EndpointGeminiModels + } + return EndpointMessages + } + + // Unknown platform — fall back to inbound. + return inbound +} + +// responsesSubpathSuffix extracts the part after "/responses" in a raw +// request path, e.g. "/openai/v1/responses/compact" → "/compact". +// Returns "" when there is no meaningful suffix. +func responsesSubpathSuffix(rawPath string) string { + trimmed := strings.TrimRight(strings.TrimSpace(rawPath), "/") + idx := strings.LastIndex(trimmed, "/responses") + if idx < 0 { + return "" + } + suffix := trimmed[idx+len("/responses"):] + if suffix == "" || suffix == "/" { + return "" + } + if !strings.HasPrefix(suffix, "/") { + return "" + } + return suffix +} + +// ────────────────────────────────────────────────────────── +// Middleware +// ────────────────────────────────────────────────────────── + +// InboundEndpointMiddleware normalizes the request path and stores the +// canonical inbound endpoint in gin.Context so that every handler in +// the chain can read it via GetInboundEndpoint. +// +// Apply this middleware to all gateway route groups. +func InboundEndpointMiddleware() gin.HandlerFunc { + return func(c *gin.Context) { + path := c.FullPath() + if path == "" && c.Request != nil && c.Request.URL != nil { + path = c.Request.URL.Path + } + c.Set(ctxKeyInboundEndpoint, NormalizeInboundEndpoint(path)) + c.Next() + } +} + +// ────────────────────────────────────────────────────────── +// Context helpers — used by handlers before building +// RecordUsageInput / RecordUsageLongContextInput. +// ────────────────────────────────────────────────────────── + +// GetInboundEndpoint returns the canonical inbound endpoint stored by +// InboundEndpointMiddleware. If the middleware did not run (e.g. in +// tests), it falls back to normalizing c.FullPath() on the fly. +func GetInboundEndpoint(c *gin.Context) string { + if v, ok := c.Get(ctxKeyInboundEndpoint); ok { + if s, ok := v.(string); ok && s != "" { + return s + } + } + // Fallback: normalize on the fly. + path := "" + if c != nil { + path = c.FullPath() + if path == "" && c.Request != nil && c.Request.URL != nil { + path = c.Request.URL.Path + } + } + return NormalizeInboundEndpoint(path) +} + +// GetUpstreamEndpoint derives the upstream endpoint from the context +// and the account platform. Handlers call this after scheduling an +// account, passing account.Platform. +func GetUpstreamEndpoint(c *gin.Context, platform string) string { + inbound := GetInboundEndpoint(c) + rawPath := "" + if c != nil && c.Request != nil && c.Request.URL != nil { + rawPath = c.Request.URL.Path + } + return DeriveUpstreamEndpoint(inbound, rawPath, platform) +} diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go new file mode 100644 index 0000000000..a3767ac499 --- /dev/null +++ b/backend/internal/handler/endpoint_test.go @@ -0,0 +1,159 @@ +package handler + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func init() { gin.SetMode(gin.TestMode) } + +// ────────────────────────────────────────────────────────── +// NormalizeInboundEndpoint +// ────────────────────────────────────────────────────────── + +func TestNormalizeInboundEndpoint(t *testing.T) { + tests := []struct { + path string + want string + }{ + // Direct canonical paths. + {"/v1/messages", EndpointMessages}, + {"/v1/chat/completions", EndpointChatCompletions}, + {"/v1/responses", EndpointResponses}, + {"/v1beta/models", EndpointGeminiModels}, + + // Prefixed paths (antigravity, openai, sora). + {"/antigravity/v1/messages", EndpointMessages}, + {"/openai/v1/responses", EndpointResponses}, + {"/openai/v1/responses/compact", EndpointResponses}, + {"/sora/v1/chat/completions", EndpointChatCompletions}, + {"/antigravity/v1beta/models/gemini:generateContent", EndpointGeminiModels}, + + // Gin route patterns with wildcards. + {"/v1beta/models/*modelAction", EndpointGeminiModels}, + {"/v1/responses/*subpath", EndpointResponses}, + + // Unknown path is returned as-is. + {"/v1/embeddings", "/v1/embeddings"}, + {"", ""}, + {" /v1/messages ", EndpointMessages}, + } + for _, tt := range tests { + t.Run(tt.path, func(t *testing.T) { + require.Equal(t, tt.want, NormalizeInboundEndpoint(tt.path)) + }) + } +} + +// ────────────────────────────────────────────────────────── +// DeriveUpstreamEndpoint +// ────────────────────────────────────────────────────────── + +func TestDeriveUpstreamEndpoint(t *testing.T) { + tests := []struct { + name string + inbound string + rawPath string + platform string + want string + }{ + // Anthropic. + {"anthropic messages", EndpointMessages, "/v1/messages", service.PlatformAnthropic, EndpointMessages}, + + // Gemini. + {"gemini models", EndpointGeminiModels, "/v1beta/models/gemini:gen", service.PlatformGemini, EndpointGeminiModels}, + + // Sora. + {"sora completions", EndpointChatCompletions, "/sora/v1/chat/completions", service.PlatformSora, EndpointChatCompletions}, + + // OpenAI — always /v1/responses. + {"openai responses root", EndpointResponses, "/v1/responses", service.PlatformOpenAI, EndpointResponses}, + {"openai responses compact", EndpointResponses, "/openai/v1/responses/compact", service.PlatformOpenAI, "/v1/responses/compact"}, + {"openai responses nested", EndpointResponses, "/openai/v1/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"}, + {"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses}, + {"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses}, + + // Antigravity — uses inbound to pick Claude vs Gemini upstream. + {"antigravity claude", EndpointMessages, "/antigravity/v1/messages", service.PlatformAntigravity, EndpointMessages}, + {"antigravity gemini", EndpointGeminiModels, "/antigravity/v1beta/models", service.PlatformAntigravity, EndpointGeminiModels}, + + // Unknown platform — passthrough. + {"unknown platform", "/v1/embeddings", "/v1/embeddings", "unknown", "/v1/embeddings"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, DeriveUpstreamEndpoint(tt.inbound, tt.rawPath, tt.platform)) + }) + } +} + +// ────────────────────────────────────────────────────────── +// responsesSubpathSuffix +// ────────────────────────────────────────────────────────── + +func TestResponsesSubpathSuffix(t *testing.T) { + tests := []struct { + raw string + want string + }{ + {"/v1/responses", ""}, + {"/v1/responses/", ""}, + {"/v1/responses/compact", "/compact"}, + {"/openai/v1/responses/compact/detail", "/compact/detail"}, + {"/v1/messages", ""}, + {"", ""}, + } + for _, tt := range tests { + t.Run(tt.raw, func(t *testing.T) { + require.Equal(t, tt.want, responsesSubpathSuffix(tt.raw)) + }) + } +} + +// ────────────────────────────────────────────────────────── +// InboundEndpointMiddleware + context helpers +// ────────────────────────────────────────────────────────── + +func TestInboundEndpointMiddleware(t *testing.T) { + router := gin.New() + router.Use(InboundEndpointMiddleware()) + + var captured string + router.POST("/v1/messages", func(c *gin.Context) { + captured = GetInboundEndpoint(c) + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, EndpointMessages, captured) +} + +func TestGetInboundEndpoint_FallbackWithoutMiddleware(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/antigravity/v1/messages", nil) + + // Middleware did not run — fallback to normalizing c.Request.URL.Path. + got := GetInboundEndpoint(c) + require.Equal(t, EndpointMessages, got) +} + +func TestGetUpstreamEndpoint_FullFlow(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses/compact", nil) + + // Simulate middleware. + c.Set(ctxKeyInboundEndpoint, NormalizeInboundEndpoint(c.Request.URL.Path)) + + got := GetUpstreamEndpoint(c, service.PlatformOpenAI) + require.Equal(t, "/v1/responses/compact", got) +} From 7bd1972f945a14fed4bf981541fbcf65ed55e4b2 Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Sun, 15 Mar 2026 22:13:42 +0800 Subject: [PATCH 04/15] refactor: migrate all handlers to shared endpoint normalization middleware - Apply InboundEndpointMiddleware to all gateway route groups - Replace normalizedOpenAIInboundEndpoint/normalizedOpenAIUpstreamEndpoint and normalizedGatewayInboundEndpoint/normalizedGatewayUpstreamEndpoint with GetInboundEndpoint/GetUpstreamEndpoint - Remove 4 old constants and 4 old normalization functions (-70 lines) - Migrate existing endpoint normalization test to new API Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-opencode) Co-authored-by: Sisyphus --- backend/internal/handler/gateway_handler.go | 10 ++- .../internal/handler/gemini_v1beta_handler.go | 4 + .../handler/openai_chat_completions.go | 4 +- ...nai_gateway_endpoint_normalization_test.go | 43 ++++++----- .../handler/openai_gateway_handler.go | 75 ++----------------- .../internal/handler/sora_gateway_handler.go | 4 + backend/internal/server/routes/gateway.go | 14 +++- 7 files changed, 56 insertions(+), 98 deletions(-) diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 09652adaa7..831029c48b 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -442,6 +442,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { userAgent := c.GetHeader("User-Agent") clientIP := ip.GetClientIP(c) requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) if result.ReasoningEffort == nil { result.ReasoningEffort = service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort) @@ -455,6 +457,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, @@ -757,6 +761,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { userAgent := c.GetHeader("User-Agent") clientIP := ip.GetClientIP(c) requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) if result.ReasoningEffort == nil { result.ReasoningEffort = service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort) @@ -770,6 +776,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { User: currentAPIKey.User, Account: account, Subscription: currentSubscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, @@ -935,7 +943,7 @@ func (h *GatewayHandler) parseUsageDateRange(c *gin.Context) (time.Time, time.Ti } if s := c.Query("end_date"); s != "" { if t, err := timezone.ParseInLocation("2006-01-02", s); err == nil { - endTime = t.Add(24*time.Hour - time.Second) // end of day + endTime = t.AddDate(0, 0, 1) // half-open range upper bound } } return startTime, endTime diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 9a16ff3a20..cfe809114b 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -504,6 +504,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) h.submitUsageRecordTask(func(ctx context.Context) { if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{ Result: result, @@ -511,6 +513,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 82b11c1065..4db5cadd71 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -261,8 +261,8 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, - InboundEndpoint: normalizedOpenAIInboundEndpoint(c, openAIInboundEndpointChatCompletions), - UpstreamEndpoint: normalizedOpenAIUpstreamEndpoint(c, openAIUpstreamEndpointResponses), + InboundEndpoint: GetInboundEndpoint(c), + UpstreamEndpoint: GetUpstreamEndpoint(c, account.Platform), UserAgent: userAgent, IPAddress: clientIP, APIKeyService: h.apiKeyService, diff --git a/backend/internal/handler/openai_gateway_endpoint_normalization_test.go b/backend/internal/handler/openai_gateway_endpoint_normalization_test.go index 6a055272ca..0dacd74dc1 100644 --- a/backend/internal/handler/openai_gateway_endpoint_normalization_test.go +++ b/backend/internal/handler/openai_gateway_endpoint_normalization_test.go @@ -5,42 +5,41 @@ import ( "net/http/httptest" "testing" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) -func TestNormalizedOpenAIUpstreamEndpoint(t *testing.T) { +// TestOpenAIUpstreamEndpoint_ViaGetUpstreamEndpoint verifies that the +// unified GetUpstreamEndpoint helper produces the same results as the +// former normalizedOpenAIUpstreamEndpoint for OpenAI platform requests. +func TestOpenAIUpstreamEndpoint_ViaGetUpstreamEndpoint(t *testing.T) { gin.SetMode(gin.TestMode) tests := []struct { - name string - path string - fallback string - want string + name string + path string + want string }{ { - name: "responses root maps to responses upstream", - path: "/v1/responses", - fallback: openAIUpstreamEndpointResponses, - want: "/v1/responses", + name: "responses root maps to responses upstream", + path: "/v1/responses", + want: EndpointResponses, }, { - name: "responses compact keeps compact suffix", - path: "/openai/v1/responses/compact", - fallback: openAIUpstreamEndpointResponses, - want: "/v1/responses/compact", + name: "responses compact keeps compact suffix", + path: "/openai/v1/responses/compact", + want: "/v1/responses/compact", }, { - name: "responses nested suffix preserved", - path: "/openai/v1/responses/compact/detail", - fallback: openAIUpstreamEndpointResponses, - want: "/v1/responses/compact/detail", + name: "responses nested suffix preserved", + path: "/openai/v1/responses/compact/detail", + want: "/v1/responses/compact/detail", }, { - name: "non responses path uses fallback", - path: "/v1/messages", - fallback: openAIUpstreamEndpointResponses, - want: "/v1/responses", + name: "non responses path uses platform fallback", + path: "/v1/messages", + want: EndpointResponses, }, } @@ -50,7 +49,7 @@ func TestNormalizedOpenAIUpstreamEndpoint(t *testing.T) { c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, tt.path, nil) - got := normalizedOpenAIUpstreamEndpoint(c, tt.fallback) + got := GetUpstreamEndpoint(c, service.PlatformOpenAI) require.Equal(t, tt.want, got) }) } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index b2aa5c504f..c681e61de1 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -37,13 +37,6 @@ type OpenAIGatewayHandler struct { cfg *config.Config } -const ( - openAIInboundEndpointResponses = "/v1/responses" - openAIInboundEndpointMessages = "/v1/messages" - openAIInboundEndpointChatCompletions = "/v1/chat/completions" - openAIUpstreamEndpointResponses = "/v1/responses" -) - // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler func NewOpenAIGatewayHandler( gatewayService *service.OpenAIGatewayService, @@ -369,8 +362,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, - InboundEndpoint: normalizedOpenAIInboundEndpoint(c, openAIInboundEndpointResponses), - UpstreamEndpoint: normalizedOpenAIUpstreamEndpoint(c, openAIUpstreamEndpointResponses), + InboundEndpoint: GetInboundEndpoint(c), + UpstreamEndpoint: GetUpstreamEndpoint(c, account.Platform), UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, @@ -747,8 +740,8 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, - InboundEndpoint: normalizedOpenAIInboundEndpoint(c, openAIInboundEndpointMessages), - UpstreamEndpoint: normalizedOpenAIUpstreamEndpoint(c, openAIUpstreamEndpointResponses), + InboundEndpoint: GetInboundEndpoint(c), + UpstreamEndpoint: GetUpstreamEndpoint(c, account.Platform), UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, @@ -1246,8 +1239,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, - InboundEndpoint: normalizedOpenAIInboundEndpoint(c, openAIInboundEndpointResponses), - UpstreamEndpoint: normalizedOpenAIUpstreamEndpoint(c, openAIUpstreamEndpointResponses), + InboundEndpoint: GetInboundEndpoint(c), + UpstreamEndpoint: GetUpstreamEndpoint(c, account.Platform), UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: service.HashUsageRequestPayload(firstMessage), @@ -1543,62 +1536,6 @@ func openAIWSIngressFallbackSessionSeed(userID, apiKeyID int64, groupID *int64) return fmt.Sprintf("openai_ws_ingress:%d:%d:%d", gid, userID, apiKeyID) } -func normalizedOpenAIInboundEndpoint(c *gin.Context, fallback string) string { - path := strings.TrimSpace(fallback) - if c != nil { - if fullPath := strings.TrimSpace(c.FullPath()); fullPath != "" { - path = fullPath - } else if c.Request != nil && c.Request.URL != nil { - if requestPath := strings.TrimSpace(c.Request.URL.Path); requestPath != "" { - path = requestPath - } - } - } - - switch { - case strings.Contains(path, openAIInboundEndpointChatCompletions): - return openAIInboundEndpointChatCompletions - case strings.Contains(path, openAIInboundEndpointMessages): - return openAIInboundEndpointMessages - case strings.Contains(path, openAIInboundEndpointResponses): - return openAIInboundEndpointResponses - default: - return path - } -} - -func normalizedOpenAIUpstreamEndpoint(c *gin.Context, fallback string) string { - base := strings.TrimSpace(fallback) - if base == "" { - base = openAIUpstreamEndpointResponses - } - base = strings.TrimRight(base, "/") - - if c == nil || c.Request == nil || c.Request.URL == nil { - return base - } - - path := strings.TrimRight(strings.TrimSpace(c.Request.URL.Path), "/") - if path == "" { - return base - } - - idx := strings.LastIndex(path, "/responses") - if idx < 0 { - return base - } - - suffix := strings.TrimSpace(path[idx+len("/responses"):]) - if suffix == "" || suffix == "/" { - return base - } - if !strings.HasPrefix(suffix, "/") { - return base - } - - return base + suffix -} - func isOpenAIWSUpgradeRequest(r *http.Request) bool { if r == nil { return false diff --git a/backend/internal/handler/sora_gateway_handler.go b/backend/internal/handler/sora_gateway_handler.go index 06abdf6037..dc301ce149 100644 --- a/backend/internal/handler/sora_gateway_handler.go +++ b/backend/internal/handler/sora_gateway_handler.go @@ -400,6 +400,8 @@ func (h *SoraGatewayHandler) ChatCompletions(c *gin.Context) { userAgent := c.GetHeader("User-Agent") clientIP := ip.GetClientIP(c) requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 h.submitUsageRecordTask(func(ctx context.Context) { @@ -409,6 +411,8 @@ func (h *SoraGatewayHandler) ChatCompletions(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, IPAddress: clientIP, RequestPayloadHash: requestPayloadHash, diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index ea40f2f189..fe82083096 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -30,6 +30,7 @@ func RegisterGatewayRoutes( soraBodyLimit := middleware.RequestBodyLimit(soraMaxBodySize) clientRequestID := middleware.ClientRequestID() opsErrorLogger := handler.OpsErrorLoggerMiddleware(opsService) + endpointNorm := handler.InboundEndpointMiddleware() // 未分组 Key 拦截中间件(按协议格式区分错误响应) requireGroupAnthropic := middleware.RequireGroupAssignment(settingService, middleware.AnthropicErrorWriter) @@ -40,6 +41,7 @@ func RegisterGatewayRoutes( gateway.Use(bodyLimit) gateway.Use(clientRequestID) gateway.Use(opsErrorLogger) + gateway.Use(endpointNorm) gateway.Use(gin.HandlerFunc(apiKeyAuth)) gateway.Use(requireGroupAnthropic) { @@ -80,6 +82,7 @@ func RegisterGatewayRoutes( gemini.Use(bodyLimit) gemini.Use(clientRequestID) gemini.Use(opsErrorLogger) + gemini.Use(endpointNorm) gemini.Use(middleware.APIKeyAuthWithSubscriptionGoogle(apiKeyService, subscriptionService, cfg)) gemini.Use(requireGroupGoogle) { @@ -90,11 +93,11 @@ func RegisterGatewayRoutes( } // OpenAI Responses API(不带v1前缀的别名) - r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.Responses) - r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.Responses) - r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ResponsesWebSocket) + r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.Responses) + r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.Responses) + r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ResponsesWebSocket) // OpenAI Chat Completions API(不带v1前缀的别名) - r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ChatCompletions) + r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ChatCompletions) // Antigravity 模型列表 r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels) @@ -104,6 +107,7 @@ func RegisterGatewayRoutes( antigravityV1.Use(bodyLimit) antigravityV1.Use(clientRequestID) antigravityV1.Use(opsErrorLogger) + antigravityV1.Use(endpointNorm) antigravityV1.Use(middleware.ForcePlatform(service.PlatformAntigravity)) antigravityV1.Use(gin.HandlerFunc(apiKeyAuth)) antigravityV1.Use(requireGroupAnthropic) @@ -118,6 +122,7 @@ func RegisterGatewayRoutes( antigravityV1Beta.Use(bodyLimit) antigravityV1Beta.Use(clientRequestID) antigravityV1Beta.Use(opsErrorLogger) + antigravityV1Beta.Use(endpointNorm) antigravityV1Beta.Use(middleware.ForcePlatform(service.PlatformAntigravity)) antigravityV1Beta.Use(middleware.APIKeyAuthWithSubscriptionGoogle(apiKeyService, subscriptionService, cfg)) antigravityV1Beta.Use(requireGroupGoogle) @@ -132,6 +137,7 @@ func RegisterGatewayRoutes( soraV1.Use(soraBodyLimit) soraV1.Use(clientRequestID) soraV1.Use(opsErrorLogger) + soraV1.Use(endpointNorm) soraV1.Use(middleware.ForcePlatform(service.PlatformSora)) soraV1.Use(gin.HandlerFunc(apiKeyAuth)) soraV1.Use(requireGroupAnthropic) From 8147866c09a6e970bf04a2caeb709cadd6630a83 Mon Sep 17 00:00:00 2001 From: Peter <1tRq4X287b7W7sfKf9GsWI+Peter@noreply.cnb.cool> Date: Mon, 16 Mar 2026 00:17:47 +0800 Subject: [PATCH 05/15] fix(admin): polish spending ranking and usage defaults --- .../handler/admin/dashboard_handler.go | 2 + .../dashboard_handler_request_type_test.go | 4 ++ .../pkg/usagestats/usage_log_types.go | 2 + backend/internal/repository/usage_log_repo.go | 14 +++- .../usage_log_repo_request_type_test.go | 10 +-- .../charts/ModelDistributionChart.vue | 69 +++++++++++++++---- .../__tests__/ModelDistributionChart.spec.ts | 52 ++++++++++++++ frontend/src/types/index.ts | 2 + frontend/src/views/admin/DashboardView.vue | 10 ++- frontend/src/views/admin/UsageView.vue | 6 +- 10 files changed, 148 insertions(+), 23 deletions(-) diff --git a/backend/internal/handler/admin/dashboard_handler.go b/backend/internal/handler/admin/dashboard_handler.go index cc4ef2d0f6..f415b48f2e 100644 --- a/backend/internal/handler/admin/dashboard_handler.go +++ b/backend/internal/handler/admin/dashboard_handler.go @@ -512,6 +512,8 @@ func (h *DashboardHandler) GetUserSpendingRanking(c *gin.Context) { payload := gin.H{ "ranking": ranking.Ranking, "total_actual_cost": ranking.TotalActualCost, + "total_requests": ranking.TotalRequests, + "total_tokens": ranking.TotalTokens, "start_date": startTime.Format("2006-01-02"), "end_date": endTime.Add(-24 * time.Hour).Format("2006-01-02"), } diff --git a/backend/internal/handler/admin/dashboard_handler_request_type_test.go b/backend/internal/handler/admin/dashboard_handler_request_type_test.go index 6b363bb5c6..9aec61d469 100644 --- a/backend/internal/handler/admin/dashboard_handler_request_type_test.go +++ b/backend/internal/handler/admin/dashboard_handler_request_type_test.go @@ -61,6 +61,8 @@ func (s *dashboardUsageRepoCapture) GetUserSpendingRanking( return &usagestats.UserSpendingRankingResponse{ Ranking: s.ranking, TotalActualCost: s.rankingTotal, + TotalRequests: 44, + TotalTokens: 1234, }, nil } @@ -164,6 +166,8 @@ func TestDashboardUsersRankingLimitAndCache(t *testing.T) { require.Equal(t, http.StatusOK, rec.Code) require.Equal(t, 50, repo.rankingLimit) require.Contains(t, rec.Body.String(), "\"total_actual_cost\":88.8") + require.Contains(t, rec.Body.String(), "\"total_requests\":44") + require.Contains(t, rec.Body.String(), "\"total_tokens\":1234") require.Equal(t, "miss", rec.Header().Get("X-Snapshot-Cache")) req2 := httptest.NewRequest(http.MethodGet, "/admin/dashboard/users-ranking?limit=100&start_date=2025-01-01&end_date=2025-01-02", nil) diff --git a/backend/internal/pkg/usagestats/usage_log_types.go b/backend/internal/pkg/usagestats/usage_log_types.go index 55a049d3e2..6b980dc88d 100644 --- a/backend/internal/pkg/usagestats/usage_log_types.go +++ b/backend/internal/pkg/usagestats/usage_log_types.go @@ -116,6 +116,8 @@ type UserSpendingRankingItem struct { type UserSpendingRankingResponse struct { Ranking []UserSpendingRankingItem `json:"ranking"` TotalActualCost float64 `json:"total_actual_cost"` + TotalRequests int64 `json:"total_requests"` + TotalTokens int64 `json:"total_tokens"` } // APIKeyUsageTrendPoint represents API key usage trend data point diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index 845f2cf020..8ee21d95ad 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -2139,7 +2139,9 @@ func (r *usageLogRepository) GetUserSpendingRanking(ctx context.Context, startTi actual_cost, requests, tokens, - COALESCE(SUM(actual_cost) OVER (), 0) as total_actual_cost + COALESCE(SUM(actual_cost) OVER (), 0) as total_actual_cost, + COALESCE(SUM(requests) OVER (), 0) as total_requests, + COALESCE(SUM(tokens) OVER (), 0) as total_tokens FROM user_spend ORDER BY actual_cost DESC, tokens DESC, user_id ASC LIMIT $3 @@ -2150,7 +2152,9 @@ func (r *usageLogRepository) GetUserSpendingRanking(ctx context.Context, startTi actual_cost, requests, tokens, - total_actual_cost + total_actual_cost, + total_requests, + total_tokens FROM ranked ORDER BY actual_cost DESC, tokens DESC, user_id ASC ` @@ -2168,9 +2172,11 @@ func (r *usageLogRepository) GetUserSpendingRanking(ctx context.Context, startTi ranking := make([]UserSpendingRankingItem, 0) totalActualCost := 0.0 + totalRequests := int64(0) + totalTokens := int64(0) for rows.Next() { var row UserSpendingRankingItem - if err = rows.Scan(&row.UserID, &row.Email, &row.ActualCost, &row.Requests, &row.Tokens, &totalActualCost); err != nil { + if err = rows.Scan(&row.UserID, &row.Email, &row.ActualCost, &row.Requests, &row.Tokens, &totalActualCost, &totalRequests, &totalTokens); err != nil { return nil, err } ranking = append(ranking, row) @@ -2182,6 +2188,8 @@ func (r *usageLogRepository) GetUserSpendingRanking(ctx context.Context, startTi return &UserSpendingRankingResponse{ Ranking: ranking, TotalActualCost: totalActualCost, + TotalRequests: totalRequests, + TotalTokens: totalTokens, }, nil } 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 bcb23717bb..f1bf1f1d81 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -255,10 +255,10 @@ func TestUsageLogRepositoryGetUserSpendingRanking(t *testing.T) { start := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) end := start.Add(24 * time.Hour) - rows := sqlmock.NewRows([]string{"user_id", "email", "actual_cost", "requests", "tokens", "total_actual_cost"}). - AddRow(int64(2), "beta@example.com", 12.5, int64(9), int64(900), 40.0). - AddRow(int64(1), "alpha@example.com", 12.5, int64(8), int64(800), 40.0). - AddRow(int64(3), "gamma@example.com", 4.25, int64(5), int64(300), 40.0) + rows := sqlmock.NewRows([]string{"user_id", "email", "actual_cost", "requests", "tokens", "total_actual_cost", "total_requests", "total_tokens"}). + AddRow(int64(2), "beta@example.com", 12.5, int64(9), int64(900), 40.0, int64(30), int64(2600)). + AddRow(int64(1), "alpha@example.com", 12.5, int64(8), int64(800), 40.0, int64(30), int64(2600)). + AddRow(int64(3), "gamma@example.com", 4.25, int64(5), int64(300), 40.0, int64(30), int64(2600)) mock.ExpectQuery("WITH user_spend AS \\("). WithArgs(start, end, 12). @@ -273,6 +273,8 @@ func TestUsageLogRepositoryGetUserSpendingRanking(t *testing.T) { {UserID: 3, Email: "gamma@example.com", ActualCost: 4.25, Requests: 5, Tokens: 300}, }, TotalActualCost: 40.0, + TotalRequests: 30, + TotalTokens: 2600, }, got) require.NoError(t, mock.ExpectationsWereMet()) } diff --git a/frontend/src/components/charts/ModelDistributionChart.vue b/frontend/src/components/charts/ModelDistributionChart.vue index 5db5a14f68..5ae9b38e63 100644 --- a/frontend/src/components/charts/ModelDistributionChart.vue +++ b/frontend/src/components/charts/ModelDistributionChart.vue @@ -127,7 +127,7 @@ > {{ t('admin.dashboard.failedToLoad') }} -
+
@@ -143,21 +143,24 @@
- #{{ index + 1 }} + {{ item.isOther ? 'Σ' : `#${index + 1}` }} - {{ getRankingUserLabel(item) }} + {{ getRankingRowLabel(item) }}
@@ -197,11 +200,14 @@ ChartJS.register(ArcElement, Tooltip, Legend) const { t } = useI18n() type DistributionMetric = 'tokens' | 'actual_cost' +type RankingDisplayItem = UserSpendingRankingItem & { isOther?: boolean } const props = withDefaults(defineProps<{ modelStats: ModelStat[] enableRankingView?: boolean rankingItems?: UserSpendingRankingItem[] rankingTotalActualCost?: number + rankingTotalRequests?: number + rankingTotalTokens?: number loading?: boolean metric?: DistributionMetric showMetricToggle?: boolean @@ -211,6 +217,8 @@ const props = withDefaults(defineProps<{ enableRankingView: false, rankingItems: () => [], rankingTotalActualCost: 0, + rankingTotalRequests: 0, + rankingTotalTokens: 0, loading: false, metric: 'tokens', showMetricToggle: false, @@ -266,14 +274,14 @@ const chartData = computed(() => { const rankingChartData = computed(() => { if (!props.rankingItems?.length) return null - const rankedTotal = props.rankingItems.reduce((sum, item) => sum + item.actual_cost, 0) - const otherActualCost = Math.max((props.rankingTotalActualCost || 0) - rankedTotal, 0) const labels = props.rankingItems.map((item, index) => `#${index + 1} ${getRankingUserLabel(item)}`) const data = props.rankingItems.map((item) => item.actual_cost) + const backgroundColor = chartColors.slice(0, props.rankingItems.length) - if (otherActualCost > 0.000001) { + if (otherRankingItem.value) { labels.push(t('admin.dashboard.spendingRankingOther')) - data.push(otherActualCost) + data.push(otherRankingItem.value.actual_cost) + backgroundColor.push('#94a3b8') } return { @@ -281,13 +289,43 @@ const rankingChartData = computed(() => { datasets: [ { data, - backgroundColor: chartColors.slice(0, data.length), + backgroundColor, borderWidth: 0 } ] } }) +const otherRankingItem = computed(() => { + if (!props.rankingItems?.length) return null + + const rankedActualCost = props.rankingItems.reduce((sum, item) => sum + item.actual_cost, 0) + const rankedRequests = props.rankingItems.reduce((sum, item) => sum + item.requests, 0) + const rankedTokens = props.rankingItems.reduce((sum, item) => sum + item.tokens, 0) + + const otherActualCost = Math.max((props.rankingTotalActualCost || 0) - rankedActualCost, 0) + const otherRequests = Math.max((props.rankingTotalRequests || 0) - rankedRequests, 0) + const otherTokens = Math.max((props.rankingTotalTokens || 0) - rankedTokens, 0) + + if (otherActualCost <= 0.000001 && otherRequests <= 0 && otherTokens <= 0) return null + + return { + user_id: 0, + email: '', + actual_cost: otherActualCost, + requests: otherRequests, + tokens: otherTokens, + isOther: true + } +}) + +const rankingDisplayItems = computed(() => { + if (!props.rankingItems?.length) return [] + return otherRankingItem.value + ? [...props.rankingItems, otherRankingItem.value] + : [...props.rankingItems] +}) + const doughnutOptions = computed(() => ({ responsive: true, maintainAspectRatio: false, @@ -351,6 +389,11 @@ const getRankingUserLabel = (item: UserSpendingRankingItem): string => { return t('admin.redeem.userPrefix', { id: item.user_id }) } +const getRankingRowLabel = (item: RankingDisplayItem): string => { + if (item.isOther) return t('admin.dashboard.spendingRankingOther') + return getRankingUserLabel(item) +} + const formatCost = (value: number): string => { if (value >= 1000) { return (value / 1000).toFixed(2) + 'K' diff --git a/frontend/src/components/charts/__tests__/ModelDistributionChart.spec.ts b/frontend/src/components/charts/__tests__/ModelDistributionChart.spec.ts index 27fb8bd459..82b6236765 100644 --- a/frontend/src/components/charts/__tests__/ModelDistributionChart.spec.ts +++ b/frontend/src/components/charts/__tests__/ModelDistributionChart.spec.ts @@ -5,6 +5,14 @@ import ModelDistributionChart from '../ModelDistributionChart.vue' const messages: Record = { 'admin.dashboard.modelDistribution': 'Model Distribution', + 'admin.dashboard.spendingRankingTitle': 'User Spending Ranking', + 'admin.dashboard.viewModelDistribution': 'Model Distribution', + 'admin.dashboard.viewSpendingRanking': 'User Spending Ranking', + 'admin.dashboard.spendingRankingUser': 'User', + 'admin.dashboard.spendingRankingRequests': 'Requests', + 'admin.dashboard.spendingRankingTokens': 'Tokens', + 'admin.dashboard.spendingRankingSpend': 'Spend', + 'admin.dashboard.spendingRankingOther': 'Others', 'admin.dashboard.model': 'Model', 'admin.dashboard.requests': 'Requests', 'admin.dashboard.tokens': 'Tokens', @@ -13,6 +21,7 @@ const messages: Record = { 'admin.dashboard.metricTokens': 'By Tokens', 'admin.dashboard.metricActualCost': 'By Actual Cost', 'admin.dashboard.noDataAvailable': 'No data available', + 'admin.redeem.userPrefix': 'User #{id}', } vi.mock('vue-i18n', async () => { @@ -116,4 +125,47 @@ describe('ModelDistributionChart', () => { }) expect(label).toBe('model-b: $1.40 (87.5%)') }) + + it('renders Others in the spending ranking table and uses a dedicated chart color', async () => { + const wrapper = mount(ModelDistributionChart, { + props: { + modelStats: [], + enableRankingView: true, + rankingItems: [ + { user_id: 1, email: 'alpha@example.com', actual_cost: 12, requests: 10, tokens: 1000 }, + { user_id: 2, email: 'beta@example.com', actual_cost: 8, requests: 6, tokens: 600 }, + ], + rankingTotalActualCost: 30, + rankingTotalRequests: 20, + rankingTotalTokens: 2000, + }, + global: { + stubs: { + LoadingSpinner: true, + }, + }, + }) + + const rankingButton = wrapper.findAll('button').find((button) => button.text() === 'User Spending Ranking') + expect(rankingButton).toBeTruthy() + await rankingButton!.trigger('click') + + const chartData = JSON.parse(wrapper.find('.chart-data').text()) + expect(chartData.labels).toEqual([ + '#1 alpha@example.com', + '#2 beta@example.com', + 'Others', + ]) + expect(chartData.datasets[0].data).toEqual([12, 8, 10]) + expect(chartData.datasets[0].backgroundColor[0]).toBe('#3b82f6') + expect(chartData.datasets[0].backgroundColor[2]).toBe('#94a3b8') + expect(chartData.datasets[0].backgroundColor[2]).not.toBe(chartData.datasets[0].backgroundColor[0]) + + const rows = wrapper.findAll('tbody tr') + expect(rows).toHaveLength(3) + expect(rows[2].text()).toContain('Others') + expect(rows[2].text()).toContain('4') + expect(rows[2].text()).toContain('400') + expect(rows[2].text()).toContain('$10.00') + }) }) diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 6f9bff765b..de9ddf61d1 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -1199,6 +1199,8 @@ export interface UserSpendingRankingItem { export interface UserSpendingRankingResponse { ranking: UserSpendingRankingItem[] total_actual_cost: number + total_requests: number + total_tokens: number start_date: string end_date: string } diff --git a/frontend/src/views/admin/DashboardView.vue b/frontend/src/views/admin/DashboardView.vue index 7d66b1f0ff..8b7ff632fe 100644 --- a/frontend/src/views/admin/DashboardView.vue +++ b/frontend/src/views/admin/DashboardView.vue @@ -241,6 +241,8 @@ :enable-ranking-view="true" :ranking-items="rankingItems" :ranking-total-actual-cost="rankingTotalActualCost" + :ranking-total-requests="rankingTotalRequests" + :ranking-total-tokens="rankingTotalTokens" :loading="chartsLoading" :ranking-loading="rankingLoading" :ranking-error="rankingError" @@ -334,6 +336,8 @@ const modelStats = ref([]) const userTrend = ref([]) const rankingItems = ref([]) const rankingTotalActualCost = ref(0) +const rankingTotalRequests = ref(0) +const rankingTotalTokens = ref(0) let chartLoadSeq = 0 let usersTrendLoadSeq = 0 let rankingLoadSeq = 0 @@ -347,7 +351,7 @@ const formatLocalDate = (date: Date): string => { const getTodayLocalDate = () => formatLocalDate(new Date()) // Date range -const granularity = ref<'day' | 'hour'>('day') +const granularity = ref<'day' | 'hour'>('hour') const startDate = ref(getTodayLocalDate()) const endDate = ref(getTodayLocalDate()) @@ -630,11 +634,15 @@ const loadUserSpendingRanking = async () => { if (currentSeq !== rankingLoadSeq) return rankingItems.value = response.ranking || [] rankingTotalActualCost.value = response.total_actual_cost || 0 + rankingTotalRequests.value = response.total_requests || 0 + rankingTotalTokens.value = response.total_tokens || 0 } catch (error) { if (currentSeq !== rankingLoadSeq) return console.error('Error loading user spending ranking:', error) rankingItems.value = [] rankingTotalActualCost.value = 0 + rankingTotalRequests.value = 0 + rankingTotalTokens.value = 0 rankingError.value = true } finally { if (currentSeq === rankingLoadSeq) { diff --git a/frontend/src/views/admin/UsageView.vue b/frontend/src/views/admin/UsageView.vue index 5a49864281..6b8130573f 100644 --- a/frontend/src/views/admin/UsageView.vue +++ b/frontend/src/views/admin/UsageView.vue @@ -107,7 +107,7 @@ const appStore = useAppStore() type DistributionMetric = 'tokens' | 'actual_cost' const route = useRoute() const usageStats = ref(null); const usageLogs = ref([]); const loading = ref(false); const exporting = ref(false) -const trendData = ref([]); const modelStats = ref([]); const groupStats = ref([]); const chartsLoading = ref(false); const granularity = ref<'day' | 'hour'>('day') +const trendData = ref([]); const modelStats = ref([]); const groupStats = ref([]); const chartsLoading = ref(false); const granularity = ref<'day' | 'hour'>('hour') const modelDistributionMetric = ref('tokens') const groupDistributionMetric = ref('tokens') let abortController: AbortController | null = null; let exportAbortController: AbortController | null = null @@ -137,6 +137,7 @@ const formatLD = (d: Date) => { return `${year}-${month}-${day}` } const getTodayLocalDate = () => formatLD(new Date()) +const getGranularityForRange = (start: string, end: string): 'day' | 'hour' => start === end ? 'hour' : 'day' const startDate = ref(getTodayLocalDate()); const endDate = ref(getTodayLocalDate()) const filters = ref({ user_id: undefined, model: undefined, group_id: undefined, request_type: undefined, billing_type: null, start_date: startDate.value, end_date: endDate.value }) const pagination = reactive({ page: 1, page_size: 20, total: 0 }) @@ -171,6 +172,7 @@ const applyRouteQueryFilters = () => { start_date: startDate.value, end_date: endDate.value } + granularity.value = getGranularityForRange(startDate.value, endDate.value) } const loadLogs = async () => { @@ -224,7 +226,7 @@ const loadChartData = async () => { } const applyFilters = () => { pagination.page = 1; loadLogs(); loadStats(); loadChartData() } const refreshData = () => { loadLogs(); loadStats(); loadChartData() } -const resetFilters = () => { startDate.value = getTodayLocalDate(); endDate.value = getTodayLocalDate(); filters.value = { start_date: startDate.value, end_date: endDate.value, request_type: undefined, billing_type: null }; granularity.value = 'day'; applyFilters() } +const resetFilters = () => { startDate.value = getTodayLocalDate(); endDate.value = getTodayLocalDate(); filters.value = { start_date: startDate.value, end_date: endDate.value, request_type: undefined, billing_type: null }; granularity.value = getGranularityForRange(startDate.value, endDate.value); applyFilters() } const handlePageChange = (p: number) => { pagination.page = p; loadLogs() } const handlePageSizeChange = (s: number) => { pagination.page_size = s; pagination.page = 1; loadLogs() } const cancelExport = () => exportAbortController?.abort() From ab4e8b2cf009c7067151b02adaa7bcd27865c8c0 Mon Sep 17 00:00:00 2001 From: QTom Date: Mon, 16 Mar 2026 10:28:11 +0800 Subject: [PATCH 06/15] =?UTF-8?q?fix(gateway):=20=E9=98=B2=E6=AD=A2=20Open?= =?UTF-8?q?AI=20Codex=20=E8=B7=A8=E7=94=A8=E6=88=B7=E4=B8=B2=E6=B5=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因:多个用户共享同一 OAuth 账号时,conversation_id/session_id 头 未做用户隔离,导致上游 chatgpt.com 将不同用户的请求关联到同一会话。 HTTP SSE 修复: - 新增 isolateOpenAISessionID(apiKeyID, raw),将 API Key ID 混入 session 标识符(xxhash),确保不同 Key 的用户产生不同上游会话 - buildUpstreamRequest: OAuth 分支先 Del 客户端透传的 session 头, 再用隔离值覆盖 - buildUpstreamRequestOpenAIPassthrough: 透传路径同样隔离 - ForwardAsAnthropic: Anthropic Messages 兼容路径同步修复 - buildOpenAIWSHeaders: WS 路径的 OAuth session 头同步隔离 --- .../service/openai_gateway_messages.go | 7 ++- .../service/openai_gateway_service.go | 57 ++++++++++++++----- ..._gateway_service_session_isolation_test.go | 50 ++++++++++++++++ .../internal/service/openai_ws_forwarder.go | 21 +++++-- .../openai_ws_forwarder_success_test.go | 9 ++- 5 files changed, 119 insertions(+), 25 deletions(-) create mode 100644 backend/internal/service/openai_gateway_service_session_isolation_test.go diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 1e40ec6f42..587145715d 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -107,10 +107,11 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return nil, fmt.Errorf("build upstream request: %w", err) } - // Override session_id with a deterministic UUID derived from the sticky - // session key (buildUpstreamRequest may have set it to the raw value). + // Override session_id with a deterministic UUID derived from the isolated + // session key, ensuring different API keys produce different upstream sessions. if promptCacheKey != "" { - upstreamReq.Header.Set("session_id", generateSessionUUID(promptCacheKey)) + apiKeyID := getAPIKeyIDFromContext(c) + upstreamReq.Header.Set("session_id", generateSessionUUID(isolateOpenAISessionID(apiKeyID, promptCacheKey))) } // 7. Send request diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 327ce91695..c8876edb12 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -24,6 +24,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" + "github.com/cespare/xxhash/v2" "github.com/gin-gonic/gin" "github.com/google/uuid" "github.com/tidwall/gjson" @@ -787,6 +788,20 @@ func getAPIKeyIDFromContext(c *gin.Context) int64 { return apiKey.ID } +// isolateOpenAISessionID 将 apiKeyID 混入 session 标识符, +// 确保不同 API Key 的用户即使使用相同的原始 session_id/conversation_id, +// 到达上游的标识符也不同,防止跨用户会话碰撞。 +func isolateOpenAISessionID(apiKeyID int64, raw string) string { + raw = strings.TrimSpace(raw) + if raw == "" { + return "" + } + h := xxhash.New() + _, _ = fmt.Fprintf(h, "k%d:", apiKeyID) + _, _ = h.WriteString(raw) + return fmt.Sprintf("%016x", h.Sum64()) +} + func logCodexCLIOnlyDetection(ctx context.Context, c *gin.Context, account *Account, apiKeyID int64, result CodexClientRestrictionDetectionResult, body []byte) { if !result.Enabled { return @@ -2501,13 +2516,17 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( if chatgptAccountID := account.GetChatGPTAccountID(); chatgptAccountID != "" { req.Header.Set("chatgpt-account-id", chatgptAccountID) } + apiKeyID := getAPIKeyIDFromContext(c) + // 先保存客户端原始值,再做 compact 补充,避免后续统一隔离时读到已处理的值。 + clientSessionID := strings.TrimSpace(req.Header.Get("session_id")) + clientConversationID := strings.TrimSpace(req.Header.Get("conversation_id")) if isOpenAIResponsesCompactPath(c) { req.Header.Set("accept", "application/json") if req.Header.Get("version") == "" { req.Header.Set("version", codexCLIVersion) } - if req.Header.Get("session_id") == "" { - req.Header.Set("session_id", resolveOpenAICompactSessionID(c)) + if clientSessionID == "" { + clientSessionID = resolveOpenAICompactSessionID(c) } } else if req.Header.Get("accept") == "" { req.Header.Set("accept", "text/event-stream") @@ -2518,13 +2537,18 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( if req.Header.Get("originator") == "" { req.Header.Set("originator", "codex_cli_rs") } - if promptCacheKey != "" { - if req.Header.Get("conversation_id") == "" { - req.Header.Set("conversation_id", promptCacheKey) - } - if req.Header.Get("session_id") == "" { - req.Header.Set("session_id", promptCacheKey) - } + // 用隔离后的 session 标识符覆盖客户端透传值,防止跨用户会话碰撞。 + if clientSessionID == "" { + clientSessionID = promptCacheKey + } + if clientConversationID == "" { + clientConversationID = promptCacheKey + } + if clientSessionID != "" { + req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, clientSessionID)) + } + if clientConversationID != "" { + req.Header.Set("conversation_id", isolateOpenAISessionID(apiKeyID, clientConversationID)) } } @@ -2887,22 +2911,27 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. } } if account.Type == AccountTypeOAuth { + // 清除客户端透传的 session 头,后续用隔离后的值重新设置,防止跨用户会话碰撞。 + req.Header.Del("conversation_id") + req.Header.Del("session_id") + req.Header.Set("OpenAI-Beta", "responses=experimental") req.Header.Set("originator", resolveOpenAIUpstreamOriginator(c, isCodexCLI)) + apiKeyID := getAPIKeyIDFromContext(c) if isOpenAIResponsesCompactPath(c) { req.Header.Set("accept", "application/json") if req.Header.Get("version") == "" { req.Header.Set("version", codexCLIVersion) } - if req.Header.Get("session_id") == "" { - req.Header.Set("session_id", resolveOpenAICompactSessionID(c)) - } + compactSession := resolveOpenAICompactSessionID(c) + req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, compactSession)) } else { req.Header.Set("accept", "text/event-stream") } if promptCacheKey != "" { - req.Header.Set("conversation_id", promptCacheKey) - req.Header.Set("session_id", promptCacheKey) + isolated := isolateOpenAISessionID(apiKeyID, promptCacheKey) + req.Header.Set("conversation_id", isolated) + req.Header.Set("session_id", isolated) } } diff --git a/backend/internal/service/openai_gateway_service_session_isolation_test.go b/backend/internal/service/openai_gateway_service_session_isolation_test.go new file mode 100644 index 0000000000..d42fbcc568 --- /dev/null +++ b/backend/internal/service/openai_gateway_service_session_isolation_test.go @@ -0,0 +1,50 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsolateOpenAISessionID(t *testing.T) { + t.Run("empty_raw_returns_empty", func(t *testing.T) { + assert.Equal(t, "", isolateOpenAISessionID(1, "")) + assert.Equal(t, "", isolateOpenAISessionID(1, " ")) + }) + + t.Run("deterministic", func(t *testing.T) { + a := isolateOpenAISessionID(42, "sess_abc123") + b := isolateOpenAISessionID(42, "sess_abc123") + assert.Equal(t, a, b) + }) + + t.Run("different_apiKeyID_different_result", func(t *testing.T) { + a := isolateOpenAISessionID(1, "same_session") + b := isolateOpenAISessionID(2, "same_session") + require.NotEqual(t, a, b, "不同 API Key 使用相同 session_id 应产生不同隔离值") + }) + + t.Run("different_raw_different_result", func(t *testing.T) { + a := isolateOpenAISessionID(1, "session_a") + b := isolateOpenAISessionID(1, "session_b") + require.NotEqual(t, a, b) + }) + + t.Run("format_is_16_hex_chars", func(t *testing.T) { + result := isolateOpenAISessionID(99, "test_session") + assert.Len(t, result, 16, "应为 16 字符的 hex 字符串") + for _, ch := range result { + assert.True(t, (ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'f'), + "应仅包含 hex 字符: %c", ch) + } + }) + + t.Run("zero_apiKeyID_still_works", func(t *testing.T) { + result := isolateOpenAISessionID(0, "session") + assert.NotEmpty(t, result) + // apiKeyID=0 与 apiKeyID=1 应产生不同结果 + other := isolateOpenAISessionID(1, "session") + assert.NotEqual(t, result, other) + }) +} diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index d4e4ea5aa7..b7ac84239a 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -1124,11 +1124,22 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( headers.Set("accept-language", v) } } - if sessionResolution.SessionID != "" { - headers.Set("session_id", sessionResolution.SessionID) - } - if sessionResolution.ConversationID != "" { - headers.Set("conversation_id", sessionResolution.ConversationID) + // OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。 + if account != nil && account.Type == AccountTypeOAuth { + apiKeyID := getAPIKeyIDFromContext(c) + if sessionResolution.SessionID != "" { + headers.Set("session_id", isolateOpenAISessionID(apiKeyID, sessionResolution.SessionID)) + } + if sessionResolution.ConversationID != "" { + headers.Set("conversation_id", isolateOpenAISessionID(apiKeyID, sessionResolution.ConversationID)) + } + } else { + if sessionResolution.SessionID != "" { + headers.Set("session_id", sessionResolution.SessionID) + } + if sessionResolution.ConversationID != "" { + headers.Set("conversation_id", sessionResolution.ConversationID) + } } if state := strings.TrimSpace(turnState); state != "" { headers.Set(openAIWSTurnStateHeader, state) diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index 912fade9d9..0d5004c070 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -454,8 +454,10 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T require.True(t, gjson.Get(requestJSON, "stream").Exists(), "WSv2 payload 应保留 stream 字段") require.True(t, gjson.Get(requestJSON, "stream").Bool(), "OAuth Codex 规范化后应强制 stream=true") require.Equal(t, openAIWSBetaV2Value, captureDialer.lastHeaders.Get("OpenAI-Beta")) - require.Equal(t, "sess-oauth-1", captureDialer.lastHeaders.Get("session_id")) - require.Equal(t, "conv-oauth-1", captureDialer.lastHeaders.Get("conversation_id")) + // OAuth 账号的 session_id/conversation_id 应被 isolateOpenAISessionID 隔离, + // 测试中未设置 api_key 到 context,apiKeyID=0。 + require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id")) + require.Equal(t, isolateOpenAISessionID(0, "conv-oauth-1"), captureDialer.lastHeaders.Get("conversation_id")) } func TestOpenAIGatewayService_Forward_WSv2_OAuthOriginatorCompatibility(t *testing.T) { @@ -596,7 +598,8 @@ func TestOpenAIGatewayService_Forward_WSv2_HeaderSessionFallbackFromPromptCacheK require.NotNil(t, result) require.Equal(t, "resp_prompt_cache_key", result.RequestID) - require.Equal(t, "pcache_123", captureDialer.lastHeaders.Get("session_id")) + // OAuth 账号的 session_id 应被 isolateOpenAISessionID 隔离(apiKeyID=0,未在 context 设置)。 + require.Equal(t, isolateOpenAISessionID(0, "pcache_123"), captureDialer.lastHeaders.Get("session_id")) require.Empty(t, captureDialer.lastHeaders.Get("conversation_id")) require.NotNil(t, captureConn.lastWrite) require.True(t, gjson.Get(requestToJSONString(captureConn.lastWrite), "stream").Exists()) From 3741617ebd5131921a9bf7e8a82d47fbce780891 Mon Sep 17 00:00:00 2001 From: QTom Date: Mon, 16 Mar 2026 10:27:57 +0800 Subject: [PATCH 07/15] =?UTF-8?q?fix(gateway):=20WS=20=E8=BF=9E=E6=8E=A5?= =?UTF-8?q?=E6=B1=A0=E6=9D=A1=E4=BB=B6=E5=BC=8F=20MarkBroken=20=E9=98=B2?= =?UTF-8?q?=E6=AD=A2=E8=B7=A8=E8=AF=B7=E6=B1=82=E4=B8=B2=E6=B5=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 正常终端事件(response.completed 等)退出后连接归还复用, 仅异常路径(读写错误、error 事件、客户端断连)MarkBroken 销毁。 Generate 模式: - 引入 cleanExit 标记,仅在 isTerminalEvent break 时设置 true - defer 中根据 cleanExit 决定是否 MarkBroken - 所有异常路径已在各自分支中提前调用 MarkBroken Ingress 模式: - 引入 lastTurnClean 标记,sendAndRelay 正常完成时设为 true - releaseSessionLease 根据 lastTurnClean 决定是否 MarkBroken - 错误路径重置 lastTurnClean = false - 客户端断连后 drain 仍保守 MarkBroken(L2916) --- .../internal/service/openai_ws_forwarder.go | 21 ++++++++++++++++--- .../openai_ws_forwarder_success_test.go | 9 ++++++-- 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index b7ac84239a..1d3d8fdffe 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -1870,7 +1870,16 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( } return nil, wrapOpenAIWSFallback(classifyOpenAIWSAcquireError(err), err) } - defer lease.Release() + // cleanExit 标记正常终端事件退出,此时上游不会再发送帧,连接可安全归还复用。 + // 所有异常路径(读写错误、error 事件等)已在各自分支中提前调用 MarkBroken, + // 因此 defer 中只需处理正常退出时不 MarkBroken 即可。 + cleanExit := false + defer func() { + if !cleanExit { + lease.MarkBroken() + } + lease.Release() + }() connID := strings.TrimSpace(lease.ConnID()) logOpenAIWSModeDebug( "connected account_id=%d account_type=%s transport=%s conn_id=%s conn_reused=%v conn_pick_ms=%d queue_wait_ms=%d has_previous_response_id=%v", @@ -2248,6 +2257,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( } if isTerminalEvent { + cleanExit = true break } } @@ -2983,12 +2993,15 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( pinnedSessionConnID = connID } } + // lastTurnClean 标记最后一轮 sendAndRelay 是否正常完成(收到终端事件且客户端未断连)。 + // 所有异常路径(读写错误、error 事件、客户端断连)已在各自分支或上层(L3403)中 MarkBroken, + // 因此 releaseSessionLease 中只需在非正常结束时 MarkBroken。 + lastTurnClean := false releaseSessionLease := func() { if sessionLease == nil { return } - if dedicatedMode { - // dedicated 会话结束后主动标记损坏,确保连接不会跨会话复用。 + if !lastTurnClean { sessionLease.MarkBroken() } unpinSessionConn(sessionConnID) @@ -3383,6 +3396,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( result, relayErr := sendAndRelay(turn, sessionLease, currentPayload, currentPayloadBytes, currentOriginalModel) if relayErr != nil { + lastTurnClean = false if recoverIngressPrevResponseNotFound(relayErr, turn, connID) { continue } @@ -3402,6 +3416,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( turnRetry = 0 turnPrevRecoveryTried = false lastTurnFinishedAt = time.Now() + lastTurnClean = true if hooks != nil && hooks.AfterTurn != nil { hooks.AfterTurn(turn, result, nil) } diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index 0d5004c070..7a76c38573 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -380,7 +380,8 @@ func TestOpenAIGatewayService_Forward_WSv2_PoolReuseNotOneToOne(t *testing.T) { require.True(t, strings.HasPrefix(result.RequestID, "resp_reuse_")) } - require.Equal(t, int64(1), upgradeCount.Load(), "多个客户端请求应复用账号连接池而不是 1:1 对等建链") + // 条件式 MarkBroken:正常终端事件退出后连接归还复用,不再无条件销毁。 + require.Equal(t, int64(1), upgradeCount.Load(), "正常完成后连接应归还复用,不应每次新建") metrics := svc.SnapshotOpenAIWSPoolMetrics() require.GreaterOrEqual(t, metrics.AcquireReuseTotal, int64(1)) require.GreaterOrEqual(t, metrics.ConnPickTotal, int64(1)) @@ -964,6 +965,10 @@ func TestOpenAIGatewayService_Forward_WSv2_TurnMetadataInPayloadOnConnReuse(t *t require.NotNil(t, result1) require.Equal(t, "resp_meta_1", result1.RequestID) + require.Len(t, captureConn.writes, 1) + firstWrite := requestToJSONString(captureConn.writes[0]) + require.Equal(t, "turn_meta_payload_1", gjson.Get(firstWrite, "client_metadata.x-codex-turn-metadata").String()) + rec2 := httptest.NewRecorder() c2, _ := gin.CreateTestContext(rec2) c2.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) @@ -977,7 +982,7 @@ func TestOpenAIGatewayService_Forward_WSv2_TurnMetadataInPayloadOnConnReuse(t *t require.Equal(t, 1, captureDialer.DialCount(), "同一账号两轮请求应复用同一 WS 连接") require.Len(t, captureConn.writes, 2) - firstWrite := requestToJSONString(captureConn.writes[0]) + firstWrite = requestToJSONString(captureConn.writes[0]) secondWrite := requestToJSONString(captureConn.writes[1]) require.Equal(t, "turn_meta_payload_1", gjson.Get(firstWrite, "client_metadata.x-codex-turn-metadata").String()) require.Equal(t, "turn_meta_payload_2", gjson.Get(secondWrite, "client_metadata.x-codex-turn-metadata").String()) From 67c050629006e8260eca16fc231ca508ac9d9312 Mon Sep 17 00:00:00 2001 From: erio Date: Mon, 16 Mar 2026 13:39:50 +0800 Subject: [PATCH 08/15] fix(billing): add window expiration check to Redis rate limit Lua script The updateRateLimitUsageScript Lua script previously performed unconditional HINCRBYFLOAT on all usage counters without checking whether the rate limit window had expired. This caused usage to accumulate across window boundaries in Redis while the DB correctly reset on expiration, leading to incorrect 429 rate limiting that could persist for up to 24 hours. The Lua script now checks each window timestamp before incrementing: - If the window has expired, usage is reset to the current cost and the window timestamp is updated (matching DB-side semantics) - If the window is still valid, usage is accumulated normally This also resolves the async race condition where stale HINCRBYFLOAT tasks from the worker queue could pollute a freshly rebuilt cache after invalidation, since the script now self-corrects expired windows. Closes #1049 --- backend/internal/repository/billing_cache.go | 48 +++++++++++++++++--- 1 file changed, 42 insertions(+), 6 deletions(-) diff --git a/backend/internal/repository/billing_cache.go b/backend/internal/repository/billing_cache.go index 4fbdae14fd..6922b4c8bb 100644 --- a/backend/internal/repository/billing_cache.go +++ b/backend/internal/repository/billing_cache.go @@ -20,6 +20,11 @@ const ( billingCacheTTL = 5 * time.Minute billingCacheJitter = 30 * time.Second rateLimitCacheTTL = 7 * 24 * time.Hour // 7 days matches the longest window + + // Rate limit window durations — must match service.RateLimitWindow* constants. + rateLimitWindow5h = 5 * time.Hour + rateLimitWindow1d = 24 * time.Hour + rateLimitWindow7d = 7 * 24 * time.Hour ) // jitteredTTL 返回带随机抖动的 TTL,防止缓存雪崩 @@ -90,17 +95,40 @@ var ( return 1 `) - // updateRateLimitUsageScript atomically increments all three rate limit usage counters. - // Returns 0 if the key doesn't exist (cache miss), 1 on success. + // updateRateLimitUsageScript atomically increments all three rate limit usage counters + // with window expiration checking. If a window has expired, its usage is reset to cost + // (instead of accumulated) and the window timestamp is updated, matching the DB-side + // IncrementRateLimitUsage semantics. + // + // ARGV: [1]=cost, [2]=ttl_seconds, [3]=now_unix, [4]=window_5h_seconds, [5]=window_1d_seconds, [6]=window_7d_seconds updateRateLimitUsageScript = redis.NewScript(` local exists = redis.call('EXISTS', KEYS[1]) if exists == 0 then return 0 end local cost = tonumber(ARGV[1]) - redis.call('HINCRBYFLOAT', KEYS[1], 'usage_5h', cost) - redis.call('HINCRBYFLOAT', KEYS[1], 'usage_1d', cost) - redis.call('HINCRBYFLOAT', KEYS[1], 'usage_7d', cost) + local now = tonumber(ARGV[3]) + local win5h = tonumber(ARGV[4]) + local win1d = tonumber(ARGV[5]) + local win7d = tonumber(ARGV[6]) + + -- Helper: check if window is expired and update usage + window accordingly + -- Returns nothing, modifies the hash in-place. + local function update_window(usage_field, window_field, window_duration) + local w = tonumber(redis.call('HGET', KEYS[1], window_field) or 0) + if w == 0 or (now - w) >= window_duration then + -- Window expired or never started: reset usage to cost, start new window + redis.call('HSET', KEYS[1], usage_field, tostring(cost)) + redis.call('HSET', KEYS[1], window_field, tostring(now)) + else + -- Window still valid: accumulate + redis.call('HINCRBYFLOAT', KEYS[1], usage_field, cost) + end + end + + update_window('usage_5h', 'window_5h', win5h) + update_window('usage_1d', 'window_1d', win1d) + update_window('usage_7d', 'window_7d', win7d) redis.call('EXPIRE', KEYS[1], ARGV[2]) return 1 `) @@ -280,7 +308,15 @@ func (c *billingCache) SetAPIKeyRateLimit(ctx context.Context, keyID int64, data func (c *billingCache) UpdateAPIKeyRateLimitUsage(ctx context.Context, keyID int64, cost float64) error { key := billingRateLimitKey(keyID) - _, err := updateRateLimitUsageScript.Run(ctx, c.rdb, []string{key}, cost, int(rateLimitCacheTTL.Seconds())).Result() + now := time.Now().Unix() + _, err := updateRateLimitUsageScript.Run(ctx, c.rdb, []string{key}, + cost, + int(rateLimitCacheTTL.Seconds()), + now, + int(rateLimitWindow5h.Seconds()), + int(rateLimitWindow1d.Seconds()), + int(rateLimitWindow7d.Seconds()), + ).Result() if err != nil && !errors.Is(err, redis.Nil) { log.Printf("Warning: update rate limit usage cache failed for api key %d: %v", keyID, err) return err From 71f72e167eb00c68ce46391f9d98f24ccfde9c0f Mon Sep 17 00:00:00 2001 From: erio Date: Mon, 16 Mar 2026 15:47:32 +0800 Subject: [PATCH 09/15] chore(antigravity): bump default User-Agent version to 1.20.5 --- backend/internal/pkg/antigravity/oauth.go | 4 ++-- backend/internal/pkg/antigravity/oauth_test.go | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/backend/internal/pkg/antigravity/oauth.go b/backend/internal/pkg/antigravity/oauth.go index 5bda31aca2..8a8bed92d9 100644 --- a/backend/internal/pkg/antigravity/oauth.go +++ b/backend/internal/pkg/antigravity/oauth.go @@ -49,8 +49,8 @@ const ( antigravityDailyBaseURL = "https://daily-cloudcode-pa.sandbox.googleapis.com" ) -// defaultUserAgentVersion 可通过环境变量 ANTIGRAVITY_USER_AGENT_VERSION 配置,默认 1.20.4 -var defaultUserAgentVersion = "1.20.4" +// defaultUserAgentVersion 可通过环境变量 ANTIGRAVITY_USER_AGENT_VERSION 配置,默认 1.20.5 +var defaultUserAgentVersion = "1.20.5" // defaultClientSecret 可通过环境变量 ANTIGRAVITY_OAUTH_CLIENT_SECRET 配置 var defaultClientSecret = "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" diff --git a/backend/internal/pkg/antigravity/oauth_test.go b/backend/internal/pkg/antigravity/oauth_test.go index f4630b093e..3a093fe657 100644 --- a/backend/internal/pkg/antigravity/oauth_test.go +++ b/backend/internal/pkg/antigravity/oauth_test.go @@ -690,7 +690,7 @@ func TestConstants_值正确(t *testing.T) { if RedirectURI != "http://localhost:8085/callback" { t.Errorf("RedirectURI 不匹配: got %s", RedirectURI) } - if GetUserAgent() != "antigravity/1.20.4 windows/amd64" { + if GetUserAgent() != "antigravity/1.20.5 windows/amd64" { t.Errorf("UserAgent 不匹配: got %s", GetUserAgent()) } if SessionTTL != 30*time.Minute { From afd72abc6ed765cda52e2c6e5876df5d28e5e5ee Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Mon, 16 Mar 2026 16:22:31 +0800 Subject: [PATCH 10/15] fix: allow empty extra payload to clear account quota limits UpdateAccount previously required len(input.Extra) > 0, causing explicit empty payloads (extra:{}) to be silently skipped. Change condition to input.Extra != nil so clearing quota keys actually persists. --- backend/internal/service/admin_service.go | 4 ++- .../service/admin_service_overages_test.go | 32 +++++++++++++++++++ 2 files changed, 35 insertions(+), 1 deletion(-) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index ea76e17199..5eeac18329 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -1530,7 +1530,9 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if len(input.Credentials) > 0 { account.Credentials = input.Credentials } - if len(input.Extra) > 0 { + // Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。 + // 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。 + if input.Extra != nil { // 保留配额用量字段,防止编辑账号时意外重置 for _, key := range []string{"quota_used", "quota_daily_used", "quota_daily_start", "quota_weekly_used", "quota_weekly_start"} { if v, ok := account.Extra[key]; ok { diff --git a/backend/internal/service/admin_service_overages_test.go b/backend/internal/service/admin_service_overages_test.go index 779b08b96d..d6380f4dcd 100644 --- a/backend/internal/service/admin_service_overages_test.go +++ b/backend/internal/service/admin_service_overages_test.go @@ -121,3 +121,35 @@ func TestUpdateAccount_EnableOveragesClearsModelRateLimitsBeforePersist(t *testi _, exists := repo.account.Extra[modelRateLimitsKey] require.False(t, exists, "开启 overages 时应在持久化前清掉旧模型限流") } + +func TestUpdateAccount_EmptyExtraPayloadCanClearQuotaLimits(t *testing.T) { + accountID := int64(103) + repo := &updateAccountOveragesRepoStub{ + account: &Account{ + ID: accountID, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{ + "quota_limit": 100.0, + "quota_daily_limit": 10.0, + "quota_weekly_limit": 40.0, + }, + }, + } + + svc := &adminServiceImpl{accountRepo: repo} + updated, err := svc.UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + // 显式空对象:语义是“清空 extra 中的可配置键”(例如关闭配额限制) + Extra: map[string]any{}, + }) + + require.NoError(t, err) + require.NotNil(t, updated) + require.Equal(t, 1, repo.updateCalls) + require.NotNil(t, repo.account.Extra) + require.NotContains(t, repo.account.Extra, "quota_limit") + require.NotContains(t, repo.account.Extra, "quota_daily_limit") + require.NotContains(t, repo.account.Extra, "quota_weekly_limit") + require.Len(t, repo.account.Extra, 0) +} From fa782e70a43916e374bd53444208fdf52f7bca8d Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Mon, 16 Mar 2026 16:22:42 +0800 Subject: [PATCH 11/15] fix: always attach OpenAI 5h/7d window stats regardless of zero values Removes hasMeaningfulWindowStats guard so the /usage endpoint consistently returns WindowStats for both time windows. The frontend now controls zero-value display filtering at the component level. --- .../internal/service/account_usage_service.go | 25 +++++-------------- 1 file changed, 6 insertions(+), 19 deletions(-) diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index f117abfda6..959c11827e 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -446,23 +446,17 @@ func (s *AccountUsageService) getOpenAIUsage(ctx context.Context, account *Accou } if stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, account.ID, now.Add(-5*time.Hour)); err == nil { - windowStats := windowStatsFromAccountStats(stats) - if hasMeaningfulWindowStats(windowStats) { - if usage.FiveHour == nil { - usage.FiveHour = &UsageProgress{Utilization: 0} - } - usage.FiveHour.WindowStats = windowStats + if usage.FiveHour == nil { + usage.FiveHour = &UsageProgress{Utilization: 0} } + usage.FiveHour.WindowStats = windowStatsFromAccountStats(stats) } if stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, account.ID, now.Add(-7*24*time.Hour)); err == nil { - windowStats := windowStatsFromAccountStats(stats) - if hasMeaningfulWindowStats(windowStats) { - if usage.SevenDay == nil { - usage.SevenDay = &UsageProgress{Utilization: 0} - } - usage.SevenDay.WindowStats = windowStats + if usage.SevenDay == nil { + usage.SevenDay = &UsageProgress{Utilization: 0} } + usage.SevenDay.WindowStats = windowStatsFromAccountStats(stats) } return usage, nil @@ -992,13 +986,6 @@ func windowStatsFromAccountStats(stats *usagestats.AccountStats) *WindowStats { } } -func hasMeaningfulWindowStats(stats *WindowStats) bool { - if stats == nil { - return false - } - return stats.Requests > 0 || stats.Tokens > 0 || stats.Cost > 0 || stats.StandardCost > 0 || stats.UserCost > 0 -} - func buildCodexUsageProgressFromExtra(extra map[string]any, window string, now time.Time) *UsageProgress { if len(extra) == 0 { return nil From 8640a62319e1764a9ec29465f568ca203c698620 Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Mon, 16 Mar 2026 16:22:51 +0800 Subject: [PATCH 12/15] refactor: extract formatCompactNumber util and add last_used_at to refresh key - Add formatCompactNumber() for consistent large-number formatting (K/M/B) - Include last_used_at in OpenAI usage refresh key for better change detection - Add .gitattributes eol=lf rules for frontend source files --- .gitattributes | 7 ++++++ .../__tests__/accountUsageRefresh.spec.ts | 24 +++++++++++++++++++ .../__tests__/formatCompactNumber.spec.ts | 22 +++++++++++++++++ frontend/src/utils/accountUsageRefresh.ts | 3 ++- frontend/src/utils/format.ts | 20 ++++++++++++++++ 5 files changed, 75 insertions(+), 1 deletion(-) create mode 100644 frontend/src/utils/__tests__/formatCompactNumber.spec.ts diff --git a/.gitattributes b/.gitattributes index 3db3b83dc7..37e3bee2cc 100644 --- a/.gitattributes +++ b/.gitattributes @@ -4,6 +4,13 @@ backend/migrations/*.sql text eol=lf # Go 源代码文件 *.go text eol=lf +# 前端 源代码文件 +*.ts text eol=lf +*.tsx text eol=lf +*.js text eol=lf +*.jsx text eol=lf +*.vue text eol=lf + # Shell 脚本 *.sh text eol=lf diff --git a/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts b/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts index ae13d690e5..aef73b0fa4 100644 --- a/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts +++ b/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts @@ -8,6 +8,7 @@ describe('buildOpenAIUsageRefreshKey', () => { platform: 'openai', type: 'oauth', updated_at: '2026-03-07T10:00:00Z', + last_used_at: '2026-03-07T09:59:00Z', extra: { codex_usage_updated_at: '2026-03-07T10:00:00Z', codex_5h_used_percent: 0, @@ -27,12 +28,35 @@ describe('buildOpenAIUsageRefreshKey', () => { expect(buildOpenAIUsageRefreshKey(base)).not.toBe(buildOpenAIUsageRefreshKey(next)) }) + it('会在 last_used_at 变化时生成不同 key', () => { + const base = { + id: 3, + platform: 'openai', + type: 'oauth', + updated_at: '2026-03-07T10:00:00Z', + last_used_at: '2026-03-07T10:00:00Z', + extra: { + codex_usage_updated_at: '2026-03-07T10:00:00Z', + codex_5h_used_percent: 12, + codex_7d_used_percent: 24 + } + } as any + + const next = { + ...base, + last_used_at: '2026-03-07T10:02:00Z' + } + + expect(buildOpenAIUsageRefreshKey(base)).not.toBe(buildOpenAIUsageRefreshKey(next)) + }) + it('非 OpenAI OAuth 账号返回空 key', () => { expect(buildOpenAIUsageRefreshKey({ id: 2, platform: 'anthropic', type: 'oauth', updated_at: '2026-03-07T10:00:00Z', + last_used_at: '2026-03-07T10:00:00Z', extra: {} } as any)).toBe('') }) diff --git a/frontend/src/utils/__tests__/formatCompactNumber.spec.ts b/frontend/src/utils/__tests__/formatCompactNumber.spec.ts new file mode 100644 index 0000000000..a5a9ed9f37 --- /dev/null +++ b/frontend/src/utils/__tests__/formatCompactNumber.spec.ts @@ -0,0 +1,22 @@ +import { describe, expect, it } from 'vitest' +import { formatCompactNumber } from '../format' + +describe('formatCompactNumber', () => { + it('formats boundary values with K/M/B', () => { + expect(formatCompactNumber(0)).toBe('0') + expect(formatCompactNumber(999)).toBe('999') + expect(formatCompactNumber(1000)).toBe('1.0K') + expect(formatCompactNumber(999999)).toBe('1000.0K') + expect(formatCompactNumber(1000000)).toBe('1.0M') + expect(formatCompactNumber(1000000000)).toBe('1.0B') + }) + + it('supports disabling billion unit (requests style)', () => { + expect(formatCompactNumber(1000000000, { allowBillions: false })).toBe('1000.0M') + }) + + it('returns 0 for nullish input', () => { + expect(formatCompactNumber(null)).toBe('0') + expect(formatCompactNumber(undefined)).toBe('0') + }) +}) diff --git a/frontend/src/utils/accountUsageRefresh.ts b/frontend/src/utils/accountUsageRefresh.ts index 219ac57f76..3406c7a504 100644 --- a/frontend/src/utils/accountUsageRefresh.ts +++ b/frontend/src/utils/accountUsageRefresh.ts @@ -5,7 +5,7 @@ const normalizeUsageRefreshValue = (value: unknown): string => { return String(value) } -export const buildOpenAIUsageRefreshKey = (account: Pick): string => { +export const buildOpenAIUsageRefreshKey = (account: Pick): string => { if (account.platform !== 'openai' || account.type !== 'oauth') { return '' } @@ -14,6 +14,7 @@ export const buildOpenAIUsageRefreshKey = (account: Pick= 1_000_000_000) return `${(num / 1_000_000_000).toFixed(1)}B` + if (abs >= 1_000_000) return `${(num / 1_000_000).toFixed(1)}M` + if (abs >= 1_000) return `${(num / 1_000).toFixed(1)}K` + return num.toString() +} + /** * 格式化倒计时(从现在到目标时间的剩余时间) * @param targetDate 目标日期字符串或 Date 对象 From fbffb08aae4a8bea2db6a6392c9ef30f777a14d0 Mon Sep 17 00:00:00 2001 From: Ethan0x0000 <3352979663@qq.com> Date: Mon, 16 Mar 2026 16:23:00 +0800 Subject: [PATCH 13/15] feat: add today-stats and manual refresh token propagation to usage cells - Pass todayStats/todayStatsLoading to AccountUsageCell for key accounts - Propagate usageManualRefreshToken to force usage reload on explicit refresh - Refresh today stats when toggling usage/today_stats columns visible --- frontend/src/views/admin/AccountsView.vue | 39 +++++++++++++++++------ 1 file changed, 30 insertions(+), 9 deletions(-) diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index dd342a5b8b..2ec5b47d37 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -203,7 +203,12 @@ @@ -423,12 +442,23 @@ import { adminAPI } from '@/api/admin' import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from '@/types' import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { resolveCodexUsageWindow } from '@/utils/codexUsage' +import { formatCompactNumber } from '@/utils/format' import UsageProgressBar from './UsageProgressBar.vue' import AccountQuotaInfo from './AccountQuotaInfo.vue' -const props = defineProps<{ - account: Account -}>() +const props = withDefaults( + defineProps<{ + account: Account + todayStats?: WindowStats | null + todayStatsLoading?: boolean + manualRefreshToken?: number + }>(), + { + todayStats: null, + todayStatsLoading: false, + manualRefreshToken: 0 + } +) const { t } = useI18n() @@ -490,26 +520,9 @@ const isActiveOpenAIRateLimited = computed(() => { return !Number.isNaN(resetAt) && resetAt > Date.now() }) -const preferFetchedOpenAIUsage = computed(() => { - return (isActiveOpenAIRateLimited.value || isOpenAICodexSnapshotStale.value) && hasOpenAIUsageFallback.value -}) - const openAIUsageRefreshKey = computed(() => buildOpenAIUsageRefreshKey(props.account)) -const isOpenAICodexSnapshotStale = computed(() => { - if (props.account.platform !== 'openai' || props.account.type !== 'oauth') return false - const extra = props.account.extra as Record | undefined - const updatedAtRaw = extra?.codex_usage_updated_at - if (!updatedAtRaw) return true - const updatedAt = Date.parse(String(updatedAtRaw)) - if (Number.isNaN(updatedAt)) return true - return Date.now() - updatedAt >= 10 * 60 * 1000 -}) - const shouldAutoLoadUsageOnMount = computed(() => { - if (props.account.platform === 'openai' && props.account.type === 'oauth') { - return isActiveOpenAIRateLimited.value || !hasCodexUsage.value || isOpenAICodexSnapshotStale.value - } return shouldFetchUsage.value }) @@ -1006,6 +1019,28 @@ const quotaTotalBar = computed((): QuotaBarInfo | null => { return makeQuotaBar(props.account.quota_used ?? 0, limit) }) +// ===== Key account today stats formatters ===== + +const formatKeyRequests = computed(() => { + if (!props.todayStats) return '' + return formatCompactNumber(props.todayStats.requests, { allowBillions: false }) +}) + +const formatKeyTokens = computed(() => { + if (!props.todayStats) return '' + return formatCompactNumber(props.todayStats.tokens) +}) + +const formatKeyCost = computed(() => { + if (!props.todayStats) return '0.00' + return props.todayStats.cost.toFixed(2) +}) + +const formatKeyUserCost = computed(() => { + if (!props.todayStats || props.todayStats.user_cost == null) return '0.00' + return props.todayStats.user_cost.toFixed(2) +}) + onMounted(() => { if (!shouldAutoLoadUsageOnMount.value) return loadUsage() @@ -1014,10 +1049,21 @@ onMounted(() => { watch(openAIUsageRefreshKey, (nextKey, prevKey) => { if (!prevKey || nextKey === prevKey) return if (props.account.platform !== 'openai' || props.account.type !== 'oauth') return - if (!isActiveOpenAIRateLimited.value && hasCodexUsage.value && !isOpenAICodexSnapshotStale.value) return loadUsage().catch((e) => { console.error('Failed to refresh OpenAI usage:', e) }) }) + +watch( + () => props.manualRefreshToken, + (nextToken, prevToken) => { + if (nextToken === prevToken) return + if (!shouldFetchUsage.value) return + + loadUsage().catch((e) => { + console.error('Failed to refresh usage after manual refresh:', e) + }) + } +) diff --git a/frontend/src/components/account/UsageProgressBar.vue b/frontend/src/components/account/UsageProgressBar.vue index cd5c991f10..5ce8bfe07a 100644 --- a/frontend/src/components/account/UsageProgressBar.vue +++ b/frontend/src/components/account/UsageProgressBar.vue @@ -2,7 +2,7 @@
@@ -12,12 +12,13 @@ {{ formatTokens }} - + A ${{ formatAccountCost }} U ${{ formatUserCost }} @@ -56,7 +57,9 @@