mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:48:43 +08:00
Merge pull request #5661 from wucm667/fix/issue-5659-openai-custom-tools
fix(openai): restore API-key custom tool calls
This commit is contained in:
@@ -21,6 +21,7 @@ import (
|
||||
func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) {
|
||||
beginUpstreamResponseModelObservation(c)
|
||||
clearGrokResponsesClientToolMapping(c)
|
||||
clearOpenAIResponsesClientToolMapping(c)
|
||||
clearOpenAIResponsesNamespaceNames(c)
|
||||
startTime := time.Now()
|
||||
// 固定渠道映射后的请求级 canonical body;账号 normalize/strip 不得改写跨 failover hint。
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -24,6 +25,84 @@ import (
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const openAIResponsesClientToolMappingContextKey = "openai_responses_client_tool_mapping"
|
||||
|
||||
func hasOpenAIResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
|
||||
return len(mapping.CustomTools) > 0 || mapping.ToolSearch || len(mapping.NamespaceTools) > 0
|
||||
}
|
||||
|
||||
func adaptOpenAIResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) {
|
||||
if !needsOpenAIResponsesClientToolAdaptation(body) {
|
||||
return body, apicompat.ResponsesClientToolMapping{}, nil
|
||||
}
|
||||
|
||||
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 OpenAI Responses client tools: %w", err)
|
||||
}
|
||||
var trailingValue any
|
||||
if err := decoder.Decode(&trailingValue); !errors.Is(err, io.EOF) {
|
||||
if err == nil {
|
||||
err = errors.New("multiple JSON values")
|
||||
}
|
||||
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode OpenAI Responses client tools trailing data: %w", err)
|
||||
}
|
||||
mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody)
|
||||
if err != nil || !changed {
|
||||
return body, mapping, err
|
||||
}
|
||||
rebuilt, err := marshalOpenAIUpstreamJSON(requestBody)
|
||||
if err != nil {
|
||||
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode OpenAI Responses client tools: %w", err)
|
||||
}
|
||||
return rebuilt, mapping, nil
|
||||
}
|
||||
|
||||
func needsOpenAIResponsesClientToolAdaptation(body []byte) bool {
|
||||
needsAdaptation := false
|
||||
var visit func(gjson.Result) bool
|
||||
visit = func(value gjson.Result) bool {
|
||||
if value.IsObject() {
|
||||
switch strings.TrimSpace(value.Get("type").String()) {
|
||||
case "custom", "custom_tool_call", "custom_tool_call_output",
|
||||
"tool_search", "tool_search_call", "tool_search_output":
|
||||
needsAdaptation = true
|
||||
return false
|
||||
}
|
||||
}
|
||||
if value.IsObject() || value.IsArray() {
|
||||
value.ForEach(func(_, child gjson.Result) bool {
|
||||
return visit(child)
|
||||
})
|
||||
}
|
||||
return !needsAdaptation
|
||||
}
|
||||
visit(gjson.ParseBytes(body))
|
||||
return needsAdaptation
|
||||
}
|
||||
|
||||
func openAIResponsesClientToolMapping(c *gin.Context) (apicompat.ResponsesClientToolMapping, bool) {
|
||||
if c == nil {
|
||||
return apicompat.ResponsesClientToolMapping{}, false
|
||||
}
|
||||
value, ok := c.Get(openAIResponsesClientToolMappingContextKey)
|
||||
mapping, typed := value.(apicompat.ResponsesClientToolMapping)
|
||||
return mapping, ok && typed && hasOpenAIResponsesClientToolMapping(mapping)
|
||||
}
|
||||
|
||||
// clearOpenAIResponsesClientToolMapping removes mapping state from the prior
|
||||
// forwarding attempt. Forward retries accounts on the same Gin context.
|
||||
func clearOpenAIResponsesClientToolMapping(c *gin.Context) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
if _, exists := c.Get(openAIResponsesClientToolMappingContextKey); exists {
|
||||
c.Set(openAIResponsesClientToolMappingContextKey, apicompat.ResponsesClientToolMapping{})
|
||||
}
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
@@ -104,6 +183,16 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
}
|
||||
}
|
||||
|
||||
if account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey &&
|
||||
!isOpenAIResponsesCompactPath(c) && needsOpenAIResponsesClientToolAdaptation(body) {
|
||||
adaptedBody, mapping, adaptErr := adaptOpenAIResponsesClientTools(body)
|
||||
if adaptErr != nil {
|
||||
return nil, adaptErr
|
||||
}
|
||||
body = adaptedBody
|
||||
c.Set(openAIResponsesClientToolMappingContextKey, mapping)
|
||||
}
|
||||
|
||||
sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -258,6 +347,13 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body, probeBody)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if mapping, ok := openAIResponsesClientToolMapping(c); ok && isEventStreamResponse(resp.Header) {
|
||||
maxLineSize := defaultMaxLineSize
|
||||
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
|
||||
maxLineSize = s.cfg.Gateway.MaxLineSize
|
||||
}
|
||||
resp.Body = newGrokResponsesClientToolStreamBody(resp.Body, mapping, maxLineSize)
|
||||
}
|
||||
|
||||
serviceTier := extractOpenAIServiceTierFromBody(body)
|
||||
|
||||
@@ -1563,6 +1659,12 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough(
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", err)
|
||||
}
|
||||
if mapping, ok := openAIResponsesClientToolMapping(c); ok && json.Valid(body) {
|
||||
body, _, err = apicompat.RestoreResponsesClientToolPayload(body, mapping)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI Responses client tools: %w", err)
|
||||
}
|
||||
}
|
||||
if !writeOpenAICompactSSEBridge(c, resp.StatusCode, body) {
|
||||
c.Data(resp.StatusCode, contentType, body)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
func openAIClientToolsRequest(stream bool) []byte {
|
||||
streamValue := "false"
|
||||
if stream {
|
||||
streamValue = "true"
|
||||
}
|
||||
return []byte(`{"model":"gpt-5.4","input":"fix it","stream":` + streamValue + `,"tools":[{"type":"custom","name":"exec"},{"type":"custom","name":"apply_patch"}]}`)
|
||||
}
|
||||
|
||||
func assertOpenAIClientToolsLowered(t *testing.T, body []byte) {
|
||||
t.Helper()
|
||||
for index, name := range []string{"exec", "apply_patch"} {
|
||||
tool := gjson.GetBytes(body, "tools."+string(rune('0'+index)))
|
||||
require.Equal(t, "function", tool.Get("type").String())
|
||||
require.Equal(t, name, tool.Get("name").String())
|
||||
require.Equal(t, "string", tool.Get("parameters.properties.input.type").String())
|
||||
}
|
||||
}
|
||||
|
||||
func openAIClientToolsTestService(upstream *httpUpstreamRecorder) *OpenAIGatewayService {
|
||||
return &OpenAIGatewayService{
|
||||
httpUpstream: upstream,
|
||||
cfg: &config.Config{Security: config.SecurityConfig{
|
||||
URLAllowlist: config.URLAllowlistConfig{Enabled: false},
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdaptOpenAIResponsesClientToolsLeavesNamespaceOnlyBodyUnchanged(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model": "gpt-5.5",
|
||||
"tools": [{"type": "namespace", "name": "code_tools", "tools": [{"type": "function", "name": "run"}]}],
|
||||
"tool_choice": "auto"
|
||||
}`)
|
||||
|
||||
adapted, mapping, err := adaptOpenAIResponsesClientTools(body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, body, adapted)
|
||||
require.Empty(t, mapping.CustomTools)
|
||||
require.Empty(t, mapping.NamespaceTools)
|
||||
require.False(t, mapping.ToolSearch)
|
||||
}
|
||||
|
||||
func TestAdaptOpenAIResponsesClientToolsRejectsTrailingData(t *testing.T) {
|
||||
tests := map[string][]byte{
|
||||
"trailing garbage": append(openAIClientToolsRequest(false), []byte(` garbage`)...),
|
||||
"second JSON document": append(openAIClientToolsRequest(false), []byte(` {"model":"other"}`)...),
|
||||
}
|
||||
|
||||
for name, body := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
adapted, mapping, err := adaptOpenAIResponsesClientTools(body)
|
||||
|
||||
require.ErrorContains(t, err, "decode OpenAI Responses client tools trailing data")
|
||||
require.Equal(t, body, adapted)
|
||||
require.Empty(t, mapping)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClearOpenAIResponsesClientToolMappingRemovesStaleContextState(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
c.Set(openAIResponsesClientToolMappingContextKey, apicompat.ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}})
|
||||
|
||||
clearOpenAIResponsesClientToolMapping(c)
|
||||
|
||||
_, ok := openAIResponsesClientToolMapping(c)
|
||||
require.False(t, ok)
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyRestoresClientToolsNonStreaming(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := openAIClientToolsRequest(false)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_tools","status":"completed","output":[
|
||||
{"type":"function_call","id":"i1","call_id":"c1","name":"exec","arguments":"{\"input\":\"pwd\"}"},
|
||||
{"type":"function_call","id":"i2","call_id":"c2","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}],"usage":{}}`)),
|
||||
}}
|
||||
svc := openAIClientToolsTestService(upstream)
|
||||
account := &Account{ID: 5659, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}}
|
||||
|
||||
result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, false, time.Now())
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assertOpenAIClientToolsLowered(t, upstream.lastBody)
|
||||
require.Equal(t, "custom_tool_call", gjson.Get(recorder.Body.String(), "output.0.type").String())
|
||||
require.Equal(t, "pwd", gjson.Get(recorder.Body.String(), "output.0.input").String())
|
||||
require.Equal(t, "custom_tool_call", gjson.Get(recorder.Body.String(), "output.1.type").String())
|
||||
require.Equal(t, "*** Begin Patch", gjson.Get(recorder.Body.String(), "output.1.input").String())
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyRestoresClientToolsStreaming(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := openAIClientToolsRequest(true)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
|
||||
sse := strings.Join([]string{
|
||||
`data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","status":"in_progress"}}`,
|
||||
`data: {"type":"response.function_call_arguments.done","sequence_number":1,"item_id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}`,
|
||||
`data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}","status":"completed"}}`,
|
||||
`data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_stream_tools","status":"completed","output":[{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}],"usage":{"input_tokens":1,"output_tokens":1}}}`,
|
||||
}, "\n\n") + "\n\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 := openAIClientToolsTestService(upstream)
|
||||
account := &Account{ID: 5660, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}}
|
||||
|
||||
result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, true, time.Now())
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assertOpenAIClientToolsLowered(t, upstream.lastBody)
|
||||
output := recorder.Body.String()
|
||||
require.Contains(t, output, `"type":"custom_tool_call"`)
|
||||
require.Contains(t, output, `"type":"response.custom_tool_call_input.done"`)
|
||||
require.Contains(t, output, `"input":"*** Begin Patch"`)
|
||||
require.NotContains(t, output, `"input":{`)
|
||||
}
|
||||
Reference in New Issue
Block a user