diff --git a/agent/reviewer.py b/agent/reviewer.py index ac7285fa..801a4ec9 100644 --- a/agent/reviewer.py +++ b/agent/reviewer.py @@ -880,7 +880,7 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: ) # 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, repo=repo_config + thread_id, github_token, expires_at=expires_at, repo=repo_config, is_bot_token=True ) github_proxy_token = github_token diff --git a/agent/utils/auth.py b/agent/utils/auth.py index c061cded..5309a386 100644 --- a/agent/utils/auth.py +++ b/agent/utils/auth.py @@ -14,7 +14,11 @@ from langgraph.graph.state import RunnableConfig from langgraph_sdk import get_client from .github_app import get_github_app_installation_token_with_expiry -from .github_token import cache_github_token_for_thread, get_github_token_from_thread +from .github_token import ( + cache_github_token_for_thread, + get_github_token_from_thread, + github_token_principal, +) from .http import DEFAULT_HTTP_TIMEOUT from .linear import comment_on_linear_issue from .slack import post_slack_thread_reply @@ -304,9 +308,21 @@ def _current_repo() -> Any: def _cache_resolved_github_token( - thread_id: str, token: str, expires_at: str | None = None + thread_id: str, + token: str, + expires_at: str | None = None, + *, + principal: str | None = None, + is_bot_token: bool = False, ) -> tuple[str, str | None]: - cache_github_token_for_thread(thread_id, token, expires_at=expires_at, repo=_current_repo()) + cache_github_token_for_thread( + thread_id, + token, + expires_at=expires_at, + repo=_current_repo(), + principal=principal, + is_bot_token=is_bot_token, + ) return token, expires_at @@ -374,7 +390,13 @@ async def resolve_token_from_email( expires_at = auth_result.get("expires_at") if isinstance(auth_result, dict) else None return _cache_resolved_github_token( - thread_id, token, expires_at=expires_at if isinstance(expires_at, str) else None + thread_id, + token, + expires_at=expires_at if isinstance(expires_at, str) else None, + principal=github_token_principal( + login=configurable.get("github_login"), + email=email, + ), ) @@ -395,7 +417,10 @@ async def _resolve_dashboard_user_token( record = await get_oauth_record(OAUTH_TOKENS_NAMESPACE, login) expires_at = record.get("token_expires_at") if isinstance(record, dict) else None return _cache_resolved_github_token( - thread_id, token, expires_at=expires_at if isinstance(expires_at, str) else None + thread_id, + token, + expires_at=expires_at if isinstance(expires_at, str) else None, + principal=github_token_principal(login=login), ) @@ -417,7 +442,9 @@ async def _resolve_bot_installation_token(thread_id: str) -> tuple[str, str | No logger.info( "Using GitHub App installation token for thread %s (bot-token-only mode)", thread_id ) - return _cache_resolved_github_token(thread_id, bot_token, expires_at=expires_at) + return _cache_resolved_github_token( + thread_id, bot_token, expires_at=expires_at, is_bot_token=True + ) async def resolve_github_token(config: RunnableConfig, thread_id: str) -> tuple[str, str | None]: @@ -476,7 +503,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, expected_repo=configurable.get("repo") + thread_id, + principal=github_token_principal(login=github_login), + expected_repo=configurable.get("repo"), ) if cached_token: return cached_token, cached_expires_at diff --git a/agent/utils/github_token.py b/agent/utils/github_token.py index b44353eb..357c8e7d 100644 --- a/agent/utils/github_token.py +++ b/agent/utils/github_token.py @@ -17,10 +17,23 @@ _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, 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]] = {} +_BOT_PRINCIPAL = "bot" +# (thread_id, principal) -> (token, token_expires_at, cached_at, bound_repo). +# ``principal`` binds user tokens to the user they were resolved for so a +# cached token can never be served to a different user on the same thread. +# ``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[tuple[str, str], tuple[str, str | None, datetime, str | None]] = {} + + +def github_token_principal(*, login: str | None = None, email: str | None = None) -> str | None: + """Return the normalized principal used to isolate cached user tokens.""" + if isinstance(login, str) and login.strip(): + return f"login:{login.strip().casefold()}" + if isinstance(email, str) and email.strip(): + return f"email:{email.strip().casefold()}" + return None class GitHubAuthError(Exception): @@ -46,13 +59,28 @@ def repo_cache_key(repo: Any) -> str | None: def cache_github_token_for_thread( - thread_id: str, token: str, expires_at: str | None = None, *, repo: Any = None + thread_id: str, + token: str, + expires_at: str | None = None, + *, + repo: Any = None, + principal: str | None = None, + is_bot_token: bool = False, ) -> None: - """Cache a GitHub token in process for the current thread.""" + """Cache a GitHub token in process for the current thread and principal.""" if not thread_id or not token: return + cache_principal = _BOT_PRINCIPAL if is_bot_token else principal + if not cache_principal: + logger.warning("Refusing to cache an unbound user GitHub token for thread %s", thread_id) + return now = datetime.now(UTC) - _GITHUB_TOKEN_CACHE[thread_id] = (token, expires_at, now, repo_cache_key(repo)) + _GITHUB_TOKEN_CACHE[(thread_id, cache_principal)] = ( + token, + expires_at, + now, + repo_cache_key(repo), + ) _evict_expired(now=now) @@ -97,39 +125,45 @@ def _entry_expired(expires_at: str | None, cached_at: datetime, *, now: datetime def _evict_expired(*, now: datetime | None = None) -> None: current = now or datetime.now(UTC) stale = [ - tid - for tid, (_token, expires_at, cached_at, _repo) in _GITHUB_TOKEN_CACHE.items() + key + for key, (_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) + for key in stale: + _GITHUB_TOKEN_CACHE.pop(key, None) def _cached_token_if_fresh( - thread_id: str | None, *, expected_repo: Any = None + thread_id: str | None, principal: 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, 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 + keys = [] + if principal: + keys.append((thread_id, principal)) + keys.append((thread_id, _BOT_PRINCIPAL)) 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 + for key in keys: + cached = _GITHUB_TOKEN_CACHE.get(key) + if not cached: + continue + token, expires_at, cached_at, bound_repo = cached + if _entry_expired(expires_at, cached_at, now=datetime.now(UTC)): + _GITHUB_TOKEN_CACHE.pop(key, None) + logger.info("Cached GitHub token for thread %s has expired; re-resolving", thread_id) + continue + if expected and bound_repo and expected != bound_repo: + _GITHUB_TOKEN_CACHE.pop(key, 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, + ) + continue + return token, expires_at + return None, None def _thread_id_from_config(run_config: Mapping[str, Any]) -> str | None: @@ -147,23 +181,36 @@ def _repo_from_config(run_config: Mapping[str, Any]) -> Any: return configurable.get("repo") +def _principal_from_config(run_config: Mapping[str, Any]) -> str | None: + configurable = run_config.get("configurable", {}) + if not isinstance(configurable, Mapping): + return None + return github_token_principal( + login=configurable.get("github_login"), + email=configurable.get("user_email"), + ) + + 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), expected_repo=_repo_from_config(resolved) + _thread_id_from_config(resolved), + _principal_from_config(resolved), + expected_repo=_repo_from_config(resolved), ) return token async def get_github_token_from_thread( - thread_id: str, *, expected_repo: Any = None + thread_id: str, *, principal: str | None = None, 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, expected_repo=expected_repo) + """Resolve the current process's cached GitHub token for a thread and principal.""" + return _cached_token_if_fresh(thread_id, principal, expected_repo=expected_repo) async def invalidate_cached_github_token(thread_id: str) -> None: - """Clear a cached GitHub token for a thread.""" - _GITHUB_TOKEN_CACHE.pop(thread_id, None) + """Clear every cached GitHub token for a thread.""" + for key in [key for key in _GITHUB_TOKEN_CACHE if key[0] == thread_id]: + _GITHUB_TOKEN_CACHE.pop(key, None) logger.info("Invalidated cached GitHub token for thread %s", thread_id) diff --git a/agent/webhooks/common.py b/agent/webhooks/common.py index edbd4399..23cb03d9 100644 --- a/agent/webhooks/common.py +++ b/agent/webhooks/common.py @@ -89,6 +89,7 @@ from ..utils.github_org_membership import INTERNAL_BOT_LOGINS, is_user_active_or from ..utils.github_token import ( cache_github_token_for_thread, get_github_token_from_thread, + github_token_principal, invalidate_cached_github_token, ) from ..utils.http import DEFAULT_HTTP_TIMEOUT @@ -1656,12 +1657,17 @@ async def _get_or_resolve_thread_github_token( 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, repo=repo) + cache_github_token_for_thread( + thread_id, bot_token, expires_at=expires_at, repo=repo, is_bot_token=True + ) 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, expected_repo=repo) + principal = github_token_principal(email=email) + github_token, _expires_at = await get_github_token_from_thread( + thread_id, principal=principal, expected_repo=repo + ) if github_token: return github_token @@ -1676,6 +1682,7 @@ async def _get_or_resolve_thread_github_token( github_token, expires_at=expires_at if isinstance(expires_at, str) else None, repo=repo, + principal=principal, ) return github_token diff --git a/tests/auth/test_auth_sources.py b/tests/auth/test_auth_sources.py index fc3d0282..90eca352 100644 --- a/tests/auth/test_auth_sources.py +++ b/tests/auth/test_auth_sources.py @@ -64,7 +64,7 @@ def _stub_dashboard_store( ) -> None: from agent.dashboard import profiles - async def fake_get_from_thread(thread_id: str): + async def fake_get_from_thread(thread_id: str, **_kwargs): return cached async def fake_get_valid(login: str): diff --git a/tests/auth/test_github_token_ttl.py b/tests/auth/test_github_token_ttl.py index 79ad980d..1041ede6 100644 --- a/tests/auth/test_github_token_ttl.py +++ b/tests/auth/test_github_token_ttl.py @@ -50,42 +50,77 @@ def test_is_expired_treats_unparseable_as_not_expired() -> None: def test_get_github_token_returns_none_for_expired_cache() -> None: past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() - github_token.cache_github_token_for_thread("tid", "ghp_secret", expires_at=past) + github_token.cache_github_token_for_thread( + "tid", "ghp_secret", expires_at=past, is_bot_token=True + ) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) is None def test_get_github_token_returns_fresh_cached_token() -> None: future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() - github_token.cache_github_token_for_thread("tid", "ghp_secret", expires_at=future) + github_token.cache_github_token_for_thread( + "tid", "ghp_secret", expires_at=future, is_bot_token=True + ) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) == "ghp_secret" def test_get_github_token_returns_cached_token_when_no_expires_at() -> None: - github_token.cache_github_token_for_thread("tid", "ghp_secret") + github_token.cache_github_token_for_thread("tid", "ghp_secret", is_bot_token=True) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) == "ghp_secret" +def test_user_token_cache_is_bound_to_github_login() -> None: + principal = github_token.github_token_principal(login=" Alice ") + github_token.cache_github_token_for_thread("shared-thread", "alice-token", principal=principal) + + alice_config = {"configurable": {"thread_id": "shared-thread", "github_login": "ALICE"}} + bob_config = {"configurable": {"thread_id": "shared-thread", "github_login": "bob"}} + + assert github_token.get_github_token(alice_config) == "alice-token" + assert github_token.get_github_token(bob_config) is None + + +def test_unbound_user_token_is_not_cached() -> None: + github_token.cache_github_token_for_thread("tid", "unbound-token") + assert github_token._GITHUB_TOKEN_CACHE == {} + + +def test_cached_bot_token_is_available_to_any_principal() -> None: + github_token.cache_github_token_for_thread("tid", "bot-token", is_bot_token=True) + config = {"configurable": {"thread_id": "tid", "github_login": "alice"}} + assert github_token.get_github_token(config) == "bot-token" + + 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, None) + github_token._GITHUB_TOKEN_CACHE[("tid", github_token._BOT_PRINCIPAL)] = ( + "ghp_secret", + far_future, + old_cached_at, + None, + ) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) is None def test_cache_write_sweeps_other_expired_entries() -> None: """Writing one entry evicts unrelated entries that have passed their expiry.""" past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() - github_token.cache_github_token_for_thread("stale", "ghp_stale", expires_at=past) - github_token.cache_github_token_for_thread("fresh", "ghp_fresh") - assert "stale" not in github_token._GITHUB_TOKEN_CACHE - assert "fresh" in github_token._GITHUB_TOKEN_CACHE + github_token.cache_github_token_for_thread( + "stale", "ghp_stale", expires_at=past, is_bot_token=True + ) + github_token.cache_github_token_for_thread("fresh", "ghp_fresh", is_bot_token=True) + assert ("stale", github_token._BOT_PRINCIPAL) not in github_token._GITHUB_TOKEN_CACHE + assert ("fresh", github_token._BOT_PRINCIPAL) in github_token._GITHUB_TOKEN_CACHE @pytest.mark.asyncio async def test_get_github_token_from_thread_skips_expired() -> None: past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() - github_token.cache_github_token_for_thread("tid", "ghp_revoked", expires_at=past) + github_token.cache_github_token_for_thread( + "tid", "ghp_revoked", expires_at=past, is_bot_token=True + ) token, expires_at = await github_token.get_github_token_from_thread("tid") assert token is None assert expires_at is None @@ -94,7 +129,9 @@ async def test_get_github_token_from_thread_skips_expired() -> None: @pytest.mark.asyncio async def test_get_github_token_from_thread_returns_fresh() -> None: future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() - github_token.cache_github_token_for_thread("tid", "ghp_live", expires_at=future) + github_token.cache_github_token_for_thread( + "tid", "ghp_live", expires_at=future, is_bot_token=True + ) token, expires_at = await github_token.get_github_token_from_thread("tid") assert token == "ghp_live" assert expires_at == future @@ -102,7 +139,7 @@ async def test_get_github_token_from_thread_returns_fresh() -> None: @pytest.mark.asyncio async def test_invalidate_cached_github_token_clears_cache() -> None: - github_token.cache_github_token_for_thread("tid-42", "ghp_live") + github_token.cache_github_token_for_thread("tid-42", "ghp_live", is_bot_token=True) await github_token.invalidate_cached_github_token("tid-42") token, expires_at = await github_token.get_github_token_from_thread("tid-42") assert token is None diff --git a/tests/reviewer/test_reviewer.py b/tests/reviewer/test_reviewer.py index 1073c356..dee50ea5 100644 --- a/tests/reviewer/test_reviewer.py +++ b/tests/reviewer/test_reviewer.py @@ -235,7 +235,11 @@ async def test_reviewer_resolves_app_installation_token_at_run_start() -> None: # 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, repo={"owner": "acme", "name": "repo"} + "reviewer-thread-id", + "app-token", + expires_at=None, + repo={"owner": "acme", "name": "repo"}, + is_bot_token=True, ) middleware = create_agent.call_args.kwargs["middleware"] assert reviewer.check_message_queue_before_model in middleware diff --git a/tests/sandbox/test_repo_binding_isolation.py b/tests/sandbox/test_repo_binding_isolation.py index 6e9c94bd..32bc5106 100644 --- a/tests/sandbox/test_repo_binding_isolation.py +++ b/tests/sandbox/test_repo_binding_isolation.py @@ -28,20 +28,50 @@ 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"} + "tid", + "ghp_repoA", + expires_at=future, + repo={"owner": "acme", "name": "alpha"}, + is_bot_token=True, ) # 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 + assert all(key[0] != "tid" for key in github_token._GITHUB_TOKEN_CACHE) + + +def test_cached_user_token_not_reused_across_repos() -> None: + """Repo binding still applies to principal-bound user tokens.""" + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + github_token.cache_github_token_for_thread( + "tid", + "ghp_userA", + expires_at=future, + repo={"owner": "acme", "name": "alpha"}, + principal=github_token.github_token_principal(login="alice"), + ) + + cfg_same_user_other_repo = { + "configurable": { + "thread_id": "tid", + "github_login": "alice", + "repo": {"owner": "evil", "name": "beta"}, + } + } + assert github_token.get_github_token(cfg_same_user_other_repo) is None + assert all(key[0] != "tid" for key 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"} + "tid", + "ghp_repoA", + expires_at=future, + repo={"owner": "acme", "name": "alpha"}, + is_bot_token=True, ) cfg_a = {"configurable": {"thread_id": "tid", "repo": {"owner": "acme", "name": "alpha"}}} assert github_token.get_github_token(cfg_a) == "ghp_repoA" @@ -49,7 +79,7 @@ def test_cached_token_reused_for_same_repo() -> None: 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") + github_token.cache_github_token_for_thread("tid", "ghp_legacy", is_bot_token=True) assert github_token.get_github_token({"configurable": {"thread_id": "tid"}}) == "ghp_legacy" @@ -57,7 +87,11 @@ def test_cached_token_reused_for_same_repo_case_insensitive() -> None: """``Org/Repo`` and ``org/repo`` are the same repo: no spurious refusal.""" 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"} + "tid", + "ghp_repoA", + expires_at=future, + repo={"owner": "Acme", "name": "Alpha"}, + is_bot_token=True, ) cfg = {"configurable": {"thread_id": "tid", "repo": {"owner": "acme", "name": "alpha"}}} assert github_token.get_github_token(cfg) == "ghp_repoA"