Skip to content
Open
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
1 change: 1 addition & 0 deletions backend/internal/service/openai_gateway_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,7 @@ type OpenAIForwardResult struct {
FirstTokenMs *int
ImageCount int
ImageSize string
HasToolCall bool
}

type OpenAIWSRetryMetricsSnapshot struct {
Expand Down
12 changes: 10 additions & 2 deletions backend/internal/service/openai_ws_forwarder.go
Original file line number Diff line number Diff line change
Expand Up @@ -2822,6 +2822,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
eventCount := 0
tokenEventCount := 0
terminalEventCount := 0
emittedToolCall := false
firstEventType := ""
lastEventType := ""
needModelReplace := false
Expand Down Expand Up @@ -2939,6 +2940,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
}
if openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
emittedToolCall = true
if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
upstreamMessage = corrected
}
Expand Down Expand Up @@ -3004,6 +3006,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
ResponseHeaders: lease.HandshakeHeaders(),
Duration: time.Since(turnStart),
FirstTokenMs: firstTokenMs,
HasToolCall: emittedToolCall,
}, nil
}
}
Expand Down Expand Up @@ -3075,6 +3078,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
lastTurnResponseID := ""
lastTurnPayload := []byte(nil)
var lastTurnStrictState *openAIWSIngressPreviousTurnStrictState
lastTurnHasToolCall := false
lastTurnReplayInput := []json.RawMessage(nil)
lastTurnReplayInputExists := false
currentTurnReplayInput := []json.RawMessage(nil)
Expand Down Expand Up @@ -3254,7 +3258,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
shouldKeepPreviousResponseID := false
strictReason := ""
var strictErr error
if lastTurnStrictState != nil {
if lastTurnHasToolCall && expectedPrev != "" && currentPreviousResponseID == expectedPrev {
shouldKeepPreviousResponseID = true
strictReason = "previous_turn_has_tool_call"
} else if lastTurnStrictState != nil {
shouldKeepPreviousResponseID, strictReason, strictErr = shouldKeepIngressPreviousResponseIDWithStrictState(
lastTurnStrictState,
currentPayload,
Expand Down Expand Up @@ -3365,7 +3372,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
truncateOpenAIWSLogValue(pingErr.Error(), openAIWSLogValueMaxLen),
)
if forcePreferredConn {
if !turnPrevRecoveryTried && currentPreviousResponseID != "" {
if !turnPrevRecoveryTried && currentPreviousResponseID != "" && !hasFunctionCallOutput {
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
reason := "not_removed"
Expand Down Expand Up @@ -3489,6 +3496,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
lastTurnPayload = cloneOpenAIWSPayloadBytes(currentPayload)
lastTurnReplayInput = cloneOpenAIWSRawMessages(currentTurnReplayInput)
lastTurnReplayInputExists = currentTurnReplayInputExists
lastTurnHasToolCall = result.HasToolCall
nextStrictState, strictStateErr := buildOpenAIWSIngressPreviousTurnStrictState(currentPayload)
if strictStateErr != nil {
lastTurnStrictState = nil
Expand Down
280 changes: 280 additions & 0 deletions backend/internal/service/openai_ws_forwarder_ingress_session_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -673,6 +673,144 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPre
require.Equal(t, "world", gjson.Get(secondWrite, "input.1.text").String())
}

func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledKeepsPreviousResponseIDAfterToolCallDespiteNonInputChange(t *testing.T) {
gin.SetMode(gin.TestMode)

cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.OAuthEnabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3

captureConn := &openAIWSCaptureConn{
events: [][]byte{
[]byte(`{"type":"response.output_item.done","response_id":"resp_tool_pending_1","item":{"type":"function_call","call_id":"call_pending_1","name":"shell","arguments":"{}"}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_tool_pending_1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_tool_pending_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
}
captureDialer := &openAIWSCaptureDialer{conn: captureConn}
pool := newOpenAIWSConnPool(cfg)
pool.setClientDialerForTest(captureDialer)

svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: &httpUpstreamRecorder{},
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
openaiWSPool: pool,
}

account := &Account{
ID: 157,
Name: "openai-ingress-session-pending-tool-call-anchor",
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,
},
}

serverErrCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
CompressionMode: coderws.CompressionContextTakeover,
})
if err != nil {
serverErrCh <- err
return
}
defer func() {
_ = conn.CloseNow()
}()

rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
req := r.Clone(r.Context())
req.Header = req.Header.Clone()
req.Header.Set("User-Agent", "unit-test-agent/1.0")
ginCtx.Request = req

readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
msgType, firstMessage, readErr := conn.Read(readCtx)
cancel()
if readErr != nil {
serverErrCh <- readErr
return
}
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
serverErrCh <- errors.New("unsupported websocket client message type")
return
}

serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
}))
defer wsServer.Close()

dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() {
_ = clientConn.CloseNow()
}()

writeMessage := func(payload string) {
writeCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
}
readMessage := func() []byte {
readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
msgType, message, readErr := clientConn.Read(readCtx)
require.NoError(t, readErr)
require.Equal(t, coderws.MessageText, msgType)
return message
}

writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"input":[{"type":"input_text","text":"hello"}]}`)
toolCallEvent := readMessage()
require.Equal(t, "response.output_item.done", gjson.GetBytes(toolCallEvent, "type").String())
firstTurn := readMessage()
require.Equal(t, "resp_tool_pending_1", gjson.GetBytes(firstTurn, "response.id").String())

writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"instructions":"changed","previous_response_id":"resp_tool_pending_1","input":[{"type":"input_text","text":"world"}]}`)
secondTurn := readMessage()
require.Equal(t, "resp_tool_pending_2", gjson.GetBytes(secondTurn, "response.id").String())

require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
select {
case serverErr := <-serverErrCh:
require.NoError(t, serverErr)
case <-time.After(5 * time.Second):
t.Fatal("等待 ingress websocket 结束超时")
}

require.Equal(t, 1, captureDialer.DialCount())
require.Len(t, captureConn.writes, 2)
secondWrite := requestToJSONString(captureConn.writes[1])
require.Equal(t, "resp_tool_pending_1", gjson.Get(secondWrite, "previous_response_id").String(), "上一轮存在 pending tool call 时不能丢弃 previous_response_id")
require.Equal(t, "changed", gjson.Get(secondWrite, "instructions").String())
require.Equal(t, "world", gjson.Get(secondWrite, "input.0.text").String())
}

func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPrevResponseStrictDropBeforePreflightPingFailReconnects(t *testing.T) {
gin.SetMode(gin.TestMode)
prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
Expand Down Expand Up @@ -825,6 +963,148 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPre
require.Equal(t, "world", gjson.Get(secondWrite, "input.1.text").String())
}

func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledFunctionCallOutputPreflightPingFailDoesNotDropPreviousResponseID(t *testing.T) {
gin.SetMode(gin.TestMode)
prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
openAIWSIngressPreflightPingIdle = 0
defer func() {
openAIWSIngressPreflightPingIdle = prevPreflightPingIdle
}()

cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.OAuthEnabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3

firstConn := &openAIWSPreflightFailConn{
events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_tool_1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
}
secondConn := &openAIWSCaptureConn{
events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_tool_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
}
dialer := &openAIWSQueueDialer{
conns: []openAIWSClientConn{firstConn, secondConn},
}
pool := newOpenAIWSConnPool(cfg)
pool.setClientDialerForTest(dialer)

svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: &httpUpstreamRecorder{},
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
openaiWSPool: pool,
}

account := &Account{
ID: 158,
Name: "openai-ingress-fco-preflight-no-drop",
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,
},
}

serverErrCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
CompressionMode: coderws.CompressionContextTakeover,
})
if err != nil {
serverErrCh <- err
return
}
defer func() {
_ = conn.CloseNow()
}()

rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
req := r.Clone(r.Context())
req.Header = req.Header.Clone()
req.Header.Set("User-Agent", "unit-test-agent/1.0")
ginCtx.Request = req

readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
msgType, firstMessage, readErr := conn.Read(readCtx)
cancel()
if readErr != nil {
serverErrCh <- readErr
return
}
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
serverErrCh <- errors.New("unsupported websocket client message type")
return
}

serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
}))
defer wsServer.Close()

dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() {
_ = clientConn.CloseNow()
}()

writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"input":[{"type":"input_text","text":"hello"}]}`))
cancelWrite()
require.NoError(t, err)

readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
msgType, firstTurn, readErr := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, readErr)
require.Equal(t, coderws.MessageText, msgType)
require.Equal(t, "resp_turn_ping_tool_1", gjson.GetBytes(firstTurn, "response.id").String())

writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"previous_response_id":"resp_turn_ping_tool_1","input":[{"type":"function_call","call_id":"call_pending_1","name":"shell","arguments":"{}"},{"type":"function_call_output","call_id":"call_pending_1","output":"ok"}]}`))
cancelWrite()
require.NoError(t, err)

select {
case serverErr := <-serverErrCh:
var closeErr *OpenAIWSClientCloseError
require.ErrorAs(t, serverErr, &closeErr)
require.Equal(t, coderws.StatusPolicyViolation, closeErr.StatusCode())
case <-time.After(5 * time.Second):
t.Fatal("等待 function_call_output 预检失败保护结束超时")
}

require.Equal(t, 1, firstConn.WriteCount(), "preflight ping 失败后不应继续向旧连接发送第二轮")
require.GreaterOrEqual(t, firstConn.PingCount(), 1, "第二轮前应执行 preflight ping")
secondConn.mu.Lock()
secondWrites := append([]map[string]any(nil), secondConn.writes...)
secondConn.mu.Unlock()
require.Empty(t, secondWrites, "function_call_output 依赖 previous_response_id 时不应删除锚点后换连重放")
}

func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreEnabledSkipsStrictPrevResponseEval(t *testing.T) {
gin.SetMode(gin.TestMode)

Expand Down
Loading