mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 09:13:14 +00:00
fix: auto-recover from expired GitHub refresh tokens (#1491)
* fix: auto-recover from expired GitHub refresh tokens When a user's GitHub OAuth refresh token was permanently dead (revoked or expired), token refresh failed but get_valid_access_token still handed back the known-stale access token, so dashboard GitHub calls kept 401ing until the user manually logged out and back in. Now we distinguish unrecoverable refresh failures (bad_refresh_token / unauthorized_client) from transient ones: on an unrecoverable failure we drop the dead stored authorization and return None, so callers prompt a clean re-login. Transient failures still fall back to the stored token. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com> * fix: don't delete fresh re-auth when stale refresh fails --------- Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
f08177aed7
commit
9a2f68e99d
3 changed files with 201 additions and 15 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue