Files
sub2api/backend/internal/service/openai_gateway_passthrough_function_args_test.go
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

259 lines
10 KiB
Go

package service
import (
"bufio"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestHandleStreamingResponsePassthroughDeduplicatesFunctionCallArguments(t *testing.T) {
gin.SetMode(gin.TestMode)
argsA := `{"cmd":"echo hi","meta":{"nested":[1,{"ok":true}],"quote":"a}b"}}`
argsB := `{"path":"/tmp/file","patch":{"ops":[{"op":"replace","value":{"lines":["x","y"]}}]}}`
upstreamBody := strings.Join([]string{
passthroughSSEData(`{"type":"response.created","response":{"id":"resp_passthrough_args","model":"gpt-5.4"}}`),
passthroughSSEData(`{"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","id":"fc_a","call_id":"call_a","name":"exec_command","arguments":"","status":"in_progress"}}`),
passthroughSSEData(functionArgsDeltaJSON(0, "fc_a", "call_a", "exec_command", `{"cmd":`)),
passthroughSSEData(functionArgsDeltaJSON(0, "fc_a", "call_a", "exec_command", `"echo hi","meta":{"nested":[1,{"ok":true}],"quote":"a}b"}}`)),
passthroughSSEData(functionArgsDoneJSON(0, "fc_a", "call_a", "exec_command", argsA+argsA)),
passthroughSSEData(outputItemDoneJSON(0, "fc_a", "call_a", "exec_command", argsA+argsA)),
passthroughSSEData(`{"type":"response.output_item.added","output_index":1,"item":{"type":"function_call","id":"fc_b","call_id":"call_b","name":"apply_patch","arguments":"","status":"in_progress"}}`),
passthroughSSEData(functionArgsDeltaJSON(1, "fc_b", "call_b", "apply_patch", `{"path":"/tmp/file",`)),
passthroughSSEData(functionArgsDeltaJSON(1, "fc_b", "call_b", "apply_patch", `"patch":{"ops":[{"op":"replace","value":{"lines":["x","y"]}}]}}`)),
passthroughSSEData(functionArgsDoneJSON(1, "fc_b", "call_b", "apply_patch", argsB+argsB)),
passthroughSSEData(outputItemDoneJSON(1, "fc_b", "call_b", "apply_patch", argsB+argsB)),
passthroughSSEData(completedWithFunctionCallsJSON(argsA+argsA, argsB+argsB)),
"data: [DONE]\n\n",
}, "")
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(upstreamBody)),
}
svc := &OpenAIGatewayService{}
result, err := svc.handleStreamingResponsePassthrough(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "gpt-5.4", "gpt-5.4")
require.NoError(t, err)
require.NotNil(t, result)
events := collectSSEDataPayloads(t, rec.Body.String())
require.Equal(t, argsA, accumulateFunctionArgumentDeltas(events, "call_a"))
require.Equal(t, argsB, accumulateFunctionArgumentDeltas(events, "call_b"))
require.Equal(t, argsA, gjson.Get(findSSEEvent(t, events, "response.function_call_arguments.done", "call_a"), "arguments").String())
require.Equal(t, argsB, gjson.Get(findSSEEvent(t, events, "response.function_call_arguments.done", "call_b"), "arguments").String())
require.Equal(t, argsA, gjson.Get(findSSEEvent(t, events, "response.output_item.done", "call_a"), "item.arguments").String())
require.Equal(t, argsB, gjson.Get(findSSEEvent(t, events, "response.output_item.done", "call_b"), "item.arguments").String())
completed := findSSEEvent(t, events, "response.completed", "")
require.Equal(t, argsA, gjson.Get(completed, "response.output.0.arguments").String())
require.Equal(t, argsB, gjson.Get(completed, "response.output.1.arguments").String())
requireJSONArgument(t, gjson.Get(completed, "response.output.0.arguments").String())
requireJSONArgument(t, gjson.Get(completed, "response.output.1.arguments").String())
}
func TestForwardResponsesChatCompletionsFallbackKeepsFunctionArgumentsSingle(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"gpt-5.4","input":"run a command","stream":true}`)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(string(body)))
c.Request.Header.Set("Content-Type", "application/json")
upstreamBody := strings.Join([]string{
passthroughSSEData(chatToolCallChunkJSON(true, "")),
"",
passthroughSSEData(chatToolCallChunkJSON(false, `{"cmd":"echo hi"}`)),
"",
`data: {"id":"chatcmpl_tool","object":"chat.completion.chunk","model":"gpt-5.4","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`,
"",
"data: [DONE]",
"",
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_fallback_tool_args"}},
Body: io.NopCloser(strings.NewReader(upstreamBody)),
}}
account := passthroughArgsFallbackAccount()
account.Extra = map[string]any{
openai_compat.ExtraKeyResponsesMode: string(openai_compat.ResponsesSupportModeForceChatCompletions),
}
svc := &OpenAIGatewayService{
cfg: passthroughArgsTestConfig(),
httpUpstream: upstream,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
const wantArgs = `{"cmd":"echo hi"}`
events := collectSSEDataPayloads(t, rec.Body.String())
require.Equal(t, wantArgs, accumulateFunctionArgumentDeltas(events, "chatcmpl-tool-a"))
require.Equal(t, wantArgs, gjson.Get(findSSEEvent(t, events, "response.function_call_arguments.done", "chatcmpl-tool-a"), "arguments").String())
require.Equal(t, wantArgs, gjson.Get(findSSEEvent(t, events, "response.output_item.done", "chatcmpl-tool-a"), "item.arguments").String())
}
func passthroughSSEData(payload string) string {
return "data: " + payload + "\n\n"
}
func functionArgsDeltaJSON(outputIndex int, itemID, callID, name, delta string) string {
return fmt.Sprintf(
`{"type":"response.function_call_arguments.delta","output_index":%d,"item_id":%s,"call_id":%s,"name":%s,"delta":%s}`,
outputIndex,
strconv.Quote(itemID),
strconv.Quote(callID),
strconv.Quote(name),
strconv.Quote(delta),
)
}
func functionArgsDoneJSON(outputIndex int, itemID, callID, name, arguments string) string {
return fmt.Sprintf(
`{"type":"response.function_call_arguments.done","output_index":%d,"item_id":%s,"call_id":%s,"name":%s,"arguments":%s}`,
outputIndex,
strconv.Quote(itemID),
strconv.Quote(callID),
strconv.Quote(name),
strconv.Quote(arguments),
)
}
func outputItemDoneJSON(outputIndex int, itemID, callID, name, arguments string) string {
return fmt.Sprintf(
`{"type":"response.output_item.done","output_index":%d,"item":{"type":"function_call","id":%s,"call_id":%s,"name":%s,"arguments":%s,"status":"completed"}}`,
outputIndex,
strconv.Quote(itemID),
strconv.Quote(callID),
strconv.Quote(name),
strconv.Quote(arguments),
)
}
func completedWithFunctionCallsJSON(argsA, argsB string) string {
return fmt.Sprintf(
`{"type":"response.completed","response":{"id":"resp_passthrough_args","status":"completed","output":[{"type":"function_call","id":"fc_a","call_id":"call_a","name":"exec_command","arguments":%s,"status":"completed"},{"type":"function_call","id":"fc_b","call_id":"call_b","name":"apply_patch","arguments":%s,"status":"completed"}],"usage":{"input_tokens":2,"output_tokens":3,"total_tokens":5}}}`,
strconv.Quote(argsA),
strconv.Quote(argsB),
)
}
func chatToolCallChunkJSON(includeIdentity bool, arguments string) string {
identity := ""
functionFields := make([]string, 0, 2)
if includeIdentity {
identity = `"id":"chatcmpl-tool-a","type":"function",`
functionFields = append(functionFields, `"name":"exec_command"`)
}
if includeIdentity || arguments != "" {
functionFields = append(functionFields, `"arguments":`+strconv.Quote(arguments))
}
return fmt.Sprintf(
`{"id":"chatcmpl_tool","object":"chat.completion.chunk","model":"gpt-5.4","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,%s"function":{%s}}]},"finish_reason":null}]}`,
identity,
strings.Join(functionFields, ","),
)
}
func passthroughArgsTestConfig() *config.Config {
return &config.Config{
Security: config.SecurityConfig{
URLAllowlist: config.URLAllowlistConfig{
Enabled: false,
AllowInsecureHTTP: true,
},
},
}
}
func passthroughArgsFallbackAccount() *Account {
return &Account{
ID: 102,
Name: "passthrough-args-openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "http://upstream.example",
},
}
}
func collectSSEDataPayloads(t *testing.T, body string) []string {
t.Helper()
scanner := bufio.NewScanner(strings.NewReader(body))
var events []string
for scanner.Scan() {
data, ok := extractOpenAISSEDataLine(scanner.Text())
if !ok {
continue
}
if strings.TrimSpace(data) == "[DONE]" {
continue
}
require.True(t, gjson.Valid(data), "invalid SSE data payload: %s", data)
events = append(events, data)
}
require.NoError(t, scanner.Err())
return events
}
func findSSEEvent(t *testing.T, events []string, eventType, callID string) string {
t.Helper()
for _, event := range events {
if gjson.Get(event, "type").String() != eventType {
continue
}
if callID == "" ||
gjson.Get(event, "call_id").String() == callID ||
gjson.Get(event, "item.call_id").String() == callID {
return event
}
}
t.Fatalf("missing event type=%s call_id=%s in %d events", eventType, callID, len(events))
return ""
}
func accumulateFunctionArgumentDeltas(events []string, callID string) string {
var b strings.Builder
for _, event := range events {
if gjson.Get(event, "type").String() != "response.function_call_arguments.delta" {
continue
}
if gjson.Get(event, "call_id").String() != callID {
continue
}
_, _ = b.WriteString(gjson.Get(event, "delta").String())
}
return b.String()
}
func requireJSONArgument(t *testing.T, arguments string) {
t.Helper()
var decoded any
require.NoError(t, json.Unmarshal([]byte(arguments), &decoded))
}