Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 21 additions & 5 deletions backend/internal/service/openai_gateway_grok_tool_protocol.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -97,7 +105,7 @@ func (b *grokResponsesClientToolStreamBody) Close() error {
return sourceErr
}

func newGrokResponsesClientToolStreamBody(
func newResponsesClientToolStreamBody(
source io.ReadCloser,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
Expand All @@ -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,
Expand Down
13 changes: 12 additions & 1 deletion backend/internal/service/openai_ws_http_bridge.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
75 changes: 75 additions & 0 deletions backend/internal/service/openai_ws_http_bridge_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down
Loading