From e3cbe437d00b36fd41a00a61243fc2625ca96b57 Mon Sep 17 00:00:00 2001 From: Luis Pater Date: Tue, 15 Sep 2026 01:23:26 +0800 Subject: [PATCH] =?UTF-8?q?feat(auth):=20add=20tests=20to=20ensure=20async?= =?UTF-8?q?=20operations=20don=E2=80=99t=20block=20unrelated=20actions?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add tests to verify that runtime hooks, auth modifications, and related updates do not block on unrelated operations such as antigravity probes or plugin virtual models. - Introduce detailed scenarios, e.g., stale disables, batch operations, and conflict handling during concurrent auth updates. - Refactor locking mechanisms to avoid unnecessary blocking during hooks and model registrations. Fixed: #5813 Closes: #5773 --- .../handlers/management/auth_files_fields.go | 78 ++- .../management/auth_files_status_sync_test.go | 114 ++++ sdk/cliproxy/builder.go | 3 +- sdk/cliproxy/service.go | 2 + sdk/cliproxy/service_auth.go | 145 ++++- sdk/cliproxy/service_auth_sync_test.go | 548 ++++++++++++++++++ sdk/cliproxy/service_models.go | 10 + sdk/cliproxy/service_plugins.go | 26 +- 8 files changed, 883 insertions(+), 43 deletions(-) diff --git a/internal/api/handlers/management/auth_files_fields.go b/internal/api/handlers/management/auth_files_fields.go index fb6404f6c..856827d4d 100644 --- a/internal/api/handlers/management/auth_files_fields.go +++ b/internal/api/handlers/management/auth_files_fields.go @@ -53,7 +53,12 @@ func (h *Handler) PatchAuthFileStatus(c *gin.Context) { } h.authStatusMu.Lock() - defer h.authStatusMu.Unlock() + locked := true + defer func() { + if locked { + h.authStatusMu.Unlock() + } + }() ctx := c.Request.Context() @@ -69,7 +74,8 @@ func (h *Handler) PatchAuthFileStatus(c *gin.Context) { c.JSON(http.StatusConflict, gin.H{"error": errPluginVirtualAuth.Error()}) return } - if errPatch := h.patchPluginVirtualSourceStatus(ctx, targetAuth, *req.Disabled); errPatch != nil { + hookAuths, errPatch := h.patchPluginVirtualSourceStatus(ctx, targetAuth, *req.Disabled) + if errPatch != nil { status := http.StatusInternalServerError if errors.Is(errPatch, errAuthFileNotFound) || os.IsNotExist(errPatch) { status = http.StatusNotFound @@ -77,6 +83,13 @@ func (h *Handler) PatchAuthFileStatus(c *gin.Context) { c.JSON(status, gin.H{"error": errPatch.Error()}) return } + locked = false + h.authStatusMu.Unlock() + if errHook := h.invokePostAuthPersistHooks(ctx, hookAuths); errHook != nil { + log.Errorf("post-auth persist hook failed for plugin virtual source status update on %s: %v", targetAuth.ID, errHook) + c.JSON(http.StatusInternalServerError, gin.H{"error": fmt.Sprintf("failed to synchronize plugin virtual auth: %v", errHook)}) + return + } c.JSON(http.StatusOK, gin.H{"status": "ok", "disabled": *req.Disabled}) return } @@ -118,16 +131,16 @@ func (h *Handler) PatchAuthFileStatus(c *gin.Context) { c.JSON(http.StatusInternalServerError, gin.H{"error": fmt.Sprintf("failed to update auth: %v", err)}) return } - if h.postAuthPersistHook != nil { - hookAuth := updatedAuth - if hookAuth == nil { - hookAuth = targetAuth - } - if errHook := h.postAuthPersistHook(ctx, hookAuth); errHook != nil { - log.Errorf("post-auth persist hook failed for status update on %s: %v", targetAuth.ID, errHook) - c.JSON(http.StatusInternalServerError, gin.H{"error": fmt.Sprintf("failed to synchronize auth runtime: %v", errHook)}) - return - } + hookAuth := updatedAuth + if hookAuth == nil { + hookAuth = targetAuth + } + locked = false + h.authStatusMu.Unlock() + if errHook := h.invokePostAuthPersistHooks(ctx, []*coreauth.Auth{hookAuth}); errHook != nil { + log.Errorf("post-auth persist hook failed for status update on %s: %v", targetAuth.ID, errHook) + c.JSON(http.StatusInternalServerError, gin.H{"error": fmt.Sprintf("failed to synchronize auth runtime: %v", errHook)}) + return } c.JSON(http.StatusOK, gin.H{"status": "ok", "disabled": *req.Disabled}) @@ -135,24 +148,25 @@ func (h *Handler) PatchAuthFileStatus(c *gin.Context) { // patchPluginVirtualSourceStatus toggles disabled on a plugin multi-auth source file and all // runtime auths expanded from it. Virtual project children cannot be toggled independently. -func (h *Handler) patchPluginVirtualSourceStatus(ctx context.Context, targetAuth *coreauth.Auth, disabled bool) error { +func (h *Handler) patchPluginVirtualSourceStatus(ctx context.Context, targetAuth *coreauth.Auth, disabled bool) ([]*coreauth.Auth, error) { if h == nil || h.authManager == nil || targetAuth == nil { - return fmt.Errorf("core auth manager unavailable") + return nil, fmt.Errorf("core auth manager unavailable") } sourcePath := strings.TrimSpace(authAttribute(targetAuth, coreauth.AttributeVirtualSource)) if sourcePath == "" { sourcePath = strings.TrimSpace(authAttribute(targetAuth, "path")) } if sourcePath == "" { - return errPluginVirtualAuth + return nil, errPluginVirtualAuth } if errWrite := setSourceAuthFileDisabled(sourcePath, disabled); errWrite != nil { if os.IsNotExist(errWrite) { - return errAuthFileNotFound + return nil, errAuthFileNotFound } - return fmt.Errorf("failed to update source auth file: %w", errWrite) + return nil, fmt.Errorf("failed to update source auth file: %w", errWrite) } now := time.Now() + hookAuths := make([]*coreauth.Auth, 0) for _, auth := range h.authManager.List() { if auth == nil { continue @@ -165,17 +179,27 @@ func (h *Handler) patchPluginVirtualSourceStatus(ctx context.Context, targetAuth auth.UpdatedAt = now updated, errUpdate := h.authManager.Update(ctx, auth) if errUpdate != nil { - return fmt.Errorf("failed to update auth %s: %w", auth.ID, errUpdate) + return nil, fmt.Errorf("failed to update auth %s: %w", auth.ID, errUpdate) } - if h.postAuthPersistHook != nil { - hookAuth := updated - if hookAuth == nil { - hookAuth = auth - } - if errHook := h.postAuthPersistHook(ctx, hookAuth); errHook != nil { - log.Errorf("post-auth persist hook failed for plugin virtual auth %s: %v", auth.ID, errHook) - return fmt.Errorf("failed to synchronize plugin virtual auth %s: %w", auth.ID, errHook) - } + hookAuth := updated + if hookAuth == nil { + hookAuth = auth + } + hookAuths = append(hookAuths, hookAuth) + } + return hookAuths, nil +} + +func (h *Handler) invokePostAuthPersistHooks(ctx context.Context, auths []*coreauth.Auth) error { + if h == nil || h.postAuthPersistHook == nil { + return nil + } + for _, auth := range auths { + if auth == nil { + continue + } + if errHook := h.postAuthPersistHook(ctx, auth); errHook != nil { + return errHook } } return nil diff --git a/internal/api/handlers/management/auth_files_status_sync_test.go b/internal/api/handlers/management/auth_files_status_sync_test.go index 7e1c40722..de03368e2 100644 --- a/internal/api/handlers/management/auth_files_status_sync_test.go +++ b/internal/api/handlers/management/auth_files_status_sync_test.go @@ -9,7 +9,9 @@ import ( "os" "path/filepath" "strings" + "sync" "testing" + "time" "github.com/gin-gonic/gin" "github.com/router-for-me/CLIProxyAPI/v7/internal/config" @@ -90,6 +92,118 @@ func TestPatchAuthFileStatusInvokesPostAuthPersistHook(t *testing.T) { } } +func TestPatchAuthFileStatusDoesNotHoldLockAcrossPersistHook(t *testing.T) { + t.Setenv("MANAGEMENT_PASSWORD", "") + + authDir := t.TempDir() + manager := coreauth.NewManager(nil, nil, nil) + for _, fileName := range []string{"codex-status-a.json", "codex-status-b.json"} { + filePath := filepath.Join(authDir, fileName) + if errWrite := os.WriteFile(filePath, []byte(`{"type":"codex","disabled":false}`), 0o600); errWrite != nil { + t.Fatalf("write auth file %s: %v", fileName, errWrite) + } + auth := &coreauth.Auth{ + ID: fileName, + FileName: fileName, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "path": filePath, + }, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth %s: %v", fileName, errRegister) + } + } + + h := NewHandlerWithoutConfigFilePath(&config.Config{AuthDir: authDir}, manager) + + hookStarted := make(chan struct{}) + unblock := make(chan struct{}) + var hookMu sync.Mutex + hookCalls := 0 + h.SetPostAuthPersistHook(func(_ context.Context, _ *coreauth.Auth) error { + hookMu.Lock() + hookCalls++ + n := hookCalls + hookMu.Unlock() + if n == 1 { + close(hookStarted) + <-unblock + } + return nil + }) + + var codeA int + finishedA := make(chan struct{}) + go func() { + defer close(finishedA) + rec := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(rec) + req := httptest.NewRequest(http.MethodPatch, "/v0/management/auth-files/status", strings.NewReader(`{"name":"codex-status-a.json","disabled":true}`)) + req.Header.Set("Content-Type", "application/json") + ctx.Request = req + h.PatchAuthFileStatus(ctx) + codeA = rec.Code + }() + + var releaseOnce sync.Once + releaseHook := func() { + releaseOnce.Do(func() { + close(unblock) + }) + } + t.Cleanup(func() { + releaseHook() + select { + case <-finishedA: + case <-time.After(2 * time.Second): + } + }) + + select { + case <-hookStarted: + case <-time.After(2 * time.Second): + t.Fatal("first persist hook did not start") + } + + doneB := make(chan struct { + code int + body string + }, 1) + go func() { + rec := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(rec) + req := httptest.NewRequest(http.MethodPatch, "/v0/management/auth-files/status", strings.NewReader(`{"name":"codex-status-b.json","disabled":true}`)) + req.Header.Set("Content-Type", "application/json") + ctx.Request = req + h.PatchAuthFileStatus(ctx) + doneB <- struct { + code int + body string + }{code: rec.Code, body: rec.Body.String()} + }() + + select { + case got := <-doneB: + if got.code != http.StatusOK { + t.Fatalf("second PATCH status = %d, want %d body=%s", got.code, http.StatusOK, got.body) + } + case <-time.After(2 * time.Second): + t.Fatal("second PATCH blocked by authStatusMu held across persist hook") + } + + releaseHook() + select { + case <-finishedA: + if codeA != http.StatusOK { + t.Fatalf("first PATCH status = %d, want %d", codeA, http.StatusOK) + } + case <-time.After(2 * time.Second): + t.Fatal("first PATCH did not finish after hook was unblocked") + } +} + func TestPatchAuthFileStatusRestoresModelsViaSyncHook(t *testing.T) { t.Setenv("MANAGEMENT_PASSWORD", "") diff --git a/sdk/cliproxy/builder.go b/sdk/cliproxy/builder.go index b4e1cecbb..19b6d37a6 100644 --- a/sdk/cliproxy/builder.go +++ b/sdk/cliproxy/builder.go @@ -328,7 +328,8 @@ func (s *Service) runtimeAuthSyncHook() coreauth.PostAuthHook { } } // Detach from request cancellation so runtime model registration always completes - // once the credential has been persisted to disk. + // once the credential has been persisted to disk. If the watcher consumer already + // claimed this revision, handleAuthUpdate waits for that registration to finish. syncCtx := coreauth.WithSkipPersist(context.Background()) s.handleAuthUpdate(syncCtx, update) return nil diff --git a/sdk/cliproxy/service.go b/sdk/cliproxy/service.go index 331325ed8..840c7a368 100644 --- a/sdk/cliproxy/service.go +++ b/sdk/cliproxy/service.go @@ -40,6 +40,8 @@ type Service struct { executorRegistrationMu sync.Mutex authUpdateMu sync.Mutex authRevisions map[string]uint64 + authRegWaitMu sync.Mutex + authRegWaiters map[string]chan struct{} configSequence uint64 appliedRoutingState *routingRuntimeState diff --git a/sdk/cliproxy/service_auth.go b/sdk/cliproxy/service_auth.go index cc7e5782b..00315b406 100644 --- a/sdk/cliproxy/service_auth.go +++ b/sdk/cliproxy/service_auth.go @@ -95,14 +95,25 @@ func (s *Service) handleAuthUpdates(ctx context.Context, updates []watcher.AuthU if s == nil { return } + // Keep this path limited to the targeted auths. Global plugin rebuilds and + // in-flight Antigravity probes can hold authUpdateMu for minutes and stall + // Management API PATCH responses. s.authUpdateMu.Lock() - defer s.authUpdateMu.Unlock() + locked := true + defer func() { + if locked { + s.authUpdateMu.Unlock() + } + }() if s.authRevisions == nil { s.authRevisions = make(map[string]uint64) } filtered := make([]watcher.AuthUpdate, 0, len(updates)) + skippedWaits := make([]chan struct{}, 0) + startedRegs := make([]authRegistrationWait, 0) + startedWaitByID := make(map[string]authRegistrationWait) for _, update := range updates { id := authUpdateID(update) if id == "" { @@ -113,13 +124,28 @@ func (s *Service) handleAuthUpdates(ctx context.Context, updates []watcher.AuthU if rev > 0 { if prevRev, exists := s.authRevisions[id]; exists && rev <= prevRev { log.Debugf("skipping stale auth update for %s: rev %d <= processed %d", id, rev, prevRev) + if ch := s.authRegistrationWaitCh(id); ch != nil { + skippedWaits = append(skippedWaits, ch) + } continue } s.authRevisions[id] = rev } + wait := s.beginAuthRegistration(id) + startedRegs = append(startedRegs, wait) + startedWaitByID[id] = wait filtered = append(filtered, update) } + registrationsFinished := false + defer func() { + if !registrationsFinished { + finishAuthRegistrations(s, startedRegs) + } + }() if len(filtered) == 0 { + locked = false + s.authUpdateMu.Unlock() + waitAuthRegistrations(skippedWaits) return } updates = coalesceAuthUpdates(filtered) @@ -127,12 +153,14 @@ func (s *Service) handleAuthUpdates(ctx context.Context, updates []watcher.AuthU cfg := s.cfg s.cfgMu.RUnlock() if cfg == nil || s.coreManager == nil { + locked = false + s.authUpdateMu.Unlock() + waitAuthRegistrations(skippedWaits) return } registrationCtx := coreauth.WithDeferredAPIKeyModelAliasRebuild(ctx) tasks := make([]modelRegistrationTask, 0, len(updates)) - needsPluginSync := false needsAliasRebuild := false for _, update := range updates { switch update.Action { @@ -146,14 +174,25 @@ func (s *Service) handleAuthUpdates(ctx context.Context, updates []watcher.AuthU } needsAliasRebuild = true authForRegistration := auth + expectedGeneration := authForRegistration.Generation + expectedDisabled := authForRegistration.Disabled + authID := authForRegistration.ID + wait := startedWaitByID[authID] tasks = append(tasks, modelRegistrationTask{ phase: modelRegistrationPhase(authForRegistration), category: modelRegistrationCategory(authForRegistration), run: func(compatCache *openAICompatibilityRegistrationCache) { + if s.shouldSkipModelRegistration(authID, expectedGeneration, expectedDisabled) { + return + } s.completeModelRegistrationForAuthWithCache(registrationCtx, authForRegistration, compatCache) }, + done: func() { + if wait.ch != nil { + finishAuthRegistrations(s, []authRegistrationWait{wait}) + } + }, }) - needsPluginSync = true case watcher.AuthUpdateActionDelete: id := update.ID if id == "" && update.Auth != nil { @@ -178,10 +217,12 @@ func (s *Service) handleAuthUpdates(ctx context.Context, updates []watcher.AuthU if needsAliasRebuild { s.coreManager.RefreshAPIKeyModelAlias() } + locked = false + s.authUpdateMu.Unlock() s.runModelRegistrationTasks(registrationCtx, tasks) - if needsPluginSync { - s.syncPluginRuntime(registrationCtx) - } + finishAuthRegistrations(s, startedRegs) + registrationsFinished = true + waitAuthRegistrations(skippedWaits) } func coalesceAuthUpdates(updates []watcher.AuthUpdate) []watcher.AuthUpdate { @@ -300,7 +341,6 @@ func (s *Service) applyCoreAuthAddOrUpdate(ctx context.Context, auth *coreauth.A return } s.completeModelRegistrationForAuth(ctx, auth) - s.syncPluginRuntime(ctx) } func (s *Service) prepareCoreAuthForModelRegistration(ctx context.Context, auth *coreauth.Auth) *coreauth.Auth { @@ -345,8 +385,91 @@ func (s *Service) prepareCoreAuthForModelRegistration(ctx context.Context, auth return auth } +type authRegistrationWait struct { + id string + ch chan struct{} +} + +func (s *Service) authRegistrationWaitCh(id string) chan struct{} { + if s == nil { + return nil + } + s.authRegWaitMu.Lock() + defer s.authRegWaitMu.Unlock() + if s.authRegWaiters == nil { + return nil + } + return s.authRegWaiters[id] +} + +func (s *Service) beginAuthRegistration(id string) authRegistrationWait { + ch := make(chan struct{}) + if s == nil || id == "" { + close(ch) + return authRegistrationWait{id: id, ch: ch} + } + s.authRegWaitMu.Lock() + if s.authRegWaiters == nil { + s.authRegWaiters = make(map[string]chan struct{}) + } + s.authRegWaiters[id] = ch + s.authRegWaitMu.Unlock() + return authRegistrationWait{id: id, ch: ch} +} + +func finishAuthRegistrations(s *Service, waits []authRegistrationWait) { + if len(waits) == 0 { + return + } + if s != nil { + s.authRegWaitMu.Lock() + for _, wait := range waits { + if wait.ch == nil { + continue + } + if s.authRegWaiters != nil && s.authRegWaiters[wait.id] == wait.ch { + delete(s.authRegWaiters, wait.id) + } + } + s.authRegWaitMu.Unlock() + } + for _, wait := range waits { + if wait.ch == nil { + continue + } + select { + case <-wait.ch: + default: + close(wait.ch) + } + } +} + +func waitAuthRegistrations(waits []chan struct{}) { + for _, ch := range waits { + if ch == nil { + continue + } + <-ch + } +} + +func (s *Service) shouldSkipModelRegistration(authID string, expectedGeneration uint64, expectedDisabled bool) bool { + if s == nil || s.coreManager == nil || strings.TrimSpace(authID) == "" { + return true + } + current, ok := s.coreManager.GetByID(authID) + if !ok || current == nil { + return !expectedDisabled + } + if expectedGeneration > 0 && current.Generation > expectedGeneration { + return true + } + return current.Disabled != expectedDisabled +} + // isStaleCoreAuth reports whether an incoming auth update is older than the current -// state in coreManager, based on registration epoch. +// state in coreManager, based on registration epoch and generation. func isStaleCoreAuth(existing, incoming *coreauth.Auth) bool { if existing == nil || incoming == nil { return false @@ -355,6 +478,11 @@ func isStaleCoreAuth(existing, incoming *coreauth.Auth) bool { if incoming.RegistrationEpoch > 0 && incoming.RegistrationEpoch < existing.RegistrationEpoch { return true } + // Versioned snapshots with an older generation must not overwrite a newer persist. + // Generation 0 is left unversioned (file-watcher synthesizer snapshots). + if incoming.Generation > 0 && incoming.Generation < existing.Generation { + return true + } return false } @@ -402,7 +530,6 @@ func (s *Service) applyCoreAuthRemoval(ctx context.Context, id string) { if strings.EqualFold(provider, "xai") { executor.CloseXAIWebsocketSessionsForAuthID(id, "auth_removed") } - s.syncPluginRuntime(ctx) } func (s *Service) applyRetryConfig(cfg *config.Config) { diff --git a/sdk/cliproxy/service_auth_sync_test.go b/sdk/cliproxy/service_auth_sync_test.go index 82ef56d25..d5c9cfc01 100644 --- a/sdk/cliproxy/service_auth_sync_test.go +++ b/sdk/cliproxy/service_auth_sync_test.go @@ -11,12 +11,14 @@ import ( "os" "path/filepath" "strings" + "sync/atomic" "testing" "time" "golang.org/x/crypto/bcrypt" "github.com/router-for-me/CLIProxyAPI/v7/internal/api" + "github.com/router-for-me/CLIProxyAPI/v7/internal/pluginhost" internalregistry "github.com/router-for-me/CLIProxyAPI/v7/internal/registry" "github.com/router-for-me/CLIProxyAPI/v7/internal/watcher" coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" @@ -227,6 +229,548 @@ func TestRuntimeAuthSyncHook_NilAndEmptyAuth(t *testing.T) { } } +func TestRuntimeAuthSyncHook_DoesNotBlockOnUnrelatedAntigravityProbes(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authID := "codex-probe-wait-auth" + reg.UnregisterClient(authID) + t.Cleanup(func() { + reg.UnregisterClient(authID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + auth := &coreauth.Auth{ + ID: authID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-probe-wait.json", + }, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + pluginHost: pluginhost.New(), + } + service.antigravityProbeWg.Add(1) + done := make(chan error, 1) + finished := make(chan struct{}) + t.Cleanup(func() { + service.antigravityProbeWg.Done() + select { + case <-finished: + case <-time.After(2 * time.Second): + } + }) + + hook := service.runtimeAuthSyncHook() + if hook == nil { + t.Fatal("runtimeAuthSyncHook() returned nil") + } + + go func() { + defer close(finished) + done <- hook(context.Background(), auth) + }() + + select { + case errHook := <-done: + if errHook != nil { + t.Fatalf("hook failed: %v", errHook) + } + case <-time.After(2 * time.Second): + t.Fatal("runtimeAuthSyncHook blocked waiting for unrelated antigravity probes") + } + + models := reg.GetModelsForClient(authID) + if len(models) == 0 { + t.Fatal("expected models registered for enabled auth, got none") + } +} + +func TestHandleAuthUpdates_DeleteDoesNotBlockOnUnrelatedAntigravityProbes(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authID := "codex-probe-wait-delete-auth" + reg.UnregisterClient(authID) + t.Cleanup(func() { + reg.UnregisterClient(authID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + auth := &coreauth.Auth{ + ID: authID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-probe-wait-delete.json", + }, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + pluginHost: pluginhost.New(), + } + service.registerModelsForAuth(context.Background(), auth) + if len(reg.GetModelsForClient(authID)) == 0 { + t.Fatal("expected models registered before delete") + } + + runAuthUpdateWithoutProbeWait(t, service, func() { + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionDelete, + ID: authID, + Auth: auth, + }) + }) + + if _, ok := manager.GetByID(authID); ok { + t.Fatal("expected auth to be removed") + } + if len(reg.GetModelsForClient(authID)) != 0 { + t.Fatal("expected models unregistered after delete") + } +} + +func TestHandleAuthUpdates_PluginVirtualModifyDoesNotBlockOnUnrelatedAntigravityProbes(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authID := "plugin-virtual-probe-wait-auth" + reg.UnregisterClient(authID) + t.Cleanup(func() { + reg.UnregisterClient(authID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + auth := &coreauth.Auth{ + ID: authID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/plugin-virtual-probe-wait.json", + }, + } + coreauth.MarkPluginVirtualAuth(auth, "/path/to/plugin-source.json", 0) + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + pluginHost: pluginhost.New(), + } + + runAuthUpdateWithoutProbeWait(t, service, func() { + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authID, + Auth: auth, + }) + }) +} + +func runAuthUpdateWithoutProbeWait(t *testing.T, service *Service, run func()) { + t.Helper() + if service == nil { + t.Fatal("service is nil") + } + service.antigravityProbeWg.Add(1) + done := make(chan struct{}) + finished := make(chan struct{}) + t.Cleanup(func() { + service.antigravityProbeWg.Done() + select { + case <-finished: + case <-time.After(2 * time.Second): + } + }) + + go func() { + defer close(finished) + run() + close(done) + }() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("auth update blocked waiting for unrelated antigravity probes") + } +} + +func TestHandleAuthUpdates_ModelRegistrationDoesNotHoldAuthUpdateLock(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authAID := "codex-lock-a-auth" + authBID := "codex-lock-b-auth" + reg.UnregisterClient(authAID) + reg.UnregisterClient(authBID) + t.Cleanup(func() { + reg.UnregisterClient(authAID) + reg.UnregisterClient(authBID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + authA := &coreauth.Auth{ + ID: authAID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-lock-a.json", + }, + } + authB := &coreauth.Auth{ + ID: authBID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-lock-b.json", + }, + } + for _, auth := range []*coreauth.Auth{authA, authB} { + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth %s: %v", auth.ID, errRegister) + } + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + } + + started := make(chan struct{}) + block := make(chan struct{}) + var first atomic.Bool + modelRegistrationTaskHook = func() { + if first.CompareAndSwap(false, true) { + close(started) + <-block + } + } + t.Cleanup(func() { + modelRegistrationTaskHook = nil + select { + case <-block: + default: + close(block) + } + }) + + finishedA := make(chan struct{}) + go func() { + defer close(finishedA) + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authAID, + Auth: authA, + }) + }() + + select { + case <-started: + case <-time.After(2 * time.Second): + t.Fatal("first auth update did not reach model registration") + } + + doneB := make(chan struct{}) + go func() { + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authBID, + Auth: authB, + }) + close(doneB) + }() + + select { + case <-doneB: + case <-time.After(2 * time.Second): + t.Fatal("second auth update blocked by first auth model registration") + } + + close(block) + select { + case <-finishedA: + case <-time.After(2 * time.Second): + t.Fatal("first auth update did not finish after registration was unblocked") + } +} + +func TestHandleAuthUpdates_OlderGenerationDoesNotOverrideNewerPersist(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authID := "codex-generation-stale-auth" + reg.UnregisterClient(authID) + t.Cleanup(func() { + reg.UnregisterClient(authID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + auth := &coreauth.Auth{ + ID: authID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-generation-stale.json", + }, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + } + + disabledAuth := auth.Clone() + disabledAuth.Disabled = true + disabledAuth.Status = coreauth.StatusDisabled + disabledSnapshot, errDisable := manager.Update(context.Background(), disabledAuth) + if errDisable != nil { + t.Fatalf("disable update: %v", errDisable) + } + + enabledAuth := disabledSnapshot.Clone() + enabledAuth.Disabled = false + enabledAuth.Status = coreauth.StatusActive + enabledSnapshot, errEnable := manager.Update(context.Background(), enabledAuth) + if errEnable != nil { + t.Fatalf("enable update: %v", errEnable) + } + + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authID, + Auth: enabledSnapshot, + }) + manager.RegisterExecutor(&syncTestExecutor{}) + if len(reg.GetModelsForClient(authID)) == 0 { + t.Fatal("expected models after enable") + } + + staleDisable := watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authID, + Auth: disabledSnapshot, + } + staleDisable.SetRevision(99) + service.handleAuthUpdate(context.Background(), staleDisable) + manager.RegisterExecutor(&syncTestExecutor{}) + + current, ok := manager.GetByID(authID) + if !ok || current == nil || current.Disabled { + t.Fatal("expected newer enabled persist to win over older disable snapshot") + } + if len(reg.GetModelsForClient(authID)) == 0 { + t.Fatal("expected models to remain registered after stale disable update") + } +} + +func TestHandleAuthUpdates_StaleDisableRegistrationDoesNotDropNewerEnable(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authID := "codex-stale-disable-task-auth" + reg.UnregisterClient(authID) + t.Cleanup(func() { + reg.UnregisterClient(authID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + auth := &coreauth.Auth{ + ID: authID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{ + "plan_type": "pro", + "path": "/path/to/codex-stale-disable-task.json", + }, + } + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth: %v", errRegister) + } + + service := &Service{ + cfg: &config.Config{}, + coreManager: manager, + } + + disabledAuth := auth.Clone() + disabledAuth.Disabled = true + disabledAuth.Status = coreauth.StatusDisabled + disabledSnapshot, errDisable := manager.Update(context.Background(), disabledAuth) + if errDisable != nil { + t.Fatalf("disable update: %v", errDisable) + } + + started := make(chan struct{}) + block := make(chan struct{}) + var first atomic.Bool + modelRegistrationTaskHook = func() { + if first.CompareAndSwap(false, true) { + close(started) + <-block + } + } + t.Cleanup(func() { + modelRegistrationTaskHook = nil + select { + case <-block: + default: + close(block) + } + }) + + finishedDisable := make(chan struct{}) + go func() { + defer close(finishedDisable) + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authID, + Auth: disabledSnapshot, + }) + }() + + select { + case <-started: + case <-time.After(2 * time.Second): + t.Fatal("disable registration did not start") + } + + enabledAuth := disabledSnapshot.Clone() + enabledAuth.Disabled = false + enabledAuth.Status = coreauth.StatusActive + enabledSnapshot, errEnable := manager.Update(context.Background(), enabledAuth) + if errEnable != nil { + t.Fatalf("enable update: %v", errEnable) + } + service.handleAuthUpdate(context.Background(), watcher.AuthUpdate{ + Action: watcher.AuthUpdateActionModify, + ID: authID, + Auth: enabledSnapshot, + }) + manager.RegisterExecutor(&syncTestExecutor{}) + if len(reg.GetModelsForClient(authID)) == 0 { + t.Fatal("expected models after enable") + } + + close(block) + select { + case <-finishedDisable: + case <-time.After(2 * time.Second): + t.Fatal("disable registration did not finish") + } + + current, ok := manager.GetByID(authID) + if !ok || current == nil || current.Disabled { + t.Fatal("expected auth to remain enabled after stale disable registration") + } + if len(reg.GetModelsForClient(authID)) == 0 { + t.Fatal("expected models to remain registered after stale disable registration") + } +} + +func TestHandleAuthUpdates_SameRevisionWaitDoesNotWaitForOtherAuthInBatch(t *testing.T) { + reg := internalregistry.GetGlobalRegistry() + authAID := "codex-batch-wait-a-auth" + authBID := "codex-batch-wait-b-auth" + reg.UnregisterClient(authAID) + reg.UnregisterClient(authBID) + t.Cleanup(func() { + reg.UnregisterClient(authAID) + reg.UnregisterClient(authBID) + }) + + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(&syncTestExecutor{}) + authA := &coreauth.Auth{ + ID: authAID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{"plan_type": "pro", "path": "/path/to/codex-batch-wait-a.json"}, + } + authB := &coreauth.Auth{ + ID: authBID, + Provider: "codex", + Status: coreauth.StatusActive, + Attributes: map[string]string{"plan_type": "pro", "path": "/path/to/codex-batch-wait-b.json"}, + } + for _, auth := range []*coreauth.Auth{authA, authB} { + if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil { + t.Fatalf("register auth %s: %v", auth.ID, errRegister) + } + } + + service := &Service{cfg: &config.Config{}, coreManager: manager} + + bStarted := make(chan struct{}) + bBlock := make(chan struct{}) + var started atomic.Int32 + modelRegistrationTaskHook = func() { + if started.Add(1) == 2 { + close(bStarted) + <-bBlock + } + } + t.Cleanup(func() { + modelRegistrationTaskHook = nil + select { + case <-bBlock: + default: + close(bBlock) + } + }) + + updateA := watcher.AuthUpdate{Action: watcher.AuthUpdateActionModify, ID: authAID, Auth: authA} + updateA.SetRevision(1) + updateB := watcher.AuthUpdate{Action: watcher.AuthUpdateActionModify, ID: authBID, Auth: authB} + updateB.SetRevision(1) + + finishedBatch := make(chan struct{}) + go func() { + defer close(finishedBatch) + service.handleAuthUpdates(context.Background(), []watcher.AuthUpdate{updateA, updateB}) + }() + + select { + case <-bStarted: + case <-time.After(2 * time.Second): + t.Fatal("second auth registration in batch did not start") + } + + doneA := make(chan struct{}) + go func() { + service.handleAuthUpdate(context.Background(), updateA) + close(doneA) + }() + select { + case <-doneA: + case <-time.After(2 * time.Second): + t.Fatal("auth A hook wait blocked on unrelated auth B registration") + } + + close(bBlock) + select { + case <-finishedBatch: + case <-time.After(2 * time.Second): + t.Fatal("batch auth update did not finish") + } +} + func TestEndToEndStatusPatch_RestoresModelsThroughServerPipeline(t *testing.T) { authDir := t.TempDir() fileName := "codex-e2e.json" @@ -626,6 +1170,10 @@ func TestRuntimeAuthSync_FileSnapshotEnqueuedThenRequestExecutionThenConsume(t * fileSnapshotAuth.Disabled = true fileSnapshotAuth.Status = coreauth.StatusDisabled fileSnapshotAuth.UpdatedAt = time.Now() + // File synthesizer snapshots are unversioned; keep this fixture aligned so an + // intermediate MarkResult generation bump cannot drop a newer watcher revision. + fileSnapshotAuth.RegistrationEpoch = 0 + fileSnapshotAuth.Generation = 0 fileUpdate := watcher.AuthUpdate{ Action: watcher.AuthUpdateActionModify, diff --git a/sdk/cliproxy/service_models.go b/sdk/cliproxy/service_models.go index eac895f89..603e92ac0 100644 --- a/sdk/cliproxy/service_models.go +++ b/sdk/cliproxy/service_models.go @@ -29,9 +29,19 @@ func (s *Service) registerModelsForAuthWithCache(ctx context.Context, a *coreaut return } if a.Disabled { + if s != nil && s.coreManager != nil { + if current, ok := s.coreManager.GetByID(a.ID); ok && current != nil && !current.Disabled { + return + } + } GlobalModelRegistry().UnregisterClient(a.ID) return } + if s != nil && s.coreManager != nil { + if current, ok := s.coreManager.GetByID(a.ID); !ok || current == nil || current.Disabled { + return + } + } authKind := a.AuthKind() // Unregister legacy client ID (if present) to avoid double counting if a.Runtime != nil { diff --git a/sdk/cliproxy/service_plugins.go b/sdk/cliproxy/service_plugins.go index 4f783ea3b..80884fe3e 100644 --- a/sdk/cliproxy/service_plugins.go +++ b/sdk/cliproxy/service_plugins.go @@ -32,6 +32,7 @@ type modelRegistrationTask struct { phase int category string run func(*openAICompatibilityRegistrationCache) + done func() } type executorRegistrationOptions struct { @@ -48,6 +49,11 @@ var registerPluginExecutors = func(host *pluginhost.Host, manager *coreauth.Mana host.RegisterExecutors(manager, registry.GetGlobalRegistry()) } +// modelRegistrationTaskHook, if set, runs after auth-update commits and before +// model registration workers start. Tests use it to prove registration no longer +// holds authUpdateMu. +var modelRegistrationTaskHook func() + // RegisterUsagePlugin registers a usage plugin on the global usage manager. // This allows external code to monitor API usage and token consumption. // @@ -244,12 +250,20 @@ func (s *Service) runModelRegistrationTaskPhase(ctx context.Context, tasks []mod go func() { defer wg.Done() for task := range taskCh { - select { - case <-ctx.Done(): - return - default: - } - task.run(compatCache) + func(task modelRegistrationTask) { + if task.done != nil { + defer task.done() + } + select { + case <-ctx.Done(): + return + default: + } + if modelRegistrationTaskHook != nil { + modelRegistrationTaskHook() + } + task.run(compatCache) + }(task) } }() }