diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 42dc13112a..79709cabb7 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -255,6 +255,48 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { if err == nil && result != nil && result.FirstTokenMs != nil { service.SetOpsLatencyMs(c, service.OpsTimeToFirstTokenMsKey, int64(*result.FirstTokenMs)) } + // #5148 对齐:错误返回携带的部分 result(流中断前上游已计量的 usage)照常 + // 入账;failover 错误恒定 result=nil,不会重复计费。 + submitChatUsage := func(res *service.OpenAIForwardResult) { + if res == nil { + return + } + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, res) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + sessionID := service.ExtractClientSessionID(c) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil + h.submitOpenAIUsageRecordTask(c.Request.Context(), res, func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ + Result: res, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + SessionID: sessionID, + ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, res.UpstreamModel), + PricingAt: pricingAt, + CyberBlocked: cyberBlocked, + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.chat_completions"), + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("model", reqModel), + zap.Int64("account_id", account.ID), + ).Error("openai_chat_completions.record_usage_failed", zap.Error(err)) + } + }) + } if err != nil { if result != nil && result.ImageCount > 0 { reqLog.Warn("openai_chat_completions.forward_partial_error_with_image_result", @@ -339,6 +381,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { zap.Bool("upstream_error_response_already_written", upstreamErrorAlreadyCommunicated), zap.Error(err), ) + submitChatUsage(result) return } } @@ -348,42 +391,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil) } - userAgent := c.GetHeader("User-Agent") - clientIP := ip.GetClientIP(c) - inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) - quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - sessionID := service.ExtractClientSessionID(c) - - cyberBlocked := service.GetOpsCyberPolicy(c) != nil - h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { - if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ - Result: result, - APIKey: apiKey, - User: apiKey.User, - Account: account, - Subscription: subscription, - InboundEndpoint: inboundEndpoint, - UpstreamEndpoint: upstreamEndpoint, - UserAgent: userAgent, - IPAddress: clientIP, - APIKeyService: h.apiKeyService, - QuotaPlatform: quotaPlatform, - SessionID: sessionID, - ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel), - PricingAt: pricingAt, - CyberBlocked: cyberBlocked, - }); err != nil { - logger.L().With( - zap.String("component", "handler.openai_gateway.chat_completions"), - zap.Int64("user_id", subject.UserID), - zap.Int64("api_key_id", apiKey.ID), - zap.Any("group_id", apiKey.GroupID), - zap.String("model", reqModel), - zap.Int64("account_id", account.ID), - ).Error("openai_chat_completions.record_usage_failed", zap.Error(err)) - } - }) + submitChatUsage(result) reqLog.Debug("openai_chat_completions.request_completed", zap.Int64("account_id", account.ID), zap.Int("switch_count", switchCount), diff --git a/backend/internal/handler/openai_gateway_cn_dispatch_test.go b/backend/internal/handler/openai_gateway_cn_dispatch_test.go new file mode 100644 index 0000000000..b7f11c5462 --- /dev/null +++ b/backend/internal/handler/openai_gateway_cn_dispatch_test.go @@ -0,0 +1,29 @@ +package handler + +// CN 分组 /v1/messages 调度闸门回归(修复:正常途径创建的 CN 分组曾恒 403): +// sanitizeGroupMessagesDispatchFields 对非 openai 平台强制 AllowMessagesDispatch +// =false,故 CN 分组必须与 grok 一样在闸门处豁免,否则原生 Anthropic 直通 +//(Claude Code 主用例)永远不可达。 + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestAllowOpenAICompatibleMessagesDispatch_CNProvidersExempt(t *testing.T) { + require.True(t, allowOpenAICompatibleMessagesDispatch(nil), "无 key 保持放行") + + for _, platform := range []string{service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformGrok} { + apiKey := &service.APIKey{Group: &service.Group{Platform: platform, AllowMessagesDispatch: false}} + require.True(t, allowOpenAICompatibleMessagesDispatch(apiKey), + "%s 分组必须豁免 allow_messages_dispatch 闸门", platform) + } + + // 非回归:openai 分组仍受开关控制。 + openaiOff := &service.APIKey{Group: &service.Group{Platform: service.PlatformOpenAI, AllowMessagesDispatch: false}} + require.False(t, allowOpenAICompatibleMessagesDispatch(openaiOff)) + openaiOn := &service.APIKey{Group: &service.Group{Platform: service.PlatformOpenAI, AllowMessagesDispatch: true}} + require.True(t, allowOpenAICompatibleMessagesDispatch(openaiOn)) +} diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index ac6deda4aa..70e61e24f1 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -78,7 +78,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { zap.Any("group_id", apiKey.GroupID), ) - if apiKey.Group != nil && !apiKey.Group.AllowMessagesDispatch { + if !allowOpenAICompatibleMessagesDispatch(apiKey) { h.anthropicErrorResponse(c, http.StatusForbidden, "permission_error", "This group does not allow /v1/messages dispatch") return diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 28515d98f0..4389395800 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -197,6 +197,13 @@ func allowOpenAICompatibleMessagesDispatch(apiKey *service.APIKey) bool { if apiKey.Group.Platform == service.PlatformGrok { return true } + // 国产供应商分组与 grok 同语义:/v1/messages 就是其主要服务形态(anthropic + // 协议账号原生直通 Claude Code),无需 allow_messages_dispatch 开关授权—— + // 该开关对非 openai 平台恒被 sanitizeGroupMessagesDispatchFields 置 false, + // 若不豁免,CN 分组将永远 403。 + if service.IsCNProvider(apiKey.Group.Platform) { + return true + } return apiKey.Group.AllowMessagesDispatch } @@ -597,6 +604,50 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { if err == nil && result != nil && result.FirstTokenMs != nil { service.SetOpsLatencyMs(c, service.OpsTimeToFirstTokenMsKey, int64(*result.FirstTokenMs)) } + // #5148 对齐:错误返回携带的部分 result(流中断前上游已计量的 usage)照常 + // 入账;failover 错误恒定 result=nil,不会重复计费。 + submitResponsesUsage := func(res *service.OpenAIForwardResult) { + if res == nil { + return + } + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, res) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + sessionID := service.ExtractClientSessionID(c) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil + h.submitOpenAIUsageRecordTask(c.Request.Context(), res, func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ + Result: res, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + SessionID: sessionID, + ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, res.UpstreamModel), + PricingAt: pricingAt, + CyberBlocked: cyberBlocked, + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.responses"), + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("model", reqModel), + zap.Int64("account_id", account.ID), + ).Error("openai.record_usage_failed", zap.Error(err)) + } + }) + } if err != nil { if result != nil && result.ImageCount > 0 { reqLog.Warn("openai.forward_partial_error_with_image_result", @@ -696,6 +747,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { zap.Bool("upstream_error_response_already_written", upstreamErrorAlreadyCommunicated), zap.Error(err), } + submitResponsesUsage(result) if shouldLogOpenAIForwardFailureAsWarn(c, wroteFallback) { reqLog.Warn("openai.forward_failed", fields...) return @@ -714,46 +766,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), nil) } - // 捕获请求信息(用于异步记录,避免在 goroutine 中访问 gin.Context) - userAgent := c.GetHeader("User-Agent") - clientIP := ip.GetClientIP(c) - requestPayloadHash := service.HashUsageRequestPayload(body) - inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) - quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - sessionID := service.ExtractClientSessionID(c) - // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 - cyberBlocked := service.GetOpsCyberPolicy(c) != nil - h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { - if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ - Result: result, - APIKey: apiKey, - User: apiKey.User, - Account: account, - Subscription: subscription, - InboundEndpoint: inboundEndpoint, - UpstreamEndpoint: upstreamEndpoint, - UserAgent: userAgent, - IPAddress: clientIP, - RequestPayloadHash: requestPayloadHash, - APIKeyService: h.apiKeyService, - QuotaPlatform: quotaPlatform, - SessionID: sessionID, - ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel), - PricingAt: pricingAt, - CyberBlocked: cyberBlocked, - }); err != nil { - logger.L().With( - zap.String("component", "handler.openai_gateway.responses"), - zap.Int64("user_id", subject.UserID), - zap.Int64("api_key_id", apiKey.ID), - zap.Any("group_id", apiKey.GroupID), - zap.String("model", reqModel), - zap.Int64("account_id", account.ID), - ).Error("openai.record_usage_failed", zap.Error(err)) - } - }) + submitResponsesUsage(result) reqLog.Debug("openai.request_completed", zap.Int64("account_id", account.ID), zap.Int("switch_count", switchCount), @@ -1145,6 +1159,51 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { if err == nil && result != nil && result.FirstTokenMs != nil { service.SetOpsLatencyMs(c, service.OpsTimeToFirstTokenMsKey, int64(*result.FirstTokenMs)) } + // Forward 与错误一起返回的部分结果:流中断/客户端断开排水前上游已计量的 + // usage 照常入账,避免上游已产生消耗的请求完全漏记(#5148,对齐 anthropic + // 网关同名修复)。failover 错误恒定 result=nil,不会重复计费。 + submitMessagesUsage := func(res *service.OpenAIForwardResult) { + if res == nil { + return + } + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, res) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + sessionID := service.ExtractClientSessionID(c) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil + h.submitOpenAIUsageRecordTask(c.Request.Context(), res, func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ + Result: res, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + SessionID: sessionID, + ChannelUsageFields: clientRequestedUsageFields(c, channelMappingMsg, reqModel, res.UpstreamModel), + PricingAt: pricingAt, + CyberBlocked: cyberBlocked, + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.messages"), + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("model", reqModel), + zap.Int64("account_id", account.ID), + ).Error("openai_messages.record_usage_failed", zap.Error(err)) + } + }) + } if err != nil { if result != nil && result.ImageCount > 0 { reqLog.Warn("openai_messages.forward_partial_error_with_image_result", @@ -1219,6 +1278,9 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { zap.Int64("account_id", account.ID), zap.Error(err), ) + // 断开排水期间上游已计量的 usage 必须入账(此前直接 return 丢弃, + // payg 上游照常计费而平台漏记)。 + submitMessagesUsage(result) return } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), false, nil) @@ -1228,6 +1290,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { zap.Bool("fallback_error_response_written", wroteFallback), zap.Error(err), ) + submitMessagesUsage(result) return } } @@ -1237,44 +1300,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), true, nil) } - userAgent := c.GetHeader("User-Agent") - clientIP := ip.GetClientIP(c) - requestPayloadHash := service.HashUsageRequestPayload(body) - inboundEndpoint := GetInboundEndpoint(c) - upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) - quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - sessionID := service.ExtractClientSessionID(c) - - cyberBlocked := service.GetOpsCyberPolicy(c) != nil - h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { - if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ - Result: result, - APIKey: apiKey, - User: apiKey.User, - Account: account, - Subscription: subscription, - InboundEndpoint: inboundEndpoint, - UpstreamEndpoint: upstreamEndpoint, - UserAgent: userAgent, - IPAddress: clientIP, - RequestPayloadHash: requestPayloadHash, - APIKeyService: h.apiKeyService, - QuotaPlatform: quotaPlatform, - SessionID: sessionID, - ChannelUsageFields: clientRequestedUsageFields(c, channelMappingMsg, reqModel, result.UpstreamModel), - PricingAt: pricingAt, - CyberBlocked: cyberBlocked, - }); err != nil { - logger.L().With( - zap.String("component", "handler.openai_gateway.messages"), - zap.Int64("user_id", subject.UserID), - zap.Int64("api_key_id", apiKey.ID), - zap.Any("group_id", apiKey.GroupID), - zap.String("model", reqModel), - zap.Int64("account_id", account.ID), - ).Error("openai_messages.record_usage_failed", zap.Error(err)) - } - }) + submitMessagesUsage(result) reqLog.Debug("openai_messages.request_completed", zap.Int64("account_id", account.ID), zap.Int("switch_count", switchCount), diff --git a/backend/internal/service/openai_gateway_cn_fixes_test.go b/backend/internal/service/openai_gateway_cn_fixes_test.go new file mode 100644 index 0000000000..ce80a1aca7 --- /dev/null +++ b/backend/internal/service/openai_gateway_cn_fixes_test.go @@ -0,0 +1,144 @@ +//go:build unit + +package service + +// 国产供应商功能修复回归测试: +// 1. CN 分组不适用 /v1/messages 调度级模型映射(openai 的 gpt-5.x 默认值发给 +// CN 上游必错); +// 2. 计费候选链对 CN 账号过滤 claude-* 兜底候选(防按 Claude 原价误计 CN 流量); +// 3. 空候选按 ErrModelPricingUnavailable 处理(零成本落账而非丢弃 usage 记录); +// 4. Responses×anthropic 流式转换器客户端断开后继续排水、usage 汇总完整。 + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestResolveMessagesDispatchModel_CNProvidersNoDispatchMapping(t *testing.T) { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek} { + g := &Group{Platform: platform} + require.Empty(t, g.ResolveMessagesDispatchModel("claude-sonnet-4-5"), + "CN 分组(%s)不得返回调度级映射模型(openai 默认值会发给 CN 上游)", platform) + require.Empty(t, g.ResolveMessagesDispatchModel("claude-opus-4-1"), platform) + } + // 非回归:openai 分组保持原有默认映射行为。 + openaiGroup := &Group{Platform: PlatformOpenAI} + require.NotEmpty(t, openaiGroup.ResolveMessagesDispatchModel("claude-sonnet-4-5"), + "openai 分组的调度默认映射不应受 CN 修复影响") +} + +func TestFilterCNProviderBillingModelCandidates(t *testing.T) { + svc := &OpenAIGatewayService{} // resolver 为 nil → 无显式分组/渠道定价 + apiKey := &APIKey{Group: &Group{ID: 1, Platform: PlatformKimi}} + + cnAccount := &Account{ID: 1, Platform: PlatformKimi} + filtered := svc.filterCNProviderBillingModelCandidates(context.Background(), cnAccount, apiKey, + []string{"kimi-k2-0905-preview", "claude-sonnet-4-5", "moonshot-v1-8k"}) + require.Equal(t, []string{"kimi-k2-0905-preview", "moonshot-v1-8k"}, filtered, + "无显式定价时 claude-* 候选必须被过滤") + + allClaude := svc.filterCNProviderBillingModelCandidates(context.Background(), cnAccount, apiKey, + []string{"claude-sonnet-4-5", "claude-sonnet-4-5"}) + require.Empty(t, allClaude, "全 claude 候选应被清空(上层走零成本+告警落账)") + + // 非 CN 账号完全不受影响。 + openaiAccount := &Account{ID: 2, Platform: PlatformOpenAI} + passthrough := svc.filterCNProviderBillingModelCandidates(context.Background(), openaiAccount, apiKey, + []string{"claude-sonnet-4-5", "gpt-5.4"}) + require.Equal(t, []string{"claude-sonnet-4-5", "gpt-5.4"}, passthrough) + + require.Nil(t, svc.filterCNProviderBillingModelCandidates(context.Background(), nil, apiKey, nil)) +} + +func TestCalculateOpenAIRecordUsageCost_EmptyCandidatesIsPricingUnavailable(t *testing.T) { + svc := &OpenAIGatewayService{} + apiKey := &APIKey{Group: &Group{ID: 1, Platform: PlatformKimi}} + + _, err := svc.calculateOpenAIRecordUsageCost( + context.Background(), nil, apiKey, nil, + 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 100}, "", nil, + ) + require.Error(t, err) + require.True(t, isUsagePricingUnavailableError(err), + "空候选必须按无价可循处理(上层零成本落账),而不是丢弃整条 usage 记录: %v", err) +} + +func TestResponsesStreamingFromNativeAnthropic_ClientDisconnectDrainsUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newNativeAnthropicHangTestService(5) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + // failAfter=0:首次写出即失败,模拟客户端断开(复用测试包既有 failingGinWriter)。 + failWriter := &failingGinWriter{ResponseWriter: c.Writer, failAfter: 0} + c.Writer = failWriter + + resp, pr, pw := newHangingUpstreamResponse() + go func() { + // 首事件触发客户端写失败后,末尾 message_delta 才携带最终 output_tokens: + // 断开即弃会把整段生成记成 1 token。 + _, _ = pw.Write([]byte(miniAnthropicSSEStream())) + _ = pw.Close() + }() + defer func() { _ = pr.Close() }() + + res, err := svc.handleResponsesStreamingFromNativeAnthropic( + resp, c, "glm-4.7", "glm-4.7", "glm-4.7", nil, time.Now(), apicompat.ResponsesClientToolMapping{}) + + require.NoError(t, err, "断开排水至上游自然结束应返回 nil error(usage 走成功路径落账)") + require.NotNil(t, res) + require.True(t, res.ClientDisconnect) + require.Equal(t, 10, res.Usage.InputTokens, "input_tokens 应来自 message_start") + require.Equal(t, 5, res.Usage.OutputTokens, + "output_tokens 必须来自排水读到的末尾 message_delta(断开即弃时会是 1)") +} + +func TestHandle403_CNProviderHTMLBodySkipsAccountPenalty(t *testing.T) { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek} { + repo := &rateLimitAccountRepoStub{} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + account := &Account{ID: 401, Platform: platform, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), + account, + http.StatusForbidden, + http.Header{}, + []byte("Access denied by CDN"), + ) + + require.False(t, shouldDisable, "%s: HTML 403(CDN/代理拦截页)不得作为账号失效证据", platform) + require.Equal(t, 0, repo.setErrorCalls, "%s: 不得永久禁用账号", platform) + require.Equal(t, 0, repo.tempCalls, "%s: 不得临时停调账号", platform) + } +} + +func TestHandle403_CNProviderStructured403TempUnschedulableFirstHit(t *testing.T) { + repo := &rateLimitAccountRepoStub{} + counter := &openAI403CounterCacheStub{counts: []int64{1}} + service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + service.SetOpenAI403CounterCache(counter) + account := &Account{ID: 402, Platform: PlatformKimi, Type: AccountTypeAPIKey} + + shouldDisable := service.HandleUpstreamError( + context.Background(), + account, + http.StatusForbidden, + http.Header{}, + []byte(`{"error":{"message":"forbidden"}}`), + ) + + require.True(t, shouldDisable) + require.Equal(t, 0, repo.setErrorCalls, "首次结构化 403 应临时停调而非永久禁用") + require.Equal(t, 1, repo.tempCalls) + require.Contains(t, repo.lastTempReason, "(1/3)") +} diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index f20b45b526..bcdaa457b8 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -92,15 +92,13 @@ func (s *OpenAIGatewayService) ForwardCountTokensAsAnthropic( return fmt.Errorf("count_tokens: missing account") } - // 国产供应商 Anthropic 协议:上游有原生 /v1/messages/count_tokens 端点, - // 直接透传(仅模型名映射),不走 /v1/responses/input_tokens 估算。 - if account.IsAnthropicProtocol() { - return s.forwardCountTokensViaNativeAnthropic(ctx, c, account, body, defaultMappedModel) - } - - // 国产供应商其余协议(chat_completions / responses):三家上游均无 - // OpenAI 兼容的 /v1/responses/input_tokens 端点,与 Grok 一样本地估算, - // 不发上游请求(Claude Code 客户端会高频调用 count_tokens)。 + // 国产供应商(全部协议,含 anthropic):一律本地估算,不发上游请求。 + // 依据(2026-08 核实):三家的 Anthropic 兼容层均未提供 + // /v1/messages/count_tokens——DeepSeek 官方 anthropic_api 文档无此端点 + // (且注明 anthropic-version 头被忽略),聚合网关 OpenModel 明确标注 + // count_tokens 为 "Anthropic only",Kimi/智谱亦无任何文档承诺。转发上游 + // 只会常态 404,且错误还会流入账号处置逻辑误伤整账号调度;Claude Code + // 高频调用此端点,本地 tiktoken 估算是与 Grok 一致的既有方案。 if account.IsCNProvider() { estimated, err := estimateAnthropicCountTokensLocally(body) if err != nil { diff --git a/backend/internal/service/openai_gateway_messages_anthropic_native.go b/backend/internal/service/openai_gateway_messages_anthropic_native.go index e785f3755f..497cfe09ab 100644 --- a/backend/internal/service/openai_gateway_messages_anthropic_native.go +++ b/backend/internal/service/openai_gateway_messages_anthropic_native.go @@ -515,95 +515,6 @@ func (s *OpenAIGatewayService) nativeAnthropicStreamResult( } } -// forwardCountTokensViaNativeAnthropic 把 Anthropic count_tokens 请求透传到 -// 国产供应商原生 Anthropic 端点({base}/v1/messages/count_tokens), -// 仅做模型名映射,不做协议转换。 -func (s *OpenAIGatewayService) forwardCountTokensViaNativeAnthropic( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - defaultMappedModel string, -) error { - originalModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) - if originalModel == "" { - writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "model is required") - return fmt.Errorf("count_tokens: missing model in request") - } - billingModel := resolveOpenAIForwardModel(account, originalModel, strings.TrimSpace(defaultMappedModel)) - upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) - if upstreamModel != originalModel { - rewritten, err := sjson.SetBytes(body, "model", upstreamModel) - if err != nil { - return fmt.Errorf("count_tokens: rewrite model: %w", err) - } - body = rewritten - } - - apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey()) - if apiKey == "" { - writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Account api_key is missing") - return fmt.Errorf("count_tokens: account %d missing api_key", account.ID) - } - targetURL, err := s.nativeAnthropicTargetURL(account) - if err != nil { - return fmt.Errorf("count_tokens: %w", err) - } - targetURL = strings.TrimSuffix(targetURL, "/v1/messages") + "/v1/messages/count_tokens" - - upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) - if err != nil { - writeAnthropicCountTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request") - return fmt.Errorf("count_tokens: build request: %w", err) - } - reqHeader := upstreamReq.Header - reqHeader.Del("authorization") - reqHeader.Del("x-api-key") - setAnthropicAPIKeyAuthHeader(reqHeader, account, apiKey) - reqHeader.Set("content-type", "application/json") - reqHeader.Set("accept", "application/json") - account.ApplyHeaderOverrides(reqHeader) - - proxyURL := "" - if account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) - if err != nil { - safeErr := sanitizeUpstreamErrorMessage(err.Error()) - setOpsUpstreamError(c, 0, safeErr, "") - writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") - return fmt.Errorf("count_tokens: upstream request failed: %s", safeErr) - } - defer func() { _ = resp.Body.Close() }() - - // count_tokens 响应体极小;与其他探测路径一致加 256KB 上限防异常上游放大内存。 - respBody, err := io.ReadAll(io.LimitReader(resp.Body, cnQuotaMaxBodyBytes)) - if err != nil { - writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to read response") - return fmt.Errorf("count_tokens: read response: %w", err) - } - if resp.StatusCode >= 400 { - if s.rateLimitService != nil { - s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) - } - upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) - setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, "") - writeAnthropicCountTokensError(c, resp.StatusCode, "upstream_error", "Upstream request failed") - return fmt.Errorf("count_tokens: upstream error: %d", resp.StatusCode) - } - - inputTokens := gjson.GetBytes(respBody, "input_tokens") - if !inputTokens.Exists() { - writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream response missing input_tokens") - return fmt.Errorf("count_tokens: response missing input_tokens field") - } - c.JSON(http.StatusOK, gin.H{ - "input_tokens": int(inputTokens.Int()), - }) - return nil -} - // claudeUsageToOpenAIUsage 把 Anthropic 格式 usage 映射到 OpenAI 网关统一的 // 用量结构(字段一一对应)。 func claudeUsageToOpenAIUsage(u *ClaudeUsage) OpenAIUsage { diff --git a/backend/internal/service/openai_gateway_responses_anthropic_native.go b/backend/internal/service/openai_gateway_responses_anthropic_native.go index b72067fc4d..04cd7fc274 100644 --- a/backend/internal/service/openai_gateway_responses_anthropic_native.go +++ b/backend/internal/service/openai_gateway_responses_anthropic_native.go @@ -377,7 +377,11 @@ func (s *OpenAIGatewayService) handleResponsesStreamingFromNativeAnthropic( return resultWithUsage(), fmt.Errorf("stream data interval timeout") } - processAnthropicEvent := func(event *apicompat.AnthropicStreamEvent) bool { + // 与 CC 姊妹路径(handleCCStreamingFromNativeAnthropic.writeChunk)同语义: + // 客户端断开后不再写出,但继续排水上游至流自然结束——Anthropic 的最终 + // output_tokens 只在末尾 message_delta 携带,提前退出会把整段生成记成 ~1 + // token,payg 上游照常计费而平台漏记。状态机照常推进以保证 finalize 一致。 + processAnthropicEvent := func(event *apicompat.AnthropicStreamEvent) { if firstChunk { firstChunk = false ms := int(time.Since(startTime).Milliseconds()) @@ -392,6 +396,9 @@ func (s *OpenAIGatewayService) handleResponsesStreamingFromNativeAnthropic( } events := apicompat.AnthropicEventToResponsesEvents(event, state) + if clientDisconnected { + return + } for _, evt := range events { payload, err := json.Marshal(evt) if err != nil { @@ -406,14 +413,13 @@ func (s *OpenAIGatewayService) handleResponsesStreamingFromNativeAnthropic( eventType := gjson.GetBytes(restored, "type").String() if _, err := fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", eventType, restored); err != nil { clientDisconnected = true - return true + return } } } if len(events) > 0 { c.Writer.Flush() } - return false } for { @@ -447,21 +453,39 @@ func (s *OpenAIGatewayService) handleResponsesStreamingFromNativeAnthropic( continue } - if processAnthropicEvent(&event) { - return resultWithUsage(), nil - } + processAnthropicEvent(&event) } - // Finalize state machine(客户端已断开时仍执行,保证 usage 汇总完整)。 - if finalEvents := apicompat.FinalizeAnthropicResponsesStream(state); len(finalEvents) > 0 { + // Finalize state machine(客户端已断开时仍推进,保证 usage 汇总完整;仅在 + // 客户端仍连接时写出)。终态帧与逐事件路径一致过工具名反转与客户端工具还原, + // 避免流截断时终态帧携带改写后的工具名。 + if finalEvents := apicompat.FinalizeAnthropicResponsesStream(state); len(finalEvents) > 0 && !clientDisconnected { + wrote := false for _, evt := range finalEvents { - sse, err := apicompat.ResponsesEventToSSE(evt) + payload, err := json.Marshal(evt) if err != nil { continue } - fmt.Fprint(c.Writer, sse) //nolint:errcheck + payload = reverseToolNamesIfPresent(c, payload) + payloads, _, err := clientToolRestorer.RestoreEvent(payload) + if err != nil { + continue + } + for _, restored := range payloads { + eventType := gjson.GetBytes(restored, "type").String() + if _, err := fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", eventType, restored); err != nil { + clientDisconnected = true + break + } + wrote = true + } + if clientDisconnected { + break + } + } + if wrote { + c.Writer.Flush() } - c.Writer.Flush() } return resultWithUsage(), nil diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index c55bcdbeac..ab2b599734 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -200,6 +200,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec result.UpstreamModel, result.Model, ) + billingModels = s.filterCNProviderBillingModelCandidates(ctx, account, apiKey, billingModels) serviceTier := "" if result.ServiceTier != nil { serviceTier = strings.TrimSpace(*result.ServiceTier) @@ -253,7 +254,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec result.AudioUsage != nil || result.SearchCount > 0, ); responseModel != "" && !strings.EqualFold(responseModel, baselineBillingModel) { if identified, responseChannelPriced := s.hasIdentifiedOpenAIResponsePricing(ctx, responseModel, apiKey); identified { - responseModels := usageBillingModelCandidates(responseModel) + responseModels := s.filterCNProviderBillingModelCandidates(ctx, account, apiKey, usageBillingModelCandidates(responseModel)) responseCost, responseErr := s.calculateOpenAIRecordUsageCost( ctx, result, apiKey, responseModels, multiplier, imageMultiplier, videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate, @@ -585,7 +586,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( if tokenCost == nil { if tokenBillingAttempted { if lastErr == nil { - lastErr = errors.New("no non-empty billing model candidates") + lastErr = fmt.Errorf("%w: no non-empty billing model candidates", ErrModelPricingUnavailable) } return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) } @@ -593,8 +594,11 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( if searchCost != nil { return searchCost, nil } + // 空候选按「无价可循」处理并携带 ErrModelPricingUnavailable:上层据此走 + // 零成本+告警落账,而不是丢弃整条 usage 记录。CN 账号的 claude-* 候选被 + // filterCNProviderBillingModelCandidates 全数过滤后即落到这里。 if lastErr == nil { - lastErr = errors.New("openai usage billing model is empty") + lastErr = fmt.Errorf("%w: openai usage billing model is empty", ErrModelPricingUnavailable) } return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) } @@ -841,6 +845,36 @@ func groupMediaPricingLooksIncomplete(group *Group) bool { group.VideoPrice480P == nil && group.VideoPrice720P == nil && group.VideoPrice1080P == nil } +// filterCNProviderBillingModelCandidates 过滤国产供应商(kimi/zhipu/deepseek) +// 账号的计费候选模型名:claude-* 候选仅在运营者显式配置了分组/渠道定价时保留。 +// +// 背景:候选链的兜底候选含客户端请求的原始模型名。CN 上游的 Anthropic 兼容端点 +// 接受 claude-* 模型名但从不真正服务 Claude 模型;若放行,目录里的 Claude 价卡 +// 与 getFallbackPricing 的 "claude"→Sonnet 统一兜底会把 CN 流量按 Claude 原价 +// (数倍~数十倍)静默误计,且 usage 日志显示的正是 claude-* 名,无从察觉。 +// 候选全部落空时走既有的零成本+告警路径(openai_usage.pricing_missing_record_ +// zero_cost),与定价层「未知型号不回退以避免误计价」的既有设计意图一致; +// 运营者的修复手段是配置账号级 model_mapping(映射到已定价的 CN 模型)或 +// 分组/渠道显式定价。 +func (s *OpenAIGatewayService) filterCNProviderBillingModelCandidates(ctx context.Context, account *Account, apiKey *APIKey, candidates []string) []string { + if account == nil || !account.IsCNProvider() { + return candidates + } + out := make([]string, 0, len(candidates)) + for _, candidate := range candidates { + trimmed := strings.TrimSpace(candidate) + if trimmed == "" { + continue + } + if strings.Contains(strings.ToLower(trimmed), "claude") && + s.resolveOpenAIChannelPricing(ctx, trimmed, apiKey) == nil { + continue + } + out = append(out, candidate) + } + return out +} + func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, billingModel string, apiKey *APIKey) *ResolvedPricing { if s.resolver == nil || apiKey == nil || apiKey.Group == nil { return nil diff --git a/backend/internal/service/openai_messages_dispatch.go b/backend/internal/service/openai_messages_dispatch.go index aedfb1b3f7..562e9583c5 100644 --- a/backend/internal/service/openai_messages_dispatch.go +++ b/backend/internal/service/openai_messages_dispatch.go @@ -79,6 +79,13 @@ func (g *Group) ResolveMessagesDispatchModel(requestedModel string) string { return xai.ModelMappingWithOptions(opts)["claude-*"] } + // 国产供应商分组:调度级模型映射不适用(其配置被 sanitize 置空,且下方的 + // gpt-5.x 默认值是 openai 专属,发给 CN 上游必错)。模型改写完全交给账号级 + // model_mapping;anthropic 协议上游本身接受 claude-* 模型名。 + if IsCNProvider(g.Platform) { + return "" + } + cfg := normalizeOpenAIMessagesDispatchModelConfig(g.MessagesDispatchModelConfig) if mappedModel := strings.TrimSpace(cfg.ExactModelMappings[requestedModel]); mappedModel != "" { return mappedModel diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 0f14491914..b361901653 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -906,7 +906,10 @@ func (s *RateLimitService) handle403(ctx context.Context, account *Account, upst if account.Platform == PlatformAntigravity { return s.handleAntigravity403(ctx, account, upstreamMsg, responseBody) } - if account.Platform == PlatformOpenAI { + // 国产供应商与 openai 同口径:HTML 403(CDN/代理拦截页)不构成账号失效证据, + // 且 403 在 failover 状态集里会被逐账号重放——直接 SetError 会让一个坏请求/ + // 一层坏代理连环永久禁用整组账号。走 HTML 豁免 + N 次累计 + 临时冷却。 + if account.Platform == PlatformOpenAI || IsCNProvider(account.Platform) { return s.handleOpenAI403(ctx, account, upstreamMsg, responseBody) } // 非 Antigravity 平台:保持原有行为