mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-07 16:19:09 +00:00
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.
This commit is contained in:
parent
ba9865e1cf
commit
e4bee74c53
11 changed files with 298 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
124
tests/test_repo_binding_isolation.py
Normal file
124
tests/test_repo_binding_isolation.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue