mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
Merge pull request #5764 from hansnow/fix/ws-http-bridge-custom-tools
fix(openai): 补齐 WS HTTP bridge 的客户端工具适配
This commit is contained in:
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user