From 5fa871a505e56d0e2c1217fd029f8ed95b8c8214 Mon Sep 17 00:00:00 2001 From: Seefs Date: Sat, 11 Jul 2026 11:42:13 +0800 Subject: [PATCH] feat: forward all Codex SSE response headers for codex sub --- relay/helper/stream_scanner.go | 21 +++++++++------ relay/helper/stream_scanner_test.go | 40 +++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 8 deletions(-) diff --git a/relay/helper/stream_scanner.go b/relay/helper/stream_scanner.go index f88bef5206fc..a96939606bad 100644 --- a/relay/helper/stream_scanner.go +++ b/relay/helper/stream_scanner.go @@ -18,14 +18,17 @@ import ( "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/bytedance/gopkg/util/gopool" + "github.com/samber/lo" "github.com/gin-gonic/gin" ) const ( - InitialScannerBufferSize = 64 << 10 // 64KB (64*1024) - DefaultMaxScannerBufferSize = 128 << 20 // 64MB (64*1024*1024) default SSE buffer size - DefaultPingInterval = 10 * time.Second + InitialScannerBufferSize = 64 << 10 // 64KB (64*1024) + DefaultMaxScannerBufferSize = 128 << 20 // 64MB (64*1024*1024) default SSE buffer size + DefaultPingInterval = 10 * time.Second + codexReasoningIncludedHeader = "X-Reasoning-Included" + codexTurnStateHeader = "X-Codex-Turn-State" // streamWriteTimeout bounds a single blocked write to a slow client so the // unconditional wg.Wait() in cleanup can always finish. Without it, a slow // but connected client (full TCP buffer, no server WriteTimeout) could hang @@ -46,13 +49,15 @@ func NewStreamScanner(reader io.Reader) *bufio.Scanner { return scanner } -func copyCodexSSEHeaders(c *gin.Context, resp *http.Response) { +func copyCodexSSEHeaders(c *gin.Context, resp *http.Response, copyAll bool) { if c == nil || c.Writer == nil || resp == nil { return } - // codex - for _, name := range []string{"X-Reasoning-Included", "X-Codex-Turn-State"} { - values := resp.Header.Values(name) + headers := resp.Header + if !copyAll { + headers = lo.PickByKeys(resp.Header, []string{codexReasoningIncludedHeader, codexTurnStateHeader}) + } + for name, values := range headers { if !service.ShouldCopyUpstreamHeader(c, name, values) { continue } @@ -141,7 +146,7 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon defer cleanup() scanner.Split(bufio.ScanLines) - copyCodexSSEHeaders(c, resp) + copyCodexSSEHeaders(c, resp, info.ChannelMeta.ChannelType == constant.ChannelTypeCodex) SetEventStreamHeaders(c) ctx = context.WithValue(ctx, "stop_chan", stopChan) diff --git a/relay/helper/stream_scanner_test.go b/relay/helper/stream_scanner_test.go index a211951287d2..4c88b213da53 100644 --- a/relay/helper/stream_scanner_test.go +++ b/relay/helper/stream_scanner_test.go @@ -99,6 +99,46 @@ func TestStreamScannerHandler_EmptyBody(t *testing.T) { assert.False(t, called.Load(), "handler should not be called for empty body") } +func TestStreamScannerHandler_CodexCopiesEligibleUpstreamHeaders(t *testing.T) { + t.Parallel() + + c, resp, info := setupStreamTest(t, strings.NewReader("")) + info.ChannelMeta.ChannelType = constant.ChannelTypeCodex + resp.Header = http.Header{ + "X-Reasoning-Included": {"true"}, + "X-Codex-Turn-State": {"turn-state"}, + "X-Upstream-Trace": {"trace-id"}, + "Set-Cookie": {"first=1", "second=2"}, + "Content-Length": {"42"}, + } + + StreamScannerHandler(c, resp, info, func(data string, sr *StreamResult) {}) + + assert.Equal(t, "true", c.Writer.Header().Get("X-Reasoning-Included")) + assert.Equal(t, "turn-state", c.Writer.Header().Get("X-Codex-Turn-State")) + assert.Equal(t, "trace-id", c.Writer.Header().Get("X-Upstream-Trace")) + assert.Equal(t, []string{"first=1", "second=2"}, c.Writer.Header().Values("Set-Cookie")) + assert.Empty(t, c.Writer.Header().Get("Content-Length")) +} + +func TestStreamScannerHandler_NonCodexPreservesLegacyHeaderBehavior(t *testing.T) { + t.Parallel() + + c, resp, info := setupStreamTest(t, strings.NewReader("")) + info.ChannelMeta.ChannelType = constant.ChannelTypeOpenAI + resp.Header = http.Header{ + "X-Reasoning-Included": {"true"}, + "X-Codex-Turn-State": {"turn-state"}, + "X-Upstream-Trace": {"trace-id"}, + } + + StreamScannerHandler(c, resp, info, func(data string, sr *StreamResult) {}) + + assert.Equal(t, "true", c.Writer.Header().Get("X-Reasoning-Included")) + assert.Equal(t, "turn-state", c.Writer.Header().Get("X-Codex-Turn-State")) + assert.Empty(t, c.Writer.Header().Get("X-Upstream-Trace")) +} + func TestStreamScannerHandler_1000Chunks(t *testing.T) { t.Parallel()