Files
sub2api/backend/internal/service/openai_capacity_shed_test.go
T

373 lines
16 KiB
Go
Raw Normal View History

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,
)
}