diff --git a/sdk/cliproxy/auth/conductor.go b/sdk/cliproxy/auth/conductor.go index 52d4c2c78..f6a123935 100644 --- a/sdk/cliproxy/auth/conductor.go +++ b/sdk/cliproxy/auth/conductor.go @@ -260,6 +260,9 @@ type Manager struct { refreshLoop *authAutoRefreshLoop requestPrepareLocks sync.Map + // refreshLocks serializes credential refresh per auth ID so concurrent + // 401 recoveries and auto-refresh workers do not race the same refresh_token. + refreshLocks sync.Map } // NewManager constructs a manager with optional custom selector and hook. @@ -1829,6 +1832,7 @@ func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor Provi } ctx = contextWithRequestedModelAlias(ctx, opts, routeModel) var lastErr error + didRefreshOnUnauthorized := false for idx, execModel := range execModels { resultModel := m.stateModelForExecution(auth, routeModel, execModel, pooled) execReq := req @@ -1843,6 +1847,18 @@ func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor Provi if errCtx := ctx.Err(); errCtx != nil { return nil, errCtx } + if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(ctx, auth, errStream, didRefreshOnUnauthorized); okRefresh { + auth = refreshed + didRefreshOnUnauthorized = true + streamResult, errStream = executor.ExecuteStream(ctx, auth, execReq, execOpts) + if errStream != nil { + if errCtx := ctx.Err(); errCtx != nil { + return nil, errCtx + } + } + } + } + if errStream != nil { rerr := &Error{Message: errStream.Error()} if se, ok := errors.AsType[cliproxyexecutor.StatusError](errStream); ok && se != nil { rerr.HTTPStatus = se.StatusCode() @@ -1863,6 +1879,24 @@ func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor Provi discardStreamChunks(streamResult.Chunks) return nil, errCtx } + if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(ctx, auth, bootstrapErr, didRefreshOnUnauthorized); okRefresh { + discardStreamChunks(streamResult.Chunks) + auth = refreshed + didRefreshOnUnauthorized = true + retryStream, retryErr := executor.ExecuteStream(ctx, auth, execReq, execOpts) + if retryErr != nil { + if errCtx := ctx.Err(); errCtx != nil { + return nil, errCtx + } + bootstrapErr = retryErr + streamResult = &cliproxyexecutor.StreamResult{} + } else { + streamResult = retryStream + buffered, closed, bootstrapErr = readStreamBootstrap(ctx, streamResult.Chunks) + } + } + } + if bootstrapErr != nil { if isRequestInvalidError(bootstrapErr) { rerr := &Error{Message: bootstrapErr.Error()} if se, ok := errors.AsType[cliproxyexecutor.StatusError](bootstrapErr); ok && se != nil { @@ -2542,6 +2576,7 @@ func (m *Manager) executeMixedOnce(ctx context.Context, providers []string, req continue } var authErr error + didRefreshOnUnauthorized := false for _, upstreamModel := range models { resultModel := m.stateModelForExecution(auth, routeModel, upstreamModel, pooled) execReq := req @@ -2552,11 +2587,23 @@ func (m *Manager) executeMixedOnce(ctx context.Context, providers []string, req execOpts := opts execReq, execOpts = applyRequestAfterAuthInterceptor(execCtx, executor, provider, execReq, execOpts, requestedModelAliasFromOptions(execOpts, routeModel)) resp, errExec := executor.Execute(execCtx, auth, execReq, execOpts) - result := Result{AuthID: auth.ID, Provider: provider, Model: resultModel, Success: errExec == nil} if errExec != nil { if errCtx := execCtx.Err(); errCtx != nil { return cliproxyexecutor.Response{}, errCtx } + if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(execCtx, auth, errExec, didRefreshOnUnauthorized); okRefresh { + auth = refreshed + didRefreshOnUnauthorized = true + resp, errExec = executor.Execute(execCtx, auth, execReq, execOpts) + if errExec != nil { + if errCtx := execCtx.Err(); errCtx != nil { + return cliproxyexecutor.Response{}, errCtx + } + } + } + } + result := Result{AuthID: auth.ID, Provider: provider, Model: resultModel, Success: errExec == nil} + if errExec != nil { result.Error = &Error{Message: errExec.Error()} if se, ok := errors.AsType[cliproxyexecutor.StatusError](errExec); ok && se != nil { result.Error.HTTPStatus = se.StatusCode() @@ -2648,6 +2695,7 @@ func (m *Manager) executeCountMixedOnce(ctx context.Context, providers []string, continue } var authErr error + didRefreshOnUnauthorized := false for _, upstreamModel := range models { resultModel := m.stateModelForExecution(auth, routeModel, upstreamModel, pooled) execReq := req @@ -2658,11 +2706,23 @@ func (m *Manager) executeCountMixedOnce(ctx context.Context, providers []string, execOpts := opts execReq, execOpts = applyRequestAfterAuthInterceptor(execCtx, executor, provider, execReq, execOpts, requestedModelAliasFromOptions(execOpts, routeModel)) resp, errExec := executor.CountTokens(execCtx, auth, execReq, execOpts) - result := Result{AuthID: auth.ID, Provider: provider, Model: resultModel, Success: errExec == nil} if errExec != nil { if errCtx := execCtx.Err(); errCtx != nil { return cliproxyexecutor.Response{}, errCtx } + if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(execCtx, auth, errExec, didRefreshOnUnauthorized); okRefresh { + auth = refreshed + didRefreshOnUnauthorized = true + resp, errExec = executor.CountTokens(execCtx, auth, execReq, execOpts) + if errExec != nil { + if errCtx := execCtx.Err(); errCtx != nil { + return cliproxyexecutor.Response{}, errCtx + } + } + } + } + result := Result{AuthID: auth.ID, Provider: provider, Model: resultModel, Success: errExec == nil} + if errExec != nil { result.Error = &Error{Message: errExec.Error()} if se, ok := errors.AsType[cliproxyexecutor.StatusError](errExec); ok && se != nil { result.Error.HTTPStatus = se.StatusCode() @@ -5696,26 +5756,114 @@ func (m *Manager) markRefreshPending(id string, now time.Time) bool { return true } +type authRefreshLock struct { + mu sync.Mutex +} + +func authAccessToken(auth *Auth) string { + if token := authMetadataString(auth, "access_token"); token != "" { + return token + } + return authMetadataString(auth, "accessToken") +} + +func authHasRefreshCredential(auth *Auth) bool { + if authMetadataString(auth, "refresh_token") != "" { + return true + } + return authMetadataString(auth, "refreshToken") != "" +} + +func clearUnauthorizedModelStates(auth *Auth, now time.Time) []string { + if auth == nil || len(auth.ModelStates) == 0 { + return nil + } + var resumed []string + for model, state := range auth.ModelStates { + if state == nil || state.LastError == nil { + continue + } + if state.LastError.StatusCode() != http.StatusUnauthorized && !strings.EqualFold(state.LastError.Code, "unauthorized") { + continue + } + resetModelState(state, now) + resumed = append(resumed, model) + } + if len(resumed) > 0 { + updateAggregatedAvailability(auth, now) + } + return resumed +} + +// tryRefreshAfterUnauthorized refreshes OAuth credentials once after a 401 so the +// current auth can be retried before fallback/suspend. +func (m *Manager) tryRefreshAfterUnauthorized(ctx context.Context, auth *Auth, execErr error, alreadyTried bool) (*Auth, bool) { + if m == nil || auth == nil || alreadyTried || execErr == nil { + return auth, false + } + if !isUnauthorizedError(execErr) || !authHasRefreshCredential(auth) { + return auth, false + } + log.Debugf("unauthorized response for %s (%s), refreshing credentials before fallback", auth.Provider, auth.ID) + refreshed, errRefresh := m.refreshAuthForRequest(ctx, auth.ID, authAccessToken(auth)) + if errRefresh != nil || refreshed == nil { + log.Debugf("credential refresh before fallback failed for %s (%s): %v", auth.Provider, auth.ID, errRefresh) + return auth, false + } + return refreshed, true +} + func (m *Manager) refreshAuth(ctx context.Context, id string) { + _, _ = m.refreshAuthForRequest(ctx, id, "") +} + +// refreshAuthForRequest performs a synchronous credential refresh for the given auth. +// failedAccessToken lets concurrent callers reuse a refresh that already replaced the +// access token that produced the unauthorized response. +func (m *Manager) refreshAuthForRequest(ctx context.Context, id, failedAccessToken string) (*Auth, error) { + if m == nil { + return nil, errors.New("auth manager is nil") + } if ctx == nil { ctx = context.Background() } + id = strings.TrimSpace(id) + if id == "" { + return nil, errors.New("auth id is empty") + } + + lockValue, _ := m.refreshLocks.LoadOrStore(id, &authRefreshLock{}) + lock, _ := lockValue.(*authRefreshLock) + if lock == nil { + lock = &authRefreshLock{} + m.refreshLocks.Store(id, lock) + } + lock.mu.Lock() + defer lock.mu.Unlock() + m.mu.RLock() auth := m.auths[id] var exec ProviderExecutor - var cloned *Auth if auth != nil { exec = m.executors[auth.Provider] - cloned = auth.Clone() } m.mu.RUnlock() if auth == nil || exec == nil { - return + return nil, errors.New("auth or executor not found") } + + // Another request may already have refreshed this credential. + if failedAccessToken != "" { + if currentToken := authAccessToken(auth); currentToken != "" && currentToken != failedAccessToken { + return auth.Clone(), nil + } + } + + cloned := auth.Clone() updated, err := exec.Refresh(ctx, cloned) if err != nil && errors.Is(err, context.Canceled) { log.Debugf("refresh canceled for %s, %s", auth.Provider, auth.ID) - return + return nil, err } log.Debugf("refreshed %s, %s, %v", auth.Provider, auth.ID, err) now := time.Now() @@ -5743,7 +5891,7 @@ func (m *Manager) refreshAuth(ctx context.Context, id string) { if shouldReschedule { m.queueRefreshReschedule(id) } - return + return nil, err } if updated == nil { updated = cloned @@ -5756,11 +5904,27 @@ func (m *Manager) refreshAuth(ctx context.Context, id string) { updated.LastRefreshedAt = now updated.NextRefreshAfter = time.Time{} updated.LastError = nil + updated.StatusMessage = "" + updated.Unavailable = false + if updated.Status == StatusError { + updated.Status = StatusActive + } updated.UpdatedAt = now + modelsToResume := clearUnauthorizedModelStates(updated, now) if m.shouldRefresh(updated, now) { updated.NextRefreshAfter = now.Add(refreshIneffectiveBackoff) } - _, _ = m.Update(ctx, updated) + saved, errUpdate := m.Update(ctx, updated) + for _, model := range modelsToResume { + registry.GetGlobalRegistry().ResumeClientModel(id, model) + } + if errUpdate != nil { + log.Debugf("persist refreshed auth %s (%s) failed: %v", auth.Provider, auth.ID, errUpdate) + } + if saved != nil { + return saved, nil + } + return updated.Clone(), nil } func (m *Manager) executorFor(provider string) ProviderExecutor { diff --git a/sdk/cliproxy/auth/conductor_unauthorized_refresh_test.go b/sdk/cliproxy/auth/conductor_unauthorized_refresh_test.go new file mode 100644 index 000000000..374519250 --- /dev/null +++ b/sdk/cliproxy/auth/conductor_unauthorized_refresh_test.go @@ -0,0 +1,336 @@ +package auth + +import ( + "context" + "net/http" + "sync" + "testing" + + "github.com/router-for-me/CLIProxyAPI/v7/internal/registry" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" +) + +type unauthorizedRefreshExecutor struct { + id string + + mu sync.Mutex + executeCalls []string + streamCalls []string + refreshCalls int + tokenInvalid map[string]struct{} + refreshFail bool + refreshTokens map[string]string +} + +func (e *unauthorizedRefreshExecutor) Identifier() string { return e.id } + +func (e *unauthorizedRefreshExecutor) Execute(_ context.Context, auth *Auth, _ cliproxyexecutor.Request, _ cliproxyexecutor.Options) (cliproxyexecutor.Response, error) { + e.mu.Lock() + e.executeCalls = append(e.executeCalls, auth.ID) + token := authAccessToken(auth) + _, invalid := e.tokenInvalid[token] + e.mu.Unlock() + if invalid { + return cliproxyexecutor.Response{}, &Error{ + HTTPStatus: http.StatusUnauthorized, + Message: "Your authentication token has been invalidated. Please try signing in again.", + } + } + return cliproxyexecutor.Response{Payload: []byte(auth.ID + ":" + token)}, nil +} + +func (e *unauthorizedRefreshExecutor) ExecuteStream(_ context.Context, auth *Auth, _ cliproxyexecutor.Request, _ cliproxyexecutor.Options) (*cliproxyexecutor.StreamResult, error) { + e.mu.Lock() + e.streamCalls = append(e.streamCalls, auth.ID) + token := authAccessToken(auth) + _, invalid := e.tokenInvalid[token] + e.mu.Unlock() + if invalid { + return nil, &Error{ + HTTPStatus: http.StatusUnauthorized, + Message: "Your authentication token has been invalidated. Please try signing in again.", + } + } + ch := make(chan cliproxyexecutor.StreamChunk, 1) + ch <- cliproxyexecutor.StreamChunk{Payload: []byte(auth.ID + ":" + token)} + close(ch) + return &cliproxyexecutor.StreamResult{Headers: http.Header{"X-Auth": {auth.ID}}, Chunks: ch}, nil +} + +func (e *unauthorizedRefreshExecutor) Refresh(_ context.Context, auth *Auth) (*Auth, error) { + e.mu.Lock() + defer e.mu.Unlock() + e.refreshCalls++ + if e.refreshFail { + return nil, &Error{HTTPStatus: http.StatusUnauthorized, Message: "refresh token invalid"} + } + if auth.Metadata == nil { + auth.Metadata = make(map[string]any) + } + next := e.refreshTokens[auth.ID] + if next == "" { + next = "refreshed-access-token" + } + auth.Metadata["access_token"] = next + return auth, nil +} + +func (e *unauthorizedRefreshExecutor) CountTokens(context.Context, *Auth, cliproxyexecutor.Request, cliproxyexecutor.Options) (cliproxyexecutor.Response, error) { + return cliproxyexecutor.Response{}, &Error{HTTPStatus: http.StatusNotImplemented, Message: "not implemented"} +} + +func (e *unauthorizedRefreshExecutor) HttpRequest(context.Context, *Auth, *http.Request) (*http.Response, error) { + return nil, nil +} + +func (e *unauthorizedRefreshExecutor) ExecuteCalls() []string { + e.mu.Lock() + defer e.mu.Unlock() + out := make([]string, len(e.executeCalls)) + copy(out, e.executeCalls) + return out +} + +func (e *unauthorizedRefreshExecutor) StreamCalls() []string { + e.mu.Lock() + defer e.mu.Unlock() + out := make([]string, len(e.streamCalls)) + copy(out, e.streamCalls) + return out +} + +func (e *unauthorizedRefreshExecutor) RefreshCalls() int { + e.mu.Lock() + defer e.mu.Unlock() + return e.refreshCalls +} + +func newUnauthorizedRefreshFixture(t *testing.T, refreshFail bool) (*Manager, *unauthorizedRefreshExecutor, *Auth, *Auth, string) { + t.Helper() + + model := "gpt-5.5" + primary := &Auth{ + ID: "aa-primary", + Provider: "codex", + Metadata: map[string]any{ + "access_token": "stale-access-token", + "refresh_token": "primary-refresh-token", + }, + } + backup := &Auth{ + ID: "bb-backup", + Provider: "codex", + Metadata: map[string]any{ + "access_token": "backup-access-token", + "refresh_token": "backup-refresh-token", + }, + } + + executor := &unauthorizedRefreshExecutor{ + id: "codex", + tokenInvalid: map[string]struct{}{ + "stale-access-token": {}, + }, + refreshFail: refreshFail, + refreshTokens: map[string]string{ + primary.ID: "fresh-access-token", + }, + } + + m := NewManager(nil, nil, nil) + m.RegisterExecutor(executor) + + reg := registry.GetGlobalRegistry() + reg.RegisterClient(primary.ID, "codex", []*registry.ModelInfo{{ID: model}}) + reg.RegisterClient(backup.ID, "codex", []*registry.ModelInfo{{ID: model}}) + t.Cleanup(func() { + reg.UnregisterClient(primary.ID) + reg.UnregisterClient(backup.ID) + }) + + if _, errRegister := m.Register(context.Background(), primary); errRegister != nil { + t.Fatalf("register primary: %v", errRegister) + } + if _, errRegister := m.Register(context.Background(), backup); errRegister != nil { + t.Fatalf("register backup: %v", errRegister) + } + + return m, executor, primary, backup, model +} + +func TestManager_Execute_UnauthorizedRefreshesCurrentAuthBeforeFallback(t *testing.T) { + m, executor, primary, backup, model := newUnauthorizedRefreshFixture(t, false) + + resp, errExecute := m.Execute(context.Background(), []string{"codex"}, cliproxyexecutor.Request{Model: model}, cliproxyexecutor.Options{}) + if errExecute != nil { + t.Fatalf("Execute error = %v, want success on refreshed primary", errExecute) + } + if got := string(resp.Payload); got != primary.ID+":fresh-access-token" { + t.Fatalf("payload = %q, want refreshed primary response", got) + } + + if got := executor.RefreshCalls(); got != 1 { + t.Fatalf("Refresh calls = %d, want 1", got) + } + if got := executor.ExecuteCalls(); len(got) != 2 || got[0] != primary.ID || got[1] != primary.ID { + t.Fatalf("Execute calls = %v, want [primary, primary]", got) + } + for _, id := range executor.ExecuteCalls() { + if id == backup.ID { + t.Fatalf("backup auth should not be used when refresh recovers primary") + } + } + + updated, ok := m.GetByID(primary.ID) + if !ok || updated == nil { + t.Fatalf("primary auth missing after refresh") + } + if got := authAccessToken(updated); got != "fresh-access-token" { + t.Fatalf("primary access_token = %q, want fresh-access-token", got) + } + if state := updated.ModelStates[model]; state != nil && state.Unavailable { + t.Fatalf("primary model should not remain suspended after successful refresh retry") + } +} + +func TestManager_ExecuteStream_UnauthorizedRefreshesCurrentAuthBeforeFallback(t *testing.T) { + m, executor, primary, backup, model := newUnauthorizedRefreshFixture(t, false) + + stream, errStream := m.ExecuteStream(context.Background(), []string{"codex"}, cliproxyexecutor.Request{Model: model}, cliproxyexecutor.Options{}) + if errStream != nil { + t.Fatalf("ExecuteStream error = %v, want success on refreshed primary", errStream) + } + if stream == nil || stream.Chunks == nil { + t.Fatalf("expected stream result") + } + chunk, ok := <-stream.Chunks + if !ok { + t.Fatalf("expected stream chunk") + } + if chunk.Err != nil { + t.Fatalf("stream chunk error = %v", chunk.Err) + } + if got := string(chunk.Payload); got != primary.ID+":fresh-access-token" { + t.Fatalf("stream payload = %q, want refreshed primary response", got) + } + + if got := executor.RefreshCalls(); got != 1 { + t.Fatalf("Refresh calls = %d, want 1", got) + } + if got := executor.StreamCalls(); len(got) != 2 || got[0] != primary.ID || got[1] != primary.ID { + t.Fatalf("Stream calls = %v, want [primary, primary]", got) + } + for _, id := range executor.StreamCalls() { + if id == backup.ID { + t.Fatalf("backup auth should not be used when refresh recovers primary") + } + } +} + +func TestManager_Execute_UnauthorizedRefreshFailureFallsBackToNextAuth(t *testing.T) { + m, executor, primary, backup, model := newUnauthorizedRefreshFixture(t, true) + + resp, errExecute := m.Execute(context.Background(), []string{"codex"}, cliproxyexecutor.Request{Model: model}, cliproxyexecutor.Options{}) + if errExecute != nil { + t.Fatalf("Execute error = %v, want success via backup", errExecute) + } + if got := string(resp.Payload); got != backup.ID+":backup-access-token" { + t.Fatalf("payload = %q, want backup response", got) + } + + if got := executor.RefreshCalls(); got != 1 { + t.Fatalf("Refresh calls = %d, want 1", got) + } + if got := executor.ExecuteCalls(); len(got) != 2 || got[0] != primary.ID || got[1] != backup.ID { + t.Fatalf("Execute calls = %v, want [primary, backup]", got) + } + + updated, ok := m.GetByID(primary.ID) + if !ok || updated == nil { + t.Fatalf("primary auth missing after failed refresh") + } + state := updated.ModelStates[model] + if state == nil || !state.Unavailable { + t.Fatalf("expected primary model to be suspended after refresh failure") + } + if state.StatusMessage != "unauthorized" && (state.LastError == nil || state.LastError.StatusCode() != http.StatusUnauthorized) { + t.Fatalf("expected unauthorized suspension, got state=%+v", state) + } +} + +func TestManager_Execute_UnauthorizedWithoutRefreshTokenDoesNotCallRefresh(t *testing.T) { + model := "gpt-5.5" + primary := &Auth{ + ID: "aa-primary-api-key", + Provider: "codex", + Metadata: map[string]any{ + "access_token": "stale-access-token", + }, + } + backup := &Auth{ + ID: "bb-backup-api-key", + Provider: "codex", + Metadata: map[string]any{ + "access_token": "backup-access-token", + }, + } + executor := &unauthorizedRefreshExecutor{ + id: "codex", + tokenInvalid: map[string]struct{}{ + "stale-access-token": {}, + }, + } + m := NewManager(nil, nil, nil) + m.RegisterExecutor(executor) + + reg := registry.GetGlobalRegistry() + reg.RegisterClient(primary.ID, "codex", []*registry.ModelInfo{{ID: model}}) + reg.RegisterClient(backup.ID, "codex", []*registry.ModelInfo{{ID: model}}) + t.Cleanup(func() { + reg.UnregisterClient(primary.ID) + reg.UnregisterClient(backup.ID) + }) + if _, errRegister := m.Register(context.Background(), primary); errRegister != nil { + t.Fatalf("register primary: %v", errRegister) + } + if _, errRegister := m.Register(context.Background(), backup); errRegister != nil { + t.Fatalf("register backup: %v", errRegister) + } + + resp, errExecute := m.Execute(context.Background(), []string{"codex"}, cliproxyexecutor.Request{Model: model}, cliproxyexecutor.Options{}) + if errExecute != nil { + t.Fatalf("Execute error = %v, want success via backup", errExecute) + } + if got := string(resp.Payload); got != backup.ID+":backup-access-token" { + t.Fatalf("payload = %q, want backup response", got) + } + if got := executor.RefreshCalls(); got != 0 { + t.Fatalf("Refresh calls = %d, want 0 when no refresh_token is present", got) + } + if got := executor.ExecuteCalls(); len(got) != 2 || got[0] != primary.ID || got[1] != backup.ID { + t.Fatalf("Execute calls = %v, want [primary, backup]", got) + } +} + +func TestManager_Execute_UnauthorizedRefreshThenRetryStillFailsFallsBackOnce(t *testing.T) { + m, executor, primary, backup, model := newUnauthorizedRefreshFixture(t, false) + // Refresh "succeeds" but hands back another invalidated token. + executor.refreshTokens[primary.ID] = "still-invalid-token" + executor.mu.Lock() + executor.tokenInvalid["still-invalid-token"] = struct{}{} + executor.mu.Unlock() + + resp, errExecute := m.Execute(context.Background(), []string{"codex"}, cliproxyexecutor.Request{Model: model}, cliproxyexecutor.Options{}) + if errExecute != nil { + t.Fatalf("Execute error = %v, want success via backup", errExecute) + } + if got := string(resp.Payload); got != backup.ID+":backup-access-token" { + t.Fatalf("payload = %q, want backup response", got) + } + if got := executor.RefreshCalls(); got != 1 { + t.Fatalf("Refresh calls = %d, want 1 (no refresh loop)", got) + } + if got := executor.ExecuteCalls(); len(got) != 3 || got[0] != primary.ID || got[1] != primary.ID || got[2] != backup.ID { + t.Fatalf("Execute calls = %v, want [primary, primary, backup]", got) + } +}