mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-28 11:50:20 +08:00
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:
@@ -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 {
|
||||
|
||||
336
sdk/cliproxy/auth/conductor_unauthorized_refresh_test.go
Normal file
336
sdk/cliproxy/auth/conductor_unauthorized_refresh_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user