diff --git a/backend/internal/infra/provider/conversation/stream.go b/backend/internal/infra/provider/conversation/stream.go index ad8899089..ea67644ee 100644 --- a/backend/internal/infra/provider/conversation/stream.go +++ b/backend/internal/infra/provider/conversation/stream.go @@ -7,12 +7,24 @@ import ( "fmt" "io" "strings" + "sync" "time" ) const ( maxDeferredSearchTextBytes = 8 << 20 maxDeferredReasoningSummaryBytes = 8 << 20 + + // contentDoomLoopThreshold 连续重复同一可见内容增量时终止流。真正的 + // 内容循环会消耗配额和客户端上下文,因此远低于推理上限;但仍需容纳 + // 合法重复:markdown 分隔线与表格边框会以相同单字符增量("-"、"="、 + // "|")连续输出。 + contentDoomLoopThreshold = 128 + + // reasoningDoomLoopThreshold 高于内容阈值:high/xhigh 推理会大量重复 + // 同一短标记("so"、"hmm"、"wait"、列表符号)。共用低阈值会过早终止 + // 有效的深度推理响应。 + reasoningDoomLoopThreshold = 256 ) // ConvertResponseStream 将 Responses SSE 转换为 Chat Completions 或 Anthropic Messages SSE。 @@ -23,11 +35,12 @@ func ConvertResponseStream(source io.ReadCloser, operation string) io.ReadCloser // ConvertResponseStreamWithOptions 按下游协议选项生成 Chat 或 Anthropic SSE。 func ConvertResponseStreamWithOptions(source io.ReadCloser, operation string, options ResponseOptions) io.ReadCloser { if operation == OperationResponses { - return source + return guardResponseStream(source) } reader, writer := io.Pipe() + stream := newStreamPipeReadCloser(reader, source) go func() { - defer source.Close() + defer stream.closeSource() converter := newStreamConverter(writer, operation, options) err := consumeSSE(source, converter.handle) if err == nil { @@ -35,7 +48,7 @@ func ConvertResponseStreamWithOptions(source io.ReadCloser, operation string, op } _ = writer.CloseWithError(err) }() - return reader + return stream } type streamConverter struct { @@ -68,6 +81,45 @@ type streamConverter struct { stopFilter *anthropicStreamStopFilter stopSequence string refused bool + repeatTracker streamRepeatTracker +} + +// streamRepeatTracker 在协议转换、缓冲和 stop filter 之前跟踪上游增量, +// 避免任一下游路径绕过循环保护。 +type streamRepeatTracker struct { + lastContentDelta string + contentRepeatCount int + lastReasonDelta string + reasonRepeatCount int +} + +// streamPipeReadCloser ensures a downstream cancellation immediately closes the +// upstream body, including while the forwarding goroutine is blocked in Read. +type streamPipeReadCloser struct { + *io.PipeReader + source io.ReadCloser + closeOnce sync.Once + closeErr error +} + +func newStreamPipeReadCloser(reader *io.PipeReader, source io.ReadCloser) *streamPipeReadCloser { + return &streamPipeReadCloser{PipeReader: reader, source: source} +} + +func (r *streamPipeReadCloser) Close() error { + readerErr := r.PipeReader.Close() + sourceErr := r.closeSource() + if readerErr != nil { + return readerErr + } + return sourceErr +} + +func (r *streamPipeReadCloser) closeSource() error { + r.closeOnce.Do(func() { + r.closeErr = r.source.Close() + }) + return r.closeErr } type streamTool struct { @@ -229,16 +281,12 @@ func (c *streamConverter) handle(event string, data []byte) error { if c.finished { return nil } - if bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { - return nil - } - var root map[string]json.RawMessage - if json.Unmarshal(data, &root) != nil { + typeName, root, ok := parseSSEEvent(event, data) + if !ok { return nil } - typeName := event - if raw := root["type"]; typeName == "" { - _ = json.Unmarshal(raw, &typeName) + if err := c.repeatTracker.trackEvent(typeName, root); err != nil { + return err } if c.stopSequence != "" && typeName != "response.completed" && typeName != "response.incomplete" && typeName != "response.failed" && typeName != "error" { return nil @@ -711,3 +759,87 @@ func consumeSSE(source io.Reader, handle func(string, []byte) error) error { } } } + +func parseSSEEvent(event string, data []byte) (string, map[string]json.RawMessage, bool) { + if bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { + return "", nil, false + } + var root map[string]json.RawMessage + if json.Unmarshal(data, &root) != nil { + return "", nil, false + } + typeName := event + if typeName == "" { + _ = json.Unmarshal(root["type"], &typeName) + } + return typeName, root, true +} + +func (t *streamRepeatTracker) trackEvent(typeName string, root map[string]json.RawMessage) error { + var delta string + switch typeName { + case "response.output_text.delta": + _ = json.Unmarshal(root["delta"], &delta) + return t.trackContent(delta) + case "response.reasoning_summary_text.delta": + _ = json.Unmarshal(root["delta"], &delta) + return t.trackReasoning(delta, "model reasoning summary loop detected") + case "response.reasoning_text.delta": + _ = json.Unmarshal(root["delta"], &delta) + return t.trackReasoning(delta, "model reasoning loop detected") + default: + return nil + } +} + +func (t *streamRepeatTracker) trackContent(delta string) error { + if delta == "" { + return nil + } + if delta != t.lastContentDelta { + t.lastContentDelta = delta + t.contentRepeatCount = 1 + return nil + } + t.contentRepeatCount++ + if t.contentRepeatCount > contentDoomLoopThreshold { + return fmt.Errorf("model output loop detected (repeated content delta %d times)", t.contentRepeatCount) + } + return nil +} + +func (t *streamRepeatTracker) trackReasoning(delta, message string) error { + if delta == "" { + return nil + } + if delta != t.lastReasonDelta { + t.lastReasonDelta = delta + t.reasonRepeatCount = 1 + return nil + } + t.reasonRepeatCount++ + if t.reasonRepeatCount > reasoningDoomLoopThreshold { + return fmt.Errorf("%s (repeated delta %d times)", message, t.reasonRepeatCount) + } + return nil +} + +// guardResponseStream 保持 native Responses SSE 的原始字节不变,同时在读取时 +// 解析事件并在检测到循环时关闭上游。 +func guardResponseStream(source io.ReadCloser) io.ReadCloser { + reader, writer := io.Pipe() + stream := newStreamPipeReadCloser(reader, source) + go func() { + defer stream.closeSource() + tracker := streamRepeatTracker{} + err := consumeSSE(io.TeeReader(source, writer), func(event string, data []byte) error { + typeName, root, ok := parseSSEEvent(event, data) + if !ok { + return nil + } + return tracker.trackEvent(typeName, root) + }) + _ = writer.CloseWithError(err) + }() + return stream +} diff --git a/backend/internal/infra/provider/conversation/stream_doomloop_test.go b/backend/internal/infra/provider/conversation/stream_doomloop_test.go new file mode 100644 index 000000000..533ce2d49 --- /dev/null +++ b/backend/internal/infra/provider/conversation/stream_doomloop_test.go @@ -0,0 +1,321 @@ +package conversation + +import ( + "fmt" + "io" + "strings" + "testing" +) + +type blockingStreamSource struct { + closed chan struct{} +} + +func (s *blockingStreamSource) Read([]byte) (int, error) { + <-s.closed + return 0, io.EOF +} + +func (s *blockingStreamSource) Close() error { + close(s.closed) + return nil +} + +// repeatSSE builds an SSE stream that emits the same delta count times using +// the given event name and payload template. +func repeatSSE(event, payloadTemplate string, count int, trailer ...string) string { + lines := []string{ + `event: response.created`, + `data: {"type":"response.created","response":{"id":"resp_1","model":"grok-4.6","status":"in_progress"}}`, "", + } + for i := 0; i < count; i++ { + lines = append(lines, "event: "+event, "data: "+payloadTemplate, "") + } + lines = append(lines, trailer...) + lines = append(lines, + `event: response.completed`, + `data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}`, "", "") + return strings.Join(lines, "\n") +} + +// A visible-content loop is a real quota burn: it must be terminated near the +// content threshold and must not be allowed to run to the reasoning threshold. +func TestConvertResponsesStreamTerminatesContentDoomLoop(t *testing.T) { + stream := repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"loop"}`, + contentDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err == nil { + t.Fatal("repeated visible content must terminate the stream") + } + if !strings.Contains(err.Error(), "model output loop detected") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestConvertResponsesStreamAllowsExactlyContentThreshold(t *testing.T) { + stream := repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"loop"}`, + contentDoomLoopThreshold) + converted, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err != nil { + t.Fatalf("the content ceiling itself must remain valid: %v", err) + } + if !strings.Contains(string(converted), "data: [DONE]") { + t.Fatalf("stream did not complete: %s", converted) + } +} + +func TestConvertResponsesStreamProtectsDeferredWebSearchText(t *testing.T) { + stream := repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"loop"}`, + contentDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStreamWithOptions( + io.NopCloser(strings.NewReader(stream)), + OperationMessages, + ResponseOptions{AnthropicWebSearch: true}, + )) + if err == nil || !strings.Contains(err.Error(), "model output loop detected") { + t.Fatalf("deferred web-search text must retain loop protection: %v", err) + } +} + +func TestConvertResponsesStreamProtectsAfterStopSequence(t *testing.T) { + stream := repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"STOP"}`, + contentDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStreamWithOptions( + io.NopCloser(strings.NewReader(stream)), + OperationChat, + ResponseOptions{StopSequences: []string{"STOP"}}, + )) + if err == nil || !strings.Contains(err.Error(), "model output loop detected") { + t.Fatalf("discarded post-stop deltas must retain loop protection: %v", err) + } +} + +// Regression: high/xhigh effort reasoning legitimately repeats short tokens +// ("so", "hmm", "wait", bullet markers) far more often than visible output. +// A shared counter at the content threshold truncated valid deep-thinking +// answers, so reasoning must survive well past contentDoomLoopThreshold. +func TestConvertResponsesStreamKeepsRepeatedReasoningBelowThreshold(t *testing.T) { + for _, event := range []string{"response.reasoning_text.delta", "response.reasoning_summary_text.delta"} { + t.Run(event, func(t *testing.T) { + // Sit between the two ceilings: this run must be fatal for visible + // content but survivable for reasoning. + repeats := (contentDoomLoopThreshold + reasoningDoomLoopThreshold) / 2 + if repeats <= contentDoomLoopThreshold || repeats >= reasoningDoomLoopThreshold { + t.Fatalf("test invariant broken: %d must sit between %d and %d", + repeats, contentDoomLoopThreshold, reasoningDoomLoopThreshold) + } + stream := repeatSSE(event, + fmt.Sprintf(`{"type":%q,"item_id":"rs_1","delta":"hmm"}`, event), + repeats, + `event: response.output_text.delta`, + `data: {"type":"response.output_text.delta","delta":"answer"}`, "") + converted, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err != nil { + t.Fatalf("deep reasoning must not be treated as a loop: %v", err) + } + if !strings.Contains(string(converted), `"content":"answer"`) { + t.Fatalf("visible answer was lost: %s", converted) + } + }) + } +} + +// The elevated reasoning threshold is a higher ceiling, not an exemption: +// a runaway reasoning loop still has to be terminated. +func TestConvertResponsesStreamTerminatesReasoningDoomLoop(t *testing.T) { + for _, testCase := range []struct { + event string + want string + }{ + {"response.reasoning_text.delta", "model reasoning loop detected"}, + {"response.reasoning_summary_text.delta", "model reasoning summary loop detected"}, + } { + t.Run(testCase.event, func(t *testing.T) { + stream := repeatSSE(testCase.event, + fmt.Sprintf(`{"type":%q,"item_id":"rs_1","delta":"hmm"}`, testCase.event), + reasoningDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err == nil { + t.Fatal("a runaway reasoning loop must still terminate the stream") + } + if !strings.Contains(err.Error(), testCase.want) { + t.Fatalf("unexpected error: %v", err) + } + }) + } +} + +func TestConvertResponsesStreamTracksSuppressedReasoning(t *testing.T) { + stream := repeatSSE("response.reasoning_text.delta", + `{"type":"response.reasoning_text.delta","item_id":"rs_1","delta":"hmm"}`, + reasoningDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationMessages)) + if err == nil || !strings.Contains(err.Error(), "model reasoning loop detected") { + t.Fatalf("reasoning hidden from the downstream protocol must still be guarded: %v", err) + } +} + +func TestConvertResponsesStreamDoesNotCountFlushedSummaryTwice(t *testing.T) { + repeats := (contentDoomLoopThreshold + reasoningDoomLoopThreshold) / 2 + lines := []string{ + `event: response.created`, + `data: {"type":"response.created","response":{"id":"resp_1","model":"grok-4.6","status":"in_progress"}}`, "", + } + for i := 0; i < repeats; i++ { + itemID := fmt.Sprintf("rs_%d", i) + lines = append(lines, + `event: response.reasoning_summary_text.delta`, + fmt.Sprintf(`data: {"type":"response.reasoning_summary_text.delta","item_id":%q,"delta":"hmm"}`, itemID), "", + `event: response.output_item.done`, + fmt.Sprintf(`data: {"type":"response.output_item.done","item":{"id":%q,"type":"reasoning"}}`, itemID), "", + ) + } + lines = append(lines, + `event: response.completed`, + `data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}`, "", "") + converted, err := io.ReadAll(ConvertResponseStream( + io.NopCloser(strings.NewReader(strings.Join(lines, "\n"))), OperationChat, + )) + if err != nil { + t.Fatalf("flushing buffered summaries must not increment the upstream repeat counter: %v", err) + } + if !strings.Contains(string(converted), "data: [DONE]") { + t.Fatalf("stream did not complete: %s", converted) + } +} + +func TestConvertResponsesStreamSharesReasoningCounterAcrossEventTypes(t *testing.T) { + lines := []string{ + `event: response.created`, + `data: {"type":"response.created","response":{"id":"resp_1","model":"grok-4.6","status":"in_progress"}}`, "", + } + for i := 0; i < reasoningDoomLoopThreshold/2; i++ { + lines = append(lines, + `event: response.reasoning_summary_text.delta`, + `data: {"type":"response.reasoning_summary_text.delta","item_id":"rs_1","delta":"hmm"}`, "") + } + for i := 0; i <= reasoningDoomLoopThreshold/2; i++ { + lines = append(lines, + `event: response.reasoning_text.delta`, + `data: {"type":"response.reasoning_text.delta","item_id":"rs_1","delta":"hmm"}`, "") + } + _, err := io.ReadAll(ConvertResponseStream( + io.NopCloser(strings.NewReader(strings.Join(lines, "\n"))), OperationChat, + )) + if err == nil || !strings.Contains(err.Error(), "model reasoning loop detected") { + t.Fatalf("summary and raw reasoning must share one upstream counter: %v", err) + } +} + +func TestConvertResponseStreamGuardsNativeResponsesWithoutRewriting(t *testing.T) { + t.Run("passthrough", func(t *testing.T) { + source := ": keep-this-comment\r\n\r\n" + repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"answer"}`, 1) + converted, err := io.ReadAll(ConvertResponseStream( + io.NopCloser(strings.NewReader(source)), OperationResponses, + )) + if err != nil { + t.Fatalf("native response passthrough failed: %v", err) + } + if string(converted) != source { + t.Fatalf("native response bytes changed:\nwant %q\n got %q", source, converted) + } + }) + + t.Run("doom loop", func(t *testing.T) { + stream := repeatSSE("response.output_text.delta", + `{"type":"response.output_text.delta","delta":"loop"}`, + contentDoomLoopThreshold+8) + _, err := io.ReadAll(ConvertResponseStream( + io.NopCloser(strings.NewReader(stream)), OperationResponses, + )) + if err == nil || !strings.Contains(err.Error(), "model output loop detected") { + t.Fatalf("native responses must retain loop protection: %v", err) + } + }) +} + +func TestConvertResponseStreamCloseImmediatelyClosesUpstream(t *testing.T) { + for _, operation := range []string{OperationChat, OperationResponses} { + t.Run(operation, func(t *testing.T) { + source := &blockingStreamSource{closed: make(chan struct{})} + stream := ConvertResponseStream(source, operation) + if err := stream.Close(); err != nil { + t.Fatalf("close converted stream: %v", err) + } + select { + case <-source.closed: + default: + t.Fatal("closing the downstream stream did not close the upstream source") + } + }) + } +} + +// Legitimate visible repetition must survive: markdown horizontal rules and +// ASCII table borders stream as long runs of an identical single-character +// delta. This is why the content threshold cannot sit near typical rule width. +func TestConvertResponsesStreamKeepsMarkdownRuleAndTableBorders(t *testing.T) { + for _, testCase := range []struct{ name, delta string }{ + {"horizontal rule", "-"}, + {"table border", "="}, + {"empty table cells", " | "}, + } { + t.Run(testCase.name, func(t *testing.T) { + // A wide rule or table border comfortably exceeds 32 characters. + stream := repeatSSE("response.output_text.delta", + fmt.Sprintf(`{"type":"response.output_text.delta","delta":%q}`, testCase.delta), + 80) + converted, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err != nil { + t.Fatalf("legitimate repeated formatting must not be treated as a loop: %v", err) + } + if !strings.Contains(string(converted), "data: [DONE]") { + t.Fatalf("stream did not complete: %s", converted) + } + }) + } +} + +// Counters are per-channel and reset on change, so alternating deltas and +// interleaved reasoning must never accumulate into a false positive. +func TestConvertResponsesStreamDoomLoopCountersResetAndStaySeparate(t *testing.T) { + lines := []string{ + `event: response.created`, + `data: {"type":"response.created","response":{"id":"resp_1","model":"grok-4.6","status":"in_progress"}}`, "", + } + // Alternating visible content never trips the content counter. + for i := 0; i < contentDoomLoopThreshold*2; i++ { + delta := "a" + if i%2 == 1 { + delta = "b" + } + lines = append(lines, + `event: response.output_text.delta`, + fmt.Sprintf(`data: {"type":"response.output_text.delta","delta":%q}`, delta), "") + } + // Reasoning repeats interleaved with distinct content must not share a + // counter: the reasoning run alone exceeds the content threshold. + for i := 0; i < contentDoomLoopThreshold*2; i++ { + lines = append(lines, + `event: response.reasoning_text.delta`, + `data: {"type":"response.reasoning_text.delta","item_id":"rs_1","delta":"hmm"}`, "", + `event: response.output_text.delta`, + fmt.Sprintf(`data: {"type":"response.output_text.delta","delta":"tick%d"}`, i), "") + } + lines = append(lines, + `event: response.completed`, + `data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}`, "", "") + stream := strings.Join(lines, "\n") + converted, err := io.ReadAll(ConvertResponseStream(io.NopCloser(strings.NewReader(stream)), OperationChat)) + if err != nil { + t.Fatalf("alternating deltas must not be treated as a loop: %v", err) + } + if !strings.Contains(string(converted), "data: [DONE]") { + t.Fatalf("stream did not complete: %s", converted) + } +}