mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-10-06 15:50:49 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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", "")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user