feat(auth): enable automatic credential refresh on unauthorized errors

- Added `tryRefreshAfterUnauthorized` to refresh OAuth credentials on 401 errors during requests.
- Implemented `refreshLocks` to prevent concurrent refreshes for the same auth ID.
- Updated auth/state handling to reset unauthorized model states and resume operations after a successful refresh.
- Enhanced refresh logic with error handling, synchronization, and state updates.

Closes: #4087
This commit is contained in:
Luis Pater
2026-07-09 03:08:02 +08:00
parent 186c87ba6d
commit ec3aba23fa
2 changed files with 508 additions and 8 deletions

View File

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

View File

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