diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index 3a8f2128c8d..7f085dcc2eb 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -1,6 +1,7 @@ package openai import ( + "errors" "fmt" "io" "net/http" @@ -18,6 +19,11 @@ import ( "github.com/gin-gonic/gin" ) +// Fallback token estimation is used only when a failed stream omits terminal +// usage. Bound retained deltas so large tool arguments cannot grow per-request +// memory and tokenizer CPU without limit. +const maxResponsesFallbackUsageBytes = 1 << 20 + func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { defer service.CloseResponseBodyGracefully(resp) @@ -83,6 +89,16 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp var usage = &dto.Usage{} var responseTextBuilder strings.Builder + var terminalErr *types.NewAPIError + var retryableTerminalFailure bool + + var responseHeaderSnapshot http.Header + var eventStreamHeadersValue any + var hadEventStreamHeaders bool + if info.ChannelType == constant.ChannelTypeCodex { + responseHeaderSnapshot = c.Writer.Header().Clone() + eventStreamHeadersValue, hadEventStreamHeaders = c.Get("event_stream_headers_set") + } helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { @@ -96,34 +112,35 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp if streamResponse.Response != nil && streamResponse.Response.ID != "" { c.Set(common.UpstreamResponseIdKey, streamResponse.Response.ID) } + if info.ChannelType == constant.ChannelTypeCodex && streamResponse.Type == "response.failed" { + applyResponsesTerminalUsage(c, usage, streamResponse.Response) + responseWritten := c.Writer.Written() + if responseWritten { + sendResponsesStreamData(c, streamResponse, data) + } + retryableTerminalFailure = !responseWritten + terminalErr = newCodexResponsesFailedError(streamResponse.Response, responseWritten) + sr.Stop(terminalErr) + return + } sendResponsesStreamData(c, streamResponse, data) switch streamResponse.Type { case "response.completed", "response.incomplete": - if streamResponse.Response != nil { - if streamResponse.Response.Usage != nil { - if streamResponse.Response.Usage.InputTokens != 0 { - usage.PromptTokens = streamResponse.Response.Usage.InputTokens - } - if streamResponse.Response.Usage.OutputTokens != 0 { - usage.CompletionTokens = streamResponse.Response.Usage.OutputTokens - } - if streamResponse.Response.Usage.TotalTokens != 0 { - usage.TotalTokens = streamResponse.Response.Usage.TotalTokens - } - if streamResponse.Response.Usage.InputTokensDetails != nil { - usage.PromptTokensDetails.CachedTokens = streamResponse.Response.Usage.InputTokensDetails.CachedTokens - usage.PromptTokensDetails.CacheWriteTokens = streamResponse.Response.Usage.InputTokensDetails.CacheWriteTokens - } - } - if streamResponse.Response.HasImageGenerationCall() { - c.Set("image_generation_call", true) - c.Set("image_generation_call_quality", streamResponse.Response.GetQuality()) - c.Set("image_generation_call_size", streamResponse.Response.GetSize()) - } + applyResponsesTerminalUsage(c, usage, streamResponse.Response) + case "response.done": + if info.ChannelType == constant.ChannelTypeCodex { + applyResponsesTerminalUsage(c, usage, streamResponse.Response) } - case "response.output_text.delta": - // 处理输出文本 - responseTextBuilder.WriteString(streamResponse.Delta) + case "response.output_text.delta", + "response.reasoning_summary_text.delta", + "response.reasoning_text.delta", + "response.function_call_arguments.delta", + "response.custom_tool_call_input.delta", + "response.mcp_call_arguments.delta", + "response.code_interpreter_call_code.delta": + // Preserve a text-equivalent fallback when a failed terminal event + // omits usage, including tool-only and reasoning-only output. + appendResponsesFallbackUsage(&responseTextBuilder, streamResponse.Delta) case dto.ResponsesOutputTypeItemDone: // 函数调用处理 if streamResponse.Item != nil { @@ -138,6 +155,14 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } } }) + if terminalErr != nil { + if retryableTerminalFailure { + restoreResponsesStreamAttemptState(c, info, responseHeaderSnapshot, eventStreamHeadersValue, hadEventStreamHeaders) + } else { + finalizeResponsesUsage(usage, &responseTextBuilder, info) + } + return usage, terminalErr + } // FRT watchdog: upstream accepted the request but never produced a data // event within constant.StreamingFirstResponseTimeout seconds. Surface as @@ -155,21 +180,154 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ) } - if usage.CompletionTokens == 0 { - // 计算输出文本的 token 数量 - tempStr := responseTextBuilder.String() - if len(tempStr) > 0 { - // 非正常结束,使用输出文本的 token 数量 - completionTokens := service.CountTextToken(tempStr, info.UpstreamModelName) - usage.CompletionTokens = completionTokens - } + finalizeResponsesUsage(usage, &responseTextBuilder, info) + + return usage, nil +} + +func appendResponsesFallbackUsage(builder *strings.Builder, delta string) { + if builder == nil || builder.Len() >= maxResponsesFallbackUsageBytes || delta == "" { + return } + remaining := maxResponsesFallbackUsageBytes - builder.Len() + if len(delta) > remaining { + delta = delta[:remaining] + } + builder.WriteString(delta) +} +func finalizeResponsesUsage(usage *dto.Usage, responseTextBuilder *strings.Builder, info *relaycommon.RelayInfo) { + if usage.CompletionTokens == 0 && responseTextBuilder.Len() > 0 { + usage.CompletionTokens = service.CountTextToken(responseTextBuilder.String(), info.UpstreamModelName) + } if usage.PromptTokens == 0 && usage.CompletionTokens != 0 { usage.PromptTokens = info.GetEstimatePromptTokens() } - usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens +} - return usage, nil +func restoreResponsesStreamAttemptState( + c *gin.Context, + info *relaycommon.RelayInfo, + headerSnapshot http.Header, + eventStreamHeadersValue any, + hadEventStreamHeaders bool, +) { + header := c.Writer.Header() + clear(header) + for key, values := range headerSnapshot { + header[key] = append([]string(nil), values...) + } + if hadEventStreamHeaders { + c.Set("event_stream_headers_set", eventStreamHeadersValue) + } else if c.Keys != nil { + delete(c.Keys, "event_stream_headers_set") + } + info.ResetStreamResponseStateForRetry() +} + +func applyResponsesTerminalUsage(c *gin.Context, usage *dto.Usage, response *dto.OpenAIResponsesResponse) { + if usage == nil || response == nil { + return + } + if response.Usage != nil { + if response.Usage.InputTokens != 0 { + usage.PromptTokens = response.Usage.InputTokens + } + if response.Usage.OutputTokens != 0 { + usage.CompletionTokens = response.Usage.OutputTokens + } + if response.Usage.TotalTokens != 0 { + usage.TotalTokens = response.Usage.TotalTokens + } + if response.Usage.InputTokensDetails != nil { + usage.PromptTokensDetails.CachedTokens = response.Usage.InputTokensDetails.CachedTokens + usage.PromptTokensDetails.CacheWriteTokens = response.Usage.InputTokensDetails.CacheWriteTokens + } + } + if response.HasImageGenerationCall() { + c.Set("image_generation_call", true) + c.Set("image_generation_call_quality", response.GetQuality()) + c.Set("image_generation_call_size", response.GetSize()) + } +} + +func newCodexResponsesFailedError(response *dto.OpenAIResponsesResponse, skipRetry bool) *types.NewAPIError { + options := make([]types.NewAPIErrorOptions, 0, 1) + if skipRetry { + options = append(options, types.ErrOptionWithSkipRetry()) + } + if response != nil { + if openAIError := response.GetOpenAIError(); openAIError != nil && openAIError.Message != "" { + return types.WithOpenAIError(*openAIError, codexResponsesFailedStatus(openAIError), options...) + } + } + return types.NewOpenAIError( + errors.New("codex upstream response failed"), + types.ErrorCodeBadResponse, + http.StatusInternalServerError, + options..., + ) +} + +func codexResponsesFailedStatus(openAIError *types.OpenAIError) int { + if openAIError == nil { + return http.StatusInternalServerError + } + errorType := strings.ToLower(strings.TrimSpace(openAIError.Type)) + errorCode := strings.ToLower(strings.TrimSpace(common.Interface2String(openAIError.Code))) + if statusCode := codexResponsesFailedIdentifierStatus(errorCode); statusCode != 0 { + return statusCode + } + if statusCode := codexResponsesFailedIdentifierStatus(errorType); statusCode != 0 { + return statusCode + } + return http.StatusInternalServerError +} + +func codexResponsesFailedIdentifierStatus(identifier string) int { + switch identifier { + case + "insufficient_quota", + "credit_balance_exhausted", + "organization_spend_limit_exceeded", + "project_spend_limit_exceeded", + "organization_usage_limit_exceeded", + "project_usage_limit_exceeded", + "rate_limit_error", + "rate_limit_exceeded": + return http.StatusTooManyRequests + case "permission_error", "permission_denied", "unsupported_country_region_territory": + return http.StatusForbidden + case "authentication_error", "invalid_api_key": + return http.StatusUnauthorized + case "invalid_request_error", + "invalid_request", + "bad_request", + "invalid_prompt", + "content_policy_violation", + "data_residency_mismatch", + "bio_policy", + "invalid_image", + "invalid_image_format", + "invalid_base64_image", + "invalid_image_url", + "image_too_large", + "image_too_small", + "image_parse_error", + "image_content_policy_violation", + "invalid_image_mode", + "image_file_too_large", + "unsupported_image_media_type", + "empty_image_file", + "failed_to_download_image", + "image_file_not_found": + return http.StatusBadRequest + case "overloaded_error", "overloaded", "service_unavailable": + return http.StatusServiceUnavailable + case "server_error": + return http.StatusInternalServerError + default: + return 0 + } } diff --git a/relay/channel/openai/relay_responses_test.go b/relay/channel/openai/relay_responses_test.go index 7ee0ada65f7..a5462b0ebeb 100644 --- a/relay/channel/openai/relay_responses_test.go +++ b/relay/channel/openai/relay_responses_test.go @@ -1,6 +1,7 @@ package openai import ( + "fmt" "io" "net/http" "net/http/httptest" @@ -10,6 +11,8 @@ import ( "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/dto" relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/service" + "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" ) @@ -64,3 +67,231 @@ func TestOaiResponsesStreamHandlerCapturesIncompleteUsage(t *testing.T) { t.Fatalf("incomplete event was not forwarded: %s", recorder.Body.String()) } } + +func TestOaiResponsesStreamHandlerResponseDoneIsCodexOnly(t *testing.T) { + upstream := strings.Join([]string{ + "event: response.done", + `data: {"type":"response.done","response":{"id":"resp_done","status":"completed","usage":{"input_tokens":71,"output_tokens":19,"total_tokens":90,"input_tokens_details":{"cached_tokens":7,"cache_write_tokens":2}}}}`, + "", + }, "\n") + + for _, test := range []struct { + name, want string + channel int + }{ + {name: "Codex captures usage", channel: constant.ChannelTypeCodex, want: "71/19/90"}, + {name: "other channels stay unchanged", channel: constant.ChannelTypeOpenAI, want: "0/0/0"}, + } { + t.Run(test.name, func(t *testing.T) { + recorder, ctx, info, resp := newResponsesStreamTest(t, upstream, test.channel) + usage, apiErr := OaiResponsesStreamHandler(ctx, info, resp) + if apiErr != nil { + t.Fatal(apiErr) + } + if got := fmt.Sprintf("%d/%d/%d", usage.PromptTokens, usage.CompletionTokens, usage.TotalTokens); got != test.want { + t.Fatalf("usage = %s, want %s", got, test.want) + } + if !strings.Contains(recorder.Body.String(), "response.done") { + t.Fatalf("done event was not forwarded: %s", recorder.Body.String()) + } + }) + } +} + +func TestOaiResponsesStreamHandlerCodexFailedBeforeCommitIsRetryable(t *testing.T) { + upstream := strings.Join([]string{ + "event: response.failed", + `data: {"type":"response.failed","response":{"id":"resp_failed","status":"failed","error":{"code":"server_error","message":"upstream blew up"}}}`, + "", + }, "\n") + + recorder, ctx, info, resp := newResponsesStreamTest(t, upstream, constant.ChannelTypeCodex) + info.ApiType = constant.APITypeCodex + ctx.Header("X-Existing", "keep") + resp.Header.Set("X-Reasoning-Included", "true") + resp.Header.Set("X-Codex-Turn-State", "state-from-failed-attempt") + _, apiErr := OaiResponsesStreamHandler(ctx, info, resp) + if apiErr == nil { + t.Fatal("expected Codex response.failed error") + } + if types.IsSkipRetryError(apiErr) { + t.Fatal("failure before response commit must remain retryable") + } + if recorder.Body.Len() != 0 || ctx.Writer.Written() { + t.Fatalf("failed event committed before retry: %s", recorder.Body.String()) + } + for _, header := range []string{"Content-Type", "Transfer-Encoding", "X-Reasoning-Included", "X-Codex-Turn-State"} { + if got := recorder.Header().Get(header); got != "" { + t.Fatalf("retryable failure retained %s %q", header, got) + } + } + if got := recorder.Header().Get("X-Existing"); got != "keep" { + t.Fatalf("pre-existing response header = %q", got) + } + if _, exists := ctx.Get("event_stream_headers_set"); exists { + t.Fatal("retryable failure retained event_stream_headers_set") + } + if info.HasSendResponse() || info.ReceivedResponseCount != 0 { + t.Fatalf("retryable failure retained first-response state: sent=%v received=%d", info.HasSendResponse(), info.ReceivedResponseCount) + } + + // A real retry reuses RelayInfo and the downstream writer. Prove the reset + // above re-arms both response observation and SSE headers. + retry := strings.Join([]string{ + "event: response.done", + `data: {"type":"response.done","response":{"id":"resp_done","status":"completed","usage":{"input_tokens":8,"output_tokens":2,"total_tokens":10}}}`, + "", + }, "\n") + resp = &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(retry)), + } + usage, apiErr := OaiResponsesStreamHandler(ctx, info, resp) + if apiErr != nil || !info.HasSendResponse() || info.ReceivedResponseCount != 1 { + t.Fatalf("retry error=%v sent=%v received=%d", apiErr, info.HasSendResponse(), info.ReceivedResponseCount) + } + if usage.PromptTokens != 8 || usage.CompletionTokens != 2 || recorder.Header().Get("Content-Type") != "text/event-stream" { + t.Fatalf("retry usage=%#v content-type=%q", *usage, recorder.Header().Get("Content-Type")) + } +} + +func TestOaiResponsesStreamHandlerCodexFailedAfterCommitSkipsRetry(t *testing.T) { + upstream := strings.Join([]string{ + "event: response.created", + `data: {"type":"response.created","response":{"id":"resp_failed","status":"in_progress"}}`, + "", + "event: response.output_text.delta", + `data: {"type":"response.output_text.delta","delta":"partial output"}`, + "", + "event: response.failed", + `data: {"type":"response.failed","response":{"id":"resp_failed","status":"failed","error":{"code":"server_error","message":"upstream blew up"},"usage":{"input_tokens":11,"output_tokens":3,"total_tokens":14}}}`, + "", + }, "\n") + + recorder, ctx, info, resp := newResponsesStreamTest(t, upstream, constant.ChannelTypeCodex) + usage, apiErr := OaiResponsesStreamHandler(ctx, info, resp) + if apiErr == nil { + t.Fatal("expected Codex response.failed error") + } + if !types.IsSkipRetryError(apiErr) { + t.Fatal("failure after response commit must skip retry") + } + if !strings.Contains(recorder.Body.String(), "response.created") || !strings.Contains(recorder.Body.String(), "response.failed") { + t.Fatalf("committed Codex events were not forwarded: %s", recorder.Body.String()) + } + if usage.PromptTokens != 11 || usage.CompletionTokens != 3 || usage.TotalTokens != 14 { + t.Fatalf("failed response usage = %#v", *usage) + } +} + +func TestOaiResponsesStreamHandlerCodexFailedAfterCommitEstimatesDeliveredUsage(t *testing.T) { + service.InitTokenEncoders() + tests := []struct { + name string + eventType string + delta string + }{ + {name: "text", eventType: "response.output_text.delta", delta: "partial output without terminal usage"}, + {name: "tool arguments", eventType: "response.function_call_arguments.delta", delta: `{"city":"Shanghai"}`}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + upstream := strings.Join([]string{ + "event: response.created", + `data: {"type":"response.created","response":{"id":"resp_failed","status":"in_progress"}}`, + "", + "event: " + test.eventType, + fmt.Sprintf(`data: {"type":%q,"item_id":"call_1","delta":%q}`, test.eventType, test.delta), + "", + "event: response.failed", + `data: {"type":"response.failed","response":{"id":"resp_failed","status":"failed","error":{"code":"server_error","message":"upstream blew up"}}}`, + "", + }, "\n") + + _, ctx, info, resp := newResponsesStreamTest(t, upstream, constant.ChannelTypeCodex) + info.SetEstimatePromptTokens(17) + usage, apiErr := OaiResponsesStreamHandler(ctx, info, resp) + if apiErr == nil || !types.IsSkipRetryError(apiErr) { + t.Fatalf("failed response error = %#v", apiErr) + } + if usage.PromptTokens != 17 || usage.CompletionTokens <= 0 || usage.TotalTokens <= 17 { + t.Fatalf("estimated failed response usage = %#v", *usage) + } + }) + } +} + +func TestAppendResponsesFallbackUsageCapsBufferedBytes(t *testing.T) { + var builder strings.Builder + delta := strings.Repeat("x", maxResponsesFallbackUsageBytes+1024) + + appendResponsesFallbackUsage(&builder, delta) + appendResponsesFallbackUsage(&builder, "ignored after cap") + + if builder.Len() != maxResponsesFallbackUsageBytes { + t.Fatalf("buffered bytes = %d, want %d", builder.Len(), maxResponsesFallbackUsageBytes) + } +} + +func TestCodexResponsesFailedStatus(t *testing.T) { + tests := []struct { + name string + errorType string + errorCode any + statusCode int + }{ + {name: "invalid request type", errorType: "invalid_request_error", statusCode: http.StatusBadRequest}, + {name: "image validation", errorCode: "unsupported_image_media_type", statusCode: http.StatusBadRequest}, + {name: "authentication", errorType: "authentication_error", statusCode: http.StatusUnauthorized}, + {name: "specific code overrides generic type", errorType: "invalid_request_error", errorCode: "permission_denied", statusCode: http.StatusForbidden}, + {name: "server code overrides generic client type", errorType: "invalid_request_error", errorCode: "server_error", statusCode: http.StatusInternalServerError}, + {name: "permanent quota", errorCode: "insufficient_quota", statusCode: http.StatusTooManyRequests}, + {name: "overloaded", errorCode: "service_unavailable", statusCode: http.StatusServiceUnavailable}, + {name: "unknown fallback", errorCode: "new_unrecognized_code", statusCode: http.StatusInternalServerError}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + response := &dto.OpenAIResponsesResponse{Error: types.OpenAIError{ + Type: test.errorType, + Code: test.errorCode, + Message: "upstream failure", + }} + apiErr := newCodexResponsesFailedError(response, false) + if apiErr.StatusCode != test.statusCode { + t.Fatalf("status = %d, want %d", apiErr.StatusCode, test.statusCode) + } + if types.IsSkipRetryError(apiErr) { + t.Fatal("pre-commit retry must remain controlled by the centralized status policy") + } + }) + } +} + +func newResponsesStreamTest(t *testing.T, upstream string, channelType int) (*httptest.ResponseRecorder, *gin.Context, *relaycommon.RelayInfo, *http.Response) { + t.Helper() + previousTimeout := constant.StreamingTimeout + if previousTimeout <= 0 { + constant.StreamingTimeout = 30 + } + t.Cleanup(func() { constant.StreamingTimeout = previousTimeout }) + + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + info := &relaycommon.RelayInfo{ + IsStream: true, + ChannelMeta: &relaycommon.ChannelMeta{ + ChannelType: channelType, + UpstreamModelName: "gpt-5.6-terra", + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstream)), + } + return recorder, ctx, info, resp +} diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index 3c05a468be2..09220f42758 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -671,6 +671,18 @@ func (info *RelayInfo) SetFirstResponseTime() { } } +// ResetStreamResponseStateForRetry restores the response-observation fields +// after an upstream attempt produced no downstream-visible response. Callers +// must only use it before retrying an uncommitted response. +func (info *RelayInfo) ResetStreamResponseStateForRetry() { + if info == nil { + return + } + info.FirstResponseTime = info.StartTime.Add(-time.Second) + info.isFirstResponse = true + info.ReceivedResponseCount = 0 +} + func (info *RelayInfo) HasSendResponse() bool { return info.FirstResponseTime.After(info.StartTime) } diff --git a/relay/responses_handler.go b/relay/responses_handler.go index cbe4f37297d..085584034c2 100644 --- a/relay/responses_handler.go +++ b/relay/responses_handler.go @@ -125,6 +125,13 @@ func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError * usage, newAPIError := adaptor.DoResponse(c, httpResp, info) if newAPIError != nil { + if info.ChannelType == appconstant.ChannelTypeCodex && c.Writer.Written() { + if usageDto, ok := usage.(*dto.Usage); ok && usageDto != nil && usageDto.TotalTokens > 0 { + if settleErr := service.PostTextConsumeQuotaOnError(c, info, usageDto, []string{"Codex 流式响应部分输出后失败,按已返回 usage 结算"}); settleErr != nil { + logger.LogError(c, "failed to settle delivered Codex response usage: "+settleErr.Error()) + } + } + } // reset status code 重置状态码 service.ResetStatusCode(newAPIError, statusCodeMappingStr) return newAPIError diff --git a/service/billing.go b/service/billing.go index 9c9fc48f222..c9cc1ed16d8 100644 --- a/service/billing.go +++ b/service/billing.go @@ -32,6 +32,11 @@ func PreConsumeBilling(c *gin.Context, preConsumedQuota int, relayInfo *relaycom // SettleBilling 执行计费结算。如果 RelayInfo 上有 BillingSession 则通过 session 结算, // 否则回退到旧的 PostConsumeQuota 路径(兼容按次计费等场景)。 func SettleBilling(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, actualQuota int) error { + _, err := settleBillingWithStatus(ctx, relayInfo, actualQuota) + return err +} + +func settleBillingWithStatus(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, actualQuota int) (bool, error) { if relayInfo.Billing != nil { preConsumed := relayInfo.Billing.GetPreConsumedQuota() delta := actualQuota - preConsumed @@ -55,7 +60,10 @@ func SettleBilling(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, actualQuo } if err := relayInfo.Billing.Settle(actualQuota); err != nil { - return err + if state, ok := relayInfo.Billing.(interface{ settlementApplied() bool }); ok { + return state.settlementApplied(), err + } + return false, err } // 发送额度通知(订阅计费使用订阅剩余额度) @@ -66,16 +74,17 @@ func SettleBilling(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, actualQuo checkAndSendWalletQuotaNonEmailNotify(relayInfo, delta, preConsumed) } } - return nil + return true, nil } // 回退:无 BillingSession 时使用旧路径 quotaDelta := actualQuota - relayInfo.FinalPreConsumedQuota if quotaDelta != 0 { - if err := PostConsumeQuota(relayInfo, quotaDelta, relayInfo.FinalPreConsumedQuota, true); err != nil { - return err + fundingApplied, err := postConsumeQuotaWithStatus(relayInfo, quotaDelta, relayInfo.FinalPreConsumedQuota, true) + if err != nil { + return fundingApplied, err } relayInfo.FinalPreConsumedQuota = actualQuota } - return nil + return true, nil } diff --git a/service/billing_session.go b/service/billing_session.go index f4bef075715..bc222355cb8 100644 --- a/service/billing_session.go +++ b/service/billing_session.go @@ -134,6 +134,15 @@ func (s *BillingSession) NeedsRefund() bool { return s.needsRefundLocked() } +// settlementApplied reports whether settlement has committed far enough that +// the original pre-consumption must not be refunded. It intentionally remains +// a service-local optional capability rather than widening BillingSettler. +func (s *BillingSession) settlementApplied() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.settled || s.fundingSettled +} + func (s *BillingSession) needsRefundLocked() bool { if s.settled || s.refunded || s.fundingSettled { // fundingSettled 时资金来源已提交结算,不能再退预扣费 diff --git a/service/billing_status_test.go b/service/billing_status_test.go index 80309e97afa..d122c53f5db 100644 --- a/service/billing_status_test.go +++ b/service/billing_status_test.go @@ -342,6 +342,26 @@ func TestPostConsumeQuotaTracksUnlimitedTokenQuota(t *testing.T) { require.Equal(t, 920, userQuota) } +func TestSettleBillingLegacyReportsCommittedWalletDebit(t *testing.T) { + const ( + userID = 10122 + tokenID = 10222 + tokenKey = "billing-status-legacy-partial-settlement" + ) + resetBillingStatusTables(t) + seedUser(t, userID, 1000) + seedToken(t, tokenID, userID, tokenKey, 1000) + blockTokenDebitForTest(t, tokenID, "legacy_token_debit_blocked") + + relayInfo := newQuotaStatusRelayInfo(userID, tokenID, tokenKey) + applied, err := settleBillingWithStatus(nil, relayInfo, 80) + + require.ErrorContains(t, err, "legacy_token_debit_blocked") + require.True(t, applied) + require.Equal(t, 920, getUserQuota(t, userID)) + require.Equal(t, 1000, getTokenRemainQuota(t, tokenID)) +} + func TestBillingSessionSettleTracksUnlimitedTokenQuota(t *testing.T) { const ( userID = 10117 @@ -724,6 +744,21 @@ func blockTokenCreditForTest(t *testing.T, tokenID int, message string) { }) } +func blockTokenDebitForTest(t *testing.T, tokenID int, message string) { + t.Helper() + triggerName := "test_block_token_debit_" + strings.ReplaceAll(message, "-", "_") + require.NoError(t, model.DB.Exec("DROP TRIGGER IF EXISTS "+triggerName).Error) + require.NoError(t, model.DB.Exec(fmt.Sprintf( + "CREATE TRIGGER %s BEFORE UPDATE OF remain_quota ON tokens "+ + "WHEN OLD.id = %d AND NEW.remain_quota < OLD.remain_quota "+ + "BEGIN SELECT RAISE(ABORT, '%s'); END", + triggerName, tokenID, message, + )).Error) + t.Cleanup(func() { + require.NoError(t, model.DB.Exec("DROP TRIGGER IF EXISTS "+triggerName).Error) + }) +} + func drainWalletAfterTokenDebitForTest(t *testing.T, userID int, tokenID int) { t.Helper() const triggerName = "test_drain_wallet_after_token_debit" diff --git a/service/quota.go b/service/quota.go index 70e42da65ef..3c0887364ff 100644 --- a/service/quota.go +++ b/service/quota.go @@ -416,18 +416,26 @@ func PreConsumeTokenQuota(relayInfo *relaycommon.RelayInfo, quota int) error { } func PostConsumeQuota(relayInfo *relaycommon.RelayInfo, quota int, preConsumedQuota int, sendEmail bool) (err error) { + _, err = postConsumeQuotaWithStatus(relayInfo, quota, preConsumedQuota, sendEmail) + return err +} + +// postConsumeQuotaWithStatus reports whether the funding-side mutation was +// committed before a later token-quota update failed. +func postConsumeQuotaWithStatus(relayInfo *relaycommon.RelayInfo, quota int, preConsumedQuota int, sendEmail bool) (fundingApplied bool, err error) { // 1) Consume from wallet quota OR subscription item if relayInfo != nil && relayInfo.BillingSource == BillingSourceSubscription { if relayInfo.SubscriptionId == 0 { - return errors.New("subscription id is missing") + return false, errors.New("subscription id is missing") } delta := int64(quota) if delta != 0 { if err := model.PostConsumeUserSubscriptionDelta(relayInfo.SubscriptionId, delta); err != nil { - return err + return false, err } relayInfo.SubscriptionPostDelta += delta + fundingApplied = true } } else { // Wallet @@ -437,8 +445,9 @@ func PostConsumeQuota(relayInfo *relaycommon.RelayInfo, quota int, preConsumedQu err = model.IncreaseUserQuota(relayInfo.UserId, -quota, false) } if err != nil { - return err + return false, err } + fundingApplied = quota != 0 } if !relayInfo.IsPlayground { @@ -448,7 +457,7 @@ func PostConsumeQuota(relayInfo *relaycommon.RelayInfo, quota int, preConsumedQu err = model.IncreaseTokenQuota(relayInfo.TokenId, relayInfo.TokenKey, -quota) } if err != nil { - return err + return fundingApplied, err } } @@ -460,7 +469,7 @@ func PostConsumeQuota(relayInfo *relaycommon.RelayInfo, quota int, preConsumedQu } } - return nil + return fundingApplied, nil } // notifyLang resolves the language for a background notification (no gin diff --git a/service/text_quota.go b/service/text_quota.go index a31d99cfd44..f3867599a97 100644 --- a/service/text_quota.go +++ b/service/text_quota.go @@ -326,6 +326,17 @@ func usageSemanticFromUsage(relayInfo *relaycommon.RelayInfo, usage *dto.Usage) } func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage, extraContent []string) { + _ = postTextConsumeQuota(ctx, relayInfo, usage, extraContent, true) +} + +// PostTextConsumeQuotaOnError settles usage that was already delivered before +// a terminal stream failure. The controller records the failed relay sample, so +// this path deliberately avoids recording a second sample as a success. +func PostTextConsumeQuotaOnError(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage, extraContent []string) error { + return postTextConsumeQuota(ctx, relayInfo, usage, extraContent, false) +} + +func postTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage *dto.Usage, extraContent []string, recordRelaySample bool) error { originUsage := usage if usage == nil { extraContent = append(extraContent, "上游无计费信息") @@ -371,13 +382,29 @@ func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, us if summary.TotalTokens == 0 { extraContent = append(extraContent, "上游没有返回计费信息,无法扣费(可能是上游超时)") logger.LogError(ctx, fmt.Sprintf("total tokens is 0, cannot consume quota, userId %d, channelId %d, tokenId %d, model %s, pre-consumed quota %d", relayInfo.UserId, relayInfo.ChannelId, relayInfo.TokenId, summary.ModelName, relayInfo.FinalPreConsumedQuota)) - } else { - model.UpdateUserUsedQuotaAndRequestCount(relayInfo.UserId, summary.Quota) - model.UpdateChannelUsedQuota(relayInfo.ChannelId, summary.Quota) } - if err := SettleBilling(ctx, relayInfo, summary.Quota); err != nil { - logger.LogError(ctx, "error settling billing: "+err.Error()) + settlementApplied, settleErr := settleBillingWithStatus(ctx, relayInfo, summary.Quota) + if settleErr != nil { + logger.LogError(ctx, "error settling billing: "+settleErr.Error()) + if !recordRelaySample && !settlementApplied && relayInfo.Billing != nil && relayInfo.Billing.NeedsRefund() { + retainedQuota := relayInfo.Billing.GetPreConsumedQuota() + if retainErr := relayInfo.Billing.Settle(retainedQuota); retainErr != nil { + logger.LogError(ctx, "error retaining pre-consumed quota after failed partial-response settlement: "+retainErr.Error()) + } else { + settlementApplied = true + summary.Quota = retainedQuota + extraContent = append(extraContent, fmt.Sprintf("实际 usage 结算失败,暂按预扣额度 %s 保留", logger.FormatQuota(retainedQuota))) + } + } + } + if !recordRelaySample && !settlementApplied { + return settleErr + } + + if summary.TotalTokens != 0 { + model.UpdateUserUsedQuotaAndRequestCount(relayInfo.UserId, summary.Quota) + model.UpdateChannelUsedQuota(relayInfo.ChannelId, summary.Quota) } logModel := summary.ModelName @@ -481,5 +508,8 @@ func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, us Other: other, }) perfmetrics.RecordChannelTokens(relayInfo, int64(summary.PromptTokens), int64(summary.CompletionTokens)) - perfmetrics.RecordRelaySample(relayInfo, true, int64(summary.CompletionTokens), nil) + if recordRelaySample { + perfmetrics.RecordRelaySample(relayInfo, true, int64(summary.CompletionTokens), nil) + } + return settleErr } diff --git a/service/text_quota_error_test.go b/service/text_quota_error_test.go new file mode 100644 index 00000000000..117f32775d9 --- /dev/null +++ b/service/text_quota_error_test.go @@ -0,0 +1,123 @@ +package service + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/model" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" +) + +type recordingBillingSettler struct { + settledQuotas []int + failures int + commitOnFailure bool + refundCalls int +} + +func (s *recordingBillingSettler) Settle(actualQuota int) error { + if s.failures > 0 { + s.failures-- + if s.commitOnFailure { + s.settledQuotas = append(s.settledQuotas, actualQuota) + } + return errors.New("injected settlement failure") + } + s.settledQuotas = append(s.settledQuotas, actualQuota) + return nil +} + +func (s *recordingBillingSettler) Refund(*gin.Context) { + if s.NeedsRefund() { + s.refundCalls++ + } +} +func (s *recordingBillingSettler) NeedsRefund() bool { return len(s.settledQuotas) == 0 } +func (s *recordingBillingSettler) GetPreConsumedQuota() int { return 100 } +func (s *recordingBillingSettler) Reserve(int) error { return nil } +func (s *recordingBillingSettler) settlementApplied() bool { return len(s.settledQuotas) > 0 } + +func newTextQuotaErrorTest(t *testing.T, billing *recordingBillingSettler, logConsume bool) (*gin.Context, *relaycommon.RelayInfo, *dto.Usage) { + t.Helper() + previousBatchUpdateEnabled := common.BatchUpdateEnabled + previousLogConsumeEnabled := common.LogConsumeEnabled + common.BatchUpdateEnabled = true + common.LogConsumeEnabled = logConsume + t.Cleanup(func() { + common.BatchUpdateEnabled = previousBatchUpdateEnabled + common.LogConsumeEnabled = previousLogConsumeEnabled + }) + + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + info := &relaycommon.RelayInfo{ + StartTime: time.Now(), + OriginModelName: "gpt-5.6-sol", + IsStream: true, + IsPlayground: true, + Billing: billing, + ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeCodex}, + PriceData: types.PriceData{ + ModelRatio: 1, + CompletionRatio: 1, + GroupRatioInfo: types.GroupRatioInfo{GroupRatio: 1}, + }, + } + return ctx, info, &dto.Usage{PromptTokens: 20, CompletionTokens: 5, TotalTokens: 25} +} + +func TestPostTextConsumeQuotaOnErrorSettlementOutcomes(t *testing.T) { + for _, test := range []struct { + name string + failures int + commitOnFailure bool + wantErr, wantApplied bool + wantRetained bool + wantConsumeLogs int + }{ + {name: "settles delivered usage", wantApplied: true, wantConsumeLogs: 1}, + {name: "retains pre-consumption after settlement failure", failures: 1, wantErr: true, wantApplied: true, wantRetained: true, wantConsumeLogs: 1}, + {name: "does not replace a committed settlement after a late error", failures: 1, commitOnFailure: true, wantErr: true, wantApplied: true, wantConsumeLogs: 1}, + {name: "does not log when settlement and retention fail", failures: 2, wantErr: true}, + } { + t.Run(test.name, func(t *testing.T) { + previousSpendHook := model.TemporaryChannelSpendHook + consumeLogCalls := 0 + model.TemporaryChannelSpendHook = func(int, string, int) { consumeLogCalls++ } + t.Cleanup(func() { model.TemporaryChannelSpendHook = previousSpendHook }) + + billing := &recordingBillingSettler{failures: test.failures, commitOnFailure: test.commitOnFailure} + ctx, info, usage := newTextQuotaErrorTest(t, billing, true) + err := PostTextConsumeQuotaOnError(ctx, info, usage, nil) + if (err != nil) != test.wantErr { + t.Fatalf("error = %v, wantErr %v", err, test.wantErr) + } + if billing.settlementApplied() != test.wantApplied || billing.NeedsRefund() == test.wantApplied { + t.Fatalf("settlement applied=%v needsRefund=%v", billing.settlementApplied(), billing.NeedsRefund()) + } + if consumeLogCalls != test.wantConsumeLogs { + t.Fatalf("consume logs = %d, want %d", consumeLogCalls, test.wantConsumeLogs) + } + if test.wantApplied { + gotRetained := billing.settledQuotas[0] == billing.GetPreConsumedQuota() + if gotRetained != test.wantRetained { + t.Fatalf("settled quotas = %v, retained=%v", billing.settledQuotas, gotRetained) + } + } + billing.Refund(ctx) + if test.wantApplied && billing.refundCalls != 0 { + t.Fatalf("refund calls = %d", billing.refundCalls) + } + }) + } +}