Merge pull request #5864 from wucm667/fix/issue-5850-http-bridge-replay

fix(openai): avoid duplicate HTTP bridge replay
This commit is contained in:
Wesley Liddick
2026-08-22 13:34:32 +08:00
committed by GitHub
6 changed files with 366 additions and 9 deletions
@@ -235,16 +235,16 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
return coverage
}
input := parseRawJSONView(body).Get("input")
if !input.IsArray() {
if !input.IsArray() && !input.IsObject() {
return coverage
}
missingCallID := false
var outputCallIDs map[string]struct{}
var contextIDs map[string]struct{}
input.ForEach(func(_, item gjson.Result) bool {
analyzeItem := func(item gjson.Result) {
if !item.IsObject() {
return true
return
}
itemType := item.Get("type").String()
switch {
@@ -253,7 +253,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
callID := strings.TrimSpace(item.Get("call_id").String())
if callID == "" {
missingCallID = true
return true
return
}
if outputCallIDs == nil {
outputCallIDs = make(map[string]struct{})
@@ -262,7 +262,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
case isCodexToolCallContextItemType(itemType):
callID := strings.TrimSpace(item.Get("call_id").String())
if callID == "" {
return true
return
}
if contextIDs == nil {
contextIDs = make(map[string]struct{})
@@ -271,15 +271,22 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex
case itemType == "item_reference":
idValue := strings.TrimSpace(item.Get("id").String())
if idValue == "" {
return true
return
}
if contextIDs == nil {
contextIDs = make(map[string]struct{})
}
contextIDs[idValue] = struct{}{}
}
return true
})
}
if input.IsArray() {
input.ForEach(func(_, item gjson.Result) bool {
analyzeItem(item)
return true
})
} else {
analyzeItem(input)
}
if !coverage.HasFunctionCallOutput || missingCallID {
return coverage
@@ -206,6 +206,14 @@ func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) {
hasOutput: false,
coversAllIDs: false,
},
{
name: "object_tool_output_requires_context_replay",
body: map[string]any{"input": map[string]any{
"type": "custom_tool_call_output", "call_id": "call_a",
}},
hasOutput: true,
coversAllIDs: false,
},
{
name: "all_outputs_covered_by_context",
body: map[string]any{"input": []any{
@@ -565,7 +565,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
bridgePayloadRaw := currentBridgePayload.payloadRaw
bridgePayloadBytes := currentBridgePayload.payloadBytes
needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw)
toolOutputCoverage := AnalyzeToolCallOutputContextCoverageBytes(currentBridgePayload.payloadRaw)
needsBridgeReplay := currentBridgePayload.previousResponseID != "" ||
(toolOutputCoverage.HasFunctionCallOutput && !toolOutputCoverage.ContextCoversAllCallIDs)
turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
bridgeReplayInput,
bridgeReplayInputExists,
@@ -809,6 +809,35 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) {
require.Equal(t, "new", gjson.GetBytes(items[0], "text").String())
})
t.Run("no_previous_response_id_custom_tool_history_does_not_accumulate", func(t *testing.T) {
previousFull := []json.RawMessage{
json.RawMessage(`{"type":"input_text","text":"stale"}`),
json.RawMessage(`{"type":"custom_tool_call","id":"stale_item","call_id":"stale_call","name":"exec","input":"stale"}`),
}
currentPayload := []byte(`{"input":[
{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"},
{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"},
{"type":"input_text","text":"continue"}
]}`)
for range 3 {
items, exists, err := buildOpenAIWSReplayInputSequence(
previousFull,
true,
currentPayload,
false,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 3)
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[0], "type").String())
require.Equal(t, "call_1", gjson.GetBytes(items[0], "call_id").String())
require.Equal(t, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String())
require.Equal(t, "call_1", gjson.GetBytes(items[1], "call_id").String())
previousFull = append(items, json.RawMessage(`{"type":"custom_tool_call","id":"replayed_item","call_id":"replayed_call","name":"exec","input":"ignored"}`))
}
})
t.Run("previous_response_id_delta_append", func(t *testing.T) {
items, exists, err := buildOpenAIWSReplayInputSequence(
lastFull,
@@ -823,6 +852,91 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) {
require.Equal(t, "world", gjson.GetBytes(items[1], "text").String())
})
t.Run("previous_response_id_filters_orphan_historical_custom_tool_call", func(t *testing.T) {
previousFull := []json.RawMessage{
json.RawMessage(`{"type":"input_text","text":"hello"}`),
json.RawMessage(`{"type":"custom_tool_call","id":"item_orphan","call_id":"call_orphan","name":"exec","input":"pwd"}`),
}
items, exists, err := buildOpenAIWSReplayInputSequence(
previousFull,
true,
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
true,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 2)
require.Equal(t, "hello", gjson.GetBytes(items[0], "text").String())
require.Equal(t, "user", gjson.GetBytes(items[1], "role").String())
})
t.Run("previous_response_id_preserves_paired_historical_function_call", func(t *testing.T) {
previousFull := []json.RawMessage{
json.RawMessage(`{"type":"function_call","id":"item_1","call_id":"call_1","name":"lookup","arguments":"{}"}`),
json.RawMessage(`{"type":"function_call_output","call_id":"call_1","output":"ok"}`),
}
items, exists, err := buildOpenAIWSReplayInputSequence(
previousFull,
true,
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
true,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 3)
require.Equal(t, "function_call", gjson.GetBytes(items[0], "type").String())
require.Equal(t, "function_call_output", gjson.GetBytes(items[1], "type").String())
})
t.Run("previous_response_id_preserves_paired_historical_custom_tool_call", func(t *testing.T) {
previousFull := []json.RawMessage{
json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`),
json.RawMessage(`{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}`),
}
items, exists, err := buildOpenAIWSReplayInputSequence(
previousFull,
true,
[]byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`),
true,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 3)
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[0], "type").String())
require.Equal(t, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String())
})
t.Run("item_reference_does_not_complete_historical_call", func(t *testing.T) {
previousFull := []json.RawMessage{
json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`),
}
items, exists, err := buildOpenAIWSReplayInputSequence(
previousFull,
true,
[]byte(`{"previous_response_id":"resp_1","input":[{"type":"item_reference","id":"call_1"},{"role":"user","content":"continue"}]}`),
true,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 2)
require.Equal(t, "item_reference", gjson.GetBytes(items[0], "type").String())
require.Equal(t, "user", gjson.GetBytes(items[1], "role").String())
})
t.Run("previous_response_id_preserves_current_orphan_custom_tool_call", func(t *testing.T) {
items, exists, err := buildOpenAIWSReplayInputSequence(
lastFull,
true,
[]byte(`{"previous_response_id":"resp_1","input":[{"type":"custom_tool_call","id":"item_live","call_id":"call_live","name":"exec","input":"pwd"}]}`),
true,
)
require.NoError(t, err)
require.True(t, exists)
require.Len(t, items, 2)
require.Equal(t, "custom_tool_call", gjson.GetBytes(items[1], "type").String())
require.Equal(t, "call_live", gjson.GetBytes(items[1], "call_id").String())
})
t.Run("previous_response_id_full_input_replace", func(t *testing.T) {
items, exists, err := buildOpenAIWSReplayInputSequence(
lastFull,
@@ -610,6 +610,40 @@ func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool
return true
}
func sanitizeOpenAIWSHistoricalReplayToolCalls(
previousItems []json.RawMessage,
currentItems []json.RawMessage,
) []json.RawMessage {
if len(previousItems) == 0 {
return cloneOpenAIWSRawMessages(previousItems)
}
outputCallIDs := make(map[string]struct{})
collectOutputCallIDs := func(items []json.RawMessage) {
for _, item := range items {
if !isCodexToolCallOutputItemType(gjson.GetBytes(item, "type").String()) {
continue
}
if callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String()); callID != "" {
outputCallIDs[callID] = struct{}{}
}
}
}
collectOutputCallIDs(previousItems)
collectOutputCallIDs(currentItems)
sanitized := make([]json.RawMessage, 0, len(previousItems))
for _, item := range previousItems {
if isCodexToolCallContextItemType(gjson.GetBytes(item, "type").String()) {
callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String())
if _, paired := outputCallIDs[callID]; !paired {
continue
}
}
sanitized = append(sanitized, append(json.RawMessage(nil), item...))
}
return sanitized
}
func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool {
if len(payload) == 0 {
return false
@@ -648,6 +682,7 @@ func buildOpenAIWSReplayInputSequence(
if !previousFullInputExists {
return cloneOpenAIWSRawMessages(currentItems), currentExists, nil
}
previousFullInput = sanitizeOpenAIWSHistoricalReplayToolCalls(previousFullInput, currentItems)
if !currentExists || len(currentItems) == 0 {
return cloneOpenAIWSRawMessages(previousFullInput), true, nil
}
@@ -416,6 +416,197 @@ func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t
require.False(t, secondInput[2].Get("id").Exists())
}
func TestOpenAIWSHTTPBridgeFullCustomToolHistoryWithoutPreviousResponseIDDoesNotReplay(t *testing.T) {
gin.SetMode(gin.TestMode)
completed := func(responseID string, output string) string {
return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n"
}
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_3", `[]`)))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_4", `[]`)))},
}}
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.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
svc := &OpenAIGatewayService{
cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(),
}
account := &Account{
ID: 9002, Name: "oauth-full-context", Platform: PlatformOpenAI, Type: AccountTypeOAuth,
Credentials: map[string]any{"access_token": "test-token"}, Extra: map[string]any{"responses_websockets_v2_enabled": true},
Concurrency: 1, Status: StatusActive, Schedulable: true,
}
errCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
errCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, firstMessage, err := conn.Read(readCtx)
cancelRead()
if err != nil {
errCh <- err
return
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
ginCtx.Request = r.Clone(r.Context())
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "test-token", 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() }()
writeAndRead := func(payload string) {
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
cancelWrite()
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, event, readErr := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, readErr)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
}
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`)
writeAndRead(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_1","input":[{"role":"user","content":"continue without tool output"}]}`)
fullContext := `{"type":"response.create","model":"gpt-5.1","input":[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"},{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"},{"role":"user","content":"continue"}]}`
writeAndRead(fullContext)
writeAndRead(fullContext)
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
select {
case proxyErr := <-errCh:
require.NoError(t, proxyErr)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for websocket bridge proxy to finish")
}
require.Len(t, upstream.bodies, 4)
orphanInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
require.Len(t, orphanInput, 2)
require.Equal(t, "run pwd", orphanInput[0].String())
require.Equal(t, "user", orphanInput[1].Get("role").String())
for _, body := range upstream.bodies[2:] {
input := gjson.GetBytes(body, "input").Array()
require.Len(t, input, 3)
require.Equal(t, "custom_tool_call", input[0].Get("type").String())
require.Equal(t, "call_1", input[0].Get("call_id").String())
require.Equal(t, "custom_tool_call_output", input[1].Get("type").String())
require.Equal(t, "call_1", input[1].Get("call_id").String())
}
}
func TestOpenAIWSHTTPBridgeObjectToolOutputWithoutPreviousResponseIDReplaysMatchingCall(t *testing.T) {
gin.SetMode(gin.TestMode)
completed := func(responseID string, output string) string {
return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n"
}
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))},
}}
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.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
svc := &OpenAIGatewayService{
cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(),
}
account := &Account{
ID: 9003, Name: "oauth-output-only", Platform: PlatformOpenAI, Type: AccountTypeOAuth,
Credentials: map[string]any{"access_token": "test-token"}, Extra: map[string]any{"responses_websockets_v2_enabled": true},
Concurrency: 1, Status: StatusActive, Schedulable: true,
}
errCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
errCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, firstMessage, err := conn.Read(readCtx)
cancelRead()
if err != nil {
errCh <- err
return
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
ginCtx.Request = r.Clone(r.Context())
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "test-token", 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() }()
writeAndRead := func(payload string) {
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
cancelWrite()
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, event, readErr := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, readErr)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
}
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`)
writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}}`)
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
select {
case proxyErr := <-errCh:
require.NoError(t, proxyErr)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for websocket bridge proxy to finish")
}
require.Len(t, upstream.bodies, 2)
secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
require.Len(t, secondInput, 3)
require.Equal(t, "custom_tool_call", secondInput[1].Get("type").String())
require.Equal(t, "call_1", secondInput[1].Get("call_id").String())
require.Equal(t, "custom_tool_call_output", secondInput[2].Get("type").String())
require.Equal(t, "call_1", secondInput[2].Get("call_id").String())
}
func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
svc := &OpenAIGatewayService{
cfg: &config.Config{