mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 14:08:14 +08:00
Merge pull request #5154 from Wei-Shaw/fix/issue-5148-stream-partial-usage-billing
fix(gateway): 流中断时保留已观测 usage 入账,修复 newapi 类上游大面积漏记(#5148)
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