From 0dce07ee8b189c1c8ce9f90e636c50e0c638a170 Mon Sep 17 00:00:00 2001 From: gebdalaoli-arch <265052293+gebdalaoli-arch@users.noreply.github.com> Date: Mon, 13 Jul 2026 03:47:39 +0800 Subject: [PATCH 1/3] fix: proxy Codex models through API key upstreams --- .../handler/openai_codex_models_handler.go | 7 +- .../service/openai_codex_models_service.go | 121 +++++++-- .../openai_codex_models_service_test.go | 257 ++++++++++++++++++ 3 files changed, 362 insertions(+), 23 deletions(-) diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index e64c555d14..87ab072fb4 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -15,10 +15,9 @@ import ( // Codex CLI and the Codex desktop app refresh their model picker from // GET {base_url}/models?client_version=... (custom provider mode) or // GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land -// here. The manifest is proxied verbatim from the ChatGPT backend with a -// schedulable OAuth account's credentials, so clients pointed at the gateway -// see the account's real, always-current model entitlements instead of a -// frozen local cache. +// here. The manifest is proxied verbatim from the selected account's ChatGPT +// backend or custom API key upstream, so clients pointed at the gateway see an +// always-current manifest instead of a frozen local cache. func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey.Group == nil { diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index 8a919fa2b0..d29b4dca59 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -2,6 +2,7 @@ package service import ( "context" + "fmt" "io" "net/http" "net/url" @@ -26,8 +27,8 @@ type CodexModelsManifest struct { NotModified bool } -// FetchCodexModelsManifest fetches the live Codex models manifest from the -// ChatGPT backend using the account's OAuth credentials. +// FetchCodexModelsManifest fetches the live Codex models manifest from either +// the ChatGPT backend for OAuth accounts or a custom upstream for API key accounts. // // The response body is passed through verbatim: the manifest schema evolves // with Codex client releases, and interpreting it here would force the gateway @@ -41,24 +42,61 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if err != nil { return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_CREDENTIALS_FAILED", "resolve credential account: %v", err) } - accessToken := credAccount.GetOpenAIAccessToken() - if accessToken == "" { - return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") - } clientVersion = strings.TrimSpace(clientVersion) if clientVersion == "" { clientVersion = openAICodexProbeVersion } - requestURL := chatgptCodexModelsURL + "?client_version=" + url.QueryEscape(clientVersion) + + requestEndpoint := chatgptCodexModelsURL + authToken := "" + useAPIKeyUpstream := false + appendModelsPath := false + switch { + case credAccount.IsOpenAIOAuth(): + authToken = strings.TrimSpace(credAccount.GetOpenAIAccessToken()) + if authToken == "" { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") + } + case credAccount.IsOpenAIApiKey(): + baseURL := strings.TrimSpace(credAccount.GetCredential("base_url")) + if baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) { + return nil, infraerrors.New( + http.StatusBadGateway, + "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED", + "Codex models manifest requires a custom API key upstream base URL", + ) + } + authToken = strings.TrimSpace(credAccount.GetOpenAIApiKey()) + if authToken == "" { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_MISSING", "account has no API key for the Codex models upstream") + } + normalizedBaseURL, validateErr := s.validateUpstreamBaseURL(baseURL) + if validateErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", validateErr) + } + requestEndpoint = normalizedBaseURL + useAPIKeyUpstream = true + appendModelsPath = true + default: + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_ACCOUNT_TYPE_UNSUPPORTED", "account type %q cannot fetch the Codex models manifest", credAccount.Type) + } + + requestURL, err := buildCodexModelsManifestURL(requestEndpoint, appendModelsPath, clientVersion) + if err != nil { + if useAPIKeyUpstream { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", err) + } + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "parse codex models request URL: %v", err) + } reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second) defer cancel() - req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL, nil) + req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL.String(), nil) if err != nil { return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) } - req.Header.Set("Authorization", "Bearer "+accessToken) + req.Header.Set("Authorization", "Bearer "+authToken) req.Header.Set("Accept", "application/json") req.Header.Set("Originator", "codex_cli_rs") req.Header.Set("Version", clientVersion) @@ -66,22 +104,35 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { req.Header.Set("If-None-Match", ifNoneMatch) } - setOpenAIChatGPTAccountHeaders(req.Header, credAccount) + if useAPIKeyUpstream { + credAccount.ApplyHeaderOverrides(req.Header) + } else { + setOpenAIChatGPTAccountHeaders(req.Header, credAccount) + } proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() } - client, err := httpclient.GetClient(httpclient.Options{ - ProxyURL: proxyURL, - Timeout: 15 * time.Second, - ResponseHeaderTimeout: 10 * time.Second, - }) - if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", err) - } - resp, err := client.Do(req) + var resp *http.Response + if useAPIKeyUpstream { + if s.httpUpstream == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_UPSTREAM_NOT_CONFIGURED", "Codex models upstream HTTP client is not configured") + } + req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) + resp, err = s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + } else { + client, clientErr := httpclient.GetClient(httpclient.Options{ + ProxyURL: proxyURL, + Timeout: 15 * time.Second, + ResponseHeaderTimeout: 10 * time.Second, + }) + if clientErr != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", clientErr) + } + resp, err = client.Do(req) + } if err != nil { return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err) } @@ -105,3 +156,35 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc } return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil } + +func isOfficialOpenAIModelsBaseURL(raw string) bool { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil { + return false + } + hostname := strings.TrimSuffix(parsed.Hostname(), ".") + return strings.EqualFold(hostname, "api.openai.com") +} + +func buildCodexModelsManifestURL(endpoint string, appendModelsPath bool, clientVersion string) (*url.URL, error) { + requestURL, err := url.Parse(endpoint) + if err != nil { + return nil, err + } + if requestURL.Fragment != "" { + return nil, fmt.Errorf("URL fragments are not supported") + } + + query := requestURL.Query() + requestURL.RawQuery = "" + requestURL.ForceQuery = false + if appendModelsPath { + requestURL, err = url.Parse(buildOpenAIModelsURL(requestURL.String())) + if err != nil { + return nil, err + } + } + query.Set("client_version", clientVersion) + requestURL.RawQuery = query.Encode() + return requestURL, nil +} diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index c9eae35629..d4d9c2d360 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -2,11 +2,52 @@ package service import ( "context" + "io" "net/http" "net/http/httptest" + "strings" "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" ) +type codexModelsHTTPUpstreamStub struct { + do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) +} + +func (s *codexModelsHTTPUpstreamStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + return s.do(req, proxyURL, accountID, accountConcurrency) +} + +func (s *codexModelsHTTPUpstreamStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return s.Do(req, proxyURL, accountID, accountConcurrency) +} + +func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{ + Enabled: false, + }}}, + httpUpstream: upstream, + } +} + +func newCodexModelsAPIKeyTestAccount(baseURL string) *Account { + credentials := map[string]any{"api_key": "sk-upstream"} + if baseURL != "" { + credentials["base_url"] = baseURL + } + return &Account{ + ID: 2, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: credentials, + Concurrency: 3, + } +} + func newCodexModelsTestAccount() *Account { return &Account{ ID: 1, @@ -136,3 +177,219 @@ func TestFetchCodexModelsManifestMissingToken(t *testing.T) { t.Fatal("expected error for missing access token, got nil") } } + +func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { + manifestBody := `{"models":[{"slug":"gpt-5.6"}]}` + var gotRequest *http.Request + var gotProxyURL string + var gotAccountID int64 + var gotConcurrency int + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + gotRequest = req + gotProxyURL = proxyURL + gotAccountID = accountID + gotConcurrency = accountConcurrency + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(manifestBody)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + + if gotRequest == nil { + t.Fatal("expected request to custom API key upstream") + } + if gotRequest.Method != http.MethodGet { + t.Errorf("method: got %q", gotRequest.Method) + } + if gotRequest.URL.String() != "https://upstream.example/v1/models?client_version=0.144.0" { + t.Errorf("request URL: got %q", gotRequest.URL.String()) + } + if gotRequest.Header.Get("Authorization") != "Bearer sk-upstream" { + t.Errorf("authorization header: got %q", gotRequest.Header.Get("Authorization")) + } + if gotRequest.Header.Get("Originator") != "codex_cli_rs" { + t.Errorf("originator header: got %q", gotRequest.Header.Get("Originator")) + } + if gotRequest.Header.Get("Version") != "0.144.0" { + t.Errorf("version header: got %q", gotRequest.Header.Get("Version")) + } + if gotRequest.Header.Get("User-Agent") != codexCLIUserAgent { + t.Errorf("user-agent header: got %q", gotRequest.Header.Get("User-Agent")) + } + if gotRequest.Header.Get("chatgpt-account-id") != "" { + t.Errorf("chatgpt-account-id must not be sent to API key upstream: got %q", gotRequest.Header.Get("chatgpt-account-id")) + } + if gotProxyURL != "" || gotAccountID != 2 || gotConcurrency != 3 { + t.Errorf("upstream routing metadata: proxy=%q account_id=%d concurrency=%d", gotProxyURL, gotAccountID, gotConcurrency) + } + if string(manifest.Body) != manifestBody { + t.Errorf("body not passed through verbatim: got %q", manifest.Body) + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } +} + +func TestFetchCodexModelsManifestAPIKeyNotModified(t *testing.T) { + var gotIfNoneMatch string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotIfNoneMatch = req.Header.Get("If-None-Match") + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusNotModified, + Header: header, + Body: http.NoBody, + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + `W/"api-key-manifest"`, + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if !manifest.NotModified { + t.Error("expected NotModified to be true") + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } + if gotIfNoneMatch != `W/"api-key-manifest"` { + t.Errorf("if-none-match header: got %q", gotIfNoneMatch) + } +} + +func TestFetchCodexModelsManifestAPIKeyPreservesBaseURLQuery(t *testing.T) { + var gotURL string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotURL = req.URL.String() + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1?tenant=acme"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if gotURL != "https://upstream.example/v1/models?client_version=0.144.0&tenant=acme" { + t.Errorf("request URL: got %q", gotURL) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsBaseURLFragment(t *testing.T) { + called := false + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + called = true + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1#models"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected invalid upstream base URL error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + if called { + t.Fatal("fragment-bearing base URL must be rejected before the upstream request") + } +} + +func TestFetchCodexModelsManifestAPIKeyUpstreamError(t *testing.T) { + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusTooManyRequests, + Status: "429 Too Many Requests", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":"rate limited"}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected error for upstream 429, got nil") + } + if infraerrors.Code(err) != http.StatusBadGateway { + t.Errorf("error status: got %d, want %d", infraerrors.Code(err), http.StatusBadGateway) + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_UPSTREAM_FAILED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T) { + tests := []struct { + name string + baseURL string + }{ + {name: "missing base URL"}, + {name: "official host", baseURL: "https://api.openai.com"}, + {name: "official versioned URL", baseURL: "https://API.OPENAI.COM:443/v1/"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := newCodexModelsAPIKeyTestService(&codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + t.Fatal("official OpenAI API key must not be used as a Codex manifest upstream") + return nil, nil + }}) + + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount(tt.baseURL), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected unsupported API key upstream error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + }) + } +} From ed31a52424ff996d9210fbad1644cd9a1a7c8698 Mon Sep 17 00:00:00 2001 From: gebdalaoli-arch <265052293+gebdalaoli-arch@users.noreply.github.com> Date: Mon, 13 Jul 2026 19:31:38 +0800 Subject: [PATCH 2/3] fix: stabilize API key Codex manifest refreshes --- .../handler/openai_codex_models_handler.go | 16 +- .../openai_codex_models_handler_test.go | 26 + .../service/openai_codex_models_service.go | 255 ++++++++- .../openai_codex_models_service_test.go | 494 +++++++++++++++++- .../service/openai_gateway_service.go | 1 + 5 files changed, 764 insertions(+), 28 deletions(-) create mode 100644 backend/internal/handler/openai_codex_models_handler_test.go diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index 87ab072fb4..d98cf590e3 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -16,9 +16,12 @@ import ( // GET {base_url}/models?client_version=... (custom provider mode) or // GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land // here. The manifest is proxied verbatim from the selected account's ChatGPT -// backend or custom API key upstream, so clients pointed at the gateway see an -// always-current manifest instead of a frozen local cache. +// backend or custom API key upstream. API key manifests use a short-lived, +// asynchronously revalidated cache to tolerate canceled client requests. func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { + if c.Request.Context().Err() != nil { + return + } apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey.Group == nil { h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required") @@ -31,15 +34,24 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "") if err != nil { + if c.Request.Context().Err() != nil { + return + } h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") return } manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) if err != nil { + if c.Request.Context().Err() != nil { + return + } h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) return } + if c.Request.Context().Err() != nil { + return + } if manifest.ETag != "" { c.Header("ETag", manifest.ETag) diff --git a/backend/internal/handler/openai_codex_models_handler_test.go b/backend/internal/handler/openai_codex_models_handler_test.go new file mode 100644 index 0000000000..5b1382fade --- /dev/null +++ b/backend/internal/handler/openai_codex_models_handler_test.go @@ -0,0 +1,26 @@ +package handler + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil).WithContext(ctx) + + h := &OpenAIGatewayHandler{} + h.CodexModels(c) + + if c.Writer.Written() { + t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String()) + } +} diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index d29b4dca59..ef5af10a99 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -2,22 +2,33 @@ package service import ( "context" + "crypto/sha256" "fmt" "io" "net/http" "net/url" + "sort" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "golang.org/x/sync/singleflight" ) // chatgptCodexModelsURL is the ChatGPT Codex models manifest endpoint. // Package-level variable so tests can point it at a stub server. var chatgptCodexModelsURL = "https://chatgpt.com/backend-api/codex/models" -const codexModelsManifestBodyLimit int64 = 8 << 20 +const ( + codexModelsManifestBodyLimit int64 = 8 << 20 + codexModelsManifestCacheBodyLimit = 1 << 20 + codexModelsManifestCacheMaxEntries = 64 + codexModelsManifestCacheTTL = 30 * time.Second + codexModelsManifestCacheStaleTTL = 5 * time.Minute + codexModelsManifestRequestTimeout = 15 * time.Second +) // CodexModelsManifest carries the raw upstream manifest payload plus caching // metadata so handlers can pass both through to the client untouched. @@ -27,6 +38,90 @@ type CodexModelsManifest struct { NotModified bool } +type codexModelsManifestRequest struct { + url string + headers http.Header + proxyURL string + accountID int64 + credentialAccountID int64 + accountConcurrency int + useAPIKeyUpstream bool +} + +type codexModelsManifestCacheEntry struct { + manifest *CodexModelsManifest + order uint64 + expiresAt time.Time + staleUntil time.Time +} + +type codexModelsManifestCacheState uint8 + +const ( + codexModelsManifestCacheMiss codexModelsManifestCacheState = iota + codexModelsManifestCacheFresh + codexModelsManifestCacheStale +) + +type codexModelsManifestCache struct { + mu sync.Mutex + entries map[string]codexModelsManifestCacheEntry + nextOrder uint64 + refresh singleflight.Group +} + +func (c *codexModelsManifestCache) get(key string, now time.Time) (*CodexModelsManifest, codexModelsManifestCacheState) { + c.mu.Lock() + defer c.mu.Unlock() + entry, ok := c.entries[key] + if !ok { + return nil, codexModelsManifestCacheMiss + } + if !now.Before(entry.staleUntil) { + delete(c.entries, key) + return nil, codexModelsManifestCacheMiss + } + if now.Before(entry.expiresAt) { + return entry.manifest, codexModelsManifestCacheFresh + } + return entry.manifest, codexModelsManifestCacheStale +} + +func (c *codexModelsManifestCache) set(key string, manifest *CodexModelsManifest, now time.Time) { + if manifest == nil || len(manifest.Body) > codexModelsManifestCacheBodyLimit { + return + } + c.mu.Lock() + defer c.mu.Unlock() + if c.entries == nil { + c.entries = make(map[string]codexModelsManifestCacheEntry) + } + if _, exists := c.entries[key]; !exists && len(c.entries) >= codexModelsManifestCacheMaxEntries { + oldestKey := "" + var oldestOrder uint64 + for candidateKey, entry := range c.entries { + if !now.Before(entry.staleUntil) { + delete(c.entries, candidateKey) + continue + } + if oldestKey == "" || entry.order < oldestOrder { + oldestKey = candidateKey + oldestOrder = entry.order + } + } + if len(c.entries) >= codexModelsManifestCacheMaxEntries && oldestKey != "" { + delete(c.entries, oldestKey) + } + } + c.nextOrder++ + c.entries[key] = codexModelsManifestCacheEntry{ + manifest: manifest, + order: c.nextOrder, + expiresAt: now.Add(codexModelsManifestCacheTTL), + staleUntil: now.Add(codexModelsManifestCacheStaleTTL), + } +} + // FetchCodexModelsManifest fetches the live Codex models manifest from either // the ChatGPT backend for OAuth accounts or a custom upstream for API key accounts. // @@ -90,24 +185,16 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "parse codex models request URL: %v", err) } - reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second) - defer cancel() - req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL.String(), nil) - if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) - } - req.Header.Set("Authorization", "Bearer "+authToken) - req.Header.Set("Accept", "application/json") - req.Header.Set("Originator", "codex_cli_rs") - req.Header.Set("Version", clientVersion) - req.Header.Set("User-Agent", codexCLIUserAgent) - if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { - req.Header.Set("If-None-Match", ifNoneMatch) - } + headers := make(http.Header) + headers.Set("Authorization", "Bearer "+authToken) + headers.Set("Accept", "application/json") + headers.Set("Originator", "codex_cli_rs") + headers.Set("Version", clientVersion) + headers.Set("User-Agent", codexCLIUserAgent) if useAPIKeyUpstream { - credAccount.ApplyHeaderOverrides(req.Header) + credAccount.ApplyHeaderOverrides(headers) } else { - setOpenAIChatGPTAccountHeaders(req.Header, credAccount) + setOpenAIChatGPTAccountHeaders(headers, credAccount) } proxyURL := "" @@ -115,17 +202,94 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc proxyURL = account.Proxy.URL() } - var resp *http.Response + request := codexModelsManifestRequest{ + url: requestURL.String(), + headers: headers, + proxyURL: proxyURL, + accountID: account.ID, + credentialAccountID: credAccount.ID, + accountConcurrency: account.Concurrency, + useAPIKeyUpstream: useAPIKeyUpstream, + } if useAPIKeyUpstream { + return s.fetchCachedAPIKeyCodexModelsManifest(ctx, request, ifNoneMatch) + } + return s.fetchCodexModelsManifestUpstream(ctx, request, ifNoneMatch) +} + +func (s *OpenAIGatewayService) fetchCachedAPIKeyCodexModelsManifest(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + cacheKey := buildCodexModelsManifestCacheKey(request) + manifest, state := s.codexModelsManifestCache.get(cacheKey, time.Now()) + if state == codexModelsManifestCacheFresh { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + resultCh := s.refreshCachedAPIKeyCodexModelsManifest(cacheKey, request) + if state == codexModelsManifestCacheStale { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-resultCh: + if result.Err != nil { + return nil, result.Err + } + manifest, ok := result.Val.(*CodexModelsManifest) + if !ok || manifest == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "invalid shared Codex models manifest result") + } + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } +} + +func (s *OpenAIGatewayService) refreshCachedAPIKeyCodexModelsManifest(cacheKey string, request codexModelsManifestRequest) <-chan singleflight.Result { + return s.codexModelsManifestCache.refresh.DoChan(cacheKey, func() (any, error) { + cached, _ := s.codexModelsManifestCache.get(cacheKey, time.Now()) + ifNoneMatch := "" + if cached != nil { + ifNoneMatch = cached.ETag + } + manifest, err := s.fetchCodexModelsManifestUpstream(context.Background(), request, ifNoneMatch) + if err != nil { + return nil, err + } + if manifest.NotModified && cached != nil { + s.codexModelsManifestCache.set(cacheKey, cached, time.Now()) + return cached, nil + } + if !manifest.NotModified { + s.codexModelsManifestCache.set(cacheKey, manifest, time.Now()) + } + return manifest, nil + }) +} + +func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + reqCtx, cancel := context.WithTimeout(ctx, codexModelsManifestRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, request.url, nil) + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) + } + req.Header = request.headers.Clone() + if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { + req.Header.Set("If-None-Match", ifNoneMatch) + } + + var resp *http.Response + if request.useAPIKeyUpstream { if s.httpUpstream == nil { return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_UPSTREAM_NOT_CONFIGURED", "Codex models upstream HTTP client is not configured") } req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) - resp, err = s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + resp, err = s.httpUpstream.Do(req, request.proxyURL, request.accountID, request.accountConcurrency) } else { client, clientErr := httpclient.GetClient(httpclient.Options{ - ProxyURL: proxyURL, - Timeout: 15 * time.Second, + ProxyURL: request.proxyURL, + Timeout: codexModelsManifestRequestTimeout, ResponseHeaderTimeout: 10 * time.Second, }) if clientErr != nil { @@ -157,6 +321,55 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil } +func buildCodexModelsManifestCacheKey(request codexModelsManifestRequest) string { + hasher := sha256.New() + _, _ = fmt.Fprintf(hasher, "%d\n%d\n%s\n%s\n", request.accountID, request.credentialAccountID, request.proxyURL, request.url) + headerNames := make([]string, 0, len(request.headers)) + for name := range request.headers { + headerNames = append(headerNames, name) + } + sort.Strings(headerNames) + for _, name := range headerNames { + _, _ = fmt.Fprintf(hasher, "%s\n", strings.ToLower(name)) + for _, value := range request.headers[name] { + _, _ = fmt.Fprintf(hasher, "%s\n", value) + } + } + return fmt.Sprintf("%x", hasher.Sum(nil)) +} + +func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch string) *CodexModelsManifest { + if manifest == nil { + return nil + } + if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { + return &CodexModelsManifest{ETag: manifest.ETag, NotModified: true} + } + return manifest +} + +func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool { + etag = strings.TrimSpace(etag) + if etag == "" { + return false + } + normalize := func(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && strings.EqualFold(value[:2], "W/") { + value = strings.TrimSpace(value[2:]) + } + return value + } + want := normalize(etag) + for _, candidate := range strings.Split(ifNoneMatch, ",") { + candidate = strings.TrimSpace(candidate) + if candidate == "*" || normalize(candidate) == want { + return true + } + } + return false +} + func isOfficialOpenAIModelsBaseURL(raw string) bool { parsed, err := url.Parse(strings.TrimSpace(raw)) if err != nil { diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index d4d9c2d360..9d628c81f0 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -2,11 +2,15 @@ package service import ( "context" + "errors" "io" "net/http" "net/http/httptest" "strings" + "sync" + "sync/atomic" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" @@ -17,6 +21,26 @@ type codexModelsHTTPUpstreamStub struct { do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) } +type codexModelsBlockingBody struct { + ctx context.Context + readStarted chan struct{} + startedOnce *sync.Once + release <-chan struct{} + body *strings.Reader +} + +func (b *codexModelsBlockingBody) Read(p []byte) (int, error) { + b.startedOnce.Do(func() { close(b.readStarted) }) + select { + case <-b.release: + return b.body.Read(p) + case <-b.ctx.Done(): + return 0, b.ctx.Err() + } +} + +func (b *codexModelsBlockingBody) Close() error { return nil } + func (s *codexModelsHTTPUpstreamStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { return s.do(req, proxyURL, accountID, accountConcurrency) } @@ -244,16 +268,449 @@ func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { } } -func TestFetchCodexModelsManifestAPIKeyNotModified(t *testing.T) { +func TestFetchCodexModelsManifestAPIKeySharedRefreshSurvivesCallerCancellation(t *testing.T) { + const manifestBody = `{"models":[{"slug":"gpt-5.6"}]}` + var calls atomic.Int32 + var readStartedOnce sync.Once + readStarted := make(chan struct{}) + deadlineRemaining := make(chan time.Duration, 1) + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + deadline, ok := req.Context().Deadline() + if !ok { + deadlineRemaining <- 0 + } else { + deadlineRemaining <- time.Until(deadline) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"shared"`}}, + Body: &codexModelsBlockingBody{ + ctx: req.Context(), + readStarted: readStarted, + startedOnce: &readStartedOnce, + release: release, + body: strings.NewReader(manifestBody), + }, + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + firstCtx, cancelFirst := context.WithCancel(context.Background()) + firstErr := make(chan error, 1) + go func() { + _, err := s.FetchCodexModelsManifest(firstCtx, account, "0.144.0", "") + firstErr <- err + }() + + select { + case <-readStarted: + case <-time.After(time.Second): + t.Fatal("upstream body read did not start") + } + remaining := <-deadlineRemaining + if remaining < 14*time.Second || remaining > codexModelsManifestRequestTimeout { + t.Errorf("detached refresh deadline: got %s, want approximately %s", remaining, codexModelsManifestRequestTimeout) + } + cancelFirst() + select { + case err := <-firstErr: + if !errors.Is(err, context.Canceled) { + t.Fatalf("first caller error: got %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("canceled caller did not return promptly") + } + + secondResult := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + secondResult <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls before shared refresh completed: got %d, want 1", got) + } + close(release) + select { + case result := <-secondResult: + if result.err != nil { + t.Fatalf("second caller returned error: %v", result.err) + } + if string(result.manifest.Body) != manifestBody { + t.Errorf("second caller body: got %q", result.manifest.Body) + } + case <-time.After(time.Second): + t.Fatal("second caller did not receive shared refresh result") + } + if got := calls.Load(); got != 1 { + t.Errorf("total upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyConcurrentRequestsShareRefresh(t *testing.T) { + const callers = 8 + var calls atomic.Int32 + started := make(chan struct{}) + var startedOnce sync.Once + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + startedOnce.Do(func() { close(started) }) + <-release + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + begin := make(chan struct{}) + errs := make(chan error, callers) + for i := 0; i < callers; i++ { + go func() { + <-begin + _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + errs <- err + }() + } + close(begin) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("upstream request did not start") + } + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("concurrent upstream calls: got %d, want 1", got) + } + close(release) + for i := 0; i < callers; i++ { + if err := <-errs; err != nil { + t.Errorf("caller %d returned error: %v", i, err) + } + } +} + +func TestFetchCodexModelsManifestAPIKeyFreshCacheHandlesETagLocally(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cache refresh must not inherit a caller's If-None-Match: got %q", got) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"cached"`}}, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", `W/"cached"`) + if err != nil { + t.Fatalf("cached fetch returned error: %v", err) + } + if !manifest.NotModified { + t.Fatal("matching cached ETag must return NotModified") + } + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheKeyIsolatesRequestIdentity(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + + base := newCodexModelsAPIKeyTestAccount("https://upstream.example") + fetch := func(account *Account, version string) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, version, ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + fetch(base, "0.144.0") + fetch(base, "0.144.0") + + differentAccount := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentAccount.ID = 3 + fetch(differentAccount, "0.144.0") + + differentToken := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentToken.Credentials["api_key"] = "sk-other" + fetch(differentToken, "0.144.0") + + differentUpstream := newCodexModelsAPIKeyTestAccount("https://other-upstream.example") + fetch(differentUpstream, "0.144.0") + fetch(base, "0.145.0") + + differentHeaders := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentHeaders.Credentials[credKeyHeaderOverrideEnabled] = true + differentHeaders.Credentials[credKeyHeaderOverrides] = map[string]any{"x-tenant": "other"} + fetch(differentHeaders, "0.144.0") + + proxyID := int64(9) + differentProxy := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentProxy.ProxyID = &proxyID + differentProxy.Proxy = &Proxy{Protocol: "http", Host: "127.0.0.1", Port: 8080} + fetch(differentProxy, "0.144.0") + fetch(differentProxy, "0.144.0") + + if got := calls.Load(); got != 7 { + t.Errorf("isolated upstream calls: got %d, want 7", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + body := `{"models":[]}` + if strings.Contains(req.URL.Host, "large") { + body = strings.Repeat("x", (1<<20)+1) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + fetch := func(account *Account) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + + small := newCodexModelsAPIKeyTestAccount("https://small.example") + fetch(small) + fetch(small) + large := newCodexModelsAPIKeyTestAccount("https://large.example") + large.ID = 3 + fetch(large) + fetch(large) + if got := calls.Load(); got != 3 { + t.Fatalf("body-size bounded cache calls: got %d, want 3", got) + } + + for i := int64(10); i < 75; i++ { + account := newCodexModelsAPIKeyTestAccount("https://bounded.example") + account.ID = i + fetch(account) + } + last := newCodexModelsAPIKeyTestAccount("https://bounded.example") + last.ID = 74 + fetch(last) + if got := calls.Load(); got != 68 { + t.Fatalf("most recent cache entry was not retained: calls=%d, want 68", got) + } + first := newCodexModelsAPIKeyTestAccount("https://bounded.example") + first.ID = 10 + fetch(first) + if got := calls.Load(); got != 69 { + t.Errorf("oldest cache entry was not evicted: calls=%d, want 69", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyServesStaleWhileRefreshing(t *testing.T) { + var calls atomic.Int32 + refreshStarted := make(chan struct{}) + releaseRefresh := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + body := `{"models":[{"slug":"old"}]}` + if call > 1 { + if call == 2 { + close(refreshStarted) + } + <-releaseRefresh + body = `{"models":[{"slug":"new"}]}` + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + resultCh := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + resultCh <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + select { + case <-refreshStarted: + case <-time.After(time.Second): + t.Fatal("background refresh did not start") + } + + var staleResult struct { + manifest *CodexModelsManifest + err error + } + select { + case staleResult = <-resultCh: + case <-time.After(100 * time.Millisecond): + t.Error("stale manifest was not returned while refresh was blocked") + close(releaseRefresh) + staleResult = <-resultCh + } + if staleResult.err != nil { + t.Fatalf("stale fetch returned error: %v", staleResult.err) + } + if got := string(staleResult.manifest.Body); got != `{"models":[{"slug":"old"}]}` { + t.Errorf("stale body: got %q", got) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls during stale refresh: got %d, want 2", got) + } + + select { + case <-releaseRefresh: + default: + close(releaseRefresh) + } + deadline := time.Now().Add(time.Second) + for { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err == nil && string(manifest.Body) == `{"models":[{"slug":"new"}]}` { + break + } + if time.Now().After(deadline) { + t.Fatalf("refreshed manifest was not cached: manifest=%v err=%v", manifest, err) + } + time.Sleep(10 * time.Millisecond) + } + if got := calls.Load(); got != 2 { + t.Errorf("stale refresh was not deduplicated: calls=%d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyRevalidatesStaleETag(t *testing.T) { + var calls atomic.Int32 + refreshDone := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + if call == 1 { + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"cached"}]}`)), + }, nil + } + if got := req.Header.Get("If-None-Match"); got != `W/"cached"` { + t.Errorf("background revalidation If-None-Match: got %q", got) + } + close(refreshDone) + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("stale fetch returned error: %v", err) + } + if got := string(manifest.Body); got != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("stale body: got %q", got) + } + select { + case <-refreshDone: + case <-time.After(time.Second): + t.Fatal("ETag revalidation did not complete") + } + + deadline := time.Now().Add(time.Second) + for { + s.codexModelsManifestCache.mu.Lock() + fresh := false + for _, entry := range s.codexModelsManifestCache.entries { + fresh = time.Now().Before(entry.expiresAt) + } + s.codexModelsManifestCache.mu.Unlock() + if fresh { + break + } + if time.Now().After(deadline) { + t.Fatal("304 revalidation did not renew the cached manifest") + } + time.Sleep(10 * time.Millisecond) + } + manifest, err = s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil || string(manifest.Body) != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("renewed cached manifest: body=%q err=%v", manifest.Body, err) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls: got %d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyColdCacheHandlesNotModifiedLocally(t *testing.T) { var gotIfNoneMatch string upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { gotIfNoneMatch = req.Header.Get("If-None-Match") header := make(http.Header) header.Set("ETag", `W/"api-key-manifest"`) return &http.Response{ - StatusCode: http.StatusNotModified, + StatusCode: http.StatusOK, Header: header, - Body: http.NoBody, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), }, nil }} @@ -273,8 +730,35 @@ func TestFetchCodexModelsManifestAPIKeyNotModified(t *testing.T) { if manifest.ETag != `W/"api-key-manifest"` { t.Errorf("etag not passed through: got %q", manifest.ETag) } - if gotIfNoneMatch != `W/"api-key-manifest"` { - t.Errorf("if-none-match header: got %q", gotIfNoneMatch) + if gotIfNoneMatch != "" { + t.Errorf("cold shared refresh must not inherit caller if-none-match: got %q", gotIfNoneMatch) + } +} + +func TestFetchCodexModelsManifestAPIKeyDoesNotCacheUnexpectedColdNotModified(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cold shared refresh If-None-Match: got %q", got) + } + header := make(http.Header) + header.Set("ETag", `W/"unexpected"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + for i := 0; i < 2; i++ { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("fetch %d returned error: %v", i, err) + } + if !manifest.NotModified { + t.Fatalf("fetch %d: expected upstream NotModified response", i) + } + } + if got := calls.Load(); got != 2 { + t.Errorf("unexpected cold 304 was cached: upstream calls=%d, want 2", got) } } diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 0d03d83dcf..3db2bf05e9 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -369,6 +369,7 @@ type OpenAIGatewayService struct { openaiWSRetryMetrics openAIWSRetryMetrics responseHeaderFilter *responseheaders.CompiledHeaderFilter codexSnapshotThrottle *accountWriteThrottle + codexModelsManifestCache codexModelsManifestCache openaiCompatSessionResponses sync.Map openaiCompatAnthropicDigestSessions sync.Map } From 3c68b2e3693272e6066cf2cb072e3fb6b41a02aa Mon Sep 17 00:00:00 2001 From: gebdalaoli-arch <265052293+gebdalaoli-arch@users.noreply.github.com> Date: Mon, 13 Jul 2026 20:36:37 +0800 Subject: [PATCH 3/3] fix: fail over Codex manifest accounts --- .../handler/openai_codex_models_handler.go | 63 +++-- .../openai_codex_models_handler_test.go | 262 ++++++++++++++++++ .../service/openai_codex_models_service.go | 107 ++++++- .../openai_codex_models_service_test.go | 70 +++++ 4 files changed, 478 insertions(+), 24 deletions(-) diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index d98cf590e3..1c1357cbfa 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -32,33 +32,54 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { return } - account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "") - if err != nil { + maxAccountSwitches := h.maxAccountSwitches + if maxAccountSwitches <= 0 { + maxAccountSwitches = 3 + } + failedAccountIDs := make(map[int64]struct{}) + switchCount := 0 + var lastUpstreamErr error + + for { + account, err := h.gatewayService.SelectAccountForModelWithExclusions(c.Request.Context(), apiKey.GroupID, "", "", failedAccountIDs) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if lastUpstreamErr != nil { + h.errorResponse(c, infraerrors.Code(lastUpstreamErr), "upstream_error", infraerrors.Message(lastUpstreamErr)) + return + } + h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") + return + } + + manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if service.IsRetryableCodexModelsManifestError(err) && switchCount < maxAccountSwitches { + failedAccountIDs[account.ID] = struct{}{} + switchCount++ + lastUpstreamErr = err + continue + } + h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) + return + } if c.Request.Context().Err() != nil { return } - h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") - return - } - manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) - if err != nil { - if c.Request.Context().Err() != nil { + if manifest.ETag != "" { + c.Header("ETag", manifest.ETag) + } + if manifest.NotModified { + c.Status(http.StatusNotModified) return } - h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) + c.Data(http.StatusOK, "application/json", manifest.Body) return } - if c.Request.Context().Err() != nil { - return - } - - if manifest.ETag != "" { - c.Header("ETag", manifest.ETag) - } - if manifest.NotModified { - c.Status(http.StatusNotModified) - return - } - c.Data(http.StatusOK, "application/json", manifest.Body) } diff --git a/backend/internal/handler/openai_codex_models_handler_test.go b/backend/internal/handler/openai_codex_models_handler_test.go index 5b1382fade..ba74a5869f 100644 --- a/backend/internal/handler/openai_codex_models_handler_test.go +++ b/backend/internal/handler/openai_codex_models_handler_test.go @@ -2,13 +2,95 @@ package handler import ( "context" + "errors" + "fmt" + "io" + "net" "net/http" "net/http/httptest" + "strings" + "sync" "testing" + "github.com/Wei-Shaw/sub2api/internal/config" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" ) +type codexModelsFailoverAccountRepo struct { + service.AccountRepository + accounts []service.Account +} + +func (r codexModelsFailoverAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + for i := range r.accounts { + if r.accounts[i].ID == id { + account := r.accounts[i] + return &account, nil + } + } + return nil, service.ErrNoAvailableAccounts +} + +func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) { + accounts := make([]service.Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform { + accounts = append(accounts, account) + } + } + return accounts, nil +} + +type codexModelsFailoverHTTPUpstream struct { + service.HTTPUpstream + mu sync.Mutex + accountIDs []int64 + firstErr error + firstStatus int + statuses map[int64]int +} + +func (u *codexModelsFailoverHTTPUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + u.mu.Lock() + u.accountIDs = append(u.accountIDs, accountID) + u.mu.Unlock() + + status, hasStatus := u.statuses[accountID] + if accountID == 1 || hasStatus { + if u.firstErr != nil { + return nil, u.firstErr + } + if !hasStatus { + status = u.firstStatus + } + if status == 0 { + status = http.StatusServiceUnavailable + } + return &http.Response{ + StatusCode: status, + Status: http.StatusText(status), + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader( + `{"error":{"message":"No available OpenAI accounts","type":"upstream_error"}}`, + )), + }, nil + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"gpt-5.6-sol"}]}`)), + }, nil +} + +func (u *codexModelsFailoverHTTPUpstream) calls() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.accountIDs...) +} + func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() @@ -24,3 +106,183 @@ func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String()) } } + +func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) { + retryableStatuses := []int{ + http.StatusTooManyRequests, + http.StatusInternalServerError, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout, + } + for _, status := range retryableStatuses { + t.Run(http.StatusText(status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want { + t.Fatalf("body: got %q, want %q", got, want) + } + }) + } +} + +func TestCodexModelsFailsOverFromUpstreamTransportError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } +} + +func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) { + statuses := []int{ + http.StatusBadRequest, + http.StatusUnauthorized, + http.StatusForbidden, + http.StatusNotFound, + 600, + } + for _, status := range statuses { + t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + }) + } +} + +func TestCodexModelsDoesNotFailOverFromUpstreamConfigurationError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = errors.New("invalid proxy URL") + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } +} + +func TestCodexModelsReturnsLastUpstreamErrorWhenAccountsAreExhausted(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusGatewayTimeout, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the last upstream error: %s", body) + } +} + +func TestCodexModelsHonorsAccountSwitchLimit(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandlerWithAccountCount(http.StatusServiceUnavailable, 4, 2) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusBadGateway, + 3: http.StatusGatewayTimeout, + 4: http.StatusInternalServerError, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2, 3}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the limit-ending upstream error: %s", body) + } +} + +func newCodexModelsFailoverTestHandler(firstStatus int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + return newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, 2, 3) +} + +func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount, maxSwitches int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + gin.SetMode(gin.TestMode) + groupID := int64(42) + accounts := make([]service.Account, 0, accountCount) + for i := 1; i <= accountCount; i++ { + accounts = append(accounts, service.Account{ + ID: int64(i), + Name: fmt.Sprintf("upstream-%d", i), + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Priority: i - 1, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": fmt.Sprintf("sk-%d", i), + "base_url": fmt.Sprintf("https://upstream-%d.example/v1", i), + }, + }) + } + upstream := &codexModelsFailoverHTTPUpstream{firstStatus: firstStatus} + cfg := &config.Config{RunMode: config.RunModeSimple} + gatewayService := service.NewOpenAIGatewayService( + codexModelsFailoverAccountRepo{accounts: accounts}, + nil, nil, nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + return &OpenAIGatewayHandler{gatewayService: gatewayService, maxAccountSwitches: maxSwitches}, upstream, groupID +} + +func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder { + t.Helper() + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + GroupID: &groupID, + Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}, + }) + + handler.CodexModels(c) + return recorder +} + +func equalInt64Slices(got, want []int64) bool { + if len(got) != len(want) { + return false + } + for i := range got { + if got[i] != want[i] { + return false + } + } + return true +} diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index ef5af10a99..a7331f64c9 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -3,8 +3,10 @@ package service import ( "context" "crypto/sha256" + "errors" "fmt" "io" + "net" "net/http" "net/url" "sort" @@ -14,6 +16,7 @@ import ( infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "golang.org/x/net/http2" "golang.org/x/sync/singleflight" ) @@ -38,6 +41,94 @@ type CodexModelsManifest struct { NotModified bool } +type codexModelsManifestUpstreamError struct { + err error + retryable bool +} + +func (e *codexModelsManifestUpstreamError) Error() string { return e.err.Error() } + +func (e *codexModelsManifestUpstreamError) Unwrap() error { return e.err } + +// IsRetryableCodexModelsManifestError reports whether another selected account +// may succeed without changing the request. Configuration and upstream 4xx +// responses, except 429, are intentionally not retried. +func IsRetryableCodexModelsManifestError(err error) bool { + var upstreamErr *codexModelsManifestUpstreamError + return errors.As(err, &upstreamErr) && upstreamErr.retryable +} + +func isRetryableCodexModelsManifestTransportError(err error) bool { + if err == nil || errors.Is(err, context.Canceled) { + return false + } + if errors.Is(err, context.DeadlineExceeded) || + errors.Is(err, io.EOF) || + errors.Is(err, io.ErrUnexpectedEOF) || + errors.Is(err, net.ErrClosed) { + return true + } + + var opErr *net.OpError + if errors.As(err, &opErr) { + return true + } + var dnsErr *net.DNSError + if errors.As(err, &dnsErr) { + return true + } + var goAwayErr http2.GoAwayError + if errors.As(err, &goAwayErr) { + return true + } + var streamErr http2.StreamError + if errors.As(err, &streamErr) { + return true + } + var connectionErr http2.ConnectionError + if errors.As(err, &connectionErr) { + return true + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + + // net/http uses unexported HTTP/2 error types, so typed matching is not + // possible for errors produced by the standard library transport. + message := strings.ToLower(err.Error()) + if strings.Contains(message, "http2:") && + (strings.Contains(message, "goaway") || + strings.Contains(message, "refused_stream") || + strings.Contains(message, "frame too large")) { + return true + } + if strings.Contains(message, "stream error: stream id ") { + return true + } + for _, code := range []http2.ErrCode{ + http2.ErrCodeNo, + http2.ErrCodeProtocol, + http2.ErrCodeInternal, + http2.ErrCodeFlowControl, + http2.ErrCodeSettingsTimeout, + http2.ErrCodeStreamClosed, + http2.ErrCodeFrameSize, + http2.ErrCodeRefusedStream, + http2.ErrCodeCancel, + http2.ErrCodeCompression, + http2.ErrCodeConnect, + http2.ErrCodeEnhanceYourCalm, + http2.ErrCodeInadequateSecurity, + http2.ErrCodeHTTP11Required, + } { + if strings.Contains(message, "connection error: "+strings.ToLower(code.String())) { + return true + } + } + return false +} + type codexModelsManifestRequest struct { url string headers http.Header @@ -298,7 +389,10 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont resp, err = client.Do(req) } if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } defer func() { _ = resp.Body.Close() }() @@ -311,12 +405,19 @@ func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Cont if message == "" { message = resp.Status } - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message), + retryable: resp.StatusCode == http.StatusTooManyRequests || + (resp.StatusCode >= http.StatusInternalServerError && resp.StatusCode < 600), + } } body, err := io.ReadAll(io.LimitReader(resp.Body, codexModelsManifestBodyLimit)) if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil } diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index 9d628c81f0..2c87861f60 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -4,8 +4,10 @@ import ( "context" "errors" "io" + "net" "net/http" "net/http/httptest" + "net/url" "strings" "sync" "sync/atomic" @@ -15,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "golang.org/x/net/http2" ) type codexModelsHTTPUpstreamStub struct { @@ -49,6 +52,73 @@ func (s *codexModelsHTTPUpstreamStub) DoWithTLS(req *http.Request, proxyURL stri return s.Do(req, proxyURL, accountID, accountConcurrency) } +func TestIsRetryableCodexModelsManifestTransportError(t *testing.T) { + tests := []struct { + name string + err error + retryable bool + }{ + {name: "nil", err: nil}, + {name: "configuration error", err: errors.New("invalid proxy URL")}, + {name: "upstream configuration error", err: errors.New("upstream error: invalid proxy")}, + {name: "proxy connection configuration error", err: errors.New("proxy connection error: invalid configuration")}, + {name: "canceled request", err: context.Canceled}, + { + name: "redirect policy error", + err: &url.Error{ + Op: "Get", + URL: "https://upstream.example/v1/models", + Err: errors.New("stopped after 10 redirects"), + }, + }, + {name: "deadline exceeded", err: context.DeadlineExceeded, retryable: true}, + {name: "unexpected EOF", err: io.ErrUnexpectedEOF, retryable: true}, + {name: "closed connection", err: net.ErrClosed, retryable: true}, + { + name: "network operation", + err: &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + }, + retryable: true, + }, + { + name: "DNS error", + err: &net.DNSError{Err: "temporary failure", Name: "upstream.example"}, + retryable: true, + }, + { + name: "typed HTTP2 GOAWAY", + err: http2.GoAwayError{ErrCode: http2.ErrCodeNo}, + retryable: true, + }, + { + name: "stdlib HTTP2 GOAWAY", + err: errors.New("http2: server sent GOAWAY and closed the connection; LastStreamID=1, ErrCode=NO_ERROR"), + retryable: true, + }, + { + name: "stdlib HTTP2 refused stream", + err: errors.New("stream error: stream ID 3; REFUSED_STREAM"), + retryable: true, + }, + { + name: "stdlib HTTP2 connection error", + err: errors.New(`Get "https://upstream.example/v1/models": connection error: PROTOCOL_ERROR`), + retryable: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isRetryableCodexModelsManifestTransportError(tt.err); got != tt.retryable { + t.Fatalf("retryable = %v, want %v", got, tt.retryable) + } + }) + } +} + func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService { return &OpenAIGatewayService{ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{