From 5e43c7faac3227baebea5d8a421506c00b0906a6 Mon Sep 17 00:00:00 2001 From: jjaw Date: Wed, 15 Jul 2026 00:00:40 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E9=80=8F=E4=BC=A0=E6=B5=81?= =?UTF-8?q?=E4=BA=8B=E4=BB=B6=E8=BE=B9=E7=95=8C=E5=88=B7=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../service/openai_gateway_passthrough.go | 14 +- .../openai_gateway_passthrough_flush_test.go | 284 ++++++++++++++++++ 2 files changed, 297 insertions(+), 1 deletion(-) create mode 100644 backend/internal/service/openai_gateway_passthrough_flush_test.go diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index 8d8c6c96f..5dc0d180b 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -864,6 +864,15 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( clientOutputStarted := false upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id")) pendingLines := make([]string, 0, 8) + flushPending := false + flushPendingOutput := func() { + if clientDisconnected || !flushPending { + return + } + flusher.Flush() + flushPending = false + } + defer flushPendingOutput() writePendingLines := func() bool { for _, pending := range pendingLines { if _, err := fmt.Fprintln(w, pending); err != nil { @@ -1013,7 +1022,10 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) } else { clientOutputStarted = true - flusher.Flush() + flushPending = true + if line == "" { + flushPendingOutput() + } } } } diff --git a/backend/internal/service/openai_gateway_passthrough_flush_test.go b/backend/internal/service/openai_gateway_passthrough_flush_test.go new file mode 100644 index 000000000..b36a03780 --- /dev/null +++ b/backend/internal/service/openai_gateway_passthrough_flush_test.go @@ -0,0 +1,284 @@ +package service + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type passthroughFlushTestWriter struct { + gin.ResponseWriter + recorder *httptest.ResponseRecorder + failAfterWrites int + successfulWrites int + failedWrites int + flushBodyLengths []int +} + +func (w *passthroughFlushTestWriter) Write(data []byte) (int, error) { + if w.failAfterWrites >= 0 && w.successfulWrites >= w.failAfterWrites { + w.failedWrites++ + return 0, errors.New("client disconnected") + } + n, err := w.ResponseWriter.Write(data) + if err == nil { + w.successfulWrites++ + } + return n, err +} + +func (w *passthroughFlushTestWriter) WriteString(data string) (int, error) { + return w.Write([]byte(data)) +} + +func (w *passthroughFlushTestWriter) Flush() { + w.ResponseWriter.Flush() + w.flushBodyLengths = append(w.flushBodyLengths, w.recorder.Body.Len()) +} + +type passthroughFlushTestErrorBody struct { + payload []byte + err error + sent bool +} + +func (r *passthroughFlushTestErrorBody) Read(p []byte) (int, error) { + if !r.sent { + r.sent = true + return copy(p, r.payload), nil + } + return 0, r.err +} + +func (r *passthroughFlushTestErrorBody) Close() error { return nil } + +func runPassthroughFlushTest( + t *testing.T, + body io.ReadCloser, + failAfterWrites int, + setups ...func(*gin.Context), +) (*openaiStreamingResultPassthrough, *httptest.ResponseRecorder, *passthroughFlushTestWriter, error) { + t.Helper() + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + writer := &passthroughFlushTestWriter{ + ResponseWriter: c.Writer, + recorder: recorder, + failAfterWrites: failAfterWrites, + } + c.Writer = writer + for _, setup := range setups { + setup(c) + } + + svc := &OpenAIGatewayService{cfg: &config.Config{ + Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, + }} + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: body, + } + result, err := svc.handleStreamingResponsePassthrough( + context.Background(), + resp, + c, + &Account{ID: 1, Platform: PlatformOpenAI, Name: "flush-test"}, + time.Now(), + "", + "", + ) + return result, recorder, writer, err +} + +func TestOpenAIStreamingPassthroughFlushesAtCompleteEventBoundaries(t *testing.T) { + firstEvent := "event: response.output_text.delta\n" + + "id: event-1\n" + + `data: {"type":"response.output_text.delta","delta":"hello"}` + "\n\n" + heartbeat := ": keepalive\n\n" + terminalEvent := "event: response.completed\n" + + `data: {"type":"response.completed","response":{"id":"resp_flush","usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}` + "\n\n" + upstream := firstEvent + heartbeat + terminalEvent + + result, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, upstream, recorder.Body.String()) + require.Equal(t, []int{ + len(firstEvent), + len(firstEvent) + len(heartbeat), + len(upstream), + }, writer.flushBodyLengths) + require.Equal(t, 3, result.usage.InputTokens) + require.Equal(t, 2, result.usage.OutputTokens) +} + +func TestOpenAIStreamingPassthroughKeepsPreamblePendingUntilFirstOutputBoundary(t *testing.T) { + preamble := "event: response.created\n" + + `data: {"type":"response.created","response":{"id":"resp_pending"}}` + "\n\n" + + ": waiting\n\n" + firstOutput := `data: {"type":"response.output_text.delta","delta":"ready"}` + "\n\n" + terminalEvent := `data: {"type":"response.completed","response":{"id":"resp_pending","usage":{"input_tokens":4,"output_tokens":1,"total_tokens":5}}}` + "\n\n" + upstream := preamble + firstOutput + terminalEvent + + _, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.NoError(t, err) + require.Equal(t, upstream, recorder.Body.String()) + require.Equal(t, []int{ + len(preamble) + len(firstOutput), + len(upstream), + }, writer.flushBodyLengths) +} + +func TestOpenAIStreamingPassthroughFlushesTerminalEventAtEOFWithoutBlankLine(t *testing.T) { + upstream := "event: response.completed\n" + + `data: {"type":"response.completed","response":{"id":"resp_eof","usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}` + wantBody := upstream + "\n" + + result, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, wantBody, recorder.Body.String()) + require.Equal(t, []int{len(wantBody)}, writer.flushBodyLengths) + require.Equal(t, 5, result.usage.InputTokens) + require.Equal(t, 2, result.usage.OutputTokens) +} + +func TestOpenAIStreamingPassthroughFailedBeforeOutputCanStillFailOverWithoutFlush(t *testing.T) { + upstream := "event: response.created\n" + + `data: {"type":"response.created","response":{"id":"resp_failover"}}` + "\n\n" + + "event: response.failed\n" + + `data: {"type":"response.failed","error":{"code":"server_error","message":"upstream processing failed"}}` + "\n\n" + + _, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Empty(t, recorder.Body.String()) + require.Empty(t, writer.flushBodyLengths) +} + +func TestOpenAIStreamingPassthroughNonRetryableFailedBeforeOutputFlushesAtBoundary(t *testing.T) { + upstream := "event: response.failed\n" + + `data: {"type":"response.failed","error":{"code":"content_policy","message":"request blocked by policy"},"usage":{"input_tokens":6,"output_tokens":0,"total_tokens":6}}` + "\n\n" + + result, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.False(t, errors.As(err, &failoverErr)) + require.NotNil(t, result) + require.Equal(t, upstream, recorder.Body.String()) + require.Equal(t, []int{len(upstream)}, writer.flushBodyLengths) + require.Equal(t, 6, result.usage.InputTokens) + require.Zero(t, result.usage.OutputTokens) +} + +func TestOpenAIStreamingPassthroughFailedAfterOutputFlushesAtBoundaryAndKeepsUsage(t *testing.T) { + firstOutput := `data: {"type":"response.output_text.delta","delta":"partial"}` + "\n\n" + failedEvent := "event: response.failed\n" + + `data: {"type":"response.failed","error":{"code":"server_error","message":"upstream processing failed"},"usage":{"input_tokens":7,"output_tokens":2,"total_tokens":9}}` + "\n\n" + upstream := firstOutput + failedEvent + + result, recorder, writer, err := runPassthroughFlushTest(t, io.NopCloser(strings.NewReader(upstream)), -1) + + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.False(t, errors.As(err, &failoverErr)) + require.NotNil(t, result) + require.Equal(t, upstream, recorder.Body.String()) + require.Equal(t, []int{len(firstOutput), len(upstream)}, writer.flushBodyLengths) + require.Equal(t, 7, result.usage.InputTokens) + require.Equal(t, 2, result.usage.OutputTokens) +} + +func TestOpenAIStreamingPassthroughClientDisconnectStillDrainsTerminalUsage(t *testing.T) { + firstOutput := `data: {"type":"response.output_text.delta","delta":"partial"}` + "\n\n" + terminalEvent := `data: {"type":"response.completed","response":{"id":"resp_drain","usage":{"input_tokens":11,"output_tokens":4,"total_tokens":15}}}` + "\n\n" + + result, recorder, writer, err := runPassthroughFlushTest( + t, + io.NopCloser(strings.NewReader(firstOutput+terminalEvent)), + 2, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, firstOutput, recorder.Body.String()) + require.Equal(t, []int{len(firstOutput)}, writer.flushBodyLengths) + require.Equal(t, 1, writer.failedWrites) + require.Equal(t, 11, result.usage.InputTokens) + require.Equal(t, 4, result.usage.OutputTokens) +} + +func TestOpenAIStreamingPassthroughScannerErrorFlushesWrittenResidual(t *testing.T) { + upstream := []byte(`data: {"type":"response.output_text.delta","delta":"partial"}`) + readErr := errors.New("upstream read failed") + + _, recorder, writer, err := runPassthroughFlushTest(t, &passthroughFlushTestErrorBody{ + payload: upstream, + err: readErr, + }, -1) + + require.ErrorIs(t, err, readErr) + wantBody := string(upstream) + "\n" + require.Equal(t, wantBody, recorder.Body.String()) + require.Equal(t, []int{len(wantBody)}, writer.flushBodyLengths) +} + +func TestOpenAIStreamingPassthroughNamespaceRestoreErrorFlushesWrittenResidualOnce(t *testing.T) { + writtenPrefix := `data: {"type":"response.output_text.delta","delta":"prefix"}` + "\n" + overflowData := `data: {"type":"response.output_text.delta","delta":"not-written","overflow":1e1000}` + + _, recorder, writer, err := runPassthroughFlushTest( + t, + io.NopCloser(strings.NewReader(writtenPrefix+overflowData)), + -1, + func(c *gin.Context) { + setOpenAIResponsesNamespaceNames(c, map[string]apicompat.ResponsesNamespaceName{ + "collaboration__spawn_agent": {Namespace: "collaboration", Name: "spawn_agent"}, + }) + }, + ) + + require.ErrorContains(t, err, "restore OpenAI passthrough namespace response") + require.Equal(t, writtenPrefix, recorder.Body.String()) + require.Equal(t, []int{len(writtenPrefix)}, writer.flushBodyLengths) +} + +func TestOpenAIStreamingPassthroughBlankWriteFailureDoesNotFlushAndStillDrainsUsage(t *testing.T) { + writtenDataLine := `data: {"type":"response.output_text.delta","delta":"partial"}` + "\n" + terminalEvent := `data: {"type":"response.completed","response":{"id":"resp_blank_failure","usage":{"input_tokens":13,"output_tokens":5,"total_tokens":18}}}` + "\n\n" + + result, recorder, writer, err := runPassthroughFlushTest( + t, + io.NopCloser(strings.NewReader(writtenDataLine+"\n"+terminalEvent)), + 1, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, writtenDataLine, recorder.Body.String()) + require.Empty(t, writer.flushBodyLengths) + require.Equal(t, 1, writer.successfulWrites) + require.Equal(t, 1, writer.failedWrites) + require.Equal(t, 13, result.usage.InputTokens) + require.Equal(t, 5, result.usage.OutputTokens) +}