mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3626 from DanisJiang/fix/claude-code-stream-keepalive
fix: prevent Claude Code stream keepalive stalls
This commit is contained in:
@@ -72,10 +72,11 @@ IMPORTANT: You must NEVER generate or guess URLs for the user unless you are con
|
||||
- Do not use a colon before tool calls. Your tool calls may not be shown directly in the output, so text like "Let me read the file:" followed by a read tool call should just be "Let me read the file." with a period.`
|
||||
maxCacheControlBlocks = 4 // Anthropic API 允许的最大 cache_control 块数量
|
||||
|
||||
defaultUserGroupRateCacheTTL = 30 * time.Second
|
||||
defaultModelsListCacheTTL = 15 * time.Second
|
||||
postUsageBillingTimeout = 15 * time.Second
|
||||
debugGatewayBodyEnv = "SUB2API_DEBUG_GATEWAY_BODY"
|
||||
defaultUserGroupRateCacheTTL = 30 * time.Second
|
||||
defaultModelsListCacheTTL = 15 * time.Second
|
||||
postUsageBillingTimeout = 15 * time.Second
|
||||
claudeCodeNoopDeltaKeepaliveMinVersion = "2.1.193"
|
||||
debugGatewayBodyEnv = "SUB2API_DEBUG_GATEWAY_BODY"
|
||||
// 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
|
||||
gatewayUpstreamErrorBodyReadLimit int64 = 512 << 10
|
||||
)
|
||||
@@ -4130,6 +4131,67 @@ func isClaudeCodeClient(userAgent string, metadataUserID string) bool {
|
||||
return ParseMetadataUserID(metadataUserID) != nil
|
||||
}
|
||||
|
||||
func shouldUseClaudeCodeNoopDeltaKeepalive(userAgent string) bool {
|
||||
version := ExtractCLIVersion(userAgent)
|
||||
if version == "" {
|
||||
return false
|
||||
}
|
||||
return CompareVersions(version, claudeCodeNoopDeltaKeepaliveMinVersion) >= 0
|
||||
}
|
||||
|
||||
func claudeCodeKeepaliveDeltaTypeForContentBlock(blockType string) string {
|
||||
switch blockType {
|
||||
case "text":
|
||||
return "text_delta"
|
||||
case "tool_use":
|
||||
return "input_json_delta"
|
||||
case "thinking":
|
||||
return "thinking_delta"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func claudeCodeKeepaliveFieldForDeltaType(deltaType string) string {
|
||||
switch deltaType {
|
||||
case "text_delta":
|
||||
return "text"
|
||||
case "input_json_delta":
|
||||
return "partial_json"
|
||||
case "thinking_delta":
|
||||
return "thinking"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func buildClaudeCodeNoopDeltaKeepalive(index int, deltaType string) (string, bool) {
|
||||
fieldName := claudeCodeKeepaliveFieldForDeltaType(deltaType)
|
||||
if fieldName == "" {
|
||||
return "", false
|
||||
}
|
||||
return fmt.Sprintf("event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":%d,\"delta\":{\"type\":\"%s\",\"%s\":\"\"}}\n\n", index, deltaType, fieldName), true
|
||||
}
|
||||
|
||||
func sseEventIndex(event map[string]any) (int, bool) {
|
||||
switch v := event["index"].(type) {
|
||||
case float64:
|
||||
return int(v), true
|
||||
case int:
|
||||
return v, true
|
||||
case int64:
|
||||
return int(v), true
|
||||
case json.Number:
|
||||
i, err := v.Int64()
|
||||
if err != nil {
|
||||
return 0, false
|
||||
}
|
||||
return int(i), true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
// normalizeSystemParam 将 json.RawMessage 类型的 system 参数转为标准 Go 类型(string / []any / nil),
|
||||
// 避免 type switch 中 json.RawMessage(底层 []byte)无法匹配 case string / case []any / case nil 的问题。
|
||||
// 这是 Go 的 typed nil 陷阱:(json.RawMessage, nil) ≠ (nil, nil)。
|
||||
@@ -5995,16 +6057,28 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
|
||||
if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 {
|
||||
keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second
|
||||
}
|
||||
var keepaliveTicker *time.Ticker
|
||||
var keepaliveTimer *time.Timer
|
||||
if keepaliveInterval > 0 {
|
||||
keepaliveTicker = time.NewTicker(keepaliveInterval)
|
||||
defer keepaliveTicker.Stop()
|
||||
keepaliveTimer = time.NewTimer(keepaliveInterval)
|
||||
defer keepaliveTimer.Stop()
|
||||
}
|
||||
var keepaliveCh <-chan time.Time
|
||||
if keepaliveTicker != nil {
|
||||
keepaliveCh = keepaliveTicker.C
|
||||
if keepaliveTimer != nil {
|
||||
keepaliveCh = keepaliveTimer.C
|
||||
}
|
||||
lastDataAt := time.Now()
|
||||
resetKeepaliveTimer := func() {
|
||||
if keepaliveTimer == nil {
|
||||
return
|
||||
}
|
||||
if !keepaliveTimer.Stop() {
|
||||
select {
|
||||
case <-keepaliveTimer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
keepaliveTimer.Reset(keepaliveInterval)
|
||||
}
|
||||
inPartialEvent := false
|
||||
|
||||
for {
|
||||
@@ -6073,6 +6147,7 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
|
||||
// 按 SSE 事件边界刷出,减少每行 flush 带来的 syscall 开销。
|
||||
flusher.Flush()
|
||||
lastDataAt = time.Now()
|
||||
resetKeepaliveTimer()
|
||||
inPartialEvent = false
|
||||
} else {
|
||||
inPartialEvent = true
|
||||
@@ -6094,10 +6169,15 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
|
||||
return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream data interval timeout")
|
||||
|
||||
case <-keepaliveCh:
|
||||
if clientDisconnected || inPartialEvent {
|
||||
if clientDisconnected {
|
||||
continue
|
||||
}
|
||||
if inPartialEvent {
|
||||
resetKeepaliveTimer()
|
||||
continue
|
||||
}
|
||||
if time.Since(lastDataAt) < keepaliveInterval {
|
||||
resetKeepaliveTimer()
|
||||
continue
|
||||
}
|
||||
if _, err := fmt.Fprint(w, "event: ping\ndata: {\"type\": \"ping\"}\n\n"); err != nil {
|
||||
@@ -6107,6 +6187,7 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
|
||||
}
|
||||
flusher.Flush()
|
||||
lastDataAt = time.Now()
|
||||
resetKeepaliveTimer()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -8207,16 +8288,28 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
|
||||
if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 {
|
||||
keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second
|
||||
}
|
||||
var keepaliveTicker *time.Ticker
|
||||
var keepaliveTimer *time.Timer
|
||||
if keepaliveInterval > 0 {
|
||||
keepaliveTicker = time.NewTicker(keepaliveInterval)
|
||||
defer keepaliveTicker.Stop()
|
||||
keepaliveTimer = time.NewTimer(keepaliveInterval)
|
||||
defer keepaliveTimer.Stop()
|
||||
}
|
||||
var keepaliveCh <-chan time.Time
|
||||
if keepaliveTicker != nil {
|
||||
keepaliveCh = keepaliveTicker.C
|
||||
if keepaliveTimer != nil {
|
||||
keepaliveCh = keepaliveTimer.C
|
||||
}
|
||||
lastDataAt := time.Now()
|
||||
resetKeepaliveTimer := func() {
|
||||
if keepaliveTimer == nil {
|
||||
return
|
||||
}
|
||||
if !keepaliveTimer.Stop() {
|
||||
select {
|
||||
case <-keepaliveTimer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
keepaliveTimer.Reset(keepaliveInterval)
|
||||
}
|
||||
|
||||
// 仅发送一次错误事件,避免多次写入导致协议混乱(写失败时尽力通知客户端)。
|
||||
// 事件格式遵循 Anthropic SSE 标准:{"type":"error","error":{"type":<reason>,"message":<message>}}
|
||||
@@ -8249,6 +8342,9 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
|
||||
needModelReplace := originalModel != mappedModel
|
||||
clientDisconnected := false // 客户端断开标志,断开后继续读取上游以获取完整usage
|
||||
sawTerminalEvent := false
|
||||
useNoopDeltaKeepalive := c != nil && c.Request != nil && shouldUseClaudeCodeNoopDeltaKeepalive(c.GetHeader("User-Agent"))
|
||||
noopDeltaKeepaliveBlockIndex := -1
|
||||
noopDeltaKeepaliveDeltaType := ""
|
||||
|
||||
pendingEventLines := make([]string, 0, 4)
|
||||
|
||||
@@ -8305,6 +8401,41 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
|
||||
}
|
||||
eventChanged := false
|
||||
|
||||
if useNoopDeltaKeepalive {
|
||||
switch eventType {
|
||||
case "content_block_start":
|
||||
if idx, ok := sseEventIndex(event); ok {
|
||||
noopDeltaKeepaliveBlockIndex = -1
|
||||
noopDeltaKeepaliveDeltaType = ""
|
||||
if contentBlock, ok := event["content_block"].(map[string]any); ok {
|
||||
blockType, _ := contentBlock["type"].(string)
|
||||
if deltaType := claudeCodeKeepaliveDeltaTypeForContentBlock(blockType); deltaType != "" {
|
||||
noopDeltaKeepaliveBlockIndex = idx
|
||||
noopDeltaKeepaliveDeltaType = deltaType
|
||||
}
|
||||
}
|
||||
}
|
||||
case "content_block_delta":
|
||||
if idx, ok := sseEventIndex(event); ok {
|
||||
if delta, ok := event["delta"].(map[string]any); ok {
|
||||
deltaType, _ := delta["type"].(string)
|
||||
if claudeCodeKeepaliveFieldForDeltaType(deltaType) != "" {
|
||||
noopDeltaKeepaliveBlockIndex = idx
|
||||
noopDeltaKeepaliveDeltaType = deltaType
|
||||
}
|
||||
}
|
||||
}
|
||||
case "content_block_stop":
|
||||
if idx, ok := sseEventIndex(event); ok && idx == noopDeltaKeepaliveBlockIndex {
|
||||
noopDeltaKeepaliveBlockIndex = -1
|
||||
noopDeltaKeepaliveDeltaType = ""
|
||||
}
|
||||
case "message_stop":
|
||||
noopDeltaKeepaliveBlockIndex = -1
|
||||
noopDeltaKeepaliveDeltaType = ""
|
||||
}
|
||||
}
|
||||
|
||||
// 兼容 Kimi cached_tokens → cache_read_input_tokens
|
||||
if eventType == "message_start" {
|
||||
if msg, ok := event["message"].(map[string]any); ok {
|
||||
@@ -8457,6 +8588,7 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
|
||||
}
|
||||
flusher.Flush()
|
||||
lastDataAt = time.Now()
|
||||
resetKeepaliveTimer()
|
||||
}
|
||||
if data != "" {
|
||||
if firstTokenMs == nil && data != "[DONE]" {
|
||||
@@ -8494,16 +8626,23 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
|
||||
continue
|
||||
}
|
||||
if time.Since(lastDataAt) < keepaliveInterval {
|
||||
resetKeepaliveTimer()
|
||||
continue
|
||||
}
|
||||
// SSE ping 事件:Anthropic 原生格式,客户端会正确处理,
|
||||
// 同时保持连接活跃防止 Cloudflare Tunnel 等代理断开
|
||||
if _, werr := fmt.Fprint(w, "event: ping\ndata: {\"type\": \"ping\"}\n\n"); werr != nil {
|
||||
keepaliveBlock := "event: ping\ndata: {\"type\": \"ping\"}\n\n"
|
||||
if useNoopDeltaKeepalive && noopDeltaKeepaliveBlockIndex >= 0 {
|
||||
if block, ok := buildClaudeCodeNoopDeltaKeepalive(noopDeltaKeepaliveBlockIndex, noopDeltaKeepaliveDeltaType); ok {
|
||||
keepaliveBlock = block
|
||||
}
|
||||
}
|
||||
if _, werr := fmt.Fprint(w, keepaliveBlock); werr != nil {
|
||||
clientDisconnected = true
|
||||
logger.LegacyPrintf("service.gateway", "Client disconnected during keepalive ping, continuing to drain upstream for billing")
|
||||
continue
|
||||
}
|
||||
flusher.Flush()
|
||||
lastDataAt = time.Now()
|
||||
resetKeepaliveTimer()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -15,19 +15,21 @@ import (
|
||||
|
||||
type upstreamContextTestKey string
|
||||
|
||||
func TestGatewayService_StreamingReusesScannerBufferAndStillParsesUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
cfg := &config.Config{
|
||||
Gateway: config.GatewayConfig{
|
||||
StreamDataIntervalTimeout: 0,
|
||||
MaxLineSize: defaultMaxLineSize,
|
||||
func newStreamingResponseTestGatewayService() *GatewayService {
|
||||
return &GatewayService{
|
||||
cfg: &config.Config{
|
||||
Gateway: config.GatewayConfig{
|
||||
StreamDataIntervalTimeout: 0,
|
||||
MaxLineSize: defaultMaxLineSize,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
svc := &GatewayService{
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{},
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayService_StreamingReusesScannerBufferAndStillParsesUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newStreamingResponseTestGatewayService()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
@@ -53,6 +55,126 @@ func TestGatewayService_StreamingReusesScannerBufferAndStillParsesUsage(t *testi
|
||||
require.Equal(t, 7, result.usage.OutputTokens)
|
||||
}
|
||||
|
||||
func TestGatewayService_StreamingKeepaliveUsesIdleTimer(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newStreamingResponseTestGatewayService()
|
||||
svc.cfg.Gateway.StreamKeepaliveInterval = 1
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: pr}
|
||||
|
||||
go func() {
|
||||
defer func() { _ = pw.Close() }()
|
||||
_, _ = pw.Write([]byte("data: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1}}}\n\n"))
|
||||
time.Sleep(1100 * time.Millisecond)
|
||||
_, _ = pw.Write([]byte("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"))
|
||||
}()
|
||||
|
||||
result, err := svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "model", "model", false)
|
||||
_ = pr.Close()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Contains(t, rec.Body.String(), "event: ping")
|
||||
}
|
||||
|
||||
func TestGatewayService_StreamingKeepaliveUsesNoopDeltaForAffectedClaudeCodeVersion(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newStreamingResponseTestGatewayService()
|
||||
svc.cfg.Gateway.StreamKeepaliveInterval = 1
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
c.Request.Header.Set("User-Agent", "claude-cli/2.1.198 (external, cli)")
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: pr}
|
||||
|
||||
go func() {
|
||||
defer func() { _ = pw.Close() }()
|
||||
_, _ = pw.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1}}}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"))
|
||||
time.Sleep(1100 * time.Millisecond)
|
||||
_, _ = pw.Write([]byte("event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"))
|
||||
}()
|
||||
|
||||
result, err := svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "model", "model", false)
|
||||
_ = pr.Close()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "event: content_block_delta")
|
||||
require.Contains(t, body, `"delta":{"type":"text_delta","text":""}`)
|
||||
}
|
||||
|
||||
func TestGatewayService_StreamingKeepaliveUsesNoopDeltaDuringToolUseForAffectedClaudeCodeVersion(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newStreamingResponseTestGatewayService()
|
||||
svc.cfg.Gateway.StreamKeepaliveInterval = 1
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
c.Request.Header.Set("User-Agent", "claude-cli/2.1.198 (external, cli)")
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: pr}
|
||||
|
||||
go func() {
|
||||
defer func() { _ = pw.Close() }()
|
||||
_, _ = pw.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1}}}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":1,\"content_block\":{\"type\":\"tool_use\",\"id\":\"toolu_1\",\"name\":\"Edit\",\"input\":{}}}\n\n"))
|
||||
time.Sleep(1100 * time.Millisecond)
|
||||
_, _ = pw.Write([]byte("event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":1}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"))
|
||||
}()
|
||||
|
||||
result, err := svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "model", "model", false)
|
||||
_ = pr.Close()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "event: content_block_delta")
|
||||
require.Contains(t, body, `"index":1`)
|
||||
require.Contains(t, body, `"delta":{"type":"input_json_delta","partial_json":""}`)
|
||||
}
|
||||
|
||||
func TestGatewayService_StreamingKeepaliveKeepsPingForOlderClaudeCodeVersion(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newStreamingResponseTestGatewayService()
|
||||
svc.cfg.Gateway.StreamKeepaliveInterval = 1
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
c.Request.Header.Set("User-Agent", "claude-cli/2.1.187 (external, cli)")
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: pr}
|
||||
|
||||
go func() {
|
||||
defer func() { _ = pw.Close() }()
|
||||
_, _ = pw.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":1}}}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"))
|
||||
time.Sleep(1100 * time.Millisecond)
|
||||
_, _ = pw.Write([]byte("event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n"))
|
||||
_, _ = pw.Write([]byte("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"))
|
||||
}()
|
||||
|
||||
result, err := svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "model", "model", false)
|
||||
_ = pr.Close()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "event: ping")
|
||||
require.NotContains(t, body, `"delta":{"type":"text_delta","text":""}`)
|
||||
}
|
||||
|
||||
func TestDetachUpstreamContextIgnoresClientCancel(t *testing.T) {
|
||||
parent, cancel := context.WithCancel(context.WithValue(context.Background(), upstreamContextTestKey("test-key"), "test-value"))
|
||||
upstreamCtx, release := detachUpstreamContext(parent)
|
||||
|
||||
Reference in New Issue
Block a user