feat(auth): add tests to ensure async operations don’t block unrelated actions

- 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
This commit is contained in:
Luis Pater
2026-09-15 01:23:26 +08:00
parent 748d576731
commit e3cbe437d0
8 changed files with 883 additions and 43 deletions

View File

@@ -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

View File

@@ -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", "")

View File

@@ -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

View File

@@ -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

View File

@@ -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) {

View File

@@ -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,

View File

@@ -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 {

View File

@@ -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)
}
}()
}