Merge pull request #5822 from hansnow/fix/ws-http-bridge-followup-client-tools

fix(openai): 修复 WS HTTP bridge 多轮客户端工具回放
This commit is contained in:
Wesley Liddick
2026-08-19 16:42:28 +08:00
committed by GitHub
5 changed files with 288 additions and 6 deletions
@@ -130,6 +130,42 @@ func AdaptResponsesClientTools(req map[string]any) (ResponsesClientToolMapping,
return adapter, changed, nil
}
// AdaptResponsesClientToolsWithInheritedMapping lowers client-tool history on
// a follow-up request that omits the session-level tools declaration. An
// explicitly present tools field, including an empty or malformed value,
// always replaces the inherited mapping and is handled by the ordinary
// declaration-driven adapter.
func AdaptResponsesClientToolsWithInheritedMapping(
req map[string]any,
inherited ResponsesClientToolMapping,
) (ResponsesClientToolMapping, bool, error) {
if req == nil {
return ResponsesClientToolMapping{}, false, nil
}
if _, toolsPresent := req["tools"]; toolsPresent {
return AdaptResponsesClientTools(req)
}
if len(inherited.CustomTools) == 0 && !inherited.ToolSearch && len(inherited.NamespaceTools) == 0 {
return ResponsesClientToolMapping{}, false, nil
}
changed := rewriteClientToolHistory(req["input"], &inherited)
if len(inherited.NamespaceTools) > 0 {
before := changed
rewriteNamespaceQualifiedCalls(req["input"], inherited.NamespaceTools)
// Namespace rewriting does not currently report whether it changed a
// value. A retained namespace mapping is only used for follow-up
// history, so conservatively rebuild the request when input exists.
if _, inputPresent := req["input"]; inputPresent && !before {
changed = true
}
}
if rewriteClientToolChoice(req, &inherited) {
changed = true
}
return inherited, changed, nil
}
func copyClientTool(tool map[string]any) map[string]any {
copy := make(map[string]any, len(tool))
for key, value := range tool {
@@ -155,10 +191,12 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo
typed["type"] = "function_call"
typed["arguments"] = customToolCallArguments(stringValue(typed["input"]))
delete(typed, "input")
dropInvalidLoweredFunctionItemID(typed)
changed = true
}
case "custom_tool_call_output":
typed["type"] = "function_call_output"
dropInvalidLoweredFunctionItemID(typed)
normalizeClientToolOutput(typed)
changed = true
case "tool_search_call":
@@ -167,11 +205,13 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo
typed["name"] = toolSearchProxyName
typed["arguments"] = rawObjectString(typed["arguments"])
delete(typed, "execution")
dropInvalidLoweredFunctionItemID(typed)
changed = true
}
case "tool_search_output":
if adapter.ToolSearch {
typed["type"] = "function_call_output"
dropInvalidLoweredFunctionItemID(typed)
normalizeClientToolOutput(typed)
changed = true
}
@@ -185,6 +225,17 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo
return changed
}
// dropInvalidLoweredFunctionItemID removes Codex client-only item IDs such as
// ctc_*, ctco_*, tsc_*, and tso_* after their item type is lowered to the
// function protocol. Function upstreams validate these IDs with the fc prefix;
// call_id, which is preserved separately, is the tool call/output pairing key.
func dropInvalidLoweredFunctionItemID(item map[string]any) {
id := strings.TrimSpace(stringValue(item["id"]))
if id != "" && !strings.HasPrefix(id, "fc") {
delete(item, "id")
}
}
func normalizeClientToolOutput(item map[string]any) {
output, exists := item["output"]
if !exists {
@@ -17,10 +17,10 @@ func TestAdaptResponsesClientTools_LowersDeclarationsHistoryChoiceAndNamespaces(
},
"tool_choice": map[string]any{"type": "custom", "name": "exec"},
"input": []any{
map[string]any{"type": "custom_tool_call", "call_id": "c1", "name": "exec", "input": "dir"},
map[string]any{"type": "custom_tool_call_output", "call_id": "c1", "output": "ok"},
map[string]any{"type": "tool_search_call", "call_id": "s1", "arguments": map[string]any{"query": "git"}},
map[string]any{"type": "tool_search_output", "call_id": "s1", "output": map[string]any{"groups": []string{"git"}}},
map[string]any{"type": "custom_tool_call", "id": "ctc_client", "call_id": "c1", "name": "exec", "input": "dir"},
map[string]any{"type": "custom_tool_call_output", "id": "ctco_client", "call_id": "c1", "output": "ok"},
map[string]any{"type": "tool_search_call", "id": "tsc_client", "call_id": "s1", "arguments": map[string]any{"query": "git"}},
map[string]any{"type": "tool_search_output", "id": "tso_client", "call_id": "s1", "output": map[string]any{"groups": []string{"git"}}},
map[string]any{"type": "function_call", "call_id": "n1", "namespace": "team", "name": "send", "arguments": "{}"},
},
}
@@ -48,15 +48,19 @@ func TestAdaptResponsesClientTools_LowersDeclarationsHistoryChoiceAndNamespaces(
input := requireResponsesClientToolValue[[]any](t, req["input"])
customCall := requireResponsesClientToolValue[map[string]any](t, input[0])
require.Equal(t, "function_call", customCall["type"])
require.NotContains(t, customCall, "id")
require.JSONEq(t, `{"input":"dir"}`, requireResponsesClientToolValue[string](t, customCall["arguments"]))
customOutput := requireResponsesClientToolValue[map[string]any](t, input[1])
require.Equal(t, "function_call_output", customOutput["type"])
require.NotContains(t, customOutput, "id")
searchCall := requireResponsesClientToolValue[map[string]any](t, input[2])
require.Equal(t, "function_call", searchCall["type"])
require.NotContains(t, searchCall, "id")
require.Equal(t, toolSearchProxyName, searchCall["name"])
require.JSONEq(t, `{"query":"git"}`, requireResponsesClientToolValue[string](t, searchCall["arguments"]))
searchOutput := requireResponsesClientToolValue[map[string]any](t, input[3])
require.Equal(t, "function_call_output", searchOutput["type"])
require.NotContains(t, searchOutput, "id")
require.JSONEq(t, `{"groups":["git"]}`, requireResponsesClientToolValue[string](t, searchOutput["output"]))
namespaceCall := requireResponsesClientToolValue[map[string]any](t, input[4])
require.Equal(t, "team__send", namespaceCall["name"])
@@ -81,6 +85,59 @@ func TestAdaptResponsesClientTools_RejectsAmbiguousNames(t *testing.T) {
}
}
func TestAdaptResponsesClientToolsWithInheritedMapping_LowersFollowupHistoryWithoutTools(t *testing.T) {
req := map[string]any{
"input": []any{
map[string]any{
"type": "custom_tool_call", "name": "exec",
"call_id": "call_1", "input": "pwd",
},
map[string]any{
"type": "custom_tool_call_output", "call_id": "call_1",
"id": "ctco_client_output_1",
"output": []any{map[string]any{"type": "input_text", "text": "ok"}},
},
},
}
inherited := ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}}
mapping, changed, err := AdaptResponsesClientToolsWithInheritedMapping(req, inherited)
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, inherited, mapping)
items := requireResponsesClientToolValue[[]any](t, req["input"])
call := requireResponsesClientToolValue[map[string]any](t, items[0])
require.Equal(t, "function_call", call["type"])
require.JSONEq(t, `{"input":"pwd"}`, requireResponsesClientToolValue[string](t, call["arguments"]))
require.NotContains(t, call, "input")
output := requireResponsesClientToolValue[map[string]any](t, items[1])
require.Equal(t, "function_call_output", output["type"])
require.NotContains(t, output, "id")
require.JSONEq(t, `[{"text":"ok","type":"input_text"}]`, requireResponsesClientToolValue[string](t, output["output"]))
}
func TestAdaptResponsesClientToolsWithInheritedMapping_ExplicitToolsReplaceInheritedMapping(t *testing.T) {
req := map[string]any{
"tools": []any{},
"input": []any{map[string]any{
"type": "custom_tool_call", "name": "exec", "input": "pwd",
}},
}
mapping, changed, err := AdaptResponsesClientToolsWithInheritedMapping(
req,
ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}},
)
require.NoError(t, err)
require.False(t, changed)
require.Empty(t, mapping)
items := requireResponsesClientToolValue[[]any](t, req["input"])
call := requireResponsesClientToolValue[map[string]any](t, items[0])
require.Equal(t, "custom_tool_call", call["type"])
}
func TestRestoreResponsesClientToolPayload_RestoresClientAndNamespaceCalls(t *testing.T) {
mapping := ResponsesClientToolMapping{
CustomTools: map[string]bool{"exec": true}, ToolSearch: true,
@@ -16,6 +16,18 @@ import (
const grokResponsesClientToolMappingContextKey = "grok_responses_client_tool_mapping"
func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) ([]byte, apicompat.ResponsesClientToolMapping, error) {
return adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body,
upstream,
apicompat.ResponsesClientToolMapping{},
)
}
func adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body []byte,
upstream string,
inherited apicompat.ResponsesClientToolMapping,
) ([]byte, apicompat.ResponsesClientToolMapping, error) {
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
var requestBody map[string]any
@@ -23,7 +35,7 @@ func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string)
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode %s Responses client tools: %w", upstream, err)
}
mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody)
mapping, changed, err := apicompat.AdaptResponsesClientToolsWithInheritedMapping(requestBody, inherited)
if err != nil {
return body, apicompat.ResponsesClientToolMapping{}, err
}
@@ -15,6 +15,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
const (
@@ -23,6 +24,30 @@ const (
openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024
)
const openAIWSHTTPBridgeToolStateContextKey = "openai_ws_http_bridge_tool_state"
type openAIWSHTTPBridgeToolState struct {
ClientMapping apicompat.ResponsesClientToolMapping
LoweredTools json.RawMessage
}
func openAIWSHTTPBridgeToolStateFromContext(c *gin.Context) (openAIWSHTTPBridgeToolState, bool) {
if c == nil {
return openAIWSHTTPBridgeToolState{}, false
}
value, ok := c.Get(openAIWSHTTPBridgeToolStateContextKey)
state, typed := value.(openAIWSHTTPBridgeToolState)
return state, ok && typed
}
func setOpenAIWSHTTPBridgeToolState(c *gin.Context, state openAIWSHTTPBridgeToolState) {
if c == nil {
return
}
state.LoweredTools = append(json.RawMessage(nil), state.LoweredTools...)
c.Set(openAIWSHTTPBridgeToolStateContextKey, state)
}
// ResolveOpenAIWSClientFirstMessageTimeout returns the effective client ingress deadline.
func ResolveOpenAIWSClientFirstMessageTimeout(cfg *config.Config) time.Duration {
seconds := config.DefaultOpenAIWSClientFirstMessageTimeoutSeconds
@@ -189,10 +214,29 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
var clientToolMapping apicompat.ResponsesClientToolMapping
if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey {
body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstream(body, "OpenAI WS HTTP bridge")
inheritedState, _ := openAIWSHTTPBridgeToolStateFromContext(c)
toolsPresent := gjson.GetBytes(body, "tools").Exists()
body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body,
"OpenAI WS HTTP bridge",
inheritedState.ClientMapping,
)
if err != nil {
return nil, fmt.Errorf("adapt OpenAI WS HTTP bridge client tools: %w", err)
}
loweredTools := inheritedState.LoweredTools
if toolsPresent {
loweredTools = json.RawMessage(gjson.GetBytes(body, "tools").Raw)
} else if len(loweredTools) > 0 {
body, err = sjson.SetRawBytes(body, "tools", loweredTools)
if err != nil {
return nil, fmt.Errorf("inherit OpenAI WS HTTP bridge tools: %w", err)
}
}
setOpenAIWSHTTPBridgeToolState(c, openAIWSHTTPBridgeToolState{
ClientMapping: clientToolMapping,
LoweredTools: loweredTools,
})
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
@@ -170,6 +170,124 @@ func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyRestoresClientToolsInResponseDone(t *t
require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String())
}
func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t *testing.T) {
gin.SetMode(gin.TestMode)
firstSSEBody := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"resp_custom_first","model":"gpt-5.6-sol","output":[{"type":"function_call","id":"fc_custom_1","call_id":"call_custom_1","name":"exec","arguments":"{\"input\":\"pwd\"}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`,
"",
}, "\n")
secondSSEBody := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"resp_custom_second","model":"gpt-5.6-sol","output":[],"usage":{"input_tokens":1,"output_tokens":1}}}`,
"",
}, "\n")
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(firstSSEBody))},
{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(secondSSEBody))},
}}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
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: 9001, Name: "api-key-custom-followup", Platform: PlatformOpenAI, Type: AccountTypeAPIKey,
Credentials: map[string]any{"api_key": "sk-upstream"}, 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, "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, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
defer cancelWrite()
require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
}
readMessage := func() []byte {
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
defer cancelRead()
messageType, event, readErr := clientConn.Read(readCtx)
require.NoError(t, readErr)
require.Equal(t, coderws.MessageText, messageType)
return event
}
writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"tools":[{"type":"custom","name":"exec"}],"input":"run pwd"}`)
firstEvent := readMessage()
require.Equal(t, "response.completed", gjson.GetBytes(firstEvent, "type").String())
require.Equal(t, "custom_tool_call", gjson.GetBytes(firstEvent, "response.output.0.type").String())
require.Equal(t, "pwd", gjson.GetBytes(firstEvent, "response.output.0.input").String())
writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"previous_response_id":"resp_custom_first","input":[{"type":"custom_tool_call_output","id":"ctco_client_output_1","call_id":"call_custom_1","output":"ok"}]}`)
secondEvent := readMessage()
require.Equal(t, "response.completed", gjson.GetBytes(secondEvent, "type").String())
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)
firstTools := gjson.GetBytes(upstream.bodies[0], "tools").Array()
require.Len(t, firstTools, 1)
require.Equal(t, "function", firstTools[0].Get("type").String())
secondTools := gjson.GetBytes(upstream.bodies[1], "tools").Array()
require.Len(t, secondTools, 1)
require.Equal(t, "function", secondTools[0].Get("type").String())
require.Equal(t, "exec", secondTools[0].Get("name").String())
secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
require.Len(t, secondInput, 3)
require.Equal(t, "run pwd", secondInput[0].String())
require.Equal(t, "function_call", secondInput[1].Get("type").String())
require.Equal(t, "fc_custom_1", secondInput[1].Get("id").String())
require.JSONEq(t, `{"input":"pwd"}`, secondInput[1].Get("arguments").String())
require.False(t, secondInput[1].Get("input").Exists())
require.Equal(t, "function_call_output", secondInput[2].Get("type").String())
require.False(t, secondInput[2].Get("id").Exists())
}
func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
svc := &OpenAIGatewayService{
cfg: &config.Config{