mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 14:52:12 +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()
|
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]:
|
async def _request_github_tokens(body: dict[str, str]) -> dict[str, Any]:
|
||||||
if not GITHUB_APP_CLIENT_ID or not GITHUB_APP_CLIENT_SECRET:
|
if not GITHUB_APP_CLIENT_ID or not GITHUB_APP_CLIENT_SECRET:
|
||||||
raise HTTPException(500, "GitHub App OAuth not configured")
|
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):
|
if not isinstance(data, dict):
|
||||||
raise HTTPException(502, "unexpected GitHub OAuth response")
|
raise HTTPException(502, "unexpected GitHub OAuth response")
|
||||||
if data.get("error"):
|
if data.get("error"):
|
||||||
raise HTTPException(
|
raise GithubOAuthError(
|
||||||
400, f"github oauth error: {data.get('error_description') or data['error']}"
|
400,
|
||||||
|
f"github oauth error: {data.get('error_description') or data['error']}",
|
||||||
|
error_code=str(data["error"]),
|
||||||
)
|
)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -18,12 +18,15 @@ from datetime import UTC, datetime, timedelta
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException
|
|
||||||
from langgraph_sdk import get_client
|
from langgraph_sdk import get_client
|
||||||
from pydantic import BaseModel, field_validator
|
from pydantic import BaseModel, field_validator
|
||||||
|
|
||||||
from ..encryption import decrypt_token, encrypt_token
|
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
|
from .options import SUPPORTED_MODEL_IDS, model_supports_effort
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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:
|
def _decrypt_access_token(record: dict[str, Any]) -> str | None:
|
||||||
encrypted = record.get("encrypted_gh_token")
|
encrypted = record.get("encrypted_gh_token")
|
||||||
if not encrypted:
|
if not encrypted:
|
||||||
|
|
@ -220,21 +236,25 @@ def _decrypt_refresh_token(record: dict[str, Any]) -> str | None:
|
||||||
return decrypt_token(encrypted) or 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)
|
refresh_token = _decrypt_refresh_token(record)
|
||||||
if not refresh_token:
|
if not refresh_token:
|
||||||
return None
|
return None, False
|
||||||
try:
|
try:
|
||||||
data = await refresh_user_access_token(refresh_token)
|
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)
|
logger.warning("GitHub token refresh failed for %s", login, exc_info=True)
|
||||||
return None
|
return None, is_unrecoverable_refresh_error(exc)
|
||||||
except Exception:
|
|
||||||
logger.warning("GitHub token refresh failed for %s", login, exc_info=True)
|
|
||||||
return None
|
|
||||||
email = record.get("email") if isinstance(record.get("email"), str) else ""
|
email = record.get("email") if isinstance(record.get("email"), str) else ""
|
||||||
await upsert_access_token_from_github_response(login, email, data)
|
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:
|
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
|
return None
|
||||||
if not force_refresh and not _token_expired(record.get("token_expires_at")):
|
if not force_refresh and not _token_expired(record.get("token_expires_at")):
|
||||||
return access_token
|
return access_token
|
||||||
refreshed = await _refresh_stored_token(login, record)
|
refreshed, refresh_token_dead = await _refresh_stored_token(login, record)
|
||||||
return refreshed or access_token
|
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:
|
async def get_access_token(login: str) -> str | None:
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,11 @@ from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
import pytest
|
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
|
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()
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_get_valid_access_token_returns_stored_when_not_expiring() -> None:
|
async def test_get_valid_access_token_returns_stored_when_not_expiring() -> None:
|
||||||
future = (datetime.now(UTC) + timedelta(hours=5)).isoformat()
|
future = (datetime.now(UTC) + timedelta(hours=5)).isoformat()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue