Files
sub2api/backend/internal/service/openai_ws_http_bridge_resume_test.go
T
李建琦 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

253 lines
9.9 KiB
Go

package service
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestBuildOpenAIWSCurrentTurnRetryPayloadRejectsOrphanToolOutput(t *testing.T) {
payload := []byte(`{"type":"response.create","model":"mapped-model","previous_response_id":"resp_old"}`)
fullInput := []json.RawMessage{
json.RawMessage(`{"type":"function_call_output","call_id":"missing_call","output":"done"}`),
}
retryPayload, retrySafe, err := buildOpenAIWSCurrentTurnRetryPayload(payload, fullInput, true, "gpt-5.6-sol")
require.NoError(t, err)
require.False(t, retrySafe)
require.Nil(t, retryPayload)
}
func TestProxyOpenAIWSHTTPBridgeTurnLaterTurn429FailsOverBeforeClientWrite(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusTooManyRequests,
Header: http.Header{"Retry-After": []string{"60"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"usage_limit_reached","message":"The usage limit has been reached"}}`)),
}}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{ID: 129, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1}
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
payload := []byte(`{"type":"response.create","model":"gpt-5.6-sol","previous_response_id":"resp_old","input":[{"role":"user","content":"continue"}]}`)
writes := 0
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "access-token", payload, len(payload),
"gpt-5.6-sol", "", "", "", "", 281,
func([]byte) error {
writes++
return nil
},
)
require.Nil(t, result)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode)
require.Zero(t, writes)
}
func TestProxyOpenAIWSHTTPBridgeTurnLaterTurnDoesNotFailOverAfterDownstreamOutput(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\"}\n\n" +
"data: {\"type\":\"error\",\"error\":{\"type\":\"rate_limit_error\",\"message\":\"limited\"}}\n\n",
)),
}}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{ID: 10, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1}
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`)
var writes [][]byte
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "sk-test", payload, len(payload),
"gpt-5", "", "", "", "", 281,
func(message []byte) error {
writes = append(writes, append([]byte(nil), message...))
return nil
},
)
require.NotNil(t, result)
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
require.Len(t, writes, 2)
require.Equal(t, "response.output_text.delta", gjson.GetBytes(writes[0], "type").String())
require.Equal(t, "error", gjson.GetBytes(writes[1], "type").String())
}
func TestOpenAIWSHTTPBridgeLaterTurn429RetriesCurrentTurnOnReplacementAccount(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.OAuthEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"text/event-stream"},
openAIWSTurnStateHeader: []string{"old-account-state"},
},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_first\",\"output\":[{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"first-ok\"}]},{\"id\":\"fc_1\",\"type\":\"function_call\",\"call_id\":\"call_1\",\"name\":\"inspect\",\"arguments\":\"{}\"}],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n",
)),
},
{
StatusCode: http.StatusTooManyRequests,
Header: http.Header{"Retry-After": []string{"60"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"usage_limit_reached","message":"The usage limit has been reached"}}`)),
},
{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_second\",\"output\":[{\"id\":\"msg_2\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"second-ok\"}]}],\"usage\":{\"input_tokens\":4,\"output_tokens\":1}}}\n\n",
)),
},
}}
svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: upstream,
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
}
account := &Account{
ID: 129, Name: "limited", Platform: PlatformOpenAI, Type: AccountTypeOAuth,
Status: StatusActive, Schedulable: true, Concurrency: 1,
Extra: map[string]any{"openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModeHTTPBridge},
}
nextAccount := *account
nextAccount.ID = 130
nextAccount.Name = "replacement"
serverErrCh := make(chan error, 1)
failoverCh := make(chan []byte, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := websocket.Accept(w, r, nil)
if err != nil {
serverErrCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
ginCtx.Request = r.Clone(r.Context())
readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
_, firstMessage, readErr := conn.Read(readCtx)
cancel()
if readErr != nil {
serverErrCh <- readErr
return
}
proxyErr := svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token-a", firstMessage, nil)
var failoverErr *UpstreamFailoverError
if !errors.As(proxyErr, &failoverErr) {
serverErrCh <- proxyErr
return
}
retryPayload, retryCurrentTurn := OpenAIWSCurrentTurnRetryPayload(proxyErr)
if !retryCurrentTurn || len(retryPayload) == 0 {
serverErrCh <- errors.New("missing current-turn retry payload")
return
}
failoverCh <- retryPayload
serverErrCh <- svc.ProxyResponsesWebSocketFromClient(
r.Context(), ginCtx, conn, &nextAccount, "access-token-b", retryPayload, nil,
)
}))
defer wsServer.Close()
dialCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := websocket.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancel()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
writeCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, websocket.MessageText, []byte(`{"type":"response.create","model":"gpt-5.6-sol","input":[{"role":"user","content":"first"}]}`))
cancel()
require.NoError(t, err)
readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
_, completed, err := clientConn.Read(readCtx)
cancel()
require.NoError(t, err)
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
writeCtx, cancel = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, websocket.MessageText, []byte(`{"type":"response.create","model":"gpt-5.6-sol","previous_response_id":"resp_first","input":[{"type":"function_call_output","call_id":"call_1","output":"second"}]}`))
cancel()
require.NoError(t, err)
readCtx, cancel = context.WithTimeout(context.Background(), 3*time.Second)
_, retriedCompleted, err := clientConn.Read(readCtx)
cancel()
require.NoError(t, err)
require.Equal(t, "response.completed", gjson.GetBytes(retriedCompleted, "type").String())
require.Equal(t, "resp_second", gjson.GetBytes(retriedCompleted, "response.id").String())
_ = clientConn.Close(websocket.StatusNormalClosure, "done")
select {
case retryPayload := <-failoverCh:
require.NotEmpty(t, retryPayload)
require.False(t, gjson.GetBytes(retryPayload, "previous_response_id").Exists())
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(retryPayload, "model").String())
input := gjson.GetBytes(retryPayload, "input")
require.True(t, input.IsArray())
require.Len(t, input.Array(), 4)
require.Contains(t, input.Raw, "first")
require.Contains(t, input.Raw, "first-ok")
require.Contains(t, input.Raw, "second")
require.Equal(t, 1, strings.Count(input.Raw, `"id":"fc_1"`))
require.Equal(t, 2, strings.Count(input.Raw, `"call_id":"call_1"`))
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for current-turn failover")
}
select {
case proxyErr := <-serverErrCh:
require.NoError(t, proxyErr)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for replacement-account completion")
}
require.Len(t, upstream.bodies, 3)
require.Contains(t, string(upstream.bodies[0]), "first")
require.NotContains(t, string(upstream.bodies[2]), "previous_response_id")
require.Contains(t, string(upstream.bodies[2]), "second")
require.Empty(t, upstream.requests[2].Header.Get(openAIWSTurnStateHeader))
}