Merge pull request #5767 from hansnow/fix/ws-http-bridge-custom-tools

fix(openai): 补齐客户端工具终止事件恢复
This commit is contained in:
Wesley Liddick
2026-08-18 16:21:50 +08:00
committed by GitHub
4 changed files with 97 additions and 7 deletions
@@ -430,7 +430,7 @@ func (r *ResponsesClientToolStreamRestorer) RestoreEvent(payload []byte) ([][]by
if err := json.Unmarshal(payload, &wire); err != nil {
return nil, false, err
}
if wire.Type == "response.completed" || wire.Type == "response.incomplete" || wire.Type == "response.failed" {
if isResponsesClientToolTerminalEvent(wire.Type) {
restored, changed, err := RestoreResponsesClientToolPayload(payload, r.adapter)
if err != nil {
return nil, false, err
@@ -465,6 +465,15 @@ func (r *ResponsesClientToolStreamRestorer) RestoreEvent(payload []byte) ([][]by
return result, true, nil
}
func isResponsesClientToolTerminalEvent(typ string) bool {
switch strings.TrimSpace(typ) {
case "response.completed", "response.done", "response.incomplete", "response.failed", "response.cancelled", "response.canceled":
return true
default:
return false
}
}
func (r *ResponsesClientToolStreamRestorer) clientToolEventPayload(payload []byte) bool {
var raw struct {
ItemID string `json:"item_id"`
@@ -174,3 +174,30 @@ func TestResponsesClientToolStreamRestorer_RawEventsPreserveUnknownFieldsAndOutp
require.Len(t, done, 2)
require.Equal(t, "pwd", done[1].Input)
}
func TestResponsesClientToolStreamRestorer_RestoresAllTerminalEvents(t *testing.T) {
for _, eventType := range []string{
"response.completed",
"response.done",
"response.incomplete",
"response.failed",
"response.cancelled",
"response.canceled",
} {
t.Run(eventType, func(t *testing.T) {
restorer := NewResponsesClientToolStreamRestorer(ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}})
payload := []byte(`{"type":"` + eventType + `","sequence_number":7,"response":{"id":"resp_tools","output":[{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}"}]}}`)
restored, changed, err := restorer.RestoreEvent(payload)
require.NoError(t, err)
require.True(t, changed)
require.Len(t, restored, 1)
require.Equal(t, eventType, gjson.GetBytes(restored[0], "type").String())
require.Equal(t, int64(7), gjson.GetBytes(restored[0], "sequence_number").Int())
require.Equal(t, "custom_tool_call", gjson.GetBytes(restored[0], "response.output.0.type").String())
require.Equal(t, "pwd", gjson.GetBytes(restored[0], "response.output.0.input").String())
require.False(t, gjson.GetBytes(restored[0], "response.output.0.arguments").Exists())
})
}
}
@@ -91,12 +91,12 @@ func restoreGrokResponsesClientToolPayload(c *gin.Context, payload []byte) ([]by
return restored, err
}
type grokResponsesClientToolStreamBody struct {
type responsesClientToolStreamBody struct {
*io.PipeReader
source io.Closer
}
func (b *grokResponsesClientToolStreamBody) Close() error {
func (b *responsesClientToolStreamBody) Close() error {
readerErr := b.PipeReader.Close()
sourceErr := b.source.Close()
if readerErr != nil {
@@ -111,8 +111,8 @@ func newResponsesClientToolStreamBody(
maxLineSize int,
) io.ReadCloser {
reader, writer := io.Pipe()
body := &grokResponsesClientToolStreamBody{PipeReader: reader, source: source}
go transformGrokResponsesClientToolStream(source, writer, mapping, maxLineSize)
body := &responsesClientToolStreamBody{PipeReader: reader, source: source}
go transformResponsesClientToolStream(source, writer, mapping, maxLineSize)
return body
}
@@ -124,7 +124,7 @@ func newGrokResponsesClientToolStreamBody(
return newResponsesClientToolStreamBody(source, mapping, maxLineSize)
}
func transformGrokResponsesClientToolStream(
func transformResponsesClientToolStream(
source io.ReadCloser,
destination *io.PipeWriter,
mapping apicompat.ResponsesClientToolMapping,
@@ -208,7 +208,7 @@ func transformGrokResponsesClientToolStream(
payloads, _, err = restorer.RestoreEvent(payload)
if err != nil {
_ = buffered.Flush()
_ = destination.CloseWithError(fmt.Errorf("restore Grok Responses client tool event: %w", err))
_ = destination.CloseWithError(fmt.Errorf("restore Responses client tool event: %w", err))
return
}
}
@@ -116,6 +116,60 @@ func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyAdaptsClientTools(t *testing.T) {
require.Equal(t, "pwd", gjson.GetBytes(result.wsReplayInput[0], "input").String())
}
func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyRestoresClientToolsInResponseDone(t *testing.T) {
gin.SetMode(gin.TestMode)
sse := strings.Join([]string{
`data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","status":"in_progress"}}`,
``,
`data: {"type":"response.function_call_arguments.done","sequence_number":1,"output_index":0,"item_id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}"}`,
``,
`data: {"type":"response.done","sequence_number":2,"response":{"id":"resp_tools","status":"completed","output":[{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}],"usage":{"input_tokens":1,"output_tokens":1}}}`,
``,
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(sse)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}},
httpUpstream: upstream,
}
account := &Account{ID: 5764, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1}
payload := []byte(`{
"type":"response.create","model":"gpt-5","stream":true,
"tools":[{"type":"custom","name":"exec","description":"Run a command"}],
"input":"run pwd"
}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
var events [][]byte
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "test-token", payload, len(payload),
"gpt-5", "", "", "", "", 1,
func(message []byte) error {
events = append(events, append([]byte(nil), message...))
return nil
},
)
require.NoError(t, err)
require.NotNil(t, result)
require.Len(t, events, 4)
terminal := events[len(events)-1]
require.Equal(t, "response.done", gjson.GetBytes(terminal, "type").String())
require.Equal(t, int64(3), gjson.GetBytes(terminal, "sequence_number").Int())
require.Equal(t, "custom_tool_call", gjson.GetBytes(terminal, "response.output.0.type").String())
require.Equal(t, "pwd", gjson.GetBytes(terminal, "response.output.0.input").String())
require.False(t, gjson.GetBytes(terminal, "response.output.0.arguments").Exists())
require.True(t, result.wsReplayInputExists)
require.Len(t, result.wsReplayInput, 1)
require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String())
}
func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
svc := &OpenAIGatewayService{
cfg: &config.Config{