mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:33:18 +08:00
fix(gateway): record observed usage when anthropic stream is interrupted
Fixes #5148: with aggregator upstreams (e.g. newapi) that end SSE streams without a proper terminal event, every such request was silently missing from usage logs and billing. Root cause (tracked via the nested audit issue): the low-level Anthropic SSE readers already return the partially collected usage together with the stream error (missing terminal event, read error, interval timeout), but Forward converted every such result to (nil, err) and the handler returned before submitting RecordUsage. Changes: - Add partialStreamUsageResult: on stream errors, wrap observed usage into a ForwardResult and return it alongside the error, for both the regular Anthropic path and the API-key passthrough path. Invariants: UpstreamFailoverError always keeps result=nil (failover retries are billed as the successful attempt, never twice), and zero observed usage returns no partial result (no phantom zero-usage records). - Messages handler: hoist the usage submission block into a closure shared by the success path and the new partial-result error path. - Usage record worker pool: distinguish pool-stopped drops (dropped_stopped) from operator-configured drop/sample overflow drops; billing tasks now fall back to inline synchronous execution only during the shutdown window, while explicit drop/sample overflow semantics are preserved. Image usage keeps its mandatory fallback for both drop kinds via the new mode.Dropped() helper. Tests: Forward-level regressions for missing-terminal / read-error / no-usage / failover-invariant on both paths, plus handler-level stopped-pool sync fallback and drop-policy preservation tests.
This commit is contained in:
@@ -827,6 +827,65 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
||||
if accountReleaseFunc != nil {
|
||||
accountReleaseFunc()
|
||||
}
|
||||
|
||||
// 提交 usage 记录。成功路径与"流中断但 Forward 已观测到 usage 的部分结果"
|
||||
// 错误路径共用:后者若不入账,上游已计量的请求会完全漏记漏计费(#5148)。
|
||||
submitForwardUsage := func(result *service.ForwardResult) {
|
||||
// 捕获请求信息(用于异步记录,避免在 goroutine 中访问 gin.Context)
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
clientIP := ip.GetClientIP(c)
|
||||
// Forward 内部可能继续改写 body,usage 去重指纹必须使用最终上游接受的当前 body。
|
||||
requestPayloadHash := service.HashUsageRequestPayload(attemptParsedReq.Body.Bytes())
|
||||
inboundEndpoint := GetInboundEndpoint(c)
|
||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||
|
||||
if result.ReasoningEffort == nil {
|
||||
result.ReasoningEffort = service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort)
|
||||
}
|
||||
// 同上(重试路径中的对称填充)。详见非重试路径同名注释。
|
||||
if result.ReasoningEffort == nil && attemptParsedReq.ThinkingEnabled {
|
||||
protocolModel := result.UpstreamModel
|
||||
if protocolModel == "" {
|
||||
protocolModel = result.Model
|
||||
}
|
||||
result.ReasoningEffort = service.DefaultEffortForThinkingEnabled(protocolModel)
|
||||
}
|
||||
|
||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
|
||||
forceCacheBilling := fs.ForceCacheBilling
|
||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
|
||||
sessionID := service.ExtractClientSessionID(c)
|
||||
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
||||
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
||||
Result: result,
|
||||
QuotaPlatform: quotaPlatform,
|
||||
APIKey: currentAPIKey,
|
||||
User: currentAPIKey.User,
|
||||
Account: account,
|
||||
Subscription: currentSubscription,
|
||||
InboundEndpoint: inboundEndpoint,
|
||||
UpstreamEndpoint: upstreamEndpoint,
|
||||
UserAgent: userAgent,
|
||||
IPAddress: clientIP,
|
||||
SessionID: sessionID,
|
||||
RequestPayloadHash: requestPayloadHash,
|
||||
ForceCacheBilling: forceCacheBilling,
|
||||
APIKeyService: h.apiKeyService,
|
||||
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
|
||||
}); err != nil {
|
||||
logger.L().With(
|
||||
zap.String("component", "handler.gateway.messages"),
|
||||
zap.Int64("user_id", subject.UserID),
|
||||
zap.Int64("api_key_id", currentAPIKey.ID),
|
||||
zap.Any("group_id", currentAPIKey.GroupID),
|
||||
zap.String("model", reqModel),
|
||||
zap.Int64("account_id", account.ID),
|
||||
).Error("gateway.record_usage_failed", zap.Error(err))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
// Beta policy block: return 400 immediately, no failover
|
||||
var betaBlockedErr *service.BetaBlockedError
|
||||
@@ -925,6 +984,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
||||
forwardFailedFields = append(forwardFailedFields, zap.Int64p("proxy_id", account.ProxyID))
|
||||
}
|
||||
reqLog.Error("gateway.forward_failed", forwardFailedFields...)
|
||||
// Forward 与错误一起返回的部分结果:流中断前上游已计量的 usage 照常入账,
|
||||
// 避免上游已产生消耗的请求完全漏记(#5148)。failover 错误恒定 result=nil,
|
||||
// 不会走到这里重复计费。
|
||||
if result != nil {
|
||||
submitForwardUsage(result)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -948,59 +1013,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
// 捕获请求信息(用于异步记录,避免在 goroutine 中访问 gin.Context)
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
clientIP := ip.GetClientIP(c)
|
||||
// Forward 内部可能继续改写 body,usage 去重指纹必须使用最终上游接受的当前 body。
|
||||
requestPayloadHash := service.HashUsageRequestPayload(attemptParsedReq.Body.Bytes())
|
||||
inboundEndpoint := GetInboundEndpoint(c)
|
||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||
|
||||
if result.ReasoningEffort == nil {
|
||||
result.ReasoningEffort = service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort)
|
||||
}
|
||||
// 同上(重试路径中的对称填充)。详见非重试路径同名注释。
|
||||
if result.ReasoningEffort == nil && attemptParsedReq.ThinkingEnabled {
|
||||
protocolModel := result.UpstreamModel
|
||||
if protocolModel == "" {
|
||||
protocolModel = result.Model
|
||||
}
|
||||
result.ReasoningEffort = service.DefaultEffortForThinkingEnabled(protocolModel)
|
||||
}
|
||||
|
||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
|
||||
forceCacheBilling := fs.ForceCacheBilling
|
||||
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
|
||||
sessionID := service.ExtractClientSessionID(c)
|
||||
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
|
||||
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
|
||||
Result: result,
|
||||
QuotaPlatform: quotaPlatform,
|
||||
APIKey: currentAPIKey,
|
||||
User: currentAPIKey.User,
|
||||
Account: account,
|
||||
Subscription: currentSubscription,
|
||||
InboundEndpoint: inboundEndpoint,
|
||||
UpstreamEndpoint: upstreamEndpoint,
|
||||
UserAgent: userAgent,
|
||||
IPAddress: clientIP,
|
||||
SessionID: sessionID,
|
||||
RequestPayloadHash: requestPayloadHash,
|
||||
ForceCacheBilling: forceCacheBilling,
|
||||
APIKeyService: h.apiKeyService,
|
||||
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
|
||||
}); err != nil {
|
||||
logger.L().With(
|
||||
zap.String("component", "handler.gateway.messages"),
|
||||
zap.Int64("user_id", subject.UserID),
|
||||
zap.Int64("api_key_id", currentAPIKey.ID),
|
||||
zap.Any("group_id", currentAPIKey.GroupID),
|
||||
zap.String("model", reqModel),
|
||||
zap.Int64("account_id", account.ID),
|
||||
).Error("gateway.record_usage_failed", zap.Error(err))
|
||||
}
|
||||
})
|
||||
submitForwardUsage(result)
|
||||
return
|
||||
}
|
||||
if !retryWithFallback {
|
||||
@@ -2336,10 +2349,16 @@ func (h *GatewayHandler) submitUsageRecordTask(parent context.Context, task serv
|
||||
}
|
||||
task = wrapUsageRecordTaskContext(parent, task)
|
||||
if h.usageRecordWorkerPool != nil {
|
||||
h.usageRecordWorkerPool.Submit(task)
|
||||
return
|
||||
if mode := h.usageRecordWorkerPool.Submit(task); mode != service.UsageRecordSubmitModeDroppedStopped {
|
||||
return
|
||||
}
|
||||
// 池已停止(进程关停窗口):计费任务不能静默丢失,降级为内联同步执行。
|
||||
// 显式配置的 drop/sample 溢出丢弃仍按配置语义保留。
|
||||
logger.L().With(
|
||||
zap.String("component", "handler.gateway.messages"),
|
||||
).Warn("gateway.usage_record_task_stopped_sync_fallback")
|
||||
}
|
||||
// 回退路径:worker 池未注入时同步执行,避免退回到无界 goroutine 模式。
|
||||
// 回退路径:worker 池未注入或已停止时同步执行,避免退回到无界 goroutine 模式。
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
defer func() {
|
||||
|
||||
@@ -2183,10 +2183,16 @@ func (h *OpenAIGatewayHandler) submitUsageRecordTask(parent context.Context, tas
|
||||
}
|
||||
task = wrapUsageRecordTaskContext(parent, task)
|
||||
if h.usageRecordWorkerPool != nil {
|
||||
h.usageRecordWorkerPool.Submit(task)
|
||||
return
|
||||
if mode := h.usageRecordWorkerPool.Submit(task); mode != service.UsageRecordSubmitModeDroppedStopped {
|
||||
return
|
||||
}
|
||||
// 池已停止(进程关停窗口):计费任务不能静默丢失,降级为内联同步执行。
|
||||
// 显式配置的 drop/sample 溢出丢弃仍按配置语义保留。
|
||||
logger.L().With(
|
||||
zap.String("component", "handler.openai_gateway.responses"),
|
||||
).Warn("openai.usage_record_task_stopped_sync_fallback")
|
||||
}
|
||||
// 回退路径:worker 池未注入时同步执行,避免退回到无界 goroutine 模式。
|
||||
// 回退路径:worker 池未注入或已停止时同步执行,避免退回到无界 goroutine 模式。
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
defer func() {
|
||||
@@ -2214,7 +2220,7 @@ func (h *OpenAIGatewayHandler) submitMandatoryUsageRecordTask(parent context.Con
|
||||
}
|
||||
task = wrapUsageRecordTaskContext(parent, task)
|
||||
if h.usageRecordWorkerPool != nil {
|
||||
if mode := h.usageRecordWorkerPool.Submit(task); mode != service.UsageRecordSubmitModeDropped {
|
||||
if mode := h.usageRecordWorkerPool.Submit(task); !mode.Dropped() {
|
||||
return
|
||||
}
|
||||
logger.L().With(
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// 本文件覆盖:worker 池已停止(进程关停窗口)时,计费任务不得静默丢失,
|
||||
// 必须降级为内联同步执行;显式配置的 drop/sample 溢出丢弃仍按配置语义保留。
|
||||
|
||||
func newStoppedUsageRecordPoolForTest() *service.UsageRecordWorkerPool {
|
||||
pool := service.NewUsageRecordWorkerPoolWithOptions(service.UsageRecordWorkerPoolOptions{
|
||||
WorkerCount: 1,
|
||||
QueueSize: 1,
|
||||
TaskTimeout: time.Second,
|
||||
OverflowPolicy: config.UsageRecordOverflowPolicySync,
|
||||
})
|
||||
pool.Stop()
|
||||
return pool
|
||||
}
|
||||
|
||||
func TestGatewayHandlerSubmitUsageRecordTask_StoppedPoolFallsBackToSync(t *testing.T) {
|
||||
h := &GatewayHandler{usageRecordWorkerPool: newStoppedUsageRecordPoolForTest()}
|
||||
|
||||
executed := false
|
||||
h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
|
||||
executed = true
|
||||
})
|
||||
require.True(t, executed, "池已停止时计费任务必须内联同步执行")
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayHandlerSubmitUsageRecordTask_StoppedPoolFallsBackToSync(t *testing.T) {
|
||||
h := &OpenAIGatewayHandler{usageRecordWorkerPool: newStoppedUsageRecordPoolForTest()}
|
||||
|
||||
executed := false
|
||||
h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
|
||||
executed = true
|
||||
})
|
||||
require.True(t, executed, "池已停止时计费任务必须内联同步执行")
|
||||
}
|
||||
|
||||
func TestGatewayHandlerSubmitUsageRecordTask_DropPolicyOverflowStillDrops(t *testing.T) {
|
||||
pool := service.NewUsageRecordWorkerPoolWithOptions(service.UsageRecordWorkerPoolOptions{
|
||||
WorkerCount: 1,
|
||||
QueueSize: 1,
|
||||
TaskTimeout: time.Minute,
|
||||
OverflowPolicy: config.UsageRecordOverflowPolicyDrop,
|
||||
})
|
||||
t.Cleanup(pool.Stop)
|
||||
h := &GatewayHandler{usageRecordWorkerPool: pool}
|
||||
|
||||
started := make(chan struct{})
|
||||
block := make(chan struct{})
|
||||
t.Cleanup(func() { close(block) })
|
||||
// 占满 worker 槽位后再填满队列,保证第三个任务触发溢出。
|
||||
require.Equal(t, service.UsageRecordSubmitModeEnqueued, pool.Submit(func(ctx context.Context) {
|
||||
close(started)
|
||||
<-block
|
||||
}))
|
||||
<-started
|
||||
require.Equal(t, service.UsageRecordSubmitModeEnqueued, pool.Submit(func(ctx context.Context) {
|
||||
<-block
|
||||
}))
|
||||
|
||||
var executed atomic.Bool
|
||||
h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
|
||||
executed.Store(true)
|
||||
})
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
require.False(t, executed.Load(), "drop 溢出策略是运维显式配置的取舍,不应被同步兜底覆盖")
|
||||
}
|
||||
@@ -269,6 +269,11 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
|
||||
if input.RequestStream {
|
||||
streamResult, err := s.handleStreamingResponseAnthropicAPIKeyPassthrough(ctx, resp, c, account, input.StartTime, input.RequestModel)
|
||||
if err != nil {
|
||||
// 流中断时保留已观测到的 usage 与错误一起返回,避免上游已计量的请求
|
||||
// 完全漏记漏计费(issue #5148)。
|
||||
if partial := partialStreamUsageResult(resp, streamResult, input.OriginalModel, input.RequestModel, input.StartTime, err); partial != nil {
|
||||
return partial, err
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
usage = streamResult.usage
|
||||
|
||||
@@ -843,6 +843,11 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
|
||||
ResponseBody: body,
|
||||
}
|
||||
}
|
||||
// 流中断(缺失 terminal 事件、读错误、数据间隔超时等)时保留已观测到的
|
||||
// usage 与错误一起返回,handler 在错误处理完成后照常提交 usage 记录。
|
||||
if partial := partialStreamUsageResult(resp, streamResult, originalModel, mappedModel, startTime, err); partial != nil {
|
||||
return partial, err
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
usage = streamResult.usage
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"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"
|
||||
)
|
||||
|
||||
// 本文件覆盖 issue #5148:流式转发中途出错(缺失 terminal 事件、读错误等)时,
|
||||
// 已观测到的上游 usage 不得随错误一起被丢弃,Forward 必须把部分结果与错误一同
|
||||
// 返回,供 handler 照常提交 usage 记录。
|
||||
|
||||
func newForwardPartialUsageServiceForTest(upstream *anthropicHTTPUpstreamRecorder) *GatewayService {
|
||||
cfg := &config.Config{
|
||||
Gateway: config.GatewayConfig{
|
||||
MaxLineSize: defaultMaxLineSize,
|
||||
},
|
||||
}
|
||||
return &GatewayService{
|
||||
cfg: cfg,
|
||||
responseHeaderFilter: compileResponseHeaderFilter(cfg),
|
||||
httpUpstream: upstream,
|
||||
rateLimitService: &RateLimitService{},
|
||||
deferredService: &DeferredService{},
|
||||
}
|
||||
}
|
||||
|
||||
func newAnthropicOAuthAccountForPartialUsageTest() *Account {
|
||||
return &Account{
|
||||
ID: 501,
|
||||
Name: "anthropic-oauth-partial-usage",
|
||||
Platform: PlatformAnthropic,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "oauth-token",
|
||||
},
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayService_Forward_StreamMissingTerminalPreservesPartialUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
body := []byte(`{"model":"claude-3-5-sonnet-latest","stream":true,"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), PlatformAnthropic)
|
||||
require.NoError(t, err)
|
||||
|
||||
// newapi 等聚合上游的典型失败形态:message_start/message_delta 携带 usage,
|
||||
// 但流在 message_stop 前直接结束。
|
||||
upstreamSSE := strings.Join([]string{
|
||||
`event: message_start`,
|
||||
`data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-3-5-sonnet-latest","content":[],"usage":{"input_tokens":11,"cache_read_input_tokens":7}}}`,
|
||||
"",
|
||||
`event: content_block_delta`,
|
||||
`data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}`,
|
||||
"",
|
||||
`event: message_delta`,
|
||||
`data: {"type":"message_delta","delta":{"stop_reason":null},"usage":{"output_tokens":5}}`,
|
||||
"",
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &anthropicHTTPUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{
|
||||
"Content-Type": []string{"text/event-stream"},
|
||||
"X-Request-Id": []string{"rid-partial"},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamSSE)),
|
||||
}}
|
||||
svc := newForwardPartialUsageServiceForTest(upstream)
|
||||
account := newAnthropicOAuthAccountForPartialUsageTest()
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, parsed)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "missing terminal event")
|
||||
require.NotNil(t, result, "流中断但已观测到 usage 时必须返回部分结果用于计费")
|
||||
require.True(t, result.Stream)
|
||||
require.Equal(t, 11, result.Usage.InputTokens)
|
||||
require.Equal(t, 7, result.Usage.CacheReadInputTokens)
|
||||
require.Equal(t, 5, result.Usage.OutputTokens)
|
||||
require.Equal(t, "rid-partial", result.RequestID)
|
||||
require.NotNil(t, result.FirstTokenMs)
|
||||
}
|
||||
|
||||
func TestGatewayService_Forward_StreamReadErrorAfterOutputPreservesPartialUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
body := []byte(`{"model":"claude-3-5-sonnet-latest","stream":true,"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), PlatformAnthropic)
|
||||
require.NoError(t, err)
|
||||
|
||||
// message_start 已写出(含 usage),随后上游连接异常中断。
|
||||
upstream := &anthropicHTTPUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: &streamReadCloser{
|
||||
payload: []byte("data: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":9,\"cache_creation_input_tokens\":4}}}\n\n"),
|
||||
err: io.ErrUnexpectedEOF,
|
||||
},
|
||||
}}
|
||||
svc := newForwardPartialUsageServiceForTest(upstream)
|
||||
account := newAnthropicOAuthAccountForPartialUsageTest()
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, parsed)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "stream read error")
|
||||
require.NotNil(t, result, "已写出内容后的读错误必须保留部分 usage")
|
||||
require.Equal(t, 9, result.Usage.InputTokens)
|
||||
require.Equal(t, 4, result.Usage.CacheCreationInputTokens)
|
||||
}
|
||||
|
||||
func TestGatewayService_Forward_StreamErrorWithoutUsageReturnsNilResult(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
body := []byte(`{"model":"claude-3-5-sonnet-latest","stream":true,"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), PlatformAnthropic)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 只有 ping、没有任何 usage 的流中断:不应产生零 usage 的幽灵账单记录。
|
||||
upstream := &anthropicHTTPUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader("event: ping\ndata: {\"type\": \"ping\"}\n\n")),
|
||||
}}
|
||||
svc := newForwardPartialUsageServiceForTest(upstream)
|
||||
account := newAnthropicOAuthAccountForPartialUsageTest()
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, parsed)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "missing terminal event")
|
||||
require.Nil(t, result, "无已观测 usage 时不应返回部分结果")
|
||||
}
|
||||
|
||||
func TestGatewayService_Forward_FailoverErrorKeepsNilResult(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
body := []byte(`{"model":"claude-3-5-sonnet-latest","stream":true,"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), PlatformAnthropic)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 未向客户端写出任何字节前的读错误会包成 UpstreamFailoverError 走换号重试。
|
||||
// 该路径必须保持 result=nil:failover 成功后按成功请求计费,双份结果会重复计费。
|
||||
upstream := &anthropicHTTPUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: &streamReadCloser{
|
||||
err: errors.New("connection reset by peer"),
|
||||
},
|
||||
}}
|
||||
svc := newForwardPartialUsageServiceForTest(upstream)
|
||||
account := newAnthropicOAuthAccountForPartialUsageTest()
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, parsed)
|
||||
require.Error(t, err)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.True(t, errors.As(err, &failoverErr))
|
||||
require.Nil(t, result, "failover 错误必须保持 result=nil,防止重试成功后双重计费")
|
||||
}
|
||||
|
||||
func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardStreamMissingTerminalPreservesPartialUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
body := []byte(`{"model":"claude-3-7-sonnet-20250219","stream":true,"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
|
||||
parsed := &ParsedRequest{
|
||||
Body: NewRequestBodyRef(body),
|
||||
Model: "claude-3-7-sonnet-20250219",
|
||||
Stream: true,
|
||||
}
|
||||
|
||||
upstreamSSE := strings.Join([]string{
|
||||
`data: {"type":"message_start","message":{"usage":{"input_tokens":9,"cache_read_input_tokens":2}}}`,
|
||||
"",
|
||||
`data: {"type":"message_delta","usage":{"output_tokens":3}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &anthropicHTTPUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{
|
||||
"Content-Type": []string{"text/event-stream"},
|
||||
"X-Request-Id": []string{"rid-pass-partial"},
|
||||
},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamSSE)),
|
||||
}}
|
||||
svc := newForwardPartialUsageServiceForTest(upstream)
|
||||
account := newAnthropicAPIKeyAccountForTest()
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, parsed)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "missing terminal event")
|
||||
require.NotNil(t, result, "透传流中断但已观测到 usage 时必须返回部分结果用于计费")
|
||||
require.True(t, result.Stream)
|
||||
require.Equal(t, 9, result.Usage.InputTokens)
|
||||
require.Equal(t, 2, result.Usage.CacheReadInputTokens)
|
||||
require.Equal(t, 3, result.Usage.OutputTokens)
|
||||
require.Equal(t, "claude-3-7-sonnet-20250219", result.Model)
|
||||
}
|
||||
@@ -647,6 +647,44 @@ type streamingResult struct {
|
||||
clientDisconnect bool // 客户端是否在流式传输过程中断开
|
||||
}
|
||||
|
||||
// hasObservedTokens 报告流式过程中是否已观测到任何上游计量的 token。
|
||||
func (u *ClaudeUsage) hasObservedTokens() bool {
|
||||
if u == nil {
|
||||
return false
|
||||
}
|
||||
return u.InputTokens > 0 || u.OutputTokens > 0 ||
|
||||
u.CacheCreationInputTokens > 0 || u.CacheReadInputTokens > 0 ||
|
||||
u.CacheCreation5mTokens > 0 || u.CacheCreation1hTokens > 0 ||
|
||||
u.ImageOutputTokens > 0
|
||||
}
|
||||
|
||||
// partialStreamUsageResult 在流式转发中途出错时,把已观测到 usage 的部分结果包装为
|
||||
// ForwardResult(与错误一起返回给 handler 记录)。上游一旦下发过 message_start,
|
||||
// input/cache token 就已计量,直接丢弃会让请求完全漏记漏计费(issue #5148)。
|
||||
// 无已观测 usage 时返回 nil。
|
||||
//
|
||||
// 不变式:UpstreamFailoverError 必须保持 result=nil——failover 重试成功后按成功请求
|
||||
// 计费,若同时返回部分 usage 会造成双重计费,此处显式拦截兜底。
|
||||
func partialStreamUsageResult(resp *http.Response, streamResult *streamingResult, model, upstreamModel string, startTime time.Time, err error) *ForwardResult {
|
||||
if streamResult == nil || !streamResult.usage.hasObservedTokens() {
|
||||
return nil
|
||||
}
|
||||
var failoverErr *UpstreamFailoverError
|
||||
if errors.As(err, &failoverErr) {
|
||||
return nil
|
||||
}
|
||||
return &ForwardResult{
|
||||
RequestID: resp.Header.Get("x-request-id"),
|
||||
Usage: *streamResult.usage,
|
||||
Model: model,
|
||||
UpstreamModel: upstreamModel,
|
||||
Stream: true,
|
||||
Duration: time.Since(startTime),
|
||||
FirstTokenMs: streamResult.firstTokenMs,
|
||||
ClientDisconnect: streamResult.clientDisconnect,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string, mimicClaudeCode bool) (*streamingResult, error) {
|
||||
// 更新5h窗口状态
|
||||
s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header)
|
||||
|
||||
@@ -43,9 +43,18 @@ type UsageRecordSubmitMode string
|
||||
const (
|
||||
UsageRecordSubmitModeEnqueued UsageRecordSubmitMode = "enqueued"
|
||||
UsageRecordSubmitModeDropped UsageRecordSubmitMode = "dropped"
|
||||
UsageRecordSubmitModeSync UsageRecordSubmitMode = "sync_fallback"
|
||||
// UsageRecordSubmitModeDroppedStopped 表示任务因池已停止(进程关停窗口)被丢弃。
|
||||
// 与显式 drop/sample 溢出策略的丢弃区分开:溢出丢弃是运维显式配置的取舍,
|
||||
// 而关停窗口丢弃不是,计费关键任务应在调用侧降级为同步执行兜底。
|
||||
UsageRecordSubmitModeDroppedStopped UsageRecordSubmitMode = "dropped_stopped"
|
||||
UsageRecordSubmitModeSync UsageRecordSubmitMode = "sync_fallback"
|
||||
)
|
||||
|
||||
// Dropped 报告任务是否未被执行(入队失败且未同步执行)。
|
||||
func (m UsageRecordSubmitMode) Dropped() bool {
|
||||
return m == UsageRecordSubmitModeDropped || m == UsageRecordSubmitModeDroppedStopped
|
||||
}
|
||||
|
||||
// UsageRecordWorkerPoolOptions 使用量记录池配置。
|
||||
type UsageRecordWorkerPoolOptions struct {
|
||||
WorkerCount int
|
||||
@@ -150,7 +159,7 @@ func (p *UsageRecordWorkerPool) Submit(task UsageRecordTask) UsageRecordSubmitMo
|
||||
if p.pool == nil || p.pool.Stopped() {
|
||||
p.droppedPoolStopped.Add(1)
|
||||
p.logDrop("stopped")
|
||||
return UsageRecordSubmitModeDropped
|
||||
return UsageRecordSubmitModeDroppedStopped
|
||||
}
|
||||
|
||||
_, ok := p.pool.TrySubmit(func() {
|
||||
@@ -163,7 +172,7 @@ func (p *UsageRecordWorkerPool) Submit(task UsageRecordTask) UsageRecordSubmitMo
|
||||
if p.pool.Stopped() {
|
||||
p.droppedPoolStopped.Add(1)
|
||||
p.logDrop("stopped")
|
||||
return UsageRecordSubmitModeDropped
|
||||
return UsageRecordSubmitModeDroppedStopped
|
||||
}
|
||||
|
||||
switch p.overflowPolicy {
|
||||
|
||||
@@ -176,7 +176,11 @@ func TestUsageRecordWorkerPool_SubmitAfterStop(t *testing.T) {
|
||||
|
||||
pool.Stop()
|
||||
mode := pool.Submit(func(ctx context.Context) {})
|
||||
require.Equal(t, UsageRecordSubmitModeDropped, mode)
|
||||
require.Equal(t, UsageRecordSubmitModeDroppedStopped, mode)
|
||||
require.True(t, mode.Dropped())
|
||||
require.True(t, UsageRecordSubmitModeDropped.Dropped())
|
||||
require.False(t, UsageRecordSubmitModeEnqueued.Dropped())
|
||||
require.False(t, UsageRecordSubmitModeSync.Dropped())
|
||||
require.GreaterOrEqual(t, pool.Stats().DroppedPoolStopped, uint64(1))
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user