mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:08:02 +08:00
Merge pull request #5912 from xuhaihan/fix/deepseek-responses-client-tools
fix(deepseek): adapt Codex client tools for Responses
This commit is contained in:
@@ -115,12 +115,22 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
}
|
||||
}
|
||||
|
||||
nativeDeepSeekResponses := account.Platform == PlatformDeepseek &&
|
||||
(account.GetAPIProtocol() == APIProtocolResponses || account.IsAdaptiveAPIProtocol())
|
||||
if nativeDeepSeekResponses && account.Type == AccountTypeAPIKey && !compactPath &&
|
||||
needsOpenAIResponsesClientToolAdaptation(body) {
|
||||
adaptedBody, mapping, adaptErr := adaptOpenAIResponsesClientTools(body)
|
||||
if adaptErr != nil {
|
||||
return nil, fmt.Errorf("adapt DeepSeek Responses client tools: %w", adaptErr)
|
||||
}
|
||||
body = adaptedBody
|
||||
setOpenAIResponsesClientToolMapping(c, mapping)
|
||||
}
|
||||
|
||||
originalBody := body
|
||||
requestView := newOpenAIRequestView(body)
|
||||
reqModel, reqStream, promptCacheKey := requestView.Model, requestView.Stream, requestView.PromptCacheKey
|
||||
originalModel := reqModel
|
||||
nativeDeepSeekResponses := account.Platform == PlatformDeepseek &&
|
||||
(account.GetAPIProtocol() == APIProtocolResponses || account.IsAdaptiveAPIProtocol())
|
||||
|
||||
if account.Platform == PlatformGrok {
|
||||
return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime)
|
||||
@@ -1054,6 +1064,14 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
}
|
||||
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 = newResponsesClientToolStreamBody(resp.Body, mapping, maxLineSize)
|
||||
}
|
||||
|
||||
serviceTier := extractOpenAIServiceTierFromBody(body)
|
||||
// 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。
|
||||
reqBody = nil
|
||||
|
||||
@@ -92,6 +92,17 @@ func openAIResponsesClientToolMapping(c *gin.Context) (apicompat.ResponsesClient
|
||||
return mapping, ok && typed && hasOpenAIResponsesClientToolMapping(mapping)
|
||||
}
|
||||
|
||||
func setOpenAIResponsesClientToolMapping(c *gin.Context, mapping apicompat.ResponsesClientToolMapping) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
if !hasOpenAIResponsesClientToolMapping(mapping) {
|
||||
clearOpenAIResponsesClientToolMapping(c)
|
||||
return
|
||||
}
|
||||
c.Set(openAIResponsesClientToolMappingContextKey, mapping)
|
||||
}
|
||||
|
||||
// clearOpenAIResponsesClientToolMapping removes mapping state from the prior
|
||||
// forwarding attempt. Forward retries accounts on the same Gin context.
|
||||
func clearOpenAIResponsesClientToolMapping(c *gin.Context) {
|
||||
@@ -103,6 +114,15 @@ func clearOpenAIResponsesClientToolMapping(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
func restoreOpenAIResponsesClientToolPayload(c *gin.Context, payload []byte) ([]byte, error) {
|
||||
mapping, ok := openAIResponsesClientToolMapping(c)
|
||||
if !ok || !bytes.Contains(payload, []byte(`"function_call"`)) || !json.Valid(payload) {
|
||||
return payload, nil
|
||||
}
|
||||
restored, _, err := apicompat.RestoreResponsesClientToolPayload(payload, mapping)
|
||||
return restored, err
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
@@ -210,7 +230,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
return nil, adaptErr
|
||||
}
|
||||
body = adaptedBody
|
||||
c.Set(openAIResponsesClientToolMappingContextKey, mapping)
|
||||
setOpenAIResponsesClientToolMapping(c, mapping)
|
||||
}
|
||||
|
||||
sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body)
|
||||
@@ -2099,11 +2119,9 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough(
|
||||
return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", err)
|
||||
}
|
||||
body = restoreCodexToolNamesFromContext(c, body)
|
||||
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)
|
||||
}
|
||||
body, err = restoreOpenAIResponsesClientToolPayload(c, body)
|
||||
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)
|
||||
|
||||
@@ -1614,6 +1614,10 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("restore Grok Responses client tool response: %w", err)
|
||||
}
|
||||
body, err = restoreOpenAIResponsesClientToolPayload(c, body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI Responses client tool response: %w", err)
|
||||
}
|
||||
body, err = restoreOpenAIResponsesNamespacePayload(c, body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI namespace response: %w", err)
|
||||
@@ -1707,6 +1711,10 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte
|
||||
if restoreErr != nil {
|
||||
return nil, fmt.Errorf("restore Grok Responses client tool response: %w", restoreErr)
|
||||
}
|
||||
restoredBody, restoreErr = restoreOpenAIResponsesClientToolPayload(c, restoredBody)
|
||||
if restoreErr != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI Responses client tool response: %w", restoreErr)
|
||||
}
|
||||
restoredBody, restoreErr = restoreOpenAIResponsesNamespacePayload(c, restoredBody)
|
||||
if restoreErr != nil {
|
||||
return nil, fmt.Errorf("restore OpenAI namespace response: %w", restoreErr)
|
||||
|
||||
@@ -108,6 +108,120 @@ func TestClearOpenAIResponsesClientToolMappingRemovesStaleContextState(t *testin
|
||||
require.False(t, ok)
|
||||
}
|
||||
|
||||
func TestDeepSeekResponsesForwardRestoresClientToolsStreaming(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":"exec","status":"in_progress"}}`,
|
||||
`data: {"type":"response.function_call_arguments.done","sequence_number":1,"item_id":"i1","call_id":"c1","name":"exec","arguments":"{\"input\":\"pwd\"}"}`,
|
||||
`data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"i1","call_id":"c1","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}}`,
|
||||
`data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_ds_tools","status":"completed","output":[{"type":"function_call","id":"i1","call_id":"c1","name":"exec","arguments":"{\"input\":\"pwd\"}"}],"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: 5661,
|
||||
Platform: PlatformDeepseek,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "test-key",
|
||||
"api_protocol": APIProtocolResponses,
|
||||
"base_url": "https://relay.example",
|
||||
},
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assertOpenAIClientToolsLowered(t, upstream.lastBody)
|
||||
require.Equal(t, "/responses", upstream.lastReq.URL.Path)
|
||||
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":"pwd"`)
|
||||
}
|
||||
|
||||
func TestDeepSeekAdaptiveResponsesForwardRestoresClientToolsNonStreaming(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_ds_adaptive_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":{"input_tokens":1,"output_tokens":1}}`)),
|
||||
}}
|
||||
svc := openAIClientToolsTestService(upstream)
|
||||
account := &Account{
|
||||
ID: 5662,
|
||||
Platform: PlatformDeepseek,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "test-key",
|
||||
"api_protocol": APIProtocolAdaptive,
|
||||
"api_base_urls": map[string]any{
|
||||
APIProtocolResponses: "https://relay.example",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assertOpenAIClientToolsLowered(t, upstream.lastBody)
|
||||
require.Equal(t, "/responses", upstream.lastReq.URL.Path)
|
||||
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 TestDeepSeekResponsesCompactSkipsClientToolAdaptation(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/compact", 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_compact","status":"completed","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`)),
|
||||
}}
|
||||
svc := openAIClientToolsTestService(upstream)
|
||||
account := &Account{
|
||||
ID: 5663,
|
||||
Platform: PlatformDeepseek,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "test-key",
|
||||
"api_protocol": APIProtocolResponses,
|
||||
"base_url": "https://relay.example",
|
||||
},
|
||||
}
|
||||
|
||||
_, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "custom", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
|
||||
require.Equal(t, "/responses/compact", upstream.lastReq.URL.Path)
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyRestoresClientToolsNonStreaming(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := openAIClientToolsRequest(false)
|
||||
|
||||
Reference in New Issue
Block a user