From e4bee74c539181d94862425499144ad25760dbe6 Mon Sep 17 00:00:00 2001 From: Adam Moussa Date: Mon, 29 Jun 2026 11:37:26 -0400 Subject: [PATCH] fix: bind sandbox and token caches to repo to prevent thread-id collision A PR head-branch name is attacker-controllable and get_thread_id_from_branch derives a thread_id from its first UUID with no repo binding (TID-COLLIDE-01). The in-memory sandbox cache and the per-thread GitHub-token cache were keyed on thread_id alone, and a cached sandbox was reused after only an echo-ping, so a different repo's webhook could bind to another thread's sandbox or token. Without changing the persistent thread-id scheme: - Persist the bound repo (owner/name) in thread metadata on sandbox creation and refuse to reuse a sandbox whose bound repo does not match the current event (SandboxRepoMismatchError); the in-memory proxy also carries the binding. - Bind the GitHub-token cache entries to their repo and evict on a cross-repo read so a colliding thread_id cannot be served another repo's token. - Thread repo through the reviewer and the webhook token resolvers. --- agent/reviewer.py | 4 +- agent/server.py | 47 +++++++++- agent/utils/auth.py | 15 +++- agent/utils/github_token.py | 59 ++++++++++--- agent/utils/sandbox_state.py | 24 ++++++ agent/webapp.py | 42 ++++++--- tests/test_github_issue_webhook.py | 8 +- tests/test_github_token_ttl.py | 4 +- tests/test_proxy_auth.py | 2 + tests/test_repo_binding_isolation.py | 124 +++++++++++++++++++++++++++ tests/test_reviewer.py | 4 +- 11 files changed, 298 insertions(+), 35 deletions(-) create mode 100644 tests/test_repo_binding_isolation.py diff --git a/agent/reviewer.py b/agent/reviewer.py index 99e28020..e3a39c23 100644 --- a/agent/reviewer.py +++ b/agent/reviewer.py @@ -839,7 +839,9 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: f"GitHub App installation token unavailable for reviewer thread {thread_id}" ) # Cache in-process so reviewer tools and the sandbox proxy can read it this run. - cache_github_token_for_thread(thread_id, github_token, expires_at=expires_at) + cache_github_token_for_thread( + thread_id, github_token, expires_at=expires_at, repo=repo_config + ) github_proxy_token = github_token github_api_token = github_token diff --git a/agent/server.py b/agent/server.py index 42e58661..5faef2f6 100644 --- a/agent/server.py +++ b/agent/server.py @@ -98,6 +98,7 @@ from .utils.github_app import ( get_github_app_installation_token_with_expiry, ) from .utils.github_proxy import record_proxy_token_expiry +from .utils.github_token import repo_cache_key from .utils.model import ( DEFAULT_LLM_REASONING, ModelKwargs, @@ -117,6 +118,7 @@ SANDBOX_POLL_INTERVAL = 1.0 from .utils.sandbox_state import ( SANDBOX_BACKENDS, + get_bound_repo_from_metadata, get_sandbox_id_from_metadata, set_sandbox_backend, unwrap_sandbox_backend, @@ -405,6 +407,24 @@ def graph_loaded_for_execution(config: RunnableConfig) -> bool: ) +class SandboxRepoMismatchError(RuntimeError): + """Raised when a thread_id is presented for a repo it is not bound to. + + A thread is bound to exactly one repo. A different repo presenting a + colliding thread_id (e.g. an attacker-named branch whose first UUID matches + another thread) must never reuse this thread's sandbox or token. + """ + + def __init__(self, thread_id: str, bound_repo: str, current_repo: str) -> None: + self.thread_id = thread_id + self.bound_repo = bound_repo + self.current_repo = current_repo + super().__init__( + f"Thread {thread_id} is bound to repo {bound_repo}, " + f"refusing to serve sandbox for {current_repo}" + ) + + async def ensure_sandbox_for_thread( thread_id: str, *, @@ -431,6 +451,22 @@ async def ensure_sandbox_for_thread( sandbox_backend = SANDBOX_BACKENDS.get(thread_id) sandbox_id = await get_sandbox_id_from_metadata(thread_id) + # Repo-binding guard (TID-COLLIDE-01): refuse to reuse a thread's sandbox for + # a repo it is not bound to, so a colliding thread_id from a different repo + # cannot bind to (or clobber) another thread's sandbox. + current_repo = repo_cache_key(repo) + bound_repo = await get_bound_repo_from_metadata(thread_id) + proxy_bound = getattr(sandbox_backend, "bound_repo", None) + effective_bound = bound_repo or (proxy_bound if isinstance(proxy_bound, str) else None) + if current_repo and effective_bound and effective_bound != current_repo: + logger.error( + "Repo mismatch for thread %s: bound=%s current=%s; refusing sandbox reuse", + thread_id, + effective_bound, + current_repo, + ) + raise SandboxRepoMismatchError(thread_id, effective_bound, current_repo) + if sandbox_id == SANDBOX_CREATING and not sandbox_backend: logger.info("Sandbox creation in progress for thread %s, waiting...", thread_id) sandbox_id = await _resolve_creating_sentinel(thread_id) @@ -493,12 +529,15 @@ async def ensure_sandbox_for_thread( sandbox_backend, thread_id, github_proxy_token, github_proxy_repositories, repo ) - sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend) + sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend, repo=current_repo) + metadata_update: dict[str, Any] = {} if sandbox_id != sandbox_backend.id: - await client.threads.update( - thread_id=thread_id, metadata={"sandbox_id": sandbox_backend.id} - ) + metadata_update["sandbox_id"] = sandbox_backend.id + if current_repo and bound_repo != current_repo: + metadata_update["bound_repo"] = current_repo + if metadata_update: + await client.threads.update(thread_id=thread_id, metadata=metadata_update) # Re-apply git identity every run: cached/reconnected sandboxes may have # lost their `--global` config (or had it overwritten), and Vercel preview diff --git a/agent/utils/auth.py b/agent/utils/auth.py index 4141b0c3..a8a0dfb1 100644 --- a/agent/utils/auth.py +++ b/agent/utils/auth.py @@ -288,10 +288,19 @@ async def leave_failure_comment( raise ValueError(f"Unknown source: {source}") +def _current_repo() -> Any: + """Best-effort read of the run's repo (owner/name) for cache binding.""" + try: + configurable = get_config().get("configurable", {}) + except Exception: + return None + return configurable.get("repo") if isinstance(configurable, dict) else None + + def _cache_resolved_github_token( thread_id: str, token: str, expires_at: str | None = None ) -> tuple[str, str | None]: - cache_github_token_for_thread(thread_id, token, expires_at=expires_at) + cache_github_token_for_thread(thread_id, token, expires_at=expires_at, repo=_current_repo()) return token, expires_at @@ -452,7 +461,9 @@ async def resolve_github_token(config: RunnableConfig, thread_id: str) -> tuple[ try: if source == "github": - cached_token, cached_expires_at = await get_github_token_from_thread(thread_id) + cached_token, cached_expires_at = await get_github_token_from_thread( + thread_id, expected_repo=configurable.get("repo") + ) if cached_token: return cached_token, cached_expires_at from ..dashboard.user_mappings import email_for_login diff --git a/agent/utils/github_token.py b/agent/utils/github_token.py index 1738bf3c..6b2006ee 100644 --- a/agent/utils/github_token.py +++ b/agent/utils/github_token.py @@ -17,22 +17,37 @@ _GITHUB_TOKEN_EXPIRY_SKEW_SECONDS = 60 # Hard cap on how long an entry stays cached regardless of the token's own # expiry, so entries for threads that are never read again don't accumulate. _GITHUB_TOKEN_MAX_TTL = timedelta(hours=24) -# thread_id -> (token, token_expires_at, cached_at) -_GITHUB_TOKEN_CACHE: dict[str, tuple[str, str | None, datetime]] = {} +# thread_id -> (token, token_expires_at, cached_at, bound_repo). ``bound_repo`` +# ("owner/name") binds the entry to the repo it was resolved for so a colliding +# thread_id originating from a different repo cannot reuse another repo's token. +_GITHUB_TOKEN_CACHE: dict[str, tuple[str, str | None, datetime, str | None]] = {} class GitHubAuthError(Exception): """Raised when a GitHub call returns 401, signalling a stale/revoked token.""" +def repo_cache_key(repo: Any) -> str | None: + """Normalize a repo dict/string to ``owner/name`` (None when unknown).""" + if isinstance(repo, str): + cleaned = repo.strip() + return cleaned or None + if isinstance(repo, Mapping): + owner = repo.get("owner") + name = repo.get("name") + if isinstance(owner, str) and isinstance(name, str) and owner and name: + return f"{owner}/{name}" + return None + + def cache_github_token_for_thread( - thread_id: str, token: str, expires_at: str | None = None + thread_id: str, token: str, expires_at: str | None = None, *, repo: Any = None ) -> None: """Cache a GitHub token in process for the current thread.""" if not thread_id or not token: return now = datetime.now(UTC) - _GITHUB_TOKEN_CACHE[thread_id] = (token, expires_at, now) + _GITHUB_TOKEN_CACHE[thread_id] = (token, expires_at, now, repo_cache_key(repo)) _evict_expired(now=now) @@ -78,24 +93,37 @@ def _evict_expired(*, now: datetime | None = None) -> None: current = now or datetime.now(UTC) stale = [ tid - for tid, (_token, expires_at, cached_at) in _GITHUB_TOKEN_CACHE.items() + for tid, (_token, expires_at, cached_at, _repo) in _GITHUB_TOKEN_CACHE.items() if _entry_expired(expires_at, cached_at, now=current) ] for tid in stale: _GITHUB_TOKEN_CACHE.pop(tid, None) -def _cached_token_if_fresh(thread_id: str | None) -> tuple[str | None, str | None]: +def _cached_token_if_fresh( + thread_id: str | None, *, expected_repo: Any = None +) -> tuple[str | None, str | None]: if not thread_id: return None, None cached = _GITHUB_TOKEN_CACHE.get(thread_id) if not cached: return None, None - token, expires_at, cached_at = cached + token, expires_at, cached_at, bound_repo = cached if _entry_expired(expires_at, cached_at, now=datetime.now(UTC)): _GITHUB_TOKEN_CACHE.pop(thread_id, None) logger.info("Cached GitHub token for thread %s has expired; re-resolving", thread_id) return None, None + expected = repo_cache_key(expected_repo) + if expected and bound_repo and expected != bound_repo: + _GITHUB_TOKEN_CACHE.pop(thread_id, None) + logger.warning( + "Cached GitHub token for thread %s is bound to repo %s, not %s; " + "refusing cross-repo reuse", + thread_id, + bound_repo, + expected, + ) + return None, None return token, expires_at @@ -107,16 +135,27 @@ def _thread_id_from_config(run_config: Mapping[str, Any]) -> str | None: return thread_id if isinstance(thread_id, str) and thread_id else None +def _repo_from_config(run_config: Mapping[str, Any]) -> Any: + configurable = run_config.get("configurable", {}) + if not isinstance(configurable, Mapping): + return None + return configurable.get("repo") + + def get_github_token(run_config: Mapping[str, Any] | None = None) -> str | None: """Resolve the current thread's GitHub token from process memory.""" resolved = run_config if run_config is not None else get_config() - token, _expires_at = _cached_token_if_fresh(_thread_id_from_config(resolved)) + token, _expires_at = _cached_token_if_fresh( + _thread_id_from_config(resolved), expected_repo=_repo_from_config(resolved) + ) return token -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, *, expected_repo: Any = None +) -> tuple[str | None, str | None]: """Resolve the current process's cached GitHub token for a thread.""" - return _cached_token_if_fresh(thread_id) + return _cached_token_if_fresh(thread_id, expected_repo=expected_repo) async def invalidate_cached_github_token(thread_id: str) -> None: diff --git a/agent/utils/sandbox_state.py b/agent/utils/sandbox_state.py index caf7b6b4..90581273 100644 --- a/agent/utils/sandbox_state.py +++ b/agent/utils/sandbox_state.py @@ -29,6 +29,9 @@ class SandboxBackendProxy(SandboxBackendProtocol): def __init__(self, backend: SandboxBackendProtocol) -> None: self._backend = backend + # "owner/name" of the repo this sandbox is bound to, used to refuse + # reuse by a different repo presenting a colliding thread_id. + self.bound_repo: str | None = None @property def current(self) -> SandboxBackendProtocol: @@ -131,21 +134,42 @@ def unwrap_sandbox_backend(sandbox_backend: SandboxBackendProtocol) -> SandboxBa def set_sandbox_backend( thread_id: str, sandbox_backend: SandboxBackendProtocol, + *, + repo: str | None = None, ) -> SandboxBackendProxy: if isinstance(sandbox_backend, SandboxBackendProxy): + if repo: + sandbox_backend.bound_repo = repo SANDBOX_BACKENDS[thread_id] = sandbox_backend return sandbox_backend existing = SANDBOX_BACKENDS.get(thread_id) if isinstance(existing, SandboxBackendProxy): existing.replace_backend(sandbox_backend) + if repo: + existing.bound_repo = repo return existing proxy = SandboxBackendProxy(sandbox_backend) + if repo: + proxy.bound_repo = repo SANDBOX_BACKENDS[thread_id] = proxy return proxy +async def get_bound_repo_from_metadata(thread_id: str) -> str | None: + """Fetch the repo (``owner/name``) this thread's sandbox is bound to.""" + try: + config = get_config() + except Exception: + return None + metadata = config.get("metadata", {}) + if not isinstance(metadata, dict): + return None + bound_repo = metadata.get("bound_repo") + return bound_repo if isinstance(bound_repo, str) and bound_repo else None + + def clear_sandbox_backend(thread_id: str) -> None: SANDBOX_BACKENDS.pop(thread_id, None) diff --git a/agent/webapp.py b/agent/webapp.py index ece49d8b..f686ba16 100644 --- a/agent/webapp.py +++ b/agent/webapp.py @@ -2946,31 +2946,36 @@ async def process_github_autofix_review(payload: dict[str, Any], event_type: str ) -async def _refresh_thread_github_token_after_401(thread_id: str, email: str) -> str | None: +async def _refresh_thread_github_token_after_401( + thread_id: str, email: str, *, repo: dict[str, str] | None = None +) -> 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) + return await _get_or_resolve_thread_github_token(thread_id, email, repo=repo) -async def _get_or_resolve_thread_github_token(thread_id: str, email: str) -> str | None: +async def _get_or_resolve_thread_github_token( + thread_id: str, email: str, *, repo: dict[str, str] | None = None +) -> str | None: """Resolve and cache a GitHub token for a thread when available. In bot-token-only mode, returns a fresh GitHub App installation token - instead of resolving per-user OAuth tokens. + instead of resolving per-user OAuth tokens. ``repo`` (owner/name) binds the + cached entry so a colliding thread_id from a different repo cannot reuse it. """ if is_bot_token_only_mode(): bot_token, expires_at = await get_github_app_installation_token_with_expiry() if bot_token: - cache_github_token_for_thread(thread_id, bot_token, expires_at=expires_at) + cache_github_token_for_thread(thread_id, bot_token, expires_at=expires_at, repo=repo) return bot_token logger.warning("Bot-token-only mode but GitHub App token unavailable") return None - github_token, _expires_at = await get_github_token_from_thread(thread_id) + github_token, _expires_at = await get_github_token_from_thread(thread_id, expected_repo=repo) if github_token: return github_token @@ -2981,7 +2986,10 @@ async def _get_or_resolve_thread_github_token(thread_id: str, email: str) -> str expires_at = auth_result.get("expires_at") cache_github_token_for_thread( - thread_id, github_token, expires_at=expires_at if isinstance(expires_at, str) else None + thread_id, + github_token, + expires_at=expires_at if isinstance(expires_at, str) else None, + repo=repo, ) return github_token @@ -3043,7 +3051,7 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) -> email = await email_for_login(github_login) or "" if email: - github_token = await _get_or_resolve_thread_github_token(thread_id, email) + github_token = await _get_or_resolve_thread_github_token(thread_id, email, repo=repo_config) else: logger.warning("No email mapping for GitHub user '%s', skipping", github_login) return @@ -3063,7 +3071,9 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) -> node_id=node_id, ) except GitHubAuthError: - github_token = await _refresh_thread_github_token_after_401(thread_id, email) + github_token = await _refresh_thread_github_token_after_401( + thread_id, email, repo=repo_config + ) if not github_token: logger.warning("Re-auth failed for thread %s after 401; skipping", thread_id) return @@ -3085,7 +3095,9 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) -> repo_config, pr_number, token=github_token ) except GitHubAuthError: - github_token = await _refresh_thread_github_token_after_401(thread_id, email) + github_token = await _refresh_thread_github_token_after_401( + thread_id, email, repo=repo_config + ) if not github_token: logger.warning("Re-auth failed for thread %s after 401; skipping", thread_id) return @@ -3323,7 +3335,7 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None thread_id = generate_thread_id_from_github_issue(issue_id) existing_thread = await _thread_exists(thread_id) - github_token = await _get_or_resolve_thread_github_token(thread_id, email) + github_token = await _get_or_resolve_thread_github_token(thread_id, email, repo=repo_config) app_token = await get_github_app_installation_token() reaction_token = github_token or app_token comment = payload.get("comment", {}) @@ -3340,7 +3352,9 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None token=reaction_token, ) except GitHubAuthError: - github_token = await _refresh_thread_github_token_after_401(thread_id, email) + github_token = await _refresh_thread_github_token_after_401( + thread_id, email, repo=repo_config + ) reaction_token = github_token or app_token reacted = False if reaction_token: @@ -3374,7 +3388,9 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None repo_config, issue_number, token=github_token or app_token ) except GitHubAuthError: - github_token = await _refresh_thread_github_token_after_401(thread_id, email) + github_token = await _refresh_thread_github_token_after_401( + thread_id, email, repo=repo_config + ) comments = await fetch_issue_comments( repo_config, issue_number, token=github_token or app_token ) diff --git a/tests/test_github_issue_webhook.py b/tests/test_github_issue_webhook.py index eee94131..7971577d 100644 --- a/tests/test_github_issue_webhook.py +++ b/tests/test_github_issue_webhook.py @@ -1127,7 +1127,9 @@ def test_process_github_pr_comment_without_email_skips( def test_process_github_issue_uses_resolved_user_token_for_reaction(monkeypatch) -> None: captured: dict[str, object] = {} - async def fake_get_or_resolve_thread_github_token(thread_id: str, email: str) -> str | None: + async def fake_get_or_resolve_thread_github_token( + thread_id: str, email: str, *, repo: dict[str, str] | None = None + ) -> str | None: captured["thread_id"] = thread_id captured["email"] = email return "user-token" @@ -1210,7 +1212,9 @@ def test_process_github_issue_uses_resolved_user_token_for_reaction(monkeypatch) def test_process_github_issue_existing_thread_uses_followup_prompt(monkeypatch) -> None: captured: dict[str, object] = {} - async def fake_get_or_resolve_thread_github_token(thread_id: str, email: str) -> str | None: + async def fake_get_or_resolve_thread_github_token( + thread_id: str, email: str, *, repo: dict[str, str] | None = None + ) -> str | None: return "user-token" async def fake_get_github_app_installation_token() -> str | None: diff --git a/tests/test_github_token_ttl.py b/tests/test_github_token_ttl.py index bd3ee158..0d880814 100644 --- a/tests/test_github_token_ttl.py +++ b/tests/test_github_token_ttl.py @@ -69,7 +69,7 @@ def test_cached_token_expires_after_max_ttl() -> None: """A token with no/far expiry is still dropped once it's older than the 24h cap.""" far_future = (datetime.now(UTC) + timedelta(days=30)).isoformat() old_cached_at = datetime.now(UTC) - timedelta(hours=25) - github_token._GITHUB_TOKEN_CACHE["tid"] = ("ghp_secret", far_future, old_cached_at) + github_token._GITHUB_TOKEN_CACHE["tid"] = ("ghp_secret", far_future, old_cached_at, None) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) is None @@ -210,7 +210,7 @@ def test_process_github_pr_comment_invalidates_and_reauths_on_401( tokens = iter(["stale-token", "fresh-token"]) - async def fake_get_or_resolve(thread_id: str, email: str) -> str | None: + async def fake_get_or_resolve(thread_id: str, email: str, *, repo: Any = None) -> str | None: token = next(tokens) resolves.append(token) return token diff --git a/tests/test_proxy_auth.py b/tests/test_proxy_auth.py index 163bd1a4..cc827bcb 100644 --- a/tests/test_proxy_auth.py +++ b/tests/test_proxy_auth.py @@ -264,6 +264,7 @@ class TestRefreshProxyOnSandboxReuse: mock_sandbox = MagicMock(id="sandbox-cached") with ( + patch("agent.server.client.threads.update", new_callable=AsyncMock), patch( "agent.server.resolve_github_token", new_callable=AsyncMock, @@ -313,6 +314,7 @@ class TestRefreshProxyOnSandboxReuse: mock_sandbox = MagicMock(id="sandbox-existing") with ( + patch("agent.server.client.threads.update", new_callable=AsyncMock), patch( "agent.server.resolve_github_token", new_callable=AsyncMock, diff --git a/tests/test_repo_binding_isolation.py b/tests/test_repo_binding_isolation.py new file mode 100644 index 00000000..3ecf288b --- /dev/null +++ b/tests/test_repo_binding_isolation.py @@ -0,0 +1,124 @@ +"""Repo-binding isolation for the sandbox and GitHub-token caches. + +Covers TID-COLLIDE-01: a second repo presenting a colliding thread_id must not +reuse the first repo's in-process sandbox or cached GitHub token. +""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from typing import Any + +import pytest + +from agent import server +from agent.utils import github_token + + +@pytest.fixture(autouse=True) +def _clear_caches() -> None: + github_token._GITHUB_TOKEN_CACHE.clear() + server.SANDBOX_BACKENDS.clear() + + +# --- token cache repo binding ------------------------------------------------ + + +def test_cached_token_not_reused_across_repos() -> None: + """A token bound to repo A is refused for a colliding thread_id from repo B.""" + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + github_token.cache_github_token_for_thread( + "tid", "ghp_repoA", expires_at=future, repo={"owner": "acme", "name": "alpha"} + ) + + # Same thread_id, different repo → must NOT serve repo A's token. + cfg_b = {"configurable": {"thread_id": "tid", "repo": {"owner": "evil", "name": "beta"}}} + assert github_token.get_github_token(cfg_b) is None + # And the poisoned entry is evicted. + assert "tid" not in github_token._GITHUB_TOKEN_CACHE + + +def test_cached_token_reused_for_same_repo() -> None: + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + github_token.cache_github_token_for_thread( + "tid", "ghp_repoA", expires_at=future, repo={"owner": "acme", "name": "alpha"} + ) + cfg_a = {"configurable": {"thread_id": "tid", "repo": {"owner": "acme", "name": "alpha"}}} + assert github_token.get_github_token(cfg_a) == "ghp_repoA" + + +def test_unbound_token_served_when_repo_unknown() -> None: + """Legacy entries with no bound repo still resolve when no repo is supplied.""" + github_token.cache_github_token_for_thread("tid", "ghp_legacy") + assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) == "ghp_legacy" + + +def test_repo_cache_key_normalizes() -> None: + assert github_token.repo_cache_key({"owner": "o", "name": "r"}) == "o/r" + assert github_token.repo_cache_key("o/r") == "o/r" + assert github_token.repo_cache_key({"owner": "o"}) is None + assert github_token.repo_cache_key(None) is None + + +# --- sandbox repo binding ---------------------------------------------------- + + +class _FakeBackend: + def __init__(self, sandbox_id: str) -> None: + self.id = sandbox_id + self.bound_repo: str | None = None + + def execute(self, *_a: Any, **_k: Any) -> Any: # pragma: no cover - must not run + raise AssertionError("colliding-repo sandbox must not be pinged/reused") + + +@pytest.mark.asyncio +async def test_ensure_sandbox_refuses_colliding_repo(monkeypatch: pytest.MonkeyPatch) -> None: + async def fake_sandbox_id(_tid: str) -> str: + return "sb-A" + + async def fake_bound_repo(_tid: str) -> str: + return "acme/alpha" + + monkeypatch.setattr(server, "get_sandbox_id_from_metadata", fake_sandbox_id) + monkeypatch.setattr(server, "get_bound_repo_from_metadata", fake_bound_repo) + backend = _FakeBackend("sb-A") + backend.bound_repo = "acme/alpha" + server.SANDBOX_BACKENDS["tid"] = backend # type: ignore[assignment] + + with pytest.raises(server.SandboxRepoMismatchError): + await server.ensure_sandbox_for_thread("tid", repo={"owner": "evil", "name": "beta"}) + + +@pytest.mark.asyncio +async def test_ensure_sandbox_allows_matching_repo(monkeypatch: pytest.MonkeyPatch) -> None: + async def fake_sandbox_id(_tid: str) -> str: + return "sb-A" + + async def fake_bound_repo(_tid: str) -> str: + return "acme/alpha" + + calls: dict[str, int] = {"git": 0} + backend = _FakeBackend("sb-A") + backend.bound_repo = "acme/alpha" + + async def fake_check(b: Any, *_a: Any, **_k: Any) -> Any: + return b + + async def fake_refresh(b: Any, *_a: Any, **_k: Any) -> Any: + return b + + async def fake_git(_b: Any) -> None: + calls["git"] += 1 + + monkeypatch.setattr(server, "get_sandbox_id_from_metadata", fake_sandbox_id) + monkeypatch.setattr(server, "get_bound_repo_from_metadata", fake_bound_repo) + monkeypatch.setattr(server, "check_or_recreate_sandbox", fake_check) + monkeypatch.setattr(server, "_refresh_github_proxy_or_recreate", fake_refresh) + monkeypatch.setattr(server, "set_sandbox_backend", lambda _tid, b, **_k: b) + monkeypatch.setattr(server, "_configure_git_identity", fake_git) + server.SANDBOX_BACKENDS["tid"] = backend # type: ignore[assignment] + + result = await server.ensure_sandbox_for_thread("tid", repo={"owner": "acme", "name": "alpha"}) + assert result is backend + assert calls["git"] == 1 diff --git a/tests/test_reviewer.py b/tests/test_reviewer.py index 109fc90c..a381a0ce 100644 --- a/tests/test_reviewer.py +++ b/tests/test_reviewer.py @@ -221,7 +221,9 @@ async def test_reviewer_resolves_app_installation_token_at_run_start() -> None: # Token is resolved in this process at run start (scoped to the repo), not read # from a cache the webhook handler populated in a different process. mock_app_token.assert_awaited_once_with(repositories=["repo"]) - mock_cache_token.assert_called_once_with("reviewer-thread-id", "app-token", expires_at=None) + mock_cache_token.assert_called_once_with( + "reviewer-thread-id", "app-token", expires_at=None, repo={"owner": "acme", "name": "repo"} + ) middleware = create_agent.call_args.kwargs["middleware"] assert reviewer.check_message_queue_before_model in middleware