Files
sub2api/backend/internal/service/openai_capacity_shed_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

373 lines
16 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// --- mock: 只记录临时不可调度写入,其余方法不应被调用 ---
type capacityShedAccountRepoStub struct {
AccountRepository // 嵌入接口,未实现的方法会 panic(不应被调用)
tempUnschedCalls int
}
func (r *capacityShedAccountRepoStub) SetTempUnschedulable(_ context.Context, _ int64, _ time.Time, _ string) error {
r.tempUnschedCalls++
return nil
}
// 上游容量降载是请求级信号:故障因素(客户端身份、模型容量)与账号无关,
// 同账号重试用尽后不得把账号临时摘掉——否则一个被降载的请求会顺着 failover
// 把整池账号逐个封禁,而每个账号都会以同一个错误失败。
func TestTempUnscheduleRetryableErrorSkipsRequestScopedTransient(t *testing.T) {
t.Run("请求级瞬时故障不写账号状态", func(t *testing.T) {
repo := &capacityShedAccountRepoStub{}
svc := &GatewayService{accountRepo: repo}
svc.TempUnscheduleRetryableError(context.Background(), 1, &UpstreamFailoverError{
StatusCode: http.StatusBadGateway,
RetryableOnSameAccount: true,
RequestScopedTransient: true,
})
require.Zero(t, repo.tempUnschedCalls)
})
// 对照组:同样的 502 在未标记请求级瞬时故障时仍按原有语义临时摘号,
// 确认上面的断言来自新增守卫而非其他前置条件。
t.Run("未标记时保持原有临时摘号语义", func(t *testing.T) {
repo := &capacityShedAccountRepoStub{}
svc := &GatewayService{accountRepo: repo}
svc.TempUnscheduleRetryableError(context.Background(), 1, &UpstreamFailoverError{
StatusCode: http.StatusBadGateway,
RetryableOnSameAccount: true,
})
require.Equal(t, 1, repo.tempUnschedCalls)
})
}
// 非池模式账号同样要先在同账号重试:换号不改变降载因素。
func TestStreamFailedEventCapacityShedRetriesOnSameAccount(t *testing.T) {
nonPool := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
for _, code := range []string{"server_is_overloaded", "slow_down"} {
payload := []byte(`{"type":"response.failed","response":{"error":{"code":"` + code + `"}}}`)
require.True(t, isOpenAIUpstreamCapacityShedEvent(payload), code)
require.True(t, openAIStreamFailedEventRetryableOnSameAccount(nonPool, payload, "overloaded"), code)
}
// 非降载的 failed 事件在非池模式下仍不做同账号重试,避免放大改动面。
other := []byte(`{"type":"response.failed","response":{"error":{"code":"server_error"}}}`)
require.False(t, isOpenAIUpstreamCapacityShedEvent(other))
require.False(t, openAIStreamFailedEventRetryableOnSameAccount(nonPool, other, "boom"))
}
func TestOpenAIHTTPCapacityShedIsRequestScopedForOAuthAccounts(t *testing.T) {
payload := []byte(`{"error":{"type":"server_error","message":"Our servers are currently overloaded. Please try again later."}}`)
failoverErr := newOpenAIUpstreamFailoverError(
http.StatusBadRequest,
http.Header{"X-Request-Id": []string{"rid-http-capacity"}},
payload,
"Our servers are currently overloaded. Please try again later.",
false,
)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
repo := &capacityShedAccountRepoStub{}
(&GatewayService{accountRepo: repo}).TempUnscheduleRetryableError(context.Background(), 1, failoverErr)
require.Zero(t, repo.tempUnschedCalls)
rateLimitService := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
gateway := &OpenAIGatewayService{rateLimitService: rateLimitService}
account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
require.False(t, gateway.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusBadRequest,
nil,
payload,
"gpt-5",
))
require.Zero(t, repo.tempUnschedCalls)
}
// 上游降载的真实序列是「event: error → event: response.failed」。error 帧不算
// 客户端输出:若把它当首输出 flushclientOutputStarted 被固化,随后的 failed
// 事件就进不了 pre-output failover 分支,只能把致命错误原样转发给客户端。
func TestOpenAIStreamErrorFrameDoesNotStartClientOutput(t *testing.T) {
cases := []struct {
data string
eventType string
want bool
}{
{`{"type":"error","error":{"code":"server_is_overloaded","message":"overloaded"}}`, "error", false},
{`{"type":"error","error":{"code":"slow_down","message":"slow down"}}`, "error", false},
{`{"type":"error","error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"limited"}}`, "error", false},
// 不可重试类错误帧维持原样转发(不进 failover),保留上游错误细节。
{`{"type":"error","error":{"type":"invalid_request_error","code":"content_policy_violation","message":"blocked"}}`, "error", true},
{`{"type":"response.failed","response":{"error":{"code":"server_is_overloaded"}}}`, "response.failed", false},
{`{"type":"response.created","response":{"id":"resp_1"}}`, "response.created", false},
{`{"type":"response.in_progress","response":{"id":"resp_1"}}`, "response.in_progress", false},
{`{"type":"response.output_item.added","item":{"type":"reasoning","summary":[]}}`, "response.output_item.added", false},
{`{"type":"response.output_item.added","item":{"type":"reasoning","encrypted_content":"ciphertext"}}`, "response.output_item.added", true},
{`{"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":""}}`, "response.reasoning_summary_part.added", false},
{`{"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":"thinking"}}`, "response.reasoning_summary_part.added", true},
{`{"type":"response.content_part.added","part":{"type":"output_text","text":""}}`, "response.content_part.added", false},
{`{"type":"response.output_text.delta","delta":"hi"}`, "response.output_text.delta", true},
{`[DONE]`, "", true},
}
for _, tc := range cases {
require.Equal(t, tc.want, openAIStreamDataStartsClientOutput(tc.data, tc.eventType), "data=%s type=%s", tc.data, tc.eventType)
}
}
func TestOpenAIStreamMetadataPreambleAndMessageOnlyOverloadFailOver(t *testing.T) {
gin.SetMode(gin.TestMode)
largeMetadata := strings.Repeat("x", 16*1024)
stream := strings.Join([]string{
"event: response.created",
`data: {"type":"response.created","response":{"id":"resp_1","metadata":{"padding":"` + largeMetadata + `"}}}`,
"",
"event: response.output_item.added",
`data: {"type":"response.output_item.added","item":{"type":"reasoning","summary":[]}}`,
"",
"event: response.reasoning_summary_part.added",
`data: {"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":""}}`,
"",
"event: error",
`data: {"type":"error","error":{"type":"service_unavailable_error","message":"Our servers are currently overloaded. Please try again later."}}`,
"",
}, "\n")
tests := []struct {
name string
run func(*OpenAIGatewayService, *gin.Context, *http.Response, *Account) error
}{
{
name: "native",
run: func(svc *OpenAIGatewayService, c *gin.Context, resp *http.Response, account *Account) error {
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, account, time.Now(), "model", "model")
return err
},
},
{
name: "passthrough",
run: func(svc *OpenAIGatewayService, c *gin.Context, resp *http.Response, account *Account) error {
_, err := svc.handleStreamingResponsePassthrough(c.Request.Context(), resp, c, account, time.Now(), "model", "model")
return err
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(stream)),
Header: http.Header{"X-Request-Id": []string{"rid-message-only-overload"}},
}
account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}
err := tt.run(svc, c, resp, account)
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
require.False(t, c.Writer.Written())
require.Empty(t, rec.Body.String())
})
}
}
// 回归用例(真实上游降载序列):created → in_progress → error 帧 → response.failed。
// 期望仍然走 pre-output failover(同账号重试 + 请求级瞬时标记),且不向客户端写出任何字节。
func TestOpenAIStreamCapacityShedErrorFramePrecedingFailedStillFailsOver(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{
Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize},
}
svc := &OpenAIGatewayService{cfg: cfg}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
"event: response.created",
`data: {"type":"response.created","response":{"id":"resp_1"},"sequence_number":0}`,
"",
"event: response.in_progress",
`data: {"type":"response.in_progress","response":{"id":"resp_1"},"sequence_number":1}`,
"",
"event: error",
`data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`,
"",
"event: response.failed",
`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`,
"",
}, "\n"))),
Header: http.Header{"X-Request-Id": []string{"rid-shed-error-then-failed"}},
}
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model")
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
require.False(t, c.Writer.Written())
require.Empty(t, rec.Body.String())
}
// 流中途(已有真实输出)降载时无法再 failover,此时必须把降载码改写为客户端
// 可重试的 server_error 再转发——Codex 对 server_is_overloaded/slow_down 判致命
// 并终止会话,对其余错误码执行内置退避重试。消息原样保留。
func TestOpenAIStreamCapacityShedAfterOutputRewritesCodeForClient(t *testing.T) {
gin.SetMode(gin.TestMode)
logSink, restore := captureStructuredLog(t)
defer restore()
cfg := &config.Config{
Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize},
}
svc := &OpenAIGatewayService{cfg: cfg}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
"event: response.created",
`data: {"type":"response.created","response":{"id":"resp_1"}}`,
"",
"event: response.output_text.delta",
`data: {"type":"response.output_text.delta","delta":"partial"}`,
"",
"event: error",
`data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`,
"",
"event: response.failed",
`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`,
"",
}, "\n"))),
Header: http.Header{"X-Request-Id": []string{"rid-shed-after-output"}},
}
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model")
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
body := rec.Body.String()
require.Contains(t, body, "partial")
require.Contains(t, body, "event: response.failed")
require.Contains(t, body, `"code":"server_error"`)
require.NotContains(t, body, "server_is_overloaded")
require.Contains(t, body, "Our servers are currently overloaded")
require.True(t, logSink.ContainsMessage("gateway.failover_suppressed_after_semantic_output"))
require.True(t, logSink.ContainsFieldValue("path", "native_sse"))
require.True(t, logSink.ContainsFieldValue("upstream_request_id", "rid-shed-after-output"))
}
// helper 单测:只有降载码被改写,其余错误码(尤其 rate_limit_exceeded,客户端
// 依赖其原码解析重试延时)必须原样保留。
func TestSanitizeOpenAICapacityShedErrorCodeForClient(t *testing.T) {
cases := []struct {
name string
payload string
wantChanged bool
wantContain string
}{
{
name: "failed事件嵌套code改写",
payload: `{"type":"response.failed","response":{"error":{"code":"server_is_overloaded","message":"overloaded"}}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "error帧裸code改写",
payload: `{"type":"error","error":{"code":"slow_down","message":"slow down"}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "failed事件只有过载文案时补充code",
payload: `{"type":"response.failed","response":{"error":{"message":"Our servers are currently overloaded. Please try again later."}}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "error帧只有过载文案时补充code",
payload: `{"type":"error","error":{"message":"Server is overloaded. Please try again later."}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "rate_limit不改写",
payload: `{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"try again in 3s"}}}`,
wantChanged: false,
wantContain: `"code":"rate_limit_exceeded"`,
},
{
name: "普通server_error不改写",
payload: `{"type":"response.failed","response":{"error":{"code":"server_error","message":"boom"}}}`,
wantChanged: false,
wantContain: `"code":"server_error"`,
},
{
name: "非JSON不改写",
payload: `not-json`,
wantChanged: false,
wantContain: `not-json`,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
out, changed := sanitizeOpenAICapacityShedErrorCodeForClient([]byte(tc.payload))
require.Equal(t, tc.wantChanged, changed)
require.Contains(t, string(out), tc.wantContain)
if changed {
require.NotContains(t, string(out), "server_is_overloaded")
require.NotContains(t, string(out), "slow_down")
}
})
}
}
// 出站身份的版本声明只能有一个来源:UA 的版本段、version 头、探针版本三处必须同源,
// 各自硬编码会漂移成互相矛盾的身份,而自相矛盾或陈旧的身份会被上游优先降载。
func TestCodexOutboundVersionHasSingleSource(t *testing.T) {
require.True(t,
strings.HasPrefix(codexCLIUserAgent, openai.CodexDefaultOriginator+"/"+codexCLIVersion+" "),
"codexCLIUserAgent=%q 必须以 codexCLIVersion=%q 作为版本段", codexCLIUserAgent, codexCLIVersion,
)
require.GreaterOrEqual(t, CompareVersions(codexCLIVersion, codexUpstreamMinVersion), 0,
"codexCLIVersion=%q 不得低于上游最低门槛 %q", codexCLIVersion, codexUpstreamMinVersion,
)
}