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

fix(openai): 补齐 WS HTTP bridge 的客户端工具适配
This commit is contained in:
Wesley Liddick
2026-08-18 15:47:59 +08:00
committed by GitHub
3 changed files with 108 additions and 6 deletions
@@ -15,12 +15,12 @@ import (
const grokResponsesClientToolMappingContextKey = "grok_responses_client_tool_mapping"
func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) {
func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) ([]byte, apicompat.ResponsesClientToolMapping, error) {
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
var requestBody map[string]any
if err := decoder.Decode(&requestBody); err != nil {
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode Grok Responses client tools: %w", err)
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode %s Responses client tools: %w", upstream, err)
}
mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody)
@@ -32,15 +32,23 @@ func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClie
}
rebuilt, err := marshalOpenAIUpstreamJSON(requestBody)
if err != nil {
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode Grok Responses client tools: %w", err)
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode %s Responses client tools: %w", upstream, err)
}
return rebuilt, mapping, nil
}
func hasGrokResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) {
return adaptResponsesClientToolsForFunctionUpstream(body, "Grok")
}
func hasResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
return len(mapping.CustomTools) > 0 || mapping.ToolSearch || len(mapping.NamespaceTools) > 0
}
func hasGrokResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
return hasResponsesClientToolMapping(mapping)
}
func setGrokResponsesClientToolMapping(c *gin.Context, mapping apicompat.ResponsesClientToolMapping) {
if c == nil {
return
@@ -97,7 +105,7 @@ func (b *grokResponsesClientToolStreamBody) Close() error {
return sourceErr
}
func newGrokResponsesClientToolStreamBody(
func newResponsesClientToolStreamBody(
source io.ReadCloser,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
@@ -108,6 +116,14 @@ func newGrokResponsesClientToolStreamBody(
return body
}
func newGrokResponsesClientToolStreamBody(
source io.ReadCloser,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
) io.ReadCloser {
return newResponsesClientToolStreamBody(source, mapping, maxLineSize)
}
func transformGrokResponsesClientToolStream(
source io.ReadCloser,
destination *io.PipeWriter,
@@ -12,6 +12,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
@@ -186,6 +187,13 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
if err != nil {
return nil, fmt.Errorf("prepare http bridge body: %w", err)
}
var clientToolMapping apicompat.ResponsesClientToolMapping
if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey {
body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstream(body, "OpenAI WS HTTP bridge")
if err != nil {
return nil, fmt.Errorf("adapt OpenAI WS HTTP bridge client tools: %w", err)
}
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
var upstreamReq *http.Request
@@ -329,11 +337,14 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
return result
}
scanner := bufio.NewScanner(resp.Body)
maxLineSize := defaultMaxLineSize
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
maxLineSize = s.cfg.Gateway.MaxLineSize
}
if hasResponsesClientToolMapping(clientToolMapping) {
resp.Body = newResponsesClientToolStreamBody(resp.Body, clientToolMapping, maxLineSize)
}
scanner := bufio.NewScanner(resp.Body)
scanBuf := getSSEScannerBuf64K()
scanner.Buffer(scanBuf[:0], maxLineSize)
defer putSSEScannerBuf64K(scanBuf)
@@ -41,6 +41,81 @@ func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) {
require.Equal(t, "hi", gjson.GetBytes(body, "input").String())
}
func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyAdaptsClientTools(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.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}}`,
``,
`data: {"type":"response.completed","sequence_number":3,"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: 5659, 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":[
{"type":"custom_tool_call","id":"previous_item","call_id":"previous_call","name":"exec","input":"echo ready"},
{"type":"custom_tool_call_output","call_id":"previous_call","output":"ready"}
]
}`)
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", "", "", "", "", 2,
func(message []byte) error {
events = append(events, append([]byte(nil), message...))
return nil
},
)
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, "function", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
require.Equal(t, "function_call", gjson.GetBytes(upstream.lastBody, "input.0.type").String())
require.JSONEq(t, `{"input":"echo ready"}`, gjson.GetBytes(upstream.lastBody, "input.0.arguments").String())
require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.input").Exists())
require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.1.type").String())
var outputDone, completed []byte
for _, event := range events {
switch gjson.GetBytes(event, "type").String() {
case "response.output_item.done":
outputDone = event
case "response.completed":
completed = event
}
}
require.NotEmpty(t, outputDone)
require.Equal(t, "custom_tool_call", gjson.GetBytes(outputDone, "item.type").String())
require.Equal(t, "pwd", gjson.GetBytes(outputDone, "item.input").String())
require.False(t, gjson.GetBytes(outputDone, "item.arguments").Exists())
require.NotEmpty(t, completed)
require.Equal(t, "custom_tool_call", gjson.GetBytes(completed, "response.output.0.type").String())
require.Equal(t, "pwd", gjson.GetBytes(completed, "response.output.0.input").String())
require.True(t, result.wsReplayInputExists)
require.Len(t, result.wsReplayInput, 1)
require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String())
require.Equal(t, "pwd", gjson.GetBytes(result.wsReplayInput[0], "input").String())
}
func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
svc := &OpenAIGatewayService{
cfg: &config.Config{