Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions internal/proxy/codex_rate_limits.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package proxy

import (
"bufio"
"bytes"
"context"
"encoding/json"
Expand Down Expand Up @@ -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:") {
Expand Down
26 changes: 16 additions & 10 deletions internal/proxy/streaming.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:") {
Expand Down Expand Up @@ -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:") {
Expand Down Expand Up @@ -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:") {
Expand Down
38 changes: 38 additions & 0 deletions internal/proxy/streaming_completion_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading