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/README.md b/README.md index 4a7bde8ef0..4f2d4c511d 100644 --- a/README.md +++ b/README.md @@ -8,27 +8,31 @@ [![Redis](https://img.shields.io/badge/Redis-7+-DC382D.svg)](https://redis.io/) [![Docker](https://img.shields.io/badge/Docker-Ready-2496ED.svg)](https://www.docker.com/) +Wei-Shaw%2Fsub2api | Trendshift + **AI API Gateway Platform for Subscription Quota Distribution** English | [中文](README_CN.md) +> **Sub2API officially uses only the domains `sub2api.org` and `pincc.ai`. Other websites using the Sub2API name may be third-party deployments or services and are not affiliated with this project. Please verify and exercise your own judgment.** + --- ## Demo -Try Sub2API online: **https://demo.sub2api.org/** +Try Sub2API online: **[https://demo.sub2api.org/](https://demo.sub2api.org/)** Demo credentials (shared demo environment; **not** created automatically for self-hosted installs): | Email | Password | |-------|----------| -| admin@sub2api.com | admin123 | +| admin@sub2api.org | admin123 | ## Overview -Sub2API is an AI API gateway platform designed to distribute and manage API quotas from AI product subscriptions (like Claude Code $200/month). Users can access upstream AI services through platform-generated API Keys, while the platform handles authentication, billing, load balancing, and request forwarding. +Sub2API is an AI API gateway platform designed to distribute and manage API quotas from AI product subscriptions. Users can access upstream AI services through platform-generated API Keys, while the platform handles authentication, billing, load balancing, and request forwarding. ## Features @@ -41,6 +45,15 @@ Sub2API is an AI API gateway platform designed to distribute and manage API quot - **Admin Dashboard** - Web interface for monitoring and management - **External System Integration** - Embed external systems (e.g. payment, ticketing) via iframe to extend the admin dashboard +## Don't Want to Self-Host? + + + + + + +
pinccPinCC is the official relay service built on Sub2API, offering stable access to Claude Code, Codex, Gemini and other popular models — ready to use, no deployment or maintenance required.
+ ## Ecosystem Community projects that extend or integrate with Sub2API: @@ -61,10 +74,15 @@ Community projects that extend or integrate with Sub2API: --- -## Documentation +## Nginx Reverse Proxy Note -- Dependency Security: `docs/dependency-security.md` -- Admin Payment Integration API: `docs/ADMIN_PAYMENT_INTEGRATION_API.md` +When using Nginx as a reverse proxy for Sub2API (or CRS) with Codex CLI, add the following to the `http` block in your Nginx configuration: + +```nginx +underscores_in_headers on; +``` + +Nginx drops headers containing underscores by default (e.g. `session_id`), which breaks sticky session routing in multi-account setups. --- diff --git a/README_CN.md b/README_CN.md index eee89b074d..849f384094 100644 --- a/README_CN.md +++ b/README_CN.md @@ -8,27 +8,30 @@ [![Redis](https://img.shields.io/badge/Redis-7+-DC382D.svg)](https://redis.io/) [![Docker](https://img.shields.io/badge/Docker-Ready-2496ED.svg)](https://www.docker.com/) +Wei-Shaw%2Fsub2api | Trendshift + **AI API 网关平台 - 订阅配额分发管理** [English](README.md) | 中文 +> **Sub2API 官方仅使用 `sub2api.org` 与 `pincc.ai` 两个域名。其他使用 Sub2API 名义的网站可能为第三方部署或服务,与本项目无关,请自行甄别。** --- ## 在线体验 -体验地址:**https://v2.pincc.ai/** +体验地址:**[https://demo.sub2api.org/](https://demo.sub2api.org/)** 演示账号(共享演示环境;自建部署不会自动创建该账号): | 邮箱 | 密码 | |------|------| -| admin@sub2api.com | admin123 | +| admin@sub2api.org | admin123 | ## 项目概述 -Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅(如 Claude Code $200/月)的 API 配额。用户通过平台生成的 API Key 调用上游 AI 服务,平台负责鉴权、计费、负载均衡和请求转发。 +Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅的 API 配额。用户通过平台生成的 API Key 调用上游 AI 服务,平台负责鉴权、计费、负载均衡和请求转发。 ## 核心功能 @@ -41,6 +44,15 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅( - **管理后台** - Web 界面进行监控和管理 - **外部系统集成** - 支持通过 iframe 嵌入外部系统(如支付、工单等),扩展管理后台功能 +## 不想自建?试试官方中转 + + + + + + +
pinccPinCC 是基于 Sub2API 搭建的官方中转服务,提供 Claude Code、Codex、Gemini 等主流模型的稳定中转,开箱即用,免去自建部署与运维烦恼。
+ ## 生态项目 围绕 Sub2API 的社区扩展与集成项目: @@ -61,17 +73,18 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅( --- -## 文档 +## Nginx 反向代理注意事项 -- 依赖安全:`docs/dependency-security.md` +通过 Nginx 反向代理 Sub2API(或 CRS 服务)并搭配 Codex CLI 使用时,需要在 Nginx 配置的 `http` 块中添加: + +```nginx +underscores_in_headers on; +``` + +Nginx 默认会丢弃名称中含下划线的请求头(如 `session_id`),这会导致多账号环境下的粘性会话功能失效。 --- -## OpenAI Responses 兼容注意事项 - -- 当请求包含 `function_call_output` 时,需要携带 `previous_response_id`,或在 `input` 中包含带 `call_id` 的 `tool_call`/`function_call`,或带非空 `id` 且与 `function_call_output.call_id` 匹配的 `item_reference`。 -- 若依赖上游历史记录,网关会强制 `store=true` 并需要复用 `previous_response_id`,以避免出现 “No tool call found for function call output” 错误。 - ## 部署方式 ### 方式一:脚本安装(推荐) diff --git a/assets/partners/logos/pincc-logo.png b/assets/partners/logos/pincc-logo.png new file mode 100644 index 0000000000..081b6c8463 Binary files /dev/null and b/assets/partners/logos/pincc-logo.png differ 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/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/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) +} diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index f73ceba1dd..cb90f49ba9 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) @@ -456,6 +458,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, @@ -759,6 +763,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) @@ -773,6 +779,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, @@ -938,7 +946,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/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/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 { diff --git a/backend/internal/pkg/usagestats/usage_log_types.go b/backend/internal/pkg/usagestats/usage_log_types.go index e9a5cae5af..99c9cda7c5 100644 --- a/backend/internal/pkg/usagestats/usage_log_types.go +++ b/backend/internal/pkg/usagestats/usage_log_types.go @@ -125,6 +125,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/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 diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index cf2d1ca57e..002281ec21 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -2161,7 +2161,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 @@ -2172,7 +2174,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 ` @@ -2190,9 +2194,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) @@ -2204,6 +2210,8 @@ func (r *usageLogRepository) GetUserSpendingRanking(ctx context.Context, startTi return &UserSpendingRankingResponse{ Ranking: ranking, TotalActualCost: totalActualCost, + TotalRequests: totalRequests, + TotalTokens: totalTokens, }, nil } @@ -3004,7 +3012,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{} 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 70b82bb079..27ae457181 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -259,10 +259,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). @@ -277,6 +277,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/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 7dd089cbce..f4bd810198 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) 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 diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index ad0c81ef8a..0c4c3072fc 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -1549,7 +1549,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) +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 6b82ef14b9..1204e71314 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -7405,6 +7405,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 误复用时的静默误去重风险 @@ -7813,6 +7815,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, @@ -7893,6 +7897,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 误复用时的静默误去重风险 @@ -7990,6 +7996,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, 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 60f8e2a464..ddce6b8ed0 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..1d3d8fdffe 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) @@ -1859,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", @@ -2237,6 +2257,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( } if isTerminalEvent { + cleanExit = true break } } @@ -2972,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) @@ -3372,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 } @@ -3391,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 912fade9d9..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)) @@ -454,8 +455,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 +599,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()) @@ -961,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) @@ -974,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()) diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index f9bf6ae54a..cc27e0cce3 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -75,7 +75,7 @@ @@ -389,8 +371,43 @@
- -
+ +
+ +
+
+ + {{ formatKeyRequests }} req + + + {{ formatKeyTokens }} + + + A ${{ formatKeyCost }} + + + U ${{ formatKeyUserCost }} + +
+
+ +
+
+
+
+
+ + + + +
-
-
-
@@ -427,12 +446,23 @@ import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from ' import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { resolveCodexUsageWindow } from '@/utils/codexUsage' import { enqueueUsageRequest } from '@/utils/usageLoadQueue' +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() @@ -497,26 +527,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 }) @@ -1021,6 +1034,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() @@ -1029,10 +1064,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 @@