mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-01 15:02:58 +08:00
fix(openai): reject malformed WS v2 events
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user