fix(openai): avoid duplicate HTTP bridge replay

This commit is contained in:
wucm667
2026-08-20 20:10:57 +08:00
parent 2bc139ab52
commit 25da02dddd
5 changed files with 240 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{
@@ -511,7 +511,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,
@@ -413,6 +413,191 @@ 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", `[]`)))},
}}
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"}`)
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, 3)
for _, body := range upstream.bodies[1:] {
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{