diff --git a/internal/proxy/codex_rate_limits.go b/internal/proxy/codex_rate_limits.go index 205ac72..bcdad85 100644 --- a/internal/proxy/codex_rate_limits.go +++ b/internal/proxy/codex_rate_limits.go @@ -1,7 +1,6 @@ package proxy import ( - "bufio" "bytes" "context" "encoding/json" @@ -566,7 +565,7 @@ func parseRateLimitHeaderKey(lower string) (limitID, window, field string, ok bo } func parseCodexRateLimitsFromEvent(eventData []byte) *storage.CodexRateLimitsData { - scanner := bufio.NewScanner(bytes.NewReader(eventData)) + scanner := newSSEScanner(bytes.NewReader(eventData)) for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data:") { diff --git a/internal/proxy/streaming.go b/internal/proxy/streaming.go index ce9acf5..f2d0c43 100644 --- a/internal/proxy/streaming.go +++ b/internal/proxy/streaming.go @@ -23,6 +23,17 @@ var ( errEndpointSwitched = errors.New("endpoint switched") ) +const ( + initialSSEBufferSize = 128 * 1024 + maxSSETokenSize = 16 * 1024 * 1024 +) + +func newSSEScanner(reader io.Reader) *bufio.Scanner { + scanner := bufio.NewScanner(reader) + scanner.Buffer(make([]byte, 0, initialSSEBufferSize), maxSSETokenSize) + return scanner +} + // handleStreamingResponse processes streaming SSE responses func (p *Proxy) handleStreamingResponse(w http.ResponseWriter, resp *http.Response, endpoint config.Endpoint, trans transformer.Transformer, transformerName string, thinkingEnabled bool, modelName string, bodyBytes []byte, credentialID int64) (int, int, string, error) { defer resp.Body.Close() @@ -72,10 +83,7 @@ func (p *Proxy) handleStreamingResponse(w http.ResponseWriter, resp *http.Respon } } - scanner := bufio.NewScanner(reader) - // Increase buffer sizes to handle large SSE events (e.g., large file reads in tool calls) - buf := make([]byte, 0, 128*1024) // 128KB initial buffer (was 64KB) - scanner.Buffer(buf, 2*1024*1024) // 2MB max buffer (was 1MB) + scanner := newSSEScanner(reader) var inputTokens, outputTokens int var buffer bytes.Buffer @@ -218,9 +226,7 @@ func (p *Proxy) handleStreamingAsNonStreaming(w http.ResponseWriter, resp *http. } defer resp.Body.Close() - scanner := bufio.NewScanner(reader) - buf := make([]byte, 0, 128*1024) - scanner.Buffer(buf, 2*1024*1024) + scanner := newSSEScanner(reader) var completedPayload []byte var lastJSONPayload []byte @@ -329,7 +335,7 @@ func (p *Proxy) transformStreamEvent(eventData []byte, trans transformer.Transfo // extractTokensFromEvent extracts token counts from SSE event func (p *Proxy) extractTokensFromEvent(eventData []byte, inputTokens, outputTokens *int) { - scanner := bufio.NewScanner(bytes.NewReader(eventData)) + scanner := newSSEScanner(bytes.NewReader(eventData)) for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data:") { @@ -393,7 +399,7 @@ func (p *Proxy) extractTokensFromEvent(eventData []byte, inputTokens, outputToke // extractTextFromEvent extracts text content from transformed event // Enhanced to support both delta.text and content_block_delta formats func (p *Proxy) extractTextFromEvent(transformedEvent []byte, outputText *strings.Builder) { - scanner := bufio.NewScanner(bytes.NewReader(transformedEvent)) + scanner := newSSEScanner(bytes.NewReader(transformedEvent)) for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data:") { @@ -454,7 +460,7 @@ func (p *Proxy) isMessageStopEvent(eventData []byte) bool { } func hasStreamEventType(eventData []byte, want string) bool { - scanner := bufio.NewScanner(bytes.NewReader(eventData)) + scanner := newSSEScanner(bytes.NewReader(eventData)) for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data:") { diff --git a/internal/proxy/streaming_completion_test.go b/internal/proxy/streaming_completion_test.go index 0840b00..e738acb 100644 --- a/internal/proxy/streaming_completion_test.go +++ b/internal/proxy/streaming_completion_test.go @@ -128,6 +128,44 @@ func TestHandleStreamingResponseRequiresResponsesCompletion(t *testing.T) { } } +func TestHandleStreamingResponseAllowsLargeSSEEvent(t *testing.T) { + endpoint := config.Endpoint{ + Name: "OpenAIResponses", + APIUrl: "https://example.com", + APIKey: "x", + AuthMode: config.AuthModeAPIKey, + Enabled: true, + Transformer: "openai2", + Model: "gpt-image-2", + } + payload := strings.Repeat("a", 3*1024*1024) + completed := `data: {"type":"response.completed","response":{"output":"` + payload + `"}}` + "\n\n" + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(completed)), + } + rec := httptest.NewRecorder() + + _, _, _, err := (&Proxy{}).handleStreamingResponse( + rec, + resp, + endpoint, + responses.NewOpenAI2Transformer(endpoint.Model), + "cx_resp_openai2", + false, + endpoint.Model, + []byte(`{}`), + 0, + ) + if err != nil { + t.Fatalf("large SSE event failed: %v", err) + } + if rec.Body.String() != completed { + t.Fatalf("large SSE event was not forwarded unchanged") + } +} + func TestHandleStreamingResponseAllowsNonCurrentSpecifiedEndpoint(t *testing.T) { endpointA := config.Endpoint{ Name: "A", APIUrl: "https://a.example", APIKey: "key-a", AuthMode: config.AuthModeAPIKey,