From ad8afc8a2e984118f25aa8aa269de536c760ca57 Mon Sep 17 00:00:00 2001 From: CHOS1N Date: Thu, 9 Jul 2026 22:52:35 +0800 Subject: [PATCH 01/27] Add parallel_tool_calls compatibility mapping Preserve Chat Completions parallel_tool_calls when converting requests to the Responses API, and map Responses parallel_tool_calls back when falling back to Chat Completions upstreams. Cover both true and explicit false values so clients can disable parallel tool calls without the field being dropped by omitempty. --- .../chatcompletions_responses_bridge.go | 1 + .../chatcompletions_responses_bridge_test.go | 20 +++++++++++++++++++ .../chatcompletions_responses_test.go | 19 ++++++++++++++++++ .../apicompat/chatcompletions_to_responses.go | 13 ++++++------ backend/internal/pkg/apicompat/types.go | 1 + 5 files changed, 48 insertions(+), 6 deletions(-) diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go index eeeedd29aa..2043246de0 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go @@ -28,6 +28,7 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR TopP: req.TopP, Stream: req.Stream, ServiceTier: req.ServiceTier, + ParallelToolCalls: req.ParallelToolCalls, } if req.Reasoning != nil { out.ReasoningEffort = req.Reasoning.Effort diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_test.go index b194d88141..2e770b4158 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_test.go @@ -127,6 +127,26 @@ func TestResponsesToChatCompletionsRequest_TextFormatJsonSchema(t *testing.T) { }`, string(out.ResponseFormat)) } +func TestResponsesToChatCompletionsRequest_ParallelToolCalls(t *testing.T) { + parallel := false + req := &ResponsesRequest{ + Model: "gpt-4o", + Input: json.RawMessage(`[ + {"role":"user","content":"Use tools"} + ]`), + ParallelToolCalls: ¶llel, + } + + out, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.NotNil(t, out.ParallelToolCalls) + assert.False(t, *out.ParallelToolCalls) + + payload, err := json.Marshal(out) + require.NoError(t, err) + assert.Contains(t, string(payload), `"parallel_tool_calls":false`) +} + func chatMessageRoles(messages []ChatMessage) []string { roles := make([]string, 0, len(messages)) for _, message := range messages { diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index b30330863c..8772ec0f4e 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -531,6 +531,25 @@ func TestChatCompletionsToResponses_ServiceTier(t *testing.T) { assert.Equal(t, "flex", resp.ServiceTier) } +func TestChatCompletionsToResponses_ParallelToolCalls(t *testing.T) { + for _, value := range []bool{false, true} { + req := &ChatCompletionsRequest{ + Model: "gpt-4o", + ParallelToolCalls: &value, + Messages: []ChatMessage{{Role: "user", Content: json.RawMessage(`"Hi"`)}}, + } + + resp, err := ChatCompletionsToResponses(req) + require.NoError(t, err) + require.NotNil(t, resp.ParallelToolCalls) + assert.Equal(t, value, *resp.ParallelToolCalls) + + payload, err := json.Marshal(resp) + require.NoError(t, err) + assert.Contains(t, string(payload), `"parallel_tool_calls":`+string(mustMarshalJSON(t, value))) + } +} + // --------------------------------------------------------------------------- // temperature / top_p stripping for reasoning models // --------------------------------------------------------------------------- diff --git a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go index 07c557ab3b..0f65d217e8 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go +++ b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go @@ -27,12 +27,13 @@ func ChatCompletionsToResponses(req *ChatCompletionsRequest) (*ResponsesRequest, } out := &ResponsesRequest{ - Model: req.Model, - Instructions: req.Instructions, - Input: inputJSON, - Stream: true, // upstream always streams - Include: []string{"reasoning.encrypted_content"}, - ServiceTier: req.ServiceTier, + Model: req.Model, + Instructions: req.Instructions, + Input: inputJSON, + Stream: true, // upstream always streams + Include: []string{"reasoning.encrypted_content"}, + ServiceTier: req.ServiceTier, + ParallelToolCalls: req.ParallelToolCalls, } // Reasoning models (gpt-5.x) do not accept sampling parameters. diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index a0fd07a0d1..f7955c70c0 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -435,6 +435,7 @@ type ChatCompletionsRequest struct { Stream bool `json:"stream,omitempty"` StreamOptions *ChatStreamOptions `json:"stream_options,omitempty"` Tools []ChatTool `json:"tools,omitempty"` + ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"` ToolChoice json.RawMessage `json:"tool_choice,omitempty"` ReasoningEffort string `json:"reasoning_effort,omitempty"` // "low" | "medium" | "high" | "xhigh" ServiceTier string `json:"service_tier,omitempty"` From 99da3081961f983f865290aa4657bfaae3530a8a Mon Sep 17 00:00:00 2001 From: Shumin <332587268@qq.com> Date: Thu, 9 Jul 2026 11:05:05 -0400 Subject: [PATCH 02/27] =?UTF-8?q?fix:=20=E5=90=8E=E5=8F=B0=E8=87=AA?= =?UTF-8?q?=E5=8A=A8=E5=88=B7=E6=96=B0=E7=BA=B3=E5=85=A5=20setup-token=20?= =?UTF-8?q?=E8=B4=A6=E5=8F=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit setup-token 的 access_token 同为 8h 短期令牌(expires_in=28800), 此前被后台刷新服务排除,到期后请求返回 401 authentication_error。 放开 CanRefresh 与候选查询的 type='oauth' 限制,与手动刷新入口 (account.IsOAuth()) 保持一致;实际刷新仍由 NeedsRefresh 基于 expires_at 门控并在分布式锁下执行,不会过度刷新。 --- backend/internal/repository/account_repo.go | 2 +- backend/internal/service/token_refresher.go | 10 ++++++---- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index a8015ddba0..c3c2f6708a 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -729,7 +729,7 @@ func (r *accountRepository) ListOAuthRefreshCandidates(ctx context.Context) ([]s FROM accounts WHERE deleted_at IS NULL AND status = 'active' - AND type = 'oauth' + AND type IN ('oauth', 'setup-token') AND platform IN ('anthropic', 'openai', 'gemini', 'antigravity') AND credentials ? 'refresh_token' AND btrim(credentials->>'refresh_token') <> '' diff --git a/backend/internal/service/token_refresher.go b/backend/internal/service/token_refresher.go index da5edb782c..637efc21cc 100644 --- a/backend/internal/service/token_refresher.go +++ b/backend/internal/service/token_refresher.go @@ -38,11 +38,13 @@ func (r *ClaudeTokenRefresher) CacheKey(account *Account) string { } // CanRefresh 检查是否能处理此账号 -// 只处理 anthropic 平台的 oauth 类型账号 -// setup-token 虽然也是OAuth,但有效期1年,不需要频繁刷新 +// 处理 anthropic 平台的 oauth 与 setup-token 类型账号。 +// 两者的 access_token 均为短期令牌(expires_in=28800,即 8h),到期都需刷新; +// setup-token 之前被排除会导致其 access_token 过期后请求 401。 +// 此处与手动刷新入口(account.IsOAuth())保持一致,实际是否刷新由 NeedsRefresh +// 基于 expires_at 门控,并在分布式锁保护下执行,不会造成过度刷新。 func (r *ClaudeTokenRefresher) CanRefresh(account *Account) bool { - return account.Platform == PlatformAnthropic && - account.Type == AccountTypeOAuth + return account.Platform == PlatformAnthropic && account.IsOAuth() } // NeedsRefresh 检查token是否需要刷新 From 4a2b10c94e91c275a43b61ee68b6fe89855d1953 Mon Sep 17 00:00:00 2001 From: benjamin Date: Fri, 10 Jul 2026 09:08:58 +0800 Subject: [PATCH 03/27] feat(openai): support GPT-5.6 cache write billing --- .../chatcompletions_responses_bridge.go | 12 ++- .../chatcompletions_responses_test.go | 21 ++++ .../apicompat/responses_to_chatcompletions.go | 16 ++- backend/internal/pkg/apicompat/types.go | 28 +++++- backend/internal/service/billing_service.go | 97 +++++++++++-------- .../service/model_pricing_resolver.go | 2 + backend/internal/service/openai_embeddings.go | 13 +-- .../openai_gateway_chat_completions_raw.go | 9 +- .../service/openai_gateway_messages.go | 12 ++- .../openai_gateway_record_usage_test.go | 45 ++++++++- .../openai_gateway_response_handling.go | 33 +++++-- .../service/openai_gateway_service_test.go | 12 ++- .../internal/service/openai_gateway_usage.go | 6 +- .../service/openai_ws_forwarder_support.go | 12 +-- .../service/openai_ws_v2/passthrough_relay.go | 20 +++- .../passthrough_relay_internal_test.go | 2 +- backend/internal/service/pricing_service.go | 5 + .../internal/service/pricing_service_test.go | 57 +++++++++++ 18 files changed, 311 insertions(+), 91 deletions(-) diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go index eeeedd29aa..13a044926b 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go @@ -608,9 +608,17 @@ func ChatUsageToResponsesUsage(usage *ChatUsage) *ResponsesUsage { if out.TotalTokens == 0 { out.TotalTokens = out.InputTokens + out.OutputTokens } - if usage.PromptTokensDetails != nil && usage.PromptTokensDetails.CachedTokens > 0 { + if usage.PromptTokensDetails != nil && (usage.PromptTokensDetails.CachedTokens > 0 || + usage.PromptTokensDetails.CacheCreationTokens > 0 || usage.PromptTokensDetails.CacheWriteTokens > 0) { out.InputTokensDetails = &ResponsesInputTokensDetails{ - CachedTokens: usage.PromptTokensDetails.CachedTokens, + CachedTokens: usage.PromptTokensDetails.CachedTokens, + CacheCreationTokens: usage.PromptTokensDetails.CacheCreationTokens, + CacheWriteTokens: usage.PromptTokensDetails.CacheWriteTokens, + } + if usage.PromptTokensDetails.CacheWriteTokens > 0 { + out.CacheCreationInputTokens = usage.PromptTokensDetails.CacheWriteTokens + } else { + out.CacheCreationInputTokens = usage.PromptTokensDetails.CacheCreationTokens } } return out diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index b30330863c..358aa9ac8a 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -32,6 +32,27 @@ func TestChatCompletionsToResponses_BasicText(t *testing.T) { assert.Equal(t, "user", items[0].Role) } +func TestUsageConversionsPreserveCacheWriteTokens(t *testing.T) { + var responsesUsage ResponsesUsage + require.NoError(t, json.Unmarshal([]byte(`{ + "input_tokens":1000, + "output_tokens":50, + "input_tokens_details":{"cached_tokens":100,"cache_write_tokens":200} + }`), &responsesUsage)) + require.NotNil(t, responsesUsage.InputTokensDetails) + require.Equal(t, 200, responsesUsage.InputTokensDetails.CacheWriteTokens) + + chatUsage := chatUsageFromResponsesUsage(&responsesUsage) + require.NotNil(t, chatUsage.PromptTokensDetails) + require.Equal(t, 100, chatUsage.PromptTokensDetails.CachedTokens) + require.Equal(t, 200, chatUsage.PromptTokensDetails.CacheWriteTokens) + + roundTrip := ChatUsageToResponsesUsage(chatUsage) + require.NotNil(t, roundTrip.InputTokensDetails) + require.Equal(t, 200, roundTrip.CacheCreationInputTokens) + require.Equal(t, 200, roundTrip.InputTokensDetails.CacheWriteTokens) +} + func TestChatCompletionsToResponses_SystemMessage(t *testing.T) { req := &ChatCompletionsRequest{ Model: "gpt-4o", diff --git a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go index 13a89ab04c..2ae6f8ac3f 100644 --- a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go +++ b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go @@ -338,6 +338,14 @@ func chatUsageFromResponsesUsage(u *ResponsesUsage) *ChatUsage { TotalTokens: u.InputTokens + u.OutputTokens, } usage.PromptTokensDetails = promptDetailsFromResponses(u.InputTokensDetails) + if u.CacheCreationInputTokens > 0 { + if usage.PromptTokensDetails == nil { + usage.PromptTokensDetails = &ChatTokenDetails{} + } + if usage.PromptTokensDetails.CacheWriteTokens == 0 && usage.PromptTokensDetails.CacheCreationTokens == 0 { + usage.PromptTokensDetails.CacheCreationTokens = u.CacheCreationInputTokens + } + } usage.CompletionTokensDetails = completionDetailsFromResponses(u.OutputTokensDetails) return usage } @@ -349,12 +357,14 @@ func promptDetailsFromResponses(src *ResponsesInputTokensDetails) *ChatTokenDeta if src == nil { return nil } - if src.CachedTokens == 0 && src.AudioTokens == 0 { + if src.CachedTokens == 0 && src.AudioTokens == 0 && src.CacheCreationTokens == 0 && src.CacheWriteTokens == 0 { return nil } return &ChatTokenDetails{ - CachedTokens: src.CachedTokens, - AudioTokens: src.AudioTokens, + CachedTokens: src.CachedTokens, + AudioTokens: src.AudioTokens, + CacheCreationTokens: src.CacheCreationTokens, + CacheWriteTokens: src.CacheWriteTokens, } } diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index a0fd07a0d1..6980d52b7c 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -320,9 +320,10 @@ type ResponsesSummary struct { // ResponsesUsage holds token counts in Responses API format. type ResponsesUsage struct { - InputTokens int `json:"input_tokens"` - OutputTokens int `json:"output_tokens"` - TotalTokens int `json:"total_tokens"` + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` + CacheCreationInputTokens int `json:"cache_creation_input_tokens,omitempty"` // Optional detailed breakdown InputTokensDetails *ResponsesInputTokensDetails `json:"input_tokens_details,omitempty"` @@ -335,6 +336,9 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { responsesUsageAlias PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` + CacheCreationTokens int `json:"cache_creation_tokens"` + CacheWriteInputTokens int `json:"cache_write_input_tokens"` + CacheWriteTokens int `json:"cache_write_tokens"` PromptTokensDetails *ResponsesInputTokensDetails `json:"prompt_tokens_details,omitempty"` CompletionTokensDetails *ResponsesOutputTokensDetails `json:"completion_tokens_details,omitempty"` } @@ -348,6 +352,16 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { if u.OutputTokens == 0 && aux.CompletionTokens != 0 { u.OutputTokens = aux.CompletionTokens } + if u.CacheCreationInputTokens == 0 { + switch { + case aux.CacheWriteInputTokens > 0: + u.CacheCreationInputTokens = aux.CacheWriteInputTokens + case aux.CacheCreationTokens > 0: + u.CacheCreationInputTokens = aux.CacheCreationTokens + case aux.CacheWriteTokens > 0: + u.CacheCreationInputTokens = aux.CacheWriteTokens + } + } if u.InputTokensDetails == nil && aux.PromptTokensDetails != nil { u.InputTokensDetails = aux.PromptTokensDetails } @@ -362,8 +376,10 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { // ResponsesInputTokensDetails breaks down input token usage. type ResponsesInputTokensDetails struct { - CachedTokens int `json:"cached_tokens,omitempty"` - AudioTokens int `json:"audio_tokens,omitempty"` + CachedTokens int `json:"cached_tokens,omitempty"` + AudioTokens int `json:"audio_tokens,omitempty"` + CacheCreationTokens int `json:"cache_creation_tokens,omitempty"` + CacheWriteTokens int `json:"cache_write_tokens,omitempty"` } // ResponsesOutputTokensDetails breaks down output token usage. @@ -545,6 +561,8 @@ type ChatUsage struct { type ChatTokenDetails struct { CachedTokens int `json:"cached_tokens,omitempty"` AudioTokens int `json:"audio_tokens,omitempty"` + CacheCreationTokens int `json:"cache_creation_tokens,omitempty"` + CacheWriteTokens int `json:"cache_write_tokens,omitempty"` ReasoningTokens int `json:"reasoning_tokens,omitempty"` AcceptedPredictionTokens int `json:"accepted_prediction_tokens,omitempty"` RejectedPredictionTokens int `json:"rejected_prediction_tokens,omitempty"` diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 8dceebc250..e82c38a14f 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -90,22 +90,23 @@ type BillingCache interface { // ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致) type ModelPricing struct { - InputPricePerToken float64 // 每token输入价格 (USD) - InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD) - ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken - OutputPricePerToken float64 // 每token输出价格 (USD) - OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD) - CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD) - CacheReadPricePerToken float64 // 缓存读取每token价格 (USD) - CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD) - CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD) - CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD) - SupportsCacheBreakdown bool // 是否支持详细的缓存分类 - LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格 - LongContextInputMultiplier float64 // 长上下文整次会话输入倍率 - LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率 - ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD) - ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退) + InputPricePerToken float64 // 每token输入价格 (USD) + InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD) + ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken + OutputPricePerToken float64 // 每token输出价格 (USD) + OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD) + CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD) + CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD) + CacheReadPricePerToken float64 // 缓存读取每token价格 (USD) + CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD) + CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD) + CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD) + SupportsCacheBreakdown bool // 是否支持详细的缓存分类 + LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格 + LongContextInputMultiplier float64 // 长上下文整次会话输入倍率 + LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率 + ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD) + ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退) } const ( @@ -122,7 +123,8 @@ func usePriorityServiceTierPricing(serviceTier string, pricing *ModelPricing) bo if pricing == nil || normalizeBillingServiceTier(serviceTier) != "priority" { return false } - return pricing.InputPricePerTokenPriority > 0 || pricing.OutputPricePerTokenPriority > 0 || pricing.CacheReadPricePerTokenPriority > 0 + return pricing.InputPricePerTokenPriority > 0 || pricing.OutputPricePerTokenPriority > 0 || + pricing.CacheCreationPricePerTokenPriority > 0 || pricing.CacheReadPricePerTokenPriority > 0 } func serviceTierCostMultiplier(serviceTier string) float64 { @@ -739,20 +741,21 @@ func (s *BillingService) GetModelPricing(model string) (*ModelPricing, error) { price1h := litellmPricing.CacheCreationInputTokenCostAbove1hr enableBreakdown := price1h > 0 && price1h > price5m return s.applyModelSpecificPricingPolicy(model, &ModelPricing{ - InputPricePerToken: litellmPricing.InputCostPerToken, - InputPricePerTokenPriority: litellmPricing.InputCostPerTokenPriority, - OutputPricePerToken: litellmPricing.OutputCostPerToken, - OutputPricePerTokenPriority: litellmPricing.OutputCostPerTokenPriority, - CacheCreationPricePerToken: litellmPricing.CacheCreationInputTokenCost, - CacheReadPricePerToken: litellmPricing.CacheReadInputTokenCost, - CacheReadPricePerTokenPriority: litellmPricing.CacheReadInputTokenCostPriority, - CacheCreation5mPrice: price5m, - CacheCreation1hPrice: price1h, - SupportsCacheBreakdown: enableBreakdown, - LongContextInputThreshold: litellmPricing.LongContextInputTokenThreshold, - LongContextInputMultiplier: litellmPricing.LongContextInputCostMultiplier, - LongContextOutputMultiplier: litellmPricing.LongContextOutputCostMultiplier, - ImageOutputPricePerToken: litellmPricing.OutputCostPerImageToken, + InputPricePerToken: litellmPricing.InputCostPerToken, + InputPricePerTokenPriority: litellmPricing.InputCostPerTokenPriority, + OutputPricePerToken: litellmPricing.OutputCostPerToken, + OutputPricePerTokenPriority: litellmPricing.OutputCostPerTokenPriority, + CacheCreationPricePerToken: litellmPricing.CacheCreationInputTokenCost, + CacheCreationPricePerTokenPriority: litellmPricing.CacheCreationInputTokenCostPriority, + CacheReadPricePerToken: litellmPricing.CacheReadInputTokenCost, + CacheReadPricePerTokenPriority: litellmPricing.CacheReadInputTokenCostPriority, + CacheCreation5mPrice: price5m, + CacheCreation1hPrice: price1h, + SupportsCacheBreakdown: enableBreakdown, + LongContextInputThreshold: litellmPricing.LongContextInputTokenThreshold, + LongContextInputMultiplier: litellmPricing.LongContextInputCostMultiplier, + LongContextOutputMultiplier: litellmPricing.LongContextOutputCostMultiplier, + ImageOutputPricePerToken: litellmPricing.OutputCostPerImageToken, }), nil } } @@ -794,6 +797,7 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing } if channelPricing.CacheWritePrice != nil { pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice + pricing.CacheCreationPricePerTokenPriority = *channelPricing.CacheWritePrice pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice } @@ -867,7 +871,7 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, // calculateTokenCost 按 token 区间计费 func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) { - totalContext := input.Tokens.InputTokens + input.Tokens.CacheReadTokens + totalContext := input.Tokens.InputTokens + input.Tokens.CacheCreationTokens + input.Tokens.CacheReadTokens pricing := input.Resolver.GetIntervalPricing(resolved, totalContext) if pricing == nil { @@ -897,6 +901,7 @@ func (s *BillingService) computeTokenBreakdown( inputPrice := pricing.InputPricePerToken outputPrice := pricing.OutputPricePerToken cacheReadPrice := pricing.CacheReadPricePerToken + cacheCreationPrice := pricing.CacheCreationPricePerToken cacheCreationMultiplier := 1.0 tierMultiplier := 1.0 @@ -910,6 +915,9 @@ func (s *BillingService) computeTokenBreakdown( if pricing.CacheReadPricePerTokenPriority > 0 { cacheReadPrice = pricing.CacheReadPricePerTokenPriority } + if pricing.CacheCreationPricePerTokenPriority > 0 { + cacheCreationPrice = pricing.CacheCreationPricePerTokenPriority + } } else { tierMultiplier = serviceTierCostMultiplier(serviceTier) } @@ -963,7 +971,7 @@ func (s *BillingService) computeTokenBreakdown( } // 缓存创建费用 - bd.CacheCreationCost = s.computeCacheCreationCost(pricing, tokens, cacheCreationMultiplier) + bd.CacheCreationCost = s.computeCacheCreationCost(pricing, tokens, cacheCreationPrice, cacheCreationMultiplier) bd.CacheReadCost = float64(tokens.CacheReadTokens) * cacheReadPrice @@ -984,7 +992,7 @@ func (s *BillingService) computeTokenBreakdown( // computeCacheCreationCost 计算缓存创建费用(支持 5m/1h 分类或标准计费)。 // multiplier 用于长上下文等场景下的整体价格缩放(普通调用传 1.0 即可)。 -func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens, multiplier float64) float64 { +func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens, price, multiplier float64) float64 { if pricing.SupportsCacheBreakdown && (pricing.CacheCreation5mPrice > 0 || pricing.CacheCreation1hPrice > 0) { if tokens.CacheCreation5mTokens == 0 && tokens.CacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 { // API 未返回 ephemeral 明细,回退到全部按 5m 单价计费 @@ -993,7 +1001,7 @@ func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens return float64(tokens.CacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier + float64(tokens.CacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier } - return float64(tokens.CacheCreationTokens) * pricing.CacheCreationPricePerToken * multiplier + return float64(tokens.CacheCreationTokens) * price * multiplier } // calculatePerRequestCost 按次/图片计费 @@ -1010,7 +1018,7 @@ func (s *BillingService) calculatePerRequestCost(resolved *ResolvedPricing, inpu } if unitPrice == 0 { - totalContext := input.Tokens.InputTokens + input.Tokens.CacheReadTokens + totalContext := input.Tokens.InputTokens + input.Tokens.CacheCreationTokens + input.Tokens.CacheReadTokens unitPrice = input.Resolver.GetRequestTierPriceByContext(resolved, totalContext) } @@ -1057,13 +1065,26 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * if pricing == nil { return nil } + normalized := normalizeKnownOpenAICodexModel(model) if !isOpenAIGPT54Model(model) { return pricing } - if pricing.LongContextInputThreshold > 0 && pricing.LongContextInputMultiplier > 0 && pricing.LongContextOutputMultiplier > 0 { + isGPT56 := normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" + needsLongContextPolicy := pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0 + needsCacheCreationPolicy := isGPT56 && (pricing.CacheCreationPricePerToken <= 0 || + (pricing.InputPricePerTokenPriority > 0 && pricing.CacheCreationPricePerTokenPriority <= 0)) + if !needsLongContextPolicy && !needsCacheCreationPolicy { return pricing } cloned := *pricing + if isGPT56 { + if cloned.CacheCreationPricePerToken <= 0 { + cloned.CacheCreationPricePerToken = cloned.InputPricePerToken + } + if cloned.CacheCreationPricePerTokenPriority <= 0 { + cloned.CacheCreationPricePerTokenPriority = cloned.InputPricePerTokenPriority + } + } if cloned.LongContextInputThreshold <= 0 { cloned.LongContextInputThreshold = openAIGPT54LongContextInputThreshold } @@ -1083,7 +1104,7 @@ func (s *BillingService) shouldApplySessionLongContextPricing(tokens UsageTokens if pricing.LongContextInputMultiplier <= 1 && pricing.LongContextOutputMultiplier <= 1 { return false } - totalInputTokens := tokens.InputTokens + tokens.CacheReadTokens + totalInputTokens := tokens.InputTokens + tokens.CacheCreationTokens + tokens.CacheReadTokens return totalInputTokens > pricing.LongContextInputThreshold } diff --git a/backend/internal/service/model_pricing_resolver.go b/backend/internal/service/model_pricing_resolver.go index 0cc7a0ac47..7ddf6ea475 100644 --- a/backend/internal/service/model_pricing_resolver.go +++ b/backend/internal/service/model_pricing_resolver.go @@ -183,6 +183,7 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin } if chPricing.CacheWritePrice != nil { resolved.BasePricing.CacheCreationPricePerToken = *chPricing.CacheWritePrice + resolved.BasePricing.CacheCreationPricePerTokenPriority = *chPricing.CacheWritePrice resolved.BasePricing.CacheCreation5mPrice = *chPricing.CacheWritePrice resolved.BasePricing.CacheCreation1hPrice = *chPricing.CacheWritePrice } @@ -251,6 +252,7 @@ func intervalToModelPricing(iv *PricingInterval, supportsCacheBreakdown bool, ch } if iv.CacheWritePrice != nil { pricing.CacheCreationPricePerToken = *iv.CacheWritePrice + pricing.CacheCreationPricePerTokenPriority = *iv.CacheWritePrice pricing.CacheCreation5mPrice = *iv.CacheWritePrice pricing.CacheCreation1hPrice = *iv.CacheWritePrice } diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go index fb2dc5ccbb..fa821d8f31 100644 --- a/backend/internal/service/openai_embeddings.go +++ b/backend/internal/service/openai_embeddings.go @@ -206,17 +206,8 @@ func extractOpenAIEmbeddingsUsage(body []byte) OpenAIUsage { usage.Get("completion_tokens"), usage.Get("output_tokens"), ) - cacheReadTokens := firstPositiveGJSONInt( - usage.Get("prompt_tokens_details.cached_tokens"), - usage.Get("input_tokens_details.cached_tokens"), - usage.Get("cache_read_tokens"), - usage.Get("cache_read_input_tokens"), - ) - cacheCreationTokens := firstPositiveGJSONInt( - usage.Get("cache_creation_tokens"), - usage.Get("cache_creation_input_tokens"), - usage.Get("input_tokens_details.cache_creation_tokens"), - ) + cacheReadTokens := openAICacheReadTokensFromUsage(usage) + cacheCreationTokens := openAICacheCreationTokensFromUsage(usage) // 多模态 embedding(如 doubao-embedding-vision)回传图文 token 拆分, // 用于图文不同价计费;纯文本 embedding 该字段为 0,行为不变。 imageInputTokens := firstPositiveGJSONInt( diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 9b31b803d2..d73bdd1284 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -383,12 +383,9 @@ func extractCCStreamUsage(payload string) *OpenAIUsage { if !usageResult.Exists() || !usageResult.IsObject() { return nil } - u := OpenAIUsage{ - InputTokens: int(gjson.Get(payload, "usage.prompt_tokens").Int()), - OutputTokens: int(gjson.Get(payload, "usage.completion_tokens").Int()), - } - if cached := gjson.Get(payload, "usage.prompt_tokens_details.cached_tokens"); cached.Exists() { - u.CacheReadInputTokens = int(cached.Int()) + u, ok := openAIUsageFromGJSON(usageResult) + if !ok { + return nil } return &u } diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index b621a96aee..52ec34f27e 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -1114,11 +1114,19 @@ func copyOpenAIUsageFromResponsesUsage(usage *apicompat.ResponsesUsage) OpenAIUs return OpenAIUsage{} } result := OpenAIUsage{ - InputTokens: usage.InputTokens, - OutputTokens: usage.OutputTokens, + InputTokens: usage.InputTokens, + OutputTokens: usage.OutputTokens, + CacheCreationInputTokens: usage.CacheCreationInputTokens, } if usage.InputTokensDetails != nil { result.CacheReadInputTokens = usage.InputTokensDetails.CachedTokens + if result.CacheCreationInputTokens == 0 { + if usage.InputTokensDetails.CacheWriteTokens > 0 { + result.CacheCreationInputTokens = usage.InputTokensDetails.CacheWriteTokens + } else { + result.CacheCreationInputTokens = usage.InputTokensDetails.CacheCreationTokens + } + } } return result } diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index b9bbd137c6..f749c9c1f7 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -247,7 +247,7 @@ func expectedOpenAICost(t *testing.T, svc *OpenAIGatewayService, model string, u t.Helper() cost, err := svc.billingService.CalculateCost(model, UsageTokens{ - InputTokens: max(usage.InputTokens-usage.CacheReadInputTokens, 0), + InputTokens: max(usage.InputTokens-usage.CacheReadInputTokens-usage.CacheCreationInputTokens, 0), OutputTokens: usage.OutputTokens, CacheCreationTokens: usage.CacheCreationInputTokens, CacheReadTokens: usage.CacheReadInputTokens, @@ -1002,6 +1002,49 @@ func TestOpenAIGatewayServiceRecordUsage_ClampsActualInputTokensToZero(t *testin require.Equal(t, 0, usageRepo.lastLog.InputTokens) } +func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndStats(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + svc.billingService = NewBillingService(svc.cfg, &PricingService{pricingData: map[string]*LiteLLMModelPricing{ + "gpt-5.6-sol": { + InputCostPerToken: 5e-6, + OutputCostPerToken: 30e-6, + CacheReadInputTokenCost: 0.5e-6, + }, + }}) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_gpt56_cache_write", + Usage: OpenAIUsage{ + InputTokens: 1000, + OutputTokens: 50, + CacheCreationInputTokens: 200, + CacheReadInputTokens: 100, + }, + Model: "gpt-5.6-sol", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 1056}, + User: &User{ID: 2056}, + Account: &Account{ID: 3056}, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, 700, usageRepo.lastLog.InputTokens) + require.Equal(t, 200, usageRepo.lastLog.CacheCreationTokens) + require.Equal(t, 100, usageRepo.lastLog.CacheReadTokens) + require.Equal(t, 1050, usageRepo.lastLog.TotalTokens()) + require.InDelta(t, 700*5e-6, usageRepo.lastLog.InputCost, 1e-12) + require.InDelta(t, 200*5e-6, usageRepo.lastLog.CacheCreationCost, 1e-12) + require.InDelta(t, 100*0.5e-6, usageRepo.lastLog.CacheReadCost, 1e-12) + require.InDelta(t, 50*30e-6, usageRepo.lastLog.OutputCost, 1e-12) + require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12) +} + func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index cb4b780cdf..77d212dbe9 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -756,10 +756,8 @@ func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) { if outputTokens == 0 { outputTokens = value.Get("completion_tokens").Int() } - cacheReadTokens := value.Get("input_tokens_details.cached_tokens").Int() - if cacheReadTokens == 0 { - cacheReadTokens = value.Get("prompt_tokens_details.cached_tokens").Int() - } + cacheReadTokens := openAICacheReadTokensFromUsage(value) + cacheCreationTokens := openAICacheCreationTokensFromUsage(value) imageOutputTokens := value.Get("output_tokens_details.image_tokens").Int() if imageOutputTokens == 0 { imageOutputTokens = value.Get("completion_tokens_details.image_tokens").Int() @@ -767,12 +765,35 @@ func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) { return OpenAIUsage{ InputTokens: int(inputTokens), OutputTokens: int(outputTokens), - CacheCreationInputTokens: int(value.Get("cache_creation_input_tokens").Int()), - CacheReadInputTokens: int(cacheReadTokens), + CacheCreationInputTokens: cacheCreationTokens, + CacheReadInputTokens: cacheReadTokens, ImageOutputTokens: int(imageOutputTokens), }, true } +func openAICacheReadTokensFromUsage(value gjson.Result) int { + return firstPositiveGJSONInt( + value.Get("input_tokens_details.cached_tokens"), + value.Get("prompt_tokens_details.cached_tokens"), + value.Get("cache_read_input_tokens"), + value.Get("cache_read_tokens"), + value.Get("cached_tokens"), + ) +} + +func openAICacheCreationTokensFromUsage(value gjson.Result) int { + return firstPositiveGJSONInt( + value.Get("cache_creation_input_tokens"), + value.Get("cache_write_input_tokens"), + value.Get("cache_creation_tokens"), + value.Get("cache_write_tokens"), + value.Get("input_tokens_details.cache_creation_tokens"), + value.Get("input_tokens_details.cache_write_tokens"), + value.Get("prompt_tokens_details.cache_creation_tokens"), + value.Get("prompt_tokens_details.cache_write_tokens"), + ) +} + func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, originalModel, mappedModel string) (*openaiNonStreamingResult, error) { body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) if err != nil { diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 9f42a82312..1628ca3d2b 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2800,17 +2800,23 @@ func TestParseSSEUsage_SelectiveParsing(t *testing.T) { } func TestExtractOpenAIUsageFromJSONBytes_AcceptsResponseAndChatUsageShapes(t *testing.T) { - usage, ok := extractOpenAIUsageFromJSONBytes([]byte(`{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":5,"input_tokens_details":{"cached_tokens":2}}}`)) + usage, ok := extractOpenAIUsageFromJSONBytes([]byte(`{"id":"resp_1","usage":{"input_tokens":9,"output_tokens":5,"input_tokens_details":{"cached_tokens":2,"cache_write_tokens":4}}}`)) require.True(t, ok) - require.Equal(t, 3, usage.InputTokens) + require.Equal(t, 9, usage.InputTokens) require.Equal(t, 5, usage.OutputTokens) require.Equal(t, 2, usage.CacheReadInputTokens) + require.Equal(t, 4, usage.CacheCreationInputTokens) - usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"type":"response.completed","response":{"usage":{"prompt_tokens":13,"completion_tokens":7,"prompt_tokens_details":{"cached_tokens":4}}}}`)) + usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"type":"response.completed","response":{"usage":{"prompt_tokens":13,"completion_tokens":7,"prompt_tokens_details":{"cached_tokens":4,"cache_creation_tokens":3}}}}`)) require.True(t, ok) require.Equal(t, 13, usage.InputTokens) require.Equal(t, 7, usage.OutputTokens) require.Equal(t, 4, usage.CacheReadInputTokens) + require.Equal(t, 3, usage.CacheCreationInputTokens) + + usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":11,"output_tokens":2,"cache_write_input_tokens":6}}`)) + require.True(t, ok) + require.Equal(t, 6, usage.CacheCreationInputTokens) } func TestExtractCodexFinalResponse_SampleReplay(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index 4a16facccd..f96b679cf6 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -119,9 +119,9 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec ApplyOpenAIImageBillingResolution(result) } - // 计算实际的新输入token(减去缓存读取的token) - // 因为 input_tokens 包含了 cache_read_tokens,而缓存读取的token不应按输入价格计费 - actualInputTokens := result.Usage.InputTokens - result.Usage.CacheReadInputTokens + // OpenAI input_tokens 是总输入,包含缓存读取和缓存写入明细。 + // 将三类 token 拆成互斥桶,避免缓存写入同时按普通输入和 cache_write 重复计费。 + actualInputTokens := result.Usage.InputTokens - result.Usage.CacheReadInputTokens - result.Usage.CacheCreationInputTokens if actualInputTokens < 0 { actualInputTokens = 0 } diff --git a/backend/internal/service/openai_ws_forwarder_support.go b/backend/internal/service/openai_ws_forwarder_support.go index a2fc28d9e7..31c63e214b 100644 --- a/backend/internal/service/openai_ws_forwarder_support.go +++ b/backend/internal/service/openai_ws_forwarder_support.go @@ -262,15 +262,9 @@ func populateOpenAIUsageFromResponseJSON(body []byte, usage *OpenAIUsage) { if usage == nil || len(body) == 0 { return } - values := gjson.GetManyBytes( - body, - "usage.input_tokens", - "usage.output_tokens", - "usage.input_tokens_details.cached_tokens", - ) - usage.InputTokens = int(values[0].Int()) - usage.OutputTokens = int(values[1].Int()) - usage.CacheReadInputTokens = int(values[2].Int()) + if parsed, ok := extractOpenAIUsageFromJSONBytes(body); ok { + *usage = parsed + } } func getOpenAIGroupIDFromContext(c *gin.Context) int64 { diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index 6aba3b7dbb..874c618a29 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -787,7 +787,7 @@ func parseUsageAndAccumulate( parsedUsage := Usage{ InputTokens: inputTokens, OutputTokens: outputTokens, - CacheCreationInputTokens: int(usageResult.Get("cache_creation_input_tokens").Int()), + CacheCreationInputTokens: openAICacheCreationTokensFromUsage(usageResult), CacheReadInputTokens: cachedTokens, ImageOutputTokens: int(imageTokens), } @@ -810,6 +810,24 @@ func parseUsageIntField(value gjson.Result, required bool) (int, bool) { return int(value.Int()), true } +func openAICacheCreationTokensFromUsage(value gjson.Result) int { + for _, field := range []string{ + "cache_creation_input_tokens", + "cache_write_input_tokens", + "cache_creation_tokens", + "cache_write_tokens", + "input_tokens_details.cache_creation_tokens", + "input_tokens_details.cache_write_tokens", + "prompt_tokens_details.cache_creation_tokens", + "prompt_tokens_details.cache_write_tokens", + } { + if tokens := int(value.Get(field).Int()); tokens > 0 { + return tokens + } + } + return 0 +} + func enrichResult(result *RelayResult, state *relayState, duration time.Duration) { if result == nil { return diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go index 13c51f663f..cb2bd9cc31 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go @@ -300,7 +300,7 @@ func TestParseUsageAndEnrichCoverage(t *testing.T) { require.Equal(t, 0, state.usage.OutputTokens) require.Equal(t, 0, state.usage.CacheReadInputTokens) - parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1},"cache_creation_input_tokens":4,"output_tokens_details":{"image_tokens":3}}}}`), "response.completed", nil) + parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1,"cache_write_tokens":4},"output_tokens_details":{"image_tokens":3}}}}`), "response.completed", nil) require.Equal(t, 2, state.usage.InputTokens) require.Equal(t, 1, state.usage.OutputTokens) require.Equal(t, 1, state.usage.CacheReadInputTokens) diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index 2ae15df507..a21cb98783 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -61,6 +61,7 @@ type LiteLLMModelPricing struct { OutputCostPerToken float64 `json:"output_cost_per_token"` OutputCostPerTokenPriority float64 `json:"output_cost_per_token_priority"` CacheCreationInputTokenCost float64 `json:"cache_creation_input_token_cost"` + CacheCreationInputTokenCostPriority float64 `json:"cache_creation_input_token_cost_priority"` CacheCreationInputTokenCostAbove1hr float64 `json:"cache_creation_input_token_cost_above_1hr"` CacheReadInputTokenCost float64 `json:"cache_read_input_token_cost"` CacheReadInputTokenCostPriority float64 `json:"cache_read_input_token_cost_priority"` @@ -93,6 +94,7 @@ type LiteLLMRawEntry struct { OutputCostPerToken *float64 `json:"output_cost_per_token"` OutputCostPerTokenPriority *float64 `json:"output_cost_per_token_priority"` CacheCreationInputTokenCost *float64 `json:"cache_creation_input_token_cost"` + CacheCreationInputTokenCostPriority *float64 `json:"cache_creation_input_token_cost_priority"` CacheCreationInputTokenCostAbove1hr *float64 `json:"cache_creation_input_token_cost_above_1hr"` CacheReadInputTokenCost *float64 `json:"cache_read_input_token_cost"` CacheReadInputTokenCostPriority *float64 `json:"cache_read_input_token_cost_priority"` @@ -406,6 +408,9 @@ func (s *PricingService) parsePricingData(body []byte) (map[string]*LiteLLMModel if entry.CacheCreationInputTokenCost != nil { pricing.CacheCreationInputTokenCost = *entry.CacheCreationInputTokenCost } + if entry.CacheCreationInputTokenCostPriority != nil { + pricing.CacheCreationInputTokenCostPriority = *entry.CacheCreationInputTokenCostPriority + } if entry.CacheCreationInputTokenCostAbove1hr != nil { pricing.CacheCreationInputTokenCostAbove1hr = *entry.CacheCreationInputTokenCostAbove1hr } diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index 4bf8f2379e..e8f00ae3b2 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -19,6 +19,7 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { "output_cost_per_token": 0.000015, "output_cost_per_token_priority": 0.00003, "cache_creation_input_token_cost": 0.0000025, + "cache_creation_input_token_cost_priority": 0.000005, "cache_read_input_token_cost": 0.00000025, "cache_read_input_token_cost_priority": 0.0000005, "supports_service_tier": true, @@ -34,10 +35,66 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { require.NotNil(t, pricing) require.InDelta(t, 5e-6, pricing.InputCostPerTokenPriority, 1e-12) require.InDelta(t, 3e-5, pricing.OutputCostPerTokenPriority, 1e-12) + require.InDelta(t, 5e-6, pricing.CacheCreationInputTokenCostPriority, 1e-12) require.InDelta(t, 5e-7, pricing.CacheReadInputTokenCostPriority, 1e-12) require.True(t, pricing.SupportsServiceTier) } +func TestBillingService_GPT56CacheWritePricingUsesInputTier(t *testing.T) { + for _, model := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { + t.Run(model, func(t *testing.T) { + pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ + model: { + InputCostPerToken: 5e-6, + InputCostPerTokenPriority: 10e-6, + OutputCostPerToken: 30e-6, + OutputCostPerTokenPriority: 60e-6, + CacheReadInputTokenCost: 0.5e-6, + CacheReadInputTokenCostPriority: 1e-6, + }, + }} + svc := NewBillingService(&config.Config{}, pricingSvc) + + pricing, err := svc.GetModelPricing(model) + require.NoError(t, err) + require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12) + require.InDelta(t, 10e-6, pricing.CacheCreationPricePerTokenPriority, 1e-12) + + tokens := UsageTokens{InputTokens: 700, OutputTokens: 50, CacheCreationTokens: 200, CacheReadTokens: 100} + standard, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "") + require.NoError(t, err) + require.InDelta(t, 200*5e-6, standard.CacheCreationCost, 1e-12) + + priority, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "priority") + require.NoError(t, err) + require.InDelta(t, 200*10e-6, priority.CacheCreationCost, 1e-12) + + flex, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "flex") + require.NoError(t, err) + require.InDelta(t, 200*2.5e-6, flex.CacheCreationCost, 1e-12) + }) + } +} + +func TestBillingService_GPT56CacheWriteContributesToLongContextThreshold(t *testing.T) { + model := "gpt-5.6-sol" + pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ + model: { + InputCostPerToken: 5e-6, + OutputCostPerToken: 30e-6, + CacheReadInputTokenCost: 0.5e-6, + }, + }} + svc := NewBillingService(&config.Config{}, pricingSvc) + tokens := UsageTokens{InputTokens: 100000, CacheCreationTokens: 173000, OutputTokens: 10} + + cost, err := svc.CalculateCost(model, tokens, 1) + require.NoError(t, err) + require.InDelta(t, 100000*10e-6, cost.InputCost, 1e-12) + require.InDelta(t, 173000*10e-6, cost.CacheCreationCost, 1e-12) + require.InDelta(t, 10*45e-6, cost.OutputCost, 1e-12) +} + func TestParsePricingData_KeepsImageOnlyPricing(t *testing.T) { svc := &PricingService{} body := []byte(`{ From 383f61d0e9aa737e33bf2201b47228338fe55fa7 Mon Sep 17 00:00:00 2001 From: benjamin Date: Fri, 10 Jul 2026 09:42:31 +0800 Subject: [PATCH 04/27] fix(openai): align GPT-5.6 billing with official pricing --- backend/internal/service/billing_service.go | 78 ++++++---- .../openai_gateway_record_usage_test.go | 2 +- .../openai_gateway_response_handling.go | 10 +- .../service/openai_gateway_service_test.go | 4 + .../service/openai_ws_v2/passthrough_relay.go | 10 +- backend/internal/service/pricing_service.go | 59 +++++++- .../internal/service/pricing_service_test.go | 135 +++++++++++++++--- .../model_prices_and_context_window.json | 65 +++++---- 8 files changed, 270 insertions(+), 93 deletions(-) diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index e82c38a14f..ae68c93ef6 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -282,10 +282,37 @@ func (s *BillingService) initFallbackPricing() { s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"] s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"] - // GPT-5.6(sol / terra / luna)暂无独立定价,回退到 GPT-5.4。 - s.fallbackPrices["gpt-5.6-sol"] = s.fallbackPrices["gpt-5.4"] - s.fallbackPrices["gpt-5.6-terra"] = s.fallbackPrices["gpt-5.4"] - s.fallbackPrices["gpt-5.6-luna"] = s.fallbackPrices["gpt-5.4"] + // OpenAI GPT-5.6 官方价格(USD/token)。缓存写入为输入价的 1.25 倍。 + s.fallbackPrices["gpt-5.6-sol"] = &ModelPricing{ + InputPricePerToken: 5e-6, + InputPricePerTokenPriority: 10e-6, + OutputPricePerToken: 30e-6, + OutputPricePerTokenPriority: 60e-6, + CacheCreationPricePerToken: 6.25e-6, + CacheCreationPricePerTokenPriority: 12.5e-6, + CacheReadPricePerToken: 0.5e-6, + CacheReadPricePerTokenPriority: 1e-6, + } + s.fallbackPrices["gpt-5.6-terra"] = &ModelPricing{ + InputPricePerToken: 2.5e-6, + InputPricePerTokenPriority: 5e-6, + OutputPricePerToken: 15e-6, + OutputPricePerTokenPriority: 30e-6, + CacheCreationPricePerToken: 3.125e-6, + CacheCreationPricePerTokenPriority: 6.25e-6, + CacheReadPricePerToken: 0.25e-6, + CacheReadPricePerTokenPriority: 0.5e-6, + } + s.fallbackPrices["gpt-5.6-luna"] = &ModelPricing{ + InputPricePerToken: 1e-6, + InputPricePerTokenPriority: 2e-6, + OutputPricePerToken: 6e-6, + OutputPricePerTokenPriority: 12e-6, + CacheCreationPricePerToken: 1.25e-6, + CacheCreationPricePerTokenPriority: 2.5e-6, + CacheReadPricePerToken: 0.1e-6, + CacheReadPricePerTokenPriority: 0.2e-6, + } s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{ InputPricePerToken: 7.5e-7, @@ -1066,11 +1093,13 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * return nil } normalized := normalizeKnownOpenAICodexModel(model) - if !isOpenAIGPT54Model(model) { + isGPT56 := isOpenAIGPT56Model(normalized) + usesLegacyLongContextPricing := usesOpenAILegacyLongContextPricing(normalized) + if !isGPT56 && !usesLegacyLongContextPricing { return pricing } - isGPT56 := normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" - needsLongContextPolicy := pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0 + needsLongContextPolicy := usesLegacyLongContextPricing && + (pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0) needsCacheCreationPolicy := isGPT56 && (pricing.CacheCreationPricePerToken <= 0 || (pricing.InputPricePerTokenPriority > 0 && pricing.CacheCreationPricePerTokenPriority <= 0)) if !needsLongContextPolicy && !needsCacheCreationPolicy { @@ -1079,20 +1108,22 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * cloned := *pricing if isGPT56 { if cloned.CacheCreationPricePerToken <= 0 { - cloned.CacheCreationPricePerToken = cloned.InputPricePerToken + cloned.CacheCreationPricePerToken = cloned.InputPricePerToken * 1.25 } if cloned.CacheCreationPricePerTokenPriority <= 0 { - cloned.CacheCreationPricePerTokenPriority = cloned.InputPricePerTokenPriority + cloned.CacheCreationPricePerTokenPriority = cloned.InputPricePerTokenPriority * 1.25 } } - if cloned.LongContextInputThreshold <= 0 { - cloned.LongContextInputThreshold = openAIGPT54LongContextInputThreshold - } - if cloned.LongContextInputMultiplier <= 0 { - cloned.LongContextInputMultiplier = openAIGPT54LongContextInputMultiplier - } - if cloned.LongContextOutputMultiplier <= 0 { - cloned.LongContextOutputMultiplier = openAIGPT54LongContextOutputMultiplier + if usesLegacyLongContextPricing { + if cloned.LongContextInputThreshold <= 0 { + cloned.LongContextInputThreshold = openAIGPT54LongContextInputThreshold + } + if cloned.LongContextInputMultiplier <= 0 { + cloned.LongContextInputMultiplier = openAIGPT54LongContextInputMultiplier + } + if cloned.LongContextOutputMultiplier <= 0 { + cloned.LongContextOutputMultiplier = openAIGPT54LongContextOutputMultiplier + } } return &cloned } @@ -1108,13 +1139,12 @@ func (s *BillingService) shouldApplySessionLongContextPricing(tokens UsageTokens return totalInputTokens > pricing.LongContextInputThreshold } -func isOpenAIGPT54Model(model string) bool { - // 仅当模型字符串实际属于已知 GPT-5/Codex 族时才做归一判定,避免 - // normalizeCodexModel 的默认兜底把非 OpenAI 模型(claude-*、gemini-*、gpt-4o) - // 误识别为 gpt-5.4。 - normalized := normalizeKnownOpenAICodexModel(model) - return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" || - normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" +func isOpenAIGPT56Model(normalized string) bool { + return normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" +} + +func usesOpenAILegacyLongContextPricing(normalized string) bool { + return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" } // CalculateCostWithConfig 使用配置中的默认倍率计算费用 diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index f749c9c1f7..81578f630f 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -1039,7 +1039,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt require.Equal(t, 100, usageRepo.lastLog.CacheReadTokens) require.Equal(t, 1050, usageRepo.lastLog.TotalTokens()) require.InDelta(t, 700*5e-6, usageRepo.lastLog.InputCost, 1e-12) - require.InDelta(t, 200*5e-6, usageRepo.lastLog.CacheCreationCost, 1e-12) + require.InDelta(t, 200*6.25e-6, usageRepo.lastLog.CacheCreationCost, 1e-12) require.InDelta(t, 100*0.5e-6, usageRepo.lastLog.CacheReadCost, 1e-12) require.InDelta(t, 50*30e-6, usageRepo.lastLog.OutputCost, 1e-12) require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12) diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 77d212dbe9..d527bb2d32 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -783,14 +783,14 @@ func openAICacheReadTokensFromUsage(value gjson.Result) int { func openAICacheCreationTokensFromUsage(value gjson.Result) int { return firstPositiveGJSONInt( + value.Get("input_tokens_details.cache_write_tokens"), + value.Get("prompt_tokens_details.cache_write_tokens"), + value.Get("input_tokens_details.cache_creation_tokens"), + value.Get("prompt_tokens_details.cache_creation_tokens"), + value.Get("cache_write_tokens"), value.Get("cache_creation_input_tokens"), value.Get("cache_write_input_tokens"), value.Get("cache_creation_tokens"), - value.Get("cache_write_tokens"), - value.Get("input_tokens_details.cache_creation_tokens"), - value.Get("input_tokens_details.cache_write_tokens"), - value.Get("prompt_tokens_details.cache_creation_tokens"), - value.Get("prompt_tokens_details.cache_write_tokens"), ) } diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 1628ca3d2b..3fdef698db 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2817,6 +2817,10 @@ func TestExtractOpenAIUsageFromJSONBytes_AcceptsResponseAndChatUsageShapes(t *te usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":11,"output_tokens":2,"cache_write_input_tokens":6}}`)) require.True(t, ok) require.Equal(t, 6, usage.CacheCreationInputTokens) + + usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":20,"output_tokens":2,"cache_creation_input_tokens":19,"input_tokens_details":{"cache_write_tokens":7}}}`)) + require.True(t, ok) + require.Equal(t, 7, usage.CacheCreationInputTokens, "官方嵌套字段应优先于兼容顶层别名") } func TestExtractCodexFinalResponse_SampleReplay(t *testing.T) { diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index 874c618a29..be19438eaa 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -812,14 +812,14 @@ func parseUsageIntField(value gjson.Result, required bool) (int, bool) { func openAICacheCreationTokensFromUsage(value gjson.Result) int { for _, field := range []string{ + "input_tokens_details.cache_write_tokens", + "prompt_tokens_details.cache_write_tokens", + "input_tokens_details.cache_creation_tokens", + "prompt_tokens_details.cache_creation_tokens", + "cache_write_tokens", "cache_creation_input_tokens", "cache_write_input_tokens", "cache_creation_tokens", - "cache_write_tokens", - "input_tokens_details.cache_creation_tokens", - "input_tokens_details.cache_write_tokens", - "prompt_tokens_details.cache_creation_tokens", - "prompt_tokens_details.cache_write_tokens", } { if tokens := int(value.Get(field).Int()); tokens > 0 { return tokens diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index a21cb98783..4528a8886f 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -35,6 +35,48 @@ var ( Mode: "chat", SupportsPromptCaching: true, } + openAIGPT56SolFallbackPricing = &LiteLLMModelPricing{ + InputCostPerToken: 5e-06, + InputCostPerTokenPriority: 1e-05, + OutputCostPerToken: 3e-05, + OutputCostPerTokenPriority: 6e-05, + CacheCreationInputTokenCost: 6.25e-06, + CacheCreationInputTokenCostPriority: 1.25e-05, + CacheReadInputTokenCost: 5e-07, + CacheReadInputTokenCostPriority: 1e-06, + SupportsServiceTier: true, + LiteLLMProvider: "openai", + Mode: "chat", + SupportsPromptCaching: true, + } + openAIGPT56TerraFallbackPricing = &LiteLLMModelPricing{ + InputCostPerToken: 2.5e-06, + InputCostPerTokenPriority: 5e-06, + OutputCostPerToken: 1.5e-05, + OutputCostPerTokenPriority: 3e-05, + CacheCreationInputTokenCost: 3.125e-06, + CacheCreationInputTokenCostPriority: 6.25e-06, + CacheReadInputTokenCost: 2.5e-07, + CacheReadInputTokenCostPriority: 5e-07, + SupportsServiceTier: true, + LiteLLMProvider: "openai", + Mode: "chat", + SupportsPromptCaching: true, + } + openAIGPT56LunaFallbackPricing = &LiteLLMModelPricing{ + InputCostPerToken: 1e-06, + InputCostPerTokenPriority: 2e-06, + OutputCostPerToken: 6e-06, + OutputCostPerTokenPriority: 1.2e-05, + CacheCreationInputTokenCost: 1.25e-06, + CacheCreationInputTokenCostPriority: 2.5e-06, + CacheReadInputTokenCost: 1e-07, + CacheReadInputTokenCostPriority: 2e-07, + SupportsServiceTier: true, + LiteLLMProvider: "openai", + Mode: "chat", + SupportsPromptCaching: true, + } openAIGPT54MiniFallbackPricing = &LiteLLMModelPricing{ InputCostPerToken: 7.5e-07, OutputCostPerToken: 4.5e-06, @@ -842,11 +884,20 @@ func (s *PricingService) matchOpenAIModel(model string) *LiteLLMModelPricing { } } - // GPT-5.6(sol / terra / luna)回退到 GPT-5.4 定价 - if strings.HasPrefix(model, "gpt-5.6") { + if strings.HasPrefix(model, "gpt-5.6-sol") { logger.With(zap.String("component", "service.pricing")). - Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.4(static)")) - return openAIGPT54FallbackPricing + Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.6-sol(static)")) + return openAIGPT56SolFallbackPricing + } + if strings.HasPrefix(model, "gpt-5.6-terra") { + logger.With(zap.String("component", "service.pricing")). + Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.6-terra(static)")) + return openAIGPT56TerraFallbackPricing + } + if strings.HasPrefix(model, "gpt-5.6-luna") { + logger.With(zap.String("component", "service.pricing")). + Info(fmt.Sprintf("[Pricing] OpenAI fallback matched %s -> %s", model, "gpt-5.6-luna(static)")) + return openAIGPT56LunaFallbackPricing } // GPT-5.5 回退到 GPT-5.4 定价 diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index e8f00ae3b2..19098df917 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -40,43 +40,57 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { require.True(t, pricing.SupportsServiceTier) } -func TestBillingService_GPT56CacheWritePricingUsesInputTier(t *testing.T) { - for _, model := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { - t.Run(model, func(t *testing.T) { +func TestBillingService_GPT56CacheWritePricingUsesOfficialMultiplier(t *testing.T) { + tests := []struct { + model string + input float64 + inputPriority float64 + output float64 + outputPriority float64 + cacheRead float64 + cacheReadPriority float64 + }{ + {model: "gpt-5.6-sol", input: 5e-6, inputPriority: 10e-6, output: 30e-6, outputPriority: 60e-6, cacheRead: 0.5e-6, cacheReadPriority: 1e-6}, + {model: "gpt-5.6-terra", input: 2.5e-6, inputPriority: 5e-6, output: 15e-6, outputPriority: 30e-6, cacheRead: 0.25e-6, cacheReadPriority: 0.5e-6}, + {model: "gpt-5.6-luna", input: 1e-6, inputPriority: 2e-6, output: 6e-6, outputPriority: 12e-6, cacheRead: 0.1e-6, cacheReadPriority: 0.2e-6}, + } + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ - model: { - InputCostPerToken: 5e-6, - InputCostPerTokenPriority: 10e-6, - OutputCostPerToken: 30e-6, - OutputCostPerTokenPriority: 60e-6, - CacheReadInputTokenCost: 0.5e-6, - CacheReadInputTokenCostPriority: 1e-6, + tt.model: { + InputCostPerToken: tt.input, + InputCostPerTokenPriority: tt.inputPriority, + OutputCostPerToken: tt.output, + OutputCostPerTokenPriority: tt.outputPriority, + CacheReadInputTokenCost: tt.cacheRead, + CacheReadInputTokenCostPriority: tt.cacheReadPriority, }, }} svc := NewBillingService(&config.Config{}, pricingSvc) - pricing, err := svc.GetModelPricing(model) + pricing, err := svc.GetModelPricing(tt.model) require.NoError(t, err) - require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12) - require.InDelta(t, 10e-6, pricing.CacheCreationPricePerTokenPriority, 1e-12) + require.InDelta(t, tt.input*1.25, pricing.CacheCreationPricePerToken, 1e-12) + require.InDelta(t, tt.inputPriority*1.25, pricing.CacheCreationPricePerTokenPriority, 1e-12) + require.Zero(t, pricing.LongContextInputThreshold) tokens := UsageTokens{InputTokens: 700, OutputTokens: 50, CacheCreationTokens: 200, CacheReadTokens: 100} - standard, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "") + standard, err := svc.CalculateCostWithServiceTier(tt.model, tokens, 1, "") require.NoError(t, err) - require.InDelta(t, 200*5e-6, standard.CacheCreationCost, 1e-12) + require.InDelta(t, 200*tt.input*1.25, standard.CacheCreationCost, 1e-12) - priority, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "priority") + priority, err := svc.CalculateCostWithServiceTier(tt.model, tokens, 1, "priority") require.NoError(t, err) - require.InDelta(t, 200*10e-6, priority.CacheCreationCost, 1e-12) + require.InDelta(t, 200*tt.inputPriority*1.25, priority.CacheCreationCost, 1e-12) - flex, err := svc.CalculateCostWithServiceTier(model, tokens, 1, "flex") + flex, err := svc.CalculateCostWithServiceTier(tt.model, tokens, 1, "flex") require.NoError(t, err) - require.InDelta(t, 200*2.5e-6, flex.CacheCreationCost, 1e-12) + require.InDelta(t, 200*tt.input*1.25*0.5, flex.CacheCreationCost, 1e-12) }) } } -func TestBillingService_GPT56CacheWriteContributesToLongContextThreshold(t *testing.T) { +func TestBillingService_GPT56DoesNotUseLegacyLongContextMultiplier(t *testing.T) { model := "gpt-5.6-sol" pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ model: { @@ -90,9 +104,84 @@ func TestBillingService_GPT56CacheWriteContributesToLongContextThreshold(t *test cost, err := svc.CalculateCost(model, tokens, 1) require.NoError(t, err) - require.InDelta(t, 100000*10e-6, cost.InputCost, 1e-12) - require.InDelta(t, 173000*10e-6, cost.CacheCreationCost, 1e-12) - require.InDelta(t, 10*45e-6, cost.OutputCost, 1e-12) + require.InDelta(t, 100000*5e-6, cost.InputCost, 1e-12) + require.InDelta(t, 173000*6.25e-6, cost.CacheCreationCost, 1e-12) + require.InDelta(t, 10*30e-6, cost.OutputCost, 1e-12) +} + +func TestDefaultPricingIncludesOfficialGPT56Rates(t *testing.T) { + data, err := os.ReadFile(filepath.Join("..", "..", "resources", "model-pricing", "model_prices_and_context_window.json")) + require.NoError(t, err) + + pricingSvc := &PricingService{} + pricingData, err := pricingSvc.parsePricingData(data) + require.NoError(t, err) + pricingSvc.pricingData = pricingData + billingSvc := NewBillingService(&config.Config{}, pricingSvc) + + tests := []struct { + model string + input, cached, cacheWrite, output float64 + inputPriority, cachedPriority, cacheWritePriority, outputPriority float64 + }{ + {model: "gpt-5.6-sol", input: 5e-6, cached: 0.5e-6, cacheWrite: 6.25e-6, output: 30e-6, inputPriority: 10e-6, cachedPriority: 1e-6, cacheWritePriority: 12.5e-6, outputPriority: 60e-6}, + {model: "gpt-5.6-terra", input: 2.5e-6, cached: 0.25e-6, cacheWrite: 3.125e-6, output: 15e-6, inputPriority: 5e-6, cachedPriority: 0.5e-6, cacheWritePriority: 6.25e-6, outputPriority: 30e-6}, + {model: "gpt-5.6-luna", input: 1e-6, cached: 0.1e-6, cacheWrite: 1.25e-6, output: 6e-6, inputPriority: 2e-6, cachedPriority: 0.2e-6, cacheWritePriority: 2.5e-6, outputPriority: 12e-6}, + } + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { + pricing, err := billingSvc.GetModelPricing(tt.model) + require.NoError(t, err) + require.InDelta(t, tt.input, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, tt.cached, pricing.CacheReadPricePerToken, 1e-12) + require.InDelta(t, tt.cacheWrite, pricing.CacheCreationPricePerToken, 1e-12) + require.InDelta(t, tt.output, pricing.OutputPricePerToken, 1e-12) + require.InDelta(t, tt.inputPriority, pricing.InputPricePerTokenPriority, 1e-12) + require.InDelta(t, tt.cachedPriority, pricing.CacheReadPricePerTokenPriority, 1e-12) + require.InDelta(t, tt.cacheWritePriority, pricing.CacheCreationPricePerTokenPriority, 1e-12) + require.InDelta(t, tt.outputPriority, pricing.OutputPricePerTokenPriority, 1e-12) + require.Zero(t, pricing.LongContextInputThreshold) + }) + } +} + +func TestGPT56DedicatedFallbacksUseOfficialRates(t *testing.T) { + tests := []struct { + model string + input, cached, cacheWrite, output float64 + }{ + {model: "gpt-5.6-sol", input: 5e-6, cached: 0.5e-6, cacheWrite: 6.25e-6, output: 30e-6}, + {model: "gpt-5.6-terra", input: 2.5e-6, cached: 0.25e-6, cacheWrite: 3.125e-6, output: 15e-6}, + {model: "gpt-5.6-luna", input: 1e-6, cached: 0.1e-6, cacheWrite: 1.25e-6, output: 6e-6}, + } + + for _, tt := range tests { + t.Run(tt.model+"/pricing_service", func(t *testing.T) { + pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ + "gpt-5.1-codex": {InputCostPerToken: 1.25e-6}, + }} + svc := NewBillingService(&config.Config{}, pricingSvc) + pricing, err := svc.GetModelPricing(tt.model + "-preview") + require.NoError(t, err) + assertGPT56FallbackPricing(t, pricing, tt.input, tt.cached, tt.cacheWrite, tt.output) + }) + + t.Run(tt.model+"/billing_service", func(t *testing.T) { + svc := NewBillingService(&config.Config{}, nil) + pricing, err := svc.GetModelPricing(tt.model) + require.NoError(t, err) + assertGPT56FallbackPricing(t, pricing, tt.input, tt.cached, tt.cacheWrite, tt.output) + }) + } +} + +func assertGPT56FallbackPricing(t *testing.T, pricing *ModelPricing, input, cached, cacheWrite, output float64) { + t.Helper() + require.InDelta(t, input, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, cached, pricing.CacheReadPricePerToken, 1e-12) + require.InDelta(t, cacheWrite, pricing.CacheCreationPricePerToken, 1e-12) + require.InDelta(t, output, pricing.OutputPricePerToken, 1e-12) + require.Zero(t, pricing.LongContextInputThreshold) } func TestParsePricingData_KeepsImageOnlyPricing(t *testing.T) { diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index b9d18f8bf5..439988a53d 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -4961,12 +4961,14 @@ "supports_xhigh_reasoning_effort": true }, "gpt-5.6-sol": { + "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_batches": 3.125e-06, + "cache_creation_input_token_cost_flex": 3.125e-06, + "cache_creation_input_token_cost_priority": 1.25e-05, "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, "cache_read_input_token_cost_flex": 2.5e-07, "cache_read_input_token_cost_priority": 1e-06, "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, "input_cost_per_token_batches": 2.5e-06, "input_cost_per_token_flex": 2.5e-06, "input_cost_per_token_priority": 1e-05, @@ -4976,7 +4978,6 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, "output_cost_per_token_batches": 1.5e-05, "output_cost_per_token_flex": 1.5e-05, "output_cost_per_token_priority": 6e-05, @@ -5009,25 +5010,26 @@ "supports_xhigh_reasoning_effort": true }, "gpt-5.6-terra": { - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_flex": 2.5e-07, - "cache_read_input_token_cost_priority": 1e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_batches": 2.5e-06, - "input_cost_per_token_flex": 2.5e-06, - "input_cost_per_token_priority": 1e-05, + "cache_creation_input_token_cost": 3.125e-06, + "cache_creation_input_token_cost_batches": 1.5625e-06, + "cache_creation_input_token_cost_flex": 1.5625e-06, + "cache_creation_input_token_cost_priority": 6.25e-06, + "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_flex": 1.25e-07, + "cache_read_input_token_cost_priority": 5e-07, + "input_cost_per_token": 2.5e-06, + "input_cost_per_token_batches": 1.25e-06, + "input_cost_per_token_flex": 1.25e-06, + "input_cost_per_token_priority": 5e-06, "litellm_provider": "openai", "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_batches": 1.5e-05, - "output_cost_per_token_flex": 1.5e-05, - "output_cost_per_token_priority": 6e-05, + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, + "output_cost_per_token_flex": 7.5e-06, + "output_cost_per_token_priority": 3e-05, "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -5057,25 +5059,26 @@ "supports_xhigh_reasoning_effort": true }, "gpt-5.6-luna": { - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_flex": 2.5e-07, - "cache_read_input_token_cost_priority": 1e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_batches": 2.5e-06, - "input_cost_per_token_flex": 2.5e-06, - "input_cost_per_token_priority": 1e-05, + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, + "cache_creation_input_token_cost_flex": 6.25e-07, + "cache_creation_input_token_cost_priority": 2.5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_flex": 5e-08, + "cache_read_input_token_cost_priority": 2e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, + "input_cost_per_token_flex": 5e-07, + "input_cost_per_token_priority": 2e-06, "litellm_provider": "openai", "max_input_tokens": 1050000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_batches": 1.5e-05, - "output_cost_per_token_flex": 1.5e-05, - "output_cost_per_token_priority": 6e-05, + "output_cost_per_token": 6e-06, + "output_cost_per_token_batches": 3e-06, + "output_cost_per_token_flex": 3e-06, + "output_cost_per_token_priority": 1.2e-05, "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", From 657c4f97d904d796c806c7bd612ac5929307dcaf Mon Sep 17 00:00:00 2001 From: li Date: Fri, 10 Jul 2026 09:56:33 +0800 Subject: [PATCH 05/27] =?UTF-8?q?fix(openai):=20=E5=8D=87=E7=BA=A7=20Codex?= =?UTF-8?q?=20=E5=AE=A2=E6=88=B7=E7=AB=AF=E7=89=88=E6=9C=AC=E8=87=B3=200.1?= =?UTF-8?q?44.1=EF=BC=8C=E4=BF=AE=E5=A4=8D=20gpt-5.6-luna=20404?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #3895 --- .../internal/service/account_usage_service.go | 2 +- .../openai_codex_version_consistency_test.go | 21 +++++++++++++++++++ .../service/openai_gateway_service.go | 4 ++-- .../service/setting_gateway_runtime.go | 2 +- 4 files changed, 25 insertions(+), 4 deletions(-) create mode 100644 backend/internal/service/openai_codex_version_consistency_test.go diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index d9f6200359..4add6310c6 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -111,7 +111,7 @@ const ( apiQueryMaxJitter = 800 * time.Millisecond // 用量查询最大随机延迟 windowStatsCacheTTL = 1 * time.Minute openAIProbeCacheTTL = 10 * time.Minute - openAICodexProbeVersion = "0.125.0" + openAICodexProbeVersion = "0.144.1" ) // UsageCache 封装账户使用量相关的缓存 diff --git a/backend/internal/service/openai_codex_version_consistency_test.go b/backend/internal/service/openai_codex_version_consistency_test.go new file mode 100644 index 0000000000..48883032aa --- /dev/null +++ b/backend/internal/service/openai_codex_version_consistency_test.go @@ -0,0 +1,21 @@ +//go:build unit + +package service + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCodexVersionConstants_Consistency(t *testing.T) { + require.Equal(t, codexCLIVersion, openAICodexProbeVersion, + "codexCLIVersion and openAICodexProbeVersion must stay in sync") + + require.True(t, strings.Contains(codexCLIUserAgent, "codex_cli_rs/"+codexCLIVersion), + "codexCLIUserAgent must embed codexCLIVersion") + + require.True(t, strings.Contains(DefaultOpenAICodexUserAgent, codexCLIVersion), + "DefaultOpenAICodexUserAgent must embed codexCLIVersion") +} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index af8ae3ffe5..0d03d83dcf 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -35,7 +35,7 @@ const ( // 与真实 Codex CLI 的 User-Agent 结构对齐: // {originator}/{version} ({OS} {OS_version}; {arch}) {terminal} // 旧值 "codex_cli_rs/0.125.0" 缺少 OS/架构/终端后缀,易被上游指纹识别为非官方客户端。 - codexCLIUserAgent = "codex_cli_rs/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color" + codexCLIUserAgent = "codex_cli_rs/0.144.1 (Ubuntu 22.4.0; x86_64) xterm-256color" // codex_cli_only 拒绝时单个请求头日志长度上限(字符) codexCLIOnlyHeaderValueMaxBytes = 256 @@ -49,7 +49,7 @@ const ( openAIWSRetryBackoffMaxDefault = 2 * time.Second openAIWSRetryJitterRatioDefault = 0.2 openAICompactSessionSeedKey = "openai_compact_session_seed" - codexCLIVersion = "0.125.0" + codexCLIVersion = "0.144.1" // Codex 限额快照仅用于后台展示/诊断,不需要每个成功请求都立即落库。 openAICodexSnapshotPersistMinInterval = 30 * time.Second // 配额自动暂停时,超过该时长仍未刷新的 used% 快照视为陈旧,不再据此暂停账号。 diff --git a/backend/internal/service/setting_gateway_runtime.go b/backend/internal/service/setting_gateway_runtime.go index ea4037d990..3ed4f39def 100644 --- a/backend/internal/service/setting_gateway_runtime.go +++ b/backend/internal/service/setting_gateway_runtime.go @@ -83,7 +83,7 @@ const antigravityUserAgentVersionErrorTTL = 5 * time.Second const antigravityUserAgentVersionDBTimeout = 5 * time.Second // DefaultOpenAICodexUserAgent OpenAI Codex 默认 User-Agent(用于规避 Cloudflare 对浏览器 UA 的质询) -const DefaultOpenAICodexUserAgent = "codex-tui/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.125.0)" +const DefaultOpenAICodexUserAgent = "codex-tui/0.144.1 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.144.1)" // cachedOpenAICodexUserAgent 缓存 OpenAI Codex UA(进程内缓存,60s TTL) type cachedOpenAICodexUserAgent struct { From 2cffe1cf5f1c20413790f109d234fa15be544191 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 10 Jul 2026 09:56:56 +0800 Subject: [PATCH 06/27] =?UTF-8?q?fix(compact):=20SSE=E2=86=92JSON=20?= =?UTF-8?q?=E4=BF=9D=E7=95=99=20raw=20output=5Fitem.done=20=E5=B9=B6?= =?UTF-8?q?=E4=B8=BA=20unary=20=E7=AD=89=E5=BE=85=E8=A1=A5=E4=B8=8B?= =?UTF-8?q?=E6=B8=B8=E5=BF=83=E8=B7=B3=EF=BC=88=E4=BF=AE=E5=A4=8D=20#3887?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit #3887 报告 v0.1.149(已含 #3880 桥接修复)remote compact 仍失败且持续 计费。核实为 #3880 明示的两个遗留: 1. 上游对 unary compact 返回 SSE 且 compaction item 只出现在 raw response.output_item.done、终态 completed.response.output 为空 (#3777 实录形态)时,reconstructResponseOutputFromSSE 只累加 text/function_call/reasoning 三类 delta,compaction item 被丢弃, 桥接合成 0 个 output_item.done,Codex 报 "expected exactly one compaction output item, got 0" 后盲目重试,每次重试重新消耗上游 compact 配额。 2. 桥接为全缓冲写回:上游 unary 完成前(大上下文可达数分钟,且上游 处理期间不发送任何字节)下游连响应头都收不到,反向代理 (Nginx/Cloudflare Tunnel)空闲超时掐断连接同样触发盲重连 (同类问题见 #2243/#2976)。 修复: - reconstructResponseOutputFromSSE 优先以 raw JSON 逐字节收集 output_item.done item(协议上的最终完整形态),不经窄结构体, encrypted_content/summary/opaque 等 compact 专属字段全部保留; 无 done 事件时退回收集 output_item.added 中的 compaction 类 item; 两者皆无才回到原 delta 重建。path-based v1 JSON 写回同样受益。 - 新增 openAICompactSSEKeepalive:body-signal 客户端流式 compact 在 上游等待期间按 gateway.stream_keepalive_interval 向下游写 SSE 注释 行心跳(eventsource 解析层直接忽略)。首拍延迟一个间隔,快速失败 仍走 JSON+状态码;首拍后状态码固化为 200,桥接/错误链路 (writeOpenAICompactSSEBridge、errorResponse、 handleStreamingAwareError、ensureForwardErrorResponse、 writeOpenAINonStreamingProtocolError)统一降级为 response.failed 终止事件并标记 ops 流内错误。 - API-key 账号的 compact 上游请求也强制 accept: application/json (#3777 期望行为 4;透传白名单原会放行客户端的 text/event-stream)。 测试:#3777 实录形态经 handleSSEToJSON / 透传 / path-based 三条链路 的修补断言;raw-done 优先不与 delta 重复;added 回退门控;心跳提交、 提交后 2xx 续写、提交后失败降级、未提交行为不变;-race 通过。 Fixes #3887 Refs #3777 #3875 #3880 --- .../handler/openai_gateway_handler.go | 31 ++++ .../service/openai_compact_sse_keepalive.go | 129 +++++++++++++++ .../openai_compact_sse_keepalive_test.go | 122 ++++++++++++++ .../service/openai_compact_stream_bridge.go | 81 +++++++++- .../openai_compact_stream_bridge_test.go | 150 ++++++++++++++++++ .../service/openai_gateway_forward.go | 4 + .../service/openai_gateway_passthrough.go | 5 + .../openai_gateway_response_handling.go | 82 +++++++++- 8 files changed, 594 insertions(+), 10 deletions(-) create mode 100644 backend/internal/service/openai_compact_sse_keepalive.go create mode 100644 backend/internal/service/openai_compact_sse_keepalive_test.go diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index feba261103..714e7906f3 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -205,6 +205,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { if !ok { return } + // body-signal compact:上游 unary 等待期间向下游发 SSE 注释行心跳,防止 + // 反向代理空闲超时掐断长压缩连接(#3887)。首拍延迟一个心跳间隔,快速 + // 失败仍走 JSON+状态码链路;未标记客户端流式或间隔为 0 时是 no-op。 + stopCompactKeepalive := service.StartOpenAICompactSSEKeepalive(c, h.openAICompactKeepaliveInterval()) + defer stopCompactKeepalive() // 校验请求体 JSON 合法性 if !gjson.ValidBytes(body) { @@ -1943,6 +1948,11 @@ func (h *OpenAIGatewayHandler) mapUpstreamError(statusCode int) (int, string, st // handleStreamingAwareError handles errors that may occur after streaming has started func (h *OpenAIGatewayHandler) handleStreamingAwareError(c *gin.Context, status int, errType, message string, streamStarted bool) { + // body-signal compact 心跳可能已把响应头提交为 200:先停心跳(建立 + // happens-before,接管 ResponseWriter),并升级为流内错误处理。 + if service.StopOpenAICompactSSEKeepaliveCommitted(c) { + streamStarted = true + } if streamStarted { // /v1/responses 的严格 SDK(Codex CLI)要求终止事件必须属于 // response.completed/failed/incomplete/cancelled 集合。 @@ -1975,6 +1985,10 @@ func (h *OpenAIGatewayHandler) ensureForwardErrorResponse(c *gin.Context, stream if c == nil || c.Writer == nil { return false } + // 先停 compact 心跳再读 Writer 状态,避免与心跳 goroutine 竞争。 + if service.StopOpenAICompactSSEKeepaliveCommitted(c) { + streamStarted = true + } if service.IsResponseCommitted(c) { return false } @@ -2036,6 +2050,14 @@ func openAIForwardErrorAlreadyCommunicated(c *gin.Context, writerSizeBeforeForwa // errorResponse returns OpenAI API format error response func (h *OpenAIGatewayHandler) errorResponse(c *gin.Context, status int, errType, message string) { + // body-signal compact 心跳可能已把响应头提交为 200:JSON 错误体会与已 + // 提交的 SSE 流交错,必须降级为 response.failed 终止事件(#3887)。 + if service.StopOpenAICompactSSEKeepaliveCommitted(c) { + service.MarkOpsStreamError(c, errType, message, status) + if writeResponsesFailedSSE(c, errType, message) { + return + } + } c.JSON(status, gin.H{ "error": gin.H{ "type": errType, @@ -2044,6 +2066,15 @@ func (h *OpenAIGatewayHandler) errorResponse(c *gin.Context, status int, errType }) } +// openAICompactKeepaliveInterval 复用流式 keepalive 配置作为 compact 下游 +// 心跳间隔;0 表示禁用(与流式路径语义一致)。 +func (h *OpenAIGatewayHandler) openAICompactKeepaliveInterval() time.Duration { + if h.cfg == nil || h.cfg.Gateway.StreamKeepaliveInterval <= 0 { + return 0 + } + return time.Duration(h.cfg.Gateway.StreamKeepaliveInterval) * time.Second +} + func setOpenAIClientTransportHTTP(c *gin.Context) { service.SetOpenAIClientTransport(c, service.OpenAIClientTransportHTTP) } diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go new file mode 100644 index 0000000000..0695d16fb3 --- /dev/null +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -0,0 +1,129 @@ +package service + +import ( + "net/http" + "sync" + "time" + + "github.com/gin-gonic/gin" +) + +// openAICompactSSEKeepaliveKey 存放 body-signal compact 请求的下游 SSE 心跳器。 +const openAICompactSSEKeepaliveKey = "openai_compact_sse_keepalive" + +// openAICompactSSEKeepalive 在 compact 上游 unary 等待期间向下游写 SSE 注释行 +// 心跳。上游 /responses/compact 在模型处理期间不发送任何字节(大上下文可长达 +// 数分钟),下游若经过反向代理(Nginx/Cloudflare Tunnel 等),零字节静默会触发 +// 代理的空闲/读超时并掐断连接,Codex 只会盲目重连并重复消耗上游 compact +// 配额(#3887)。SSE 注释行在 eventsource 解析层被直接忽略,不会进入客户端 +// 事件流。 +// +// 首拍延迟一个 interval:绝大多数硬错误(鉴权/参数/限流)在此之前返回,仍走 +// 原 JSON+状态码链路(Codex 按 HTTP 状态码重试);首拍之后状态码固化为 200, +// 后续错误由写回方降级为 response.failed 流内终止事件。 +type openAICompactSSEKeepalive struct { + mu sync.Mutex + writer gin.ResponseWriter + started bool + stopped bool + stop chan struct{} +} + +// StartOpenAICompactSSEKeepalive 为已标记 body-signal 客户端流式的 compact +// 请求启动下游心跳,返回幂等的停止函数。interval<=0 或请求未标记时为 no-op。 +func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func() { + if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) { + return func() {} + } + k := &openAICompactSSEKeepalive{ + writer: c.Writer, + stop: make(chan struct{}), + } + c.Set(openAICompactSSEKeepaliveKey, k) + + var reqDone <-chan struct{} + if c.Request != nil { + reqDone = c.Request.Context().Done() + } + go func() { + timer := time.NewTimer(interval) + defer timer.Stop() + for { + select { + case <-k.stop: + return + case <-reqDone: + return + case <-timer.C: + } + if !k.beat() { + return + } + timer.Reset(interval) + } + }() + return k.Stop +} + +// beat 在锁内提交(首次)响应头并写出一条 SSE 注释行;返回 false 表示心跳已 +// 停止或下游写入失败,goroutine 应退出。 +func (k *openAICompactSSEKeepalive) beat() bool { + k.mu.Lock() + defer k.mu.Unlock() + if k.stopped { + return false + } + if !k.started { + header := k.writer.Header() + header.Set("Content-Type", "text/event-stream") + header.Set("Cache-Control", "no-cache") + header.Set("Connection", "keep-alive") + header.Set("X-Accel-Buffering", "no") + k.writer.WriteHeader(http.StatusOK) + k.started = true + } + if _, err := k.writer.Write([]byte(": keepalive\n\n")); err != nil { + k.stopped = true + return false + } + k.writer.Flush() + return true +} + +// Stop 停止心跳;幂等,可与写回路径并发调用。 +func (k *openAICompactSSEKeepalive) Stop() { + k.mu.Lock() + k.markStoppedLocked() + k.mu.Unlock() +} + +func (k *openAICompactSSEKeepalive) markStoppedLocked() { + if k.stopped { + return + } + k.stopped = true + close(k.stop) +} + +// StopOpenAICompactSSEKeepaliveCommitted 停止当前请求的 compact 心跳(若有) +// 并报告响应头是否已被心跳提交为 200。写回方以此决定继续走原 JSON/状态码 +// 链路,还是降级为流内终止事件。调用后不会再有心跳字节写出,且经由互斥锁 +// 与心跳 goroutine 建立 happens-before,调用方可安全接管 ResponseWriter。 +func StopOpenAICompactSSEKeepaliveCommitted(c *gin.Context) bool { + if c == nil { + return false + } + value, ok := c.Get(openAICompactSSEKeepaliveKey) + if !ok { + return false + } + k, ok := value.(*openAICompactSSEKeepalive) + if !ok || k == nil { + return false + } + k.mu.Lock() + k.markStoppedLocked() + committed := k.started + k.mu.Unlock() + return committed +} diff --git a/backend/internal/service/openai_compact_sse_keepalive_test.go b/backend/internal/service/openai_compact_sse_keepalive_test.go new file mode 100644 index 0000000000..4abff79fc1 --- /dev/null +++ b/backend/internal/service/openai_compact_sse_keepalive_test.go @@ -0,0 +1,122 @@ +package service + +import ( + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +const keepaliveTestInterval = 10 * time.Millisecond + +// waitForKeepaliveBeats 等待至少一次心跳写出。读取 recorder 前必须先经 +// StopOpenAICompactSSEKeepaliveCommitted 停拍建立 happens-before。 +func waitForKeepaliveBeats() { + time.Sleep(20 * keepaliveTestInterval) +} + +// stripKeepaliveComments 去掉 SSE 注释块,返回真实事件文本。 +func stripKeepaliveComments(body string) string { + var blocks []string + for _, block := range strings.Split(strings.TrimSpace(body), "\n\n") { + if strings.HasPrefix(strings.TrimSpace(block), ":") { + continue + } + blocks = append(blocks, block) + } + return strings.Join(blocks, "\n\n") +} + +func TestStartOpenAICompactSSEKeepalive_NoopWhenUnmarkedOrDisabled(t *testing.T) { + // 未标记 client stream:不启动。 + c, rec := newCompactBridgeTestContext(t, false) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + waitForKeepaliveBeats() + stop() + require.Zero(t, rec.Body.Len()) + require.False(t, StopOpenAICompactSSEKeepaliveCommitted(c)) + + // interval=0(配置禁用):不启动。 + c, rec = newCompactBridgeTestContext(t, true) + stop = StartOpenAICompactSSEKeepalive(c, 0) + waitForKeepaliveBeats() + stop() + require.Zero(t, rec.Body.Len()) + require.False(t, StopOpenAICompactSSEKeepaliveCommitted(c)) +} + +func TestOpenAICompactSSEKeepalive_CommitsHeadersAndComments(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + waitForKeepaliveBeats() + + require.True(t, StopOpenAICompactSSEKeepaliveCommitted(c)) + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + require.Equal(t, "no", rec.Header().Get("X-Accel-Buffering")) + require.Contains(t, rec.Body.String(), ": keepalive\n\n") +} + +func TestOpenAICompactSSEKeepalive_StopBeforeFirstBeatKeepsWriterUntouched(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, time.Hour) + stop() + waitForKeepaliveBeats() + require.Zero(t, rec.Body.Len()) + require.False(t, StopOpenAICompactSSEKeepaliveCommitted(c)) +} + +// 心跳已提交后,2xx 桥接续写事件而不重复提交响应头。 +func TestWriteOpenAICompactSSEBridge_AfterKeepaliveCommitAppendsEvents(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + waitForKeepaliveBeats() + + finalResponse := []byte(`{"id":"resp_ka_1","output":[{"id":"cmp_ka","type":"compaction","encrypted_content":"x"}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`) + require.True(t, writeOpenAICompactSSEBridge(c, http.StatusOK, finalResponse)) + + require.Equal(t, http.StatusOK, rec.Code) + events := parseCompactBridgeSSE(t, stripKeepaliveComments(rec.Body.String())) + require.Len(t, events, 2) + require.Equal(t, "response.output_item.done", events[0][0]) + require.Equal(t, "compaction", gjson.Get(events[0][1], "item.type").String()) + require.Equal(t, "response.completed", events[1][0]) + require.Equal(t, "resp_ka_1", gjson.Get(events[1][1], "response.id").String()) +} + +// 心跳已提交后上游非 2xx:状态码无法回传,必须以 response.failed 终止事件 +// 收尾(Codex 将其作为终止事件处理),并标记流内错误供 ops 采集。 +func TestWriteOpenAICompactSSEBridge_AfterKeepaliveCommitFailureEmitsFailedEvent(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + waitForKeepaliveBeats() + + require.True(t, writeOpenAICompactSSEBridge(c, http.StatusBadGateway, []byte(`{"error":{"message":"upstream exploded"}}`))) + + events := parseCompactBridgeSSE(t, stripKeepaliveComments(rec.Body.String())) + require.Len(t, events, 1) + require.Equal(t, "response.failed", events[0][0]) + require.Equal(t, "failed", gjson.Get(events[0][1], "response.status").String()) + require.Contains(t, gjson.Get(events[0][1], "response.error.message").String(), "upstream exploded") + require.NotEmpty(t, gjson.Get(events[0][1], "response.id").String()) + + streamErr, ok := GetOpsStreamError(c) + require.True(t, ok) + require.Equal(t, http.StatusBadGateway, streamErr.IntendedStatus) +} + +// 心跳未提交时非 2xx 行为不变:返回 false,调用方按原 JSON+状态码写回。 +func TestWriteOpenAICompactSSEBridge_BeforeKeepaliveCommitFailureKeepsJSONPath(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, time.Hour) + stop() + + require.False(t, writeOpenAICompactSSEBridge(c, http.StatusBadGateway, []byte(`{"error":{"message":"fast fail"}}`))) + require.Zero(t, rec.Body.Len()) +} diff --git a/backend/internal/service/openai_compact_stream_bridge.go b/backend/internal/service/openai_compact_stream_bridge.go index 33ca543850..98240f4229 100644 --- a/backend/internal/service/openai_compact_stream_bridge.go +++ b/backend/internal/service/openai_compact_stream_bridge.go @@ -3,6 +3,8 @@ package service import ( "bytes" "encoding/json" + "net/http" + "strconv" "strings" "github.com/gin-gonic/gin" @@ -48,25 +50,90 @@ func openAICompactClientWantsStream(c *gin.Context) bool { // compact v2 的消费协议合成为最小 Responses SSE 流写回客户端。仅当请求被标记 // 为 body-signal 客户端流式、状态码为 2xx 且 body 是合法 JSON 对象时生效; // 返回 false 表示未写出任何内容,调用方应按原路径写回。 +// +// 若下游心跳已把响应头提交为 200(见 openAICompactSSEKeepalive),则本函数 +// 必须接管一切写回:非 2xx 或不可合成的响应降级为 response.failed 终止事件, +// 不能再返回 false(否则调用方的 JSON 写回会与已提交的 SSE 流交错)。 func writeOpenAICompactSSEBridge(c *gin.Context, statusCode int, finalResponse []byte) bool { - if c == nil || statusCode < 200 || statusCode >= 300 || !openAICompactClientWantsStream(c) { + if c == nil || !openAICompactClientWantsStream(c) { + return false + } + // 先停心跳再写回,避免注释行与最终事件交错;停止后经互斥锁与心跳 + // goroutine 建立 happens-before,可安全接管 ResponseWriter。 + committed := StopOpenAICompactSSEKeepaliveCommitted(c) + if statusCode < 200 || statusCode >= 300 { + if committed { + writeOpenAICompactSSEFailure(c, statusCode, finalResponse) + return true + } return false } payload, ok := buildOpenAICompactSSEPayload(finalResponse) if !ok { + if committed { + writeOpenAICompactSSEFailure(c, http.StatusBadGateway, finalResponse) + return true + } return false } - header := c.Writer.Header() - header.Set("Content-Type", "text/event-stream") - header.Set("Cache-Control", "no-cache") - header.Set("Connection", "keep-alive") - header.Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(statusCode) + if !committed { + header := c.Writer.Header() + header.Set("Content-Type", "text/event-stream") + header.Set("Cache-Control", "no-cache") + header.Set("Connection", "keep-alive") + header.Set("X-Accel-Buffering", "no") + c.Writer.WriteHeader(statusCode) + } _, _ = c.Writer.Write(payload) c.Writer.Flush() return true } +// writeOpenAICompactSSEFailure 从上游错误 body 提取错误消息后,以 +// response.failed 终止事件回传。仅用于心跳已提交 200、无法再按 HTTP 状态码 +// 回传错误的场景。 +func writeOpenAICompactSSEFailure(c *gin.Context, statusCode int, errorBody []byte) { + message := "" + if len(errorBody) > 0 { + message = sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(errorBody))) + } + if message == "" { + message = "Upstream compact request failed with HTTP " + strconv.Itoa(statusCode) + } + writeOpenAICompactSSEFailureMessage(c, statusCode, message) +} + +// writeOpenAICompactSSEFailureMessage 写出 response.failed 终止事件。Codex 对 +// 流式 Responses 请求把 response.failed 作为合法终止事件处理(普通 error 帧 +// 不被识别,会退化为 "stream closed before response.completed" 盲重连)。 +// 同时标记流内错误,保证挂在 200 流上的失败仍进入 ops 错误看板。 +func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, message string) { + if c == nil { + return + } + MarkOpsStreamError(c, "upstream_error", message, statusCode) + payload, err := json.Marshal(map[string]any{ + "type": "response.failed", + "response": map[string]any{ + "id": "resp_" + strings.ReplaceAll(uuid.NewString(), "-", ""), + "object": "response", + "status": "failed", + "output": []any{}, + "error": map[string]any{ + "code": "upstream_error", + "message": message, + }, + }, + }) + if err != nil { + return + } + _, _ = c.Writer.Write([]byte("event: response.failed\ndata: ")) + _, _ = c.Writer.Write(payload) + _, _ = c.Writer.Write([]byte("\n\n")) + c.Writer.Flush() +} + // buildOpenAICompactSSEPayload 把 compact 的 Response JSON 转成 SSE 事件序列: // 每个 output[] item 一条 response.output_item.done,最后一条 response.completed // 携带完整 response 对象。Codex 的 SSE 解析只从 output_item.done 收集 item, diff --git a/backend/internal/service/openai_compact_stream_bridge_test.go b/backend/internal/service/openai_compact_stream_bridge_test.go index 2eb5ba80ee..1f15084fd9 100644 --- a/backend/internal/service/openai_compact_stream_bridge_test.go +++ b/backend/internal/service/openai_compact_stream_bridge_test.go @@ -258,6 +258,156 @@ func TestHandleSSEToJSON_CompactClientStreamBridgesToSSE(t *testing.T) { require.Equal(t, "resp_compact_sse", gjson.Get(events[1][1], "response.id").String()) } +// 回归 #3887(#3777 问题 2):上游对 compact 返回 SSE,compaction item 只在 +// raw output_item.done 中、终态 response.completed 的 output 为空。SSE→JSON +// 提取必须保留 raw item 修补终态 output,否则桥接合成 0 个 output_item.done, +// Codex 报 "expected exactly one compaction output item, got 0" 并盲目重试, +// 每次重试都重新计费。fixture 取自 #3777 的上游实录形态。 +func TestHandleSSEToJSON_CompactRawOutputItemDoneRepairsEmptyTerminalOutput(t *testing.T) { + svc := newCompactBridgeTestService() + c, rec := newCompactBridgeTestContext(t, true) + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"cmp_1","type":"compaction_summary","status":"completed","summary":[{"type":"summary_text","text":"compact summary"}],"encrypted_content":"compact-payload","opaque":{"kept":true}}}`, + ``, + `data: {"type":"response.completed","response":{"id":"resp_compact","object":"response","model":"gpt-5.1-codex","status":"completed","output":[],"usage":{"input_tokens":9,"output_tokens":4,"total_tokens":13}}}`, + ``, + }, "\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + } + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, &Account{ID: 1, Type: AccountTypeOAuth}, "gpt-5.5", "gpt-5.5") + require.NoError(t, err) + require.NotNil(t, result) + + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + events := parseCompactBridgeSSE(t, rec.Body.String()) + require.Len(t, events, 2) + require.Equal(t, "response.output_item.done", events[0][0]) + item := gjson.Get(events[0][1], "item") + require.Equal(t, "compaction_summary", item.Get("type").String()) + require.Equal(t, "cmp_1", item.Get("id").String()) + require.Equal(t, "compact-payload", item.Get("encrypted_content").String()) + require.Equal(t, "compact summary", item.Get("summary.0.text").String()) + require.True(t, item.Get("opaque.kept").Bool(), "raw item 字段必须逐字节保留") + require.Equal(t, "response.completed", events[1][0]) + require.Equal(t, "resp_compact", gjson.Get(events[1][1], "response.id").String()) + require.Len(t, gjson.Get(events[1][1], "response.output").Array(), 1) + require.Equal(t, int64(13), gjson.Get(events[1][1], "response.usage.total_tokens").Int()) + + require.NotNil(t, result.usage) + require.Equal(t, 9, result.usage.InputTokens) + require.Equal(t, 4, result.usage.OutputTokens) +} + +// 同一形态经透传分支(handlePassthroughSSEToJSON)也必须修补。 +func TestHandlePassthroughSSEToJSON_CompactRawOutputItemDoneRepairsEmptyTerminalOutput(t *testing.T) { + svc := newCompactBridgeTestService() + c, rec := newCompactBridgeTestContext(t, true) + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"cmp_pt_1","type":"compaction","status":"completed","encrypted_content":"compact-pt-raw"}}`, + ``, + `data: {"type":"response.completed","response":{"id":"resp_compact_pt_raw","object":"response","status":"completed","output":[],"usage":{"input_tokens":6,"output_tokens":2,"total_tokens":8}}}`, + ``, + }, "\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + } + + result, err := svc.handleNonStreamingResponsePassthrough(context.Background(), resp, c, "gpt-5.5", "") + require.NoError(t, err) + require.NotNil(t, result) + + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + events := parseCompactBridgeSSE(t, rec.Body.String()) + require.Len(t, events, 2) + require.Equal(t, "compaction", gjson.Get(events[0][1], "item.type").String()) + require.Equal(t, "compact-pt-raw", gjson.Get(events[0][1], "item.encrypted_content").String()) + require.Len(t, gjson.Get(events[1][1], "response.output").Array(), 1) +} + +// path-based(Codex v1 unary、链式 sub2api)未标记 client stream:同一上游 +// 形态修补后仍按 JSON 写回,output 中必须包含 compaction item。 +func TestHandleSSEToJSON_PathBasedCompactRawOutputItemDoneRepairsJSON(t *testing.T) { + svc := newCompactBridgeTestService() + c, rec := newCompactBridgeTestContext(t, false) + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"cmp_v1","type":"compaction_summary","encrypted_content":"compact-v1-raw"}}`, + ``, + `data: {"type":"response.completed","response":{"id":"resp_compact_v1","object":"response","status":"completed","output":[],"usage":{"input_tokens":5,"output_tokens":1,"total_tokens":6}}}`, + ``, + }, "\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + } + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, &Account{ID: 1, Type: AccountTypeOAuth}, "gpt-5.5", "gpt-5.5") + require.NoError(t, err) + require.NotNil(t, result) + + // 写回 body 必须是修补后的 JSON 文档(非 SSE 事件流)。 + body := rec.Body.String() + require.NotContains(t, body, "event:") + require.NotContains(t, body, "data:") + require.Equal(t, "resp_compact_v1", gjson.Get(body, "id").String()) + require.Equal(t, "compaction_summary", gjson.Get(body, "output.0.type").String()) + require.Equal(t, "compact-v1-raw", gjson.Get(body, "output.0.encrypted_content").String()) +} + +// raw done item 是协议上的最终完整形态,优先于 delta 重建且不得重复计入。 +func TestReconstructResponseOutputFromSSE_PrefersRawDoneItems(t *testing.T) { + bodyText := strings.Join([]string{ + `data: {"type":"response.output_text.delta","delta":"hel"}`, + `data: {"type":"response.output_text.delta","delta":"lo"}`, + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"msg_1","type":"message","role":"assistant","status":"completed","content":[{"type":"output_text","text":"hello"}]}}`, + `data: {"type":"response.completed","response":{"id":"resp_1","output":[]}}`, + }, "\n") + + outputJSON, ok := reconstructResponseOutputFromSSE(bodyText) + require.True(t, ok) + items := gjson.ParseBytes(outputJSON).Array() + require.Len(t, items, 1, "raw done item 与 delta 重建不得重复") + require.Equal(t, "msg_1", items[0].Get("id").String()) + require.Equal(t, "hello", items[0].Get("content.0.text").String()) +} + +// 无任何 done 事件时,退回收集 output_item.added 中的 compaction 类 item。 +func TestReconstructResponseOutputFromSSE_CompactionAddedFallback(t *testing.T) { + bodyText := strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"id":"cmp_add","type":"compaction","encrypted_content":"added-only"}}`, + `data: {"type":"response.completed","response":{"id":"resp_1","output":[]}}`, + }, "\n") + + outputJSON, ok := reconstructResponseOutputFromSSE(bodyText) + require.True(t, ok) + items := gjson.ParseBytes(outputJSON).Array() + require.Len(t, items, 1) + require.Equal(t, "compaction", items[0].Get("type").String()) + require.Equal(t, "added-only", items[0].Get("encrypted_content").String()) +} + +// 非 compaction 的 output_item.added 不参与回退收集(added 阶段的 message +// 通常是空壳),仍走 delta 重建。 +func TestReconstructResponseOutputFromSSE_NonCompactionAddedStillUsesDeltas(t *testing.T) { + bodyText := strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"id":"msg_1","type":"message","content":[]}}`, + `data: {"type":"response.output_text.delta","delta":"hi"}`, + `data: {"type":"response.completed","response":{"id":"resp_1","output":[]}}`, + }, "\n") + + outputJSON, ok := reconstructResponseOutputFromSSE(bodyText) + require.True(t, ok) + items := gjson.ParseBytes(outputJSON).Array() + require.Len(t, items, 1) + require.Equal(t, "hi", items[0].Get("content.0.text").String()) +} + // 透传分支(OAuth passthrough)同样命中桥接。 func TestHandleNonStreamingResponsePassthrough_CompactClientStreamBridgesToSSE(t *testing.T) { svc := newCompactBridgeTestService() diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 6fdabeb90f..3bbe1393f9 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -902,6 +902,10 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. req.Header.Set("conversation_id", isolated) } } + } else if isOpenAIResponsesCompactPath(c) { + // compact 上游是 unary JSON 协议:API-key 账号也显式声明 Accept, + // 避免 OpenAI 兼容网关按 SSE 返回(#3777 期望行为 4)。 + req.Header.Set("accept", "application/json") } // Apply custom User-Agent if configured diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index 3a56ab5c61..f4a456d7fc 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -382,6 +382,11 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( if clientConversationID != "" { req.Header.Set("conversation_id", isolateOpenAISessionID(apiKeyID, clientConversationID)) } + } else if isOpenAIResponsesCompactPath(c) { + // 透传白名单会放行客户端的 Accept: text/event-stream;compact 上游是 + // unary JSON 协议,API-key 账号同样强制 Accept,避免上游按 SSE 返回 + // (#3777 期望行为 4)。 + req.Header.Set("accept", "application/json") } // 透传模式也支持账户自定义 User-Agent 与 ForceCodexCLI 兜底。 diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index cb4b780cdf..b4acfbfb35 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -1012,6 +1012,12 @@ func (s *OpenAIGatewayService) writeOpenAINonStreamingProtocolError(resp *http.R message = "Upstream returned an invalid non-streaming response" } setOpsUpstreamError(c, http.StatusBadGateway, message, "") + // body-signal compact 心跳可能已把响应头提交为 200,此时只能以 + // response.failed 终止事件回传错误,不能再写 JSON+状态码。 + if openAICompactClientWantsStream(c) && StopOpenAICompactSSEKeepaliveCommitted(c) { + writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, message) + return fmt.Errorf("non-streaming openai protocol error: %s", message) + } responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8") c.JSON(http.StatusBadGateway, gin.H{ @@ -1081,10 +1087,80 @@ func responsesStreamEventMayContributeToOutput(eventType string) bool { } } -// reconstructResponseOutputFromSSE scans raw SSE body text for delta events and -// returns a JSON-encoded output array reconstructed from accumulated deltas. -// Returns (nil, false) if no content was found in deltas. +// collectRawResponsesOutputItemsFromSSE 按到达顺序收集 SSE 流中 +// response.output_item.done 携带的原始 item。item 以 raw JSON 逐字节保留, +// 避免经窄结构体重建时丢弃 encrypted_content/summary/opaque 等 compact +// 专属或未来新增字段(#3777 问题 2)。若整条流没有任何 done 事件,退回 +// 收集 output_item.added 中的 compaction 类 item——compaction 结果没有 +// delta 事件,部分上游只在 added 事件中携带完整 item。 +func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) { + var items []json.RawMessage + seen := make(map[string]struct{}) + appendItem := func(item gjson.Result) { + if !item.Exists() || !item.IsObject() { + return + } + key := strings.TrimSpace(item.Get("id").String()) + if key == "" { + key = item.Raw + } + if _, dup := seen[key]; dup { + return + } + seen[key] = struct{}{} + items = append(items, json.RawMessage(item.Raw)) + } + forEachOpenAISSEDataPayload(bodyText, func(data []byte) { + if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.done" { + return + } + appendItem(gjson.GetBytes(data, "item")) + }) + if len(items) == 0 { + forEachOpenAISSEDataPayload(bodyText, func(data []byte) { + if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.added" { + return + } + item := gjson.GetBytes(data, "item") + if !isResponsesCompactionItemType(item.Get("type").String()) { + return + } + appendItem(item) + }) + } + if len(items) == 0 { + return nil, false + } + outputJSON, err := json.Marshal(items) + if err != nil { + return nil, false + } + return outputJSON, true +} + +// isResponsesCompactionItemType reports whether the item type is the Codex +// remote-compact result item ("compaction", upstream alias "compaction_summary"). +func isResponsesCompactionItemType(itemType string) bool { + switch strings.TrimSpace(itemType) { + case "compaction", "compaction_summary": + return true + default: + return false + } +} + +// reconstructResponseOutputFromSSE scans raw SSE body text and returns a +// JSON-encoded output array for a terminal event whose output is empty. +// Raw output_item.done items are preferred: per the Responses protocol they +// are the authoritative final form of each item. Delta accumulation only +// covers text/function_call/reasoning content and silently drops unknown +// item types such as compaction — Codex remote compact v2 then fails with +// "expected exactly one compaction output item, got 0" (#3887). +// Returns (nil, false) if nothing could be reconstructed. func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) { + if outputJSON, ok := collectRawResponsesOutputItemsFromSSE(bodyText); ok { + return outputJSON, true + } acc := apicompat.NewBufferedResponseAccumulator() imageOutputs := make([]json.RawMessage, 0, 1) seenImages := make(map[string]struct{}) From 062af81fb5964c78f64ee31195bed88ab3e858ad Mon Sep 17 00:00:00 2001 From: benjamin Date: Fri, 10 Jul 2026 10:00:47 +0800 Subject: [PATCH 07/27] fix(openai): preserve explicit GPT-5.6 cache write prices --- backend/internal/service/billing_service.go | 4 +- .../service/model_pricing_resolver.go | 2 + .../service/model_pricing_resolver_test.go | 46 +++++++++++++++++++ 3 files changed, 51 insertions(+), 1 deletion(-) diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index ae68c93ef6..3c8b52af97 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -97,6 +97,7 @@ type ModelPricing struct { OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD) CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD) CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD) + CacheCreationPriceExplicit bool // 是否由渠道/区间定价显式设定(为 true 时即使 == 0 也不回退) CacheReadPricePerToken float64 // 缓存读取每token价格 (USD) CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD) CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD) @@ -825,6 +826,7 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing if channelPricing.CacheWritePrice != nil { pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice pricing.CacheCreationPricePerTokenPriority = *channelPricing.CacheWritePrice + pricing.CacheCreationPriceExplicit = true pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice } @@ -1100,7 +1102,7 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * } needsLongContextPolicy := usesLegacyLongContextPricing && (pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0) - needsCacheCreationPolicy := isGPT56 && (pricing.CacheCreationPricePerToken <= 0 || + needsCacheCreationPolicy := isGPT56 && !pricing.CacheCreationPriceExplicit && (pricing.CacheCreationPricePerToken <= 0 || (pricing.InputPricePerTokenPriority > 0 && pricing.CacheCreationPricePerTokenPriority <= 0)) if !needsLongContextPolicy && !needsCacheCreationPolicy { return pricing diff --git a/backend/internal/service/model_pricing_resolver.go b/backend/internal/service/model_pricing_resolver.go index 7ddf6ea475..a9f603113b 100644 --- a/backend/internal/service/model_pricing_resolver.go +++ b/backend/internal/service/model_pricing_resolver.go @@ -184,6 +184,7 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin if chPricing.CacheWritePrice != nil { resolved.BasePricing.CacheCreationPricePerToken = *chPricing.CacheWritePrice resolved.BasePricing.CacheCreationPricePerTokenPriority = *chPricing.CacheWritePrice + resolved.BasePricing.CacheCreationPriceExplicit = true resolved.BasePricing.CacheCreation5mPrice = *chPricing.CacheWritePrice resolved.BasePricing.CacheCreation1hPrice = *chPricing.CacheWritePrice } @@ -253,6 +254,7 @@ func intervalToModelPricing(iv *PricingInterval, supportsCacheBreakdown bool, ch if iv.CacheWritePrice != nil { pricing.CacheCreationPricePerToken = *iv.CacheWritePrice pricing.CacheCreationPricePerTokenPriority = *iv.CacheWritePrice + pricing.CacheCreationPriceExplicit = true pricing.CacheCreation5mPrice = *iv.CacheWritePrice pricing.CacheCreation1hPrice = *iv.CacheWritePrice } diff --git a/backend/internal/service/model_pricing_resolver_test.go b/backend/internal/service/model_pricing_resolver_test.go index d4169e5f0e..71ba9904c2 100644 --- a/backend/internal/service/model_pricing_resolver_test.go +++ b/backend/internal/service/model_pricing_resolver_test.go @@ -114,6 +114,52 @@ func TestGetIntervalPricing_NoMatch_FallsBackToBase(t *testing.T) { require.Equal(t, basePricing, result) } +func TestGPT56ExplicitZeroCacheWritePriceIsPreserved(t *testing.T) { + bs := &BillingService{} + resolver := NewModelPricingResolver(nil, bs) + zero := 0.0 + + t.Run("flat channel price", func(t *testing.T) { + resolved := &ResolvedPricing{ + Mode: BillingModeToken, + BasePricing: &ModelPricing{ + InputPricePerToken: 5e-6, + OutputPricePerToken: 30e-6, + }, + } + resolver.applyTokenOverrides(&ChannelModelPricing{CacheWritePrice: &zero}, resolved) + + require.True(t, resolved.BasePricing.CacheCreationPriceExplicit) + cost, err := bs.CalculateCostUnified(CostInput{ + Model: "gpt-5.6-sol", + Tokens: UsageTokens{CacheCreationTokens: 100}, + RateMultiplier: 1, + Resolver: resolver, + Resolved: resolved, + }) + require.NoError(t, err) + require.Zero(t, cost.CacheCreationCost) + }) + + t.Run("interval price", func(t *testing.T) { + pricing := intervalToModelPricing(&PricingInterval{CacheWritePrice: &zero}, false, nil) + require.True(t, pricing.CacheCreationPriceExplicit) + + cost, err := bs.CalculateCostUnified(CostInput{ + Model: "gpt-5.6-sol", + Tokens: UsageTokens{CacheCreationTokens: 100}, + RateMultiplier: 1, + Resolver: resolver, + Resolved: &ResolvedPricing{ + Mode: BillingModeToken, + BasePricing: pricing, + }, + }) + require.NoError(t, err) + require.Zero(t, cost.CacheCreationCost) + }) +} + func TestGetRequestTierPrice(t *testing.T) { bs := newTestBillingServiceForResolver() r := NewModelPricingResolver(&ChannelService{}, bs) From 0a5f34a2ef713826d4f3571a3362f60a91fdc84f Mon Sep 17 00:00:00 2001 From: benjamin Date: Fri, 10 Jul 2026 10:11:21 +0800 Subject: [PATCH 08/27] fix(openai): recognize Windows websocket resets --- backend/internal/service/openai_ws_forwarder_ingress_test.go | 1 + backend/internal/service/openai_ws_forwarder_logutil.go | 1 + 2 files changed, 2 insertions(+) diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index ca7c36aaa7..1d18c46fca 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -33,6 +33,7 @@ func TestIsOpenAIWSClientDisconnectError(t *testing.T) { {name: "ws_policy_violation", err: coderws.CloseError{Code: coderws.StatusPolicyViolation}, want: false}, {name: "wrapped_eof_message", err: errors.New("failed to get reader: failed to read frame header: EOF"), want: true}, {name: "connection_reset_by_peer", err: errors.New("failed to read frame header: read tcp 127.0.0.1:1234->127.0.0.1:5678: read: connection reset by peer"), want: true}, + {name: "windows_connection_reset", err: errors.New("failed to get reader: failed to read frame header: read tcp 127.0.0.1:1234->127.0.0.1:5678: wsarecv: An existing connection was forcibly closed by the remote host."), want: true}, {name: "broken_pipe", err: errors.New("write tcp 127.0.0.1:1234->127.0.0.1:5678: write: broken pipe"), want: true}, } diff --git a/backend/internal/service/openai_ws_forwarder_logutil.go b/backend/internal/service/openai_ws_forwarder_logutil.go index 611938091d..04d1fb4a75 100644 --- a/backend/internal/service/openai_ws_forwarder_logutil.go +++ b/backend/internal/service/openai_ws_forwarder_logutil.go @@ -660,6 +660,7 @@ func isOpenAIWSClientDisconnectError(err error) bool { strings.Contains(message, "use of closed network connection") || strings.Contains(message, "connection reset by peer") || strings.Contains(message, "broken pipe") || + strings.Contains(message, "an existing connection was forcibly closed by the remote host") || strings.Contains(message, "an established connection was aborted") } From ae9a01d85206e6a3527d37a7ad3e9e0e9bb9afa5 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 10 Jul 2026 10:18:16 +0800 Subject: [PATCH 09/27] =?UTF-8?q?fix(compact):=20=E4=BA=8C=E8=BD=AE?= =?UTF-8?q?=E5=AE=A1=E8=AE=A1=E5=8A=A0=E5=9B=BA=E2=80=94=E2=80=94=E5=BF=83?= =?UTF-8?q?=E8=B7=B3=E5=AD=97=E8=8A=82=E4=B8=8D=E5=BE=97=E6=B1=A1=E6=9F=93?= =?UTF-8?q?=20failover=20=E5=88=A4=E5=AE=9A=E4=B8=8E=E5=B9=B6=E5=8F=91?= =?UTF-8?q?=E5=86=99=E5=AE=89=E5=85=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 对首轮修复的对抗式审计发现并修复以下问题: 1. failover 判定污染(真实回归风险):handler 以 "Forward 前后 c.Writer.Size() 是否变化" 判定响应是否已写出并据此放弃换号。心跳注释 字节会使该判定恒真,compact 请求一旦在上游等待期间发过心跳,上游 429/5xx 将不再 failover。新增 OpenAICompactKeepaliveAdjustedWrittenSize(扣除心跳字节、互斥锁下 一致读取、仅心跳字节归一化为未写哨兵 -1),快照、failover 比较与 openAIForwardErrorAlreadyCommunicated 三处判定统一改用该口径;无心跳 请求完全等价于原 c.Writer.Size()。 2. 并发写竞争:心跳 goroutine 与未被显式拦截的写回路径(Forward 内部 本地拒绝等)存在 ResponseWriter 数据竞争。StartOpenAICompactSSE- Keepalive 现将 c.Writer 替换为 openAICompactKeepaliveWriter:写侧 方法(Header/Write/WriteString/WriteHeader/WriteHeaderNow/Flush) 先在互斥锁下停拍,读侧(Status/Size/Written)仅加锁不停拍——任何 请求侧响应构造与心跳从构造上互斥,热路径状态读取不误杀心跳。 3. 语义拦截补齐:rejectIfCyberSessionBlocked(在用户槽位长等待之后 执行的直接 c.JSON)与 writeOpenAIFastPolicyBlockedResponse 在心跳 提交后降级为 response.failed 终止事件;未提交时先停拍再写 JSON, 状态码语义不变。失败事件 errType 参数化(permission_error 等)。 4. reconstruct 混合形态:done 事件存在但 compaction 只在 output_item.added 中时也要补入;done 已含 compaction 时跳过 added, 避免无 id 可去重时收集两份(Codex 要求恰好一个)。 5. 观测性:logOpenAIRemoteCompactOutcome 对心跳提交后的失败(wire 200) 以 GetOpsStreamError 纠正 outcome,不再误记 succeeded。 新增 5 个测试:failover 口径不变式、包装器停拍语义(-race)、fast policy 提交前后两态、混合 done/added 形态(含去重)。 Refs #3887 #3777 --- .../handler/openai_gateway_handler.go | 27 ++++- .../service/openai_compact_sse_keepalive.go | 108 +++++++++++++++++- .../openai_compact_sse_keepalive_test.go | 70 ++++++++++++ .../service/openai_compact_stream_bridge.go | 8 +- .../openai_compact_stream_bridge_test.go | 30 +++++ .../service/openai_gateway_request_body.go | 7 ++ .../openai_gateway_response_handling.go | 11 +- 7 files changed, 250 insertions(+), 11 deletions(-) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 714e7906f3..83e644d857 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -407,7 +407,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { // Forward request service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) forwardStart := time.Now() - writerSizeBeforeForward := c.Writer.Size() + // 用扣除 compact 心跳字节的口径快照:心跳注释不构成语义响应, + // 不能因心跳字节变化而放弃 failover 换号(#3887)。 + writerSizeBeforeForward := service.OpenAICompactKeepaliveAdjustedWrittenSize(c) result, err := func() (*service.OpenAIForwardResult, error) { defer func() { if accountReleaseFunc != nil { @@ -441,7 +443,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { } else { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { - if c.Writer.Size() != writerSizeBeforeForward { + if service.OpenAICompactKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeForward { h.handleFailoverExhausted(c, failoverErr, true) return } @@ -642,6 +644,13 @@ func (h *OpenAIGatewayHandler) logOpenAIRemoteCompactOutcome(c *gin.Context, sta if status >= 200 && status < 300 { outcome = "succeeded" } + // compact 心跳提交后失败的 wire 状态码固化为 200,真实结局以流内错误 + // 标记为准(response.failed 降级路径会 MarkOpsStreamError)。 + if outcome == "succeeded" && c != nil { + if _, hasStreamErr := service.GetOpsStreamError(c); hasStreamErr { + outcome = "failed" + } + } latencyMs := time.Since(startedAt).Milliseconds() if latencyMs < 0 { latencyMs = 0 @@ -2024,7 +2033,9 @@ func openAIForwardErrorAlreadyCommunicated(c *gin.Context, writerSizeBeforeForwa if err == nil || c == nil || c.Writer == nil { return false } - if c.Writer.Size() == writerSizeBeforeForward { + // 与快照同口径:排除 compact 心跳字节,避免"仅心跳写出"被误判为 + // 响应已写出(#3887)。 + if service.OpenAICompactKeepaliveAdjustedWrittenSize(c) == writerSizeBeforeForward { return false } @@ -2337,6 +2348,16 @@ func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKe if !h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), key) { return false } + // body-signal compact 心跳可能已把响应头提交为 200(cyber 检查在用户槽位 + // 长等待之后执行):以 response.failed 终止事件回传;未提交时停拍后照常 + // 写 JSON(#3887)。 + if service.StopOpenAICompactSSEKeepaliveCommitted(c) { + service.MarkOpsStreamError(c, "permission_error", cyberSessionBlockedClientMsg, http.StatusForbidden) + if writeResponsesFailedSSE(c, "permission_error", cyberSessionBlockedClientMsg) { + h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, model, key) + return true + } + } switch format { case cyberBlockFormatAnthropic: c.JSON(http.StatusForbidden, gin.H{"type": "error", "error": gin.H{ diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go index 0695d16fb3..70ef3fc01a 100644 --- a/backend/internal/service/openai_compact_sse_keepalive.go +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -26,11 +26,19 @@ type openAICompactSSEKeepalive struct { writer gin.ResponseWriter started bool stopped bool - stop chan struct{} + // bytes 是心跳已写出的注释字节数。心跳不构成语义响应,handler 的 + // "Forward 期间是否已写响应"判定(failover 放弃换号的依据)必须扣除 + // 这部分字节,见 OpenAICompactKeepaliveAdjustedWrittenSize。 + bytes int + stop chan struct{} } // StartOpenAICompactSSEKeepalive 为已标记 body-signal 客户端流式的 compact // 请求启动下游心跳,返回幂等的停止函数。interval<=0 或请求未标记时为 no-op。 +// +// 同时把 c.Writer 替换为 openAICompactKeepaliveWriter:请求 goroutine 的任何 +// 响应构造都会先在心跳互斥锁下停拍,未被显式拦截的写回路径(如 Forward +// 内部的本地拒绝)也不会与心跳 goroutine 产生数据竞争或字节交错。 func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func() { if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) { return func() {} @@ -40,6 +48,7 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func stop: make(chan struct{}), } c.Set(openAICompactSSEKeepaliveKey, k) + c.Writer = &openAICompactKeepaliveWriter{ResponseWriter: c.Writer, k: k} var reqDone <-chan struct{} if c.Request != nil { @@ -82,7 +91,9 @@ func (k *openAICompactSSEKeepalive) beat() bool { k.writer.WriteHeader(http.StatusOK) k.started = true } - if _, err := k.writer.Write([]byte(": keepalive\n\n")); err != nil { + n, err := k.writer.Write([]byte(": keepalive\n\n")) + k.bytes += n + if err != nil { k.stopped = true return false } @@ -127,3 +138,96 @@ func StopOpenAICompactSSEKeepaliveCommitted(c *gin.Context) bool { k.mu.Unlock() return committed } + +// OpenAICompactKeepaliveAdjustedWrittenSize 返回排除 compact 心跳注释字节后 +// 的响应已写字节数;无心跳的请求等价于 c.Writer.Size()。心跳字节不构成语义 +// 响应——handler 以"Forward 前后 Size 是否变化"判定是否已向客户端写出响应 +// (变化则放弃 failover 换号),该判定不得被心跳污染,否则 compact 请求 +// 一旦在上游等待期间发过心跳,上游 429/5xx 就不再换号(#3887 加固审计)。 +// 仅心跳字节时归一化为 -1(gin 的"未写出"哨兵值),与提交前的快照可比。 +func OpenAICompactKeepaliveAdjustedWrittenSize(c *gin.Context) int { + if c == nil || c.Writer == nil { + return -1 + } + value, ok := c.Get(openAICompactSSEKeepaliveKey) + if !ok { + return c.Writer.Size() + } + k, ok := value.(*openAICompactSSEKeepalive) + if !ok || k == nil { + return c.Writer.Size() + } + k.mu.Lock() + defer k.mu.Unlock() + size := k.writer.Size() + if size < 0 { + return size + } + if real := size - k.bytes; real > 0 { + return real + } + return -1 +} + +// openAICompactKeepaliveWriter 包装 gin.ResponseWriter:写侧方法先停拍心跳 +// (互斥锁下建立 happens-before),读侧方法仅加锁不停拍——热路径的状态读取 +// (如 Forward 前的 Size 快照)不能误杀心跳。心跳 goroutine 直接写内层 +// writer(k.writer),不经过本包装器,不会递归。 +type openAICompactKeepaliveWriter struct { + gin.ResponseWriter + k *openAICompactSSEKeepalive +} + +// suspend 停拍心跳;幂等。任何响应构造(含 Header 访问——写响应必先操作 +// 响应头)都视为请求侧接管 ResponseWriter。 +func (w *openAICompactKeepaliveWriter) suspend() { + w.k.Stop() +} + +func (w *openAICompactKeepaliveWriter) Header() http.Header { + w.suspend() + return w.ResponseWriter.Header() +} + +func (w *openAICompactKeepaliveWriter) Write(data []byte) (int, error) { + w.suspend() + return w.ResponseWriter.Write(data) +} + +func (w *openAICompactKeepaliveWriter) WriteString(s string) (int, error) { + w.suspend() + return w.ResponseWriter.WriteString(s) +} + +func (w *openAICompactKeepaliveWriter) WriteHeader(code int) { + w.suspend() + w.ResponseWriter.WriteHeader(code) +} + +func (w *openAICompactKeepaliveWriter) WriteHeaderNow() { + w.suspend() + w.ResponseWriter.WriteHeaderNow() +} + +func (w *openAICompactKeepaliveWriter) Flush() { + w.suspend() + w.ResponseWriter.Flush() +} + +func (w *openAICompactKeepaliveWriter) Status() int { + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Status() +} + +func (w *openAICompactKeepaliveWriter) Size() int { + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Size() +} + +func (w *openAICompactKeepaliveWriter) Written() bool { + w.k.mu.Lock() + defer w.k.mu.Unlock() + return w.ResponseWriter.Written() +} diff --git a/backend/internal/service/openai_compact_sse_keepalive_test.go b/backend/internal/service/openai_compact_sse_keepalive_test.go index 4abff79fc1..3b217a0718 100644 --- a/backend/internal/service/openai_compact_sse_keepalive_test.go +++ b/backend/internal/service/openai_compact_sse_keepalive_test.go @@ -120,3 +120,73 @@ func TestWriteOpenAICompactSSEBridge_BeforeKeepaliveCommitFailureKeepsJSONPath(t require.False(t, writeOpenAICompactSSEBridge(c, http.StatusBadGateway, []byte(`{"error":{"message":"fast fail"}}`))) require.Zero(t, rec.Body.Len()) } + +// 未被显式拦截的写回路径(直接操作 c.Writer)也必须与心跳互斥:包装器在 +// 请求侧任何响应构造时停拍。-race 下验证无数据竞争,且停拍后不再有心跳 +// 字节写出。 +func TestOpenAICompactKeepaliveWriter_RequestSideWriteSuspendsBeats(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + waitForKeepaliveBeats() + + // 模拟未拦截路径的直接写回(如 Forward 内部本地拒绝的 c.JSON)。 + _, err := c.Writer.Write([]byte(`{"error":"local reject"}`)) + require.NoError(t, err) + + lenAfterWrite := rec.Body.Len() + waitForKeepaliveBeats() + require.Equal(t, lenAfterWrite, rec.Body.Len(), "请求侧写回后心跳必须停止") + require.Contains(t, rec.Body.String(), ": keepalive\n\n") + require.Contains(t, rec.Body.String(), `{"error":"local reject"}`) +} + +// fast policy block 在心跳提交后必须降级为 response.failed 终止事件。 +func TestWriteOpenAIFastPolicyBlockedResponse_AfterKeepaliveCommit(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + waitForKeepaliveBeats() + + writeOpenAIFastPolicyBlockedResponse(c, &OpenAIFastBlockedError{Message: "tier blocked"}) + + require.Equal(t, http.StatusOK, rec.Code) + events := parseCompactBridgeSSE(t, stripKeepaliveComments(rec.Body.String())) + require.Len(t, events, 1) + require.Equal(t, "response.failed", events[0][0]) + require.Equal(t, "permission_error", gjson.Get(events[0][1], "response.error.code").String()) + require.Contains(t, gjson.Get(events[0][1], "response.error.message").String(), "tier blocked") +} + +// failover"是否已写响应"判定的口径:心跳字节必须被排除,否则 compact 在 +// 上游等待期间发过心跳后,可换号的 failover 会被误判放弃;真实响应字节 +// 写出后口径必须变化。 +func TestOpenAICompactKeepaliveAdjustedWrittenSize_ExcludesHeartbeatBytes(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + // 无心跳的请求:等价于 c.Writer.Size()。 + require.Equal(t, c.Writer.Size(), OpenAICompactKeepaliveAdjustedWrittenSize(c)) + + stop := StartOpenAICompactSSEKeepalive(c, keepaliveTestInterval) + defer stop() + before := OpenAICompactKeepaliveAdjustedWrittenSize(c) + waitForKeepaliveBeats() + require.Equal(t, before, OpenAICompactKeepaliveAdjustedWrittenSize(c), "仅心跳字节不得改变判定口径") + + // 真实响应字节写出(经包装器,先停拍再写)后口径必须变化。 + _, err := c.Writer.Write([]byte("real-bytes")) + require.NoError(t, err) + require.Equal(t, len("real-bytes"), OpenAICompactKeepaliveAdjustedWrittenSize(c)) + require.Contains(t, rec.Body.String(), ": keepalive\n\n") +} + +// fast policy block 在心跳未提交时保持 403 JSON 原语义。 +func TestWriteOpenAIFastPolicyBlockedResponse_BeforeKeepaliveCommit(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, time.Hour) + defer stop() + + writeOpenAIFastPolicyBlockedResponse(c, &OpenAIFastBlockedError{Message: "tier blocked"}) + + require.Equal(t, http.StatusForbidden, rec.Code) + require.Equal(t, "permission_error", gjson.Get(rec.Body.String(), "error.type").String()) +} diff --git a/backend/internal/service/openai_compact_stream_bridge.go b/backend/internal/service/openai_compact_stream_bridge.go index 98240f4229..bcac8304c2 100644 --- a/backend/internal/service/openai_compact_stream_bridge.go +++ b/backend/internal/service/openai_compact_stream_bridge.go @@ -100,18 +100,18 @@ func writeOpenAICompactSSEFailure(c *gin.Context, statusCode int, errorBody []by if message == "" { message = "Upstream compact request failed with HTTP " + strconv.Itoa(statusCode) } - writeOpenAICompactSSEFailureMessage(c, statusCode, message) + writeOpenAICompactSSEFailureMessage(c, statusCode, "upstream_error", message) } // writeOpenAICompactSSEFailureMessage 写出 response.failed 终止事件。Codex 对 // 流式 Responses 请求把 response.failed 作为合法终止事件处理(普通 error 帧 // 不被识别,会退化为 "stream closed before response.completed" 盲重连)。 // 同时标记流内错误,保证挂在 200 流上的失败仍进入 ops 错误看板。 -func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, message string) { +func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, errType, message string) { if c == nil { return } - MarkOpsStreamError(c, "upstream_error", message, statusCode) + MarkOpsStreamError(c, errType, message, statusCode) payload, err := json.Marshal(map[string]any{ "type": "response.failed", "response": map[string]any{ @@ -120,7 +120,7 @@ func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, message "status": "failed", "output": []any{}, "error": map[string]any{ - "code": "upstream_error", + "code": errType, "message": message, }, }, diff --git a/backend/internal/service/openai_compact_stream_bridge_test.go b/backend/internal/service/openai_compact_stream_bridge_test.go index 1f15084fd9..32b86380a9 100644 --- a/backend/internal/service/openai_compact_stream_bridge_test.go +++ b/backend/internal/service/openai_compact_stream_bridge_test.go @@ -392,6 +392,36 @@ func TestReconstructResponseOutputFromSSE_CompactionAddedFallback(t *testing.T) require.Equal(t, "added-only", items[0].Get("encrypted_content").String()) } +// 混合形态:其他 item 有 done、compaction 只在 added 中——compaction 必须 +// 被补入;done 已含 compaction 时 added 不得重复计入。 +func TestReconstructResponseOutputFromSSE_MixedDoneAndCompactionAdded(t *testing.T) { + bodyText := strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"id":"cmp_mixed","type":"compaction","encrypted_content":"mixed"}}`, + `data: {"type":"response.output_item.done","output_index":1,"item":{"id":"msg_1","type":"message","content":[{"type":"output_text","text":"hi"}]}}`, + `data: {"type":"response.completed","response":{"id":"resp_1","output":[]}}`, + }, "\n") + + outputJSON, ok := reconstructResponseOutputFromSSE(bodyText) + require.True(t, ok) + items := gjson.ParseBytes(outputJSON).Array() + require.Len(t, items, 2) + require.Equal(t, "msg_1", items[0].Get("id").String()) + require.Equal(t, "cmp_mixed", items[1].Get("id").String()) + + // done 已含 compaction:added 中的同一 item(无 id 可去重的最坏情况用 + // 不同 raw 表达)不得再收集,Codex 要求恰好一个 compaction item。 + bodyText = strings.Join([]string{ + `data: {"type":"response.output_item.added","output_index":0,"item":{"type":"compaction","status":"in_progress"}}`, + `data: {"type":"response.output_item.done","output_index":0,"item":{"type":"compaction","status":"completed","encrypted_content":"final"}}`, + `data: {"type":"response.completed","response":{"id":"resp_1","output":[]}}`, + }, "\n") + outputJSON, ok = reconstructResponseOutputFromSSE(bodyText) + require.True(t, ok) + items = gjson.ParseBytes(outputJSON).Array() + require.Len(t, items, 1) + require.Equal(t, "final", items[0].Get("encrypted_content").String()) +} + // 非 compaction 的 output_item.added 不参与回退收集(added 阶段的 message // 通常是空壳),仍走 delta 重建。 func TestReconstructResponseOutputFromSSE_NonCompactionAddedStillUsesDeltas(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index b48b7510eb..4824ce11e2 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -772,6 +772,13 @@ func writeOpenAIFastPolicyBlockedResponse(c *gin.Context, err *OpenAIFastBlocked return } MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied) + // body-signal compact 心跳可能已把响应头提交为 200(长排队后才进入 + // Forward),此时以 response.failed 终止事件回传;未提交时先停拍再写 + // JSON,保持原状态码语义(#3887)。 + if StopOpenAICompactSSEKeepaliveCommitted(c) { + writeOpenAICompactSSEFailureMessage(c, http.StatusForbidden, "permission_error", err.Message) + return + } c.JSON(http.StatusForbidden, gin.H{ "error": gin.H{ "type": "permission_error", diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index b4acfbfb35..9c9ea0c1a3 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -1015,7 +1015,7 @@ func (s *OpenAIGatewayService) writeOpenAINonStreamingProtocolError(resp *http.R // body-signal compact 心跳可能已把响应头提交为 200,此时只能以 // response.failed 终止事件回传错误,不能再写 JSON+状态码。 if openAICompactClientWantsStream(c) && StopOpenAICompactSSEKeepaliveCommitted(c) { - writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, message) + writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, "upstream_error", message) return fmt.Errorf("non-streaming openai protocol error: %s", message) } responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) @@ -1096,6 +1096,7 @@ func responsesStreamEventMayContributeToOutput(eventType string) bool { func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) { var items []json.RawMessage seen := make(map[string]struct{}) + hasCompactionItem := false appendItem := func(item gjson.Result) { if !item.Exists() || !item.IsObject() { return @@ -1108,6 +1109,9 @@ func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) { return } seen[key] = struct{}{} + if isResponsesCompactionItemType(item.Get("type").String()) { + hasCompactionItem = true + } items = append(items, json.RawMessage(item.Raw)) } forEachOpenAISSEDataPayload(bodyText, func(data []byte) { @@ -1116,7 +1120,10 @@ func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) { } appendItem(gjson.GetBytes(data, "item")) }) - if len(items) == 0 { + // done 事件未携带 compaction item 时再看 added:覆盖"其他 item 有 done、 + // compaction 只在 added 中"的混合形态;done 已含 compaction 时跳过, + // 避免同一 item 在无 id 可去重时被收集两份(Codex 要求恰好一个)。 + if !hasCompactionItem { forEachOpenAISSEDataPayload(bodyText, func(data []byte) { if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.added" { return From 5c15d32ff7c4cff7c243f20c02be455ca0fefac0 Mon Sep 17 00:00:00 2001 From: benjamin Date: Fri, 10 Jul 2026 10:18:50 +0800 Subject: [PATCH 10/27] test(repository): stabilize expired lock reconciliation --- .../repository/user_msg_queue_cache_integration_test.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/backend/internal/repository/user_msg_queue_cache_integration_test.go b/backend/internal/repository/user_msg_queue_cache_integration_test.go index e683b44aa3..f57206efe8 100644 --- a/backend/internal/repository/user_msg_queue_cache_integration_test.go +++ b/backend/internal/repository/user_msg_queue_cache_integration_test.go @@ -53,11 +53,11 @@ func (s *UserMsgQueueCacheSuite) TestReconcileExpiredLockCandidatesRemovesNatura require.NoError(s.T(), err) require.True(s.T(), acquired) - score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "702").Result() + _, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "702").Result() require.NoError(s.T(), err) require.Eventually(s.T(), func() bool { - nowMs, err := s.cache.GetCurrentTimeMs(s.ctx) - return err == nil && nowMs >= int64(score) + _, err := s.rdb.Get(s.ctx, umqLockKey(accountID)).Result() + return errors.Is(err, redis.Nil) }, time.Second, 10*time.Millisecond) cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000) From 000f6dc655849f03b95d1244c9c5decf4c2bb406 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 10 Jul 2026 10:25:09 +0800 Subject: [PATCH 11/27] =?UTF-8?q?fix(compact):=20=E7=BB=88=E6=80=81=20outp?= =?UTF-8?q?ut=20=E9=9D=9E=E7=A9=BA=E4=BD=86=E7=BC=BA=20compaction=20?= =?UTF-8?q?=E6=97=B6=E4=BB=8E=E4=BA=8B=E4=BB=B6=E6=B5=81=E8=A1=A5=E5=85=A5?= =?UTF-8?q?=EF=BC=88146=20=E7=AD=89=E4=BB=B7=E6=80=A7=E6=94=B6=E5=B0=BE?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 146 纯流式透传下 Codex 只从 raw output_item.done 收集 item,无论终态 output 形态如何都能拿到 compaction;SSE→JSON 提取链路此前仅在终态 output 为空时修补。补齐最后一处不等价:compact 请求终态 output 非空但缺 compaction、事件流中存在(done 优先 added 兜底)时以 raw JSON 追加。 门控:仅 compact 路径 + 仅缺失时触发,已含 compaction 或非 compact 请求逐字节不变。 Refs #3887 #3777 --- .../openai_compact_stream_bridge_test.go | 60 +++++++++++++++++ .../service/openai_gateway_passthrough.go | 1 + .../openai_gateway_response_handling.go | 66 +++++++++++++++++++ 3 files changed, 127 insertions(+) diff --git a/backend/internal/service/openai_compact_stream_bridge_test.go b/backend/internal/service/openai_compact_stream_bridge_test.go index 32b86380a9..bf3416192e 100644 --- a/backend/internal/service/openai_compact_stream_bridge_test.go +++ b/backend/internal/service/openai_compact_stream_bridge_test.go @@ -422,6 +422,66 @@ func TestReconstructResponseOutputFromSSE_MixedDoneAndCompactionAdded(t *testing require.Equal(t, "final", items[0].Get("encrypted_content").String()) } +// 上游不一致形态:终态 output 非空(含 message)但 compaction 只在 raw +// output_item.done 中。146 纯流式透传下 Codex 直接读事件流能拿到 compaction, +// SSE→JSON 提取必须补入等价结果。 +func TestHandleSSEToJSON_CompactSupplementsMissingCompactionIntoNonEmptyOutput(t *testing.T) { + svc := newCompactBridgeTestService() + c, rec := newCompactBridgeTestContext(t, true) + upstreamSSE := strings.Join([]string{ + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"cmp_sup","type":"compaction","encrypted_content":"supplement"}}`, + ``, + `data: {"type":"response.completed","response":{"id":"resp_sup","object":"response","status":"completed","output":[{"id":"msg_sup","type":"message","role":"assistant","content":[{"type":"output_text","text":"note"}]}],"usage":{"input_tokens":2,"output_tokens":1,"total_tokens":3}}}`, + ``, + }, "\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + } + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, &Account{ID: 1, Type: AccountTypeOAuth}, "gpt-5.5", "gpt-5.5") + require.NoError(t, err) + require.NotNil(t, result) + + events := parseCompactBridgeSSE(t, rec.Body.String()) + require.Len(t, events, 3) + itemTypes := []string{ + gjson.Get(events[0][1], "item.type").String(), + gjson.Get(events[1][1], "item.type").String(), + } + require.Contains(t, itemTypes, "compaction") + require.Contains(t, itemTypes, "message") + require.Equal(t, "response.completed", events[2][0]) + require.Len(t, gjson.Get(events[2][1], "response.output").Array(), 2) +} + +// 补全逻辑的门控:非 compact 请求原样返回;终态已含 compaction 不重复补入。 +func TestSupplementCompactionItemFromSSE_Gating(t *testing.T) { + bodyText := `data: {"type":"response.output_item.done","item":{"id":"cmp_g","type":"compaction","encrypted_content":"g"}}` + "\n" + + // 非 compact 路径:不补入。 + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + finalResponse := []byte(`{"id":"r1","output":[{"type":"message"}]}`) + require.Equal(t, string(finalResponse), string(supplementCompactionItemFromSSE(c, finalResponse, bodyText))) + + // compact 路径 + 终态已含 compaction:不重复补入。 + c2, _ := newCompactBridgeTestContext(t, false) + already := []byte(`{"id":"r2","output":[{"type":"compaction","encrypted_content":"x"}]}`) + require.Equal(t, string(already), string(supplementCompactionItemFromSSE(c2, already, bodyText))) + + // compact 路径 + 终态非空缺 compaction:补入到末尾。 + missing := []byte(`{"id":"r3","output":[{"type":"message"}]}`) + patched := supplementCompactionItemFromSSE(c2, missing, bodyText) + items := gjson.GetBytes(patched, "output").Array() + require.Len(t, items, 2) + require.Equal(t, "compaction", items[1].Get("type").String()) + require.Equal(t, "g", items[1].Get("encrypted_content").String()) +} + // 非 compaction 的 output_item.added 不参与回退收集(added 阶段的 message // 通常是空壳),仍走 delta 重建。 func TestReconstructResponseOutputFromSSE_NonCompactionAddedStillUsesDeltas(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index f4a456d7fc..12eb248fcf 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -1127,6 +1127,7 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c } } } + finalResponse = supplementCompactionItemFromSSE(c, finalResponse, bodyText) body = finalResponse if originalModel != "" && mappedModel != "" && originalModel != mappedModel { body = s.replaceModelInResponseBody(body, mappedModel, originalModel) diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 9c9ea0c1a3..45c948f56f 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -878,6 +878,7 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte } } } + finalResponse = supplementCompactionItemFromSSE(c, finalResponse, bodyText) body = finalResponse if originalModel != mappedModel { body = s.replaceModelInResponseBody(body, mappedModel, originalModel) @@ -1156,6 +1157,71 @@ func isResponsesCompactionItemType(itemType string) bool { } } +// supplementCompactionItemFromSSE 保证 compact 请求的终态 output 携带 +// compaction item:终态 output 非空但缺失 compaction、而原始事件流的 +// output_item.done(或 added)中存在时(上游不一致形态),以 raw JSON 补入。 +// Codex remote compact v2 只从 output_item.done 收集 item 且要求恰好一个 +// compaction item——纯流式透传(v0.1.146)下客户端直接读事件流天然拿得到, +// SSE→JSON 提取链路必须给出等价结果。非 compact 请求原样返回。 +func supplementCompactionItemFromSSE(c *gin.Context, finalResponse []byte, bodyText string) []byte { + if !isOpenAIResponsesCompactPath(c) { + return finalResponse + } + if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 { + // 空 output 由 reconstructResponseOutputFromSSE 整体修补,不在此处理。 + return finalResponse + } + if responsesOutputHasCompactionItem(finalResponse) { + return finalResponse + } + item, found := findRawCompactionItemFromSSE(bodyText) + if !found { + return finalResponse + } + patched, err := sjson.SetRawBytes(finalResponse, "output.-1", item) + if err != nil { + return finalResponse + } + return patched +} + +// responsesOutputHasCompactionItem reports whether the response JSON already +// carries a compaction item in its output array. +func responsesOutputHasCompactionItem(response []byte) bool { + for _, item := range gjson.GetBytes(response, "output").Array() { + if isResponsesCompactionItemType(item.Get("type").String()) { + return true + } + } + return false +} + +// findRawCompactionItemFromSSE 从原始 SSE 事件流中提取第一个 compaction 类 +// item 的 raw JSON:output_item.done 优先,output_item.added 兜底。 +func findRawCompactionItemFromSSE(bodyText string) (json.RawMessage, bool) { + var found json.RawMessage + pick := func(eventType string) { + forEachOpenAISSEDataPayload(bodyText, func(data []byte) { + if found != nil { + return + } + if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != eventType { + return + } + item := gjson.GetBytes(data, "item") + if !item.IsObject() || !isResponsesCompactionItemType(item.Get("type").String()) { + return + } + found = json.RawMessage(item.Raw) + }) + } + pick("response.output_item.done") + if found == nil { + pick("response.output_item.added") + } + return found, found != nil +} + // reconstructResponseOutputFromSSE scans raw SSE body text and returns a // JSON-encoded output array for a terminal event whose output is empty. // Raw output_item.done items are preferred: per the Responses protocol they From 80b3d4c1f8a0b708e4f40757a66ecc428b714c4d Mon Sep 17 00:00:00 2001 From: psyche <123316733+githubbzxs@users.noreply.github.com> Date: Fri, 10 Jul 2026 10:39:16 +0800 Subject: [PATCH 12/27] =?UTF-8?q?fix:=20=E5=85=BC=E5=AE=B9=20GPT-5.6=20max?= =?UTF-8?q?=20=E6=8E=A8=E7=90=86=E5=BC=BA=E5=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../openai_gateway_chat_completions_raw.go | 4 +- ...penai_gateway_chat_completions_raw_test.go | 33 +++ .../service/openai_gateway_forward.go | 15 +- .../openai_gateway_messages_chat_fallback.go | 2 +- .../service/openai_gateway_request_body.go | 55 +++- .../openai_gateway_responses_chat_fallback.go | 2 +- .../internal/service/openai_gpt56_max_test.go | 266 ++++++++++++++++++ .../service/openai_ws_forwarder_ingress.go | 2 +- .../service/openai_ws_forwarder_v2.go | 2 +- .../internal/service/openai_ws_http_bridge.go | 2 +- 10 files changed, 367 insertions(+), 16 deletions(-) create mode 100644 backend/internal/service/openai_gpt56_max_test.go diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 9b31b803d2..2f46a1c76c 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -72,13 +72,13 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } clientStream := gjson.GetBytes(body, "stream").Bool() - // 1b. Extract reasoning effort and service tier from the raw body before any transformation. - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel) + // 1b. Extract service tier from the raw body before any transformation. serviceTier := extractOpenAIServiceTierFromBody(body) // 2. Resolve model mapping (same as ForwardAsChatCompletions) billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) // 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) diff --git a/backend/internal/service/openai_gateway_chat_completions_raw_test.go b/backend/internal/service/openai_gateway_chat_completions_raw_test.go index 5a6ee3f9b9..5f4b5bcd2d 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw_test.go @@ -122,6 +122,39 @@ func TestForwardAsRawChatCompletions_ForcesStreamUsageUpstreamAndPassesUsageDown require.Contains(t, rec.Body.String(), "data: [DONE]") } +func TestForwardAsRawChatCompletions_PreservesMappedGPT56MaxEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"sol","messages":[{"role":"user","content":"hello"}],"reasoning_effort":"max","stream":false}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"chatcmpl_max","object":"chat.completion","model":"gpt-5.6-sol","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`, + )), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + } + account := rawChatCompletionsTestAccount() + account.Credentials["model_mapping"] = map[string]any{"sol": "gpt-5.6-sol"} + + result, err := svc.forwardAsRawChatCompletions(context.Background(), c, account, body, "") + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning_effort").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} + func TestForwardAsRawChatCompletions_PreservesDeepSeekReasoningContentNonStreaming(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 6fdabeb90f..2b19578bd8 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -35,6 +35,14 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco return nil, errors.New("codex_cli_only restriction: only codex official clients are allowed") } + normalizedBody, normalized, err := normalizeOpenAICodexCompactReasoningEffortForAccount(c, account, body) + if err != nil { + return nil, err + } + if normalized { + body = normalizedBody + } + originalBody := body requestView := newOpenAIRequestView(body) reqModel, reqStream, promptCacheKey := requestView.Model, requestView.Stream, requestView.PromptCacheKey @@ -88,9 +96,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco passthroughEnabled := account.IsOpenAIPassthroughEnabled() if passthroughEnabled { // 透传分支只需要轻量提取字段,避免热路径全量 Unmarshal。 - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, reqModel) + mappedModel := account.GetMappedModel(reqModel) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, mappedModel) // 国产模型默认 effort 补充:也要用 mappedModel 判定是否是 passback-required 上游。 - reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, account.GetMappedModel(reqModel)) + reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, mappedModel) return s.forwardOpenAIPassthrough(ctx, c, account, originalBody, reqModel, reasoningEffort, reqStream, startTime) } @@ -746,7 +755,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } defer func() { _ = resp.Body.Close() }() - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) // 国产模型默认 effort 补充:此处 reqModel 已被 mapping 重写为 billingModel(见 // line 2510-2515 的 GetMappedModel + reqModel 赋值),可直接作为 mappedModel。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, reqModel) diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback.go b/backend/internal/service/openai_gateway_messages_chat_fallback.go index 4596177468..97eecf014b 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback.go @@ -72,7 +72,7 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions( chatReq.StreamOptions = &apicompat.ChatStreamOptions{IncludeUsage: true} } - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) serviceTier := extractOpenAIServiceTierFromBody(body) diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index b48b7510eb..cead23d262 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -199,6 +199,32 @@ func normalizeOpenAICompactRequestBody(body []byte) ([]byte, bool, error) { return normalized, true, nil } +func normalizeOpenAICodexCompactReasoningEffortForAccount(c *gin.Context, account *Account, body []byte) ([]byte, bool, error) { + if account == nil || !account.IsOpenAIOAuth() || !isOpenAIResponsesCompactPath(c) { + return body, false, nil + } + + requestedModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + effectiveModel := account.GetMappedModel(requestedModel) + return normalizeOpenAICodexCompactReasoningEffort(body, effectiveModel) +} + +func normalizeOpenAICodexCompactReasoningEffort(body []byte, effectiveModel string) ([]byte, bool, error) { + if !isOpenAIGPT56Model(effectiveModel) || + !strings.EqualFold(strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()), "max") { + return body, false, nil + } + + // Codex Ultra 在客户端编排层会下发 max;ChatGPT compact 端点目前只接受到 + // xhigh。这里只降级 OpenAI OAuth 的 GPT-5.6 compact 子请求,普通 Responses、 + // API Key 请求和其他平台的 OAuth 请求保留 max。 + normalized, err := sjson.SetBytes(body, "reasoning.effort", "xhigh") + if err != nil { + return body, false, fmt.Errorf("normalize codex compact reasoning effort: %w", err) + } + return normalized, true, nil +} + func resolveOpenAICompactSessionID(c *gin.Context) string { if c != nil { if sessionID := strings.TrimSpace(c.GetHeader("session_id")); sessionID != "" { @@ -259,7 +285,7 @@ func (s *OpenAIGatewayService) replaceModelInResponseBody(body []byte, fromModel return body } -func getOpenAIReasoningEffortFromReqBody(reqBody map[string]any) (value string, present bool) { +func getOpenAIReasoningEffortFromReqBody(reqBody map[string]any, requestedModel string) (value string, present bool) { if reqBody == nil { return "", false } @@ -267,13 +293,13 @@ func getOpenAIReasoningEffortFromReqBody(reqBody map[string]any) (value string, // Primary: reasoning.effort if reasoning, ok := reqBody["reasoning"].(map[string]any); ok { if effort, ok := reasoning["effort"].(string); ok { - return normalizeOpenAIReasoningEffort(effort), true + return normalizeOpenAIReasoningEffortForModel(effort, requestedModel), true } } // Fallback: some clients may use a flat field. if effort, ok := reqBody["reasoning_effort"].(string); ok { - return normalizeOpenAIReasoningEffort(effort), true + return normalizeOpenAIReasoningEffortForModel(effort, requestedModel), true } return "", false @@ -302,7 +328,7 @@ func deriveOpenAIReasoningEffortFromModel(model string) string { return "" } - return normalizeOpenAIReasoningEffort(parts[len(parts)-1]) + return normalizeOpenAIReasoningEffortForModel(parts[len(parts)-1], modelID) } type openAIRequestView struct { @@ -551,7 +577,7 @@ func extractOpenAIReasoningEffortFromBody(body []byte, requestedModel string) *s reasoningEffort = strings.TrimSpace(gjson.GetBytes(body, "reasoning_effort").String()) } if reasoningEffort != "" { - normalized := normalizeOpenAIReasoningEffort(reasoningEffort) + normalized := normalizeOpenAIReasoningEffortForModel(reasoningEffort, requestedModel) if normalized == "" { return nil } @@ -1127,7 +1153,7 @@ func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error } func extractOpenAIReasoningEffort(reqBody map[string]any, requestedModel string) *string { - if value, present := getOpenAIReasoningEffortFromReqBody(reqBody); present { + if value, present := getOpenAIReasoningEffortFromReqBody(reqBody, requestedModel); present { if value == "" { return nil } @@ -1162,3 +1188,20 @@ func normalizeOpenAIReasoningEffort(raw string) string { return "" } } + +func normalizeOpenAIReasoningEffortForModel(raw, model string) string { + if strings.EqualFold(strings.TrimSpace(raw), "max") && isOpenAIGPT56Model(model) { + return "max" + } + return normalizeOpenAIReasoningEffort(raw) +} + +func isOpenAIGPT56Model(model string) bool { + normalized := canonicalizeOpenAIModelAliasSpelling(model) + for _, prefix := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { + if normalized == prefix || strings.HasPrefix(normalized, prefix+"-") { + return true + } + } + return false +} diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index 16e552f21f..9f2cb84864 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -38,7 +38,6 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( } clientStream := responsesReq.Stream - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel) serviceTier := extractOpenAIServiceTierFromBody(body) chatReq, err := apicompat.ResponsesToChatCompletionsRequest(&responsesReq) @@ -49,6 +48,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( billingModel := resolveOpenAIForwardModel(account, originalModel, "") upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) // 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) chatReq.Model = upstreamModel diff --git a/backend/internal/service/openai_gpt56_max_test.go b/backend/internal/service/openai_gpt56_max_test.go new file mode 100644 index 0000000000..272eb16ff0 --- /dev/null +++ b/backend/internal/service/openai_gpt56_max_test.go @@ -0,0 +1,266 @@ +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestNormalizeOpenAIReasoningEffortForGPT56(t *testing.T) { + tests := []struct { + name string + raw string + model string + want string + }{ + {name: "Sol 保留 max", raw: "max", model: "gpt-5.6-sol", want: "max"}, + {name: "Terra 保留 max", raw: "max", model: "openai/gpt-5.6-terra", want: "max"}, + {name: "Luna 后缀保留 max", raw: "max", model: "gpt-5.6-luna-2026-07-09", want: "max"}, + {name: "其他模型沿用 xhigh", raw: "max", model: "deepseek-v4-pro", want: "xhigh"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, normalizeOpenAIReasoningEffortForModel(tt.raw, tt.model)) + }) + } +} + +func TestNormalizeOpenAICodexCompactReasoningEffortDowngradesMax(t *testing.T) { + body := []byte(`{"model":"gpt-5.6-sol","input":"compact me","reasoning":{"effort":"max","summary":"auto"}}`) + + normalized, changed, err := normalizeOpenAICodexCompactReasoningEffort(body, "gpt-5.6-sol") + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(normalized, "model").String()) + require.Equal(t, "xhigh", gjson.GetBytes(normalized, "reasoning.effort").String()) + require.Equal(t, "auto", gjson.GetBytes(normalized, "reasoning.summary").String()) +} + +func TestNormalizeOpenAICodexCompactReasoningEffortForAccountScopesCompatibility(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"gpt-5.6-sol","input":"compact me","reasoning":{"effort":"max"}}`) + + tests := []struct { + name string + path string + account *Account + changed bool + want string + }{ + { + name: "OpenAI OAuth compact 降级", + path: "/openai/v1/responses/compact", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + changed: true, + want: "xhigh", + }, + { + name: "OpenAI OAuth 普通请求保留", + path: "/openai/v1/responses", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + want: "max", + }, + { + name: "OpenAI API Key compact 保留", + path: "/openai/v1/responses/compact", + account: &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + want: "max", + }, + { + name: "Grok OAuth compact 保留", + path: "/openai/v1/responses/compact", + account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, + want: "max", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, tt.path, nil) + + normalized, changed, err := normalizeOpenAICodexCompactReasoningEffortForAccount(c, tt.account, body) + + require.NoError(t, err) + require.Equal(t, tt.changed, changed) + require.Equal(t, tt.want, gjson.GetBytes(normalized, "reasoning.effort").String()) + }) + } +} + +func TestOpenAIGatewayServiceForwardPreservesGPT56MaxEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 7, + Name: "openai-apikey", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.6-sol","stream":false,"reasoning":{"effort":"max"},"input":"hello"}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} + +func TestOpenAIGatewayServiceForwardPreservesMappedGPT56MaxEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 9, + Name: "openai-apikey-mapped", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com", + "model_mapping": map[string]any{ + "sol": "gpt-5.6-sol", + }, + }, + Extra: map[string]any{"use_responses_api": true}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"sol","stream":false,"reasoning":{"effort":"max"},"input":"hello"}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} + +func TestOpenAIGatewayServiceForwardOAuthCompactDowngradesMaxEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 8, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + Status: StatusActive, + Schedulable: true, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses/compact", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.6-sol","instructions":"compact-test","input":"hello","reasoning":{"effort":"max"}}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.Equal(t, chatgptCodexURL+"/compact", upstream.lastReq.URL.String()) + require.Equal(t, "xhigh", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "xhigh", *result.ReasoningEffort) +} + +func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 10, + Name: "openai-oauth-responses", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + Status: StatusActive, + Schedulable: true, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.6-sol","instructions":"response-test","input":"hello","reasoning":{"effort":"max"}}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String()) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 5fb9395461..9dea4118da 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -920,7 +920,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( Model: originalModel, UpstreamModel: mappedModel, ServiceTier: extractOpenAIServiceTierFromBody(payload), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, originalModel), payload, mappedModel), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, firstNonEmpty(mappedModel, originalModel)), payload, mappedModel), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: lease.HandshakeHeaders(), diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index ea8c0a0722..40a3cc0194 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -693,7 +693,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( ImageCount: imageCounter.Count(), ImageOutputSizes: imageCounter.Sizes(), ServiceTier: extractOpenAIServiceTier(reqBody), - ReasoningEffort: extractOpenAIReasoningEffort(reqBody, originalModel), + ReasoningEffort: extractOpenAIReasoningEffort(reqBody, firstNonEmpty(mappedModel, originalModel)), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: lease.HandshakeHeaders(), diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 0879d7e8fb..df1a8067cb 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -263,7 +263,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( Model: originalModel, UpstreamModel: mappedModel, ServiceTier: extractOpenAIServiceTierFromBody(body), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, originalModel), body, mappedModel), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(mappedModel, originalModel)), body, mappedModel), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: cloneHeader(resp.Header), From e984b4e2e1c96a1ab523dffefecb89bd79e6e87d Mon Sep 17 00:00:00 2001 From: li Date: Fri, 10 Jul 2026 10:45:44 +0800 Subject: [PATCH 13/27] =?UTF-8?q?fix(i18n):=20=E8=A1=A5=E9=BD=90=20en=20?= =?UTF-8?q?=E8=AF=AD=E8=A8=80=E5=8C=85=20overview=20=E5=92=8C=20resources?= =?UTF-8?q?=20=E7=BC=BA=E5=A4=B1=20key?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #3909 --- .../src/i18n/locales/en/admin/overview.ts | 89 ++++++++++++++++++- .../src/i18n/locales/en/admin/resources.ts | 68 +++++++++++++- 2 files changed, 153 insertions(+), 4 deletions(-) diff --git a/frontend/src/i18n/locales/en/admin/overview.ts b/frontend/src/i18n/locales/en/admin/overview.ts index 0d4a0d0455..7012911065 100644 --- a/frontend/src/i18n/locales/en/admin/overview.ts +++ b/frontend/src/i18n/locales/en/admin/overview.ts @@ -4,15 +4,26 @@ export default { title: 'Admin Dashboard', description: 'System overview and real-time statistics', apiKeys: 'API Keys', + totalApiKeys: 'Total API Keys', + activeApiKeys: 'Active Keys', accounts: 'Accounts', + totalAccounts: 'Total Accounts', + activeAccounts: 'Active Accounts', users: 'Users', + totalUsers: 'Total Users', todayRequests: 'Today Requests', + totalRequests: 'Total Requests', + todayCost: 'Today Cost', + totalCost: 'Total Cost', newUsersToday: 'New Users Today', todayTokens: 'Today Tokens', totalTokens: 'Total Tokens', + input: 'Input', + output: 'Output', cacheToday: 'Cache (Today)', performance: 'Performance', avgResponse: 'Avg Response', + averageTime: 'Average Time', active: 'active', ok: 'ok', err: 'err', @@ -33,6 +44,7 @@ export default { noGroup: 'No Group', requests: 'Requests', tokens: 'Tokens', + cache: 'Cache', actual: 'Actual', standard: 'Standard', accountCost: 'Cost', @@ -50,11 +62,21 @@ export default { spendShort: 'Spend', requestsShort: 'Req', tokensShort: 'Tok', + last7Days: 'Last 7 days', + noUsageRecords: 'No usage records', + startUsingApi: 'Once you start using the API, your usage history will appear here.', + viewAllUsage: 'View all', quickActions: 'Quick Actions', + manageUsers: 'Manage Users', + viewUserAccounts: 'View and manage user accounts', + manageAccounts: 'Manage Accounts', + configureAiAccounts: 'Configure AI platform accounts', batchImage: 'Batch Image', batchImageDesc: 'Submit jobs and copy agent instructions', groupPricing: 'Group Pricing', groupPricingDesc: 'Configure batch discount and hold ratio', + systemSettings: 'System Settings', + configureSystem: 'Configure system settings', failedToLoad: 'Failed to load dashboard statistics' }, @@ -398,7 +420,10 @@ export default { createUser: 'Create User', editUser: 'Edit User', deleteUser: 'Delete User', - searchUsers: 'Search by email, username, notes, or API key...', + deleteConfirmMessage: "Are you sure you want to delete user '{email}'? This action cannot be undone.", + searchPlaceholder: 'Search by email, username, notes, or API key...', + searchUsers: 'Search by email, username, notes, or API key', + roleFilter: 'Role Filter', allRoles: 'All Roles', allStatus: 'All Status', allGroups: 'All Groups', @@ -414,6 +439,8 @@ export default { searchAuthorizedGroups: 'Search authorized groups...', allApiKeyGroups: 'All API Key Groups', searchApiKeyGroups: 'Search API Key groups...', + statusFilter: 'Status Filter', + allStatuses: 'All Status', admin: 'Admin', user: 'User', disabled: 'Disabled', @@ -433,7 +460,21 @@ export default { creating: 'Creating...', updating: 'Updating...', form: { + emailLabel: 'Email', + emailPlaceholder: 'Enter email', + usernameLabel: 'Username', + usernamePlaceholder: 'Enter username (optional)', + notesLabel: 'Notes', + notesPlaceholder: 'Enter notes (admin only)', + notesHint: 'This note is only visible to administrators', + passwordLabel: 'Password', + passwordPlaceholder: 'Enter password (leave empty to keep unchanged)', roleLabel: 'Role', + selectRole: 'Select role', + balanceLabel: 'Balance', + concurrencyLabel: 'Concurrency', + statusLabel: 'Status', + selectStatus: 'Select status', rpmLimit: 'Requests Per Minute (RPM)', rpmLimitPlaceholder: '0 = unlimited', rpmLimitHint: 'Max requests per minute for this user; 0 = unlimited. Acts as a fallback only when the group has no rpm_limit set.' @@ -504,6 +545,21 @@ export default { soraStorageQuotaHint: 'In GB, 0 means use group or system default quota', amountRequired: 'Please enter a valid amount', insufficientBalance: 'Insufficient balance', + adjustBalance: 'Adjust Balance', + adjustConcurrency: 'Adjust Concurrency', + adjustmentAmount: 'Adjustment Amount', + adjustmentAmountHint: 'Positive to add, negative to subtract', + currentConcurrency: 'Current Concurrency', + saving: 'Saving...', + noUsers: 'No users yet', + noUsersDescription: 'Create your first user to get started.', + userCreatedSuccess: 'User created successfully', + userUpdatedSuccess: 'User updated successfully', + userDeletedSuccess: 'User deleted successfully', + balanceAdjustedSuccess: 'Balance adjusted successfully', + concurrencyAdjustedSuccess: 'Concurrency adjusted successfully', + failedToSave: 'Failed to save user', + failedToAdjust: 'Adjustment failed', deleteConfirm: "Are you sure you want to delete '{email}'? This action cannot be undone.", setAllowedGroups: 'Set Allowed Groups', allowedGroupsHint: @@ -695,6 +751,7 @@ export default { allPlatforms: 'All Platforms', allStatus: 'All Status', allGroups: 'All Groups', + exclusiveFilter: 'Exclusive', exclusive: 'Exclusive', nonExclusive: 'Non-Exclusive', public: 'Public', @@ -706,7 +763,10 @@ export default { rpmOverrideHint: 'Per-user RPM cap in this group; empty = group default; 0 = unlimited', rateDefault: 'default', rpmDefault: 'default', + exclusive: 'Exclusive', type: 'Type', + priority: 'Priority', + apiKeys: 'API Keys', accounts: 'Accounts', capacity: 'Capacity', usage: 'Usage', @@ -742,14 +802,39 @@ export default { rateMultiplier: 'Rate Multiplier', status: 'Status', exclusive: 'Exclusive Group', + nameLabel: 'Group Name', + namePlaceholder: 'Enter group name', + descriptionLabel: 'Description', + descriptionPlaceholder: 'Enter description (optional)', + rateMultiplierLabel: 'Rate Multiplier', + rateMultiplierHint: '1.0 = standard rate, 0.5 = half price, 2.0 = double', rpmLimit: 'Requests Per Minute (RPM)', rpmLimitPlaceholder: '0 = unlimited', - rpmLimitHint: 'Max requests per minute for each user in this group; 0 = unlimited. Once set, it takes over per-user rate limiting in this group (overrides the user-level rpm_limit fallback).' + rpmLimitHint: 'Max requests per minute for each user in this group; 0 = unlimited. Once set, it takes over per-user rate limiting in this group (overrides the user-level rpm_limit fallback).', + exclusiveLabel: 'Exclusive Group', + exclusiveHint: 'Exclusive group, can be manually assigned to users', + platformLabel: 'Platform Restriction', + platformPlaceholder: 'Select platform (leave empty for no restriction)', + accountsLabel: 'Designated Accounts', + accountsPlaceholder: 'Select accounts (leave empty for no restriction)', + priorityLabel: 'Priority', + priorityHint: 'Lower value means higher priority, used for account scheduling', + statusLabel: 'Status' + }, + exclusiveObj: { + yes: 'Yes', + no: 'No' }, enterGroupName: 'Enter group name', optionalDescription: 'Optional description', platformHint: 'Select the platform this group is associated with', platformNotEditable: 'Platform cannot be changed after creation', + saving: 'Saving...', + noGroups: 'No groups yet', + noGroupsDescription: 'Create a group to better manage API keys and rates.', + groupCreatedSuccess: 'Group created successfully', + groupUpdatedSuccess: 'Group updated successfully', + groupDeletedSuccess: 'Group deleted successfully', rateMultiplierHint: 'Cost multiplier for this group (e.g., 1.5 = 150% of base cost)', exclusiveHint: 'Exclusive group, manually assign to specific users', exclusiveTooltip: { diff --git a/frontend/src/i18n/locales/en/admin/resources.ts b/frontend/src/i18n/locales/en/admin/resources.ts index 8afd35c60a..8e9e34c0ec 100644 --- a/frontend/src/i18n/locales/en/admin/resources.ts +++ b/frontend/src/i18n/locales/en/admin/resources.ts @@ -50,6 +50,8 @@ export default { ad: { inline: 'Need proxy IP?' }, + deleteConfirmMessage: "Are you sure you want to delete proxy '{name}'?", + testProxy: 'Test Proxy', dataImport: 'Import', dataExportSelected: 'Export Selected', dataImportTitle: 'Import Proxies', @@ -93,7 +95,27 @@ export default { latency: 'Latency', expiry: 'Validity', createdAt: 'Created', - actions: 'Actions' + actions: 'Actions', + nameLabel: 'Name', + namePlaceholder: 'Enter proxy name', + protocolLabel: 'Protocol', + selectProtocol: 'Select protocol', + hostLabel: 'Host', + hostPlaceholder: 'Enter host address', + portLabel: 'Port', + portPlaceholder: 'Enter port', + usernameLabel: 'Username (Optional)', + usernamePlaceholder: 'Enter username', + passwordLabel: 'Password (Optional)', + passwordPlaceholder: 'Enter password', + priorityLabel: 'Priority', + statusLabel: 'Status' + }, + filters: { + protocol: 'Protocol', + allProtocols: 'All Protocols', + status: 'Status', + allStatuses: 'All Status' }, testConnection: 'Test Connection', qualityCheck: 'Quality Check', @@ -150,8 +172,12 @@ export default { batchImportAllSkipped: 'All {skipped} proxies already exist, skipped import', failedToImport: 'Failed to batch import', // Other messages + saving: 'Saving...', + testing: 'Testing...', creating: 'Creating...', updating: 'Updating...', + noProxies: 'No proxies yet', + noProxiesDescription: 'Add a proxy server to improve API access stability.', proxyCreated: 'Proxy created successfully', proxyUpdated: 'Proxy updated successfully', proxyDeleted: 'Proxy deleted successfully', @@ -181,7 +207,12 @@ export default { qualityStatusFail: 'Fail', qualityStatusChallenge: 'Challenge', qualityTargetBase: 'Base Connectivity', + proxyCreatedSuccess: 'Proxy created successfully', + proxyUpdatedSuccess: 'Proxy updated successfully', + proxyDeletedSuccess: 'Proxy deleted successfully', + testSuccess: 'Proxy test passed', failedToLoad: 'Failed to load proxies', + failedToSave: 'Failed to save proxy', failedToCreate: 'Failed to create proxy', failedToUpdate: 'Failed to update proxy', failedToDelete: 'Failed to delete proxy', @@ -229,6 +260,7 @@ export default { status: 'Status', usedBy: 'Used By', usedAt: 'Used At', + createdAt: 'Created At', expiresAt: 'Expires At', actions: 'Actions' }, @@ -304,7 +336,39 @@ export default { used: 'Used', expired: 'Expired', disabled: 'Disabled' - } + }, + form: { + typeLabel: 'Type', + selectType: 'Select type', + valueLabel: 'Value', + valuePlaceholder: 'Enter value', + balanceHint: 'Balance amount (USD)', + concurrencyHint: 'Concurrency increment', + countLabel: 'Count', + countPlaceholder: 'Enter count', + countHint: 'Number of redeem codes to generate', + prefixLabel: 'Prefix (Optional)', + prefixPlaceholder: 'e.g., GIFT', + expiresLabel: 'Expires At (Optional)' + }, + filters: { + type: 'Type', + allTypes: 'All Types', + status: 'Status', + allStatuses: 'All Status', + search: 'Search codes' + }, + copyCode: 'Copy', + disableCode: 'Disable', + enableCode: 'Enable', + deleteConfirmMessage: 'Are you sure you want to delete this redeem code?', + noCodes: 'No redeem codes yet', + noCodesDescription: 'Generate redeem codes to distribute balance or concurrency to users.', + codesGeneratedSuccess: 'Redeem codes generated successfully, {count} total', + codeDisabledSuccess: 'Redeem code disabled', + codeEnabledSuccess: 'Redeem code enabled', + codeDeletedSuccess: 'Redeem code deleted successfully', + failedToUpdate: 'Failed to update redeem code' }, // Announcements From fc66a30ffc2ccbc4caa7478598095b93cbf6d2e4 Mon Sep 17 00:00:00 2001 From: superman2003 <2112076433zcr@gmail.com> Date: Fri, 10 Jul 2026 10:45:27 +0800 Subject: [PATCH 14/27] fix: harden billing concurrency and payment recovery --- Makefile | 12 +- backend/internal/repository/user_repo.go | 44 +++ .../repository/user_repo_integration_test.go | 47 ++++ .../user_repo_redeem_adjustment_test.go | 56 ++++ .../repository/user_subscription_repo.go | 72 ++++- ...user_subscription_repo_integration_test.go | 49 +++- backend/internal/server/api_contract_test.go | 9 +- .../server/middleware/api_key_auth.go | 15 +- .../server/middleware/api_key_auth_google.go | 14 +- .../middleware/api_key_auth_google_test.go | 13 +- .../server/middleware/api_key_auth_test.go | 97 ++++++- .../internal/service/payment_fulfillment.go | 266 ++++++++++++++---- .../service/payment_fulfillment_test.go | 233 +++++++++++++++ backend/internal/service/redeem_service.go | 31 +- .../subscription_assign_idempotency_test.go | 47 +++- .../subscription_expiry_service_test.go | 10 +- .../service/subscription_reset_quota_test.go | 43 ++- .../internal/service/subscription_service.go | 94 ++++--- backend/internal/service/user_service.go | 8 + .../user_subscription_daily_quota_test.go | 2 +- .../service/user_subscription_port.go | 7 +- .../router/__tests__/feature-access.spec.ts | 177 ++++++++++++ frontend/src/router/index.ts | 29 +- frontend/src/stores/__tests__/app.spec.ts | 127 +++++++++ frontend/src/stores/app.ts | 49 +++- 25 files changed, 1353 insertions(+), 198 deletions(-) create mode 100644 backend/internal/repository/user_repo_redeem_adjustment_test.go create mode 100644 frontend/src/router/__tests__/feature-access.spec.ts diff --git a/Makefile b/Makefile index d00d0c4f5e..c878526f96 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: build build-backend build-frontend build-datamanagementd test test-backend test-frontend test-frontend-critical test-datamanagementd secret-scan +.PHONY: build build-backend build-frontend test test-backend test-frontend test-frontend-critical FRONTEND_CRITICAL_VITEST := \ src/views/auth/__tests__/LinuxDoCallbackView.spec.ts \ @@ -19,10 +19,6 @@ build-backend: build-frontend: @pnpm --dir frontend run build -# 编译 datamanagementd(宿主机数据管理进程) -build-datamanagementd: - @cd datamanagement && go build -o datamanagementd ./cmd/datamanagementd - # 运行测试(后端 + 前端) test: test-backend test-frontend @@ -36,9 +32,3 @@ test-frontend: test-frontend-critical: @pnpm --dir frontend exec vitest run $(FRONTEND_CRITICAL_VITEST) - -test-datamanagementd: - @cd datamanagement && go test ./... - -secret-scan: - @python3 tools/secret_scan.py diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go index 3ac8dcfbf8..d4c0d0dc28 100644 --- a/backend/internal/repository/user_repo.go +++ b/backend/internal/repository/user_repo.go @@ -32,6 +32,8 @@ type userRepository struct { sql sqlExecutor } +var _ service.RedeemUserAdjustmentRepository = (*userRepository)(nil) + func NewUserRepository(client *dbent.Client, sqlDB *sql.DB) service.UserRepository { return newUserRepositoryWithSQL(client, sqlDB) } @@ -751,6 +753,27 @@ func (r *userRepository) UpdateBalance(ctx context.Context, id int64, amount flo return nil } +func (r *userRepository) ApplyRedeemBalanceAdjustment(ctx context.Context, id int64, delta float64) error { + const updateSQL = ` + UPDATE users + SET balance = GREATEST(balance + $1, 0), updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL + ` + client := clientFromContext(ctx, r.client) + result, err := client.ExecContext(ctx, updateSQL, delta, id) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrUserNotFound + } + return nil +} + // DeductBalance 扣除用户余额 // 透支策略:允许余额变为负数,确保当前请求能够完成 // 中间件会阻止余额 <= 0 的用户发起后续请求 @@ -792,6 +815,27 @@ func (r *userRepository) UpdateConcurrency(ctx context.Context, id int64, amount return nil } +func (r *userRepository) ApplyRedeemConcurrencyAdjustment(ctx context.Context, id int64, delta int) error { + const updateSQL = ` + UPDATE users + SET concurrency = GREATEST(concurrency + $1, 0), updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL + ` + client := clientFromContext(ctx, r.client) + result, err := client.ExecContext(ctx, updateSQL, delta, id) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrUserNotFound + } + return nil +} + func (r *userRepository) BatchSetConcurrency(ctx context.Context, userIDs []int64, value int) (int, error) { if len(userIDs) == 0 { return 0, nil diff --git a/backend/internal/repository/user_repo_integration_test.go b/backend/internal/repository/user_repo_integration_test.go index 13a605a2f5..42d10af632 100644 --- a/backend/internal/repository/user_repo_integration_test.go +++ b/backend/internal/repository/user_repo_integration_test.go @@ -4,6 +4,7 @@ package repository import ( "context" + "sync" "testing" "time" @@ -353,6 +354,29 @@ func (s *UserRepoSuite) TestUpdateBalance_Negative() { s.Require().InDelta(7.0, got.Balance, 1e-6) } +func (s *UserRepoSuite) TestApplyRedeemBalanceAdjustment_ConcurrentNeverNegative() { + user := s.mustCreateUser(&service.User{Email: "redeem-bal-concurrent@test.com", Balance: 10}) + + var wg sync.WaitGroup + errs := make(chan error, 2) + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + errs <- s.repo.ApplyRedeemBalanceAdjustment(context.Background(), user.ID, -7) + }() + } + wg.Wait() + close(errs) + for err := range errs { + s.Require().NoError(err) + } + + got, err := s.repo.GetByID(s.ctx, user.ID) + s.Require().NoError(err) + s.Require().InDelta(0, got.Balance, 1e-6) +} + func (s *UserRepoSuite) TestDeductBalance() { user := s.mustCreateUser(&service.User{Email: "deduct@test.com", Balance: 10}) @@ -425,6 +449,29 @@ func (s *UserRepoSuite) TestUpdateConcurrency_Negative() { s.Require().Equal(3, got.Concurrency) } +func (s *UserRepoSuite) TestApplyRedeemConcurrencyAdjustment_ConcurrentNeverNegative() { + user := s.mustCreateUser(&service.User{Email: "redeem-concurrency-concurrent@test.com", Concurrency: 10}) + + var wg sync.WaitGroup + errs := make(chan error, 2) + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + errs <- s.repo.ApplyRedeemConcurrencyAdjustment(context.Background(), user.ID, -7) + }() + } + wg.Wait() + close(errs) + for err := range errs { + s.Require().NoError(err) + } + + got, err := s.repo.GetByID(s.ctx, user.ID) + s.Require().NoError(err) + s.Require().Equal(0, got.Concurrency) +} + // --- ExistsByEmail --- func (s *UserRepoSuite) TestExistsByEmail() { diff --git a/backend/internal/repository/user_repo_redeem_adjustment_test.go b/backend/internal/repository/user_repo_redeem_adjustment_test.go new file mode 100644 index 0000000000..2c0d21e4bd --- /dev/null +++ b/backend/internal/repository/user_repo_redeem_adjustment_test.go @@ -0,0 +1,56 @@ +package repository + +import ( + "context" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +func newRedeemAdjustmentRepoMock(t *testing.T) (*userRepository, sqlmock.Sqlmock) { + t.Helper() + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + return newUserRepositoryWithSQL(client, db), mock +} + +func TestApplyRedeemBalanceAdjustment_UsesAtomicFloor(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET balance = GREATEST\(balance \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-7.0, int64(42)). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.ApplyRedeemBalanceAdjustment(context.Background(), 42, -7)) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestApplyRedeemConcurrencyAdjustment_UsesAtomicFloor(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET concurrency = GREATEST\(concurrency \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-7, int64(42)). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.ApplyRedeemConcurrencyAdjustment(context.Background(), 42, -7)) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestApplyRedeemAdjustment_MissingUser(t *testing.T) { + repo, mock := newRedeemAdjustmentRepoMock(t) + mock.ExpectExec(`UPDATE users SET balance = GREATEST\(balance \+ \$1, 0\), updated_at = NOW\(\) WHERE id = \$2 AND deleted_at IS NULL`). + WithArgs(-1.0, int64(404)). + WillReturnResult(sqlmock.NewResult(0, 0)) + + err := repo.ApplyRedeemBalanceAdjustment(context.Background(), 404, -1) + require.ErrorIs(t, err, service.ErrUserNotFound) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/user_subscription_repo.go b/backend/internal/repository/user_subscription_repo.go index 6326c9711f..37f06a038b 100644 --- a/backend/internal/repository/user_subscription_repo.go +++ b/backend/internal/repository/user_subscription_repo.go @@ -366,31 +366,85 @@ func (r *userSubscriptionRepository) ActivateWindows(ctx context.Context, id int return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) } -func (r *userSubscriptionRepository) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + update := client.UserSubscription.UpdateOneID(id) + if resetDaily { + update.SetDailyUsageUsd(0).SetDailyWindowStart(newWindowStart) + } + if resetWeekly { + update.SetWeeklyUsageUsd(0).SetWeeklyWindowStart(newWindowStart) + } + if resetMonthly { + update.SetMonthlyUsageUsd(0).SetMonthlyWindowStart(newWindowStart) + } + _, err := update.Save(ctx) + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) +} + +func (r *userSubscriptionRepository) ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { + client := clientFromContext(ctx, r.client) + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.DailyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.DailyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetDailyUsageUsd(0). SetDailyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) } -func (r *userSubscriptionRepository) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.WeeklyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.WeeklyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetWeeklyUsageUsd(0). SetWeeklyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) } -func (r *userSubscriptionRepository) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { client := clientFromContext(ctx, r.client) - _, err := client.UserSubscription.UpdateOneID(id). + query := client.UserSubscription.Update().Where(usersubscription.IDEQ(id)) + if expectedWindowStart == nil { + query = query.Where(usersubscription.MonthlyWindowStartIsNil()) + } else { + query = query.Where(usersubscription.MonthlyWindowStartEQ(*expectedWindowStart)) + } + n, err := query. SetMonthlyUsageUsd(0). SetMonthlyWindowStart(newWindowStart). Save(ctx) - return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + return r.translateConditionalWindowReset(ctx, client, id, n, err) +} + +func (r *userSubscriptionRepository) translateConditionalWindowReset(ctx context.Context, client *dbent.Client, id int64, affected int, err error) error { + if err != nil { + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + } + if affected > 0 { + return nil + } + + // A stale reset is an expected no-op: another request already advanced the + // window. Preserve not-found semantics for callers that target a missing row. + exists, err := client.UserSubscription.Query().Where(usersubscription.IDEQ(id)).Exist(ctx) + if err != nil { + return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) + } + if !exists { + return service.ErrSubscriptionNotFound + } + return nil } // IncrementUsage 原子性地累加订阅用量。 diff --git a/backend/internal/repository/user_subscription_repo_integration_test.go b/backend/internal/repository/user_subscription_repo_integration_test.go index caa88cc640..96eead494e 100644 --- a/backend/internal/repository/user_subscription_repo_integration_test.go +++ b/backend/internal/repository/user_subscription_repo_integration_test.go @@ -472,7 +472,7 @@ func (s *UserSubscriptionRepoSuite) TestResetDailyUsage() { }) resetAt := time.Date(2025, 1, 2, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetDailyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetDailyUsage(s.ctx, sub.ID, sub.DailyWindowStart, resetAt) s.Require().NoError(err, "ResetDailyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -483,6 +483,47 @@ func (s *UserSubscriptionRepoSuite) TestResetDailyUsage() { s.Require().WithinDuration(resetAt, *got.DailyWindowStart, time.Microsecond) } +func (s *UserSubscriptionRepoSuite) TestResetDailyUsage_StaleResetDoesNotClearNewWindowUsage() { + user := s.mustCreateUser("resetd-cas@test.com", service.RoleUser) + group := s.mustCreateGroup("g-resetd-cas") + oldWindowStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + sub := s.mustCreateSubscription(user.ID, group.ID, func(c *dbent.UserSubscriptionCreate) { + c.SetDailyWindowStart(oldWindowStart) + c.SetDailyUsageUsd(10) + }) + + newWindowStart := oldWindowStart.Add(24 * time.Hour) + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + s.Require().NoError(s.repo.IncrementUsage(s.ctx, sub.ID, 3)) + // Simulate a second request carrying the stale old-window snapshot. + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + + got, err := s.repo.GetByID(s.ctx, sub.ID) + s.Require().NoError(err) + s.Require().InDelta(3, got.DailyUsageUSD, 1e-6) + s.Require().WithinDuration(newWindowStart, *got.DailyWindowStart, time.Microsecond) +} + +func (s *UserSubscriptionRepoSuite) TestResetUsageWindows_ClearsUsageAfterAutomaticWindowAdvance() { + user := s.mustCreateUser("admin-reset-current@test.com", service.RoleUser) + group := s.mustCreateGroup("g-admin-reset-current") + oldWindowStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + sub := s.mustCreateSubscription(user.ID, group.ID, func(c *dbent.UserSubscriptionCreate) { + c.SetDailyWindowStart(oldWindowStart) + c.SetDailyUsageUsd(10) + }) + + newWindowStart := oldWindowStart.Add(24 * time.Hour) + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) + s.Require().NoError(s.repo.IncrementUsage(s.ctx, sub.ID, 3)) + s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, false, false, newWindowStart)) + + got, err := s.repo.GetByID(s.ctx, sub.ID) + s.Require().NoError(err) + s.Require().InDelta(0, got.DailyUsageUSD, 1e-6) + s.Require().WithinDuration(newWindowStart, *got.DailyWindowStart, time.Microsecond) +} + func (s *UserSubscriptionRepoSuite) TestResetWeeklyUsage() { user := s.mustCreateUser("resetw@test.com", service.RoleUser) group := s.mustCreateGroup("g-resetw") @@ -492,7 +533,7 @@ func (s *UserSubscriptionRepoSuite) TestResetWeeklyUsage() { }) resetAt := time.Date(2025, 1, 6, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetWeeklyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetWeeklyUsage(s.ctx, sub.ID, sub.WeeklyWindowStart, resetAt) s.Require().NoError(err, "ResetWeeklyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -511,7 +552,7 @@ func (s *UserSubscriptionRepoSuite) TestResetMonthlyUsage() { }) resetAt := time.Date(2025, 2, 1, 0, 0, 0, 0, time.UTC) - err := s.repo.ResetMonthlyUsage(s.ctx, sub.ID, resetAt) + err := s.repo.ResetMonthlyUsage(s.ctx, sub.ID, sub.MonthlyWindowStart, resetAt) s.Require().NoError(err, "ResetMonthlyUsage") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -723,7 +764,7 @@ func (s *UserSubscriptionRepoSuite) TestActiveExpiredBoundaries_UsageAndReset_Ba s.Require().NotNil(after.MonthlyWindowStart, "expected MonthlyWindowStart activated") resetAt := time.Now().Truncate(time.Microsecond) // truncate to microsecond for DB precision - s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, active.ID, resetAt), "ResetDailyUsage") + s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, active.ID, after.DailyWindowStart, resetAt), "ResetDailyUsage") afterReset, err := s.repo.GetByID(s.ctx, active.ID) s.Require().NoError(err, "GetByID after reset") s.Require().InDelta(0.0, afterReset.DailyUsageUSD, 1e-6) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index f48ecb060a..d260afe738 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -2123,13 +2123,16 @@ func (stubUserSubscriptionRepo) UpdateNotes(ctx context.Context, subscriptionID func (stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, start time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { + return errors.New("not implemented") +} +func (stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { return errors.New("not implemented") } func (stubUserSubscriptionRepo) IncrementUsage(ctx context.Context, id int64, costUSD float64) error { diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go index 1610390440..04a09862b5 100644 --- a/backend/internal/server/middleware/api_key_auth.go +++ b/backend/internal/server/middleware/api_key_auth.go @@ -193,6 +193,15 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti // 订阅模式:验证订阅限额 if subscription != nil { needsMaintenance, validateErr := subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + if needsMaintenance { + refreshed, maintenanceErr := subscriptionService.EnsureWindowMaintenance(c.Request.Context(), subscription) + if maintenanceErr != nil { + AbortWithError(c, 500, "SUBSCRIPTION_MAINTENANCE_FAILED", "Failed to maintain subscription usage windows") + return + } + subscription = refreshed + _, validateErr = subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + } if validateErr != nil { code := "SUBSCRIPTION_INVALID" status := 403 @@ -205,12 +214,6 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti AbortWithError(c, status, code, validateErr.Error()) return } - - // 窗口维护异步化(不阻塞请求) - if needsMaintenance { - maintenanceCopy := *subscription - subscriptionService.DoWindowMaintenance(&maintenanceCopy) - } } else { // 非订阅模式 或 订阅模式但 subscriptionService 未注入:回退到余额检查 if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { diff --git a/backend/internal/server/middleware/api_key_auth_google.go b/backend/internal/server/middleware/api_key_auth_google.go index c75d5b99f2..b910dd9fa4 100644 --- a/backend/internal/server/middleware/api_key_auth_google.go +++ b/backend/internal/server/middleware/api_key_auth_google.go @@ -141,6 +141,15 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs } needsMaintenance, err := subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + if needsMaintenance { + refreshed, maintenanceErr := subscriptionService.EnsureWindowMaintenance(c.Request.Context(), subscription) + if maintenanceErr != nil { + abortWithGoogleError(c, 500, "Failed to maintain subscription usage windows") + return + } + subscription = refreshed + _, err = subscriptionService.ValidateAndCheckLimits(subscription, apiKey.Group) + } if err != nil { status := 403 if errors.Is(err, service.ErrDailyLimitExceeded) || @@ -153,11 +162,6 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs } c.Set(string(ContextKeySubscription), subscription) - - if needsMaintenance { - maintenanceCopy := *subscription - subscriptionService.DoWindowMaintenance(&maintenanceCopy) - } } else { if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { abortWithGoogleError(c, 403, "Insufficient account balance") diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go index 45ddb0bf93..746d238ca5 100644 --- a/backend/internal/server/middleware/api_key_auth_google_test.go +++ b/backend/internal/server/middleware/api_key_auth_google_test.go @@ -24,6 +24,7 @@ type fakeAPIKeyRepo struct { } type fakeGoogleSubscriptionRepo struct { + getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error activateWindow func(ctx context.Context, id int64, start time.Time) error @@ -115,6 +116,9 @@ func (f fakeGoogleSubscriptionRepo) Create(ctx context.Context, sub *service.Use return errors.New("not implemented") } func (f fakeGoogleSubscriptionRepo) GetByID(ctx context.Context, id int64) (*service.UserSubscription, error) { + if f.getByID != nil { + return f.getByID(ctx, id) + } return nil, errors.New("not implemented") } func (f fakeGoogleSubscriptionRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.UserSubscription, error) { @@ -174,19 +178,22 @@ func (f fakeGoogleSubscriptionRepo) ActivateWindows(ctx context.Context, id int6 } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + return errors.New("not implemented") +} +func (f fakeGoogleSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetDaily != nil { return f.resetDaily(ctx, id, start) } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetWeekly != nil { return f.resetWeekly(ctx, id, start) } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { if f.resetMonthly != nil { return f.resetMonthly(ctx, id, start) } diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index d5fbc46098..abb84ab852 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -58,7 +58,7 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { }, } - t.Run("standard_mode_needs_maintenance_does_not_block_request", func(t *testing.T) { + t.Run("standard_mode_completes_maintenance_before_request", func(t *testing.T) { cfg := &config.Config{RunMode: config.RunModeStandard} cfg.SubscriptionMaintenance.WorkerCount = 1 cfg.SubscriptionMaintenance.QueueSize = 1 @@ -67,16 +67,22 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { past := time.Now().Add(-48 * time.Hour) sub := &service.UserSubscription{ - ID: 55, - UserID: user.ID, - GroupID: group.ID, - Status: service.SubscriptionStatusActive, - ExpiresAt: time.Now().Add(24 * time.Hour), - DailyWindowStart: &past, - DailyUsageUSD: 0, + ID: 55, + UserID: user.ID, + GroupID: group.ID, + Status: service.SubscriptionStatusActive, + ExpiresAt: time.Now().Add(24 * time.Hour), + DailyWindowStart: &past, + WeeklyWindowStart: &past, + MonthlyWindowStart: &past, + DailyUsageUSD: 0, } maintenanceCalled := make(chan struct{}, 1) subscriptionRepo := &stubUserSubscriptionRepo{ + getByID: func(ctx context.Context, id int64) (*service.UserSubscription, error) { + clone := *sub + return &clone, nil + }, getActive: func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) { clone := *sub return &clone, nil @@ -84,11 +90,19 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { updateStatus: func(ctx context.Context, subscriptionID int64, status string) error { return nil }, activateWindow: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetDaily: func(ctx context.Context, id int64, start time.Time) error { + sub.DailyWindowStart = &start + sub.DailyUsageUSD = 0 maintenanceCalled <- struct{}{} return nil }, - resetWeekly: func(ctx context.Context, id int64, start time.Time) error { return nil }, - resetMonthly: func(ctx context.Context, id int64, start time.Time) error { return nil }, + resetWeekly: func(ctx context.Context, id int64, start time.Time) error { + sub.WeeklyWindowStart = &start + return nil + }, + resetMonthly: func(ctx context.Context, id int64, start time.Time) error { + sub.MonthlyWindowStart = &start + return nil + }, } subscriptionService := service.NewSubscriptionService(nil, subscriptionRepo, nil, nil, cfg) t.Cleanup(subscriptionService.Stop) @@ -105,10 +119,57 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { case <-maintenanceCalled: // ok case <-time.After(time.Second): - t.Fatalf("expected maintenance to be scheduled") + t.Fatalf("expected maintenance to complete before response") } }) + t.Run("standard_mode_revalidates_cas_loser_from_database", func(t *testing.T) { + cfg := &config.Config{RunMode: config.RunModeStandard} + apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) + + past := time.Now().Add(-48 * time.Hour) + current := time.Now() + stale := &service.UserSubscription{ + ID: 56, + UserID: user.ID, + GroupID: group.ID, + Status: service.SubscriptionStatusActive, + ExpiresAt: current.Add(24 * time.Hour), + DailyWindowStart: &past, + WeeklyWindowStart: &past, + MonthlyWindowStart: &past, + DailyUsageUSD: 10, + } + fresh := *stale + fresh.DailyWindowStart = ¤t + fresh.WeeklyWindowStart = ¤t + fresh.MonthlyWindowStart = ¤t + fresh.DailyUsageUSD = 2 + + subscriptionRepo := &stubUserSubscriptionRepo{ + getActive: func(context.Context, int64, int64) (*service.UserSubscription, error) { + clone := *stale + return &clone, nil + }, + getByID: func(context.Context, int64) (*service.UserSubscription, error) { + clone := fresh + return &clone, nil + }, + resetDaily: func(context.Context, int64, time.Time) error { return nil }, + resetWeekly: func(context.Context, int64, time.Time) error { return nil }, + resetMonthly: func(context.Context, int64, time.Time) error { return nil }, + } + subscriptionService := service.NewSubscriptionService(nil, subscriptionRepo, nil, nil, cfg) + router := newAuthTestRouter(apiKeyService, subscriptionService, cfg) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/t", nil) + req.Header.Set("x-api-key", apiKey.Key) + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusTooManyRequests, w.Code) + }) + t.Run("simple_mode_bypasses_quota_check", func(t *testing.T) { cfg := &config.Config{RunMode: config.RunModeSimple} apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) @@ -1210,6 +1271,7 @@ func (r *stubApiKeyRepo) GetRateLimitData(ctx context.Context, id int64) (*servi } type stubUserSubscriptionRepo struct { + getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error activateWindow func(ctx context.Context, id int64, start time.Time) error @@ -1258,6 +1320,9 @@ func (r *stubUserSubscriptionRepo) Create(ctx context.Context, sub *service.User } func (r *stubUserSubscriptionRepo) GetByID(ctx context.Context, id int64) (*service.UserSubscription, error) { + if r.getByID != nil { + return r.getByID(ctx, id) + } return nil, errors.New("not implemented") } @@ -1334,21 +1399,25 @@ func (r *stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64 return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + return errors.New("not implemented") +} + +func (r *stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetDaily != nil { return r.resetDaily(ctx, id, newWindowStart) } return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetWeeklyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetWeekly != nil { return r.resetWeekly(ctx, id, newWindowStart) } return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error { +func (r *stubUserSubscriptionRepo) ResetMonthlyUsage(ctx context.Context, id int64, _ *time.Time, newWindowStart time.Time) error { if r.resetMonthly != nil { return r.resetMonthly(ctx, id, newWindowStart) } diff --git a/backend/internal/service/payment_fulfillment.go b/backend/internal/service/payment_fulfillment.go index 51ecefe250..4d442f3d1e 100644 --- a/backend/internal/service/payment_fulfillment.go +++ b/backend/internal/service/payment_fulfillment.go @@ -28,6 +28,12 @@ import ( // misconfigured to point at us, or when our orders table has been wiped). var ErrOrderNotFound = errors.New("payment order not found") +const paymentFulfillmentLeaseDuration = 5 * time.Minute + +type paymentFulfillmentLease struct { + version time.Time +} + // --- Payment Notification & Fulfillment --- func (s *PaymentService) HandlePaymentNotification(ctx context.Context, n *payment.PaymentNotification, pk string) error { @@ -188,10 +194,8 @@ func (s *PaymentService) alreadyProcessed(ctx context.Context, o *dbent.PaymentO switch cur.Status { case OrderStatusCompleted, OrderStatusRefunded: return nil - case OrderStatusFailed: + case OrderStatusFailed, OrderStatusPaid, OrderStatusRecharging: return s.executeFulfillment(ctx, o.ID) - case OrderStatusPaid, OrderStatusRecharging: - return fmt.Errorf("order %d is being processed", o.ID) case OrderStatusExpired: slog.Warn("webhook payment success for expired order beyond grace period", "orderID", o.ID, @@ -231,23 +235,74 @@ func (s *PaymentService) ExecuteBalanceFulfillment(ctx context.Context, oid int6 if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot fulfill") } - if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed { + if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed && o.Status != OrderStatusRecharging { return infraerrors.BadRequest("INVALID_STATUS", "order cannot fulfill in status "+o.Status) } - c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed)).SetStatus(OrderStatusRecharging).Save(ctx) + lease, err := s.acquirePaymentFulfillmentLease(ctx, o) if err != nil { - return fmt.Errorf("lock: %w", err) + return err } - if c == 0 { + if lease == nil { return nil } - if err := s.doBalance(ctx, o); err != nil { - s.markFailed(ctx, oid, err) + if err := s.doBalance(ctx, o, lease); err != nil { + s.markFailed(ctx, oid, lease, err) return err } return nil } +func (s *PaymentService) acquirePaymentFulfillmentLease(ctx context.Context, o *dbent.PaymentOrder) (*paymentFulfillmentLease, error) { + if o == nil { + return nil, infraerrors.BadRequest("INVALID_STATUS", "nil payment order") + } + + now := time.Now().UTC().Truncate(time.Microsecond) + staleBefore := now.Add(-paymentFulfillmentLeaseDuration) + updated, err := s.entClient.PaymentOrder.Update(). + Where( + paymentorder.IDEQ(o.ID), + paymentorder.Or( + paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed), + paymentorder.And( + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtLTE(staleBefore), + ), + ), + ). + SetStatus(OrderStatusRecharging). + SetUpdatedAt(now). + ClearFailedAt(). + ClearFailedReason(). + Save(ctx) + if err != nil { + return nil, fmt.Errorf("acquire fulfillment lease: %w", err) + } + if updated == 0 { + current, getErr := s.entClient.PaymentOrder.Get(ctx, o.ID) + if getErr != nil { + return nil, fmt.Errorf("reload fulfillment lease: %w", getErr) + } + if current.Status == OrderStatusCompleted { + return nil, nil + } + if current.Status == OrderStatusRecharging { + return nil, infraerrors.Conflict("CONFLICT", "order is being processed") + } + return nil, infraerrors.Conflict("CONFLICT", "order status changed while acquiring fulfillment lease") + } + + // Reload the persisted timestamp instead of trusting application clock precision. + claimed, err := s.entClient.PaymentOrder.Get(ctx, o.ID) + if err != nil { + return nil, fmt.Errorf("reload acquired fulfillment lease: %w", err) + } + if claimed.Status != OrderStatusRecharging { + return nil, infraerrors.Conflict("CONFLICT", "fulfillment lease was lost") + } + return &paymentFulfillmentLease{version: claimed.UpdatedAt}, nil +} + // redeemAction represents the idempotency decision for balance fulfillment. type redeemAction int @@ -272,7 +327,7 @@ func resolveRedeemAction(existing *RedeemCode, lookupErr error) redeemAction { return redeemActionRedeem } -func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) error { +func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease) error { // Idempotency: check if redeem code already exists (from a previous partial run) existing, lookupErr := s.redeemService.GetByCode(ctx, o.RechargeCode) action := resolveRedeemAction(existing, lookupErr) @@ -283,7 +338,7 @@ func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) e return err } // Code already created and redeemed — just mark completed - return s.markCompleted(ctx, o, "RECHARGE_SUCCESS") + return s.markCompleted(ctx, o, lease, "RECHARGE_SUCCESS") case redeemActionCreate: rc := &RedeemCode{Code: o.RechargeCode, Type: RedeemTypeBalance, Value: o.Amount, Status: StatusUnused} if err := s.redeemService.CreateCode(ctx, rc); err != nil { @@ -298,21 +353,37 @@ func (s *PaymentService) doBalance(ctx context.Context, o *dbent.PaymentOrder) e if err := s.applyAffiliateRebateForOrder(ctx, o); err != nil { return err } - return s.markCompleted(ctx, o, "RECHARGE_SUCCESS") + return s.markCompleted(ctx, o, lease, "RECHARGE_SUCCESS") } -func (s *PaymentService) markCompleted(ctx context.Context, o *dbent.PaymentOrder, auditAction string) error { +func (s *PaymentService) markCompleted(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease, auditAction string) error { + if lease == nil { + return errors.New("missing payment fulfillment lease") + } now := time.Now() - _, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(o.ID), paymentorder.StatusEQ(OrderStatusRecharging)).SetStatus(OrderStatusCompleted).SetCompletedAt(now).Save(ctx) + updated, err := s.entClient.PaymentOrder.Update().Where( + paymentorder.IDEQ(o.ID), + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtEQ(lease.version), + ).SetStatus(OrderStatusCompleted).SetCompletedAt(now).Save(ctx) if err != nil { return fmt.Errorf("mark completed: %w", err) } - s.writeAuditLog(ctx, o.ID, auditAction, "system", map[string]any{ - "rechargeCode": o.RechargeCode, - "creditedAmount": o.Amount, - "payAmount": o.PayAmount, - }) - s.dispatchPaymentFulfillmentNotification(o, auditAction) + if updated == 0 { + current, getErr := s.entClient.PaymentOrder.Get(ctx, o.ID) + if getErr == nil && current.Status == OrderStatusCompleted { + return nil + } + return infraerrors.Conflict("CONFLICT", "fulfillment lease was lost before completion") + } + if !s.hasAuditLog(ctx, o.ID, auditAction) { + s.writeAuditLog(ctx, o.ID, auditAction, "system", map[string]any{ + "rechargeCode": o.RechargeCode, + "creditedAmount": o.Amount, + "payAmount": o.PayAmount, + }) + s.dispatchPaymentFulfillmentNotification(o, auditAction) + } return nil } @@ -404,51 +475,138 @@ func (s *PaymentService) ExecuteSubscriptionFulfillment(ctx context.Context, oid if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot fulfill") } - if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed { + if o.Status != OrderStatusPaid && o.Status != OrderStatusFailed && o.Status != OrderStatusRecharging { return infraerrors.BadRequest("INVALID_STATUS", "order cannot fulfill in status "+o.Status) } if o.SubscriptionGroupID == nil || o.SubscriptionDays == nil { return infraerrors.BadRequest("INVALID_STATUS", "missing subscription info") } - c, err := s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusPaid, OrderStatusFailed)).SetStatus(OrderStatusRecharging).Save(ctx) + lease, err := s.acquirePaymentFulfillmentLease(ctx, o) if err != nil { - return fmt.Errorf("lock: %w", err) + return err } - if c == 0 { + if lease == nil { return nil } - if err := s.doSub(ctx, o); err != nil { - s.markFailed(ctx, oid, err) + if err := s.doSub(ctx, o, lease); err != nil { + s.markFailed(ctx, oid, lease, err) return err } return nil } -func (s *PaymentService) doSub(ctx context.Context, o *dbent.PaymentOrder) error { +func (s *PaymentService) doSub(ctx context.Context, o *dbent.PaymentOrder, lease *paymentFulfillmentLease) error { gid := *o.SubscriptionGroupID days := *o.SubscriptionDays g, err := s.groupRepo.GetByID(ctx, gid) if err != nil || g.Status != payment.EntityStatusActive { return fmt.Errorf("group %d no longer exists or inactive", gid) } - assigned := s.hasAuditLog(ctx, o.ID, "SUBSCRIPTION_ASSIGNED") || s.hasAuditLog(ctx, o.ID, "SUBSCRIPTION_SUCCESS") - if !assigned { - orderNote := fmt.Sprintf("payment order %d", o.ID) - _, _, err = s.subscriptionSvc.AssignOrExtendSubscription(ctx, &AssignSubscriptionInput{UserID: o.UserID, GroupID: gid, ValidityDays: days, AssignedBy: 0, Notes: orderNote}) - if err != nil { - return fmt.Errorf("assign subscription: %w", err) - } - s.writeAuditLog(ctx, o.ID, "SUBSCRIPTION_ASSIGNED", "system", map[string]any{ - "groupID": gid, - "validityDays": days, - }) - } else { - slog.Info("subscription already assigned for order, skipping", "orderID", o.ID, "groupID", gid) + if err := s.ensurePaymentSubscriptionAssigned(ctx, o, gid, days); err != nil { + return err } if err := s.applyAffiliateRebateForOrder(ctx, o); err != nil { return err } - return s.markCompleted(ctx, o, "SUBSCRIPTION_SUCCESS") + return s.markCompleted(ctx, o, lease, "SUBSCRIPTION_SUCCESS") +} + +func (s *PaymentService) ensurePaymentSubscriptionAssigned(ctx context.Context, o *dbent.PaymentOrder, groupID int64, days int) error { + if s.subscriptionSvc == nil { + return errors.New("subscription service is unavailable") + } + + tx, err := s.entClient.Tx(ctx) + if err != nil { + return fmt.Errorf("begin subscription fulfillment tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + + txCtx := dbent.NewTxContext(ctx, tx) + txClient := tx.Client() + alreadyAssigned, err := hasPaymentSubscriptionAssignmentAudit(txCtx, txClient, o.ID) + if err != nil { + return fmt.Errorf("check subscription assignment audit: %w", err) + } + + recoveredFromNote := false + if !alreadyAssigned { + orderNote := paymentSubscriptionOrderNote(o.ID) + existing, lookupErr := s.subscriptionSvc.userSubRepo.GetByUserIDAndGroupID(txCtx, o.UserID, groupID) + switch { + case lookupErr == nil && existing != nil && hasPaymentSubscriptionOrderNote(existing.Notes, orderNote): + recoveredFromNote = true + case lookupErr != nil && !errors.Is(lookupErr, ErrSubscriptionNotFound): + return fmt.Errorf("check existing subscription assignment: %w", lookupErr) + default: + if _, _, err := s.subscriptionSvc.assignOrExtendSubscription(txCtx, &AssignSubscriptionInput{ + UserID: o.UserID, + GroupID: groupID, + ValidityDays: days, + AssignedBy: 0, + Notes: orderNote, + }, true); err != nil { + return fmt.Errorf("assign subscription: %w", err) + } + } + + detail, _ := json.Marshal(map[string]any{ + "groupID": groupID, + "validityDays": days, + "recoveredFromNote": recoveredFromNote, + }) + if _, err := txClient.PaymentAuditLog.Create(). + SetOrderID(strconv.FormatInt(o.ID, 10)). + SetAction("SUBSCRIPTION_ASSIGNED"). + SetDetail(string(detail)). + SetOperator("system"). + Save(txCtx); err != nil { + if dbent.IsConstraintError(err) { + _ = tx.Rollback() + claimed, checkErr := hasPaymentSubscriptionAssignmentAudit(ctx, s.entClient, o.ID) + if checkErr == nil && claimed { + return s.subscriptionSvc.invalidateSubscriptionCaches(o.UserID, groupID) + } + } + return fmt.Errorf("record subscription assignment audit: %w", err) + } + } else { + slog.Info("subscription already assigned for order, skipping", "orderID", o.ID, "groupID", groupID) + } + + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit subscription fulfillment tx: %w", err) + } + // Assignment cache invalidation is deferred while this transaction is open, + // then performed synchronously against the committed subscription. + if err := s.subscriptionSvc.invalidateSubscriptionCaches(o.UserID, groupID); err != nil { + return fmt.Errorf("invalidate subscription cache after fulfillment: %w", err) + } + return nil +} + +func hasPaymentSubscriptionAssignmentAudit(ctx context.Context, client *dbent.Client, orderID int64) (bool, error) { + count, err := client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(orderID, 10)), + paymentauditlog.ActionIn("SUBSCRIPTION_ASSIGNED", "SUBSCRIPTION_SUCCESS"), + ). + Limit(1). + Count(ctx) + return count > 0, err +} + +func paymentSubscriptionOrderNote(orderID int64) string { + return fmt.Sprintf("payment order %d", orderID) +} + +func hasPaymentSubscriptionOrderNote(notes string, orderNote string) bool { + for _, line := range strings.Split(strings.ReplaceAll(notes, "\r\n", "\n"), "\n") { + if strings.TrimSpace(line) == orderNote { + return true + } + } + return false } func (s *PaymentService) hasAuditLog(ctx context.Context, orderID int64, action string) bool { @@ -642,13 +800,20 @@ func (s *PaymentService) updateClaimedAffiliateRebateAudit(ctx context.Context, return nil } -func (s *PaymentService) markFailed(ctx context.Context, oid int64, cause error) { +func (s *PaymentService) markFailed(ctx context.Context, oid int64, lease *paymentFulfillmentLease, cause error) { + if lease == nil { + slog.Error("mark FAILED without fulfillment lease", "orderID", oid) + return + } now := time.Now() r := psErrMsg(cause) - // Only mark FAILED if still in RECHARGING state — prevents overwriting - // a COMPLETED order when markCompleted failed but fulfillment succeeded. + // The lease version prevents a stale worker from overwriting a newer owner. c, e := s.entClient.PaymentOrder.Update(). - Where(paymentorder.IDEQ(oid), paymentorder.StatusEQ(OrderStatusRecharging)). + Where( + paymentorder.IDEQ(oid), + paymentorder.StatusEQ(OrderStatusRecharging), + paymentorder.UpdatedAtEQ(lease.version), + ). SetStatus(OrderStatusFailed).SetFailedAt(now).SetFailedReason(r).Save(ctx) if e != nil { slog.Error("mark FAILED", "orderID", oid, "error", e) @@ -669,18 +834,11 @@ func (s *PaymentService) RetryFulfillment(ctx context.Context, oid int64) error if psIsRefundStatus(o.Status) { return infraerrors.BadRequest("INVALID_STATUS", "refund-related order cannot retry") } - if o.Status == OrderStatusRecharging { - return infraerrors.Conflict("CONFLICT", "order is being processed") - } if o.Status == OrderStatusCompleted { return infraerrors.BadRequest("INVALID_STATUS", "order already completed") } - if o.Status != OrderStatusFailed && o.Status != OrderStatusPaid { - return infraerrors.BadRequest("INVALID_STATUS", "only paid and failed orders can retry") - } - _, err = s.entClient.PaymentOrder.Update().Where(paymentorder.IDEQ(oid), paymentorder.StatusIn(OrderStatusFailed, OrderStatusPaid)).SetStatus(OrderStatusPaid).ClearFailedAt().ClearFailedReason().Save(ctx) - if err != nil { - return fmt.Errorf("reset for retry: %w", err) + if o.Status != OrderStatusFailed && o.Status != OrderStatusPaid && o.Status != OrderStatusRecharging { + return infraerrors.BadRequest("INVALID_STATUS", "only paid, failed, and recoverable recharging orders can retry") } s.writeAuditLog(ctx, oid, "RECHARGE_RETRY", "admin", map[string]any{"detail": "admin manual retry"}) return s.executeFulfillment(ctx, oid) diff --git a/backend/internal/service/payment_fulfillment_test.go b/backend/internal/service/payment_fulfillment_test.go index a8c78d713c..d040095d63 100644 --- a/backend/internal/service/payment_fulfillment_test.go +++ b/backend/internal/service/payment_fulfillment_test.go @@ -13,6 +13,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/paymentauditlog" "github.com/Wei-Shaw/sub2api/internal/payment" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -586,6 +587,238 @@ func TestPaymentAmountToleranceForThreeDecimalCurrency(t *testing.T) { assert.InDelta(t, 0.0005, paymentAmountToleranceForCurrency("KWD"), 1e-12) } +func TestRetryFulfillmentRejectsFreshRechargingLease(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, time.Now()) + + svc := &PaymentService{entClient: client} + err := svc.RetryFulfillment(ctx, order.ID) + require.Error(t, err) + require.Equal(t, "CONFLICT", infraerrors.Reason(err)) + + reloaded, getErr := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, getErr) + require.Equal(t, OrderStatusRecharging, reloaded.Status) +} + +func TestAlreadyProcessedRecoversStaleRechargingLease(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + order := createPaymentFulfillmentSubscriptionOrder( + t, + ctx, + client, + OrderStatusRecharging, + time.Now().Add(-paymentFulfillmentLeaseDuration-time.Minute), + ) + _, err := client.PaymentAuditLog.Create(). + SetOrderID(strconv.FormatInt(order.ID, 10)). + SetAction("SUBSCRIPTION_ASSIGNED"). + SetDetail(`{"groupID":7,"validityDays":30}`). + SetOperator("system"). + Save(ctx) + require.NoError(t, err) + + groupRepo := &subscriptionGroupRepoStub{ + group: &Group{ID: 7, Status: payment.EntityStatusActive, SubscriptionType: SubscriptionTypeSubscription}, + } + svc := &PaymentService{ + entClient: client, + groupRepo: groupRepo, + subscriptionSvc: NewSubscriptionService(groupRepo, userSubRepoNoop{}, nil, nil, nil), + } + + require.NoError(t, svc.alreadyProcessed(ctx, order)) + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusCompleted, reloaded.Status) +} + +func TestFulfillmentLeaseVersionRejectsStaleWorker(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + svc := &PaymentService{entClient: client} + + firstLease, err := svc.acquirePaymentFulfillmentLease(ctx, order) + require.NoError(t, err) + require.NotNil(t, firstLease) + + _, err = client.PaymentOrder.UpdateOneID(order.ID).SetUpdatedAt(staleAt).Save(ctx) + require.NoError(t, err) + time.Sleep(time.Millisecond) + staleOrder, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + secondLease, err := svc.acquirePaymentFulfillmentLease(ctx, staleOrder) + require.NoError(t, err) + require.NotNil(t, secondLease) + require.False(t, firstLease.version.Equal(secondLease.version)) + + err = svc.markCompleted(ctx, order, firstLease, "SUBSCRIPTION_SUCCESS") + require.Error(t, err) + require.Equal(t, "CONFLICT", infraerrors.Reason(err)) + svc.markFailed(ctx, order.ID, firstLease, errors.New("stale worker failure")) + + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusRecharging, reloaded.Status) + require.NoError(t, svc.markCompleted(ctx, order, secondLease, "SUBSCRIPTION_SUCCESS")) +} + +func TestExecuteBalanceFulfillmentRecoversAfterRedeemWithoutCreditingAgain(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + order, err := client.PaymentOrder.UpdateOneID(order.ID). + SetOrderType(payment.OrderTypeBalance). + ClearPlanID(). + ClearSubscriptionGroupID(). + ClearSubscriptionDays(). + SetUpdatedAt(staleAt). + Save(ctx) + require.NoError(t, err) + + redeemRepo := &redeemCodeRepoStub{codesByCode: map[string]*RedeemCode{ + order.RechargeCode: { + ID: 101, + Code: order.RechargeCode, + Type: RedeemTypeBalance, + Value: order.Amount, + Status: StatusUsed, + }, + }} + svc := &PaymentService{ + entClient: client, + redeemService: &RedeemService{redeemRepo: redeemRepo}, + } + + require.NoError(t, svc.ExecuteBalanceFulfillment(ctx, order.ID)) + require.Empty(t, redeemRepo.useCalls, "an already-used order code must not be redeemed again") + reloaded, err := client.PaymentOrder.Get(ctx, order.ID) + require.NoError(t, err) + require.Equal(t, OrderStatusCompleted, reloaded.Status) +} + +func TestExecuteSubscriptionFulfillmentRecoversCommittedAssignmentWithoutExtendingAgain(t *testing.T) { + ctx := context.Background() + client := newPaymentConfigServiceTestClient(t) + ensurePaymentAuditOrderActionUniqueIndex(t, ctx, client) + staleAt := time.Now().Add(-paymentFulfillmentLeaseDuration - time.Minute) + order := createPaymentFulfillmentSubscriptionOrder(t, ctx, client, OrderStatusRecharging, staleAt) + + expiresAt := time.Now().Add(30 * 24 * time.Hour).Truncate(time.Second) + subRepo := newSubscriptionUserSubRepoStub() + subRepo.seed(&UserSubscription{ + ID: 99, + UserID: order.UserID, + GroupID: *order.SubscriptionGroupID, + StartsAt: time.Now().Add(-time.Hour), + ExpiresAt: expiresAt, + Status: SubscriptionStatusActive, + Notes: "manual note\n" + paymentSubscriptionOrderNote(order.ID) + "\nretained note", + }) + groupRepo := &subscriptionGroupRepoStub{ + group: &Group{ID: 7, Status: payment.EntityStatusActive, SubscriptionType: SubscriptionTypeSubscription}, + } + svc := &PaymentService{ + entClient: client, + groupRepo: groupRepo, + subscriptionSvc: NewSubscriptionService(groupRepo, subRepo, nil, nil, nil), + } + + require.NoError(t, svc.ExecuteSubscriptionFulfillment(ctx, order.ID)) + assertPaymentSubscriptionExpiry(t, subRepo, order, expiresAt) + + assignmentAuditCount, err := client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), + paymentauditlog.ActionEQ("SUBSCRIPTION_ASSIGNED"), + ). + Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, assignmentAuditCount) + + // Simulate another stale recovery attempt after completion. The durable audit + // must make replay a no-op for the subscription entitlement. + _, err = client.PaymentOrder.UpdateOneID(order.ID). + SetStatus(OrderStatusRecharging). + SetUpdatedAt(staleAt). + ClearCompletedAt(). + Save(ctx) + require.NoError(t, err) + require.NoError(t, svc.ExecuteSubscriptionFulfillment(ctx, order.ID)) + assertPaymentSubscriptionExpiry(t, subRepo, order, expiresAt) + + assignmentAuditCount, err = client.PaymentAuditLog.Query(). + Where( + paymentauditlog.OrderIDEQ(strconv.FormatInt(order.ID, 10)), + paymentauditlog.ActionEQ("SUBSCRIPTION_ASSIGNED"), + ). + Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, assignmentAuditCount) +} + +func TestHasPaymentSubscriptionOrderNoteRequiresIndependentExactLine(t *testing.T) { + t.Parallel() + require.True(t, hasPaymentSubscriptionOrderNote("before\r\npayment order 42\r\nafter", "payment order 42")) + require.False(t, hasPaymentSubscriptionOrderNote("payment order 420", "payment order 42")) + require.False(t, hasPaymentSubscriptionOrderNote("prefix payment order 42 suffix", "payment order 42")) +} + +func createPaymentFulfillmentSubscriptionOrder( + t *testing.T, + ctx context.Context, + client *dbent.Client, + status string, + updatedAt time.Time, +) *dbent.PaymentOrder { + t.Helper() + user, err := client.User.Create(). + SetEmail("fulfillment-" + strconv.FormatInt(time.Now().UnixNano(), 10) + "@example.com"). + SetPasswordHash("hash"). + SetUsername("payment-fulfillment-user"). + Save(ctx) + require.NoError(t, err) + + order, err := client.PaymentOrder.Create(). + SetUserID(user.ID). + SetUserEmail(user.Email). + SetUserName(user.Username). + SetAmount(80). + SetPayAmount(80). + SetFeeRate(0). + SetRechargeCode("PAY-SUB-" + strconv.FormatInt(time.Now().UnixNano(), 10)). + SetOutTradeNo("sub2_fulfillment_" + strconv.FormatInt(time.Now().UnixNano(), 10)). + SetPaymentType(payment.TypeAlipay). + SetPaymentTradeNo("trade-fulfillment"). + SetOrderType(payment.OrderTypeSubscription). + SetPlanID(100). + SetSubscriptionGroupID(7). + SetSubscriptionDays(30). + SetStatus(status). + SetPaidAt(time.Now().Add(-time.Hour)). + SetExpiresAt(time.Now().Add(time.Hour)). + SetClientIP("127.0.0.1"). + SetSrcHost("api.example.com"). + SetUpdatedAt(updatedAt). + Save(ctx) + require.NoError(t, err) + return order +} + +func assertPaymentSubscriptionExpiry(t *testing.T, repo *subscriptionUserSubRepoStub, order *dbent.PaymentOrder, expected time.Time) { + t.Helper() + sub, err := repo.GetByUserIDAndGroupID(context.Background(), order.UserID, *order.SubscriptionGroupID) + require.NoError(t, err) + require.True(t, sub.ExpiresAt.Equal(expected), "subscription expiry changed from %s to %s", expected, sub.ExpiresAt) +} + func TestExecuteSubscriptionFulfillmentAppliesAffiliateRebate(t *testing.T) { ctx := context.Background() client := newPaymentConfigServiceTestClient(t) diff --git a/backend/internal/service/redeem_service.go b/backend/internal/service/redeem_service.go index 2d1962dd3c..8794872e3d 100644 --- a/backend/internal/service/redeem_service.go +++ b/backend/internal/service/redeem_service.go @@ -135,6 +135,7 @@ type RedeemCodeBatchUpdateResult struct { type RedeemService struct { redeemRepo RedeemCodeRepository userRepo UserRepository + redeemUserRepo RedeemUserAdjustmentRepository subscriptionService *SubscriptionService cache RedeemCache billingCacheService *BillingCacheService @@ -154,9 +155,11 @@ func NewRedeemService( authCacheInvalidator APIKeyAuthCacheInvalidator, affiliateService *AffiliateService, ) *RedeemService { + redeemUserRepo, _ := userRepo.(RedeemUserAdjustmentRepository) return &RedeemService{ redeemRepo: redeemRepo, userRepo: userRepo, + redeemUserRepo: redeemUserRepo, subscriptionService: subscriptionService, cache: cache, billingCacheService: billingCacheService, @@ -426,7 +429,7 @@ func (s *RedeemService) Redeem(ctx context.Context, userID int64, code string) ( } // 获取用户信息 - user, err := s.userRepo.GetByID(ctx, userID) + _, err = s.userRepo.GetByID(ctx, userID) if err != nil { return nil, fmt.Errorf("get user: %w", err) } @@ -454,21 +457,27 @@ func (s *RedeemService) Redeem(ctx context.Context, userID int64, code string) ( switch redeemCode.Type { case RedeemTypeBalance: amount := redeemCode.Value - // 负数为退款扣减,余额最低为 0 - if amount < 0 && user.Balance+amount < 0 { - amount = -user.Balance - } - if err := s.userRepo.UpdateBalance(txCtx, userID, amount); err != nil { + if amount < 0 { + if s.redeemUserRepo == nil { + return nil, errors.New("user repository does not support atomic redeem balance adjustments") + } + if err := s.redeemUserRepo.ApplyRedeemBalanceAdjustment(txCtx, userID, amount); err != nil { + return nil, fmt.Errorf("update user balance: %w", err) + } + } else if err := s.userRepo.UpdateBalance(txCtx, userID, amount); err != nil { return nil, fmt.Errorf("update user balance: %w", err) } case RedeemTypeConcurrency: delta := int(redeemCode.Value) - // 负数为退款扣减,并发数最低为 0 - if delta < 0 && user.Concurrency+delta < 0 { - delta = -user.Concurrency - } - if err := s.userRepo.UpdateConcurrency(txCtx, userID, delta); err != nil { + if delta < 0 { + if s.redeemUserRepo == nil { + return nil, errors.New("user repository does not support atomic redeem concurrency adjustments") + } + if err := s.redeemUserRepo.ApplyRedeemConcurrencyAdjustment(txCtx, userID, delta); err != nil { + return nil, fmt.Errorf("update user concurrency: %w", err) + } + } else if err := s.userRepo.UpdateConcurrency(txCtx, userID, delta); err != nil { return nil, fmt.Errorf("update user concurrency: %w", err) } diff --git a/backend/internal/service/subscription_assign_idempotency_test.go b/backend/internal/service/subscription_assign_idempotency_test.go index 8e249af52b..d4913f8994 100644 --- a/backend/internal/service/subscription_assign_idempotency_test.go +++ b/backend/internal/service/subscription_assign_idempotency_test.go @@ -6,11 +6,49 @@ import ( "testing" "time" + dbent "github.com/Wei-Shaw/sub2api/ent" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/dgraph-io/ristretto" "github.com/stretchr/testify/require" ) +func TestWithSubscriptionUpdateTx_ReusesExistingTransaction(t *testing.T) { + existingTx := &dbent.Tx{} + ctx := dbent.NewTxContext(context.Background(), existingTx) + svc := &SubscriptionService{entClient: &dbent.Client{}} + + called := false + err := svc.withSubscriptionUpdateTx(ctx, func(txCtx context.Context) error { + called = true + require.Same(t, existingTx, dbent.TxFromContext(txCtx)) + return nil + }) + + require.NoError(t, err) + require.True(t, called) +} + +func TestMaybeInvalidateAssignmentCaches_DefersForOuterTransactionOwner(t *testing.T) { + cache, err := ristretto.NewCache(&ristretto.Config{NumCounters: 1_000, MaxCost: 100, BufferItems: 64}) + require.NoError(t, err) + t.Cleanup(cache.Close) + + svc := &SubscriptionService{subCacheL1: cache} + key := subCacheKey(7, 9) + require.True(t, cache.Set(key, &UserSubscription{ID: 42}, 1)) + cache.Wait() + + svc.maybeInvalidateAssignmentCaches(7, 9, true) + _, cachedBeforeCommit := cache.Get(key) + require.True(t, cachedBeforeCommit, "outer transaction must retain caches until its owner commits") + + svc.maybeInvalidateAssignmentCaches(7, 9, false) + cache.Wait() + _, cachedAfterCommit := cache.Get(key) + require.False(t, cachedAfterCommit, "post-commit invalidation must remove the cached subscription") +} + type groupRepoNoop struct{} func (groupRepoNoop) Create(context.Context, *Group) error { panic("unexpected Create call") } @@ -119,13 +157,16 @@ func (userSubRepoNoop) UpdateNotes(context.Context, int64, string) error { func (userSubRepoNoop) ActivateWindows(context.Context, int64, time.Time) error { panic("unexpected ActivateWindows call") } -func (userSubRepoNoop) ResetDailyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { + panic("unexpected ResetUsageWindows call") +} +func (userSubRepoNoop) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetDailyUsage call") } -func (userSubRepoNoop) ResetWeeklyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetWeeklyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetWeeklyUsage call") } -func (userSubRepoNoop) ResetMonthlyUsage(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ResetMonthlyUsage(context.Context, int64, *time.Time, time.Time) error { panic("unexpected ResetMonthlyUsage call") } func (userSubRepoNoop) IncrementUsage(context.Context, int64, float64) error { diff --git a/backend/internal/service/subscription_expiry_service_test.go b/backend/internal/service/subscription_expiry_service_test.go index 056315a289..7db642c076 100644 --- a/backend/internal/service/subscription_expiry_service_test.go +++ b/backend/internal/service/subscription_expiry_service_test.go @@ -87,15 +87,19 @@ func (r *subscriptionExpiryRepoStub) ActivateWindows(context.Context, int64, tim return nil } -func (r *subscriptionExpiryRepoStub) ResetDailyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { return nil } -func (r *subscriptionExpiryRepoStub) ResetWeeklyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { return nil } -func (r *subscriptionExpiryRepoStub) ResetMonthlyUsage(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetWeeklyUsage(context.Context, int64, *time.Time, time.Time) error { + return nil +} + +func (r *subscriptionExpiryRepoStub) ResetMonthlyUsage(context.Context, int64, *time.Time, time.Time) error { return nil } diff --git a/backend/internal/service/subscription_reset_quota_test.go b/backend/internal/service/subscription_reset_quota_test.go index 3bbc217073..e4ed45ec45 100644 --- a/backend/internal/service/subscription_reset_quota_test.go +++ b/backend/internal/service/subscription_reset_quota_test.go @@ -11,7 +11,7 @@ import ( "github.com/stretchr/testify/require" ) -// resetQuotaUserSubRepoStub 支持 GetByID、ResetDailyUsage、ResetWeeklyUsage、ResetMonthlyUsage, +// resetQuotaUserSubRepoStub 支持 GetByID、ResetUsageWindows, // 其余方法继承 userSubRepoNoop(panic)。 type resetQuotaUserSubRepoStub struct { userSubRepoNoop @@ -34,7 +34,38 @@ func (r *resetQuotaUserSubRepoStub) GetByID(_ context.Context, id int64) (*UserS return &cp, nil } -func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, windowStart time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetUsageWindows(_ context.Context, _ int64, resetDaily, resetWeekly, resetMonthly bool, windowStart time.Time) error { + r.resetDailyCalled = resetDaily + r.resetWeeklyCalled = resetWeekly + r.resetMonthlyCalled = resetMonthly + if resetDaily && r.resetDailyErr != nil { + return r.resetDailyErr + } + if resetWeekly && r.resetWeeklyErr != nil { + return r.resetWeeklyErr + } + if resetMonthly && r.resetMonthlyErr != nil { + return r.resetMonthlyErr + } + if r.sub == nil { + return nil + } + if resetDaily { + r.sub.DailyUsageUSD = 0 + r.sub.DailyWindowStart = &windowStart + } + if resetWeekly { + r.sub.WeeklyUsageUSD = 0 + r.sub.WeeklyWindowStart = &windowStart + } + if resetMonthly { + r.sub.MonthlyUsageUSD = 0 + r.sub.MonthlyWindowStart = &windowStart + } + return nil +} + +func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, _ *time.Time, windowStart time.Time) error { r.resetDailyCalled = true if r.resetDailyErr == nil && r.sub != nil { r.sub.DailyUsageUSD = 0 @@ -43,12 +74,12 @@ func (r *resetQuotaUserSubRepoStub) ResetDailyUsage(_ context.Context, _ int64, return r.resetDailyErr } -func (r *resetQuotaUserSubRepoStub) ResetWeeklyUsage(_ context.Context, _ int64, _ time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetWeeklyUsage(_ context.Context, _ int64, _ *time.Time, _ time.Time) error { r.resetWeeklyCalled = true return r.resetWeeklyErr } -func (r *resetQuotaUserSubRepoStub) ResetMonthlyUsage(_ context.Context, _ int64, _ time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetMonthlyUsage(_ context.Context, _ int64, _ *time.Time, _ time.Time) error { r.resetMonthlyCalled = true return r.resetMonthlyErr } @@ -140,7 +171,7 @@ func TestAdminResetQuota_ResetDailyUsageError(t *testing.T) { require.ErrorIs(t, err, dbErr) require.True(t, stub.resetDailyCalled) - require.False(t, stub.resetWeeklyCalled, "daily 失败后不应继续调用 weekly") + require.True(t, stub.resetWeeklyCalled, "原子重置应在一次调用中提交所选窗口") } func TestAdminResetQuota_ResetWeeklyUsageError(t *testing.T) { @@ -200,7 +231,7 @@ func TestAdminResetQuota_ReturnsRefreshedSub(t *testing.T) { result, err := svc.AdminResetQuota(context.Background(), 6, true, false, false) require.NoError(t, err) - // ResetDailyUsage stub 会将 sub.DailyUsageUSD 归零, + // ResetUsageWindows stub 会将 sub.DailyUsageUSD 归零, // 服务应返回第二次 GetByID 的刷新值而非初始的 99.9 require.Equal(t, float64(0), result.DailyUsageUSD, "返回的订阅应反映已归零的用量") require.True(t, stub.resetDailyCalled) diff --git a/backend/internal/service/subscription_service.go b/backend/internal/service/subscription_service.go index 0a4fc7b757..ea1fd091d9 100644 --- a/backend/internal/service/subscription_service.go +++ b/backend/internal/service/subscription_service.go @@ -212,6 +212,10 @@ func (s *SubscriptionService) AssignSubscription(ctx context.Context, input *Ass // // 如果没有订阅:创建新订阅 func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, input *AssignSubscriptionInput) (*UserSubscription, bool, error) { + return s.assignOrExtendSubscription(ctx, input, false) +} + +func (s *SubscriptionService) assignOrExtendSubscription(ctx context.Context, input *AssignSubscriptionInput, deferCacheInvalidation bool) (*UserSubscription, bool, error) { // 检查分组是否存在且为订阅类型 group, err := s.groupRepo.GetByID(ctx, input.GroupID) if err != nil { @@ -260,15 +264,7 @@ func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, in } // 失效订阅缓存 - s.InvalidateSubCache(input.UserID, input.GroupID) - if s.billingCacheService != nil { - userID, groupID := input.UserID, input.GroupID - go func() { - cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - _ = s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID) - }() - } + s.maybeInvalidateAssignmentCaches(input.UserID, input.GroupID, deferCacheInvalidation) // 返回更新后的订阅 sub, err := s.userSubRepo.GetByID(ctx, existingSub.ID) @@ -282,17 +278,27 @@ func (s *SubscriptionService) AssignOrExtendSubscription(ctx context.Context, in } // 失效订阅缓存 - s.InvalidateSubCache(input.UserID, input.GroupID) + s.maybeInvalidateAssignmentCaches(input.UserID, input.GroupID, deferCacheInvalidation) + + return sub, false, nil // false 表示是新建 +} + +func (s *SubscriptionService) maybeInvalidateAssignmentCaches(userID, groupID int64, deferred bool) { + // Payment fulfillment owns an outer transaction and performs a synchronous + // invalidation after commit. Invalidating inside that transaction can reload + // the pre-commit subscription into cache. + if deferred { + return + } + + s.InvalidateSubCache(userID, groupID) if s.billingCacheService != nil { - userID, groupID := input.UserID, input.GroupID go func() { cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _ = s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID) }() } - - return sub, false, nil // false 表示是新建 } func (s *SubscriptionService) updateExistingSubscriptionTerm( @@ -336,6 +342,9 @@ func (s *SubscriptionService) updateExistingSubscriptionTerm( } func (s *SubscriptionService) withSubscriptionUpdateTx(ctx context.Context, fn func(context.Context) error) error { + if dbent.TxFromContext(ctx) != nil { + return fn(ctx) + } if s.entClient == nil { return fn(ctx) } @@ -834,20 +843,8 @@ func (s *SubscriptionService) AdminResetQuota(ctx context.Context, subscriptionI return nil, err } windowStart := startOfDay(time.Now()) - if resetDaily { - if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } - } - if resetWeekly { - if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } - } - if resetMonthly { - if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, windowStart); err != nil { - return nil, err - } + if err := s.userSubRepo.ResetUsageWindows(ctx, sub.ID, resetDaily, resetWeekly, resetMonthly, windowStart); err != nil { + return nil, err } // Invalidate L1 ristretto cache. Ristretto's Del() is asynchronous by design, // so call Wait() immediately after to flush pending operations and guarantee @@ -868,7 +865,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 日窗口重置(24小时) if sub.NeedsDailyReset() { - if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.DailyWindowStart + if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.DailyWindowStart = &windowStart @@ -878,7 +876,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 周窗口重置(7天) if sub.NeedsWeeklyReset() { - if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.WeeklyWindowStart + if err := s.userSubRepo.ResetWeeklyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.WeeklyWindowStart = &windowStart @@ -888,7 +887,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use // 月窗口重置(30天) if sub.NeedsMonthlyReset() { - if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, windowStart); err != nil { + expectedWindowStart := sub.MonthlyWindowStart + if err := s.userSubRepo.ResetMonthlyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err } sub.MonthlyWindowStart = &windowStart @@ -907,6 +907,32 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use return nil } +// EnsureWindowMaintenance advances expired usage windows before a request is +// allowed to proceed. It returns a fresh database snapshot because a competing +// request may have won one of the conditional resets. +func (s *SubscriptionService) EnsureWindowMaintenance(ctx context.Context, sub *UserSubscription) (*UserSubscription, error) { + if sub == nil { + return nil, ErrSubscriptionNilInput + } + if !sub.IsWindowActivated() { + if err := s.CheckAndActivateWindow(ctx, sub); err != nil { + return nil, err + } + } + if err := s.CheckAndResetWindows(ctx, sub); err != nil { + return nil, err + } + + // GetByID bypasses the service caches. This prevents a stale loser of the + // CAS from validating limits against zeroed in-memory usage. + refreshed, err := s.userSubRepo.GetByID(ctx, sub.ID) + if err != nil { + return nil, err + } + s.InvalidateSubCacheSync(sub.UserID, sub.GroupID) + return refreshed, nil +} + // CheckUsageLimits 检查使用限额(返回错误如果超限) // 用于中间件的快速预检查,additionalCost 通常为 0 func (s *SubscriptionService) CheckUsageLimits(ctx context.Context, sub *UserSubscription, group *Group, additionalCost float64) error { @@ -923,8 +949,8 @@ func (s *SubscriptionService) CheckUsageLimits(ctx context.Context, sub *UserSub } // ValidateAndCheckLimits 合并验证+限额检查(中间件热路径专用) -// 仅做内存检查,不触发 DB 写入。窗口重置的 DB 写入由 DoWindowMaintenance 异步完成。 -// 返回 needsMaintenance 表示是否需要异步执行窗口维护。 +// 仅做内存检查,不触发 DB 写入。调用方必须在放行请求前同步完成窗口维护。 +// 返回 needsMaintenance 表示是否需要执行窗口维护并回读数据库快照。 func (s *SubscriptionService) ValidateAndCheckLimits(sub *UserSubscription, group *Group) (needsMaintenance bool, err error) { // 1. 验证订阅状态 if sub.Status == SubscriptionStatusExpired { @@ -937,8 +963,8 @@ func (s *SubscriptionService) ValidateAndCheckLimits(sub *UserSubscription, grou return false, ErrSubscriptionExpired } - // 2. 内存中修正过期窗口的用量,确保 CheckUsageLimits 不会误拒绝用户 - // 实际的 DB 窗口重置由 DoWindowMaintenance 异步完成 + // 2. 内存中修正过期窗口的用量,确保预检查不会误拒绝用户。 + // 调用方随后同步推进 DB 窗口,并用回读快照重新校验。 if sub.NeedsDailyReset() { sub.DailyUsageUSD = 0 needsMaintenance = true diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index 2c87221401..98b0c8f32b 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -122,6 +122,14 @@ type UserRepository interface { DisableTotp(ctx context.Context, userID int64) error } +// RedeemUserAdjustmentRepository provides the atomic, floor-at-zero updates +// used by negative-value redeem codes. It is intentionally narrower than +// UserRepository because normal usage billing is allowed to overdraw. +type RedeemUserAdjustmentRepository interface { + ApplyRedeemBalanceAdjustment(ctx context.Context, id int64, delta float64) error + ApplyRedeemConcurrencyAdjustment(ctx context.Context, id int64, delta int) error +} + type UserAuthIdentityRecord struct { ProviderType string ProviderKey string diff --git a/backend/internal/service/user_subscription_daily_quota_test.go b/backend/internal/service/user_subscription_daily_quota_test.go index 3738bdd698..bf58de7f4c 100644 --- a/backend/internal/service/user_subscription_daily_quota_test.go +++ b/backend/internal/service/user_subscription_daily_quota_test.go @@ -15,7 +15,7 @@ type dailyResetTrackingUserSubRepo struct { resetDailyCalled bool } -func (r *dailyResetTrackingUserSubRepo) ResetDailyUsage(context.Context, int64, time.Time) error { +func (r *dailyResetTrackingUserSubRepo) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { r.resetDailyCalled = true return nil } diff --git a/backend/internal/service/user_subscription_port.go b/backend/internal/service/user_subscription_port.go index 43d41d6dd7..eeee0275f0 100644 --- a/backend/internal/service/user_subscription_port.go +++ b/backend/internal/service/user_subscription_port.go @@ -29,9 +29,10 @@ type UserSubscriptionRepository interface { UpdateNotes(ctx context.Context, subscriptionID int64, notes string) error ActivateWindows(ctx context.Context, id int64, start time.Time) error - ResetDailyUsage(ctx context.Context, id int64, newWindowStart time.Time) error - ResetWeeklyUsage(ctx context.Context, id int64, newWindowStart time.Time) error - ResetMonthlyUsage(ctx context.Context, id int64, newWindowStart time.Time) error + ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error + ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error + ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error + ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error IncrementUsage(ctx context.Context, id int64, costUSD float64) error BatchUpdateExpiredStatus(ctx context.Context) (int64, error) diff --git a/frontend/src/router/__tests__/feature-access.spec.ts b/frontend/src/router/__tests__/feature-access.spec.ts new file mode 100644 index 0000000000..3c98425f90 --- /dev/null +++ b/frontend/src/router/__tests__/feature-access.spec.ts @@ -0,0 +1,177 @@ +import { beforeAll, beforeEach, describe, expect, it, vi } from 'vitest' + +type NavigationGuard = ( + to: Record, + from: Record, + next: ReturnType +) => Promise + +const routerHarness = vi.hoisted(() => ({ + guard: null as NavigationGuard | null, +})) + +const authStore = vi.hoisted(() => ({ + checkAuth: vi.fn(), + isAuthenticated: true, + isAdmin: false, + isSimpleMode: false, + hasPendingAuthSession: false, +})) + +const appStore = vi.hoisted(() => ({ + siteName: 'Sub2API', + backendModeEnabled: false, + publicSettingsLoaded: false, + cachedPublicSettings: null as null | { + payment_enabled?: boolean + risk_control_enabled?: boolean + custom_menu_items?: [] + }, + fetchPublicSettings: vi.fn(), +})) + +vi.mock('vue-router', () => ({ + createWebHistory: vi.fn(() => ({})), + createRouter: vi.fn(() => ({ + beforeEach: vi.fn((guard: NavigationGuard) => { + routerHarness.guard = guard + }), + afterEach: vi.fn(), + onError: vi.fn(), + })), +})) + +vi.mock('@/stores/auth', () => ({ + useAuthStore: () => authStore, +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => appStore, +})) + +vi.mock('@/stores/adminSettings', () => ({ + useAdminSettingsStore: () => ({ customMenuItems: [] }), +})) + +vi.mock('@/stores/adminCompliance', () => ({ + useAdminComplianceStore: () => ({ + initialized: true, + fetchStatus: vi.fn(), + requireAcknowledgement: vi.fn(), + }), +})) + +vi.mock('@/composables/useNavigationLoading', () => ({ + useNavigationLoadingState: () => ({ + startNavigation: vi.fn(), + endNavigation: vi.fn(), + isLoading: { value: false }, + }), +})) + +vi.mock('@/composables/useRoutePrefetch', () => ({ + useRoutePrefetch: () => ({ + triggerPrefetch: vi.fn(), + cancelPendingPrefetch: vi.fn(), + resetPrefetchState: vi.fn(), + }), +})) + +function createDeferred() { + let resolve!: (value: T | PromiseLike) => void + const promise = new Promise((resolvePromise) => { + resolve = resolvePromise + }) + return { promise, resolve } +} + +function runGuard(meta: Record, path: string) { + if (!routerHarness.guard) { + throw new Error('router guard was not registered') + } + + const next = vi.fn() + const navigation = routerHarness.guard( + { + path, + fullPath: path, + name: 'FeatureRoute', + params: {}, + meta: { requiresAuth: true, ...meta }, + }, + {}, + next + ) + return { navigation, next } +} + +describe('feature route guard', () => { + beforeAll(async () => { + await import('@/router') + }) + + beforeEach(() => { + authStore.isAuthenticated = true + authStore.isAdmin = false + authStore.isSimpleMode = false + appStore.publicSettingsLoaded = false + appStore.cachedPublicSettings = null + appStore.fetchPublicSettings.mockReset() + }) + + it('waits for the first public-settings request before deciding payment access', async () => { + const deferred = createDeferred<{ payment_enabled: boolean }>() + appStore.fetchPublicSettings.mockImplementation(async () => { + const settings = await deferred.promise + appStore.cachedPublicSettings = settings + appStore.publicSettingsLoaded = true + return settings + }) + + const { navigation, next } = runGuard({ requiresPayment: true }, '/purchase') + + await vi.waitFor(() => expect(appStore.fetchPublicSettings).toHaveBeenCalledTimes(1)) + expect(next).not.toHaveBeenCalled() + + deferred.resolve({ payment_enabled: true }) + await navigation + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith() + }) + + it.each([ + ['payment', { requiresPayment: true }, '/purchase'], + ['risk control', { requiresRiskControl: true }, '/admin/risk-control'], + ])('does not treat a failed %s settings load as explicitly disabled', async (_name, meta, path) => { + authStore.isAdmin = meta.requiresRiskControl === true + appStore.fetchPublicSettings.mockResolvedValue(null) + + const { navigation, next } = runGuard(meta, path) + await navigation + + expect(appStore.publicSettingsLoaded).toBe(false) + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith() + }) + + it.each([ + ['payment', { requiresPayment: true }, { payment_enabled: false }, '/dashboard'], + [ + 'risk control', + { requiresRiskControl: true }, + { risk_control_enabled: false }, + '/admin/settings', + ], + ])('redirects when loaded settings explicitly disable %s', async (_name, meta, settings, target) => { + authStore.isAdmin = meta.requiresRiskControl === true + appStore.cachedPublicSettings = settings + appStore.publicSettingsLoaded = true + + const { navigation, next } = runGuard(meta, '/feature') + await navigation + + expect(appStore.fetchPublicSettings).not.toHaveBeenCalled() + expect(next).toHaveBeenCalledOnce() + expect(next).toHaveBeenCalledWith(target) + }) +}) diff --git a/frontend/src/router/index.ts b/frontend/src/router/index.ts index 306a0eac30..e108d9d7b9 100644 --- a/frontend/src/router/index.ts +++ b/frontend/src/router/index.ts @@ -837,21 +837,24 @@ router.beforeEach(async (to, _from, next) => { } } - // Check payment requirement (internal payment system only) - if (to.meta.requiresPayment) { - const paymentEnabled = appStore.cachedPublicSettings?.payment_enabled - if (!paymentEnabled) { - next(authStore.isAdmin ? '/admin/dashboard' : '/dashboard') - return - } + // Only an explicit value from successfully loaded settings can disable a route. + // A transient settings failure is unknown state, not a confirmed feature toggle. + if ( + to.meta.requiresPayment && + appStore.publicSettingsLoaded && + appStore.cachedPublicSettings?.payment_enabled === false + ) { + next(authStore.isAdmin ? '/admin/dashboard' : '/dashboard') + return } - if (to.meta.requiresRiskControl) { - const riskControlEnabled = appStore.cachedPublicSettings?.risk_control_enabled === true - if (!riskControlEnabled) { - next(authStore.isAdmin ? '/admin/settings' : '/dashboard') - return - } + if ( + to.meta.requiresRiskControl && + appStore.publicSettingsLoaded && + appStore.cachedPublicSettings?.risk_control_enabled === false + ) { + next(authStore.isAdmin ? '/admin/settings' : '/dashboard') + return } // 简易模式下限制访问某些页面 diff --git a/frontend/src/stores/__tests__/app.spec.ts b/frontend/src/stores/__tests__/app.spec.ts index 803dad0e90..d9e7b17f57 100644 --- a/frontend/src/stores/__tests__/app.spec.ts +++ b/frontend/src/stores/__tests__/app.spec.ts @@ -2,6 +2,63 @@ import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest' import { setActivePinia, createPinia } from 'pinia' import { useAppStore } from '@/stores/app' import { getPublicSettings } from '@/api/auth' +import type { PublicSettings } from '@/types' + +function createDeferred() { + let resolve!: (value: T | PromiseLike) => void + let reject!: (reason?: unknown) => void + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise + reject = rejectPromise + }) + + return { promise, resolve, reject } +} + +function createPublicSettings(overrides: Partial = {}): PublicSettings { + return { + registration_enabled: false, + email_verify_enabled: false, + force_email_on_third_party_signup: false, + registration_email_suffix_whitelist: [], + promo_code_enabled: true, + password_reset_enabled: false, + invitation_code_enabled: false, + turnstile_enabled: false, + turnstile_site_key: '', + site_name: 'Test Site', + site_logo: '', + site_subtitle: '', + api_base_url: '', + contact_info: '', + doc_url: '', + home_content: '', + hide_ccs_import_button: false, + payment_enabled: false, + risk_control_enabled: false, + table_default_page_size: 20, + table_page_size_options: [10, 20, 50, 100], + custom_menu_items: [], + custom_endpoints: [], + linuxdo_oauth_enabled: false, + wechat_oauth_enabled: false, + oidc_oauth_enabled: false, + oidc_oauth_provider_name: 'OIDC', + github_oauth_enabled: false, + google_oauth_enabled: false, + backend_mode_enabled: false, + version: '1.0.0', + balance_low_notify_enabled: false, + account_quota_notify_enabled: false, + balance_low_notify_threshold: 0, + channel_monitor_enabled: true, + channel_monitor_default_interval_seconds: 60, + available_channels_enabled: false, + service_quota_enabled: false, + affiliate_enabled: false, + ...overrides, + } +} // Mock API 模块 vi.mock('@/api/admin/system', () => ({ @@ -17,6 +74,7 @@ describe('useAppStore', () => { setActivePinia(createPinia()) vi.useFakeTimers() localStorage.clear() + vi.mocked(getPublicSettings).mockReset() // 清除 window.__APP_CONFIG__ delete (window as any).__APP_CONFIG__ }) @@ -263,6 +321,75 @@ describe('useAppStore', () => { // --- 公开设置 --- describe('公开设置加载', () => { + it('并发调用复用并等待同一个请求,包括 force 调用', async () => { + const deferred = createDeferred() + vi.mocked(getPublicSettings).mockReturnValue(deferred.promise) + const settings = createPublicSettings({ payment_enabled: true }) + const store = useAppStore() + + const first = store.fetchPublicSettings() + const second = store.fetchPublicSettings() + const forced = store.fetchPublicSettings(true) + + expect(getPublicSettings).toHaveBeenCalledTimes(1) + + const settled = vi.fn() + void first.then(settled) + void second.then(settled) + void forced.then(settled) + await Promise.resolve() + expect(settled).not.toHaveBeenCalled() + + deferred.resolve(settings) + await expect(Promise.all([first, second, forced])).resolves.toEqual([ + settings, + settings, + settings, + ]) + expect(store.publicSettingsLoaded).toBe(true) + expect(store.cachedPublicSettings?.payment_enabled).toBe(true) + }) + + it('force 在无活动请求时绕过缓存,刷新期间的普通调用等待刷新结果', async () => { + const initial = createPublicSettings({ site_name: 'Initial Site' }) + vi.mocked(getPublicSettings).mockResolvedValueOnce(initial) + const store = useAppStore() + await store.fetchPublicSettings() + + const deferred = createDeferred() + const updated = createPublicSettings({ site_name: 'Updated Site' }) + vi.mocked(getPublicSettings).mockReturnValueOnce(deferred.promise) + + const refresh = store.fetchPublicSettings(true) + const duringRefresh = store.fetchPublicSettings() + + expect(getPublicSettings).toHaveBeenCalledTimes(2) + + deferred.resolve(updated) + await expect(Promise.all([refresh, duringRefresh])).resolves.toEqual([updated, updated]) + expect(store.siteName).toBe('Updated Site') + + await expect(store.fetchPublicSettings()).resolves.toEqual(updated) + expect(getPublicSettings).toHaveBeenCalledTimes(2) + }) + + it('并发请求失败时所有调用得到 null,且不会标记设置已加载', async () => { + const deferred = createDeferred() + vi.mocked(getPublicSettings).mockReturnValue(deferred.promise) + const consoleError = vi.spyOn(console, 'error').mockImplementation(() => undefined) + const store = useAppStore() + + const first = store.fetchPublicSettings() + const second = store.fetchPublicSettings() + deferred.reject(new Error('network unavailable')) + + await expect(Promise.all([first, second])).resolves.toEqual([null, null]) + expect(getPublicSettings).toHaveBeenCalledTimes(1) + expect(store.publicSettingsLoaded).toBe(false) + expect(store.cachedPublicSettings).toBeNull() + consoleError.mockRestore() + }) + it('从 window.__APP_CONFIG__ 初始化', () => { const windowAny = window as any windowAny.__APP_CONFIG__ = { diff --git a/frontend/src/stores/app.ts b/frontend/src/stores/app.ts index 20d580f6b9..51fff5c0a6 100644 --- a/frontend/src/stores/app.ts +++ b/frontend/src/stores/app.ts @@ -33,6 +33,7 @@ export const useAppStore = defineStore('app', () => { const apiBaseUrl = ref('') const docUrl = ref('') const cachedPublicSettings = ref(null) + let publicSettingsRequest: Promise | null = null // Version cache state const versionLoaded = ref(false) @@ -306,19 +307,25 @@ export const useAppStore = defineStore('app', () => { * Fetch public settings (uses cache unless force=true) * @param force - Force refresh from API */ - async function fetchPublicSettings(force = false): Promise { + function fetchPublicSettings(force = false): Promise { + // An active request always wins over cache/force semantics so every caller observes + // the same refresh result and no older request can overwrite a newer one. + if (publicSettingsRequest) { + return publicSettingsRequest + } + // Check for injected config from server (eliminates flash) if (!publicSettingsLoaded.value && !force && window.__APP_CONFIG__) { applySettings(window.__APP_CONFIG__) - return window.__APP_CONFIG__ + return Promise.resolve(window.__APP_CONFIG__) } // Return cached data if available and not forcing refresh if (publicSettingsLoaded.value && !force) { if (cachedPublicSettings.value) { - return { ...cachedPublicSettings.value } + return Promise.resolve({ ...cachedPublicSettings.value }) } - return { + return Promise.resolve({ registration_enabled: false, email_verify_enabled: false, force_email_on_third_party_signup: false, @@ -362,25 +369,37 @@ export const useAppStore = defineStore('app', () => { service_quota_enabled: false, affiliate_enabled: false, allow_user_view_error_requests: false, - } - } - - // Prevent duplicate requests - if (publicSettingsLoading.value) { - return null + }) } publicSettingsLoading.value = true + let apiRequest: Promise try { - const data = await fetchPublicSettingsAPI() - applySettings(data) - return data + apiRequest = fetchPublicSettingsAPI() } catch (error) { console.error('Failed to fetch public settings:', error) - return null - } finally { publicSettingsLoading.value = false + return Promise.resolve(null) } + + const request = apiRequest + .then((data) => { + applySettings(data) + return data + }) + .catch((error) => { + console.error('Failed to fetch public settings:', error) + return null + }) + .finally(() => { + if (publicSettingsRequest === request) { + publicSettingsRequest = null + publicSettingsLoading.value = false + } + }) + + publicSettingsRequest = request + return request } /** From c3ae5fc3c5ad084730ca494b9a4a567ad5f2be32 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 10 Jul 2026 11:09:06 +0800 Subject: [PATCH 15/27] =?UTF-8?q?fix(usage):=20effort=20=E6=8F=90=E5=8F=96?= =?UTF-8?q?=E6=94=B9=E7=94=A8=E6=A8=A1=E5=9E=8B=E5=80=99=E9=80=89=E5=88=97?= =?UTF-8?q?=E8=A1=A8=EF=BC=8C=E4=BF=AE=E5=A4=8D=E5=90=8E=E7=BC=80=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E5=85=83=E6=95=B0=E6=8D=AE=E4=B8=A2=E5=A4=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit #3909 将 effort 提取的模型参数改为 firstNonEmpty(upstreamModel, ...), 而 OAuth 的 normalizeCodexModel 会剥掉模型 effort 后缀 (gpt-5.4-xhigh -> gpt-5.4),导致后缀式模型且 body 无显式 reasoning 字段的请求在 usage 元数据中丢失 effort。 extractOpenAIReasoningEffortFromBody/extractOpenAIReasoningEffort 改为接受模型候选变参:显式 effort 的 max 保留判定沿用第一个非空 候选(映射后模型),后缀推导回退依次尝试每个候选(原始模型名兜底)。 纯元数据修复,不改变任何转发行为。 --- .../openai_gateway_chat_completions_raw.go | 2 +- .../service/openai_gateway_forward.go | 2 +- .../openai_gateway_messages_chat_fallback.go | 2 +- .../service/openai_gateway_request_body.go | 28 +++- .../openai_gateway_responses_chat_fallback.go | 2 +- ...openai_reasoning_effort_candidates_test.go | 122 ++++++++++++++++++ .../service/openai_ws_forwarder_ingress.go | 2 +- .../service/openai_ws_forwarder_v2.go | 2 +- .../internal/service/openai_ws_http_bridge.go | 2 +- 9 files changed, 151 insertions(+), 13 deletions(-) create mode 100644 backend/internal/service/openai_reasoning_effort_candidates_test.go diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 2f46a1c76c..4e3cfb440e 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -78,7 +78,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( // 2. Resolve model mapping (same as ForwardAsChatCompletions) billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) // 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 96f33c20f4..974f8eefaa 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -755,7 +755,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } defer func() { _ = resp.Body.Close() }() - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) // 国产模型默认 effort 补充:此处 reqModel 已被 mapping 重写为 billingModel(见 // line 2510-2515 的 GetMappedModel + reqModel 赋值),可直接作为 mappedModel。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, reqModel) diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback.go b/backend/internal/service/openai_gateway_messages_chat_fallback.go index 97eecf014b..f66fbc9ada 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback.go @@ -72,7 +72,7 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions( chatReq.StreamOptions = &apicompat.ChatStreamOptions{IncludeUsage: true} } - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) serviceTier := extractOpenAIServiceTierFromBody(body) diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 606a0d716e..f1e353b2af 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -331,6 +331,17 @@ func deriveOpenAIReasoningEffortFromModel(model string) string { return normalizeOpenAIReasoningEffortForModel(parts[len(parts)-1], modelID) } +// deriveOpenAIReasoningEffortFromModelCandidates 依次对每个候选模型做后缀推导, +// 返回第一个非空结果。 +func deriveOpenAIReasoningEffortFromModelCandidates(models []string) string { + for _, model := range models { + if value := deriveOpenAIReasoningEffortFromModel(model); value != "" { + return value + } + } + return "" +} + type openAIRequestView struct { body []byte Model string @@ -571,20 +582,24 @@ func detectOpenAIPassthroughInstructionsRejectReason(reqModel string, body []byt return "" } -func extractOpenAIReasoningEffortFromBody(body []byte, requestedModel string) *string { +// extractOpenAIReasoningEffortFromBody 按优先级传入模型候选(如 upstreamModel, +// billingModel, originalModel):显式 effort 的模型归一化(max 保留判定)用第一个 +// 非空候选;body 未携带 effort 时的模型后缀推导依次尝试每个候选——OAuth 的 +// normalizeCodexModel 会剥掉 upstreamModel 的 effort 后缀,只有原始模型名还留着。 +func extractOpenAIReasoningEffortFromBody(body []byte, modelCandidates ...string) *string { reasoningEffort := strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()) if reasoningEffort == "" { reasoningEffort = strings.TrimSpace(gjson.GetBytes(body, "reasoning_effort").String()) } if reasoningEffort != "" { - normalized := normalizeOpenAIReasoningEffortForModel(reasoningEffort, requestedModel) + normalized := normalizeOpenAIReasoningEffortForModel(reasoningEffort, firstNonEmpty(modelCandidates...)) if normalized == "" { return nil } return &normalized } - value := deriveOpenAIReasoningEffortFromModel(requestedModel) + value := deriveOpenAIReasoningEffortFromModelCandidates(modelCandidates) if value == "" { return nil } @@ -1159,15 +1174,16 @@ func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error return reqBody, nil } -func extractOpenAIReasoningEffort(reqBody map[string]any, requestedModel string) *string { - if value, present := getOpenAIReasoningEffortFromReqBody(reqBody, requestedModel); present { +// extractOpenAIReasoningEffort 的模型候选语义同 extractOpenAIReasoningEffortFromBody。 +func extractOpenAIReasoningEffort(reqBody map[string]any, modelCandidates ...string) *string { + if value, present := getOpenAIReasoningEffortFromReqBody(reqBody, firstNonEmpty(modelCandidates...)); present { if value == "" { return nil } return &value } - value := deriveOpenAIReasoningEffortFromModel(requestedModel) + value := deriveOpenAIReasoningEffortFromModelCandidates(modelCandidates) if value == "" { return nil } diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index 9f2cb84864..247c195f24 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -48,7 +48,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( billingModel := resolveOpenAIForwardModel(account, originalModel, "") upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) - reasoningEffort := extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(upstreamModel, billingModel, originalModel)) + reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel) // 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。 reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel) chatReq.Model = upstreamModel diff --git a/backend/internal/service/openai_reasoning_effort_candidates_test.go b/backend/internal/service/openai_reasoning_effort_candidates_test.go new file mode 100644 index 0000000000..66a3e16dfb --- /dev/null +++ b/backend/internal/service/openai_reasoning_effort_candidates_test.go @@ -0,0 +1,122 @@ +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestExtractOpenAIReasoningEffortFromBodyModelCandidates(t *testing.T) { + bodyWithoutEffort := []byte(`{"model":"whatever","input":"hello"}`) + bodyWithMax := []byte(`{"model":"sol","reasoning":{"effort":"max"},"input":"hello"}`) + + tests := []struct { + name string + body []byte + candidates []string + want string // "" 表示期望 nil + }{ + { + name: "后缀推导回退到原始模型(OAuth 上游模型已剥后缀)", + body: bodyWithoutEffort, + candidates: []string{"gpt-5.4", "gpt-5.4", "gpt-5.4-xhigh"}, + want: "xhigh", + }, + { + name: "GPT-5.6 后缀 max 经原始模型推导保留", + body: bodyWithoutEffort, + candidates: []string{"gpt-5.6-sol", "gpt-5.6-sol", "gpt-5.6-sol-max"}, + want: "max", + }, + { + name: "显式 max 用第一个非空候选(映射后模型)判定", + body: bodyWithMax, + candidates: []string{"gpt-5.6-sol", "sol"}, + want: "max", + }, + { + name: "显式 max 非 5.6 首候选仍折叠为 xhigh", + body: bodyWithMax, + candidates: []string{"gpt-5.4", "sol"}, + want: "xhigh", + }, + { + name: "所有候选均无后缀时返回 nil", + body: bodyWithoutEffort, + candidates: []string{"gpt-5.4", "gpt-5.4", "gpt-5.4"}, + want: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := extractOpenAIReasoningEffortFromBody(tt.body, tt.candidates...) + if tt.want == "" { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, tt.want, *got) + }) + } +} + +func TestExtractOpenAIReasoningEffortModelCandidates(t *testing.T) { + reqBody := map[string]any{"model": "gpt-5.3-codex-high", "input": "hello"} + + got := extractOpenAIReasoningEffort(reqBody, "gpt-5.3-codex", "gpt-5.3-codex-high") + + require.NotNil(t, got) + require.Equal(t, "high", *got) +} + +// 回归:OAuth 账号请求后缀式模型(无显式 reasoning 字段)时,上游模型被 +// normalizeCodexModel 剥掉 effort 后缀,用量元数据的 effort 必须仍能从 +// 原始模型名后缀推导出来。 +func TestOpenAIGatewayServiceForwardOAuthDerivesEffortFromSuffixModel(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 11, + Name: "openai-oauth-suffix", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + Status: StatusActive, + Schedulable: true, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.3-codex-xhigh","instructions":"suffix-test","input":"hello","stream":false}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "gpt-5.3-codex", gjson.GetBytes(upstream.lastBody, "model").String()) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "xhigh", *result.ReasoningEffort) +} diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 9dea4118da..4a5fa20774 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -920,7 +920,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( Model: originalModel, UpstreamModel: mappedModel, ServiceTier: extractOpenAIServiceTierFromBody(payload), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, firstNonEmpty(mappedModel, originalModel)), payload, mappedModel), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, mappedModel, originalModel), payload, mappedModel), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: lease.HandshakeHeaders(), diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index 40a3cc0194..24853d4a6c 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -693,7 +693,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( ImageCount: imageCounter.Count(), ImageOutputSizes: imageCounter.Sizes(), ServiceTier: extractOpenAIServiceTier(reqBody), - ReasoningEffort: extractOpenAIReasoningEffort(reqBody, firstNonEmpty(mappedModel, originalModel)), + ReasoningEffort: extractOpenAIReasoningEffort(reqBody, mappedModel, originalModel), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: lease.HandshakeHeaders(), diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index df1a8067cb..dce0c7b6db 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -263,7 +263,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( Model: originalModel, UpstreamModel: mappedModel, ServiceTier: extractOpenAIServiceTierFromBody(body), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, firstNonEmpty(mappedModel, originalModel)), body, mappedModel), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel), Stream: reqStream, OpenAIWSMode: true, ResponseHeaders: cloneHeader(resp.Header), From dda8f78733958495d10b4ff17510df6aea83426c Mon Sep 17 00:00:00 2001 From: li Date: Fri, 10 Jul 2026 11:43:44 +0800 Subject: [PATCH 16/27] =?UTF-8?q?fix(admin):=20GetUserBreakdown=20?= =?UTF-8?q?=E4=BD=BF=E7=94=A8=20ParseUsageRequestType=20=E8=A7=A3=E6=9E=90?= =?UTF-8?q?=20request=5Ftype?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #3920 --- .../handler/admin/dashboard_handler.go | 11 ++++-- .../dashboard_handler_user_breakdown_test.go | 39 +++++++++++++++++++ 2 files changed, 46 insertions(+), 4 deletions(-) diff --git a/backend/internal/handler/admin/dashboard_handler.go b/backend/internal/handler/admin/dashboard_handler.go index b42b395d33..8f55fb4165 100644 --- a/backend/internal/handler/admin/dashboard_handler.go +++ b/backend/internal/handler/admin/dashboard_handler.go @@ -657,11 +657,14 @@ func (h *DashboardHandler) GetUserBreakdown(c *gin.Context) { dim.AccountID = id } } - if v := c.Query("request_type"); v != "" { - if rt, err := strconv.ParseInt(v, 10, 16); err == nil { - rtVal := int16(rt) - dim.RequestType = &rtVal + if v := strings.TrimSpace(c.Query("request_type")); v != "" { + parsed, err := service.ParseUsageRequestType(v) + if err != nil { + response.BadRequest(c, err.Error()) + return } + rtVal := int16(parsed) + dim.RequestType = &rtVal } if v := c.Query("stream"); v != "" { if s, err := strconv.ParseBool(v); err == nil { diff --git a/backend/internal/handler/admin/dashboard_handler_user_breakdown_test.go b/backend/internal/handler/admin/dashboard_handler_user_breakdown_test.go index 3065eee3c4..2381364a94 100644 --- a/backend/internal/handler/admin/dashboard_handler_user_breakdown_test.go +++ b/backend/internal/handler/admin/dashboard_handler_user_breakdown_test.go @@ -241,3 +241,42 @@ func TestGetUserBreakdown_NoFilters(t *testing.T) { require.Empty(t, repo.capturedDim.Model) require.Empty(t, repo.capturedDim.Endpoint) } + +func TestGetUserBreakdown_RequestTypeStringFilter(t *testing.T) { + cases := []struct { + name string + value string + want int16 + }{ + {"ws_v2", "ws_v2", int16(service.RequestTypeWSV2)}, + {"stream", "stream", int16(service.RequestTypeStream)}, + {"sync", "sync", int16(service.RequestTypeSync)}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + repo := &userBreakdownRepoCapture{} + router := newUserBreakdownRouter(repo) + + req := httptest.NewRequest(http.MethodGet, + "/admin/dashboard/user-breakdown?start_date=2026-03-01&end_date=2026-03-16&request_type="+tc.value, nil) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + require.NotNil(t, repo.capturedDim.RequestType, "request_type=%s should set filter", tc.value) + require.Equal(t, tc.want, *repo.capturedDim.RequestType) + }) + } +} + +func TestGetUserBreakdown_InvalidRequestType(t *testing.T) { + repo := &userBreakdownRepoCapture{} + router := newUserBreakdownRouter(repo) + + req := httptest.NewRequest(http.MethodGet, + "/admin/dashboard/user-breakdown?start_date=2026-03-01&end_date=2026-03-16&request_type=bogus", nil) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusBadRequest, w.Code) +} From b9b013a0881981fd01f59dec45a2d5f0735aa37a Mon Sep 17 00:00:00 2001 From: li Date: Fri, 10 Jul 2026 11:56:47 +0800 Subject: [PATCH 17/27] =?UTF-8?q?fix(usage):=20WS=20passthrough=20effort?= =?UTF-8?q?=20=E6=8F=90=E5=8F=96=E8=A1=A5=E5=85=A5=E6=98=A0=E5=B0=84?= =?UTF-8?q?=E5=90=8E=E6=A8=A1=E5=9E=8B=E5=80=99=E9=80=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #3924 --- .../service/openai_fast_policy_ws_test.go | 4 +- .../openai_ws_v2_passthrough_adapter.go | 12 +++--- ...i_ws_v2_passthrough_adapter_effort_test.go | 42 +++++++++++++++++++ 3 files changed, 50 insertions(+), 8 deletions(-) create mode 100644 backend/internal/service/openai_ws_v2_passthrough_adapter_effort_test.go diff --git a/backend/internal/service/openai_fast_policy_ws_test.go b/backend/internal/service/openai_fast_policy_ws_test.go index ac5b363cc0..a802540879 100644 --- a/backend/internal/service/openai_fast_policy_ws_test.go +++ b/backend/internal/service/openai_fast_policy_ws_test.go @@ -1015,7 +1015,7 @@ func TestPassthroughUsageMeta_TracksReasoningEffortAcrossTurns(t *testing.T) { firstOut, firstBlocked, firstErr := svc.applyOpenAIFastPolicyToWSResponseCreate(context.Background(), account, capturedSessionModel, firstFrame) require.NoError(t, firstErr) require.Nil(t, firstBlocked) - meta.initFromFirstFrame(firstOut) + meta.initFromFirstFrame(firstOut, capturedSessionModel) require.NotNil(t, meta.reasoningEffort.Load()) require.Equal(t, "medium", *meta.reasoningEffort.Load()) @@ -1032,7 +1032,7 @@ func TestPassthroughUsageMeta_TracksReasoningEffortAcrossTurns(t *testing.T) { out, blocked, policyErr := svc.applyOpenAIFastPolicyToWSResponseCreate(context.Background(), account, model, payload) if policyErr == nil && blocked == nil && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { - meta.updateFromResponseCreate(out, requestModelForThisFrame) + meta.updateFromResponseCreate(out, model, requestModelForThisFrame) } return out, blocked, policyErr } diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index 7241f11ed1..d822d181cd 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -142,12 +142,12 @@ func newOpenAIWSPassthroughUsageMeta(initialRequestModel string, firstFrame []by return meta } -func (m *openAIWSPassthroughUsageMeta) initFromFirstFrame(policyOutput []byte) { +func (m *openAIWSPassthroughUsageMeta) initFromFirstFrame(policyOutput []byte, mappedModel string) { if m == nil { return } m.serviceTier.Store(extractOpenAIServiceTierFromBody(policyOutput)) - m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, m.sessionRequestModel)) + m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, mappedModel, m.sessionRequestModel)) } func (m *openAIWSPassthroughUsageMeta) updateSessionRequestModel(payload []byte) { @@ -169,12 +169,12 @@ func (m *openAIWSPassthroughUsageMeta) requestModelForFrame(payload []byte) stri return m.sessionRequestModel } -func (m *openAIWSPassthroughUsageMeta) updateFromResponseCreate(policyOutput []byte, requestModelForFrame string) { +func (m *openAIWSPassthroughUsageMeta) updateFromResponseCreate(policyOutput []byte, mappedModel string, requestModelForFrame string) { if m == nil { return } m.serviceTier.Store(extractOpenAIServiceTierFromBody(policyOutput)) - m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, requestModelForFrame)) + m.reasoningEffort.Store(extractOpenAIReasoningEffortFromBody(policyOutput, mappedModel, requestModelForFrame)) } func openAIWSPassthroughRequestModelForFrame(payload []byte) string { @@ -311,7 +311,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // 因此使用 atomic.Pointer[string] 在 filter(runClientToUpstream // goroutine)和 OnTurnComplete / final result(runUpstreamToClient // goroutine)之间同步当前 turn 的 usage metadata。 - usageMeta.initFromFirstFrame(firstClientMessage) + usageMeta.initFromFirstFrame(firstClientMessage, capturedSessionModel) promptCacheKey := strings.TrimSpace(gjson.GetBytes(firstClientMessage, "prompt_cache_key").String()) wsURL, err := s.buildOpenAIResponsesWSURL(account) @@ -455,7 +455,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // service_tier 时按 default 处理,billing 应如实反映。 if policyErr == nil && blocked == nil && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { - usageMeta.updateFromResponseCreate(out, requestModelForThisFrame) + usageMeta.updateFromResponseCreate(out, model, requestModelForThisFrame) } return out, blocked, policyErr }, diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter_effort_test.go b/backend/internal/service/openai_ws_v2_passthrough_adapter_effort_test.go new file mode 100644 index 0000000000..82c9eaac85 --- /dev/null +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter_effort_test.go @@ -0,0 +1,42 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestWSPassthroughUsageMeta_InitFromFirstFrame_MappedModelCandidate(t *testing.T) { + body := []byte(`{"type":"response.create","model":"sol","reasoning":{"effort":"max"}}`) + + meta := newOpenAIWSPassthroughUsageMeta("sol", body) + meta.initFromFirstFrame(body, "gpt-5.6-sol") + + got := meta.reasoningEffort.Load() + require.NotNil(t, got, "reasoning effort should be set") + require.Equal(t, "max", *got, "mapped model gpt-5.6-sol should preserve max") +} + +func TestWSPassthroughUsageMeta_InitFromFirstFrame_NonGPT56FallsBackToXHigh(t *testing.T) { + body := []byte(`{"type":"response.create","model":"gpt-5.4","reasoning":{"effort":"max"}}`) + + meta := newOpenAIWSPassthroughUsageMeta("gpt-5.4", body) + meta.initFromFirstFrame(body, "gpt-5.4") + + got := meta.reasoningEffort.Load() + require.NotNil(t, got) + require.Equal(t, "xhigh", *got, "non-5.6 model should normalize max to xhigh") +} + +func TestWSPassthroughUsageMeta_UpdateFromResponseCreate_MappedModelCandidate(t *testing.T) { + body := []byte(`{"type":"response.create","model":"sol","reasoning":{"effort":"max"}}`) + + meta := newOpenAIWSPassthroughUsageMeta("sol", body) + meta.updateFromResponseCreate(body, "gpt-5.6-sol", "sol") + + got := meta.reasoningEffort.Load() + require.NotNil(t, got) + require.Equal(t, "max", *got, "mapped model should preserve max on multi-turn update") +} From ea9f40b63fff8a98ccfab7e1f2f51feab3bdf95d Mon Sep 17 00:00:00 2001 From: li Date: Fri, 10 Jul 2026 12:38:21 +0800 Subject: [PATCH 18/27] =?UTF-8?q?fix(frontend):=20UserBreakdownParams.requ?= =?UTF-8?q?est=5Ftype=20=E7=B1=BB=E5=9E=8B=E4=BB=8E=20number=20=E6=94=B9?= =?UTF-8?q?=E4=B8=BA=20UsageRequestType?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- frontend/src/api/admin/dashboard.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/frontend/src/api/admin/dashboard.ts b/frontend/src/api/admin/dashboard.ts index 97c16aa362..ae20d33f8e 100644 --- a/frontend/src/api/admin/dashboard.ts +++ b/frontend/src/api/admin/dashboard.ts @@ -173,7 +173,7 @@ export interface UserBreakdownParams { user_id?: number api_key_id?: number account_id?: number - request_type?: number + request_type?: UsageRequestType stream?: boolean billing_type?: number | null } From 0dec1ad2922ff8c9d27b67f8a31dfb35bce1902b Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 10 Jul 2026 14:09:07 +0800 Subject: [PATCH 19/27] =?UTF-8?q?fix(service):=20=E6=B6=88=E9=99=A4=20isOp?= =?UTF-8?q?enAIGPT56Model=20=E9=87=8D=E5=A4=8D=E5=A3=B0=E6=98=8E=EF=BC=8C?= =?UTF-8?q?=E7=BB=9F=E4=B8=80=E5=88=B0=20openai=5Fmodel=5Falias.go?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit PR #3898 与 #3909 各自新增了同名函数(billing 精确匹配版 / gateway 前缀匹配版),合并后 internal/service 编译失败。保留前缀匹配语义作为 唯一实现:对 normalizeKnownOpenAICodexModel 的输出(gpt-5.6-* 精确 基名)canonicalize 幂等且精确分支命中,计费行为不变;同时保住 gpt-5.6-*-max 等后缀变体的识别。 --- backend/internal/service/billing_service.go | 4 ---- .../internal/service/openai_gateway_request_body.go | 10 ---------- backend/internal/service/openai_model_alias.go | 12 ++++++++++++ 3 files changed, 12 insertions(+), 14 deletions(-) diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 3c8b52af97..71d74ae6e0 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -1141,10 +1141,6 @@ func (s *BillingService) shouldApplySessionLongContextPricing(tokens UsageTokens return totalInputTokens > pricing.LongContextInputThreshold } -func isOpenAIGPT56Model(normalized string) bool { - return normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna" -} - func usesOpenAILegacyLongContextPricing(normalized string) bool { return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" } diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index f1e353b2af..0e888e145b 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -1218,13 +1218,3 @@ func normalizeOpenAIReasoningEffortForModel(raw, model string) string { } return normalizeOpenAIReasoningEffort(raw) } - -func isOpenAIGPT56Model(model string) bool { - normalized := canonicalizeOpenAIModelAliasSpelling(model) - for _, prefix := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { - if normalized == prefix || strings.HasPrefix(normalized, prefix+"-") { - return true - } - } - return false -} diff --git a/backend/internal/service/openai_model_alias.go b/backend/internal/service/openai_model_alias.go index 4e3d3b2d9a..338de9e8a5 100644 --- a/backend/internal/service/openai_model_alias.go +++ b/backend/internal/service/openai_model_alias.go @@ -98,6 +98,18 @@ func normalizeKnownOpenAICodexModel(model string) string { } } +// isOpenAIGPT56Model 判断是否 GPT-5.6 系列模型;入参可为原始模型名 +// (含大小写/路径/后缀变体)或已归一化的基名,两者均能正确识别。 +func isOpenAIGPT56Model(model string) bool { + normalized := canonicalizeOpenAIModelAliasSpelling(model) + for _, prefix := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { + if normalized == prefix || strings.HasPrefix(normalized, prefix+"-") { + return true + } + } + return false +} + func appendUsageBillingModelCandidate(candidates []string, seen map[string]struct{}, model string) []string { trimmed := strings.TrimSpace(model) if trimmed == "" { From 9a2f11b4e21763cb7003ea29921d9a672ab50b1f Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Fri, 10 Jul 2026 06:27:11 +0000 Subject: [PATCH 20/27] chore: sync VERSION to 0.1.150 [skip ci] --- backend/cmd/server/VERSION | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 010fbfb884..23c64fda6f 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.149 +0.1.150 From a495d5e3024f93419d981a36fcd983b9732926d7 Mon Sep 17 00:00:00 2001 From: Shumin <332587268@qq.com> Date: Fri, 10 Jul 2026 02:27:33 -0400 Subject: [PATCH 21/27] =?UTF-8?q?test:=20=E6=9B=B4=E6=96=B0=E6=96=AD?= =?UTF-8?q?=E8=A8=80=E4=BB=A5=E8=A6=86=E7=9B=96=20setup-token=20=E7=BA=B3?= =?UTF-8?q?=E5=85=A5=E5=90=8E=E5=8F=B0=E5=88=B7=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - ListOAuthRefreshCandidates SQL 断言由 type = 'oauth' 改为 type IN ('oauth', 'setup-token')。 - ClaudeTokenRefresher.CanRefresh 增加 anthropic setup-token → true 用例。 --- .../internal/repository/account_repo_temp_unsched_test.go | 3 ++- backend/internal/service/token_refresher_test.go | 6 ++++++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/backend/internal/repository/account_repo_temp_unsched_test.go b/backend/internal/repository/account_repo_temp_unsched_test.go index eb123ee6a3..9c3e4cbdf5 100644 --- a/backend/internal/repository/account_repo_temp_unsched_test.go +++ b/backend/internal/repository/account_repo_temp_unsched_test.go @@ -43,7 +43,8 @@ func TestAccountRepository_ListOAuthRefreshCandidates_SQLFilter(t *testing.T) { normalized := normalizeSQLWhitespace(capturedSQL) require.Contains(t, normalized, "deleted_at IS NULL") require.Contains(t, normalized, "status = 'active'") - require.Contains(t, normalized, "type = 'oauth'") + // setup-token 的 access_token 同为 8h 短期令牌,必须与 oauth 一起纳入后台刷新候选 + require.Contains(t, normalized, "type IN ('oauth', 'setup-token')") require.Contains(t, normalized, "platform IN ('anthropic', 'openai', 'gemini', 'antigravity')") require.Contains(t, normalized, "credentials ? 'refresh_token'") require.Contains(t, normalized, "btrim(credentials->>'refresh_token') <> ''") diff --git a/backend/internal/service/token_refresher_test.go b/backend/internal/service/token_refresher_test.go index d1cf88983b..84cce1de72 100644 --- a/backend/internal/service/token_refresher_test.go +++ b/backend/internal/service/token_refresher_test.go @@ -194,6 +194,12 @@ func TestClaudeTokenRefresher_CanRefresh(t *testing.T) { accType: AccountTypeOAuth, want: true, }, + { + name: "anthropic setup-token - can refresh", + platform: PlatformAnthropic, + accType: AccountTypeSetupToken, + want: true, + }, { name: "anthropic api-key - cannot refresh", platform: PlatformAnthropic, From d3a1835ed76fde8860b9c149613ad8975016615c Mon Sep 17 00:00:00 2001 From: Tassoi Date: Fri, 10 Jul 2026 06:37:06 +0000 Subject: [PATCH 22/27] fix(image): strip Codex image_gen namespace declarations --- .../service/image_generation_intent.go | 58 ++++---- .../service/image_generation_intent_test.go | 55 ++++++++ .../service/openai_codex_transform.go | 128 +++++++++++++----- .../service/openai_codex_transform_test.go | 104 ++++++++++++++ .../service/openai_gateway_forward.go | 23 +++- .../openai_image_generation_controls_test.go | 59 ++++++++ .../service/openai_ws_forwarder_ingress.go | 2 +- .../openai_ws_forwarder_ingress_test.go | 87 +++++++++--- .../service/openai_ws_forwarder_v2.go | 20 +-- 9 files changed, 437 insertions(+), 99 deletions(-) diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go index 12590f3f63..5a063a7d71 100644 --- a/backend/internal/service/image_generation_intent.go +++ b/backend/internal/service/image_generation_intent.go @@ -93,7 +93,7 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { } found := false tools.ForEach(func(_, item gjson.Result) bool { - if openAIJSONString(item.Get("type")) == "image_generation" { + if isOpenAIImageGenerationType(openAIJSONString(item.Get("type"))) { found = true return false } @@ -106,12 +106,20 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { return found } +func isOpenAIImageGenerationType(value string) bool { + return strings.TrimSpace(value) == "image_generation" +} + +func isOpenAIImageGenNamespaceName(value string) bool { + return strings.TrimSpace(value) == "image_gen" +} + // isImageGenNamespaceTool detects the Codex namespace-style image generation // tool declaration: { "type": "namespace", "name": "image_gen", ... }. // Codex /image uses this instead of the flat { "type": "image_generation" }. func isImageGenNamespaceTool(tool gjson.Result) bool { return openAIJSONString(tool.Get("type")) == "namespace" && - openAIJSONString(tool.Get("name")) == "image_gen" + isOpenAIImageGenNamespaceName(openAIJSONString(tool.Get("name"))) } // openAIJSONInputContainsImageGenTool scans Responses input items for @@ -127,27 +135,19 @@ func openAIJSONInputContainsImageGenTool(input gjson.Result) bool { if openAIJSONString(item.Get("type")) != "additional_tools" { return true } - tools := item.Get("tools") - if !tools.IsArray() { - return true - } - tools.ForEach(func(_, tool gjson.Result) bool { - if isImageGenNamespaceTool(tool) { - found = true - return false - } - return true - }) + found = openAIJSONToolsContainImageGeneration(item.Get("tools")) return !found }) return found } -func openAIRequestBodyHasImageGenerationTool(body []byte) bool { +func openAIRequestBodyHasImageGenerationDeclaration(body []byte) bool { if len(body) == 0 || !gjson.ValidBytes(body) { return false } - return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools")) + return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools")) || + openAIJSONInputContainsImageGenTool(gjson.GetBytes(body, "input")) || + openAIJSONToolChoiceSelectsImageGeneration(gjson.GetBytes(body, "tool_choice")) } func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool { @@ -178,18 +178,24 @@ func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool { return false } if choice.Type == gjson.String { - return strings.TrimSpace(choice.String()) == "image_generation" + return isOpenAIImageGenerationType(choice.String()) } if !choice.IsObject() { return false } - if strings.TrimSpace(choice.Get("type").String()) == "image_generation" { + choiceType := openAIJSONString(choice.Get("type")) + if isOpenAIImageGenerationType(choiceType) { return true } - if strings.TrimSpace(choice.Get("tool.type").String()) == "image_generation" { + if choiceType == "namespace" && + (isOpenAIImageGenNamespaceName(openAIJSONString(choice.Get("name"))) || + isOpenAIImageGenNamespaceName(openAIJSONString(choice.Get("namespace")))) { return true } - if strings.TrimSpace(choice.Get("function.name").String()) == "image_generation" { + if tool := choice.Get("tool"); tool.IsObject() && openAIJSONToolChoiceSelectsImageGeneration(tool) { + return true + } + if isOpenAIImageGenerationType(openAIJSONString(choice.Get("function.name"))) { return true } return false @@ -198,15 +204,21 @@ func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool { func openAIAnyToolChoiceSelectsImageGeneration(choice any) bool { switch v := choice.(type) { case string: - return strings.TrimSpace(v) == "image_generation" + return isOpenAIImageGenerationType(v) case map[string]any: - if strings.TrimSpace(firstNonEmptyString(v["type"])) == "image_generation" { + choiceType := strings.TrimSpace(firstNonEmptyString(v["type"])) + if isOpenAIImageGenerationType(choiceType) { return true } - if tool, ok := v["tool"].(map[string]any); ok && strings.TrimSpace(firstNonEmptyString(tool["type"])) == "image_generation" { + if choiceType == "namespace" && + (isOpenAIImageGenNamespaceName(firstNonEmptyString(v["name"])) || + isOpenAIImageGenNamespaceName(firstNonEmptyString(v["namespace"]))) { return true } - if fn, ok := v["function"].(map[string]any); ok && strings.TrimSpace(firstNonEmptyString(fn["name"])) == "image_generation" { + if tool, ok := v["tool"].(map[string]any); ok && openAIAnyToolChoiceSelectsImageGeneration(tool) { + return true + } + if fn, ok := v["function"].(map[string]any); ok && isOpenAIImageGenerationType(firstNonEmptyString(fn["name"])) { return true } } diff --git a/backend/internal/service/image_generation_intent_test.go b/backend/internal/service/image_generation_intent_test.go index 1a32318cce..7f3ac8a611 100644 --- a/backend/internal/service/image_generation_intent_test.go +++ b/backend/internal/service/image_generation_intent_test.go @@ -41,6 +41,20 @@ func TestIsImageGenerationIntent(t *testing.T) { body: []byte(`{"model":"gpt-5.4","tool_choice":{"type":"image_generation"}}`), want: true, }, + { + name: "namespace image_gen tool choice", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tool_choice":{"type":"namespace","name":"image_gen"}}`), + want: true, + }, + { + name: "custom imagegen function tool choice is not image intent", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tool_choice":{"function":{"name":"imagegen"}}}`), + want: false, + }, { name: "required tool choice alone is text", endpoint: "/v1/responses", @@ -62,6 +76,13 @@ func TestIsImageGenerationIntent(t *testing.T) { body: []byte(`{"model":"gpt-5.5","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]}`), want: true, }, + { + name: "custom namespace with nested imagegen function is not image intent", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tools":[{"type":"namespace","name":"media_tools","tools":[{"type":"function","name":"imagegen"}]}]}`), + want: false, + }, { name: "namespace image_gen in input additional_tools (Responses Lite)", endpoint: "/v1/responses", @@ -118,6 +139,40 @@ func TestIsImageGenerationIntentMap_NamespaceImageGen(t *testing.T) { }, want: true, }, + { + name: "custom namespace with nested imagegen function is not image intent", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "media_tools", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + want: false, + }, + { + name: "namespace image_gen tool choice", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tool_choice": map[string]any{"type": "namespace", "name": "image_gen"}, + }, + want: true, + }, + { + name: "custom imagegen function tool choice is not image intent", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tool_choice": map[string]any{ + "function": map[string]any{"name": "imagegen"}, + }, + }, + want: false, + }, { name: "non-image namespace not flagged", reqBody: map[string]any{ diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 36426e7377..99355628f2 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -596,7 +596,7 @@ func hasOpenAIImageGenerationTool(reqBody map[string]any) bool { if toolsContainImageGeneration(reqBody["tools"]) { return true } - return inputContainsImageGenNamespace(reqBody["input"]) + return inputContainsImageGenerationTool(reqBody["input"]) } func toolsContainImageGeneration(rawTools any) bool { @@ -612,22 +612,24 @@ func toolsContainImageGeneration(rawTools any) bool { if !ok { continue } - if strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { - return true - } - if isImageGenNamespaceToolMap(toolMap) { + if isOpenAIImageGenerationToolMap(toolMap) { return true } } return false } -func isImageGenNamespaceToolMap(tool map[string]any) bool { - return strings.TrimSpace(firstNonEmptyString(tool["type"])) == "namespace" && - strings.TrimSpace(firstNonEmptyString(tool["name"])) == "image_gen" +func isOpenAIImageGenerationToolMap(tool map[string]any) bool { + return isOpenAIImageGenerationType(firstNonEmptyString(tool["type"])) || + isImageGenNamespaceToolMap(tool) } -func inputContainsImageGenNamespace(rawInput any) bool { +func isImageGenNamespaceToolMap(tool map[string]any) bool { + return strings.TrimSpace(firstNonEmptyString(tool["type"])) == "namespace" && + isOpenAIImageGenNamespaceName(firstNonEmptyString(tool["name"])) +} + +func inputContainsImageGenerationTool(rawInput any) bool { input, ok := rawInput.([]any) if !ok { return false @@ -647,54 +649,110 @@ func inputContainsImageGenNamespace(rawInput any) bool { return false } +// stripOpenAIImageGenerationTools keeps account-level strip policy symmetric +// across standard Responses tools, Responses Lite additional_tools, and tool_choice. func stripOpenAIImageGenerationTools(reqBody map[string]any) bool { - rawTools, ok := reqBody["tools"] + if reqBody == nil { + return false + } + modified := stripOpenAIImageGenerationToolList(reqBody, "tools") + if stripOpenAIImageGenerationToolsFromInput(reqBody) { + modified = true + } + if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + delete(reqBody, "tool_choice") + modified = true + } + return modified +} + +func stripOpenAIImageGenerationToolList(container map[string]any, key string) bool { + rawTools, ok := container[key] if !ok || rawTools == nil { - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") - return true - } return false } tools, ok := rawTools.([]any) if !ok { - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") - return true - } return false } filtered := make([]any, 0, len(tools)) removed := false for _, rawTool := range tools { - if toolMap, ok := rawTool.(map[string]any); ok && - strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { + if toolMap, ok := rawTool.(map[string]any); ok && isOpenAIImageGenerationToolMap(toolMap) { removed = true continue } filtered = append(filtered, rawTool) } - if !removed && !openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + if !removed { return false } - if removed { - if len(filtered) == 0 { - delete(reqBody, "tools") - } else { - reqBody["tools"] = filtered - } - } - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") + if len(filtered) == 0 { + delete(container, key) + } else { + container[key] = filtered } return true } -// stripCodexSparkImageGenerationTools removes image_generation tool entries from -// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400 -// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so -// it must be dropped for spark. When the tools list becomes empty the key is removed. -// Returns true when the body was modified. +func stripOpenAIImageGenerationToolsFromInput(reqBody map[string]any) bool { + input, ok := reqBody["input"].([]any) + if !ok { + return false + } + + filteredInput := make([]any, 0, len(input)) + modified := false + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) != "additional_tools" { + filteredInput = append(filteredInput, rawItem) + continue + } + if !stripOpenAIImageGenerationToolList(item, "tools") { + filteredInput = append(filteredInput, rawItem) + continue + } + modified = true + if _, hasTools := item["tools"]; hasTools { + filteredInput = append(filteredInput, rawItem) + } + // An empty additional_tools carrier is not useful upstream; drop the item + // after its only declared capability has been removed. + } + if modified { + reqBody["input"] = filteredInput + } + return modified +} + +// stripOpenAIImageGenerationToolsFromRawPayload is the shared adapter for paths +// that forward raw HTTP or WebSocket payloads without the normal request map. +func stripOpenAIImageGenerationToolsFromRawPayload(payload []byte) ([]byte, bool, error) { + if !openAIRequestBodyHasImageGenerationDeclaration(payload) { + if json.Valid(payload) { + return payload, false, nil + } + var invalidPayload map[string]any + return payload, false, json.Unmarshal(payload, &invalidPayload) + } + payloadMap := make(map[string]any) + if err := json.Unmarshal(payload, &payloadMap); err != nil { + return payload, false, err + } + if !stripOpenAIImageGenerationTools(payloadMap) { + return payload, false, nil + } + rebuilt, err := json.Marshal(payloadMap) + if err != nil { + return payload, false, err + } + return rebuilt, true, nil +} + +// stripCodexSparkImageGenerationTools removes image tool declarations and choices. +// gpt-5.3-codex-spark rejects those capabilities upstream, while Codex clients may +// advertise them by default. func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool { return stripOpenAIImageGenerationTools(reqBody) } diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index bb27c81352..b226655eeb 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -796,6 +796,110 @@ func TestApplyCodexOAuthTransform_StripsImageGenerationToolForSparkAlias(t *test require.False(t, hasTools) } +func TestStripOpenAIImageGenerationTools_StripsNamespaceFormats(t *testing.T) { + imageNamespace := func() map[string]any { + return map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + } + } + codeNamespace := func() map[string]any { + return map[string]any{ + "type": "namespace", + "name": "code_tools", + "tools": []any{ + map[string]any{"type": "function", "name": "run"}, + }, + } + } + + reqBody := map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{"type": "function", "name": "shell"}, + imageNamespace(), + codeNamespace(), + }, + "input": []any{ + map[string]any{"type": "message", "role": "user", "content": "hello"}, + map[string]any{ + "type": "additional_tools", + "tools": []any{imageNamespace(), codeNamespace()}, + }, + map[string]any{ + "type": "additional_tools", + "tools": []any{imageNamespace()}, + }, + }, + "tool_choice": map[string]any{"type": "namespace", "name": "image_gen"}, + } + + require.True(t, stripOpenAIImageGenerationTools(reqBody)) + require.False(t, hasOpenAIImageGenerationTool(reqBody)) + require.NotContains(t, reqBody, "tool_choice") + + tools, ok := reqBody["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 2) + firstTool, ok := tools[0].(map[string]any) + require.True(t, ok) + secondTool, ok := tools[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "shell", firstTool["name"]) + require.Equal(t, "code_tools", secondTool["name"]) + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + require.Len(t, input, 2) + message, ok := input[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "message", message["type"]) + additionalToolsItem, ok := input[1].(map[string]any) + require.True(t, ok) + additionalTools, ok := additionalToolsItem["tools"].([]any) + require.True(t, ok) + require.Len(t, additionalTools, 1) + additionalTool, ok := additionalTools[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "code_tools", additionalTool["name"]) + require.False(t, stripOpenAIImageGenerationTools(reqBody), "stripping should be idempotent") +} + +func TestStripOpenAIImageGenerationTools_KeepsNonImageNamespaces(t *testing.T) { + reqBody := map[string]any{ + "tools": []any{ + map[string]any{"type": "namespace", "name": "code_tools"}, + }, + "input": []any{ + map[string]any{ + "type": "additional_tools", + "tools": []any{ + map[string]any{"type": "namespace", "name": "browser_tools"}, + }, + }, + }, + "tool_choice": "auto", + } + + require.False(t, stripOpenAIImageGenerationTools(reqBody)) + require.Equal(t, "auto", reqBody["tool_choice"]) + require.False(t, hasOpenAIImageGenerationTool(reqBody)) +} + +func TestStripOpenAIImageGenerationTools_KeepsCustomImagegenFunctionChoice(t *testing.T) { + reqBody := map[string]any{ + "tool_choice": map[string]any{ + "function": map[string]any{"name": "imagegen"}, + }, + } + + require.False(t, stripOpenAIImageGenerationTools(reqBody)) + require.Contains(t, reqBody, "tool_choice") +} + // Non-spark Codex models support image_generation; the tool must be preserved. func TestApplyCodexOAuthTransform_KeepsImageGenerationToolForNonSpark(t *testing.T) { reqBody := map[string]any{ diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 974f8eefaa..ee18302056 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -61,6 +61,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco setOpenAICompatMessagesBridgeContext(c, compatMessagesBridge) isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI) + codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow + if isCodexCLI { + codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() + } wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account) clientTransport := GetOpenAIClientTransport(c) // 仅允许 WS 入站请求走 WS 上游,避免出现 HTTP -> WS 协议混用。 @@ -95,6 +99,17 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } passthroughEnabled := account.IsOpenAIPassthroughEnabled() if passthroughEnabled { + if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { + strippedBody, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(body) + if stripErr != nil { + return nil, stripErr + } + if changed { + body = strippedBody + originalBody = strippedBody + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Stripped /responses image_generation tool for Codex client by account policy") + } + } // 透传分支只需要轻量提取字段,避免热路径全量 Unmarshal。 mappedModel := account.GetMappedModel(reqModel) reasoningEffort := extractOpenAIReasoningEffortFromBody(body, mappedModel) @@ -159,10 +174,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if apiKey != nil { imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group) } - codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow - if isCodexCLI { - codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() - } codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) var imageIntent bool if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { @@ -272,7 +283,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco markDecodedModified() logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions") } - } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) { + } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationDeclaration(body) { // 完整 image_generation tool 只做 raw 计费读取,校验/桥接/旧字段迁移命中时才展开大 input map。 logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type) } @@ -292,7 +303,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco // gpt-5.3-codex-spark also rejects the image_generation tool (HTTP 400, // param=tools). Strip it here so both APIKey and OAuth /responses paths are // covered regardless of the image-generation feature gate. - if isCodexSparkModel(upstreamModel) && openAIRequestBodyHasImageGenerationTool(body) { + if isCodexSparkModel(upstreamModel) && openAIRequestBodyHasImageGenerationDeclaration(body) { decoded, decodeErr := ensureReqBody() if decodeErr != nil { return nil, decodeErr diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index 31edd36097..af0cdf669c 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "encoding/json" "io" "net/http" "net/http/httptest" @@ -191,6 +192,64 @@ func TestOpenAIGatewayServiceForward_AccountPolicyStripsExplicitImageTool(t *tes require.NotContains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_AccountPolicyStripsImageNamespaceTools(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + passthrough bool + }{ + {name: "managed forwarding"}, + {name: "passthrough forwarding", passthrough: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_stripped_namespace","model":"gpt-5.5","usage":{"input_tokens":2,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + c, _ := newOpenAIImageGenerationControlTestContext(false, "codex_cli_rs/0.144.1") + account := newOpenAIImageGenerationControlTestAccount() + account.Extra = map[string]any{ + featureKeyCodexImageGenerationExplicitToolPolicy: codexImageGenerationExplicitToolPolicyStrip, + "openai_passthrough": tt.passthrough, + } + body := []byte(`{ + "model":"gpt-5.5", + "stream":false, + "tools":[ + {"type":"function","name":"shell","parameters":{"type":"object"}}, + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}, + {"type":"namespace","name":"code_tools","tools":[{"type":"function","name":"run"}]} + ], + "input":[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"write code"}]}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]} + ], + "tool_choice":"auto" + }`) + + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + var forwarded map[string]any + require.NoError(t, json.Unmarshal(upstream.lastBody, &forwarded)) + require.False(t, hasOpenAIImageGenerationTool(forwarded)) + require.Equal(t, "auto", forwarded["tool_choice"]) + require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(name=="shell")`).Exists()) + require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(name=="code_tools")`).Exists()) + require.Equal(t, "write code", gjson.GetBytes(upstream.lastBody, "input.0.content.0.text").String()) + }) + } +} + func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 4a5fa20774..45e004b6d0 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -270,7 +270,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( normalized = next } if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { - if stripped, changed, stripErr := stripOpenAIImageGenerationToolFromRawPayload(normalized); stripErr != nil { + if stripped, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(normalized); stripErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr) } else if changed { normalized = stripped diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index 1d18c46fca..7753ea9598 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -153,6 +153,24 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) { require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) }) + t.Run("strips_namespace_tools_for_spark", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.3-codex-spark", + "input":[ + {"type":"message","role":"user","content":"hello"}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen"}]} + ], + "tool_choice":{"type":"namespace","name":"image_gen"} + }`) + updated, changed, err := stripCodexSparkImageGenerationToolFromRawPayload(payload, "gpt-5.3-codex-spark") + require.NoError(t, err) + require.True(t, changed) + require.False(t, IsImageGenerationIntent(openAIResponsesEndpoint, "gpt-5.3-codex-spark", updated)) + require.Equal(t, "hello", gjson.GetBytes(updated, "input.0.content").String()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + t.Run("keeps_image_generation_for_non_spark", func(t *testing.T) { payload := []byte(`{"type":"response.create","model":"gpt-5.3-codex","tools":[{"type":"image_generation","output_format":"png"}]}`) updated, changed, err := stripCodexSparkImageGenerationToolFromRawPayload(payload, "gpt-5.3-codex") @@ -170,24 +188,61 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) { }) } -func TestStripOpenAIImageGenerationToolFromRawPayload(t *testing.T) { - payload := []byte(`{ - "type":"response.create", - "model":"gpt-5.4", - "tools":[ - {"type":"function","name":"shell"}, - {"type":"image_generation","output_format":"png"} - ], - "tool_choice":{"type":"image_generation"} - }`) +func TestStripOpenAIImageGenerationToolsFromRawPayload(t *testing.T) { + t.Run("flat image tool", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.4", + "tools":[ + {"type":"function","name":"shell"}, + {"type":"image_generation","output_format":"png"} + ], + "tool_choice":{"type":"image_generation"} + }`) - updated, changed, err := stripOpenAIImageGenerationToolFromRawPayload(payload) + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) - require.NoError(t, err) - require.True(t, changed) - require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists()) - require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) - require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + require.NoError(t, err) + require.True(t, changed) + require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists()) + require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + + t.Run("namespace and Responses Lite tools", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.5", + "tools":[ + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}, + {"type":"namespace","name":"code_tools","tools":[{"type":"function","name":"run"}]} + ], + "input":[ + {"type":"message","role":"user","content":"hello"}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen"}]} + ], + "tool_choice":{"type":"namespace","name":"image_gen"} + }`) + + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) + + require.NoError(t, err) + require.True(t, changed) + require.False(t, IsImageGenerationIntent(openAIResponsesEndpoint, "gpt-5.5", updated)) + require.True(t, gjson.GetBytes(updated, `tools.#(name=="code_tools")`).Exists()) + require.Equal(t, "hello", gjson.GetBytes(updated, "input.0.content").String()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + + t.Run("non-image namespace is unchanged", func(t *testing.T) { + payload := []byte(`{"type":"response.create","model":"gpt-5.5","tools":[{"type":"namespace","name":"code_tools"}]}`) + + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) + + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, payload, updated) + }) } func TestAlignStoreDisabledPreviousResponseID(t *testing.T) { diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index 24853d4a6c..f0f71648dd 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -3,7 +3,6 @@ package service import ( "bytes" "context" - "encoding/json" "errors" "fmt" "net/http" @@ -710,23 +709,8 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( // Codex clients advertise it by default. Returns the (possibly unchanged) payload, // whether it changed, and any JSON decode error. func stripCodexSparkImageGenerationToolFromRawPayload(payload []byte, model string) ([]byte, bool, error) { - if !isCodexSparkModel(model) || !openAIRequestBodyHasImageGenerationTool(payload) { + if !isCodexSparkModel(model) { return payload, false, nil } - return stripOpenAIImageGenerationToolFromRawPayload(payload) -} - -func stripOpenAIImageGenerationToolFromRawPayload(payload []byte) ([]byte, bool, error) { - payloadMap := make(map[string]any) - if err := json.Unmarshal(payload, &payloadMap); err != nil { - return payload, false, err - } - if !stripOpenAIImageGenerationTools(payloadMap) { - return payload, false, nil - } - rebuilt, err := json.Marshal(payloadMap) - if err != nil { - return payload, false, err - } - return rebuilt, true, nil + return stripOpenAIImageGenerationToolsFromRawPayload(payload) } From de28eba3c8cdc0c9c235e68b3b9540252d4cb955 Mon Sep 17 00:00:00 2001 From: superman2003 <2112076433zcr@gmail.com> Date: Fri, 10 Jul 2026 15:13:11 +0800 Subject: [PATCH 23/27] fix(openai): harden GPT-5.6 billing and usage --- .../chatcompletions_responses_test.go | 20 ++++ backend/internal/pkg/apicompat/types.go | 25 +++++ backend/internal/pkg/openai/constants.go | 1 + backend/internal/pkg/openai/constants_test.go | 11 ++ backend/internal/repository/usage_log_repo.go | 16 ++- .../usage_log_repo_breakdown_test.go | 29 +++++ .../repository/usage_log_repo_trend.go | 5 +- backend/internal/service/billing_service.go | 15 ++- .../openai_gateway_chat_completions_raw.go | 13 +-- ...penai_gateway_chat_completions_raw_test.go | 51 +++++++++ .../service/openai_gateway_messages.go | 7 -- .../openai_gateway_messages_usage_test.go | 28 +++++ .../openai_gateway_response_handling.go | 18 ++- .../service/openai_gateway_service_test.go | 8 ++ .../internal/service/openai_model_alias.go | 14 +++ .../service/openai_model_alias_test.go | 36 ++++++ .../service/openai_model_mapping_test.go | 18 +++ .../service/openai_ws_v2/passthrough_relay.go | 7 ++ .../passthrough_relay_internal_test.go | 7 ++ backend/internal/service/pricing_service.go | 27 +++++ .../internal/service/pricing_service_test.go | 106 +++++++++++++++--- ...allow_cyber_blocked_usage_request_type.sql | 8 ++ ...tity_payment_migrations_regression_test.go | 27 +++++ .../model_prices_and_context_window.json | 9 ++ frontend/src/components/keys/UseKeyModal.vue | 26 ++++- .../keys/__tests__/UseKeyModal.spec.ts | 37 ++++++ .../__tests__/useModelWhitelist.spec.ts | 1 + frontend/src/composables/useModelWhitelist.ts | 3 +- 28 files changed, 524 insertions(+), 49 deletions(-) create mode 100644 backend/internal/pkg/openai/constants_test.go create mode 100644 backend/internal/service/openai_gateway_messages_usage_test.go create mode 100644 backend/internal/service/openai_model_alias_test.go create mode 100644 backend/migrations/173_allow_cyber_blocked_usage_request_type.sql diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index d63ce664a7..4075f57791 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -53,6 +53,26 @@ func TestUsageConversionsPreserveCacheWriteTokens(t *testing.T) { require.Equal(t, 200, roundTrip.InputTokensDetails.CacheWriteTokens) } +func TestResponsesUsageNestedCacheWritePresenceOverridesTopLevelAlias(t *testing.T) { + tests := []struct { + name string + nestedJSON string + want int + }{ + {name: "explicit zero", nestedJSON: `{"cache_write_tokens":0}`, want: 0}, + {name: "nonzero", nestedJSON: `{"cache_write_tokens":7}`, want: 7}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var usage ResponsesUsage + payload := []byte(`{"input_tokens":20,"output_tokens":2,"cache_creation_input_tokens":19,"input_tokens_details":` + tt.nestedJSON + `}`) + require.NoError(t, json.Unmarshal(payload, &usage)) + require.Equal(t, tt.want, usage.CacheCreationInputTokens) + }) + } +} + func TestChatCompletionsToResponses_SystemMessage(t *testing.T) { req := &ChatCompletionsRequest{ Model: "gpt-4o", diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index c020be94f8..8d96a1d3d3 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -332,6 +332,10 @@ type ResponsesUsage struct { func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { type responsesUsageAlias ResponsesUsage + type cacheTokenPresence struct { + CacheCreationTokens *int `json:"cache_creation_tokens"` + CacheWriteTokens *int `json:"cache_write_tokens"` + } var aux struct { responsesUsageAlias PromptTokens int `json:"prompt_tokens"` @@ -345,6 +349,13 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { if err := json.Unmarshal(data, &aux); err != nil { return err } + var nestedPresence struct { + InputTokensDetails *cacheTokenPresence `json:"input_tokens_details"` + PromptTokensDetails *cacheTokenPresence `json:"prompt_tokens_details"` + } + if err := json.Unmarshal(data, &nestedPresence); err != nil { + return err + } *u = ResponsesUsage(aux.responsesUsageAlias) if u.InputTokens == 0 && aux.PromptTokens != 0 { u.InputTokens = aux.PromptTokens @@ -368,6 +379,20 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { if u.OutputTokensDetails == nil && aux.CompletionTokensDetails != nil { u.OutputTokensDetails = aux.CompletionTokensDetails } + var canonicalCacheCreationTokens *int + switch { + case nestedPresence.InputTokensDetails != nil && nestedPresence.InputTokensDetails.CacheWriteTokens != nil: + canonicalCacheCreationTokens = nestedPresence.InputTokensDetails.CacheWriteTokens + case nestedPresence.PromptTokensDetails != nil && nestedPresence.PromptTokensDetails.CacheWriteTokens != nil: + canonicalCacheCreationTokens = nestedPresence.PromptTokensDetails.CacheWriteTokens + case nestedPresence.InputTokensDetails != nil && nestedPresence.InputTokensDetails.CacheCreationTokens != nil: + canonicalCacheCreationTokens = nestedPresence.InputTokensDetails.CacheCreationTokens + case nestedPresence.PromptTokensDetails != nil && nestedPresence.PromptTokensDetails.CacheCreationTokens != nil: + canonicalCacheCreationTokens = nestedPresence.PromptTokensDetails.CacheCreationTokens + } + if canonicalCacheCreationTokens != nil { + u.CacheCreationInputTokens = max(*canonicalCacheCreationTokens, 0) + } if u.TotalTokens == 0 && (u.InputTokens != 0 || u.OutputTokens != 0) { u.TotalTokens = u.InputTokens + u.OutputTokens } diff --git a/backend/internal/pkg/openai/constants.go b/backend/internal/pkg/openai/constants.go index c9d391df4e..79d799d774 100644 --- a/backend/internal/pkg/openai/constants.go +++ b/backend/internal/pkg/openai/constants.go @@ -18,6 +18,7 @@ type Model struct { // DefaultModels OpenAI models list var DefaultModels = []Model{ + {ID: "gpt-5.6", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 (Sol)"}, {ID: "gpt-5.6-sol", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Sol"}, {ID: "gpt-5.6-terra", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Terra"}, {ID: "gpt-5.6-luna", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Luna"}, diff --git a/backend/internal/pkg/openai/constants_test.go b/backend/internal/pkg/openai/constants_test.go new file mode 100644 index 0000000000..f1f59b70f2 --- /dev/null +++ b/backend/internal/pkg/openai/constants_test.go @@ -0,0 +1,11 @@ +package openai + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestDefaultModelsIncludeBareGPT56Alias(t *testing.T) { + require.Contains(t, DefaultModelIDs(), "gpt-5.6") +} diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index 341bdff57d..4359edc2ca 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -199,16 +199,24 @@ func appendRequestTypeOrStreamQueryFilter(query string, args []any, requestType // buildRequestTypeFilterCondition 在 request_type 过滤时兼容 legacy 字段,避免历史数据漏查。 func buildRequestTypeFilterCondition(startArgIndex int, requestType int16) (string, []any) { + return buildRequestTypeFilterConditionWithAlias(startArgIndex, requestType, "") +} + +func buildRequestTypeFilterConditionWithAlias(startArgIndex int, requestType int16, alias string) (string, []any) { normalized := service.RequestTypeFromInt16(requestType) requestTypeArg := int16(normalized) + prefix := "" + if alias != "" { + prefix = alias + "." + } switch normalized { case service.RequestTypeSync: - return fmt.Sprintf("(request_type = $%d OR (request_type = %d AND stream = FALSE AND openai_ws_mode = FALSE))", startArgIndex, int16(service.RequestTypeUnknown)), []any{requestTypeArg} + return fmt.Sprintf("(%srequest_type = $%d OR (%srequest_type = %d AND %sstream = FALSE AND %sopenai_ws_mode = FALSE))", prefix, startArgIndex, prefix, int16(service.RequestTypeUnknown), prefix, prefix), []any{requestTypeArg} case service.RequestTypeStream: - return fmt.Sprintf("(request_type = $%d OR (request_type = %d AND stream = TRUE AND openai_ws_mode = FALSE))", startArgIndex, int16(service.RequestTypeUnknown)), []any{requestTypeArg} + return fmt.Sprintf("(%srequest_type = $%d OR (%srequest_type = %d AND %sstream = TRUE AND %sopenai_ws_mode = FALSE))", prefix, startArgIndex, prefix, int16(service.RequestTypeUnknown), prefix, prefix), []any{requestTypeArg} case service.RequestTypeWSV2: - return fmt.Sprintf("(request_type = $%d OR (request_type = %d AND openai_ws_mode = TRUE))", startArgIndex, int16(service.RequestTypeUnknown)), []any{requestTypeArg} + return fmt.Sprintf("(%srequest_type = $%d OR (%srequest_type = %d AND %sopenai_ws_mode = TRUE))", prefix, startArgIndex, prefix, int16(service.RequestTypeUnknown), prefix), []any{requestTypeArg} default: - return fmt.Sprintf("request_type = $%d", startArgIndex), []any{requestTypeArg} + return fmt.Sprintf("%srequest_type = $%d", prefix, startArgIndex), []any{requestTypeArg} } } diff --git a/backend/internal/repository/usage_log_repo_breakdown_test.go b/backend/internal/repository/usage_log_repo_breakdown_test.go index da62e8dd56..1fb99b0d64 100644 --- a/backend/internal/repository/usage_log_repo_breakdown_test.go +++ b/backend/internal/repository/usage_log_repo_breakdown_test.go @@ -3,9 +3,14 @@ package repository import ( + "context" + "regexp" "testing" + "time" + "github.com/DATA-DOG/go-sqlmock" "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/stretchr/testify/require" ) @@ -48,3 +53,27 @@ func TestResolveModelDimensionExpression(t *testing.T) { }) } } + +func TestGetUserBreakdownStatsRequestTypeIncludesLegacyFallback(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + start := time.Date(2026, 7, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(24 * time.Hour) + requestType := int16(service.RequestTypeStream) + + legacyFilter := `(ul.request_type = $3 OR (ul.request_type = 0 AND ul.stream = TRUE AND ul.openai_ws_mode = FALSE))` + mock.ExpectQuery(regexp.QuoteMeta(legacyFilter)). + WithArgs(start, end, requestType). + WillReturnRows(sqlmock.NewRows([]string{ + "user_id", "email", "requests", "input_tokens", "output_tokens", + "cache_tokens", "total_tokens", "cost", "actual_cost", "account_cost", + })) + + rows, err := repo.GetUserBreakdownStats(context.Background(), start, end, usagestats.UserBreakdownDimension{ + RequestType: &requestType, + }, 0) + + require.NoError(t, err) + require.Empty(t, rows) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/usage_log_repo_trend.go b/backend/internal/repository/usage_log_repo_trend.go index 7302296683..f2e0045ff5 100644 --- a/backend/internal/repository/usage_log_repo_trend.go +++ b/backend/internal/repository/usage_log_repo_trend.go @@ -642,8 +642,9 @@ func (r *usageLogRepository) GetUserBreakdownStats(ctx context.Context, startTim args = append(args, dim.AccountID) } if dim.RequestType != nil { - query += fmt.Sprintf(" AND ul.request_type = $%d", len(args)+1) - args = append(args, *dim.RequestType) + condition, conditionArgs := buildRequestTypeFilterConditionWithAlias(len(args)+1, *dim.RequestType, "ul") + query += " AND " + condition + args = append(args, conditionArgs...) } if dim.Stream != nil { query += fmt.Sprintf(" AND ul.stream = $%d", len(args)+1) diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 71d74ae6e0..3dfd500b05 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -293,6 +293,9 @@ func (s *BillingService) initFallbackPricing() { CacheCreationPricePerTokenPriority: 12.5e-6, CacheReadPricePerToken: 0.5e-6, CacheReadPricePerTokenPriority: 1e-6, + LongContextInputThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, } s.fallbackPrices["gpt-5.6-terra"] = &ModelPricing{ InputPricePerToken: 2.5e-6, @@ -303,6 +306,9 @@ func (s *BillingService) initFallbackPricing() { CacheCreationPricePerTokenPriority: 6.25e-6, CacheReadPricePerToken: 0.25e-6, CacheReadPricePerTokenPriority: 0.5e-6, + LongContextInputThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, } s.fallbackPrices["gpt-5.6-luna"] = &ModelPricing{ InputPricePerToken: 1e-6, @@ -313,6 +319,9 @@ func (s *BillingService) initFallbackPricing() { CacheCreationPricePerTokenPriority: 2.5e-6, CacheReadPricePerToken: 0.1e-6, CacheReadPricePerTokenPriority: 0.2e-6, + LongContextInputThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, } s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{ @@ -1100,7 +1109,7 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * if !isGPT56 && !usesLegacyLongContextPricing { return pricing } - needsLongContextPolicy := usesLegacyLongContextPricing && + needsLongContextPolicy := (isGPT56 || usesLegacyLongContextPricing) && (pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0) needsCacheCreationPolicy := isGPT56 && !pricing.CacheCreationPriceExplicit && (pricing.CacheCreationPricePerToken <= 0 || (pricing.InputPricePerTokenPriority > 0 && pricing.CacheCreationPricePerTokenPriority <= 0)) @@ -1108,7 +1117,7 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * return pricing } cloned := *pricing - if isGPT56 { + if isGPT56 && !cloned.CacheCreationPriceExplicit { if cloned.CacheCreationPricePerToken <= 0 { cloned.CacheCreationPricePerToken = cloned.InputPricePerToken * 1.25 } @@ -1116,7 +1125,7 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing * cloned.CacheCreationPricePerTokenPriority = cloned.InputPricePerTokenPriority * 1.25 } } - if usesLegacyLongContextPricing { + if isGPT56 || usesLegacyLongContextPricing { if cloned.LongContextInputThreshold <= 0 { cloned.LongContextInputThreshold = openAIGPT54LongContextInputThreshold } diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 28c95f7ff8..6def67c7ec 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -2,14 +2,12 @@ package service import ( "context" - "encoding/json" "errors" "fmt" "net/http" "strings" "time" - "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" @@ -411,16 +409,9 @@ func (s *OpenAIGatewayService) bufferRawChatCompletions( return nil, fmt.Errorf("read upstream body: %w", err) } - var ccResp apicompat.ChatCompletionsResponse var usage OpenAIUsage - if err := json.Unmarshal(respBody, &ccResp); err == nil && ccResp.Usage != nil { - usage = OpenAIUsage{ - InputTokens: ccResp.Usage.PromptTokens, - OutputTokens: ccResp.Usage.CompletionTokens, - } - if ccResp.Usage.PromptTokensDetails != nil { - usage.CacheReadInputTokens = ccResp.Usage.PromptTokensDetails.CachedTokens - } + if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(respBody); ok { + usage = parsedUsage } if s.responseHeaderFilter != nil { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw_test.go b/backend/internal/service/openai_gateway_chat_completions_raw_test.go index 5f4b5bcd2d..61663e636a 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw_test.go @@ -155,6 +155,57 @@ func TestForwardAsRawChatCompletions_PreservesMappedGPT56MaxEffort(t *testing.T) require.Equal(t, "max", *result.ReasoningEffort) } +func TestForwardAsRawChatCompletions_NonStreamingCapturesCacheWriteUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + usageJSON string + wantWrite int + }{ + { + name: "positive cache write", + usageJSON: `{"prompt_tokens":12,"completion_tokens":3,"total_tokens":15,"prompt_tokens_details":{"cached_tokens":4,"cache_write_tokens":6}}`, + wantWrite: 6, + }, + { + name: "nested zero overrides legacy alias", + usageJSON: `{"prompt_tokens":12,"completion_tokens":3,"total_tokens":15,"cache_creation_input_tokens":19,"prompt_tokens_details":{"cached_tokens":4,"cache_write_tokens":0}}`, + wantWrite: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(`{"model":"gpt-5.6","messages":[{"role":"user","content":"hello"}],"stream":false}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"chatcmpl_cache","object":"chat.completion","model":"gpt-5.6","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":` + tt.usageJSON + `}`, + )), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + } + + result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "") + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 12, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.CacheReadInputTokens) + require.Equal(t, tt.wantWrite, result.Usage.CacheCreationInputTokens) + }) + } +} + func TestForwardAsRawChatCompletions_PreservesDeepSeekReasoningContentNonStreaming(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 52ec34f27e..b1e99f8da9 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -1120,13 +1120,6 @@ func copyOpenAIUsageFromResponsesUsage(usage *apicompat.ResponsesUsage) OpenAIUs } if usage.InputTokensDetails != nil { result.CacheReadInputTokens = usage.InputTokensDetails.CachedTokens - if result.CacheCreationInputTokens == 0 { - if usage.InputTokensDetails.CacheWriteTokens > 0 { - result.CacheCreationInputTokens = usage.InputTokensDetails.CacheWriteTokens - } else { - result.CacheCreationInputTokens = usage.InputTokensDetails.CacheCreationTokens - } - } } return result } diff --git a/backend/internal/service/openai_gateway_messages_usage_test.go b/backend/internal/service/openai_gateway_messages_usage_test.go new file mode 100644 index 0000000000..f884e74595 --- /dev/null +++ b/backend/internal/service/openai_gateway_messages_usage_test.go @@ -0,0 +1,28 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/stretchr/testify/require" +) + +func TestCopyOpenAIUsageFromResponsesUsageTrustsCanonicalCacheCreationValue(t *testing.T) { + usage := &apicompat.ResponsesUsage{ + InputTokens: 20, + OutputTokens: 2, + CacheCreationInputTokens: 0, + InputTokensDetails: &apicompat.ResponsesInputTokensDetails{ + CachedTokens: 3, + CacheWriteTokens: 19, + }, + } + + got := copyOpenAIUsageFromResponsesUsage(usage) + + require.Equal(t, 20, got.InputTokens) + require.Equal(t, 3, got.CacheReadInputTokens) + require.Zero(t, got.CacheCreationInputTokens) +} diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 313a255341..5841ce6556 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -772,9 +772,16 @@ func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) { } func openAICacheReadTokensFromUsage(value gjson.Result) int { - return firstPositiveGJSONInt( + for _, nested := range []gjson.Result{ value.Get("input_tokens_details.cached_tokens"), value.Get("prompt_tokens_details.cached_tokens"), + } { + if nested.Exists() { + return max(int(nested.Int()), 0) + } + } + + return firstPositiveGJSONInt( value.Get("cache_read_input_tokens"), value.Get("cache_read_tokens"), value.Get("cached_tokens"), @@ -782,11 +789,18 @@ func openAICacheReadTokensFromUsage(value gjson.Result) int { } func openAICacheCreationTokensFromUsage(value gjson.Result) int { - return firstPositiveGJSONInt( + for _, nested := range []gjson.Result{ value.Get("input_tokens_details.cache_write_tokens"), value.Get("prompt_tokens_details.cache_write_tokens"), value.Get("input_tokens_details.cache_creation_tokens"), value.Get("prompt_tokens_details.cache_creation_tokens"), + } { + if nested.Exists() { + return max(int(nested.Int()), 0) + } + } + + return firstPositiveGJSONInt( value.Get("cache_write_tokens"), value.Get("cache_creation_input_tokens"), value.Get("cache_write_input_tokens"), diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 3fdef698db..b33918652f 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2821,6 +2821,14 @@ func TestExtractOpenAIUsageFromJSONBytes_AcceptsResponseAndChatUsageShapes(t *te usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":20,"output_tokens":2,"cache_creation_input_tokens":19,"input_tokens_details":{"cache_write_tokens":7}}}`)) require.True(t, ok) require.Equal(t, 7, usage.CacheCreationInputTokens, "官方嵌套字段应优先于兼容顶层别名") + + usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":20,"output_tokens":2,"cache_creation_input_tokens":19,"input_tokens_details":{"cache_write_tokens":0}}}`)) + require.True(t, ok) + require.Zero(t, usage.CacheCreationInputTokens, "官方嵌套字段显式为零时仍应优先于兼容顶层别名") + + usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":20,"output_tokens":2,"cache_read_input_tokens":19,"input_tokens_details":{"cached_tokens":0}}}`)) + require.True(t, ok) + require.Zero(t, usage.CacheReadInputTokens, "官方嵌套缓存读取字段显式为零时仍应优先于兼容顶层别名") } func TestExtractCodexFinalResponse_SampleReplay(t *testing.T) { diff --git a/backend/internal/service/openai_model_alias.go b/backend/internal/service/openai_model_alias.go index 338de9e8a5..4ddf8110c3 100644 --- a/backend/internal/service/openai_model_alias.go +++ b/backend/internal/service/openai_model_alias.go @@ -71,6 +71,14 @@ func normalizeKnownOpenAICodexModel(model string) string { return "gpt-5.6-terra" case strings.Contains(normalized, "gpt-5.6-luna"): return "gpt-5.6-luna" + case normalized == "gpt-5.6": + return "gpt-5.6-sol" + case strings.HasPrefix(normalized, "gpt-5.6-"): + suffix := strings.TrimPrefix(normalized, "gpt-5.6-") + if suffix == "max" || isKnownCodexModelSuffix(suffix) { + return "gpt-5.6-sol" + } + return "" case strings.Contains(normalized, "gpt-5.5-pro"): return "gpt-5.5-pro" case strings.Contains(normalized, "gpt-5.5"): @@ -102,6 +110,12 @@ func normalizeKnownOpenAICodexModel(model string) string { // (含大小写/路径/后缀变体)或已归一化的基名,两者均能正确识别。 func isOpenAIGPT56Model(model string) bool { normalized := canonicalizeOpenAIModelAliasSpelling(model) + if normalized == "gpt-5.6" { + return true + } + if suffix, ok := strings.CutPrefix(normalized, "gpt-5.6-"); ok && (suffix == "max" || isKnownCodexModelSuffix(suffix)) { + return true + } for _, prefix := range []string{"gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna"} { if normalized == prefix || strings.HasPrefix(normalized, prefix+"-") { return true diff --git a/backend/internal/service/openai_model_alias_test.go b/backend/internal/service/openai_model_alias_test.go new file mode 100644 index 0000000000..4cb27206c9 --- /dev/null +++ b/backend/internal/service/openai_model_alias_test.go @@ -0,0 +1,36 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNormalizeKnownOpenAICodexModel_BareGPT56RoutesToSol(t *testing.T) { + tests := map[string]string{ + "gpt-5.6": "gpt-5.6-sol", + "openai/gpt-5.6": "gpt-5.6-sol", + "gpt5.6": "gpt-5.6-sol", + "gpt-5.6-high": "gpt-5.6-sol", + "gpt-5.6-max": "gpt-5.6-sol", + "gpt-5.6-2026-07-09": "gpt-5.6-sol", + "openai/gpt-5.6-max": "gpt-5.6-sol", + } + + for input, expected := range tests { + t.Run(input, func(t *testing.T) { + require.Equal(t, expected, normalizeKnownOpenAICodexModel(input)) + }) + } +} + +func TestUsageBillingModelCandidates_BareGPT56IncludesSol(t *testing.T) { + require.Equal(t, + []string{"gpt-5.6", "gpt-5.6-sol"}, + usageBillingModelCandidates("gpt-5.6"), + ) + require.Equal(t, + []string{"openai/gpt-5.6", "gpt-5.6", "gpt-5.6-sol"}, + usageBillingModelCandidates("openai/gpt-5.6"), + ) +} diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index 0f65740be7..f2ceb3551c 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -252,6 +252,18 @@ func TestNormalizeOpenAIModelForUpstream(t *testing.T) { model string want string }{ + { + name: "oauth routes bare GPT-5.6 alias to Sol", + account: &Account{Type: AccountTypeOAuth}, + model: "gpt-5.6", + want: "gpt-5.6-sol", + }, + { + name: "oauth routes provider-prefixed GPT-5.6 alias to Sol", + account: &Account{Type: AccountTypeOAuth}, + model: "openai/gpt-5.6", + want: "gpt-5.6-sol", + }, { name: "oauth preserves unknown non codex model", account: &Account{Type: AccountTypeOAuth}, @@ -282,6 +294,12 @@ func TestNormalizeOpenAIModelForUpstream(t *testing.T) { model: "codex-auto-review", want: "codex-auto-review", }, + { + name: "apikey preserves official bare GPT-5.6 alias", + account: &Account{Type: AccountTypeAPIKey}, + model: "gpt-5.6", + want: "gpt-5.6", + }, { name: "apikey preserves custom compatible model", account: &Account{Type: AccountTypeAPIKey}, diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index be19438eaa..d41abaac3a 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -816,6 +816,13 @@ func openAICacheCreationTokensFromUsage(value gjson.Result) int { "prompt_tokens_details.cache_write_tokens", "input_tokens_details.cache_creation_tokens", "prompt_tokens_details.cache_creation_tokens", + } { + result := value.Get(field) + if result.Exists() { + return max(int(result.Int()), 0) + } + } + for _, field := range []string{ "cache_write_tokens", "cache_creation_input_tokens", "cache_write_input_tokens", diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go index cb2bd9cc31..f92117a2bf 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go @@ -335,6 +335,13 @@ func TestParseUsageAndAccumulateAcceptsChatUsageAliases(t *testing.T) { require.Equal(t, got, state.usage) } +func TestOpenAICacheCreationTokensFromUsageNestedZeroWins(t *testing.T) { + t.Parallel() + + usage := gjson.Parse(`{"input_tokens_details":{"cache_write_tokens":0},"cache_creation_input_tokens":19}`) + require.Zero(t, openAICacheCreationTokensFromUsage(usage)) +} + func TestEmitTurnCompleteCoverage(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index 4528a8886f..d0ae3c7d96 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -44,6 +44,9 @@ var ( CacheCreationInputTokenCostPriority: 1.25e-05, CacheReadInputTokenCost: 5e-07, CacheReadInputTokenCostPriority: 1e-06, + LongContextInputTokenThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputCostMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputCostMultiplier: openAIGPT54LongContextOutputMultiplier, SupportsServiceTier: true, LiteLLMProvider: "openai", Mode: "chat", @@ -58,6 +61,9 @@ var ( CacheCreationInputTokenCostPriority: 6.25e-06, CacheReadInputTokenCost: 2.5e-07, CacheReadInputTokenCostPriority: 5e-07, + LongContextInputTokenThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputCostMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputCostMultiplier: openAIGPT54LongContextOutputMultiplier, SupportsServiceTier: true, LiteLLMProvider: "openai", Mode: "chat", @@ -72,6 +78,9 @@ var ( CacheCreationInputTokenCostPriority: 2.5e-06, CacheReadInputTokenCost: 1e-07, CacheReadInputTokenCostPriority: 2e-07, + LongContextInputTokenThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputCostMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputCostMultiplier: openAIGPT54LongContextOutputMultiplier, SupportsServiceTier: true, LiteLLMProvider: "openai", Mode: "chat", @@ -140,6 +149,9 @@ type LiteLLMRawEntry struct { CacheCreationInputTokenCostAbove1hr *float64 `json:"cache_creation_input_token_cost_above_1hr"` CacheReadInputTokenCost *float64 `json:"cache_read_input_token_cost"` CacheReadInputTokenCostPriority *float64 `json:"cache_read_input_token_cost_priority"` + LongContextInputTokenThreshold *int `json:"long_context_input_token_threshold"` + LongContextInputCostMultiplier *float64 `json:"long_context_input_cost_multiplier"` + LongContextOutputCostMultiplier *float64 `json:"long_context_output_cost_multiplier"` SupportsServiceTier bool `json:"supports_service_tier"` LiteLLMProvider string `json:"litellm_provider"` Mode string `json:"mode"` @@ -462,6 +474,15 @@ func (s *PricingService) parsePricingData(body []byte) (map[string]*LiteLLMModel if entry.CacheReadInputTokenCostPriority != nil { pricing.CacheReadInputTokenCostPriority = *entry.CacheReadInputTokenCostPriority } + if entry.LongContextInputTokenThreshold != nil { + pricing.LongContextInputTokenThreshold = *entry.LongContextInputTokenThreshold + } + if entry.LongContextInputCostMultiplier != nil { + pricing.LongContextInputCostMultiplier = *entry.LongContextInputCostMultiplier + } + if entry.LongContextOutputCostMultiplier != nil { + pricing.LongContextOutputCostMultiplier = *entry.LongContextOutputCostMultiplier + } if entry.OutputCostPerImage != nil { pricing.OutputCostPerImage = *entry.OutputCostPerImage } @@ -713,6 +734,12 @@ func normalizeModelNameForPricing(model string) string { model = strings.TrimLeft(model, "/") if canonical := canonicalizeOpenAIModelAliasSpelling(model); canonical != "" { + if canonical == "gpt-5.6" { + return "gpt-5.6-sol" + } + if suffix, ok := strings.CutPrefix(canonical, "gpt-5.6-"); ok && (suffix == "max" || isKnownCodexModelSuffix(suffix)) { + return "gpt-5.6-sol" + } return canonical } return model diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index 19098df917..b68d034638 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -22,6 +22,9 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { "cache_creation_input_token_cost_priority": 0.000005, "cache_read_input_token_cost": 0.00000025, "cache_read_input_token_cost_priority": 0.0000005, + "long_context_input_token_threshold": 272000, + "long_context_input_cost_multiplier": 2, + "long_context_output_cost_multiplier": 1.5, "supports_service_tier": true, "supports_prompt_caching": true, "litellm_provider": "openai", @@ -37,6 +40,9 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { require.InDelta(t, 3e-5, pricing.OutputCostPerTokenPriority, 1e-12) require.InDelta(t, 5e-6, pricing.CacheCreationInputTokenCostPriority, 1e-12) require.InDelta(t, 5e-7, pricing.CacheReadInputTokenCostPriority, 1e-12) + require.Equal(t, 272000, pricing.LongContextInputTokenThreshold) + require.InDelta(t, 2.0, pricing.LongContextInputCostMultiplier, 1e-12) + require.InDelta(t, 1.5, pricing.LongContextOutputCostMultiplier, 1e-12) require.True(t, pricing.SupportsServiceTier) } @@ -72,7 +78,9 @@ func TestBillingService_GPT56CacheWritePricingUsesOfficialMultiplier(t *testing. require.NoError(t, err) require.InDelta(t, tt.input*1.25, pricing.CacheCreationPricePerToken, 1e-12) require.InDelta(t, tt.inputPriority*1.25, pricing.CacheCreationPricePerTokenPriority, 1e-12) - require.Zero(t, pricing.LongContextInputThreshold) + require.Equal(t, 272000, pricing.LongContextInputThreshold) + require.InDelta(t, 2.0, pricing.LongContextInputMultiplier, 1e-12) + require.InDelta(t, 1.5, pricing.LongContextOutputMultiplier, 1e-12) tokens := UsageTokens{InputTokens: 700, OutputTokens: 50, CacheCreationTokens: 200, CacheReadTokens: 100} standard, err := svc.CalculateCostWithServiceTier(tt.model, tokens, 1, "") @@ -90,25 +98,87 @@ func TestBillingService_GPT56CacheWritePricingUsesOfficialMultiplier(t *testing. } } -func TestBillingService_GPT56DoesNotUseLegacyLongContextMultiplier(t *testing.T) { - model := "gpt-5.6-sol" - pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ - model: { - InputCostPerToken: 5e-6, - OutputCostPerToken: 30e-6, - CacheReadInputTokenCost: 0.5e-6, - }, - }} - svc := NewBillingService(&config.Config{}, pricingSvc) - tokens := UsageTokens{InputTokens: 100000, CacheCreationTokens: 173000, OutputTokens: 10} +func TestBillingService_GPT56UsesLongContextPricingAcrossModelsAndTiers(t *testing.T) { + models := []struct { + name string + input, cached float64 + cacheWrite, output float64 + }{ + {name: "gpt-5.6-sol", input: 5e-6, cached: 0.5e-6, cacheWrite: 6.25e-6, output: 30e-6}, + {name: "gpt-5.6-terra", input: 2.5e-6, cached: 0.25e-6, cacheWrite: 3.125e-6, output: 15e-6}, + {name: "gpt-5.6-luna", input: 1e-6, cached: 0.1e-6, cacheWrite: 1.25e-6, output: 6e-6}, + } + tiers := []struct { + name string + priceScale float64 + }{ + {name: "standard", priceScale: 1}, + {name: "priority", priceScale: 2}, + {name: "flex", priceScale: 0.5}, + } + tokens := UsageTokens{ + InputTokens: 100000, + CacheCreationTokens: 100000, + CacheReadTokens: 73000, + OutputTokens: 10, + } - cost, err := svc.CalculateCost(model, tokens, 1) + for _, model := range models { + for _, tier := range tiers { + t.Run(model.name+"/"+tier.name, func(t *testing.T) { + svc := NewBillingService(&config.Config{}, nil) + serviceTier := "" + if tier.name != "standard" { + serviceTier = tier.name + } + cost, err := svc.CalculateCostWithServiceTier(model.name, tokens, 1, serviceTier) + require.NoError(t, err) + require.InDelta(t, float64(tokens.InputTokens)*model.input*tier.priceScale*2, cost.InputCost, 1e-12) + require.InDelta(t, float64(tokens.CacheCreationTokens)*model.cacheWrite*tier.priceScale*2, cost.CacheCreationCost, 1e-12) + require.InDelta(t, float64(tokens.CacheReadTokens)*model.cached*tier.priceScale*2, cost.CacheReadCost, 1e-12) + require.InDelta(t, float64(tokens.OutputTokens)*model.output*tier.priceScale*1.5, cost.OutputCost, 1e-12) + }) + } + } +} + +func TestBillingService_GPT56LongContextBoundaryIsExclusive(t *testing.T) { + svc := NewBillingService(&config.Config{}, nil) + tokens := UsageTokens{InputTokens: 100000, CacheCreationTokens: 100000, CacheReadTokens: 72000, OutputTokens: 10} + + cost, err := svc.CalculateCost("gpt-5.6-sol", tokens, 1) require.NoError(t, err) require.InDelta(t, 100000*5e-6, cost.InputCost, 1e-12) - require.InDelta(t, 173000*6.25e-6, cost.CacheCreationCost, 1e-12) + require.InDelta(t, 100000*6.25e-6, cost.CacheCreationCost, 1e-12) + require.InDelta(t, 72000*0.5e-6, cost.CacheReadCost, 1e-12) require.InDelta(t, 10*30e-6, cost.OutputCost, 1e-12) } +func TestPricingService_BareGPT56AliasDeterministicallyUsesSol(t *testing.T) { + pricingSvc := &PricingService{pricingData: map[string]*LiteLLMModelPricing{ + "gpt-5.6-sol": {InputCostPerToken: 5e-6}, + "gpt-5.6-terra": {InputCostPerToken: 2.5e-6}, + "gpt-5.6-luna": {InputCostPerToken: 1e-6}, + "gpt-5.4": {InputCostPerToken: 2.5e-6}, + }} + + for i := 0; i < 100; i++ { + for _, alias := range []string{"gpt-5.6", "openai/gpt-5.6"} { + pricing := pricingSvc.GetModelPricing(alias) + require.NotNil(t, pricing) + require.InDelta(t, 5e-6, pricing.InputCostPerToken, 1e-12, "iteration=%d alias=%s", i, alias) + } + } + + billingSvc := NewBillingService(&config.Config{}, pricingSvc) + for _, alias := range []string{"gpt-5.6", "openai/gpt-5.6"} { + pricing, err := billingSvc.GetModelPricing(alias) + require.NoError(t, err) + require.InDelta(t, 5e-6, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, 6.25e-6, pricing.CacheCreationPricePerToken, 1e-12) + } +} + func TestDefaultPricingIncludesOfficialGPT56Rates(t *testing.T) { data, err := os.ReadFile(filepath.Join("..", "..", "resources", "model-pricing", "model_prices_and_context_window.json")) require.NoError(t, err) @@ -140,7 +210,9 @@ func TestDefaultPricingIncludesOfficialGPT56Rates(t *testing.T) { require.InDelta(t, tt.cachedPriority, pricing.CacheReadPricePerTokenPriority, 1e-12) require.InDelta(t, tt.cacheWritePriority, pricing.CacheCreationPricePerTokenPriority, 1e-12) require.InDelta(t, tt.outputPriority, pricing.OutputPricePerTokenPriority, 1e-12) - require.Zero(t, pricing.LongContextInputThreshold) + require.Equal(t, 272000, pricing.LongContextInputThreshold) + require.InDelta(t, 2.0, pricing.LongContextInputMultiplier, 1e-12) + require.InDelta(t, 1.5, pricing.LongContextOutputMultiplier, 1e-12) }) } } @@ -181,7 +253,9 @@ func assertGPT56FallbackPricing(t *testing.T, pricing *ModelPricing, input, cach require.InDelta(t, cached, pricing.CacheReadPricePerToken, 1e-12) require.InDelta(t, cacheWrite, pricing.CacheCreationPricePerToken, 1e-12) require.InDelta(t, output, pricing.OutputPricePerToken, 1e-12) - require.Zero(t, pricing.LongContextInputThreshold) + require.Equal(t, 272000, pricing.LongContextInputThreshold) + require.InDelta(t, 2.0, pricing.LongContextInputMultiplier, 1e-12) + require.InDelta(t, 1.5, pricing.LongContextOutputMultiplier, 1e-12) } func TestParsePricingData_KeepsImageOnlyPricing(t *testing.T) { diff --git a/backend/migrations/173_allow_cyber_blocked_usage_request_type.sql b/backend/migrations/173_allow_cyber_blocked_usage_request_type.sql new file mode 100644 index 0000000000..57b1672fa8 --- /dev/null +++ b/backend/migrations/173_allow_cyber_blocked_usage_request_type.sql @@ -0,0 +1,8 @@ +-- Cyber-policy blocks are recorded as request_type=4 so they remain visible in +-- usage audits without being confused with legacy request_type=0 rows. +ALTER TABLE usage_logs + DROP CONSTRAINT IF EXISTS usage_logs_request_type_check; + +ALTER TABLE usage_logs + ADD CONSTRAINT usage_logs_request_type_check + CHECK (request_type IN (0, 1, 2, 3, 4)) NOT VALID; diff --git a/backend/migrations/auth_identity_payment_migrations_regression_test.go b/backend/migrations/auth_identity_payment_migrations_regression_test.go index 14dab99160..68538cdd07 100644 --- a/backend/migrations/auth_identity_payment_migrations_regression_test.go +++ b/backend/migrations/auth_identity_payment_migrations_regression_test.go @@ -213,3 +213,30 @@ func TestMigration154aAddsSparkShadowIndexesConcurrently(t *testing.T) { require.Contains(t, sql, "quota_dimension = 'spark'") require.Contains(t, sql, "deleted_at IS NULL") } + +func TestMigration173AllowsCyberBlockedUsageRequestType(t *testing.T) { + entries, err := FS.ReadDir(".") + require.NoError(t, err) + + previousIndex := -1 + currentIndex := -1 + for i, entry := range entries { + switch entry.Name() { + case "172_video_per_second_billing_metadata.sql": + previousIndex = i + case "173_allow_cyber_blocked_usage_request_type.sql": + currentIndex = i + } + } + require.NotEqual(t, -1, previousIndex) + require.NotEqual(t, -1, currentIndex) + require.Less(t, previousIndex, currentIndex) + + content, err := FS.ReadFile("173_allow_cyber_blocked_usage_request_type.sql") + require.NoError(t, err) + + sql := string(content) + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS usage_logs_request_type_check") + require.Contains(t, sql, "ADD CONSTRAINT usage_logs_request_type_check") + require.Contains(t, sql, "CHECK (request_type IN (0, 1, 2, 3, 4)) NOT VALID") +} diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index 439988a53d..a762b2145f 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -4972,6 +4972,9 @@ "input_cost_per_token_batches": 2.5e-06, "input_cost_per_token_flex": 2.5e-06, "input_cost_per_token_priority": 1e-05, + "long_context_input_token_threshold": 272000, + "long_context_input_cost_multiplier": 2.0, + "long_context_output_cost_multiplier": 1.5, "litellm_provider": "openai", "max_input_tokens": 1050000, "max_output_tokens": 128000, @@ -5021,6 +5024,9 @@ "input_cost_per_token_batches": 1.25e-06, "input_cost_per_token_flex": 1.25e-06, "input_cost_per_token_priority": 5e-06, + "long_context_input_token_threshold": 272000, + "long_context_input_cost_multiplier": 2.0, + "long_context_output_cost_multiplier": 1.5, "litellm_provider": "openai", "max_input_tokens": 1050000, "max_output_tokens": 128000, @@ -5070,6 +5076,9 @@ "input_cost_per_token_batches": 5e-07, "input_cost_per_token_flex": 5e-07, "input_cost_per_token_priority": 2e-06, + "long_context_input_token_threshold": 272000, + "long_context_input_cost_multiplier": 2.0, + "long_context_output_cost_multiplier": 1.5, "litellm_provider": "openai", "max_input_tokens": 1050000, "max_output_tokens": 128000, diff --git a/frontend/src/components/keys/UseKeyModal.vue b/frontend/src/components/keys/UseKeyModal.vue index 5900644814..d3bca5188a 100644 --- a/frontend/src/components/keys/UseKeyModal.vue +++ b/frontend/src/components/keys/UseKeyModal.vue @@ -636,6 +636,23 @@ function generateOpenCodeConfig(platform: string, baseUrl: string, apiKey: strin xhigh: {} } }, + 'gpt-5.6': { + name: 'GPT-5.6 (Sol)', + limit: { + context: 1050000, + output: 128000 + }, + options: { + store: false + }, + variants: { + low: {}, + medium: {}, + high: {}, + xhigh: {}, + max: {} + } + }, 'gpt-5.6-sol': { name: 'GPT-5.6 Sol', limit: { @@ -649,7 +666,8 @@ function generateOpenCodeConfig(platform: string, baseUrl: string, apiKey: strin low: {}, medium: {}, high: {}, - xhigh: {} + xhigh: {}, + max: {} } }, 'gpt-5.6-terra': { @@ -665,7 +683,8 @@ function generateOpenCodeConfig(platform: string, baseUrl: string, apiKey: strin low: {}, medium: {}, high: {}, - xhigh: {} + xhigh: {}, + max: {} } }, 'gpt-5.6-luna': { @@ -681,7 +700,8 @@ function generateOpenCodeConfig(platform: string, baseUrl: string, apiKey: strin low: {}, medium: {}, high: {}, - xhigh: {} + xhigh: {}, + max: {} } }, 'gpt-5.5': { diff --git a/frontend/src/components/keys/__tests__/UseKeyModal.spec.ts b/frontend/src/components/keys/__tests__/UseKeyModal.spec.ts index 8ad9731baa..f8395dc75f 100644 --- a/frontend/src/components/keys/__tests__/UseKeyModal.spec.ts +++ b/frontend/src/components/keys/__tests__/UseKeyModal.spec.ts @@ -123,6 +123,43 @@ describe('UseKeyModal', () => { expect(codeBlock.text()).not.toContain('"name": "GPT-5.4 Nano"') }) + it('renders GPT-5.6 alias and max variants in OpenCode config', async () => { + const wrapper = mount(UseKeyModal, { + props: { + show: true, + apiKey: 'sk-test', + baseUrl: 'https://example.com/v1', + platform: 'openai' + }, + global: { + stubs: { + BaseDialog: { + template: '
' + }, + Icon: { + template: '' + } + } + } + }) + + const opencodeTab = wrapper.findAll('button').find((button) => + button.text().includes('keys.useKeyModal.cliTabs.opencode') + ) + expect(opencodeTab).toBeDefined() + await opencodeTab!.trigger('click') + await nextTick() + + const parsed = JSON.parse(wrapper.find('pre code').text()) + const models = parsed.provider.openai.models + for (const model of ['gpt-5.6', 'gpt-5.6-sol', 'gpt-5.6-terra', 'gpt-5.6-luna']) { + expect(models[model]).toBeDefined() + expect(models[model].variants).toHaveProperty('max') + expect(models[model].variants).toHaveProperty('xhigh') + } + expect(models['gpt-5.6'].name).toBe('GPT-5.6 (Sol)') + }) + it('renders Claude Fable 5 OpenCode config with adaptive thinking', async () => { const wrapper = mount(UseKeyModal, { props: { diff --git a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts index a9de605c28..cd186213d6 100644 --- a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts +++ b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts @@ -14,6 +14,7 @@ describe('useModelWhitelist', () => { expect(models).toContain('gpt-5.4-mini') expect(models).toContain('gpt-5.4-2026-03-05') expect(models).toContain('codex-auto-review') + expect(models).toContain('gpt-5.6') }) it('openai 模型列表不再暴露已下线的 ChatGPT 登录 Codex 模型', () => { diff --git a/frontend/src/composables/useModelWhitelist.ts b/frontend/src/composables/useModelWhitelist.ts index cce3df6a52..bbd98711fd 100644 --- a/frontend/src/composables/useModelWhitelist.ts +++ b/frontend/src/composables/useModelWhitelist.ts @@ -8,7 +8,7 @@ const openaiModels = [ 'gpt-5.2', 'gpt-5.2-2025-12-11', 'gpt-5.2-chat-latest', 'gpt-5.2-pro', 'gpt-5.2-pro-2025-12-11', // GPT-5.6 系列 - 'gpt-5.6-sol', 'gpt-5.6-terra', 'gpt-5.6-luna', + 'gpt-5.6', 'gpt-5.6-sol', 'gpt-5.6-terra', 'gpt-5.6-luna', // GPT-5.5 系列 'gpt-5.5', // GPT-5.4 系列 @@ -280,6 +280,7 @@ const openaiPresetMappings = [ { label: 'o3', from: 'o3', to: 'o3', color: 'bg-emerald-100 text-emerald-700 hover:bg-emerald-200 dark:bg-emerald-900/30 dark:text-emerald-400' }, { label: 'GPT-5.3 Codex Spark', from: 'gpt-5.3-codex-spark', to: 'gpt-5.3-codex-spark', color: 'bg-teal-100 text-teal-700 hover:bg-teal-200 dark:bg-teal-900/30 dark:text-teal-400' }, { label: 'GPT-5.2', from: 'gpt-5.2', to: 'gpt-5.2', color: 'bg-red-100 text-red-700 hover:bg-red-200 dark:bg-red-900/30 dark:text-red-400' }, + { label: 'GPT-5.6', from: 'gpt-5.6', to: 'gpt-5.6', color: 'bg-amber-100 text-amber-700 hover:bg-amber-200 dark:bg-amber-900/30 dark:text-amber-400' }, { label: 'GPT-5.6 Sol', from: 'gpt-5.6-sol', to: 'gpt-5.6-sol', color: 'bg-orange-100 text-orange-700 hover:bg-orange-200 dark:bg-orange-900/30 dark:text-orange-400' }, { label: 'GPT-5.6 Terra', from: 'gpt-5.6-terra', to: 'gpt-5.6-terra', color: 'bg-lime-100 text-lime-700 hover:bg-lime-200 dark:bg-lime-900/30 dark:text-lime-400' }, { label: 'GPT-5.6 Luna', from: 'gpt-5.6-luna', to: 'gpt-5.6-luna', color: 'bg-sky-100 text-sky-700 hover:bg-sky-200 dark:bg-sky-900/30 dark:text-sky-400' }, From f2966530c52a2aa4b4f4b507592f4bd258837eb3 Mon Sep 17 00:00:00 2001 From: Lyonle <214648221+lyon-le@users.noreply.github.com> Date: Fri, 10 Jul 2026 15:16:09 +0800 Subject: [PATCH 24/27] =?UTF-8?q?feat(openai):=20=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E7=94=A8=E6=88=B7=E7=BA=A7=20Fast/Flex=20=E7=AD=96=E7=95=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../handler/admin/admin_helpers_test.go | 2 + backend/internal/handler/dto/settings.go | 1 + backend/internal/pkg/ctxkey/ctxkey.go | 4 + .../server/middleware/api_key_auth.go | 2 + .../server/middleware/api_key_auth_test.go | 5 + .../openai_fast_policy_forwarding_test.go | 189 ++++++++++++++++++ .../service/openai_fast_policy_test.go | 86 ++++++++ .../service/openai_fast_policy_ws_test.go | 34 ++++ .../service/openai_gateway_request_body.go | 69 +++++-- backend/internal/service/setting_features.go | 10 + backend/internal/service/settings_view.go | 1 + frontend/src/api/admin/settings.ts | 1 + .../src/i18n/locales/en/admin/settings.ts | 5 + .../src/i18n/locales/zh/admin/settings.ts | 5 + frontend/src/views/admin/SettingsView.vue | 85 ++++++++ 15 files changed, 482 insertions(+), 17 deletions(-) create mode 100644 backend/internal/server/middleware/openai_fast_policy_forwarding_test.go diff --git a/backend/internal/handler/admin/admin_helpers_test.go b/backend/internal/handler/admin/admin_helpers_test.go index 6df4915486..c0775db6b6 100644 --- a/backend/internal/handler/admin/admin_helpers_test.go +++ b/backend/internal/handler/admin/admin_helpers_test.go @@ -265,10 +265,12 @@ func TestOpenAIFastPolicySettingsFromDTO_NormalizesServiceTier(t *testing.T) { ServiceTier: "PRIORITY", Action: "filter", Scope: "all", + UserIDs: []int64{42}, }}, } out := openaiFastPolicySettingsFromDTO(in) require.Equal(t, service.OpenAIFastTierPriority, out.Rules[0].ServiceTier) + require.Equal(t, []int64{42}, out.Rules[0].UserIDs) }) t.Run("non-empty values pass through (lowercased)", func(t *testing.T) { diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 99fba54980..5b7e657f26 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -422,6 +422,7 @@ type OpenAIFastPolicyRule struct { ServiceTier string `json:"service_tier"` Action string `json:"action"` Scope string `json:"scope"` + UserIDs []int64 `json:"user_ids,omitempty"` ErrorMessage string `json:"error_message,omitempty"` ModelWhitelist []string `json:"model_whitelist,omitempty"` FallbackAction string `json:"fallback_action,omitempty"` diff --git a/backend/internal/pkg/ctxkey/ctxkey.go b/backend/internal/pkg/ctxkey/ctxkey.go index dacd1bd1cf..9ba397a12c 100644 --- a/backend/internal/pkg/ctxkey/ctxkey.go +++ b/backend/internal/pkg/ctxkey/ctxkey.go @@ -41,6 +41,10 @@ const ( // Group 认证后的分组信息,由 API Key 认证中间件设置 Group Key = "ctx_group" + // UserID 认证后的 Sub2API 用户 ID,由 API Key 认证中间件设置。 + // 供 service 层执行用户级策略,不能使用客户端请求体中的 user 标识替代。 + UserID Key = "ctx_user_id" + // IsMaxTokensOneHaikuRequest 标识当前请求是否为 max_tokens=1 + haiku 模型的探测请求 // 用于 ClaudeCodeOnly 验证绕过(绕过 system prompt 检查,但仍需验证 User-Agent) IsMaxTokensOneHaikuRequest Key = "ctx_is_max_tokens_one_haiku" diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go index 04a09862b5..082766bfd5 100644 --- a/backend/internal/server/middleware/api_key_auth.go +++ b/backend/internal/server/middleware/api_key_auth.go @@ -126,6 +126,8 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti if abortIfAPIKeyGroupNotAllowed(c, apiKey) { return } + ctx := context.WithValue(c.Request.Context(), ctxkey.UserID, apiKey.User.ID) + c.Request = c.Request.WithContext(ctx) // ── 4. SimpleMode → early return ───────────────────────────── diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index abb84ab852..d9b5c7efc1 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -286,6 +286,11 @@ func TestAPIKeyAuthSetsGroupContext(t *testing.T) { c.JSON(http.StatusInternalServerError, gin.H{"ok": false}) return } + userIDFromCtx, ok := c.Request.Context().Value(ctxkey.UserID).(int64) + if !ok || userIDFromCtx != user.ID { + c.JSON(http.StatusInternalServerError, gin.H{"ok": false}) + return + } c.JSON(http.StatusOK, gin.H{"ok": true}) }) diff --git a/backend/internal/server/middleware/openai_fast_policy_forwarding_test.go b/backend/internal/server/middleware/openai_fast_policy_forwarding_test.go new file mode 100644 index 0000000000..5fea5547a3 --- /dev/null +++ b/backend/internal/server/middleware/openai_fast_policy_forwarding_test.go @@ -0,0 +1,189 @@ +package middleware + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestAPIKeyAuthForwardsUserScopedOpenAIFastPolicyToUpstream(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstreamBodies := make(chan []byte, 2) + upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, "read request body", http.StatusInternalServerError) + return + } + upstreamBodies <- body + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"resp_test","object":"response","model":"gpt-5","status":"completed","usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`)) + })) + defer upstreamServer.Close() + + settings := &service.OpenAIFastPolicySettings{ + Rules: []service.OpenAIFastPolicyRule{ + { + ServiceTier: service.OpenAIFastTierPriority, + Action: service.BetaPolicyActionFilter, + Scope: service.BetaPolicyScopeAll, + }, + { + ServiceTier: service.OpenAIFastTierPriority, + Action: service.BetaPolicyActionPass, + Scope: service.BetaPolicyScopeAll, + UserIDs: []int64{42}, + }, + }, + } + settingsJSON, err := json.Marshal(settings) + require.NoError(t, err) + + cfg := &config.Config{RunMode: config.RunModeSimple} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + + settingService := service.NewSettingService(&openAIFastPolicyForwardingSettingRepo{ + value: string(settingsJSON), + }, cfg) + gatewayService := service.NewOpenAIGatewayService( + nil, nil, nil, nil, nil, nil, nil, cfg, + nil, nil, nil, nil, nil, &openAIFastPolicyForwardingHTTPUpstream{client: upstreamServer.Client()}, + nil, nil, nil, nil, nil, nil, settingService, nil, + ) + + groupID := int64(101) + group := &service.Group{ + ID: groupID, + Name: "openai", + Status: service.StatusActive, + Platform: service.PlatformOpenAI, + Hydrated: true, + } + apiKeys := map[string]*service.APIKey{ + "key-user-42": newOpenAIFastPolicyForwardingAPIKey(1, "key-user-42", 42, groupID, group), + "key-user-43": newOpenAIFastPolicyForwardingAPIKey(2, "key-user-43", 43, groupID, group), + } + apiKeyService := service.NewAPIKeyService(&openAIFastPolicyForwardingAPIKeyRepo{apiKeys: apiKeys}, nil, nil, nil, nil, nil, cfg) + account := &service.Account{ + ID: 900, + Name: "openai-upstream", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": upstreamServer.URL, + }, + Extra: map[string]any{"use_responses_api": true}, + } + + router := gin.New() + router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg))) + router.POST("/v1/responses", func(c *gin.Context) { + body, readErr := io.ReadAll(c.Request.Body) + if readErr != nil { + c.Status(http.StatusBadRequest) + return + } + service.SetOpenAIClientTransport(c, service.OpenAIClientTransportHTTP) + if _, forwardErr := gatewayService.Forward(c.Request.Context(), c, account, body); forwardErr != nil { + c.Status(http.StatusBadGateway) + return + } + c.Status(http.StatusOK) + }) + + send := func(apiKey string) { + request := httptest.NewRequest( + http.MethodPost, + "/v1/responses", + bytes.NewBufferString(`{"model":"gpt-5","stream":false,"service_tier":"priority","input":"hi"}`), + ) + request.Header.Set("Content-Type", "application/json") + request.Header.Set("x-api-key", apiKey) + response := httptest.NewRecorder() + router.ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code) + } + + send("key-user-42") + send("key-user-43") + + allowedUserBody := <-upstreamBodies + otherUserBody := <-upstreamBodies + require.Equal(t, service.OpenAIFastTierPriority, gjson.GetBytes(allowedUserBody, "service_tier").String()) + require.False(t, gjson.GetBytes(otherUserBody, "service_tier").Exists()) +} + +func newOpenAIFastPolicyForwardingAPIKey(id int64, key string, userID, groupID int64, group *service.Group) *service.APIKey { + return &service.APIKey{ + ID: id, + UserID: userID, + Key: key, + Status: service.StatusActive, + GroupID: &groupID, + User: &service.User{ + ID: userID, + Role: service.RoleUser, + Status: service.StatusActive, + Balance: 10, + Concurrency: 1, + }, + Group: group, + } +} + +type openAIFastPolicyForwardingAPIKeyRepo struct { + service.APIKeyRepository + apiKeys map[string]*service.APIKey +} + +func (r *openAIFastPolicyForwardingAPIKeyRepo) GetByKeyForAuth(_ context.Context, key string) (*service.APIKey, error) { + apiKey, ok := r.apiKeys[key] + if !ok { + return nil, service.ErrAPIKeyNotFound + } + clone := *apiKey + return &clone, nil +} + +func (r *openAIFastPolicyForwardingAPIKeyRepo) UpdateLastUsed(context.Context, int64, time.Time) error { + return nil +} + +type openAIFastPolicyForwardingSettingRepo struct { + service.SettingRepository + value string +} + +func (r *openAIFastPolicyForwardingSettingRepo) GetValue(context.Context, string) (string, error) { + return r.value, nil +} + +type openAIFastPolicyForwardingHTTPUpstream struct { + client *http.Client +} + +func (u *openAIFastPolicyForwardingHTTPUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return u.client.Do(req) +} + +func (u *openAIFastPolicyForwardingHTTPUpstream) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return u.Do(req, proxyURL, accountID, accountConcurrency) +} diff --git a/backend/internal/service/openai_fast_policy_test.go b/backend/internal/service/openai_fast_policy_test.go index d0be963e5e..c5144b6a02 100644 --- a/backend/internal/service/openai_fast_policy_test.go +++ b/backend/internal/service/openai_fast_policy_test.go @@ -7,6 +7,7 @@ import ( "testing" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -138,6 +139,37 @@ func TestEvaluateOpenAIFastPolicy_ScopeFiltersOAuth(t *testing.T) { require.Equal(t, BetaPolicyActionPass, action) } +func TestEvaluateOpenAIFastPolicy_UserScopedRuleOverridesGlobalRule(t *testing.T) { + settings := &OpenAIFastPolicySettings{ + Rules: []OpenAIFastPolicyRule{ + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionFilter, + Scope: BetaPolicyScopeAll, + }, + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionPass, + Scope: BetaPolicyScopeAll, + UserIDs: []int64{42}, + }, + }, + } + svc := newOpenAIGatewayServiceWithSettings(t, settings) + account := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + + allowedUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(42)) + action, _ := svc.evaluateOpenAIFastPolicy(allowedUserCtx, account, "gpt-5.5", OpenAIFastTierPriority) + require.Equal(t, BetaPolicyActionPass, action) + + otherUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(43)) + action, _ = svc.evaluateOpenAIFastPolicy(otherUserCtx, account, "gpt-5.5", OpenAIFastTierPriority) + require.Equal(t, BetaPolicyActionFilter, action) + + action, _ = svc.evaluateOpenAIFastPolicy(context.Background(), account, "gpt-5.5", OpenAIFastTierPriority) + require.Equal(t, BetaPolicyActionFilter, action) +} + func TestApplyOpenAIFastPolicyToBody_DefaultPassesPriorityAndFast(t *testing.T) { svc := newOpenAIGatewayServiceWithSettings(t, DefaultOpenAIFastPolicySettings()) account := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} @@ -179,6 +211,37 @@ func TestApplyOpenAIFastPolicyToBody_ExplicitFilterRemovesField(t *testing.T) { require.NotContains(t, string(updated), `"service_tier"`) } +func TestApplyOpenAIFastPolicyToBody_UserScopedRuleOverridesGlobalRule(t *testing.T) { + settings := &OpenAIFastPolicySettings{ + Rules: []OpenAIFastPolicyRule{ + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionFilter, + Scope: BetaPolicyScopeAll, + }, + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionPass, + Scope: BetaPolicyScopeAll, + UserIDs: []int64{42}, + }, + }, + } + svc := newOpenAIGatewayServiceWithSettings(t, settings) + account := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + body := []byte(`{"model":"gpt-5.5","service_tier":"priority"}`) + + allowedUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(42)) + updated, err := svc.applyOpenAIFastPolicyToBody(allowedUserCtx, account, "gpt-5.5", body) + require.NoError(t, err) + require.Equal(t, "priority", gjson.GetBytes(updated, "service_tier").String()) + + otherUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(43)) + updated, err = svc.applyOpenAIFastPolicyToBody(otherUserCtx, account, "gpt-5.5", body) + require.NoError(t, err) + require.NotContains(t, string(updated), `"service_tier"`) +} + func TestApplyOpenAIFastPolicyToBody_ForcePriorityRewritesKnownTier(t *testing.T) { settings := &OpenAIFastPolicySettings{ Rules: []OpenAIFastPolicyRule{{ @@ -309,12 +372,34 @@ func TestSetOpenAIFastPolicySettings_Validation(t *testing.T) { }) require.Error(t, err) + // Non-positive and duplicate user IDs are rejected. + err = svc.SetOpenAIFastPolicySettings(context.Background(), &OpenAIFastPolicySettings{ + Rules: []OpenAIFastPolicyRule{{ + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionPass, + Scope: BetaPolicyScopeAll, + UserIDs: []int64{0}, + }}, + }) + require.Error(t, err) + + err = svc.SetOpenAIFastPolicySettings(context.Background(), &OpenAIFastPolicySettings{ + Rules: []OpenAIFastPolicyRule{{ + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionPass, + Scope: BetaPolicyScopeAll, + UserIDs: []int64{42, 42}, + }}, + }) + require.Error(t, err) + // Valid settings persisted err = svc.SetOpenAIFastPolicySettings(context.Background(), &OpenAIFastPolicySettings{ Rules: []OpenAIFastPolicyRule{{ ServiceTier: OpenAIFastTierPriority, Action: OpenAIFastPolicyActionForcePriority, Scope: BetaPolicyScopeAll, + UserIDs: []int64{42, 43}, }}, }) require.NoError(t, err) @@ -324,4 +409,5 @@ func TestSetOpenAIFastPolicySettings_Validation(t *testing.T) { require.Len(t, got.Rules, 1) require.Equal(t, OpenAIFastTierPriority, got.Rules[0].ServiceTier) require.Equal(t, OpenAIFastPolicyActionForcePriority, got.Rules[0].Action) + require.Equal(t, []int64{42, 43}, got.Rules[0].UserIDs) } diff --git a/backend/internal/service/openai_fast_policy_ws_test.go b/backend/internal/service/openai_fast_policy_ws_test.go index a802540879..f624a0080b 100644 --- a/backend/internal/service/openai_fast_policy_ws_test.go +++ b/backend/internal/service/openai_fast_policy_ws_test.go @@ -14,6 +14,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/claude" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" coderws "github.com/coder/websocket" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -67,6 +68,39 @@ func TestWSResponseCreate_ExplicitFilterStripsServiceTier(t *testing.T) { require.NotContains(t, string(updated), `"service_tier"`) } +func TestWSResponseCreate_UserScopedRuleOverridesGlobalRule(t *testing.T) { + settings := &OpenAIFastPolicySettings{ + Rules: []OpenAIFastPolicyRule{ + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionFilter, + Scope: BetaPolicyScopeAll, + }, + { + ServiceTier: OpenAIFastTierPriority, + Action: BetaPolicyActionPass, + Scope: BetaPolicyScopeAll, + UserIDs: []int64{42}, + }, + }, + } + svc := newOpenAIGatewayServiceWithSettings(t, settings) + account := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + frame := []byte(`{"type":"response.create","model":"gpt-5.5","service_tier":"priority"}`) + + allowedUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(42)) + updated, blocked, err := svc.applyOpenAIFastPolicyToWSResponseCreate(allowedUserCtx, account, "gpt-5.5", frame) + require.NoError(t, err) + require.Nil(t, blocked) + require.Equal(t, "priority", gjson.GetBytes(updated, "service_tier").String()) + + otherUserCtx := context.WithValue(context.Background(), ctxkey.UserID, int64(43)) + updated, blocked, err = svc.applyOpenAIFastPolicyToWSResponseCreate(otherUserCtx, account, "gpt-5.5", frame) + require.NoError(t, err) + require.Nil(t, blocked) + require.NotContains(t, string(updated), `"service_tier"`) +} + func TestWSResponseCreate_ForcePriorityRewritesKnownTier(t *testing.T) { settings := &OpenAIFastPolicySettings{ Rules: []OpenAIFastPolicyRule{{ diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 0e888e145b..935a32f58b 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -9,6 +9,7 @@ import ( "net/http" "strings" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/gin-gonic/gin" "github.com/google/uuid" @@ -659,9 +660,12 @@ func (e *OpenAIFastBlockedError) Error() string { return e.Message } // // Matching rules: // - Scope filters by account type (all / oauth / apikey / bedrock) +// - UserIDs, when present, filters by the trusted Sub2API user that owns the API key // - ServiceTier must be empty (= any), "all", or equal the normalized tier // - ModelWhitelist narrows the rule to specific models; FallbackAction // handles the non-matching case (default: pass) +// - User-specific rules take precedence over global rules; each group keeps +// the configured first-match order // // 与 Claude BetaPolicy 的差异(保留首条匹配 short-circuit): // - BetaPolicy 处理的是 anthropic-beta header 中的 token 集合,不同 @@ -687,39 +691,70 @@ func (s *OpenAIGatewayService) evaluateOpenAIFastPolicy(ctx context.Context, acc } settings = fetched } - return evaluateOpenAIFastPolicyWithSettings(settings, account, model, tier) + return evaluateOpenAIFastPolicyWithSettings(settings, openAIFastPolicyUserID(ctx), account, model, tier) } // evaluateOpenAIFastPolicyWithSettings is the pure-function core extracted so // long-lived sessions (e.g. WS) can prefetch settings once and avoid hitting // the settingService on every frame. See WSSession entry and // openAIFastPolicySettingsFromContext for the caching glue. -func evaluateOpenAIFastPolicyWithSettings(settings *OpenAIFastPolicySettings, account *Account, model, tier string) (action, errMsg string) { +func evaluateOpenAIFastPolicyWithSettings(settings *OpenAIFastPolicySettings, userID int64, account *Account, model, tier string) (action, errMsg string) { if settings == nil { return BetaPolicyActionPass, "" } isOAuth := account != nil && account.IsOAuth() isBedrock := account != nil && account.IsBedrock() - for _, rule := range settings.Rules { - if !betaPolicyScopeMatches(rule.Scope, isOAuth, isBedrock) { - continue + + // 用户专属规则先于全局规则。规则组内仍按配置顺序首条命中,允许 + // 管理员为某位用户配置例外,而不被先出现的全局规则覆盖。 + for _, userScoped := range []bool{true, false} { + for _, rule := range settings.Rules { + if (len(rule.UserIDs) > 0) != userScoped || !openAIFastPolicyUserMatches(rule.UserIDs, userID) { + continue + } + if !betaPolicyScopeMatches(rule.Scope, isOAuth, isBedrock) { + continue + } + ruleTier := strings.ToLower(strings.TrimSpace(rule.ServiceTier)) + if ruleTier != "" && ruleTier != OpenAIFastTierAny && ruleTier != tier { + continue + } + eff := BetaPolicyRule{ + Action: rule.Action, + ErrorMessage: rule.ErrorMessage, + ModelWhitelist: rule.ModelWhitelist, + FallbackAction: rule.FallbackAction, + FallbackErrorMessage: rule.FallbackErrorMessage, + } + return resolveRuleAction(eff, model) } - ruleTier := strings.ToLower(strings.TrimSpace(rule.ServiceTier)) - if ruleTier != "" && ruleTier != OpenAIFastTierAny && ruleTier != tier { - continue - } - eff := BetaPolicyRule{ - Action: rule.Action, - ErrorMessage: rule.ErrorMessage, - ModelWhitelist: rule.ModelWhitelist, - FallbackAction: rule.FallbackAction, - FallbackErrorMessage: rule.FallbackErrorMessage, - } - return resolveRuleAction(eff, model) } return BetaPolicyActionPass, "" } +func openAIFastPolicyUserID(ctx context.Context) int64 { + if ctx == nil { + return 0 + } + userID, _ := ctx.Value(ctxkey.UserID).(int64) + if userID <= 0 { + return 0 + } + return userID +} + +func openAIFastPolicyUserMatches(ruleUserIDs []int64, userID int64) bool { + if len(ruleUserIDs) == 0 { + return true + } + for _, ruleUserID := range ruleUserIDs { + if ruleUserID == userID { + return true + } + } + return false +} + // openAIFastPolicyCtxKey 是 context 中预取的 OpenAIFastPolicySettings 缓存 // 键,仅用于 WebSocket 长会话内多帧复用同一份策略快照,避免每帧 DB 命中。 // diff --git a/backend/internal/service/setting_features.go b/backend/internal/service/setting_features.go index 23dc32efa9..57c036861e 100644 --- a/backend/internal/service/setting_features.go +++ b/backend/internal/service/setting_features.go @@ -801,6 +801,16 @@ func (s *SettingService) SetOpenAIFastPolicySettings(ctx context.Context, settin if !validScopes[rule.Scope] { return fmt.Errorf("rule[%d]: invalid scope %q", i, rule.Scope) } + seenUserIDs := make(map[int64]struct{}, len(rule.UserIDs)) + for j, userID := range rule.UserIDs { + if userID <= 0 { + return fmt.Errorf("rule[%d]: user_ids[%d] must be positive", i, j) + } + if _, exists := seenUserIDs[userID]; exists { + return fmt.Errorf("rule[%d]: user_ids[%d] duplicates user_id %d", i, j, userID) + } + seenUserIDs[userID] = struct{}{} + } for j, pattern := range rule.ModelWhitelist { trimmed := strings.TrimSpace(pattern) if trimmed == "" { diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 9357c0c1e1..285df80ba3 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -586,6 +586,7 @@ type OpenAIFastPolicyRule struct { ServiceTier string `json:"service_tier"` // "priority" | "flex" | "auto" | "default" | "scale" | "all" Action string `json:"action"` // "pass" | "filter" | "block" | "force_priority" Scope string `json:"scope"` // "all" | "oauth" | "apikey" | "bedrock" + UserIDs []int64 `json:"user_ids,omitempty"` // 空=所有 Sub2API 用户;非空=仅指定 API Key 所属用户 ErrorMessage string `json:"error_message,omitempty"` // 自定义错误消息 (action=block 时生效) ModelWhitelist []string `json:"model_whitelist,omitempty"` // 模型匹配模式列表(为空=对所有模型生效) FallbackAction string `json:"fallback_action,omitempty"` // 未匹配白名单的模型的处理方式 diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 6f69163994..55c8088be1 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -1275,6 +1275,7 @@ export interface OpenAIFastPolicyRule { service_tier: "all" | "priority" | "flex"; action: "pass" | "filter" | "block" | "force_priority"; scope: "all" | "oauth" | "apikey" | "bedrock"; + user_ids?: number[]; error_message?: string; model_whitelist?: string[]; fallback_action?: "pass" | "filter" | "block" | "force_priority"; diff --git a/frontend/src/i18n/locales/en/admin/settings.ts b/frontend/src/i18n/locales/en/admin/settings.ts index d3d3f41ac9..37dc8aa1c8 100644 --- a/frontend/src/i18n/locales/en/admin/settings.ts +++ b/frontend/src/i18n/locales/en/admin/settings.ts @@ -979,6 +979,11 @@ export default { scopeOAuth: 'OAuth only', scopeAPIKey: 'API Key only', scopeBedrock: 'Bedrock only', + userIds: 'Specific user IDs', + userIdsHint: 'Leave empty to apply to all Sub2API users. Specified users match requests from their API keys and take precedence over global rules.', + userIdPlaceholder: 'e.g., 1001', + addUserId: 'Add user ID', + removeUserId: 'Remove user ID', errorMessage: 'Error message', errorMessagePlaceholder: 'Custom error message when blocked', errorMessageHint: 'Leave empty for default message', diff --git a/frontend/src/i18n/locales/zh/admin/settings.ts b/frontend/src/i18n/locales/zh/admin/settings.ts index bf848d1cc3..5c0d874b57 100644 --- a/frontend/src/i18n/locales/zh/admin/settings.ts +++ b/frontend/src/i18n/locales/zh/admin/settings.ts @@ -974,6 +974,11 @@ export default { scopeOAuth: '仅 OAuth 账号', scopeAPIKey: '仅 API Key 账号', scopeBedrock: '仅 Bedrock 账号', + userIds: '指定用户 ID', + userIdsHint: '留空表示对全部 Sub2API 用户生效。指定后仅匹配这些用户的 API Key 请求,且优先于全局规则。', + userIdPlaceholder: '例如: 1001', + addUserId: '添加用户 ID', + removeUserId: '移除用户 ID', errorMessage: '错误消息', errorMessagePlaceholder: '拦截时返回的自定义错误消息', errorMessageHint: '留空则使用默认错误消息', diff --git a/frontend/src/views/admin/SettingsView.vue b/frontend/src/views/admin/SettingsView.vue index 71db9cd4b0..a87cea9060 100644 --- a/frontend/src/views/admin/SettingsView.vue +++ b/frontend/src/views/admin/SettingsView.vue @@ -1189,6 +1189,72 @@ + +
+ +

+ {{ t("admin.settings.openaiFastPolicy.userIdsHint") }} +

+
+ + +
+ +
+