fix(openai): reject malformed WS v2 events

This commit is contained in:
王鹏
2026-07-15 20:59:56 +08:00
parent d515c3045c
commit 716fcc6f3d
3 changed files with 189 additions and 0 deletions
@@ -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()
@@ -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"}`
@@ -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)