mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-01 20:13:17 +00:00
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] <open-swe@users.noreply.github.com> Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
This commit is contained in:
parent
6a984d86c4
commit
85343fab63
14 changed files with 692 additions and 81 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
127
agent/webapp.py
127
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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
330
tests/test_github_token_ttl.py
Normal file
330
tests/test_github_token_ttl.py
Normal file
|
|
@ -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"
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue