diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index d855f364a7..66bffe5a83 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -221,6 +221,43 @@ func TestOpenAIEnsureForwardErrorResponse_ResponsesRouteAfterWrittenEmitsRespons assert.Contains(t, body, "Upstream request failed") } +func TestOpenAIEnsureForwardErrorResponse_AfterDeltaAppendsSingleValidResponseFailed(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil) + + delta := `{"type":"response.output_text.delta","delta":"ok","sequence_number":1}` + _, err := c.Writer.WriteString("event: response.output_text.delta\ndata: " + delta + "\n\n") + require.NoError(t, err) + + h := &OpenAIGatewayHandler{} + require.True(t, h.ensureForwardErrorResponse(c, true)) + + frames := strings.Split(strings.TrimSuffix(w.Body.String(), "\n\n"), "\n\n") + require.Len(t, frames, 2) + errorEvents := 0 + for _, frame := range frames { + lines := strings.Split(frame, "\n") + require.Len(t, lines, 2) + require.True(t, strings.HasPrefix(lines[0], "event: ")) + require.True(t, strings.HasPrefix(lines[1], "data: ")) + + eventType := strings.TrimPrefix(lines[0], "event: ") + data := strings.TrimPrefix(lines[1], "data: ") + require.True(t, json.Valid([]byte(data)), "each downstream SSE frame must contain valid JSON") + var event struct { + Type string `json:"type"` + } + require.NoError(t, json.Unmarshal([]byte(data), &event)) + require.Equal(t, eventType, event.Type) + if eventType == "response.failed" { + errorEvents++ + } + } + require.Equal(t, 1, errorEvents) +} + func TestOpenAIEnsureForwardErrorResponse_ImageJSONKeepaliveWritesSingleJSONFallback(t *testing.T) { gin.SetMode(gin.TestMode) w := httptest.NewRecorder() diff --git a/backend/internal/service/openai_sse_concatenated_json_test.go b/backend/internal/service/openai_sse_concatenated_json_test.go index 9b06ef9a39..27c057157e 100644 --- a/backend/internal/service/openai_sse_concatenated_json_test.go +++ b/backend/internal/service/openai_sse_concatenated_json_test.go @@ -90,6 +90,138 @@ func TestOpenAIWSv2StreamingRepairsConcatenatedJSONDocumentsInSingleMessage(t *t }) } +func TestOpenAIWSv2RejectsMalformedTypedEventBeforeWritingDownstream(t *testing.T) { + largeInProgress, _, _ := openAIConcatenatedJSONTestEvents(t) + testOpenAIWSv2RejectsMalformedEventBeforeWritingDownstream(t, []byte(largeInProgress+"unexpected-tail")) +} + +func TestOpenAIWSv2RejectsMalformedUntypedMessageBeforeWritingDownstream(t *testing.T) { + testOpenAIWSv2RejectsMalformedEventBeforeWritingDownstream(t, []byte("not-json")) +} + +func TestOpenAIWSv2RejectsMalformedEventAfterWritingDownstream(t *testing.T) { + gin.SetMode(gin.TestMode) + + outputTextDelta := `{"type":"response.output_text.delta","delta":"ok","sequence_number":1}` + malformedMessage := `{"type":"response.in_progress"}unexpected-tail` + captureConn := &openAIWSCaptureConn{events: [][]byte{ + []byte(outputTextDelta), + []byte(malformedMessage), + }} + + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 5 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + pool := newOpenAIWSConnPool(cfg) + pool.setClientDialerForTest(&openAIWSCaptureDialer{conn: captureConn}) + svc := &OpenAIGatewayService{ + cfg: cfg, + cache: &stubGatewayCache{}, + httpUpstream: &httpUpstreamRecorder{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + openaiWSPool: pool, + toolCorrector: NewCodexToolCorrector(), + } + account := &Account{ + ID: 5, + Name: "ws-malformed-event-after-output", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{"responses_websockets_v2_enabled": true}, + } + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + groupID := int64(1) + c.Set("api_key", &APIKey{GroupID: &groupID}) + + result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.6-sol","stream":true,"input":"hello"}`)) + require.Error(t, err) + require.Contains(t, err.Error(), "after downstream output") + require.Nil(t, result) + require.True(t, captureConn.closed) + require.Contains(t, recorder.Body.String(), `"delta":"ok"`) + require.NotContains(t, recorder.Body.String(), "unexpected-tail") + require.NotContains(t, recorder.Body.String(), "response.in_progress") + assertOpenAISSEFrames(t, recorder.Body.String(), []string{"response.output_text.delta"}) +} + +func testOpenAIWSv2RejectsMalformedEventBeforeWritingDownstream(t *testing.T, malformedMessage []byte) { + t.Helper() + gin.SetMode(gin.TestMode) + + _, _, completed := openAIConcatenatedJSONTestEvents(t) + outputTextDelta := `{"type":"response.output_text.delta","delta":"ok","sequence_number":3}` + captureConn := &openAIWSCaptureConn{events: [][]byte{ + malformedMessage, + []byte(outputTextDelta), + []byte(completed), + }} + + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 5 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + pool := newOpenAIWSConnPool(cfg) + pool.setClientDialerForTest(&openAIWSCaptureDialer{conn: captureConn}) + svc := &OpenAIGatewayService{ + cfg: cfg, + cache: &stubGatewayCache{}, + httpUpstream: &httpUpstreamRecorder{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + openaiWSPool: pool, + toolCorrector: NewCodexToolCorrector(), + } + account := &Account{ + ID: 4, + Name: "ws-malformed-event", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{"responses_websockets_v2_enabled": true}, + } + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + groupID := int64(1) + c.Set("api_key", &APIKey{GroupID: &groupID}) + + result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.6-sol","stream":true,"input":"hello"}`)) + require.Error(t, err) + var fallbackErr *openAIWSFallbackError + require.ErrorAs(t, err, &fallbackErr) + require.Equal(t, "invalid_event_json", fallbackErr.Reason) + require.Nil(t, result) + require.Empty(t, recorder.Body.String()) + require.True(t, captureConn.closed) +} + func TestSplitOpenAIConcatenatedJSONDocumentsRejectsPayloadOverRepairLimit(t *testing.T) { first := `{"type":"response.in_progress","padding":"` + strings.Repeat("x", 16*1024*1024) + `"}` second := `{"type":"response.completed"}` diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index 9054e60502..eefadbb82d 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -3,6 +3,7 @@ package service import ( "bytes" "context" + "encoding/json" "errors" "fmt" "net/http" @@ -454,6 +455,25 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( } } } + if readErr == nil && !json.Valid(message) { + eventType, _, _ := parseOpenAIWSEventEnvelope(message) + if eventType == "" { + eventType = "unknown" + } + lease.MarkBroken() + logOpenAIWSModeInfo( + "invalid_event_json account_id=%d conn_id=%s event_type=%s bytes=%d wrote_downstream=%v", + account.ID, + truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen), + truncateOpenAIWSLogValue(eventType, openAIWSLogValueMaxLen), + len(message), + wroteDownstream, + ) + if !wroteDownstream { + return nil, wrapOpenAIWSFallback("invalid_event_json", errors.New("upstream websocket returned malformed Responses event JSON")) + } + return nil, errors.New("upstream websocket returned malformed Responses event JSON after downstream output") + } if readErr != nil { lease.MarkBroken() closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)