fix: scope public reviewer tokens (#1389)

* fix: scope public reviewer tokens

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

* refactor: simplify reviewer token wiring; fix push re-scope + red test

- Remove the redundant _check_or_recreate_sandbox_for_proxy /
  _refresh_github_proxy_or_recreate_for_proxy wrappers and call the
  underlying functions directly (they already default the token to None).
- process_github_push_event: re-scope the GitHub App token when the push
  payload lacked repo privacy/id but PR metadata reveals a public repo, so
  reviewer.py never proxies a full-installation token for a public PR.
- Clarify the two-token sequence in trigger_pr_review_from_ref.
- Fix pre-existing failing test test_proxy_refresh_failure_recreates_sandbox
  and add coverage for _reviewer_token_for_repo + push-event scoping.

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Johannes du Plessis 2026-06-03 09:25:17 -07:00 • committed by GitHub
parent 9e3406bd29
commit 0a2e682364
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 448 additions and 50 deletions

View file

@ -637,11 +637,16 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel:
config["metadata"]["github_token_expires_at"] = new_expires_at config["metadata"]["github_token_expires_at"] = new_expires_at
github_token = _token github_token = _token
sandbox_backend = await ensure_sandbox_for_thread(thread_id) repo_config = config["configurable"].get("repo") or {}
repo_private = config["configurable"].get("repo_private")
github_proxy_token = github_token if repo_private is False else None
sandbox_backend = await ensure_sandbox_for_thread(
thread_id,
github_proxy_token=github_proxy_token,
)
work_dir = await aresolve_sandbox_work_dir(sandbox_backend) work_dir = await aresolve_sandbox_work_dir(sandbox_backend)
repo_config = config["configurable"].get("repo") or {}
repo_owner = str(repo_config.get("owner", "")) repo_owner = str(repo_config.get("owner", ""))
repo_name = str(repo_config.get("name", "")) repo_name = str(repo_config.get("name", ""))
base_sha = str(config["configurable"].get("base_sha", "") or "") base_sha = str(config["configurable"].get("base_sha", "") or "")

View file

@ -119,36 +119,35 @@ async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProt
await asyncio.to_thread(sandbox.start) await asyncio.to_thread(sandbox.start)
async def _create_sandbox_with_proxy() -> SandboxBackendProtocol: async def _create_sandbox_with_proxy(
"""Create a new sandbox with GitHub proxy auth configured. github_proxy_token: str | None = None,
) -> SandboxBackendProtocol:
Uses create_sandbox (generic factory) so non-langsmith providers still work. """Create a new sandbox with GitHub proxy auth configured."""
For langsmith sandboxes, configures the proxy with the installation token.
"""
sandbox_backend = await asyncio.to_thread(create_sandbox) sandbox_backend = await asyncio.to_thread(create_sandbox)
sandbox_type = os.getenv("SANDBOX_TYPE", "langsmith") sandbox_type = os.getenv("SANDBOX_TYPE", "langsmith")
if sandbox_type == "langsmith": if sandbox_type == "langsmith":
installation_token = await get_github_app_installation_token() token = github_proxy_token or await get_github_app_installation_token()
if not installation_token: if not token:
msg = "Cannot configure proxy: GitHub App installation token is unavailable" msg = "Cannot configure proxy: GitHub App installation token is unavailable"
logger.error(msg) logger.error(msg)
raise ValueError(msg) raise ValueError(msg)
await _start_langsmith_sandbox_if_needed(sandbox_backend) await _start_langsmith_sandbox_if_needed(sandbox_backend)
await asyncio.to_thread(_configure_github_proxy, sandbox_backend.id, installation_token) await asyncio.to_thread(_configure_github_proxy, sandbox_backend.id, token)
return sandbox_backend return sandbox_backend
async def _refresh_github_proxy( async def _refresh_github_proxy(
sandbox_backend: SandboxBackendProtocol, sandbox_backend: SandboxBackendProtocol,
github_proxy_token: str | None = None,
) -> None: ) -> None:
"""Refresh GitHub proxy credentials for reused LangSmith sandboxes.""" """Refresh GitHub proxy credentials for reused LangSmith sandboxes."""
if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith": if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith":
return return
installation_token = await get_github_app_installation_token() token = github_proxy_token or await get_github_app_installation_token()
if not installation_token: if not token:
logger.warning( logger.warning(
"Skipping GitHub proxy refresh for sandbox %s: installation token unavailable", "Skipping GitHub proxy refresh for sandbox %s: installation token unavailable",
sandbox_backend.id, sandbox_backend.id,
@ -157,16 +156,17 @@ async def _refresh_github_proxy(
current_backend = unwrap_sandbox_backend(sandbox_backend) current_backend = unwrap_sandbox_backend(sandbox_backend)
await _start_langsmith_sandbox_if_needed(current_backend) await _start_langsmith_sandbox_if_needed(current_backend)
await asyncio.to_thread(_configure_github_proxy, current_backend.id, installation_token) await asyncio.to_thread(_configure_github_proxy, current_backend.id, token)
async def _refresh_github_proxy_or_recreate( async def _refresh_github_proxy_or_recreate(
sandbox_backend: SandboxBackendProtocol, sandbox_backend: SandboxBackendProtocol,
thread_id: str, thread_id: str,
github_proxy_token: str | None = None,
) -> SandboxBackendProtocol: ) -> SandboxBackendProtocol:
"""Refresh proxy credentials, recreating stale LangSmith sandboxes on failure.""" """Refresh proxy credentials, recreating stale LangSmith sandboxes on failure."""
try: try:
await _refresh_github_proxy(sandbox_backend) await _refresh_github_proxy(sandbox_backend, github_proxy_token)
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.warning( logger.warning(
"Failed to refresh GitHub proxy for sandbox %s on thread %s, recreating sandbox", "Failed to refresh GitHub proxy for sandbox %s on thread %s, recreating sandbox",
@ -174,7 +174,7 @@ async def _refresh_github_proxy_or_recreate(
thread_id, thread_id,
exc_info=True, exc_info=True,
) )
return await _recreate_sandbox(thread_id) return await _recreate_sandbox(thread_id, github_proxy_token=github_proxy_token)
return sandbox_backend return sandbox_backend
@ -186,7 +186,11 @@ async def _configure_git_identity(sandbox_backend: SandboxBackendProtocol) -> No
) )
async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol: async def _recreate_sandbox(
thread_id: str,
*,
github_proxy_token: str | None = None,
) -> SandboxBackendProtocol:
"""Recreate a sandbox after a connection failure. """Recreate a sandbox after a connection failure.
Sets the SANDBOX_CREATING sentinel and creates a fresh sandbox Sets the SANDBOX_CREATING sentinel and creates a fresh sandbox
@ -195,7 +199,10 @@ async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol:
""" """
await client.threads.update(thread_id=thread_id, metadata=_creating_metadata()) await client.threads.update(thread_id=thread_id, metadata=_creating_metadata())
try: try:
sandbox_backend = set_sandbox_backend(thread_id, await _create_sandbox_with_proxy()) sandbox_backend = set_sandbox_backend(
thread_id,
await _create_sandbox_with_proxy(github_proxy_token),
)
except Exception: except Exception:
logger.exception("Failed to recreate sandbox after connection failure") logger.exception("Failed to recreate sandbox after connection failure")
await client.threads.update(thread_id=thread_id, metadata=_RESET_METADATA) await client.threads.update(thread_id=thread_id, metadata=_RESET_METADATA)
@ -204,7 +211,9 @@ async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol:
async def check_or_recreate_sandbox( async def check_or_recreate_sandbox(
sandbox_backend: SandboxBackendProtocol, thread_id: str sandbox_backend: SandboxBackendProtocol,
thread_id: str,
github_proxy_token: str | None = None,
) -> SandboxBackendProtocol: ) -> SandboxBackendProtocol:
"""Check if a cached sandbox is reachable; recreate it if not. """Check if a cached sandbox is reachable; recreate it if not.
@ -221,7 +230,7 @@ async def check_or_recreate_sandbox(
"Cached sandbox is no longer reachable for thread %s, recreating", "Cached sandbox is no longer reachable for thread %s, recreating",
thread_id, thread_id,
) )
sandbox_backend = await _recreate_sandbox(thread_id) sandbox_backend = await _recreate_sandbox(thread_id, github_proxy_token=github_proxy_token)
return sandbox_backend return sandbox_backend
@ -272,7 +281,11 @@ def graph_loaded_for_execution(config: RunnableConfig) -> bool:
) )
async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol: async def ensure_sandbox_for_thread(
thread_id: str,
*,
github_proxy_token: str | None = None,
) -> SandboxBackendProtocol:
"""Get-or-create a healthy sandbox bound to ``thread_id``. """Get-or-create a healthy sandbox bound to ``thread_id``.
Implements the four-state lifecycle described in AGENTS.md: Implements the four-state lifecycle described in AGENTS.md:
@ -297,14 +310,18 @@ async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol:
if sandbox_backend: if sandbox_backend:
logger.info("Using cached sandbox backend for thread %s", thread_id) logger.info("Using cached sandbox backend for thread %s", thread_id)
original_sandbox_id = sandbox_backend.id original_sandbox_id = sandbox_backend.id
sandbox_backend = await check_or_recreate_sandbox(sandbox_backend, thread_id) sandbox_backend = await check_or_recreate_sandbox(
sandbox_backend, thread_id, github_proxy_token
)
if sandbox_backend.id == original_sandbox_id: if sandbox_backend.id == original_sandbox_id:
sandbox_backend = await _refresh_github_proxy_or_recreate(sandbox_backend, thread_id) sandbox_backend = await _refresh_github_proxy_or_recreate(
sandbox_backend, thread_id, github_proxy_token
)
elif sandbox_id is None: elif sandbox_id is None:
logger.info("Creating new sandbox for thread %s", thread_id) logger.info("Creating new sandbox for thread %s", thread_id)
await client.threads.update(thread_id=thread_id, metadata=_creating_metadata()) await client.threads.update(thread_id=thread_id, metadata=_creating_metadata())
try: try:
sandbox_backend = await _create_sandbox_with_proxy() sandbox_backend = await _create_sandbox_with_proxy(github_proxy_token)
logger.info("Sandbox created: %s", sandbox_backend.id) logger.info("Sandbox created: %s", sandbox_backend.id)
except Exception: except Exception:
logger.exception("Failed to create sandbox") logger.exception("Failed to create sandbox")
@ -322,7 +339,7 @@ async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol:
logger.warning("Failed to connect to existing sandbox %s, creating new one", sandbox_id) logger.warning("Failed to connect to existing sandbox %s, creating new one", sandbox_id)
await client.threads.update(thread_id=thread_id, metadata=_creating_metadata()) await client.threads.update(thread_id=thread_id, metadata=_creating_metadata())
try: try:
sandbox_backend = await _create_sandbox_with_proxy() sandbox_backend = await _create_sandbox_with_proxy(github_proxy_token)
created_replacement_sandbox = True created_replacement_sandbox = True
except Exception: except Exception:
logger.exception("Failed to create replacement sandbox") logger.exception("Failed to create replacement sandbox")
@ -330,10 +347,12 @@ async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol:
raise raise
if not created_replacement_sandbox: if not created_replacement_sandbox:
original_sandbox_id = sandbox_backend.id original_sandbox_id = sandbox_backend.id
sandbox_backend = await check_or_recreate_sandbox(sandbox_backend, thread_id) sandbox_backend = await check_or_recreate_sandbox(
sandbox_backend, thread_id, github_proxy_token
)
if sandbox_backend.id == original_sandbox_id: if sandbox_backend.id == original_sandbox_id:
sandbox_backend = await _refresh_github_proxy_or_recreate( sandbox_backend = await _refresh_github_proxy_or_recreate(
sandbox_backend, thread_id sandbox_backend, thread_id, github_proxy_token
) )
sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend) sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend)

View file

@ -5,6 +5,8 @@ from __future__ import annotations
import logging import logging
import os import os
import time import time
from collections.abc import Sequence
from typing import Any
import httpx import httpx
import jwt import jwt
@ -28,26 +30,35 @@ def _generate_app_jwt() -> str:
return jwt.encode(payload, private_key, algorithm="RS256") return jwt.encode(payload, private_key, algorithm="RS256")
async def get_github_app_installation_token() -> str | None: async def get_github_app_installation_token(
"""Exchange the GitHub App JWT for an installation access token. *,
repository_ids: Sequence[int] | None = None,
Returns: repositories: Sequence[str] | None = None,
Installation access token string, or None if unavailable. ) -> str | None:
""" """Exchange the GitHub App JWT for an installation access token."""
token, _ = await get_github_app_installation_token_with_expiry() token, _ = await get_github_app_installation_token_with_expiry(
repository_ids=repository_ids,
repositories=repositories,
)
return token return token
async def get_github_app_installation_token_with_expiry() -> tuple[str | None, str | None]: async def get_github_app_installation_token_with_expiry(
"""Exchange the GitHub App JWT for an installation access token and its expiry. *,
repository_ids: Sequence[int] | None = None,
Returns ``(token, expires_at)`` where ``expires_at`` is the ISO-8601 string repositories: Sequence[str] | None = None,
returned by GitHub (typically 1 hour out). Either value may be ``None``. ) -> tuple[str | None, str | None]:
""" """Exchange the GitHub App JWT for an installation access token and its expiry."""
if not GITHUB_APP_ID or not GITHUB_APP_PRIVATE_KEY or not GITHUB_APP_INSTALLATION_ID: 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") logger.debug("GitHub App env vars not fully configured, skipping app token")
return None, None return None, None
body: dict[str, Any] = {}
if repository_ids:
body["repository_ids"] = list(repository_ids)
elif repositories:
body["repositories"] = list(repositories)
try: try:
app_jwt = _generate_app_jwt() app_jwt = _generate_app_jwt()
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
@ -58,6 +69,7 @@ async def get_github_app_installation_token_with_expiry() -> tuple[str | None, s
"Accept": "application/vnd.github+json", "Accept": "application/vnd.github+json",
"X-GitHub-Api-Version": "2022-11-28", "X-GitHub-Api-Version": "2022-11-28",
}, },
json=body or None,
) )
response.raise_for_status() response.raise_for_status()
data = response.json() data = response.json()

View file

@ -1668,6 +1668,46 @@ async def fetch_github_pr_metadata(pr_ref: GitHubPrRef, *, token: str) -> dict[s
return data if isinstance(data, dict) else None return data if isinstance(data, dict) else None
def _repo_private_from_pr_metadata(pr_metadata: dict[str, Any]) -> bool | None:
repo = pr_metadata.get("base", {}).get("repo")
if isinstance(repo, dict) and isinstance(repo.get("private"), bool):
return repo["private"]
return None
def _repo_id_from_pr_metadata(pr_metadata: dict[str, Any]) -> int | None:
repo = pr_metadata.get("base", {}).get("repo")
repo_id = repo.get("id") if isinstance(repo, dict) else None
return repo_id if isinstance(repo_id, int) else None
def _repo_private_from_payload(payload: dict[str, Any]) -> bool | None:
repo = payload.get("repository")
private = repo.get("private") if isinstance(repo, dict) else None
return private if isinstance(private, bool) else None
def _repo_id_from_payload(payload: dict[str, Any]) -> int | None:
repo = payload.get("repository")
repo_id = repo.get("id") if isinstance(repo, dict) else None
return repo_id if isinstance(repo_id, int) else None
async def _reviewer_token_for_repo(
repo_config: dict[str, str],
*,
repo_private: bool | None,
repo_id: int | None = None,
) -> tuple[str | None, str | None]:
if repo_private is False:
if repo_id is not None:
return await get_github_app_installation_token_with_expiry(repository_ids=[repo_id])
repo_name = repo_config.get("name")
if repo_name:
return await get_github_app_installation_token_with_expiry(repositories=[repo_name])
return await get_github_app_installation_token_with_expiry()
async def trigger_pr_review_from_ref( async def trigger_pr_review_from_ref(
pr_ref: GitHubPrRef, pr_ref: GitHubPrRef,
*, *,
@ -1681,6 +1721,8 @@ async def trigger_pr_review_from_ref(
if not await _is_repo_enabled_for_review(repo_config): if not await _is_repo_enabled_for_review(repo_config):
return {"success": False, "error": "Repository not enabled for review"} return {"success": False, "error": "Repository not enabled for review"}
# Full token to read PR metadata (privacy/id aren't in the trigger ref);
# re-scoped below once we know whether the repo is public.
app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry()
if not app_token: if not app_token:
logger.warning("No GitHub App token available for PR reviewer request") logger.warning("No GitHub App token available for PR reviewer request")
@ -1690,6 +1732,17 @@ async def trigger_pr_review_from_ref(
if not pr_metadata: if not pr_metadata:
return {"success": False, "error": "Could not fetch pull request metadata"} return {"success": False, "error": "Could not fetch pull request metadata"}
repo_private = _repo_private_from_pr_metadata(pr_metadata)
repo_id = _repo_id_from_pr_metadata(pr_metadata)
app_token, app_token_expires_at = await _reviewer_token_for_repo(
repo_config,
repo_private=repo_private,
repo_id=repo_id,
)
if not app_token:
logger.warning("No GitHub App token available for PR reviewer request")
return {"success": False, "error": "No GitHub App token available"}
base_sha = pr_metadata.get("base", {}).get("sha", "") base_sha = pr_metadata.get("base", {}).get("sha", "")
head = pr_metadata.get("head", {}) head = pr_metadata.get("head", {})
head_sha = head.get("sha", "") head_sha = head.get("sha", "")
@ -1742,6 +1795,7 @@ async def trigger_pr_review_from_ref(
base_sha=base_sha, base_sha=base_sha,
head_sha=head_sha, head_sha=head_sha,
branch_name=branch_name, branch_name=branch_name,
repo_private=repo_private,
slack_channel_id=slack_channel_id, slack_channel_id=slack_channel_id,
slack_thread_ts=slack_thread_ts, slack_thread_ts=slack_thread_ts,
) )
@ -1781,6 +1835,7 @@ def _build_reviewer_configurable(
base_sha: str, base_sha: str,
head_sha: str, head_sha: str,
branch_name: str, branch_name: str,
repo_private: bool | None = None,
re_review: bool = False, re_review: bool = False,
last_reviewed_sha: str = "", last_reviewed_sha: str = "",
slack_channel_id: str = "", slack_channel_id: str = "",
@ -1801,6 +1856,8 @@ def _build_reviewer_configurable(
} }
if branch_name: if branch_name:
configurable["branch_name"] = branch_name configurable["branch_name"] = branch_name
if repo_private is not None:
configurable["repo_private"] = repo_private
if last_reviewed_sha: if last_reviewed_sha:
configurable["last_reviewed_sha"] = last_reviewed_sha configurable["last_reviewed_sha"] = last_reviewed_sha
if slack_channel_id and slack_thread_ts: if slack_channel_id and slack_thread_ts:
@ -1836,6 +1893,8 @@ async def _dispatch_first_review_from_pr_payload(payload: dict[str, Any], *, sou
"owner": repo.get("owner", {}).get("login", ""), "owner": repo.get("owner", {}).get("login", ""),
"name": repo.get("name", ""), "name": repo.get("name", ""),
} }
repo_private = _repo_private_from_payload(payload)
repo_id = _repo_id_from_payload(payload)
pr_number = pull_request.get("number") pr_number = pull_request.get("number")
pr_url = pull_request.get("html_url", "") or pull_request.get("url", "") pr_url = pull_request.get("html_url", "") or pull_request.get("url", "")
branch_name = pull_request.get("head", {}).get("ref", "") branch_name = pull_request.get("head", {}).get("ref", "")
@ -1881,7 +1940,11 @@ async def _dispatch_first_review_from_pr_payload(payload: dict[str, Any], *, sou
return return
last_reviewed_sha = existing_last_reviewed_sha last_reviewed_sha = existing_last_reviewed_sha
app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() app_token, app_token_expires_at = await _reviewer_token_for_repo(
repo_config,
repo_private=repo_private,
repo_id=repo_id,
)
if not app_token: if not app_token:
logger.warning("No GitHub App token available for reviewer dispatch") logger.warning("No GitHub App token available for reviewer dispatch")
return return
@ -1917,6 +1980,7 @@ async def _dispatch_first_review_from_pr_payload(payload: dict[str, Any], *, sou
base_sha=base_sha, base_sha=base_sha,
head_sha=head_sha, head_sha=head_sha,
branch_name=branch_name, branch_name=branch_name,
repo_private=repo_private,
re_review=is_re_review, re_review=is_re_review,
last_reviewed_sha=last_reviewed_sha, last_reviewed_sha=last_reviewed_sha,
) )
@ -2204,6 +2268,8 @@ async def process_github_push_event(payload: dict[str, Any]) -> None:
"owner": repo.get("owner", {}).get("login", "") or repo.get("owner", {}).get("name", ""), "owner": repo.get("owner", {}).get("login", "") or repo.get("owner", {}).get("name", ""),
"name": repo.get("name", ""), "name": repo.get("name", ""),
} }
repo_private = _repo_private_from_payload(payload)
repo_id = _repo_id_from_payload(payload)
if not repo_config["owner"] or not repo_config["name"]: if not repo_config["owner"] or not repo_config["name"]:
logger.warning("Push to %s ignored: repository owner/name missing from payload", head_ref) logger.warning("Push to %s ignored: repository owner/name missing from payload", head_ref)
return return
@ -2216,7 +2282,11 @@ async def process_github_push_event(payload: dict[str, Any]) -> None:
) )
return return
app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() app_token, app_token_expires_at = await _reviewer_token_for_repo(
repo_config,
repo_private=repo_private,
repo_id=repo_id,
)
if not app_token: if not app_token:
logger.warning("No GitHub App token for push re-review on %s", head_ref) logger.warning("No GitHub App token for push re-review on %s", head_ref)
return return
@ -2231,6 +2301,21 @@ async def process_github_push_event(payload: dict[str, Any]) -> None:
) )
return return
# Push payloads normally carry repo privacy/id; fall back to PR metadata.
# If the repo turns out public, re-scope the token so reviewer.py doesn't
# proxy a full-installation token for a public PR.
if repo_private is None:
repo_private = _repo_private_from_pr_metadata(pr)
repo_id = repo_id or _repo_id_from_pr_metadata(pr)
if repo_private is False:
app_token, app_token_expires_at = await _reviewer_token_for_repo(
repo_config,
repo_private=repo_private,
repo_id=repo_id,
)
if not app_token:
logger.warning("No GitHub App token for push re-review on %s", head_ref)
return
pr_number = pr.get("number") pr_number = pr.get("number")
pr_url = pr.get("html_url") or pr.get("url") or "" pr_url = pr.get("html_url") or pr.get("url") or ""
base_sha = pr.get("base", {}).get("sha", "") base_sha = pr.get("base", {}).get("sha", "")
@ -2332,6 +2417,7 @@ async def process_github_push_event(payload: dict[str, Any]) -> None:
base_sha=base_sha, base_sha=base_sha,
head_sha=head_sha, head_sha=head_sha,
branch_name=head_ref, branch_name=head_ref,
repo_private=repo_private,
re_review=True, re_review=True,
last_reviewed_sha=last_reviewed_sha if isinstance(last_reviewed_sha, str) else "", last_reviewed_sha=last_reviewed_sha if isinstance(last_reviewed_sha, str) else "",
) )
@ -2599,6 +2685,8 @@ async def process_github_review_finding_reply(payload: dict[str, Any]) -> None:
"owner": repo.get("owner", {}).get("login", ""), "owner": repo.get("owner", {}).get("login", ""),
"name": repo.get("name", ""), "name": repo.get("name", ""),
} }
repo_private = _repo_private_from_payload(payload)
repo_id = _repo_id_from_payload(payload)
pr_number = pull_request.get("number") pr_number = pull_request.get("number")
if not isinstance(pr_number, int): if not isinstance(pr_number, int):
return return
@ -2610,7 +2698,11 @@ async def process_github_review_finding_reply(payload: dict[str, Any]) -> None:
if metadata is None or metadata.get("kind") != REVIEWER_THREAD_KIND: if metadata is None or metadata.get("kind") != REVIEWER_THREAD_KIND:
return return
app_token, app_token_expires_at = await get_github_app_installation_token_with_expiry() app_token, app_token_expires_at = await _reviewer_token_for_repo(
repo_config,
repo_private=repo_private,
repo_id=repo_id,
)
if not app_token: if not app_token:
return return
try: try:
@ -2669,6 +2761,7 @@ async def process_github_review_finding_reply(payload: dict[str, Any]) -> None:
base_sha=base_sha, base_sha=base_sha,
head_sha=head_sha, head_sha=head_sha,
branch_name=branch_name, branch_name=branch_name,
repo_private=repo_private,
re_review=True, re_review=True,
) )
configurable.update( configurable.update(

65
tests/test_github_app.py Normal file
View file

@ -0,0 +1,65 @@
from __future__ import annotations
from typing import Any
import pytest
from agent.utils import github_app
class _FakeResponse:
def raise_for_status(self) -> None:
pass
def json(self) -> dict[str, str]:
return {"token": "token", "expires_at": "expires"}
class _FakeAsyncClient:
last_post: dict[str, Any] | None = None
async def __aenter__(self) -> _FakeAsyncClient:
return self
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
return None
async def post(self, url: str, **kwargs: Any) -> _FakeResponse:
type(self).last_post = {"url": url, **kwargs}
return _FakeResponse()
@pytest.mark.asyncio
async def test_installation_token_can_be_scoped_to_repository_ids(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(github_app, "GITHUB_APP_ID", "1")
monkeypatch.setattr(github_app, "GITHUB_APP_PRIVATE_KEY", "key")
monkeypatch.setattr(github_app, "GITHUB_APP_INSTALLATION_ID", "2")
monkeypatch.setattr(github_app, "_generate_app_jwt", lambda: "jwt")
monkeypatch.setattr(github_app.httpx, "AsyncClient", _FakeAsyncClient)
token, expires_at = await github_app.get_github_app_installation_token_with_expiry(
repository_ids=[123]
)
assert token == "token"
assert expires_at == "expires"
assert _FakeAsyncClient.last_post is not None
assert _FakeAsyncClient.last_post["json"] == {"repository_ids": [123]}
@pytest.mark.asyncio
async def test_installation_token_omits_scope_for_full_installation(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(github_app, "GITHUB_APP_ID", "1")
monkeypatch.setattr(github_app, "GITHUB_APP_PRIVATE_KEY", "key")
monkeypatch.setattr(github_app, "GITHUB_APP_INSTALLATION_ID", "2")
monkeypatch.setattr(github_app, "_generate_app_jwt", lambda: "jwt")
monkeypatch.setattr(github_app.httpx, "AsyncClient", _FakeAsyncClient)
await github_app.get_github_app_installation_token_with_expiry()
assert _FakeAsyncClient.last_post is not None
assert _FakeAsyncClient.last_post["json"] is None

View file

@ -10,10 +10,19 @@ import pytest
from agent import webapp from agent import webapp
def _pr_payload(*, action: str, draft: bool, author: str = "alice") -> dict[str, Any]: def _pr_payload(
*,
action: str,
draft: bool,
author: str = "alice",
private: bool | None = None,
) -> dict[str, Any]:
repository: dict[str, Any] = {"owner": {"login": "lc"}, "name": "repo", "id": 123}
if private is not None:
repository["private"] = private
return { return {
"action": action, "action": action,
"repository": {"owner": {"login": "lc"}, "name": "repo"}, "repository": repository,
"pull_request": { "pull_request": {
"number": 7, "number": 7,
"html_url": "https://github.com/lc/repo/pull/7", "html_url": "https://github.com/lc/repo/pull/7",
@ -56,6 +65,55 @@ async def test_pr_ready_non_draft_triggers_run(monkeypatch: pytest.MonkeyPatch)
assert kwargs["config"]["configurable"]["pr_number"] == 7 assert kwargs["config"]["configurable"]["pr_number"] == 7
@pytest.mark.asyncio
async def test_pr_ready_public_repo_uses_scoped_reviewer_token(
monkeypatch: pytest.MonkeyPatch,
) -> None:
fake_client = MagicMock()
fake_client.runs.create = AsyncMock()
get_token = AsyncMock(return_value=("scoped-token", "expires"))
monkeypatch.setattr(webapp, "get_github_app_installation_token_with_expiry", get_token)
monkeypatch.setattr(webapp, "_ensure_thread_exists_for_metadata", AsyncMock(return_value=True))
persist_token = AsyncMock(return_value="enc")
monkeypatch.setattr(webapp, "persist_encrypted_github_token", persist_token)
monkeypatch.setattr(webapp, "set_reviewer_thread_metadata", AsyncMock())
monkeypatch.setattr(webapp, "is_thread_active", AsyncMock(return_value=False))
monkeypatch.setattr(webapp, "get_client", lambda url: fake_client)
monkeypatch.setattr(webapp, "get_profile", AsyncMock(return_value=None))
monkeypatch.setattr(webapp, "get_team_settings", AsyncMock(return_value={}))
await webapp.process_github_pr_ready(_pr_payload(action="opened", draft=False, private=False))
get_token.assert_awaited_once_with(repository_ids=[123])
persist_token.assert_awaited_once()
assert persist_token.await_args.args[1] == "scoped-token"
_, kwargs = fake_client.runs.create.await_args
assert kwargs["config"]["configurable"]["repo_private"] is False
@pytest.mark.asyncio
async def test_pr_ready_private_repo_uses_full_reviewer_token(
monkeypatch: pytest.MonkeyPatch,
) -> None:
fake_client = MagicMock()
fake_client.runs.create = AsyncMock()
get_token = AsyncMock(return_value=("full-token", "expires"))
monkeypatch.setattr(webapp, "get_github_app_installation_token_with_expiry", get_token)
monkeypatch.setattr(webapp, "_ensure_thread_exists_for_metadata", AsyncMock(return_value=True))
monkeypatch.setattr(webapp, "persist_encrypted_github_token", AsyncMock(return_value="enc"))
monkeypatch.setattr(webapp, "set_reviewer_thread_metadata", AsyncMock())
monkeypatch.setattr(webapp, "is_thread_active", AsyncMock(return_value=False))
monkeypatch.setattr(webapp, "get_client", lambda url: fake_client)
monkeypatch.setattr(webapp, "get_profile", AsyncMock(return_value=None))
monkeypatch.setattr(webapp, "get_team_settings", AsyncMock(return_value={}))
await webapp.process_github_pr_ready(_pr_payload(action="opened", draft=False, private=True))
get_token.assert_awaited_once_with()
_, kwargs = fake_client.runs.create.await_args
assert kwargs["config"]["configurable"]["repo_private"] is True
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pr_ready_for_review_triggers_run(monkeypatch: pytest.MonkeyPatch) -> None: async def test_pr_ready_for_review_triggers_run(monkeypatch: pytest.MonkeyPatch) -> None:
fake_client = MagicMock() fake_client = MagicMock()

View file

@ -327,7 +327,7 @@ class TestRefreshProxyOnSandboxReuse:
assert sandbox is replacement_sandbox assert sandbox is replacement_sandbox
mock_proxy.assert_called_once_with("sandbox-stale", "ghs_fresh") mock_proxy.assert_called_once_with("sandbox-stale", "ghs_fresh")
mock_recreate.assert_awaited_once_with("thread-123") mock_recreate.assert_awaited_once_with("thread-123", github_proxy_token=None)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_starts_stopped_langsmith_sandbox_before_proxy_refresh(self) -> None: async def test_starts_stopped_langsmith_sandbox_before_proxy_refresh(self) -> None:

View file

@ -3,18 +3,31 @@
from __future__ import annotations from __future__ import annotations
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest import pytest
from agent import webapp from agent import webapp
def _push_payload(*, ref: str, after: str, owner: str = "lc", name: str = "repo") -> dict[str, Any]: def _push_payload(
*,
ref: str,
after: str,
owner: str = "lc",
name: str = "repo",
private: bool | None = None,
repo_id: int | None = None,
) -> dict[str, Any]:
repository: dict[str, Any] = {"owner": {"login": owner}, "name": name}
if private is not None:
repository["private"] = private
if repo_id is not None:
repository["id"] = repo_id
return { return {
"ref": ref, "ref": ref,
"after": after, "after": after,
"repository": {"owner": {"login": owner}, "name": name}, "repository": repository,
"sender": {"login": "alice", "id": 7}, "sender": {"login": "alice", "id": 7},
} }
@ -309,6 +322,139 @@ async def test_push_event_idempotent_when_head_unchanged() -> None:
fake_client.runs.create.assert_not_called() fake_client.runs.create.assert_not_called()
@pytest.mark.asyncio
async def test_reviewer_token_for_repo_public_scopes_by_id() -> None:
get_token = AsyncMock(return_value=("scoped", "exp"))
with patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token):
token, expires = await webapp._reviewer_token_for_repo(
{"owner": "lc", "name": "repo"}, repo_private=False, repo_id=123
)
assert (token, expires) == ("scoped", "exp")
get_token.assert_awaited_once_with(repository_ids=[123])
@pytest.mark.asyncio
async def test_reviewer_token_for_repo_public_scopes_by_name_without_id() -> None:
get_token = AsyncMock(return_value=("scoped", "exp"))
with patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token):
await webapp._reviewer_token_for_repo(
{"owner": "lc", "name": "repo"}, repo_private=False, repo_id=None
)
get_token.assert_awaited_once_with(repositories=["repo"])
@pytest.mark.asyncio
async def test_reviewer_token_for_repo_private_uses_full_token() -> None:
get_token = AsyncMock(return_value=("full", "exp"))
with patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token):
await webapp._reviewer_token_for_repo(
{"owner": "lc", "name": "repo"}, repo_private=True, repo_id=123
)
get_token.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_reviewer_token_for_repo_unknown_privacy_uses_full_token() -> None:
get_token = AsyncMock(return_value=("full", "exp"))
with patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token):
await webapp._reviewer_token_for_repo(
{"owner": "lc", "name": "repo"}, repo_private=None, repo_id=123
)
get_token.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_push_event_public_repo_uses_scoped_token() -> None:
payload = _push_payload(ref="refs/heads/feat-x", after="newsha", private=False, repo_id=123)
pr = {
"number": 7,
"html_url": "https://github.com/lc/repo/pull/7",
"title": "T",
"head": {"sha": "newsha", "ref": "feat-x"},
"base": {"sha": "basesha", "ref": "main"},
}
fake_client = MagicMock()
fake_client.runs.create = AsyncMock()
get_token = AsyncMock(return_value=("scoped-token", "exp"))
persist = AsyncMock(return_value="enc")
with (
patch(
"agent.webapp._is_repo_enabled_for_review", new_callable=AsyncMock, return_value=True
),
patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token),
patch("agent.webapp._fetch_open_pr_for_branch", new_callable=AsyncMock, return_value=pr),
patch(
"agent.webapp._get_thread_metadata_safe",
new_callable=AsyncMock,
return_value={"kind": "reviewer", "watch": True},
),
patch(
"agent.webapp._ensure_thread_exists_for_metadata",
new_callable=AsyncMock,
return_value=True,
),
patch("agent.webapp.persist_encrypted_github_token", persist),
patch("agent.webapp.fetch_pr_review_threads", new_callable=AsyncMock, return_value=[]),
patch("agent.webapp.reconcile_findings_with_review_threads", new_callable=AsyncMock),
patch("agent.webapp.set_reviewer_thread_metadata", new_callable=AsyncMock),
patch("agent.webapp.is_thread_active", new_callable=AsyncMock, return_value=False),
patch("agent.webapp.get_client", return_value=fake_client),
):
await webapp.process_github_push_event(payload)
get_token.assert_awaited_once_with(repository_ids=[123])
assert persist.await_args.args[1] == "scoped-token"
_, kwargs = fake_client.runs.create.await_args
assert kwargs["config"]["configurable"]["repo_private"] is False
@pytest.mark.asyncio
async def test_push_event_rescopes_token_when_pr_metadata_reveals_public() -> None:
payload = _push_payload(ref="refs/heads/feat-x", after="newsha")
pr = {
"number": 7,
"html_url": "https://github.com/lc/repo/pull/7",
"title": "T",
"head": {"sha": "newsha", "ref": "feat-x"},
"base": {"sha": "basesha", "ref": "main", "repo": {"private": False, "id": 456}},
}
fake_client = MagicMock()
fake_client.runs.create = AsyncMock()
get_token = AsyncMock(side_effect=[("full-token", "e1"), ("scoped-token", "e2")])
persist = AsyncMock(return_value="enc")
with (
patch(
"agent.webapp._is_repo_enabled_for_review", new_callable=AsyncMock, return_value=True
),
patch("agent.webapp.get_github_app_installation_token_with_expiry", get_token),
patch("agent.webapp._fetch_open_pr_for_branch", new_callable=AsyncMock, return_value=pr),
patch(
"agent.webapp._get_thread_metadata_safe",
new_callable=AsyncMock,
return_value={"kind": "reviewer", "watch": True},
),
patch(
"agent.webapp._ensure_thread_exists_for_metadata",
new_callable=AsyncMock,
return_value=True,
),
patch("agent.webapp.persist_encrypted_github_token", persist),
patch("agent.webapp.fetch_pr_review_threads", new_callable=AsyncMock, return_value=[]),
patch("agent.webapp.reconcile_findings_with_review_threads", new_callable=AsyncMock),
patch("agent.webapp.set_reviewer_thread_metadata", new_callable=AsyncMock),
patch("agent.webapp.is_thread_active", new_callable=AsyncMock, return_value=False),
patch("agent.webapp.get_client", return_value=fake_client),
):
await webapp.process_github_push_event(payload)
assert get_token.await_args_list == [call(), call(repository_ids=[456])]
assert persist.await_args.args[1] == "scoped-token"
_, kwargs = fake_client.runs.create.await_args
assert kwargs["config"]["configurable"]["repo_private"] is False
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pr_close_disables_watch() -> None: async def test_pr_close_disables_watch() -> None:
captured: list[Any] = [] captured: list[Any] = []

View file

@ -67,7 +67,7 @@ async def test_fresh_sandbox_creating_waits_for_other_worker() -> None:
{"metadata": {"sandbox_id": "sandbox-existing", "sandbox_creating_at": fresh_at}}, {"metadata": {"sandbox_id": "sandbox-existing", "sandbox_creating_at": fresh_at}},
] ]
async def passthrough(sb, _thread_id): async def passthrough(sb, _thread_id, _github_proxy_token=None):
return sb return sb
with ( with (