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 帧不算 // 客户端输出:若把它当首输出 flush,clientOutputStarted 被固化,随后的 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, ) }