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 }