From 1e38a3a544ec8bfc4206f3a975749a344b4d96a6 Mon Sep 17 00:00:00 2001 From: Supra4E8C Date: Fri, 31 Jul 2026 20:17:30 +0800 Subject: [PATCH] fix: retry Home OAuth requests after unauthorized --- internal/api/server_routes.go | 131 ++++--- internal/api/server_test.go | 45 ++- internal/client/codex/live/live.go | 74 ++-- internal/client/codex/live/live_test.go | 131 ++++++- internal/client/codex/live/sideband.go | 70 ++-- internal/home/client.go | 6 +- internal/home/requests.go | 5 +- internal/redisqueue/plugin.go | 2 + internal/redisqueue/plugin_test.go | 2 + .../runtime/executor/helps/home_refresh.go | 67 +++- .../executor/helps/home_refresh_test.go | 90 ++++- .../runtime/executor/helps/usage_helpers.go | 39 +- sdk/cliproxy/auth/conductor_execution.go | 12 +- sdk/cliproxy/auth/conductor_home.go | 71 +++- sdk/cliproxy/auth/conductor_home_execution.go | 24 +- sdk/cliproxy/auth/conductor_refresh.go | 73 +++- sdk/cliproxy/auth/conductor_selection.go | 7 +- sdk/cliproxy/auth/conductor_stream.go | 45 ++- sdk/cliproxy/auth/home_concurrency.go | 3 +- sdk/cliproxy/auth/home_concurrency_test.go | 12 + sdk/cliproxy/auth/home_selection.go | 34 +- sdk/cliproxy/auth/home_selection_test.go | 64 ++++ .../auth/home_unauthorized_refresh_test.go | 340 ++++++++++++++++++ sdk/cliproxy/usage/manager.go | 6 +- 24 files changed, 1183 insertions(+), 170 deletions(-) create mode 100644 sdk/cliproxy/auth/home_unauthorized_refresh_test.go diff --git a/internal/api/server_routes.go b/internal/api/server_routes.go index bb35e37e..37cc170e 100644 --- a/internal/api/server_routes.go +++ b/internal/api/server_routes.go @@ -333,64 +333,52 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { } logging.SetGinCPATraceID(c, selected.EnsureIndex()) - headers := make(http.Header) - headers.Set("Content-Type", "application/json") - headers.Set("Accept", "application/json") - headers.Set("Originator", "codex_cli_rs") + baseHeaders := make(http.Header) + baseHeaders.Set("Content-Type", "application/json") + baseHeaders.Set("Accept", "application/json") + baseHeaders.Set("Originator", "codex_cli_rs") for _, name := range []string{"Version", "User-Agent", "Session_id", "X-Client-Request-Id"} { if value := strings.TrimSpace(c.GetHeader(name)); value != "" { - headers.Set(name, value) + baseHeaders.Set(name, value) } } - if accountID, ok := selected.Metadata["account_id"].(string); ok && strings.TrimSpace(accountID) != "" { - headers.Set("Chatgpt-Account-Id", accountID) - } - upstreamURL := "https://chatgpt.com/backend-api/codex/alpha/search" - if selected.AuthKind() == auth.AuthKindAPIKey { - baseURL := "" - if selected.Attributes != nil { - baseURL = strings.TrimSpace(selected.Attributes["base_url"]) + errMissingBaseURL := errors.New("Codex Alpha Search API key base URL unavailable") + performRequest := func(current *auth.Auth) (*http.Response, error) { + headers := baseHeaders.Clone() + if accountID, ok := current.Metadata["account_id"].(string); ok && strings.TrimSpace(accountID) != "" { + headers.Set("Chatgpt-Account-Id", accountID) } - if baseURL == "" { - if selection != nil { - selection.End("missing_base_url") + upstreamURL := "https://chatgpt.com/backend-api/codex/alpha/search" + if current.AuthKind() == auth.AuthKindAPIKey { + baseURL := "" + if current.Attributes != nil { + baseURL = strings.TrimSpace(current.Attributes["base_url"]) } - c.JSON(http.StatusServiceUnavailable, gin.H{"error": "Codex Alpha Search API key base URL unavailable"}) - return - } - upstreamURL = strings.TrimRight(baseURL, "/") + "/alpha/search" - } - req, err := s.handlers.AuthManager.NewHttpRequest( - ctx, selected, http.MethodPost, upstreamURL, upstreamRequestBody, headers, - ) - if err != nil { - if selection != nil { - selection.End("request_build_failed") - } - c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) - return + if baseURL == "" { + return nil, errMissingBaseURL + } + upstreamURL = strings.TrimRight(baseURL, "/") + "/alpha/search" + } + req, errRequest := s.handlers.AuthManager.NewHttpRequest(ctx, current, http.MethodPost, upstreamURL, upstreamRequestBody, headers) + if errRequest != nil { + return nil, errRequest + } + authType, authValue := current.AccountInfo() + helps.RecordAPIRequest(ctx, s.cfg, helps.UpstreamRequestLog{ + URL: upstreamURL, + Method: http.MethodPost, + Headers: req.Header.Clone(), + Body: upstreamRequestBody, + Provider: "codex", + AuthID: current.ID, + AuthLabel: current.Label, + AuthType: authType, + AuthValue: authValue, + }) + return s.handlers.AuthManager.HttpRequest(ctx, current, req) } - var authID, authLabel, authType, authValue string - if selected != nil { - authID = selected.ID - authLabel = selected.Label - authType, authValue = selected.AccountInfo() - } - helpHeaders := req.Header.Clone() - helps.RecordAPIRequest(ctx, s.cfg, helps.UpstreamRequestLog{ - URL: upstreamURL, - Method: http.MethodPost, - Headers: helpHeaders, - Body: upstreamRequestBody, - Provider: "codex", - AuthID: authID, - AuthLabel: authLabel, - AuthType: authType, - AuthValue: authValue, - }) - if errCtx := ctx.Err(); errCtx != nil { if selection != nil { selection.End("attempt_canceled") @@ -398,8 +386,15 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { c.JSON(http.StatusRequestTimeout, gin.H{"error": errCtx.Error()}) return } - resp, err := s.handlers.AuthManager.HttpRequest(ctx, selected, req) + resp, err := performRequest(selected) if err != nil { + if errors.Is(err, errMissingBaseURL) { + if selection != nil { + selection.End("missing_base_url") + } + c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) + return + } if selection != nil { selection.End("request_failed") } @@ -407,6 +402,42 @@ func (s *Server) codexAlphaSearch(c *gin.Context) { c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) return } + if selection != nil && resp.StatusCode == http.StatusUnauthorized { + helps.RecordAPIResponseMetadata(ctx, s.cfg, resp.StatusCode, resp.Header.Clone()) + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + if errClose := resp.Body.Close(); errClose != nil { + log.Errorf("codex alpha search: close unauthorized response body error: %v", errClose) + } + refreshed, didRefresh, errRefresh := s.handlers.AuthManager.RefreshHomeSelectionAfterUnauthorized(ctx, selection, selected) + if errRefresh != nil { + selection.End("refresh_failed") + status := http.StatusServiceUnavailable + if statusError, ok := errRefresh.(interface{ StatusCode() int }); ok && statusError.StatusCode() > 0 { + status = statusError.StatusCode() + } + c.JSON(status, gin.H{"error": errRefresh.Error()}) + return + } + if !didRefresh || refreshed == nil { + selection.End("refresh_unavailable") + c.JSON(http.StatusUnauthorized, gin.H{"error": "Codex credential unauthorized"}) + return + } + selected = refreshed + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + resp, err = performRequest(selected) + if err != nil { + if errors.Is(err, errMissingBaseURL) { + selection.End("missing_base_url") + c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) + return + } + selection.End("retry_failed") + helps.RecordAPIResponseError(ctx, s.cfg, err) + c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) + return + } + } closeResponseBody := func() error { errClose := resp.Body.Close() if errClose != nil { diff --git a/internal/api/server_test.go b/internal/api/server_test.go index 511a94f7..9ff764aa 100644 --- a/internal/api/server_test.go +++ b/internal/api/server_test.go @@ -38,6 +38,9 @@ type codexSearchCaptureExecutor struct { prepareErr error httpErr error responseBody io.ReadCloser + statuses []int + refreshCalls int + httpCalls int } func (e *codexSearchCaptureExecutor) Identifier() string { return "codex" } @@ -51,7 +54,13 @@ func (e *codexSearchCaptureExecutor) ExecuteStream(context.Context, *auth.Auth, } func (e *codexSearchCaptureExecutor) Refresh(_ context.Context, a *auth.Auth) (*auth.Auth, error) { - return a, nil + e.refreshCalls++ + updated := a.Clone() + if updated.Metadata == nil { + updated.Metadata = make(map[string]any) + } + updated.Metadata["access_token"] = "refreshed-home-search-token" + return updated, nil } func (e *codexSearchCaptureExecutor) CountTokens(context.Context, *auth.Auth, coreexecutor.Request, coreexecutor.Options) (coreexecutor.Response, error) { @@ -115,6 +124,7 @@ func (e *codexSearchCaptureExecutor) HttpRequest(_ context.Context, selected *au } e.request = req.Clone(req.Context()) e.authIDs = append(e.authIDs, selected.ID) + e.httpCalls++ body, err := io.ReadAll(req.Body) if err != nil { return nil, err @@ -124,8 +134,12 @@ func (e *codexSearchCaptureExecutor) HttpRequest(_ context.Context, selected *au if responseBody == nil { responseBody = io.NopCloser(strings.NewReader(`{"results":[{"url":"https://example.com"}]}`)) } + statusCode := http.StatusOK + if e.httpCalls <= len(e.statuses) && e.statuses[e.httpCalls-1] > 0 { + statusCode = e.statuses[e.httpCalls-1] + } return &http.Response{ - StatusCode: http.StatusOK, + StatusCode: statusCode, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: responseBody, }, nil @@ -310,6 +324,33 @@ func TestAuditHomeCodexSearchBodyCloseBeforeRelease(t *testing.T) { } } +func TestHomeCodexAlphaSearchRefreshesUnauthorizedSelectionOnce(t *testing.T) { + server := newTestServer(t) + dispatcher := &codexSearchHomeDispatcher{} + server.handlers.AuthManager.SetConfig(&proxyconfig.Config{Home: proxyconfig.HomeConfig{Enabled: true}}) + server.handlers.AuthManager.PublishHomeDispatch(dispatcher, executionregistry.New(), 1) + executor := &codexSearchCaptureExecutor{statuses: []int{http.StatusUnauthorized, http.StatusOK}} + server.handlers.AuthManager.RegisterExecutor(executor) + + req := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"id":"home-search-refresh","model":"gpt-5-codex","query":"test"}`)) + req.Header.Set("Authorization", "Bearer test-key") + rr := httptest.NewRecorder() + server.engine.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d; body=%s", rr.Code, http.StatusOK, rr.Body.String()) + } + if executor.refreshCalls != 1 || executor.httpCalls != 2 { + t.Fatalf("refresh/http calls = %d/%d, want 1/2", executor.refreshCalls, executor.httpCalls) + } + if got := executor.request.Header.Get("Authorization"); got != "Bearer refreshed-home-search-token" { + t.Fatalf("retry Authorization = %q, want refreshed token", got) + } + if got := dispatcher.calls.Load(); got != 1 { + t.Fatalf("Home RPOP calls = %d, want 1", got) + } +} + func TestHomeCodexAlphaSearchEndsSelectionAcrossDirectHTTPPaths(t *testing.T) { tests := []struct { name string diff --git a/internal/client/codex/live/live.go b/internal/client/codex/live/live.go index 3a699ca2..f66b9e98 100644 --- a/internal/client/codex/live/live.go +++ b/internal/client/codex/live/live.go @@ -263,31 +263,30 @@ func (h *Handler) Handle(c *gin.Context) { } } - headers := protocolHeaders(c.Request.Header) - headers.Set("Content-Type", upstreamContentType) - setAccountHeader(headers, selected) - req, errRequest := h.authManager.NewHttpRequest(ctx, selected, http.MethodPost, upstreamCallURL, upstreamBody, headers) - if errRequest != nil { - if selection != nil { - selection.End("request_build_failed") - } - c.JSON(http.StatusBadGateway, gin.H{"error": errRequest.Error()}) - return + baseHeaders := protocolHeaders(c.Request.Header) + baseHeaders.Set("Content-Type", upstreamContentType) + performRequest := func(current *auth.Auth) (*http.Response, error) { + headers := baseHeaders.Clone() + setAccountHeader(headers, current) + req, errRequest := h.authManager.NewHttpRequest(ctx, current, http.MethodPost, upstreamCallURL, upstreamBody, headers) + if errRequest != nil { + return nil, errRequest + } + authType, authValue := current.AccountInfo() + helps.RecordAPIRequest(ctx, runtimeConfig, helps.UpstreamRequestLog{ + URL: upstreamCallURL, + Method: http.MethodPost, + Headers: headersForLogging(req.Header), + Body: upstreamBody, + Provider: "codex", + AuthID: current.ID, + AuthLabel: current.Label, + AuthType: authType, + AuthValue: authValue, + }) + return h.authManager.HttpRequest(ctx, current, req) } - authType, authValue := selected.AccountInfo() - helps.RecordAPIRequest(ctx, runtimeConfig, helps.UpstreamRequestLog{ - URL: upstreamCallURL, - Method: http.MethodPost, - Headers: headersForLogging(req.Header), - Body: upstreamBody, - Provider: "codex", - AuthID: selected.ID, - AuthLabel: selected.Label, - AuthType: authType, - AuthValue: authValue, - }) - if errContext := ctx.Err(); errContext != nil { if selection != nil { selection.End("attempt_canceled") @@ -295,7 +294,7 @@ func (h *Handler) Handle(c *gin.Context) { c.JSON(http.StatusRequestTimeout, gin.H{"error": errContext.Error()}) return } - resp, errRequest := h.authManager.HttpRequest(ctx, selected, req) + resp, errRequest := performRequest(selected) if errRequest != nil { if selection != nil { selection.End("request_failed") @@ -304,6 +303,33 @@ func (h *Handler) Handle(c *gin.Context) { c.JSON(http.StatusBadGateway, gin.H{"error": errRequest.Error()}) return } + if selection != nil && resp.StatusCode == http.StatusUnauthorized { + helps.RecordAPIResponseMetadata(ctx, runtimeConfig, resp.StatusCode, callResponseHeaders(resp.Header)) + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + if errClose := resp.Body.Close(); errClose != nil { + log.Errorf("codex live: close unauthorized response body error: %v", errClose) + } + refreshed, didRefresh, errRefresh := h.authManager.RefreshHomeSelectionAfterUnauthorized(ctx, selection, selected) + if errRefresh != nil { + selection.End("refresh_failed") + writeSelectionError(c, errRefresh) + return + } + if !didRefresh || refreshed == nil { + selection.End("refresh_unavailable") + c.JSON(http.StatusUnauthorized, gin.H{"error": "Codex credential unauthorized"}) + return + } + selected = refreshed + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + resp, errRequest = performRequest(selected) + if errRequest != nil { + selection.End("retry_failed") + helps.RecordAPIResponseError(ctx, runtimeConfig, errRequest) + c.JSON(http.StatusBadGateway, gin.H{"error": errRequest.Error()}) + return + } + } var closeResponseOnce sync.Once var closeResponseErr error diff --git a/internal/client/codex/live/live_test.go b/internal/client/codex/live/live_test.go index a4f54b2e..3dcbff76 100644 --- a/internal/client/codex/live/live_test.go +++ b/internal/client/codex/live/live_test.go @@ -40,6 +40,9 @@ type captureExecutor struct { selectedAuth *auth.Auth responseBody io.ReadCloser statusCode int + statuses []int + httpCalls atomic.Int32 + refreshCalls atomic.Int32 } func (*captureExecutor) Identifier() string { return "codex" } @@ -52,8 +55,14 @@ func (*captureExecutor) ExecuteStream(context.Context, *auth.Auth, coreexecutor. return nil, nil } -func (*captureExecutor) Refresh(_ context.Context, credential *auth.Auth) (*auth.Auth, error) { - return credential, nil +func (e *captureExecutor) Refresh(_ context.Context, credential *auth.Auth) (*auth.Auth, error) { + e.refreshCalls.Add(1) + updated := credential.Clone() + if updated.Metadata == nil { + updated.Metadata = make(map[string]any) + } + updated.Metadata["access_token"] = "refreshed-home-live-token" + return updated, nil } func (*captureExecutor) CountTokens(context.Context, *auth.Auth, coreexecutor.Request, coreexecutor.Options) (coreexecutor.Response, error) { @@ -69,15 +78,23 @@ func (*captureExecutor) PrepareRequest(req *http.Request, credential *auth.Auth) func (e *captureExecutor) HttpRequest(_ context.Context, credential *auth.Auth, req *http.Request) (*http.Response, error) { e.request = req.Clone(req.Context()) e.selectedAuth = credential.Clone() + httpCall := int(e.httpCalls.Add(1)) body, errRead := io.ReadAll(req.Body) if errRead != nil { return nil, errRead } e.body = body statusCode := e.statusCode + if httpCall <= len(e.statuses) && e.statuses[httpCall-1] > 0 { + statusCode = e.statuses[httpCall-1] + } if statusCode == 0 { statusCode = http.StatusCreated } + responseBody := e.responseBody + if statusCode == http.StatusUnauthorized && httpCall < len(e.statuses) { + responseBody = io.NopCloser(strings.NewReader("unauthorized")) + } return &http.Response{ StatusCode: statusCode, Header: http.Header{ @@ -88,7 +105,7 @@ func (e *captureExecutor) HttpRequest(_ context.Context, credential *auth.Auth, "X-Connection-Secret": []string{"secret"}, "X-Live-Session": []string{"live-session-123"}, }, - Body: e.responseBody, + Body: responseBody, }, nil } @@ -589,6 +606,40 @@ func TestHandlerClosesMediaWhenResponseWriteFails(t *testing.T) { } } +func TestHandlerRefreshesUnauthorizedHomeSelectionOnce(t *testing.T) { + gin.SetMode(gin.TestMode) + manager := auth.NewManager(nil, nil, nil) + manager.SetConfig(&config.Config{Home: config.HomeConfig{Enabled: true}}) + registry := executionregistry.New() + manager.PublishHomeDispatch(&homeDispatcher{}, registry, 1) + executor := &captureExecutor{ + statuses: []int{http.StatusUnauthorized, http.StatusCreated}, + responseBody: &trackedResponseBody{Reader: strings.NewReader("v=0\r\n")}, + } + manager.RegisterExecutor(executor) + handler := NewHandler(manager, nil) + router := gin.New() + router.POST("/v1/live", handler.Handle) + + req := httptest.NewRequest(http.MethodPost, "/v1/live", strings.NewReader(`{"model":"gpt-live-1-codex","sdp":"v=0"}`)) + req.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + + if recorder.Code != http.StatusCreated { + t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusCreated, recorder.Body.String()) + } + if executor.refreshCalls.Load() != 1 || executor.httpCalls.Load() != 2 { + t.Fatalf("refresh/http calls = %d/%d, want 1/2", executor.refreshCalls.Load(), executor.httpCalls.Load()) + } + if got := executor.request.Header.Get("Authorization"); got != "Bearer refreshed-home-live-token" { + t.Fatalf("retry Authorization = %q, want refreshed token", got) + } + if errDrain := registry.Drain(context.Background()); errDrain != nil { + t.Fatalf("Drain() error = %v", errDrain) + } +} + func TestHandlerUsesLiveModelForHomeDispatch(t *testing.T) { gin.SetMode(gin.TestMode) @@ -767,6 +818,80 @@ func TestHandleSidebandPinsAuthAndRelaysBidirectionally(t *testing.T) { } } +func TestHandleSidebandRefreshesUnauthorizedHomeHandshakeOnce(t *testing.T) { + gin.SetMode(gin.TestMode) + var upstreamCalls atomic.Int32 + upstreamHeaders := make(chan http.Header, 2) + upstreamServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + upstreamCalls.Add(1) + upstreamHeaders <- request.Header.Clone() + if request.Header.Get("Authorization") != "Bearer refreshed-home-live-token" { + writer.WriteHeader(http.StatusUnauthorized) + return + } + upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} + conn, errUpgrade := upgrader.Upgrade(writer, request, nil) + if errUpgrade != nil { + return + } + defer func() { _ = conn.Close() }() + messageType, payload, errRead := conn.ReadMessage() + if errRead == nil { + _ = conn.WriteMessage(messageType, append([]byte("echo:"), payload...)) + } + })) + defer upstreamServer.Close() + + manager := auth.NewManager(nil, nil, nil) + manager.SetConfig(&config.Config{Home: config.HomeConfig{Enabled: true}}) + registry := executionregistry.New() + manager.PublishHomeDispatch(&homeDispatcher{}, registry, 1) + executor := &captureExecutor{} + manager.RegisterExecutor(executor) + selection, errSelect := manager.SelectHomeAuthByKind(context.Background(), "codex", defaultLiveModel, auth.AuthKindOAuth, coreexecutor.Options{}) + if errSelect != nil { + t.Fatalf("SelectHomeAuthByKind() error = %v", errSelect) + } + selection.Retain() + defer selection.End("test_complete") + + handler := NewHandler(manager, nil) + handler.sidebandAPIBaseURL = "ws" + strings.TrimPrefix(upstreamServer.URL, "http") + "/v1" + handler.sessions.put("call-home-refresh", liveSession{authID: "home-codex-live", model: defaultLiveModel, homeSelection: selection}) + router := gin.New() + router.GET("/v1/live/:call_id", handler.HandleSideband) + downstreamServer := httptest.NewServer(router) + defer downstreamServer.Close() + + wsURL := "ws" + strings.TrimPrefix(downstreamServer.URL, "http") + "/v1/live/call-home-refresh" + client, response, errDial := websocket.DefaultDialer.Dial(wsURL, nil) + if errDial != nil { + if response != nil && response.Body != nil { + _ = response.Body.Close() + } + t.Fatalf("dial downstream sideband: %v", errDial) + } + if response != nil && response.Body != nil { + _ = response.Body.Close() + } + defer func() { _ = client.Close() }() + if errWrite := client.WriteMessage(websocket.TextMessage, []byte("ping")); errWrite != nil { + t.Fatalf("write sideband message: %v", errWrite) + } + _, payload, errRead := client.ReadMessage() + if errRead != nil || string(payload) != "echo:ping" { + t.Fatalf("read sideband message = %q, %v", string(payload), errRead) + } + if executor.refreshCalls.Load() != 1 || upstreamCalls.Load() != 2 { + t.Fatalf("refresh/upstream calls = %d/%d, want 1/2", executor.refreshCalls.Load(), upstreamCalls.Load()) + } + first := <-upstreamHeaders + second := <-upstreamHeaders + if first.Get("Authorization") != "Bearer home-live-token" || second.Get("Authorization") != "Bearer refreshed-home-live-token" { + t.Fatalf("upstream Authorization sequence = %q, %q", first.Get("Authorization"), second.Get("Authorization")) + } +} + func TestPrepareCallRequestRewritesMultipart(t *testing.T) { const boundary = "live-model-boundary" body := multipartBody(boundary, "v=0-offer", `{"model":"future-live-model","instructions":"hi"}`) diff --git a/internal/client/codex/live/sideband.go b/internal/client/codex/live/sideband.go index 1e2b679a..7ff5840f 100644 --- a/internal/client/codex/live/sideband.go +++ b/internal/client/codex/live/sideband.go @@ -380,33 +380,53 @@ func (h *Handler) HandleSideband(c *gin.Context) { upstreamURL := buildSidebandURL(h.sidebandAPIBaseURL, style, callID) upstreamHTTPURL := websocketHTTPURL(upstreamURL) - req, errRequest := http.NewRequestWithContext(ctx, http.MethodGet, upstreamHTTPURL, nil) - if errRequest != nil { - c.JSON(http.StatusBadGateway, gin.H{"error": errRequest.Error()}) - return - } - req.Header = protocolHeaders(c.Request.Header) - setAccountHeader(req.Header, selected) - if errPrepare := h.authManager.PrepareHttpRequest(ctx, selected, req); errPrepare != nil { - c.JSON(http.StatusBadGateway, gin.H{"error": errPrepare.Error()}) - return + dialUpstream := func(current *auth.Auth) (*websocket.Conn, *http.Response, error) { + req, errRequest := http.NewRequestWithContext(ctx, http.MethodGet, upstreamHTTPURL, nil) + if errRequest != nil { + return nil, nil, errRequest + } + req.Header = protocolHeaders(c.Request.Header) + setAccountHeader(req.Header, current) + if errPrepare := h.authManager.PrepareHttpRequest(ctx, current, req); errPrepare != nil { + return nil, nil, errPrepare + } + authType, authValue := current.AccountInfo() + helps.RecordAPIWebsocketRequest(ctx, runtimeConfig, helps.UpstreamRequestLog{ + URL: upstreamURL, + Method: "WEBSOCKET", + Headers: headersForLogging(req.Header), + Provider: "codex", + AuthID: current.ID, + AuthLabel: current.Label, + AuthType: authType, + AuthValue: authValue, + }) + dialer := newProxyAwareSidebandDialer(runtimeConfig, current) + dialer.Subprotocols = websocket.Subprotocols(c.Request) + return dialer.DialContext(ctx, upstreamURL, req.Header) } - authType, authValue := selected.AccountInfo() - helps.RecordAPIWebsocketRequest(ctx, runtimeConfig, helps.UpstreamRequestLog{ - URL: upstreamURL, - Method: "WEBSOCKET", - Headers: headersForLogging(req.Header), - Provider: "codex", - AuthID: selected.ID, - AuthLabel: selected.Label, - AuthType: authType, - AuthValue: authValue, - }) - - dialer := newProxyAwareSidebandDialer(runtimeConfig, selected) - dialer.Subprotocols = websocket.Subprotocols(c.Request) - upstream, handshakeResponse, errDial := dialer.DialContext(ctx, upstreamURL, req.Header) + upstream, handshakeResponse, errDial := dialUpstream(selected) + if errDial != nil && selection != nil && handshakeResponse != nil && handshakeResponse.StatusCode == http.StatusUnauthorized { + helps.RecordAPIWebsocketHandshake(ctx, runtimeConfig, handshakeResponse.StatusCode, callResponseHeaders(handshakeResponse.Header)) + if handshakeResponse.Body != nil { + if errClose := handshakeResponse.Body.Close(); errClose != nil { + log.Errorf("codex live sideband: close unauthorized handshake body error: %v", errClose) + } + } + refreshed, didRefresh, errRefresh := h.authManager.RefreshHomeSelectionAfterUnauthorized(ctx, selection, selected) + if errRefresh != nil { + writeSelectionError(c, errRefresh) + return + } + if !didRefresh || refreshed == nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "Codex credential unauthorized"}) + return + } + selected = refreshed + logging.SetGinCPATraceID(c, selected.EnsureIndex()) + upstream, handshakeResponse, errDial = dialUpstream(selected) + } if errDial != nil { handleSidebandDialError(c, ctx, runtimeConfig, handshakeResponse, errDial) return diff --git a/internal/home/client.go b/internal/home/client.go index d07d4d42..62964f18 100644 --- a/internal/home/client.go +++ b/internal/home/client.go @@ -41,6 +41,7 @@ const ( homeReconnectInterval = time.Second homeReconnectFailoverThreshold = 3 homeRedisOperationTimeout = 3 * time.Second + homeRefreshOperationTimeout = 35 * time.Second homePluginSyncOperationTimeout = 2 * time.Minute homeSubscriptionReceiveTimeout = 3 * time.Second credentialConcurrencyNodeHeartbeatTimeout = 20 * time.Second @@ -1337,7 +1338,7 @@ func isAmbiguousIssuedRPopAuthError(err error) bool { return !errors.As(err, &redisErr) } -func (c *Client) GetRefreshAuth(ctx context.Context, authIndex string) ([]byte, error) { +func (c *Client) GetRefreshAuth(ctx context.Context, authIndex string, accessTokenSHA256 string) ([]byte, error) { cmd, errClient := c.commandClient() if errClient != nil { return nil, errClient @@ -1350,12 +1351,13 @@ func (c *Client) GetRefreshAuth(ctx context.Context, authIndex string) ([]byte, Type: "refresh", AuthIndex: authIndex, } + req.ObservedAccessTokenSHA256 = strings.TrimSpace(accessTokenSHA256) keyBytes, err := json.Marshal(&req) if err != nil { return nil, err } - raw, err := cmd.Get(ctx, string(keyBytes)).Bytes() + raw, err := cmd.WithTimeout(homeRefreshOperationTimeout).Get(ctx, string(keyBytes)).Bytes() if errors.Is(err, redis.Nil) { return nil, ErrAuthNotFound } diff --git a/internal/home/requests.go b/internal/home/requests.go index 655fc601..e013eb43 100644 --- a/internal/home/requests.go +++ b/internal/home/requests.go @@ -19,8 +19,9 @@ type modelsRequest struct { } type refreshRequest struct { - Type string `json:"type"` - AuthIndex string `json:"auth_index"` + Type string `json:"type"` + AuthIndex string `json:"auth_index"` + ObservedAccessTokenSHA256 string `json:"access_token_sha256,omitempty"` } type InFlightFrameKind string diff --git a/internal/redisqueue/plugin.go b/internal/redisqueue/plugin.go index 915f8894..d91c8a28 100644 --- a/internal/redisqueue/plugin.go +++ b/internal/redisqueue/plugin.go @@ -90,6 +90,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec TTFTMs: record.TTFT.Milliseconds(), Source: record.Source, AuthIndex: record.AuthIndex, + AccessTokenHash: record.AccessTokenSHA256, ClientIP: clientRequestMetadata.ClientIP, XForwardedFor: clientRequestMetadata.XForwardedFor, UserAgent: clientRequestMetadata.UserAgent, @@ -145,6 +146,7 @@ type requestDetail struct { TTFTMs int64 `json:"ttft_ms"` Source string `json:"source"` AuthIndex string `json:"auth_index"` + AccessTokenHash string `json:"access_token_sha256,omitempty"` ClientIP string `json:"client_ip"` XForwardedFor string `json:"x_forwarded_for"` UserAgent string `json:"user_agent"` diff --git a/internal/redisqueue/plugin_test.go b/internal/redisqueue/plugin_test.go index 34234eb4..c1a1f010 100644 --- a/internal/redisqueue/plugin_test.go +++ b/internal/redisqueue/plugin_test.go @@ -36,6 +36,7 @@ func TestUsageQueuePluginPayloadIncludesStableFieldsAndSuccess(t *testing.T) { Alias: "client-gpt", APIKey: "test-key", AuthIndex: "0", + AccessTokenSHA256: "token-version-hash", AuthType: "apikey", Source: "user@example.com", ReasoningEffort: "medium", @@ -60,6 +61,7 @@ func TestUsageQueuePluginPayloadIncludesStableFieldsAndSuccess(t *testing.T) { requireStringField(t, payload, "alias", "client-gpt") requireStringField(t, payload, "endpoint", "POST /v1/chat/completions") requireStringField(t, payload, "auth_type", "apikey") + requireStringField(t, payload, "access_token_sha256", "token-version-hash") requireMissingField(t, payload, "user_api_key") requireStringField(t, payload, "request_id", "ctx-request-id") requireStringField(t, payload, "client_ip", "192.0.2.10") diff --git a/internal/runtime/executor/helps/home_refresh.go b/internal/runtime/executor/helps/home_refresh.go index 7c971992..af444396 100644 --- a/internal/runtime/executor/helps/home_refresh.go +++ b/internal/runtime/executor/helps/home_refresh.go @@ -2,7 +2,10 @@ package helps import ( "context" + "crypto/sha256" + "encoding/hex" "encoding/json" + "errors" "fmt" "net/http" "strings" @@ -43,7 +46,7 @@ type homeErrorDetail struct { type homeRefreshClient interface { HeartbeatOK() bool - GetRefreshAuth(ctx context.Context, authIndex string) ([]byte, error) + GetRefreshAuth(ctx context.Context, authIndex string, accessTokenSHA256 string) ([]byte, error) } var currentHomeRefreshClient = func() homeRefreshClient { @@ -77,9 +80,12 @@ func RefreshAuthViaHome(ctx context.Context, cfg *config.Config, auth *cliproxya return nil, true, homeStatusErr{code: http.StatusBadGateway, msg: "home refresh: auth_index is empty"} } - raw, err := client.GetRefreshAuth(ctx, authIndex) + raw, err := client.GetRefreshAuth(ctx, authIndex, authAccessTokenSHA256(auth)) if err != nil { - return nil, true, homeStatusErr{code: http.StatusBadGateway, msg: err.Error()} + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return nil, true, err + } + return nil, true, homeStatusErr{code: http.StatusServiceUnavailable, msg: "home refresh temporarily unavailable"} } var env homeErrorEnvelope @@ -88,11 +94,15 @@ func RefreshAuthViaHome(ctx context.Context, cfg *config.Config, auth *cliproxya if code == "" { code = strings.TrimSpace(env.Error.Code) } - msg := strings.TrimSpace(env.Error.Message) - if msg == "" { - msg = "home returned error" + statusCode := statusFromHomeErrorCode(code) + message := "credential refresh temporarily unavailable" + switch statusCode { + case http.StatusUnauthorized: + message = "credential unauthorized" + case http.StatusNotFound: + message = "credential refresh target not found" } - return nil, true, homeStatusErr{code: statusFromHomeErrorCode(code), msg: msg} + return nil, true, homeStatusErr{code: statusCode, msg: message} } updated, returnedIndex, errParse := parseHomeRefreshAuth(raw) @@ -107,6 +117,43 @@ func RefreshAuthViaHome(ctx context.Context, cfg *config.Config, auth *cliproxya return updated, true, nil } +func authAccessTokenSHA256(auth *cliproxyauth.Auth) string { + accessToken := authAccessTokenForFingerprint(auth) + if accessToken == "" { + return "" + } + digest := sha256.Sum256([]byte(accessToken)) + return hex.EncodeToString(digest[:]) +} + +func authAccessTokenForFingerprint(auth *cliproxyauth.Auth) string { + if auth == nil || auth.Metadata == nil { + return "" + } + for _, key := range []string{"access_token", "accessToken"} { + if value, ok := auth.Metadata[key].(string); ok && strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + for _, key := range []string{"token", "Token"} { + switch token := auth.Metadata[key].(type) { + case map[string]any: + for _, tokenKey := range []string{"access_token", "accessToken"} { + if value, ok := token[tokenKey].(string); ok && strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + case map[string]string: + for _, tokenKey := range []string{"access_token", "accessToken"} { + if value := strings.TrimSpace(token[tokenKey]); value != "" { + return value + } + } + } + } + return "" +} + func parseHomeRefreshAuth(raw []byte) (*cliproxyauth.Auth, string, error) { var rawObject map[string]json.RawMessage if errUnmarshal := json.Unmarshal(raw, &rawObject); errUnmarshal != nil { @@ -128,11 +175,13 @@ func parseHomeRefreshAuth(raw []byte) (*cliproxyauth.Auth, string, error) { func statusFromHomeErrorCode(code string) int { switch strings.ToLower(strings.TrimSpace(code)) { - case "authentication_error", "unauthorized": + case "authentication_error", "unauthorized", "invalid_grant", "refresh_token_expired", "refresh_token_revoked", "refresh_token_reused": return http.StatusUnauthorized case "model_not_found": return http.StatusNotFound + case "auth_not_found", "auth_unavailable", "refresh_temporarily_unavailable", "refresh_unsupported", "home_unavailable": + return http.StatusServiceUnavailable default: - return http.StatusBadGateway + return http.StatusServiceUnavailable } } diff --git a/internal/runtime/executor/helps/home_refresh_test.go b/internal/runtime/executor/helps/home_refresh_test.go index ca758273..be33016d 100644 --- a/internal/runtime/executor/helps/home_refresh_test.go +++ b/internal/runtime/executor/helps/home_refresh_test.go @@ -3,7 +3,9 @@ package helps import ( "context" "encoding/json" + "errors" "net/http" + "strings" "sync/atomic" "testing" @@ -18,22 +20,96 @@ func TestStatusFromHomeErrorCodeMapsAuthenticationErrorToUnauthorized(t *testing if got := statusFromHomeErrorCode("unauthorized"); got != http.StatusUnauthorized { t.Fatalf("statusFromHomeErrorCode(unauthorized) = %d, want %d", got, http.StatusUnauthorized) } + for _, code := range []string{"auth_not_found", "auth_unavailable", "refresh_temporarily_unavailable", "refresh_unsupported"} { + if got := statusFromHomeErrorCode(code); got != http.StatusServiceUnavailable { + t.Fatalf("statusFromHomeErrorCode(%s) = %d, want %d", code, got, http.StatusServiceUnavailable) + } + } } type fakeHomeRefreshClient struct { - calls atomic.Int32 - authIndex string - raw []byte + calls atomic.Int32 + authIndex string + accessTokenHash string + raw []byte + err error } func (c *fakeHomeRefreshClient) HeartbeatOK() bool { return true } -func (c *fakeHomeRefreshClient) GetRefreshAuth(_ context.Context, authIndex string) ([]byte, error) { +func (c *fakeHomeRefreshClient) GetRefreshAuth(_ context.Context, authIndex string, accessTokenHash string) ([]byte, error) { c.calls.Add(1) c.authIndex = authIndex - return c.raw, nil + c.accessTokenHash = accessTokenHash + return c.raw, c.err +} + +func TestRefreshAuthViaHomePreservesContextErrors(t *testing.T) { + client := &fakeHomeRefreshClient{err: context.DeadlineExceeded} + oldCurrentHomeRefreshClient := currentHomeRefreshClient + currentHomeRefreshClient = func() homeRefreshClient { return client } + t.Cleanup(func() { currentHomeRefreshClient = oldCurrentHomeRefreshClient }) + + cfg := &config.Config{Home: config.HomeConfig{Enabled: true}} + auth := &cliproxyauth.Auth{ID: "home-auth", Index: "home-auth", Provider: "codex"} + _, handled, errRefresh := RefreshAuthViaHome(context.Background(), cfg, auth) + if !handled || !errors.Is(errRefresh, context.DeadlineExceeded) { + t.Fatalf("RefreshAuthViaHome() = handled %v err %v, want true/context.DeadlineExceeded", handled, errRefresh) + } +} + +func TestRefreshAuthViaHomeMapsTransportFailureToRedacted503(t *testing.T) { + client := &fakeHomeRefreshClient{err: errors.New("dial failed with provider-secret")} + oldCurrentHomeRefreshClient := currentHomeRefreshClient + currentHomeRefreshClient = func() homeRefreshClient { return client } + t.Cleanup(func() { currentHomeRefreshClient = oldCurrentHomeRefreshClient }) + + cfg := &config.Config{Home: config.HomeConfig{Enabled: true}} + auth := &cliproxyauth.Auth{ID: "home-auth", Index: "home-auth", Provider: "codex"} + _, handled, errRefresh := RefreshAuthViaHome(context.Background(), cfg, auth) + statusErr, okStatus := errRefresh.(interface{ StatusCode() int }) + if !handled || !okStatus || statusErr.StatusCode() != http.StatusServiceUnavailable { + t.Fatalf("RefreshAuthViaHome() = handled %v err %v, want redacted 503", handled, errRefresh) + } + if strings.Contains(errRefresh.Error(), "provider-secret") { + t.Fatalf("refresh error leaked transport detail: %v", errRefresh) + } +} + +func TestRefreshAuthViaHomeRedactsLegacyErrorEnvelope(t *testing.T) { + client := &fakeHomeRefreshClient{raw: []byte(`{"error":{"type":"error","message":"provider response: refresh_token=provider-secret"}}`)} + oldCurrentHomeRefreshClient := currentHomeRefreshClient + currentHomeRefreshClient = func() homeRefreshClient { return client } + t.Cleanup(func() { currentHomeRefreshClient = oldCurrentHomeRefreshClient }) + + cfg := &config.Config{Home: config.HomeConfig{Enabled: true}} + auth := &cliproxyauth.Auth{ID: "home-auth", Index: "home-auth", Provider: "codex"} + _, handled, errRefresh := RefreshAuthViaHome(context.Background(), cfg, auth) + statusErr, okStatus := errRefresh.(interface{ StatusCode() int }) + if !handled || !okStatus || statusErr.StatusCode() != http.StatusServiceUnavailable { + t.Fatalf("RefreshAuthViaHome() = handled %v err %v, want redacted 503", handled, errRefresh) + } + if strings.Contains(errRefresh.Error(), "provider-secret") { + t.Fatalf("refresh error leaked legacy Home detail: %v", errRefresh) + } +} + +func TestAuthAccessTokenSHA256SupportsKnownMetadataShapes(t *testing.T) { + want := authAccessTokenSHA256(&cliproxyauth.Auth{Metadata: map[string]any{"access_token": "same-token"}}) + cases := map[string]*cliproxyauth.Auth{ + "camel case": {Metadata: map[string]any{"accessToken": "same-token"}}, + "nested any map": {Metadata: map[string]any{"token": map[string]any{"access_token": "same-token"}}}, + "nested string map": {Metadata: map[string]any{"Token": map[string]string{"accessToken": "same-token"}}}, + } + for name, auth := range cases { + t.Run(name, func(t *testing.T) { + if got := authAccessTokenSHA256(auth); got == "" || got != want { + t.Fatalf("token hash = %q, want %q", got, want) + } + }) + } } func TestRefreshAuthViaHomeAcceptsAuthEnvelope(t *testing.T) { @@ -69,6 +145,7 @@ func TestRefreshAuthViaHomeAcceptsAuthEnvelope(t *testing.T) { Provider: "antigravity", Index: "home-index-1", Metadata: map[string]any{ + "access_token": "old-access-token", "refresh_token": "refresh-token", }, } @@ -86,6 +163,9 @@ func TestRefreshAuthViaHomeAcceptsAuthEnvelope(t *testing.T) { if client.authIndex != "home-index-1" { t.Fatalf("home refresh auth_index = %q, want home-index-1", client.authIndex) } + if client.accessTokenHash != authAccessTokenSHA256(auth) { + t.Fatalf("home refresh access token hash = %q, want %q", client.accessTokenHash, authAccessTokenSHA256(auth)) + } if updated == nil { t.Fatal("updated auth = nil") } diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index 52e1687f..39a320fe 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -22,24 +22,25 @@ import ( ) type UsageReporter struct { - provider string - executorType string - model string - alias string - authID string - authIndex string - authType string - apiKey string - source string - reasoning string - serviceTier string - generate bool - requestedAt time.Time - ttftMu sync.RWMutex - ttft time.Duration - ttftStart time.Time - ttftSet bool - once sync.Once + provider string + executorType string + model string + alias string + authID string + authIndex string + accessTokenHash string + authType string + apiKey string + source string + reasoning string + serviceTier string + generate bool + requestedAt time.Time + ttftMu sync.RWMutex + ttft time.Duration + ttftStart time.Time + ttftSet bool + once sync.Once } type usageExecutor interface { @@ -77,6 +78,7 @@ func NewUsageReporter(ctx context.Context, provider, model string, auth *cliprox if auth != nil { reporter.authID = auth.ID reporter.authIndex = auth.EnsureIndex() + reporter.accessTokenHash = authAccessTokenSHA256(auth) } return reporter } @@ -264,6 +266,7 @@ func (r *UsageReporter) buildRecordForModel(model string, detail usage.Detail, f APIKey: r.apiKey, AuthID: r.authID, AuthIndex: r.authIndex, + AccessTokenSHA256: r.accessTokenHash, AuthType: r.authType, ReasoningEffort: r.reasoning, ServiceTier: r.serviceTier, diff --git a/sdk/cliproxy/auth/conductor_execution.go b/sdk/cliproxy/auth/conductor_execution.go index 442e5d2f..fae7a058 100644 --- a/sdk/cliproxy/auth/conductor_execution.go +++ b/sdk/cliproxy/auth/conductor_execution.go @@ -546,6 +546,7 @@ func (m *Manager) executeStreamMixedOnce(ctx context.Context, providers []string homeAuthCount := 1 tried := make(map[string]struct{}) attempted := make(map[string]struct{}) + unauthorizedRefreshTried := make(map[string]struct{}) var lastErr error for { if !homeMode && maxRetryCredentials > 0 && len(attempted) >= maxRetryCredentials { @@ -586,6 +587,15 @@ func (m *Manager) executeStreamMixedOnce(ctx context.Context, providers []string } return nil, &Error{Code: "executor_not_found", Message: "executor not registered"} } + if selection != nil { + if _, refreshedAlready := unauthorizedRefreshTried[auth.ID]; refreshedAlready { + selection.End("repeated_refresh_auth") + if lastErr != nil { + return nil, lastErr + } + return nil, repeatedHomeAuthError() + } + } entry := logEntryWithRequestID(ctx) debugLogAuthSelection(entry, auth, provider, routeModel) @@ -666,7 +676,7 @@ func (m *Manager) executeStreamMixedOnce(ctx context.Context, providers []string models = models[:1] pooled = false } - streamResult, errStream := m.executeStreamWithModelPool(execCtx, executor, auth, provider, execReq, execOpts, routeModel, streamExecutionModel, models, pooled, aliasResult, routing, !homeMode, selection != nil) + streamResult, errStream := m.executeStreamWithModelPool(execCtx, executor, auth, provider, execReq, execOpts, routeModel, streamExecutionModel, models, pooled, aliasResult, routing, !homeMode || selection != nil, selection != nil, unauthorizedRefreshTried) if errStream != nil { if selection != nil { releaseAttempt() diff --git a/sdk/cliproxy/auth/conductor_home.go b/sdk/cliproxy/auth/conductor_home.go index fb576f93..9da18647 100644 --- a/sdk/cliproxy/auth/conductor_home.go +++ b/sdk/cliproxy/auth/conductor_home.go @@ -144,7 +144,19 @@ func shouldReturnLastErrorOnPickFailure(homeMode bool, lastErr error, errPick er if !homeMode { return true } - return isHomeRequestRetryExceededError(errPick) + if isHomeRequestRetryExceededError(errPick) { + return true + } + var authErr *Error + if !errors.As(errPick, &authErr) || authErr == nil { + return false + } + switch strings.ToLower(strings.TrimSpace(authErr.Code)) { + case "auth_not_found", "auth_unavailable": + return true + default: + return false + } } func homeAuthAlreadyTried(tried map[string]struct{}, authID string) bool { @@ -435,14 +447,18 @@ func (m *Manager) endHomeSelectionBeforeRedispatch(ctx context.Context, selectio } func (m *Manager) retainHomeWebsocketSelection(ctx context.Context, opts cliproxyexecutor.Options, model string, selection *HomeDispatchSelection) bool { - if m == nil || selection == nil || !selection.Retained() || !cliproxyexecutor.DownstreamWebsocket(ctx) || selection.Auth == nil { + if m == nil || selection == nil || !selection.Retained() || !cliproxyexecutor.DownstreamWebsocket(ctx) { + return false + } + selectionAuth := selection.CloneAuth() + if selectionAuth == nil { return false } sessionID := homeExecutionSessionIDFromMetadata(opts.Metadata) - credentialID := strings.TrimSpace(selection.Auth.ID) + credentialID := strings.TrimSpace(selectionAuth.ID) routeModel, validRouteModel := validCanonicalHomeConcurrencyModelKey(model) if selection.accountedModel == "" { - selection.accountedModel, _ = m.predictedHomeConcurrencyModel(selection.Auth, model) + selection.accountedModel, _ = m.predictedHomeConcurrencyModel(selectionAuth, model) } if sessionID == "" || credentialID == "" || !validRouteModel || selection.accountedModel == "" { return false @@ -461,7 +477,7 @@ func (m *Manager) retainHomeWebsocketSelection(ctx context.Context, opts cliprox previous := selections[key] selections[key] = selection m.mu.Unlock() - m.rememberHomeRuntimeAuth(sessionID, selection.Auth) + m.rememberHomeRuntimeAuth(sessionID, selectionAuth) if previous != nil && previous != selection { previous.End("target_replaced") } @@ -537,11 +553,15 @@ func (m *Manager) clearHomeRuntimeAuthsForSessionLocked(sessionID string) { } func (m *Manager) bindHomeSelectionRuntimeAuth(ctx context.Context, opts cliproxyexecutor.Options, selection *HomeDispatchSelection) error { - if m == nil || selection == nil || !cliproxyexecutor.DownstreamWebsocket(ctx) || selection.Auth == nil || !authWebsocketsEnabled(selection.Auth) { + if m == nil || selection == nil || !cliproxyexecutor.DownstreamWebsocket(ctx) { + return nil + } + selectionAuth := selection.CloneAuth() + if selectionAuth == nil || !authWebsocketsEnabled(selectionAuth) { return nil } sessionID := homeExecutionSessionIDFromMetadata(opts.Metadata) - authID := strings.TrimSpace(selection.Auth.ID) + authID := strings.TrimSpace(selectionAuth.ID) if sessionID == "" || authID == "" || !selection.runtimeAuthBound.CompareAndSwap(false, true) { return nil } @@ -558,11 +578,15 @@ func (m *Manager) bindHomeSelectionRuntimeAuth(ctx context.Context, opts cliprox } func (m *Manager) rememberHomeSelectionRuntimeAuth(sessionID string, selection *HomeDispatchSelection) { - if m == nil || selection == nil || selection.Auth == nil { + if m == nil || selection == nil { + return + } + selectionAuth := selection.CloneAuth() + if selectionAuth == nil { return } sessionID = strings.TrimSpace(sessionID) - authID := strings.TrimSpace(selection.Auth.ID) + authID := strings.TrimSpace(selectionAuth.ID) if sessionID == "" || authID == "" { return } @@ -579,11 +603,33 @@ func (m *Manager) rememberHomeSelectionRuntimeAuth(sessionID string, selection * if m.homeRuntimeAuthOwners[sessionID] == nil { m.homeRuntimeAuthOwners[sessionID] = make(map[string]*HomeDispatchSelection) } - m.homeRuntimeAuths[sessionID][authID] = selection.Auth.Clone() + m.homeRuntimeAuths[sessionID][authID] = selectionAuth m.homeRuntimeAuthOwners[sessionID][authID] = selection m.mu.Unlock() } +func (m *Manager) replaceHomeSelectionAuth(selection *HomeDispatchSelection, auth *Auth) { + if m == nil || selection == nil || auth == nil { + return + } + m.mu.Lock() + selection.ReplaceAuth(auth) + updated := selection.CloneAuth() + if updated == nil { + m.mu.Unlock() + return + } + for sessionID, owners := range m.homeRuntimeAuthOwners { + for authID, owner := range owners { + if owner != selection || m.homeRuntimeAuths[sessionID] == nil { + continue + } + m.homeRuntimeAuths[sessionID][authID] = updated.Clone() + } + } + m.mu.Unlock() +} + func (m *Manager) forgetHomeRuntimeAuth(sessionID string, authID string, owner *HomeDispatchSelection) { sessionID = strings.TrimSpace(sessionID) authID = strings.TrimSpace(authID) @@ -669,7 +715,8 @@ func (m *Manager) pickNextViaHome(ctx context.Context, model string, opts clipro if errSelection != nil { return nil, nil, "", errSelection } - if selection.Auth == nil || homeAuthAlreadyTried(tried, selection.Auth.ID) { + selectionAuth := selection.CloneAuth() + if selectionAuth == nil || homeAuthAlreadyTried(tried, selectionAuth.ID) { selection.End("repeated_auth") return nil, nil, "", repeatedHomeAuthError() } @@ -1118,7 +1165,7 @@ func (m *Manager) tryAntigravityCreditsExecuteStream(ctx context.Context, req cl if len(models) == 0 { continue } - result, errStream := m.executeStreamWithModelPool(creditsCtx, c.executor, c.auth, c.provider, req, creditsOpts, routeModel, "", models, pooled, aliasResult, routing, true, false) + result, errStream := m.executeStreamWithModelPool(creditsCtx, c.executor, c.auth, c.provider, req, creditsOpts, routeModel, "", models, pooled, aliasResult, routing, true, false, nil) if errStream != nil { continue } diff --git a/sdk/cliproxy/auth/conductor_home_execution.go b/sdk/cliproxy/auth/conductor_home_execution.go index dfd14ee0..a3a0ba8f 100644 --- a/sdk/cliproxy/auth/conductor_home_execution.go +++ b/sdk/cliproxy/auth/conductor_home_execution.go @@ -21,7 +21,7 @@ func (m *Manager) executeHome(ctx context.Context, providers []string, req clipr for homeAuthCount := 1; ; homeAuthCount++ { selection, errSelection := m.pickHomeDispatchSelection(ctx, routeModel, withHomeAuthCount(opts, homeAuthCount)) if errSelection != nil { - if lastErr != nil && isHomeRequestRetryExceededError(errSelection) { + if shouldReturnLastErrorOnPickFailure(true, lastErr, errSelection) { return cliproxyexecutor.Response{}, lastErr } return cliproxyexecutor.Response{}, errSelection @@ -81,6 +81,7 @@ func (m *Manager) executeHome(ctx context.Context, providers []string, req clipr lastErr = errPrepare continue } + didRefreshOnUnauthorized := false for _, upstreamModel := range models { resultModel := m.stateModelForExecution(preparedAuth, routeModel, upstreamModel, pooled) execReq := req @@ -107,10 +108,23 @@ func (m *Manager) executeHome(ctx context.Context, providers []string, req clipr } var response cliproxyexecutor.Response var errExecute error - if countTokens { - response, errExecute = selection.Executor.CountTokens(execCtx, preparedAuth, execReq, execOpts) - } else { - response, errExecute = selection.Executor.Execute(execCtx, preparedAuth, execReq, execOpts) + execute := func() (cliproxyexecutor.Response, error) { + if countTokens { + return selection.Executor.CountTokens(execCtx, preparedAuth, execReq, execOpts) + } + return selection.Executor.Execute(execCtx, preparedAuth, execReq, execOpts) + } + response, errExecute = execute() + if errExecute != nil { + if refreshed, okRefresh, errRefresh := m.tryRefreshExecutionAuthAfterUnauthorized(execCtx, selection.Executor, preparedAuth, errExecute, didRefreshOnUnauthorized, true); errRefresh != nil { + errExecute = errRefresh + } else if okRefresh { + preparedAuth = refreshed + m.replaceHomeSelectionAuth(selection, preparedAuth) + didRefreshOnUnauthorized = true + publishSelectedAuthMetadata(opts.Metadata, preparedAuth) + response, errExecute = execute() + } } result := Result{AuthID: preparedAuth.ID, Provider: selection.Provider, Model: resultModel, Success: errExecute == nil} if errExecute == nil { diff --git a/sdk/cliproxy/auth/conductor_refresh.go b/sdk/cliproxy/auth/conductor_refresh.go index 9d95577d..28f9b831 100644 --- a/sdk/cliproxy/auth/conductor_refresh.go +++ b/sdk/cliproxy/auth/conductor_refresh.go @@ -377,8 +377,77 @@ func clearUnauthorizedModelStates(auth *Auth, now time.Time) []string { return resumed } -// tryRefreshAfterUnauthorized refreshes OAuth credentials once after a 401 so the -// current auth can be retried before fallback/suspend. +// tryRefreshExecutionAuthAfterUnauthorized refreshes OAuth credentials once for +// either a local auth or an ephemeral Home dispatch auth. +func (m *Manager) tryRefreshExecutionAuthAfterUnauthorized(ctx context.Context, executor ProviderExecutor, auth *Auth, execErr error, alreadyTried bool, homeDispatch bool) (*Auth, bool, error) { + if !homeDispatch { + refreshed, ok := m.tryRefreshAfterUnauthorized(ctx, auth, execErr, alreadyTried) + return refreshed, ok, nil + } + if m == nil || executor == nil || auth == nil || alreadyTried || execErr == nil { + return auth, false, nil + } + if !isUnauthorizedError(execErr) || auth.AuthKind() != AuthKindOAuth { + return auth, false, nil + } + + log.Debugf("unauthorized Home response for %s (%s), refreshing credentials before redispatch", auth.Provider, auth.ID) + target := auth.Clone() + updated, errRefresh := executor.Refresh(ctx, target) + if errRefresh != nil { + log.Debugf("Home credential refresh before redispatch failed for %s (%s)", auth.Provider, auth.ID) + return auth, false, errRefresh + } + if updated == nil { + updated = target + } + if updated.ID == "" { + updated.ID = auth.ID + } + if updated.Index == "" { + updated.Index = auth.Index + } + if updated.Provider == "" { + updated.Provider = auth.Provider + } + if updated.Runtime == nil { + updated.Runtime = auth.Runtime + } + preserveHomeRoutingAttributes(updated, auth) + return updated, true, nil +} + +// RefreshHomeSelectionAfterUnauthorized refreshes the credential snapshot that +// received a 401, or reuses a newer token already installed on the selection. +func (m *Manager) RefreshHomeSelectionAfterUnauthorized(ctx context.Context, selection *HomeDispatchSelection, failedAuth *Auth) (*Auth, bool, error) { + if m == nil || selection == nil { + return nil, false, nil + } + current := selection.CloneAuth() + if failedAuth == nil { + failedAuth = current + } + if current != nil && failedAuth != nil && current.ID == failedAuth.ID { + currentToken := authAccessToken(current) + failedToken := authAccessToken(failedAuth) + if currentToken != "" && failedToken != "" && currentToken != failedToken { + return current, true, nil + } + } + refreshed, okRefresh, errRefresh := m.tryRefreshExecutionAuthAfterUnauthorized(ctx, selection.Executor, failedAuth, &Error{HTTPStatus: http.StatusUnauthorized, Message: "upstream unauthorized"}, false, true) + if errRefresh != nil || !okRefresh { + return current, false, errRefresh + } + m.replaceHomeSelectionAuth(selection, refreshed) + updated := selection.CloneAuth() + if updated == nil { + return nil, false, &Error{Code: "auth_not_found", Message: "refreshed Home auth is unavailable", HTTPStatus: http.StatusServiceUnavailable} + } + return updated, true, nil +} + +// tryRefreshAfterUnauthorized refreshes local 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 diff --git a/sdk/cliproxy/auth/conductor_selection.go b/sdk/cliproxy/auth/conductor_selection.go index 9578e178..6a8562d1 100644 --- a/sdk/cliproxy/auth/conductor_selection.go +++ b/sdk/cliproxy/auth/conductor_selection.go @@ -1180,14 +1180,15 @@ func (m *Manager) SelectHomeAuthByKind(ctx context.Context, provider string, mod return nil, errSelection } providerMatches := strings.TrimSpace(provider) == "" || strings.EqualFold(strings.TrimSpace(selection.Provider), strings.TrimSpace(provider)) - kindMatches := selection.Auth != nil && selection.Auth.AuthKind() == requiredKind + selectionAuth := selection.CloneAuth() + kindMatches := selectionAuth != nil && selectionAuth.AuthKind() == requiredKind if providerMatches && kindMatches { return selection, nil } authID := "" - if selection.Auth != nil { - authID = strings.TrimSpace(selection.Auth.ID) + if selectionAuth != nil { + authID = strings.TrimSpace(selectionAuth.ID) } reason := "auth_kind_mismatch" if !providerMatches { diff --git a/sdk/cliproxy/auth/conductor_stream.go b/sdk/cliproxy/auth/conductor_stream.go index f3209963..2d79d225 100644 --- a/sdk/cliproxy/auth/conductor_stream.go +++ b/sdk/cliproxy/auth/conductor_stream.go @@ -180,13 +180,24 @@ func (m *Manager) wrapStreamResult(ctx context.Context, auth *Auth, provider, re return &cliproxyexecutor.StreamResult{Headers: headers, Chunks: out} } -func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor ProviderExecutor, auth *Auth, provider string, req cliproxyexecutor.Request, opts cliproxyexecutor.Options, routeModel, executionModel string, execModels []string, pooled bool, aliasResult OAuthModelAliasResult, routing *apiKeyModelRoutingSnapshot, allowRetry bool, ephemeralResult bool) (*cliproxyexecutor.StreamResult, error) { +func (m *Manager) replaceHomeExecutionLifecycleAuth(lifecycle cliproxyexecutor.ExecutionLifecycle, auth *Auth) { + selection, ok := lifecycle.(*HomeDispatchSelection) + if !ok || selection == nil { + return + } + m.replaceHomeSelectionAuth(selection, auth) +} + +func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor ProviderExecutor, auth *Auth, provider string, req cliproxyexecutor.Request, opts cliproxyexecutor.Options, routeModel, executionModel string, execModels []string, pooled bool, aliasResult OAuthModelAliasResult, routing *apiKeyModelRoutingSnapshot, allowRetry bool, ephemeralResult bool, unauthorizedRefreshTried map[string]struct{}) (*cliproxyexecutor.StreamResult, error) { if executor == nil { return nil, &Error{Code: "executor_not_found", Message: "executor not registered"} } ctx = contextWithRequestedModelAlias(ctx, opts, routeModel) var lastErr error didRefreshOnUnauthorized := false + if auth != nil && unauthorizedRefreshTried != nil { + _, didRefreshOnUnauthorized = unauthorizedRefreshTried[auth.ID] + } for idx, execModel := range execModels { resultModel := m.stateModelForExecution(auth, routeModel, execModel, pooled) execReq := req @@ -212,8 +223,21 @@ func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor Provi return nil, errCtx } if allowRetry { - if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(ctx, auth, errStream, didRefreshOnUnauthorized); okRefresh { + alreadyTried := didRefreshOnUnauthorized + willAttemptHomeRefresh := ephemeralResult && !alreadyTried && auth != nil && auth.AuthKind() == AuthKindOAuth && isUnauthorizedError(errStream) + refreshed, okRefresh, errRefresh := m.tryRefreshExecutionAuthAfterUnauthorized(ctx, executor, auth, errStream, alreadyTried, ephemeralResult) + if willAttemptHomeRefresh { + didRefreshOnUnauthorized = true + if unauthorizedRefreshTried != nil { + unauthorizedRefreshTried[auth.ID] = struct{}{} + } + } + if errRefresh != nil { + errStream = errRefresh + } else if okRefresh { auth = refreshed + m.replaceHomeExecutionLifecycleAuth(execOpts.ExecutionLifecycle, auth) + publishSelectedAuthMetadata(execOpts.Metadata, auth) didRefreshOnUnauthorized = true streamResult, errStream = executor.ExecuteStream(ctx, auth, execReq, execOpts) if errStream != nil { @@ -251,9 +275,24 @@ func (m *Manager) executeStreamWithModelPool(ctx context.Context, executor Provi return nil, errCtx } if allowRetry { - if refreshed, okRefresh := m.tryRefreshAfterUnauthorized(ctx, auth, bootstrapErr, didRefreshOnUnauthorized); okRefresh { + alreadyTried := didRefreshOnUnauthorized + willAttemptHomeRefresh := ephemeralResult && !alreadyTried && auth != nil && auth.AuthKind() == AuthKindOAuth && isUnauthorizedError(bootstrapErr) + refreshed, okRefresh, errRefresh := m.tryRefreshExecutionAuthAfterUnauthorized(ctx, executor, auth, bootstrapErr, alreadyTried, ephemeralResult) + if willAttemptHomeRefresh { + didRefreshOnUnauthorized = true + if unauthorizedRefreshTried != nil { + unauthorizedRefreshTried[auth.ID] = struct{}{} + } + } + if errRefresh != nil { + discardStreamChunks(streamResult.Chunks) + bootstrapErr = errRefresh + streamResult = &cliproxyexecutor.StreamResult{} + } else if okRefresh { discardStreamChunks(streamResult.Chunks) auth = refreshed + m.replaceHomeExecutionLifecycleAuth(execOpts.ExecutionLifecycle, auth) + publishSelectedAuthMetadata(execOpts.Metadata, auth) didRefreshOnUnauthorized = true retryStream, retryErr := executor.ExecuteStream(ctx, auth, execReq, execOpts) if retryErr != nil { diff --git a/sdk/cliproxy/auth/home_concurrency.go b/sdk/cliproxy/auth/home_concurrency.go index cc961ec2..d3f91774 100644 --- a/sdk/cliproxy/auth/home_concurrency.go +++ b/sdk/cliproxy/auth/home_concurrency.go @@ -242,7 +242,8 @@ func decodeHomeDispatchError(raw []byte) error { case "credential_concurrency_exceeded", "credential_model_concurrency_exceeded": result.HTTPStatus = http.StatusTooManyRequests return newHomeConcurrencyBusyError(result, time.Duration(detail.RetryAfterMS)*time.Millisecond) - case "concurrency_protocol_required", "concurrency_tracker_unavailable", "concurrency_node_unavailable": + case "auth_not_found", "auth_unavailable", "refresh_temporarily_unavailable", "home_unavailable", + "concurrency_protocol_required", "concurrency_tracker_unavailable", "concurrency_node_unavailable": result.HTTPStatus = http.StatusServiceUnavailable } return result diff --git a/sdk/cliproxy/auth/home_concurrency_test.go b/sdk/cliproxy/auth/home_concurrency_test.go index 5bd42ca7..7408fe6b 100644 --- a/sdk/cliproxy/auth/home_concurrency_test.go +++ b/sdk/cliproxy/auth/home_concurrency_test.go @@ -384,6 +384,18 @@ func TestHomeBusyErrorMaps429AndRetryAfter(t *testing.T) { } } +func TestHomeNoCandidateErrorsMapToServiceUnavailable(t *testing.T) { + for _, code := range []string{"auth_not_found", "auth_unavailable"} { + t.Run(code, func(t *testing.T) { + errDispatch := decodeHomeDispatchError([]byte(fmt.Sprintf(`{"error":{"type":%q,"message":"no auth available"}}`, code))) + var authErr *Error + if !errors.As(errDispatch, &authErr) || authErr.Code != code || authErr.HTTPStatus != http.StatusServiceUnavailable { + t.Fatalf("decodeHomeDispatchError(%s) = %#v, want 503", code, errDispatch) + } + }) + } +} + func TestHomeConcurrencyTupleAuthMismatchEndsScope(t *testing.T) { dispatcher := &fixtureHomeDispatcher{payload: []byte(`{"concurrency":{"accounted":true,"credential_id":"cred-1","model":"gpt"},"auth_index":"other","auth":{"id":"cred-1","provider":"codex"}}`)} manager := newHomeSelectionTestManager(t, dispatcher) diff --git a/sdk/cliproxy/auth/home_selection.go b/sdk/cliproxy/auth/home_selection.go index 01a39b32..33584be6 100644 --- a/sdk/cliproxy/auth/home_selection.go +++ b/sdk/cliproxy/auth/home_selection.go @@ -141,6 +141,7 @@ type HomeDispatchSelection struct { Executor ProviderExecutor Provider string + authMu sync.RWMutex scope *executionregistry.Scope accountedModel string resources *executionResources @@ -249,9 +250,40 @@ func (s *HomeDispatchSelection) EndWithRelease(reason string) *executionregistry return s.scope.EndWithRelease("") } +// ReplaceAuth updates the selection after Home returns refreshed credentials. +func (s *HomeDispatchSelection) ReplaceAuth(auth *Auth) { + if s == nil || auth == nil { + return + } + updated := auth.Clone() + s.authMu.Lock() + defer s.authMu.Unlock() + preserveHomeRoutingAttributes(updated, s.Auth) + s.Auth = updated +} + +func preserveHomeRoutingAttributes(updated, previous *Auth) { + if updated == nil || previous == nil { + return + } + if updated.Attributes == nil { + updated.Attributes = make(map[string]string) + } + for _, key := range []string{homeUpstreamModelAttributeKey, homeForceMappingAttributeKey, homeOriginalAliasAttributeKey} { + if value := strings.TrimSpace(previous.Attributes[key]); value != "" { + updated.Attributes[key] = value + } + } +} + // CloneAuth returns a standalone auth copy without the selection handle. func (s *HomeDispatchSelection) CloneAuth() *Auth { - if s == nil || s.Auth == nil { + if s == nil { + return nil + } + s.authMu.RLock() + defer s.authMu.RUnlock() + if s.Auth == nil { return nil } return s.Auth.Clone() diff --git a/sdk/cliproxy/auth/home_selection_test.go b/sdk/cliproxy/auth/home_selection_test.go index 1f02fc5c..56cbe29d 100644 --- a/sdk/cliproxy/auth/home_selection_test.go +++ b/sdk/cliproxy/auth/home_selection_test.go @@ -42,6 +42,70 @@ func TestHomeDispatchSelectionOwnsScopeOutsideAuth(t *testing.T) { } } +func TestHomeDispatchSelectionReplaceAuthPreservesRoutingAttributes(t *testing.T) { + selection := &HomeDispatchSelection{Auth: &Auth{ + ID: "cred-1", + Provider: "codex", + Attributes: map[string]string{ + homeUpstreamModelAttributeKey: "gpt-5-upstream", + homeForceMappingAttributeKey: "true", + homeOriginalAliasAttributeKey: "team/gpt-5", + }, + Metadata: map[string]any{"access_token": "old"}, + }} + + selection.ReplaceAuth(&Auth{ + ID: "cred-1", + Provider: "codex", + Attributes: map[string]string{AttributeAuthKind: AuthKindOAuth}, + Metadata: map[string]any{"access_token": "fresh"}, + }) + + updated := selection.CloneAuth() + if updated == nil || updated.Metadata["access_token"] != "fresh" { + t.Fatalf("updated auth = %#v", updated) + } + if updated.Attributes[homeUpstreamModelAttributeKey] != "gpt-5-upstream" || updated.Attributes[homeForceMappingAttributeKey] != "true" || updated.Attributes[homeOriginalAliasAttributeKey] != "team/gpt-5" { + t.Fatalf("routing attributes were not preserved: %#v", updated.Attributes) + } +} + +func TestHomeDispatchSelectionReplaceAuthConcurrentClone(t *testing.T) { + selection := &HomeDispatchSelection{Auth: &Auth{ID: "cred-1", Metadata: map[string]any{"access_token": "old"}}} + done := make(chan struct{}) + go func() { + defer close(done) + for i := 0; i < 1000; i++ { + selection.ReplaceAuth(&Auth{ID: "cred-1", Metadata: map[string]any{"access_token": "fresh"}}) + } + }() + for i := 0; i < 1000; i++ { + if auth := selection.CloneAuth(); auth == nil || auth.ID != "cred-1" { + t.Fatalf("CloneAuth() = %#v", auth) + } + } + <-done +} + +func TestReplaceHomeSelectionAuthUpdatesRetainedRuntimeAuth(t *testing.T) { + selection := &HomeDispatchSelection{Auth: &Auth{ID: "cred-1", Provider: "codex", Metadata: map[string]any{"access_token": "old"}}} + manager := &Manager{ + homeRuntimeAuths: map[string]map[string]*Auth{ + "session-1": {"cred-1": selection.Auth.Clone()}, + }, + homeRuntimeAuthOwners: map[string]map[string]*HomeDispatchSelection{ + "session-1": {"cred-1": selection}, + }, + } + + manager.replaceHomeSelectionAuth(selection, &Auth{ID: "cred-1", Provider: "codex", Metadata: map[string]any{"access_token": "fresh"}}) + + retained := manager.homeRuntimeAuths["session-1"]["cred-1"] + if retained == nil || retained.Metadata["access_token"] != "fresh" { + t.Fatalf("retained runtime auth = %#v, want fresh token", retained) + } +} + func TestHomeDispatchSelectionDrainsResourcesAddedDuringEnd(t *testing.T) { registry := executionregistry.New() pending, errBegin := registry.BeginDispatch() diff --git a/sdk/cliproxy/auth/home_unauthorized_refresh_test.go b/sdk/cliproxy/auth/home_unauthorized_refresh_test.go new file mode 100644 index 00000000..80d8f7f9 --- /dev/null +++ b/sdk/cliproxy/auth/home_unauthorized_refresh_test.go @@ -0,0 +1,340 @@ +package auth + +import ( + "context" + "encoding/json" + "net/http" + "sync/atomic" + "testing" + + internalconfig "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executionregistry" + cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor" +) + +const homeUnauthorizedRefreshProvider = "home-unauthorized-refresh" + +type homeUnauthorizedRefreshDispatcher struct { + calls atomic.Int32 +} + +func (*homeUnauthorizedRefreshDispatcher) HeartbeatOK() bool { return true } + +func (d *homeUnauthorizedRefreshDispatcher) RPopAuth(context.Context, string, string, http.Header, int) ([]byte, error) { + d.calls.Add(1) + return json.Marshal(homeAuthDispatchResponse{Auth: Auth{ + ID: "home-refresh-auth", + Provider: homeUnauthorizedRefreshProvider, + Status: StatusActive, + Attributes: map[string]string{ + AttributeAuthKind: AuthKindOAuth, + "websockets": "true", + }, + Metadata: map[string]any{ + "access_token": "stale-access-token", + }, + }}) +} + +func (*homeUnauthorizedRefreshDispatcher) AbortAmbiguousDispatch() {} + +type homeUnauthorizedRefreshExecutor struct { + streamMode string + refreshErr error + keepStale bool + retainSelection bool + executeCalls atomic.Int32 + countCalls atomic.Int32 + streamCalls atomic.Int32 + refreshCalls atomic.Int32 +} + +func (*homeUnauthorizedRefreshExecutor) Identifier() string { return homeUnauthorizedRefreshProvider } + +func (e *homeUnauthorizedRefreshExecutor) Execute(_ context.Context, auth *Auth, _ cliproxyexecutor.Request, opts cliproxyexecutor.Options) (cliproxyexecutor.Response, error) { + e.executeCalls.Add(1) + if e.retainSelection { + if lifecycle, ok := opts.ExecutionLifecycle.(interface{ Retain() }); ok { + lifecycle.Retain() + } + } + if authAccessToken(auth) == "stale-access-token" { + return cliproxyexecutor.Response{}, &Error{HTTPStatus: http.StatusUnauthorized, Message: "expired access token"} + } + return cliproxyexecutor.Response{Payload: []byte("ok")}, nil +} + +func (e *homeUnauthorizedRefreshExecutor) ExecuteStream(_ context.Context, auth *Auth, _ cliproxyexecutor.Request, _ cliproxyexecutor.Options) (*cliproxyexecutor.StreamResult, error) { + e.streamCalls.Add(1) + if authAccessToken(auth) == "stale-access-token" { + switch e.streamMode { + case "bootstrap": + chunks := make(chan cliproxyexecutor.StreamChunk, 1) + chunks <- cliproxyexecutor.StreamChunk{Err: &Error{HTTPStatus: http.StatusUnauthorized, Message: "expired access token"}} + close(chunks) + return &cliproxyexecutor.StreamResult{Chunks: chunks}, nil + case "started": + chunks := make(chan cliproxyexecutor.StreamChunk, 2) + chunks <- cliproxyexecutor.StreamChunk{Payload: []byte("started")} + chunks <- cliproxyexecutor.StreamChunk{Err: &Error{HTTPStatus: http.StatusUnauthorized, Message: "expired access token"}} + close(chunks) + return &cliproxyexecutor.StreamResult{Chunks: chunks}, nil + default: + return nil, &Error{HTTPStatus: http.StatusUnauthorized, Message: "expired access token"} + } + } + chunks := make(chan cliproxyexecutor.StreamChunk, 1) + chunks <- cliproxyexecutor.StreamChunk{Payload: []byte("ok")} + close(chunks) + return &cliproxyexecutor.StreamResult{Chunks: chunks}, nil +} + +func (e *homeUnauthorizedRefreshExecutor) Refresh(_ context.Context, auth *Auth) (*Auth, error) { + e.refreshCalls.Add(1) + if e.refreshErr != nil { + return nil, e.refreshErr + } + updated := auth.Clone() + if e.keepStale { + return updated, nil + } + if updated.Metadata == nil { + updated.Metadata = make(map[string]any) + } + updated.Metadata["access_token"] = "fresh-access-token" + return updated, nil +} + +func (e *homeUnauthorizedRefreshExecutor) CountTokens(_ context.Context, auth *Auth, _ cliproxyexecutor.Request, _ cliproxyexecutor.Options) (cliproxyexecutor.Response, error) { + e.countCalls.Add(1) + if authAccessToken(auth) == "stale-access-token" { + return cliproxyexecutor.Response{}, &Error{HTTPStatus: http.StatusUnauthorized, Message: "expired access token"} + } + return cliproxyexecutor.Response{Payload: []byte("ok")}, nil +} + +func (*homeUnauthorizedRefreshExecutor) HttpRequest(context.Context, *Auth, *http.Request) (*http.Response, error) { + return nil, nil +} + +func newHomeUnauthorizedRefreshManager(dispatcher *homeUnauthorizedRefreshDispatcher, executor *homeUnauthorizedRefreshExecutor) *Manager { + manager := NewManager(nil, nil, nil) + manager.SetConfig(&internalconfig.Config{Home: internalconfig.HomeConfig{Enabled: true}}) + manager.PublishHomeDispatch(dispatcher, executionregistry.New(), 1) + manager.RegisterExecutor(executor) + return manager +} + +func TestHomeUnauthorizedRefreshesSameSelectionBeforeRedispatch(t *testing.T) { + for _, test := range []struct { + name string + run func(*Manager) error + }{ + { + name: "execute", + run: func(manager *Manager) error { + _, errExecute := manager.Execute(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{}) + return errExecute + }, + }, + { + name: "count_tokens", + run: func(manager *Manager) error { + _, errCount := manager.ExecuteCount(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{}) + return errCount + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + if errRun := test.run(manager); errRun != nil { + t.Fatalf("execution error = %v", errRun) + } + if got := dispatcher.calls.Load(); got != 1 { + t.Fatalf("Home dispatch calls = %d, want 1", got) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want 1", got) + } + if test.name == "execute" && executor.executeCalls.Load() != 2 { + t.Fatalf("execute calls = %d, want 2", executor.executeCalls.Load()) + } + if test.name == "count_tokens" && executor.countCalls.Load() != 2 { + t.Fatalf("count calls = %d, want 2", executor.countCalls.Load()) + } + }) + } +} + +func TestHomeUnauthorizedRefreshUpdatesRetainedSelection(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{retainSelection: true} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + ctx := cliproxyexecutor.WithDownstreamWebsocket(context.Background()) + opts := cliproxyexecutor.Options{Metadata: map[string]any{ + cliproxyexecutor.ExecutionSessionMetadataKey: "refresh-session", + cliproxyexecutor.PinnedAuthMetadataKey: "home-refresh-auth", + }} + + for range 2 { + if _, errExecute := manager.Execute(ctx, []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, opts); errExecute != nil { + t.Fatalf("Execute() error = %v", errExecute) + } + } + if got := dispatcher.calls.Load(); got != 1 { + t.Fatalf("Home dispatch calls = %d, want one retained selection", got) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want refreshed token reused by retained selection", got) + } + if got := executor.executeCalls.Load(); got != 3 { + t.Fatalf("execute calls = %d, want stale attempt, retry, and retained reuse", got) + } +} + +func TestRefreshHomeSelectionReusesConcurrentNewerToken(t *testing.T) { + executor := &homeUnauthorizedRefreshExecutor{} + selection := &HomeDispatchSelection{ + Auth: &Auth{ID: "home-refresh-auth", Provider: homeUnauthorizedRefreshProvider, Attributes: map[string]string{AttributeAuthKind: AuthKindOAuth}, Metadata: map[string]any{"access_token": "fresh-access-token"}}, + Executor: executor, + Provider: homeUnauthorizedRefreshProvider, + } + failed := &Auth{ID: "home-refresh-auth", Provider: homeUnauthorizedRefreshProvider, Attributes: map[string]string{AttributeAuthKind: AuthKindOAuth}, Metadata: map[string]any{"access_token": "stale-access-token"}} + manager := NewManager(nil, nil, nil) + + updated, reused, errRefresh := manager.RefreshHomeSelectionAfterUnauthorized(context.Background(), selection, failed) + if errRefresh != nil || !reused || authAccessToken(updated) != "fresh-access-token" { + t.Fatalf("RefreshHomeSelectionAfterUnauthorized() = %#v, %v, %v", updated, reused, errRefresh) + } + if got := executor.refreshCalls.Load(); got != 0 { + t.Fatalf("refresh calls = %d, want 0 when selection already has a newer token", got) + } +} + +func TestHomeUnauthorizedRefreshIsAttemptedAtMostOnce(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{keepStale: true} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + _, errExecute := manager.Execute(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{}) + if statusCodeFromError(errExecute) != http.StatusUnauthorized { + t.Fatalf("Execute() error = %v, want original 401", errExecute) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want exactly 1", got) + } + if got := executor.executeCalls.Load(); got != 2 { + t.Fatalf("execute calls = %d, want initial attempt and one retry", got) + } +} + +func TestHomeNoCandidateAfterRefreshFailurePreservesRefreshError(t *testing.T) { + refreshErr := &Error{Code: "refresh_temporarily_unavailable", HTTPStatus: http.StatusServiceUnavailable, Message: "refresh unavailable"} + noCandidate := &Error{Code: "auth_not_found", HTTPStatus: http.StatusServiceUnavailable, Message: "no auth available"} + if !shouldReturnLastErrorOnPickFailure(true, refreshErr, noCandidate) { + t.Fatal("Home no-candidate error would overwrite the original refresh error") + } +} + +func TestHomeUnauthorizedTransientRefreshFailureIsReturned(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{ + refreshErr: &Error{HTTPStatus: http.StatusServiceUnavailable, Message: "Home refresh temporarily unavailable"}, + } + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + _, errExecute := manager.Execute(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{}) + if statusCodeFromError(errExecute) != http.StatusServiceUnavailable { + t.Fatalf("Execute() error = %v, want transient 503", errExecute) + } + if got := executor.executeCalls.Load(); got != 1 { + t.Fatalf("execute calls = %d, want 1", got) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want 1", got) + } +} + +func TestHomeUnauthorizedStreamRefreshesAtMostOnceAcrossRedispatch(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{keepStale: true} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + _, errStream := manager.ExecuteStream(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{Stream: true}) + if statusCodeFromError(errStream) != http.StatusUnauthorized { + t.Fatalf("ExecuteStream() error = %v, want original 401", errStream) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want exactly 1", got) + } + if got := executor.streamCalls.Load(); got != 2 { + t.Fatalf("stream calls = %d, want initial attempt and one retry", got) + } +} + +func TestHomeUnauthorizedStartedStreamDoesNotReplay(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{streamMode: "started"} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + result, errStream := manager.ExecuteStream(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{Stream: true}) + if errStream != nil { + t.Fatalf("ExecuteStream() error = %v", errStream) + } + sawPayload := false + sawUnauthorized := false + for chunk := range result.Chunks { + if string(chunk.Payload) == "started" { + sawPayload = true + } + if statusCodeFromError(chunk.Err) == http.StatusUnauthorized { + sawUnauthorized = true + } + } + if !sawPayload || !sawUnauthorized { + t.Fatalf("stream results = payload %v unauthorized %v, want both", sawPayload, sawUnauthorized) + } + if got := executor.refreshCalls.Load(); got != 0 { + t.Fatalf("refresh calls = %d, want 0 after stream started", got) + } + if got := executor.streamCalls.Load(); got != 1 { + t.Fatalf("stream calls = %d, want 1", got) + } +} + +func TestHomeUnauthorizedStreamRefreshesBeforeRedispatch(t *testing.T) { + for _, mode := range []string{"synchronous", "bootstrap"} { + t.Run(mode, func(t *testing.T) { + dispatcher := &homeUnauthorizedRefreshDispatcher{} + executor := &homeUnauthorizedRefreshExecutor{streamMode: mode} + manager := newHomeUnauthorizedRefreshManager(dispatcher, executor) + + result, errStream := manager.ExecuteStream(context.Background(), []string{homeUnauthorizedRefreshProvider}, cliproxyexecutor.Request{Model: "model-a"}, cliproxyexecutor.Options{Stream: true}) + if errStream != nil { + t.Fatalf("ExecuteStream() error = %v", errStream) + } + var payload string + for chunk := range result.Chunks { + if chunk.Err != nil { + t.Fatalf("stream chunk error = %v", chunk.Err) + } + payload += string(chunk.Payload) + } + if payload != "ok" { + t.Fatalf("stream payload = %q, want ok", payload) + } + if got := dispatcher.calls.Load(); got != 1 { + t.Fatalf("Home dispatch calls = %d, want 1", got) + } + if got := executor.refreshCalls.Load(); got != 1 { + t.Fatalf("refresh calls = %d, want 1", got) + } + if got := executor.streamCalls.Load(); got != 2 { + t.Fatalf("stream calls = %d, want 2", got) + } + }) + } +} diff --git a/sdk/cliproxy/usage/manager.go b/sdk/cliproxy/usage/manager.go index 7fa60416..ca36dc55 100644 --- a/sdk/cliproxy/usage/manager.go +++ b/sdk/cliproxy/usage/manager.go @@ -28,8 +28,10 @@ type Record struct { APIKey string AuthID string AuthIndex string - AuthType string - Source string + // AccessTokenSHA256 identifies the OAuth token version without exposing the token. + AccessTokenSHA256 string + AuthType string + Source string // ReasoningEffort stores the translated upstream thinking level for request event logs. ReasoningEffort string // ServiceTier stores the client-requested service tier. -- 2.51.2