diff --git a/agent/dashboard/oauth.py b/agent/dashboard/oauth.py index 835de659..d7b16f90 100644 --- a/agent/dashboard/oauth.py +++ b/agent/dashboard/oauth.py @@ -254,6 +254,28 @@ def expires_at_from_github_response(data: dict[str, Any], *, field: str) -> str return (datetime.now(UTC) + timedelta(seconds=int(raw))).isoformat() +class GithubOAuthError(HTTPException): + """A GitHub OAuth token endpoint error, carrying GitHub's ``error`` code.""" + + def __init__(self, status_code: int, detail: str, *, error_code: str | None = None) -> None: + super().__init__(status_code, detail) + self.error_code = error_code + + +# Error codes GitHub returns when a refresh token can never mint a new access +# token again (the user must re-authorize). Anything else is treated as +# transient so we don't needlessly drop a usable authorization. +UNRECOVERABLE_REFRESH_ERROR_CODES = frozenset({"bad_refresh_token", "unauthorized_client"}) + + +def is_unrecoverable_refresh_error(exc: BaseException) -> bool: + """Whether ``exc`` means the stored refresh token is permanently dead.""" + return ( + isinstance(exc, GithubOAuthError) + and (exc.error_code or "") in UNRECOVERABLE_REFRESH_ERROR_CODES + ) + + async def _request_github_tokens(body: dict[str, str]) -> dict[str, Any]: if not GITHUB_APP_CLIENT_ID or not GITHUB_APP_CLIENT_SECRET: raise HTTPException(500, "GitHub App OAuth not configured") @@ -268,8 +290,10 @@ async def _request_github_tokens(body: dict[str, str]) -> dict[str, Any]: if not isinstance(data, dict): raise HTTPException(502, "unexpected GitHub OAuth response") if data.get("error"): - raise HTTPException( - 400, f"github oauth error: {data.get('error_description') or data['error']}" + raise GithubOAuthError( + 400, + f"github oauth error: {data.get('error_description') or data['error']}", + error_code=str(data["error"]), ) return data diff --git a/agent/dashboard/profiles.py b/agent/dashboard/profiles.py index 59a81c2c..af006245 100644 --- a/agent/dashboard/profiles.py +++ b/agent/dashboard/profiles.py @@ -18,12 +18,15 @@ from datetime import UTC, datetime, timedelta from typing import Any import httpx -from fastapi import HTTPException from langgraph_sdk import get_client from pydantic import BaseModel, field_validator from ..encryption import decrypt_token, encrypt_token -from .oauth import expires_at_from_github_response, refresh_user_access_token +from .oauth import ( + expires_at_from_github_response, + is_unrecoverable_refresh_error, + refresh_user_access_token, +) from .options import SUPPORTED_MODEL_IDS, model_supports_effort logger = logging.getLogger(__name__) @@ -206,6 +209,19 @@ async def upsert_access_token_from_github_response( ) +async def delete_access_token(login: str) -> None: + """Drop the user's stored OAuth tokens. + + Used when a refresh token is permanently dead so we stop handing out a + known-stale access token and callers prompt a clean re-login instead. + """ + try: + await _client().store.delete_item(OAUTH_TOKENS_NAMESPACE, login) + except httpx.HTTPStatusError as exc: + if exc.response.status_code != 404: + raise + + def _decrypt_access_token(record: dict[str, Any]) -> str | None: encrypted = record.get("encrypted_gh_token") if not encrypted: @@ -220,21 +236,25 @@ def _decrypt_refresh_token(record: dict[str, Any]) -> str | None: return decrypt_token(encrypted) or None -async def _refresh_stored_token(login: str, record: dict[str, Any]) -> str | None: +async def _refresh_stored_token(login: str, record: dict[str, Any]) -> tuple[str | None, bool]: + """Refresh the stored token, returning ``(access_token, refresh_token_dead)``. + + ``refresh_token_dead`` is True when GitHub says the refresh token can never + mint a new token again, so the caller should drop the stored authorization + rather than keep serving a stale access token. + """ refresh_token = _decrypt_refresh_token(record) if not refresh_token: - return None + return None, False try: data = await refresh_user_access_token(refresh_token) - except HTTPException: + except Exception as exc: # noqa: BLE001 logger.warning("GitHub token refresh failed for %s", login, exc_info=True) - return None - except Exception: - logger.warning("GitHub token refresh failed for %s", login, exc_info=True) - return None + return None, is_unrecoverable_refresh_error(exc) email = record.get("email") if isinstance(record.get("email"), str) else "" await upsert_access_token_from_github_response(login, email, data) - return data.get("access_token") if isinstance(data.get("access_token"), str) else None + access_token = data.get("access_token") + return (access_token if isinstance(access_token, str) else None), False async def get_valid_access_token(login: str, *, force_refresh: bool = False) -> str | None: @@ -262,8 +282,25 @@ async def get_valid_access_token(login: str, *, force_refresh: bool = False) -> return None if not force_refresh and not _token_expired(record.get("token_expires_at")): return access_token - refreshed = await _refresh_stored_token(login, record) - return refreshed or access_token + refreshed, refresh_token_dead = await _refresh_stored_token(login, record) + if refreshed: + return refreshed + if refresh_token_dead: + # The refresh token is permanently invalid (revoked / expired), so + # the cached access token is dead too. Drop it so callers prompt a + # clean re-login instead of repeatedly handing out a stale token. + # The OAuth callback can write a fresh authorization while the + # refresh request is in flight (it doesn't take this lock), so only + # delete if the stored record is still the one that failed. + latest = await _get_value(OAUTH_TOKENS_NAMESPACE, login) + if latest and latest.get("encrypted_gh_refresh_token") != record.get( + "encrypted_gh_refresh_token" + ): + return _decrypt_access_token(latest) + logger.info("Dropping dead GitHub authorization for %s; re-login required", login) + await delete_access_token(login) + return None + return access_token async def get_access_token(login: str) -> str | None: diff --git a/tests/test_github_oauth_refresh.py b/tests/test_github_oauth_refresh.py index 4c58bc3e..17d466ef 100644 --- a/tests/test_github_oauth_refresh.py +++ b/tests/test_github_oauth_refresh.py @@ -5,7 +5,11 @@ from unittest.mock import AsyncMock, patch import pytest -from agent.dashboard.oauth import expires_at_from_github_response +from agent.dashboard.oauth import ( + GithubOAuthError, + expires_at_from_github_response, + is_unrecoverable_refresh_error, +) from agent.dashboard.profiles import _token_expired, get_valid_access_token @@ -62,6 +66,127 @@ async def test_get_valid_access_token_refreshes_when_near_expiry() -> None: mock_upsert.assert_awaited_once() +def test_is_unrecoverable_refresh_error() -> None: + assert is_unrecoverable_refresh_error( + GithubOAuthError(400, "x", error_code="bad_refresh_token") + ) + assert is_unrecoverable_refresh_error( + GithubOAuthError(400, "x", error_code="unauthorized_client") + ) + assert not is_unrecoverable_refresh_error(GithubOAuthError(400, "x", error_code="slow_down")) + assert not is_unrecoverable_refresh_error(GithubOAuthError(400, "x")) + assert not is_unrecoverable_refresh_error(RuntimeError("boom")) + + +@pytest.mark.asyncio +async def test_get_valid_access_token_drops_record_on_dead_refresh_token() -> None: + soon = (datetime.now(UTC) + timedelta(minutes=1)).isoformat() + record = { + "email": "u@example.com", + "encrypted_gh_token": "enc-access", + "encrypted_gh_refresh_token": "enc-refresh", + "token_expires_at": soon, + } + with ( + patch( + "agent.dashboard.profiles._get_value", + new_callable=AsyncMock, + return_value=record, + ), + patch("agent.dashboard.profiles._decrypt_access_token", return_value="stale-access"), + patch("agent.dashboard.profiles._decrypt_refresh_token", return_value="ghr_dead"), + patch( + "agent.dashboard.profiles.refresh_user_access_token", + new_callable=AsyncMock, + side_effect=GithubOAuthError( + 400, "github oauth error: bad refresh token", error_code="bad_refresh_token" + ), + ), + patch( + "agent.dashboard.profiles.delete_access_token", + new_callable=AsyncMock, + ) as mock_delete, + ): + token = await get_valid_access_token("octo") + assert token is None + mock_delete.assert_awaited_once_with("octo") + + +@pytest.mark.asyncio +async def test_get_valid_access_token_keeps_fresh_reauth_on_dead_refresh_token() -> None: + soon = (datetime.now(UTC) + timedelta(minutes=1)).isoformat() + stale = { + "email": "u@example.com", + "encrypted_gh_token": "enc-access", + "encrypted_gh_refresh_token": "enc-refresh-dead", + "token_expires_at": soon, + } + reauthed = { + "email": "u@example.com", + "encrypted_gh_token": "enc-access-new", + "encrypted_gh_refresh_token": "enc-refresh-new", + "token_expires_at": (datetime.now(UTC) + timedelta(hours=8)).isoformat(), + } + with ( + patch( + "agent.dashboard.profiles._get_value", + new_callable=AsyncMock, + side_effect=[stale, stale, reauthed], + ), + patch( + "agent.dashboard.profiles._decrypt_access_token", + side_effect=lambda r: "fresh-access" if r is reauthed else "stale-access", + ), + patch("agent.dashboard.profiles._decrypt_refresh_token", return_value="ghr_dead"), + patch( + "agent.dashboard.profiles.refresh_user_access_token", + new_callable=AsyncMock, + side_effect=GithubOAuthError( + 400, "github oauth error: bad refresh token", error_code="bad_refresh_token" + ), + ), + patch( + "agent.dashboard.profiles.delete_access_token", + new_callable=AsyncMock, + ) as mock_delete, + ): + token = await get_valid_access_token("octo") + assert token == "fresh-access" + mock_delete.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_valid_access_token_keeps_record_on_transient_refresh_failure() -> None: + soon = (datetime.now(UTC) + timedelta(minutes=1)).isoformat() + record = { + "email": "u@example.com", + "encrypted_gh_token": "enc-access", + "encrypted_gh_refresh_token": "enc-refresh", + "token_expires_at": soon, + } + with ( + patch( + "agent.dashboard.profiles._get_value", + new_callable=AsyncMock, + return_value=record, + ), + patch("agent.dashboard.profiles._decrypt_access_token", return_value="still-usable"), + patch("agent.dashboard.profiles._decrypt_refresh_token", return_value="ghr_ok"), + patch( + "agent.dashboard.profiles.refresh_user_access_token", + new_callable=AsyncMock, + side_effect=GithubOAuthError(503, "github oauth temporarily unavailable"), + ), + patch( + "agent.dashboard.profiles.delete_access_token", + new_callable=AsyncMock, + ) as mock_delete, + ): + token = await get_valid_access_token("octo") + assert token == "still-usable" + mock_delete.assert_not_awaited() + + @pytest.mark.asyncio async def test_get_valid_access_token_returns_stored_when_not_expiring() -> None: future = (datetime.now(UTC) + timedelta(hours=5)).isoformat()