From 85343fab6365cc10f04dd62631ae93e8b3af43cf Mon Sep 17 00:00:00 2001 From: "open-swe[bot]" <215916821+open-swe[bot]@users.noreply.github.com> Date: Fri, 8 May 2026 22:57:01 +0000 Subject: [PATCH] feat: TTL and revocation handling for cached GitHub OAuth tokens [closes AB-2322] (#1280) * feat: TTL and revocation handling for cached GitHub OAuth tokens [closes AB-2322] Persist github_token_expires_at alongside github_token_encrypted, treat expired cache entries as missing so we re-resolve before kicking off runs, and invalidate the cached ciphertext on a downstream 401 so the next invocation gets a fresh token instead of replaying a revoked one. Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> * webapp: forward installation-token expiry to reviewer cache writes The three reviewer-thread persist sites in webapp.py were calling get_github_app_installation_token() (no expiry) and persist_encrypted_github_token without expires_at, so cached App tokens were treated as never-expiring even though they actually expire in ~1 hour. --------- Co-authored-by: open-swe[bot] Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> Co-authored-by: Johannes du Plessis --- agent/reviewer.py | 8 +- agent/reviewer_publish.py | 7 + agent/server.py | 3 +- agent/tools/publish_review.py | 41 ++-- agent/utils/auth.py | 64 ++++-- agent/utils/github_app.py | 17 +- agent/utils/github_comments.py | 32 +++ agent/utils/github_token.py | 107 +++++++++- agent/webapp.py | 127 ++++++++--- tests/test_github_issue_webhook.py | 26 ++- tests/test_github_token_ttl.py | 330 +++++++++++++++++++++++++++++ tests/test_proxy_auth.py | 4 +- tests/test_reviewer.py | 2 +- tests/test_reviewer_watch.py | 5 + 14 files changed, 692 insertions(+), 81 deletions(-) create mode 100644 tests/test_github_token_ttl.py diff --git a/agent/reviewer.py b/agent/reviewer.py index eaf77c2e..6553cf2c 100644 --- a/agent/reviewer.py +++ b/agent/reviewer.py @@ -271,13 +271,17 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: return create_deep_agent(system_prompt="", tools=[]).with_config(config) if config["configurable"].get("source"): - cached_token, cached_encrypted = await get_github_token_from_thread(thread_id) + cached_token, cached_encrypted, cached_expires_at = await get_github_token_from_thread( + thread_id + ) if cached_token and cached_encrypted: config["metadata"]["github_token_encrypted"] = cached_encrypted + config["metadata"]["github_token_expires_at"] = cached_expires_at del cached_token else: - _token, new_encrypted = await resolve_github_token(config, thread_id) + _token, new_encrypted, new_expires_at = await resolve_github_token(config, thread_id) config["metadata"]["github_token_encrypted"] = new_encrypted + config["metadata"]["github_token_expires_at"] = new_expires_at del _token sandbox_backend = await ensure_sandbox_for_thread(thread_id) diff --git a/agent/reviewer_publish.py b/agent/reviewer_publish.py index f608efb1..9d419a77 100644 --- a/agent/reviewer_publish.py +++ b/agent/reviewer_publish.py @@ -26,6 +26,7 @@ from typing import Any import httpx from .reviewer_findings import Finding +from .utils.github_token import GitHubAuthError logger = logging.getLogger(__name__) @@ -121,7 +122,13 @@ async def post_pull_request_review( async with httpx.AsyncClient() as client: try: response = await client.post(url, headers=headers, json=payload, timeout=30) + if response.status_code == 401: + raise GitHubAuthError( + f"GitHub returned 401 posting PR review for {owner}/{repo}#{pr_number}" + ) response.raise_for_status() + except GitHubAuthError: + raise except httpx.HTTPStatusError as e: body = (e.response.text or "")[:500] logger.exception( diff --git a/agent/server.py b/agent/server.py index 06d66f97..740f4711 100644 --- a/agent/server.py +++ b/agent/server.py @@ -341,8 +341,9 @@ async def get_agent(config: RunnableConfig) -> Pregel: tools=[], ).with_config(config) - github_token, new_encrypted = await resolve_github_token(config, thread_id) + github_token, new_encrypted, new_expires_at = await resolve_github_token(config, thread_id) config["metadata"]["github_token_encrypted"] = new_encrypted + config["metadata"]["github_token_expires_at"] = new_expires_at triggering_user_identity = await asyncio.to_thread( resolve_triggering_user_identity, config, github_token ) diff --git a/agent/tools/publish_review.py b/agent/tools/publish_review.py index 5f32f00a..5cc5979b 100644 --- a/agent/tools/publish_review.py +++ b/agent/tools/publish_review.py @@ -27,7 +27,11 @@ from ..reviewer_publish import ( render_review_body, resolve_review_thread, ) -from ..utils.github_token import get_github_token +from ..utils.github_token import ( + GitHubAuthError, + get_github_token, + invalidate_cached_github_token, +) from ..utils.slack import post_slack_thread_reply @@ -90,18 +94,31 @@ def publish_review( if not token: return {"success": False, "error": "No GitHub token available"} - return asyncio.run( - _publish_review_async( - owner=str(repo_config["owner"]), - repo=str(repo_config["name"]), - pr_number=pr_number, - head_sha=head_sha, - token=token, - severity_threshold=_cast_severity(severity_threshold), - cap=cap, - is_re_review=is_re_review, + try: + return asyncio.run( + _publish_review_async( + owner=str(repo_config["owner"]), + repo=str(repo_config["name"]), + pr_number=pr_number, + head_sha=head_sha, + token=token, + severity_threshold=_cast_severity(severity_threshold), + cap=cap, + is_re_review=is_re_review, + ) ) - ) + except GitHubAuthError as exc: + thread_id = get_thread_id_from_runtime() + if thread_id: + asyncio.run(invalidate_cached_github_token(thread_id)) + return { + "success": False, + "error": ( + "GitHub returned 401 — the cached OAuth token is invalid or revoked. " + "Please re-authenticate and trigger the review again." + ), + "auth_error": str(exc), + } def _cast_severity(value: str) -> Severity: diff --git a/agent/utils/auth.py b/agent/utils/auth.py index 9f7e8047..eaad44de 100644 --- a/agent/utils/auth.py +++ b/agent/utils/auth.py @@ -14,7 +14,7 @@ from langgraph.graph.state import RunnableConfig from langgraph_sdk import get_client from ..encryption import encrypt_token -from .github_app import get_github_app_installation_token +from .github_app import get_github_app_installation_token_with_expiry from .github_token import get_github_token_from_thread from .github_user_email_map import GITHUB_USER_EMAIL_MAP from .linear import comment_on_linear_issue @@ -122,6 +122,19 @@ async def get_ls_user_id_from_email(email: str) -> dict[str, str | None]: return {"ls_user_id": None, "tenant_id": None} +def _extract_expires_at(response_data: dict[str, Any]) -> str | None: + """Pull an expiry from a LangSmith auth response in any of its known shapes.""" + expires_at = response_data.get("expires_at") or response_data.get("expiresAt") + if isinstance(expires_at, str) and expires_at: + return expires_at + if isinstance(expires_at, int | float): + return datetime.fromtimestamp(float(expires_at), tz=UTC).isoformat() + expires_in = response_data.get("expires_in") or response_data.get("expiresIn") + if isinstance(expires_in, int | float) and expires_in > 0: + return (datetime.now(UTC) + timedelta(seconds=int(expires_in))).isoformat() + return None + + async def get_github_token_for_user(ls_user_id: str, tenant_id: str) -> dict[str, Any]: """Get GitHub OAuth token for a user via LangSmith agent auth.""" if not GITHUB_OAUTH_PROVIDER_ID: @@ -159,7 +172,11 @@ async def get_github_token_for_user(ls_user_id: str, tenant_id: str) -> dict[str auth_url = response_data.get("url") if token: - return {"token": token} + result: dict[str, Any] = {"token": token} + expires_at = _extract_expires_at(response_data) + if expires_at: + result["expires_at"] = expires_at + return result if auth_url: return {"auth_url": auth_url} return {"error": f"Unexpected auth result: {response_data}"} @@ -265,20 +282,23 @@ async def leave_failure_comment( raise ValueError(f"Unknown source: {source}") -async def persist_encrypted_github_token(thread_id: str, token: str) -> str: - """Encrypt a GitHub token and store it on the thread metadata.""" +async def persist_encrypted_github_token( + thread_id: str, token: str, expires_at: str | None = None +) -> str: + """Encrypt a GitHub token and store it (and its expiry) on the thread metadata.""" encrypted = encrypt_token(token) - await client.threads.update( - thread_id=thread_id, - metadata={"github_token_encrypted": encrypted}, - ) + metadata: dict[str, Any] = { + "github_token_encrypted": encrypted, + "github_token_expires_at": expires_at, + } + await client.threads.update(thread_id=thread_id, metadata=metadata) return encrypted async def save_encrypted_token_from_email( email: str | None, source: str, -) -> tuple[str, str]: +) -> tuple[str, str, str | None]: """Resolve, encrypt, and store a GitHub token based on user email.""" config = get_config() configurable = config.get("configurable", {}) @@ -337,13 +357,14 @@ async def save_encrypted_token_from_email( await leave_failure_comment(source, message) raise ValueError(f"No token found: {error}") - encrypted = await persist_encrypted_github_token(thread_id, token) - return token, encrypted + expires_at = auth_result.get("expires_at") if isinstance(auth_result, dict) else None + encrypted = await persist_encrypted_github_token(thread_id, token, expires_at=expires_at) + return token, encrypted, expires_at -async def _resolve_bot_installation_token(thread_id: str) -> tuple[str, str]: +async def _resolve_bot_installation_token(thread_id: str) -> tuple[str, str, str | None]: """Get a GitHub App installation token and persist it for the thread.""" - bot_token = await get_github_app_installation_token() + bot_token, expires_at = await get_github_app_installation_token_with_expiry() if not bot_token: raise RuntimeError( "Bot-token-only mode is active (LANGSMITH_API_KEY_PROD set without " @@ -353,11 +374,13 @@ async def _resolve_bot_installation_token(thread_id: str) -> tuple[str, str]: logger.info( "Using GitHub App installation token for thread %s (bot-token-only mode)", thread_id ) - encrypted = await persist_encrypted_github_token(thread_id, bot_token) - return bot_token, encrypted + encrypted = await persist_encrypted_github_token(thread_id, bot_token, expires_at=expires_at) + return bot_token, encrypted, expires_at -async def resolve_github_token(config: RunnableConfig, thread_id: str) -> tuple[str, str]: +async def resolve_github_token( + config: RunnableConfig, thread_id: str +) -> tuple[str, str, str | None]: """Resolve a GitHub token from the run config based on the source. Routes to the correct auth method depending on whether the run was @@ -368,7 +391,8 @@ async def resolve_github_token(config: RunnableConfig, thread_id: str) -> tuple[ for all operations instead of per-user OAuth tokens. Returns: - (github_token, new_encrypted) tuple. + (github_token, new_encrypted, expires_at) tuple. ``expires_at`` is the + ISO-8601 expiry persisted alongside the ciphertext, or ``None``. Raises: RuntimeError: If source is missing or token resolution fails. @@ -384,9 +408,11 @@ async def resolve_github_token(config: RunnableConfig, thread_id: str) -> tuple[ try: if source == "github": - cached_token, cached_encrypted = await get_github_token_from_thread(thread_id) + cached_token, cached_encrypted, cached_expires_at = await get_github_token_from_thread( + thread_id + ) if cached_token and cached_encrypted: - return cached_token, cached_encrypted + return cached_token, cached_encrypted, cached_expires_at github_login = configurable.get("github_login") email = GITHUB_USER_EMAIL_MAP.get(github_login or "") if not email: diff --git a/agent/utils/github_app.py b/agent/utils/github_app.py index 6c92772c..ba708aa6 100644 --- a/agent/utils/github_app.py +++ b/agent/utils/github_app.py @@ -34,9 +34,19 @@ async def get_github_app_installation_token() -> str | None: Returns: Installation access token string, or None if unavailable. """ + token, _ = await get_github_app_installation_token_with_expiry() + return token + + +async def get_github_app_installation_token_with_expiry() -> tuple[str | None, str | None]: + """Exchange the GitHub App JWT for an installation access token and its expiry. + + Returns ``(token, expires_at)`` where ``expires_at`` is the ISO-8601 string + returned by GitHub (typically 1 hour out). Either value may be ``None``. + """ if not GITHUB_APP_ID or not GITHUB_APP_PRIVATE_KEY or not GITHUB_APP_INSTALLATION_ID: logger.debug("GitHub App env vars not fully configured, skipping app token") - return None + return None, None try: app_jwt = _generate_app_jwt() @@ -50,7 +60,8 @@ async def get_github_app_installation_token() -> str | None: }, ) response.raise_for_status() - return response.json().get("token") + data = response.json() + return data.get("token"), data.get("expires_at") except Exception: logger.exception("Failed to get GitHub App installation token") - return None + return None, None diff --git a/agent/utils/github_comments.py b/agent/utils/github_comments.py index 9966216b..425ee04e 100644 --- a/agent/utils/github_comments.py +++ b/agent/utils/github_comments.py @@ -11,10 +11,28 @@ from typing import Any import httpx +from .github_token import GitHubAuthError from .github_user_email_map import GITHUB_USER_EMAIL_MAP logger = logging.getLogger(__name__) +__all__ = [ + "GitHubAuthError", + "OPEN_SWE_TAGS", + "build_pr_prompt", + "extract_pr_context", + "fetch_issue_comments", + "fetch_pr_branch", + "fetch_pr_comments_since_last_tag", + "format_github_comment_body_for_prompt", + "get_thread_id_from_branch", + "parse_github_review_command", + "post_github_comment", + "react_to_github_comment", + "sanitize_github_comment_body", + "verify_github_signature", +] + OPEN_SWE_TAGS = ("@openswe", "@open-swe", "@openswe-dev") _OPEN_SWE_MENTION_RE = re.compile(r"(?i)@(?:openswe-dev|open-swe|openswe)\b") _REVIEW_COMMAND_RE = re.compile(r"(?i)\Areview(?:\s+(https?://\S+))?\s*\Z") @@ -134,8 +152,12 @@ async def react_to_github_comment( }, json={"content": "eyes"}, ) + if response.status_code == 401: + raise GitHubAuthError(f"GitHub returned 401 reacting to comment {comment_id}") # 200 = already reacted, 201 = just created return response.status_code in (200, 201) + except GitHubAuthError: + raise except Exception: logger.exception("Failed to react to GitHub comment %s", comment_id) return False @@ -161,11 +183,17 @@ async def _react_via_graphql(node_id: str | None, *, token: str) -> bool: headers={"Authorization": f"Bearer {token}"}, json={"query": query, "variables": {"subjectId": node_id}}, ) + if response.status_code == 401: + raise GitHubAuthError( + f"GitHub returned 401 reacting via GraphQL for node {node_id}" + ) data = response.json() if "errors" in data: logger.warning("GraphQL reaction errors: %s", data["errors"]) return False return True + except GitHubAuthError: + raise except Exception: logger.exception("Failed to react via GraphQL for node_id %s", node_id) return False @@ -458,6 +486,8 @@ async def _fetch_paginated( while True: try: response = await client.get(url, headers=headers, params=params) + if response.status_code == 401: + raise GitHubAuthError(f"GitHub returned 401 fetching {url}") if response.status_code != 200: # noqa: PLR2004 logger.warning("GitHub API returned %s for %s", response.status_code, url) break @@ -468,6 +498,8 @@ async def _fetch_paginated( if len(page_data) < 100: # noqa: PLR2004 break params["page"] += 1 + except GitHubAuthError: + raise except Exception: logger.exception("Failed to fetch %s", url) break diff --git a/agent/utils/github_token.py b/agent/utils/github_token.py index afe5217e..3957c3ed 100644 --- a/agent/utils/github_token.py +++ b/agent/utils/github_token.py @@ -4,6 +4,7 @@ from __future__ import annotations import logging from collections.abc import Mapping +from datetime import UTC, datetime from typing import Any from langgraph.config import get_config @@ -15,6 +16,16 @@ from ..encryption import decrypt_token logger = logging.getLogger(__name__) _GITHUB_TOKEN_METADATA_KEY = "github_token_encrypted" +_GITHUB_TOKEN_EXPIRES_AT_METADATA_KEY = "github_token_expires_at" + + +class GitHubAuthError(Exception): + """Raised when a GitHub call returns 401, signalling a stale/revoked token.""" + + +# Treat tokens with <= this many seconds remaining as expired so we re-auth +# before kicking off long agent runs. +_GITHUB_TOKEN_EXPIRY_SKEW_SECONDS = 60 client = get_client() @@ -31,34 +42,112 @@ def _decrypt_github_token(encrypted_token: str | None) -> str | None: return decrypt_token(encrypted_token) +def _is_expired(expires_at: Any, *, now: datetime | None = None) -> bool: + """Return True when ``expires_at`` is past (or close to) ``now``. + + Accepts ISO-8601 strings (with or without trailing Z) and unix timestamps. + Unparseable values are treated as not expired so we don't break callers + that haven't started persisting an expiry yet. + """ + if expires_at is None: + return False + + parsed: datetime | None = None + if isinstance(expires_at, int | float): + try: + parsed = datetime.fromtimestamp(float(expires_at), tz=UTC) + except (OverflowError, OSError, ValueError): + return False + elif isinstance(expires_at, str): + raw = expires_at.strip() + if not raw: + return False + if raw.endswith("Z"): + raw = raw[:-1] + "+00:00" + try: + parsed = datetime.fromisoformat(raw) + except ValueError: + return False + + if parsed is None: + return False + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=UTC) + + current = (now or datetime.now(UTC)).astimezone(UTC) + return (parsed - current).total_seconds() <= _GITHUB_TOKEN_EXPIRY_SKEW_SECONDS + + +def _read_token_if_fresh(metadata: dict[str, Any]) -> str | None: + """Decrypt the cached token only if it has not expired.""" + encrypted = _read_encrypted_github_token(metadata) + if not encrypted: + return None + if _is_expired(metadata.get(_GITHUB_TOKEN_EXPIRES_AT_METADATA_KEY)): + return None + return _decrypt_github_token(encrypted) + + def get_github_token(run_config: Mapping[str, Any] | None = None) -> str | None: """Resolve a GitHub token from run metadata. Pass ``run_config`` when LangGraph runnable config is already available (e.g. after ``get_config()`` in callers). Omit to read from ``get_config()`` (required runnable - context). + context). Returns ``None`` for tokens whose ``github_token_expires_at`` is past. """ resolved = run_config if run_config is not None else get_config() - return _decrypt_github_token(_read_encrypted_github_token(resolved.get("metadata", {}))) + return _read_token_if_fresh(resolved.get("metadata", {})) -async def get_github_token_from_thread(thread_id: str) -> tuple[str | None, str | None]: +async def get_github_token_from_thread( + thread_id: str, +) -> tuple[str | None, str | None, str | None]: """Resolve a GitHub token from LangGraph thread metadata. - Returns: - A `(token, encrypted_token)` tuple. Either value may be `None`. + Returns ``(None, None, None)`` when no token is cached or when the cached + token's ``github_token_expires_at`` has elapsed — callers must treat the + cache as missing in that case and re-resolve. On a fresh hit, returns the + decrypted token, its ciphertext, and the persisted expiry (or ``None``). """ try: thread = await client.threads.get(thread_id) except NotFoundError: logger.debug("Thread %s not found while looking up GitHub token", thread_id) - return None, None + return None, None, None except Exception: # noqa: BLE001 logger.exception("Failed to fetch thread metadata for %s", thread_id) - return None, None + return None, None, None + + metadata = (thread or {}).get("metadata", {}) + encrypted_token = _read_encrypted_github_token(metadata) + if not encrypted_token: + return None, None, None + expires_at_raw = metadata.get(_GITHUB_TOKEN_EXPIRES_AT_METADATA_KEY) + if _is_expired(expires_at_raw): + logger.info("Cached GitHub token for thread %s has expired; re-resolving", thread_id) + return None, None, None - encrypted_token = _read_encrypted_github_token((thread or {}).get("metadata", {})) token = _decrypt_github_token(encrypted_token) if token: logger.info("Found GitHub token in thread metadata for thread %s", thread_id) - return token, encrypted_token + expires_at = expires_at_raw if isinstance(expires_at_raw, str) else None + return token, encrypted_token, expires_at + + +async def invalidate_cached_github_token(thread_id: str) -> None: + """Clear a cached GitHub token from thread metadata. + + Called when a downstream GitHub API call returns 401, so the next run + re-resolves a fresh token instead of replaying the revoked ciphertext. + """ + try: + await client.threads.update( + thread_id=thread_id, + metadata={ + _GITHUB_TOKEN_METADATA_KEY: None, + _GITHUB_TOKEN_EXPIRES_AT_METADATA_KEY: None, + }, + ) + logger.info("Invalidated cached GitHub token for thread %s", thread_id) + except Exception: + logger.exception("Failed to invalidate cached GitHub token for thread %s", thread_id) diff --git a/agent/webapp.py b/agent/webapp.py index f51c6014..193ab9de 100644 --- a/agent/webapp.py +++ b/agent/webapp.py @@ -29,9 +29,13 @@ from .utils.auth import ( ) from .utils.authorship import OPEN_SWE_BOT_NAME from .utils.comments import get_recent_comments -from .utils.github_app import get_github_app_installation_token +from .utils.github_app import ( + get_github_app_installation_token, + get_github_app_installation_token_with_expiry, +) from .utils.github_comments import ( OPEN_SWE_TAGS, + GitHubAuthError, build_pr_prompt, extract_pr_context, fetch_issue_comments, @@ -44,7 +48,7 @@ from .utils.github_comments import ( verify_github_signature, ) from .utils.github_org_membership import INTERNAL_BOT_LOGINS, is_user_active_org_member -from .utils.github_token import get_github_token_from_thread +from .utils.github_token import get_github_token_from_thread, invalidate_cached_github_token from .utils.github_user_email_map import GITHUB_USER_EMAIL_MAP from .utils.linear import post_linear_trace_comment from .utils.linear_team_repo_map import LINEAR_TEAM_TO_REPO @@ -1514,7 +1518,7 @@ async def trigger_pr_review_from_ref( if not _is_repo_allowed_for_reviewer(repo_config): return {"success": False, "error": "Repository not allowed for reviewer"} - app_token = await get_github_app_installation_token() + app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() if not app_token: logger.warning("No GitHub App token available for PR reviewer request") return {"success": False, "error": "No GitHub App token available"} @@ -1540,7 +1544,7 @@ async def trigger_pr_review_from_ref( return {"success": False, "error": "Could not create reviewer thread"} try: - await persist_encrypted_github_token(thread_id, app_token) + await persist_encrypted_github_token(thread_id, app_token, expires_at=app_token_expires_at) except Exception: logger.warning("Could not persist bot token for reviewer thread %s", thread_id) return {"success": False, "error": "Could not persist reviewer token"} @@ -1663,7 +1667,7 @@ async def process_github_pr_review_request(payload: dict[str, Any]) -> None: repo_config.get("owner", ""), repo_config.get("name", ""), pr_number ) - app_token = await get_github_app_installation_token() + app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() if not app_token: logger.warning("No GitHub App token available for PR reviewer request") return @@ -1673,7 +1677,7 @@ async def process_github_pr_review_request(payload: dict[str, Any]) -> None: return try: - await persist_encrypted_github_token(thread_id, app_token) + await persist_encrypted_github_token(thread_id, app_token, expires_at=app_token_expires_at) except Exception: logger.warning("Could not persist bot token for reviewer thread %s", thread_id) return @@ -1896,7 +1900,7 @@ async def process_github_push_event(payload: dict[str, Any]) -> None: ) return - app_token = await get_github_app_installation_token() + app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() if not app_token: logger.warning("No GitHub App token for push re-review on %s", head_ref) return @@ -1951,7 +1955,7 @@ async def process_github_push_event(payload: dict[str, Any]) -> None: if not await _ensure_thread_exists_for_metadata(thread_id, langgraph_client): return try: - await persist_encrypted_github_token(thread_id, app_token) + await persist_encrypted_github_token(thread_id, app_token, expires_at=app_token_expires_at) except Exception: logger.warning("Could not persist bot token for reviewer thread %s", thread_id) return @@ -2002,24 +2006,35 @@ async def process_github_push_event(payload: dict[str, Any]) -> None: ) +async def _refresh_thread_github_token_after_401(thread_id: str, email: str) -> str | None: + """Invalidate the cached token after a 401 and try to resolve a fresh one.""" + logger.warning( + "GitHub returned 401 for thread %s; invalidating cached token and re-resolving", + thread_id, + ) + await invalidate_cached_github_token(thread_id) + return await _get_or_resolve_thread_github_token(thread_id, email) + + async def _get_or_resolve_thread_github_token(thread_id: str, email: str) -> str | None: """Resolve and persist a GitHub token for a thread when available. + Skips the cached ciphertext when its ``github_token_expires_at`` is past. In bot-token-only mode, returns a fresh GitHub App installation token instead of resolving per-user OAuth tokens. """ if is_bot_token_only_mode(): - bot_token = await get_github_app_installation_token() + bot_token, expires_at = await get_github_app_installation_token_with_expiry() if bot_token: try: - await persist_encrypted_github_token(thread_id, bot_token) + await persist_encrypted_github_token(thread_id, bot_token, expires_at=expires_at) except Exception: logger.warning("Could not persist bot token for thread %s", thread_id) return bot_token logger.warning("Bot-token-only mode but GitHub App token unavailable") return None - github_token, _encrypted_token = await get_github_token_from_thread(thread_id) + github_token, _encrypted_token, _expires_at = await get_github_token_from_thread(thread_id) if github_token: return github_token @@ -2029,7 +2044,9 @@ async def _get_or_resolve_thread_github_token(thread_id: str, email: str) -> str return None try: - await persist_encrypted_github_token(thread_id, github_token) + await persist_encrypted_github_token( + thread_id, github_token, expires_at=auth_result.get("expires_at") + ) except Exception: logger.warning("Could not persist GitHub token for thread %s", thread_id) return github_token @@ -2101,20 +2118,45 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) -> return if comment_id: - await react_to_github_comment( - repo_config, - comment_id, - event_type=event_type, - token=github_token, - pull_number=pr_number, - node_id=node_id, - ) + try: + await react_to_github_comment( + repo_config, + comment_id, + event_type=event_type, + token=github_token, + pull_number=pr_number, + node_id=node_id, + ) + except GitHubAuthError: + github_token = await _refresh_thread_github_token_after_401(thread_id, email) + if not github_token: + logger.warning("Re-auth failed for thread %s after 401; skipping", thread_id) + return + await react_to_github_comment( + repo_config, + comment_id, + event_type=event_type, + token=github_token, + pull_number=pr_number, + node_id=node_id, + ) if not pr_number: logger.warning("No PR number found in payload, skipping") return - comments = await fetch_pr_comments_since_last_tag(repo_config, pr_number, token=github_token) + try: + comments = await fetch_pr_comments_since_last_tag( + repo_config, pr_number, token=github_token + ) + except GitHubAuthError: + github_token = await _refresh_thread_github_token_after_401(thread_id, email) + if not github_token: + logger.warning("Re-auth failed for thread %s after 401; skipping", thread_id) + return + comments = await fetch_pr_comments_since_last_tag( + repo_config, pr_number, token=github_token + ) if not comments: logger.info("No comments found since last @open-swe tag for PR %s", pr_number) return @@ -2176,12 +2218,31 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None if not reaction_token: logger.warning("No GitHub token available to react to issue comment %s", comment_id) else: - reacted = await react_to_github_comment( - repo_config, - comment_id, - event_type="issue_comment", - token=reaction_token, - ) + try: + reacted = await react_to_github_comment( + repo_config, + comment_id, + event_type="issue_comment", + token=reaction_token, + ) + except GitHubAuthError: + github_token = await _refresh_thread_github_token_after_401(thread_id, email) + reaction_token = github_token or app_token + reacted = False + if reaction_token: + try: + reacted = await react_to_github_comment( + repo_config, + comment_id, + event_type="issue_comment", + token=reaction_token, + ) + except GitHubAuthError: + logger.warning( + "Re-auth still produced 401 reacting to issue comment %s", + comment_id, + ) + reacted = False if not reacted: logger.warning("Failed to react to GitHub issue comment %s", comment_id) @@ -2194,9 +2255,15 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None else: prompt = build_github_issue_update_prompt(github_login, title, description) else: - comments = await fetch_issue_comments( - repo_config, issue_number, token=github_token or app_token - ) + try: + comments = await fetch_issue_comments( + repo_config, issue_number, token=github_token or app_token + ) + except GitHubAuthError: + github_token = await _refresh_thread_github_token_after_401(thread_id, email) + comments = await fetch_issue_comments( + repo_config, issue_number, token=github_token or app_token + ) if comment_id and not any(item.get("comment_id") == comment_id for item in comments): comments.append( { diff --git a/tests/test_github_issue_webhook.py b/tests/test_github_issue_webhook.py index 48e40925..051bcbfb 100644 --- a/tests/test_github_issue_webhook.py +++ b/tests/test_github_issue_webhook.py @@ -651,9 +651,15 @@ def test_process_github_pr_review_request_creates_reviewer_run(monkeypatch) -> N async def fake_get_github_app_installation_token() -> str | None: return "app-token" - async def fake_persist_encrypted_github_token(thread_id: str, token: str) -> str: + async def fake_get_github_app_installation_token_with_expiry() -> tuple[str | None, str | None]: + return "app-token", None + + async def fake_persist_encrypted_github_token( + thread_id: str, token: str, *, expires_at: str | None = None + ) -> str: captured["persist_thread_id"] = thread_id captured["persist_token"] = token + captured["persist_expires_at"] = expires_at return "encrypted-token" async def fake_is_thread_active(thread_id: str) -> bool: @@ -680,6 +686,11 @@ def test_process_github_pr_review_request_creates_reviewer_run(monkeypatch) -> N monkeypatch.setattr( webapp, "get_github_app_installation_token", fake_get_github_app_installation_token ) + monkeypatch.setattr( + webapp, + "get_github_app_installation_token_with_expiry", + fake_get_github_app_installation_token_with_expiry, + ) monkeypatch.setattr( webapp, "persist_encrypted_github_token", fake_persist_encrypted_github_token ) @@ -730,6 +741,9 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None: async def fake_get_github_app_installation_token() -> str | None: return "app-token" + async def fake_get_github_app_installation_token_with_expiry() -> tuple[str | None, str | None]: + return "app-token", None + async def fake_fetch_github_pr_metadata( pr_ref: GitHubPrRef, *, token: str ) -> dict[str, object]: @@ -740,9 +754,12 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None: "head": {"sha": "head-sha", "ref": "feature-branch"}, } - async def fake_persist_encrypted_github_token(thread_id: str, token: str) -> str: + async def fake_persist_encrypted_github_token( + thread_id: str, token: str, *, expires_at: str | None = None + ) -> str: captured["persist_thread_id"] = thread_id captured["persist_token"] = token + captured["persist_expires_at"] = expires_at return "encrypted-token" async def fake_is_thread_active(thread_id: str) -> bool: @@ -769,6 +786,11 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None: monkeypatch.setattr( webapp, "get_github_app_installation_token", fake_get_github_app_installation_token ) + monkeypatch.setattr( + webapp, + "get_github_app_installation_token_with_expiry", + fake_get_github_app_installation_token_with_expiry, + ) monkeypatch.setattr(webapp, "fetch_github_pr_metadata", fake_fetch_github_pr_metadata) monkeypatch.setattr( webapp, "persist_encrypted_github_token", fake_persist_encrypted_github_token diff --git a/tests/test_github_token_ttl.py b/tests/test_github_token_ttl.py new file mode 100644 index 00000000..b2df9e79 --- /dev/null +++ b/tests/test_github_token_ttl.py @@ -0,0 +1,330 @@ +"""Tests for TTL + revocation handling on cached GitHub OAuth tokens. + +Covers: +- (a) expired-cache reads return None / fall through to re-auth +- (b) 401 on a downstream GitHub call invalidates the cached ciphertext and + triggers a fresh resolve in the webapp +- (c) ``publish_review`` invalidates the cached token and returns a clean + failure when GitHub responds 401 +""" + +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime, timedelta +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from agent.utils import github_comments, github_token + +_TEST_FERNET_KEY = "GMI8FNqVnhFzVfKDUTpGAUq8a2cm14kU0SyXzMTM4Yc=" + + +@pytest.fixture(autouse=True) +def _set_encryption_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("TOKEN_ENCRYPTION_KEY", _TEST_FERNET_KEY) + + +def _encrypted(token: str) -> str: + from agent.encryption import encrypt_token + + return encrypt_token(token) + + +# (a) expired-cache reads ----------------------------------------------------- + + +def test_is_expired_handles_iso_zulu_strings() -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat().replace("+00:00", "Z") + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat().replace("+00:00", "Z") + assert github_token._is_expired(past) is True + assert github_token._is_expired(future) is False + + +def test_is_expired_handles_unix_timestamps() -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).timestamp() + future = (datetime.now(UTC) + timedelta(hours=1)).timestamp() + assert github_token._is_expired(past) is True + assert github_token._is_expired(future) is False + + +def test_is_expired_treats_unparseable_as_not_expired() -> None: + assert github_token._is_expired(None) is False + assert github_token._is_expired("") is False + assert github_token._is_expired("not-a-date") is False + + +def test_get_github_token_returns_none_for_expired_run_metadata() -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + metadata = { + "github_token_encrypted": _encrypted("ghp_secret"), + "github_token_expires_at": past, + } + assert github_token.get_github_token({"metadata": metadata}) is None + + +def test_get_github_token_returns_decrypted_for_fresh_metadata() -> None: + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + metadata = { + "github_token_encrypted": _encrypted("ghp_secret"), + "github_token_expires_at": future, + } + assert github_token.get_github_token({"metadata": metadata}) == "ghp_secret" + + +def test_get_github_token_returns_decrypted_when_no_expires_at() -> None: + metadata = {"github_token_encrypted": _encrypted("ghp_secret")} + assert github_token.get_github_token({"metadata": metadata}) == "ghp_secret" + + +@pytest.mark.asyncio +async def test_get_github_token_from_thread_skips_expired() -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + fake_client = AsyncMock() + fake_client.threads.get.return_value = { + "metadata": { + "github_token_encrypted": _encrypted("ghp_revoked"), + "github_token_expires_at": past, + } + } + with patch.object(github_token, "client", fake_client): + token, encrypted, expires_at = await github_token.get_github_token_from_thread("tid") + assert token is None + assert encrypted is None + assert expires_at is None + + +@pytest.mark.asyncio +async def test_get_github_token_from_thread_returns_fresh() -> None: + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + enc = _encrypted("ghp_live") + fake_client = AsyncMock() + fake_client.threads.get.return_value = { + "metadata": { + "github_token_encrypted": enc, + "github_token_expires_at": future, + } + } + with patch.object(github_token, "client", fake_client): + token, encrypted, expires_at = await github_token.get_github_token_from_thread("tid") + assert token == "ghp_live" + assert encrypted == enc + assert expires_at == future + + +@pytest.mark.asyncio +async def test_invalidate_cached_github_token_clears_metadata() -> None: + fake_client = AsyncMock() + with patch.object(github_token, "client", fake_client): + await github_token.invalidate_cached_github_token("tid-42") + fake_client.threads.update.assert_awaited_once_with( + thread_id="tid-42", + metadata={"github_token_encrypted": None, "github_token_expires_at": None}, + ) + + +# (b) 401 on a downstream GitHub call ----------------------------------------- + + +class _MockResponse: + def __init__(self, status_code: int, json_data: Any | None = None) -> None: + self.status_code = status_code + self._json = json_data or {} + + def json(self) -> Any: + return self._json + + +class _MockHttpxClient: + def __init__(self, status_code: int, json_data: Any | None = None) -> None: + self.status_code = status_code + self.json_data = json_data + self.posts: list[dict[str, Any]] = [] + self.gets: list[dict[str, Any]] = [] + + async def __aenter__(self) -> _MockHttpxClient: + return self + + async def __aexit__(self, *args: Any) -> None: + return None + + async def post(self, url: str, **kwargs: Any) -> _MockResponse: + self.posts.append({"url": url, **kwargs}) + return _MockResponse(self.status_code, self.json_data) + + async def get(self, url: str, **kwargs: Any) -> _MockResponse: + self.gets.append({"url": url, **kwargs}) + return _MockResponse(self.status_code, self.json_data) + + +def test_react_to_github_comment_raises_on_401(monkeypatch: pytest.MonkeyPatch) -> None: + mock_client = _MockHttpxClient(status_code=401) + monkeypatch.setattr(httpx, "AsyncClient", lambda *a, **kw: mock_client) + + async def _run() -> None: + await github_comments.react_to_github_comment( + {"owner": "o", "name": "r"}, + comment_id=1, + event_type="issue_comment", + token="revoked", + ) + + with pytest.raises(github_token.GitHubAuthError): + asyncio.run(_run()) + + +def test_fetch_pr_comments_since_last_tag_raises_on_401( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mock_client = _MockHttpxClient(status_code=401) + monkeypatch.setattr(httpx, "AsyncClient", lambda *a, **kw: mock_client) + + async def _run() -> None: + await github_comments.fetch_pr_comments_since_last_tag( + {"owner": "o", "name": "r"}, + pr_number=42, + token="revoked", + ) + + with pytest.raises(github_token.GitHubAuthError): + asyncio.run(_run()) + + +def test_fetch_issue_comments_raises_on_401(monkeypatch: pytest.MonkeyPatch) -> None: + mock_client = _MockHttpxClient(status_code=401) + monkeypatch.setattr(httpx, "AsyncClient", lambda *a, **kw: mock_client) + + async def _run() -> None: + await github_comments.fetch_issue_comments( + {"owner": "o", "name": "r"}, + issue_number=42, + token="revoked", + ) + + with pytest.raises(github_token.GitHubAuthError): + asyncio.run(_run()) + + +# (c) successful re-auth following stale-cache invalidation ------------------- + + +def test_process_github_pr_comment_invalidates_and_reauths_on_401( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """End-to-end check: a 401 on react triggers invalidate + re-resolve.""" + from agent import webapp + + invalidated: dict[str, int] = {"calls": 0} + resolves: list[str] = [] + react_calls: list[str] = [] + fetch_calls: list[str] = [] + + async def fake_invalidate(thread_id: str) -> None: + invalidated["calls"] += 1 + + tokens = iter(["stale-token", "fresh-token"]) + + async def fake_get_or_resolve(thread_id: str, email: str) -> str | None: + token = next(tokens) + resolves.append(token) + return token + + async def fake_react( + repo_config: dict[str, str], + comment_id: int, + *, + event_type: str, + token: str, + pull_number: int | None = None, + node_id: str | None = None, + ) -> bool: + react_calls.append(token) + if token == "stale-token": + raise github_comments.GitHubAuthError("revoked") + return True + + async def fake_fetch_pr_comments( + repo_config: dict[str, str], pr_number: int, *, token: str + ) -> list[dict[str, Any]]: + fetch_calls.append(token) + return [ + {"body": "@openswe please look", "author": "octo", "created_at": "2026-01-01T00:00:00Z"} + ] + + async def fake_extract_pr_context( + payload: dict[str, Any], event_type: str + ) -> tuple[dict[str, str], int, str, str, str, int, str | None]: + return ( + {"owner": "o", "name": "r"}, + 7, + "open-swe/00000000-0000-0000-0000-000000000001", + "octo", + "https://github.com/o/r/pull/7", + 42, + None, + ) + + async def fake_trigger_or_queue_run(*args: Any, **kwargs: Any) -> None: + return None + + monkeypatch.setattr(webapp, "extract_pr_context", fake_extract_pr_context) + monkeypatch.setattr(webapp, "_get_or_resolve_thread_github_token", fake_get_or_resolve) + monkeypatch.setattr(webapp, "invalidate_cached_github_token", fake_invalidate) + monkeypatch.setattr(webapp, "react_to_github_comment", fake_react) + monkeypatch.setattr(webapp, "fetch_pr_comments_since_last_tag", fake_fetch_pr_comments) + monkeypatch.setattr(webapp, "_trigger_or_queue_run", fake_trigger_or_queue_run) + monkeypatch.setattr(webapp, "GITHUB_USER_EMAIL_MAP", {"octo": "octo@example.com"}) + + asyncio.run( + webapp.process_github_pr_comment( + {"sender": {"login": "octo", "id": 1}}, + "issue_comment", + ) + ) + + assert invalidated["calls"] == 1 + assert resolves == ["stale-token", "fresh-token"] + assert react_calls == ["stale-token", "fresh-token"] + assert fetch_calls == ["fresh-token"] + + +def test_publish_review_invalidates_cached_token_on_401( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import importlib + + publish_review_module = importlib.import_module("agent.tools.publish_review") + + invalidated: dict[str, int] = {"calls": 0} + + async def fake_invalidate(thread_id: str) -> None: + invalidated["calls"] += 1 + invalidated["thread_id"] = thread_id # type: ignore[assignment] + + async def fake_publish(*args: Any, **kwargs: Any) -> dict[str, Any]: + raise github_token.GitHubAuthError("401 from PR review") + + monkeypatch.setattr( + publish_review_module, + "get_config", + lambda: { + "configurable": { + "repo": {"owner": "o", "name": "r"}, + "pr_number": 7, + "head_sha": "deadbeef", + }, + }, + ) + monkeypatch.setattr(publish_review_module, "get_github_token", lambda: "revoked-token") + monkeypatch.setattr(publish_review_module, "invalidate_cached_github_token", fake_invalidate) + monkeypatch.setattr(publish_review_module, "_publish_review_async", fake_publish) + monkeypatch.setattr(publish_review_module, "get_thread_id_from_runtime", lambda: "thread-xyz") + + result = publish_review_module.publish_review() + assert result["success"] is False + assert "401" in result["error"] + assert invalidated["calls"] == 1 + assert invalidated.get("thread_id") == "thread-xyz" diff --git a/tests/test_proxy_auth.py b/tests/test_proxy_auth.py index abf423ac..bf9b9d71 100644 --- a/tests/test_proxy_auth.py +++ b/tests/test_proxy_auth.py @@ -210,7 +210,7 @@ class TestRefreshProxyOnSandboxReuse: patch( "agent.server.resolve_github_token", new_callable=AsyncMock, - return_value=("ghp", "enc"), + return_value=("ghp", "enc", None), ), patch( "agent.server.get_sandbox_id_from_metadata", @@ -259,7 +259,7 @@ class TestRefreshProxyOnSandboxReuse: patch( "agent.server.resolve_github_token", new_callable=AsyncMock, - return_value=("ghp", "enc"), + return_value=("ghp", "enc", None), ), patch( "agent.server.get_sandbox_id_from_metadata", diff --git a/tests/test_reviewer.py b/tests/test_reviewer.py index b713c213..fffb558f 100644 --- a/tests/test_reviewer.py +++ b/tests/test_reviewer.py @@ -31,7 +31,7 @@ async def test_reviewer_uses_cached_thread_token_for_slack_review_request() -> N patch( "agent.reviewer.get_github_token_from_thread", new_callable=AsyncMock, - return_value=("app-token", "encrypted-token"), + return_value=("app-token", "encrypted-token", None), ) as mock_get_thread_token, patch("agent.reviewer.resolve_github_token", new_callable=AsyncMock) as mock_resolve_token, patch( diff --git a/tests/test_reviewer_watch.py b/tests/test_reviewer_watch.py index 50aff82e..192a2597 100644 --- a/tests/test_reviewer_watch.py +++ b/tests/test_reviewer_watch.py @@ -94,6 +94,11 @@ async def test_push_event_triggers_re_review_run_when_watching() -> None: new_callable=AsyncMock, return_value="t", ), + patch( + "agent.webapp.get_github_app_installation_token_with_expiry", + new_callable=AsyncMock, + return_value=("t", None), + ), patch( "agent.webapp._fetch_open_pr_for_branch", new_callable=AsyncMock,