diff --git a/agent/middleware/__init__.py b/agent/middleware/__init__.py index 9e4b369d..0318d530 100644 --- a/agent/middleware/__init__.py +++ b/agent/middleware/__init__.py @@ -3,6 +3,7 @@ from .ensure_no_empty_msg import ensure_no_empty_msg from .exclude_tools import ExcludeToolsMiddleware from .model_fallback import ModelFallbackMiddleware from .notify_step_limit import notify_step_limit_reached +from .refresh_github_proxy import refresh_github_proxy_before_model from .refresh_slack_status import SlackAssistantStatusMiddleware from .sandbox_circuit_breaker import SandboxCircuitBreakerMiddleware from .sanitize_thinking_blocks import SanitizeThinkingBlocksMiddleware @@ -23,5 +24,6 @@ __all__ = [ "check_message_queue_before_model", "ensure_no_empty_msg", "notify_step_limit_reached", + "refresh_github_proxy_before_model", "settle_review_check_on_exit", ] diff --git a/agent/middleware/refresh_github_proxy.py b/agent/middleware/refresh_github_proxy.py new file mode 100644 index 00000000..84adbb2f --- /dev/null +++ b/agent/middleware/refresh_github_proxy.py @@ -0,0 +1,47 @@ +"""Before-model middleware that keeps the sandbox GitHub proxy token fresh. + +The LangSmith sandbox proxy is configured with a GitHub App installation token +that expires after exactly one hour. Long runs would otherwise hit 401s on +every ``gh``/``git`` call once that snapshot goes stale. This hook re-configures +the proxy with a fresh token before each model call when the recorded token is +near expiry. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from langchain.agents.middleware import AgentState, before_model +from langgraph.config import get_config +from langgraph.runtime import Runtime + +from ..utils.github_proxy import maybe_refresh_proxy_token + +logger = logging.getLogger(__name__) + + +@before_model +async def refresh_github_proxy_before_model( + state: AgentState, # noqa: ARG001 + runtime: Runtime, # noqa: ARG001 +) -> dict[str, Any] | None: + """Refresh the sandbox proxy's GitHub token before it expires mid-run.""" + try: + config = get_config() + thread_id = config.get("configurable", {}).get("thread_id") + except Exception: # noqa: BLE001 + return None + + if not thread_id: + return None + + try: + await maybe_refresh_proxy_token(thread_id) + except Exception: # noqa: BLE001 + logger.warning( + "Failed to refresh GitHub proxy token for thread %s", + thread_id, + exc_info=True, + ) + return None diff --git a/agent/reviewer.py b/agent/reviewer.py index 982d5a26..e3df0d0f 100644 --- a/agent/reviewer.py +++ b/agent/reviewer.py @@ -38,6 +38,7 @@ from .middleware import ( SlackAssistantStatusMiddleware, ToolErrorMiddleware, check_message_queue_before_model, + refresh_github_proxy_before_model, settle_review_check_on_exit, ) from .reviewer_diff import compute_diff_line_set, fetch_pr_diff, fetch_pr_metadata @@ -709,9 +710,11 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: github_proxy_token = github_token github_api_token = github_token + repo_name_for_scope = str(repo_config.get("name") or "") sandbox_backend = await ensure_sandbox_for_thread( thread_id, github_proxy_token=github_proxy_token, + github_proxy_repositories=[repo_name_for_scope] if repo_name_for_scope else None, ) work_dir = await aresolve_sandbox_work_dir(sandbox_backend) @@ -994,6 +997,7 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: SanitizeToolInputsMiddleware(), ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"), ToolErrorMiddleware(), + refresh_github_proxy_before_model, check_message_queue_before_model, SlackAssistantStatusMiddleware(), SanitizeThinkingBlocksMiddleware(), diff --git a/agent/server.py b/agent/server.py index 265a077a..94481921 100644 --- a/agent/server.py +++ b/agent/server.py @@ -7,6 +7,7 @@ import logging import os import time import warnings +from collections.abc import Sequence from typing import Any logger = logging.getLogger(__name__) @@ -56,6 +57,7 @@ from .middleware import ( check_message_queue_before_model, ensure_no_empty_msg, notify_step_limit_reached, + refresh_github_proxy_before_model, ) from .prompt import construct_system_prompt from .tools import ( @@ -80,7 +82,10 @@ from .utils.authorship import ( OPEN_SWE_BOT_NAME, resolve_triggering_user_identity, ) -from .utils.github_app import get_github_app_installation_token +from .utils.github_app import ( + get_github_app_installation_token_with_expiry, +) +from .utils.github_proxy import record_proxy_token_expiry from .utils.model import ( DEFAULT_LLM_REASONING, ModelKwargs, @@ -162,21 +167,37 @@ async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProt await asyncio.to_thread(sandbox.start) +async def _resolve_proxy_token(github_proxy_token: str | None) -> tuple[str | None, str | None]: + """Resolve the proxy token and its expiry. + + An explicitly supplied token has no known expiry; otherwise we mint a fresh + GitHub App installation token and keep its ``expires_at`` so the proxy can + be refreshed before the (hard 1h) expiry. + """ + if github_proxy_token: + return github_proxy_token, None + return await get_github_app_installation_token_with_expiry() + + async def _create_sandbox_with_proxy( github_proxy_token: str | None = None, + *, + thread_id: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> SandboxBackendProtocol: """Create a new sandbox with GitHub proxy auth configured.""" sandbox_backend = await asyncio.to_thread(create_sandbox) sandbox_type = os.getenv("SANDBOX_TYPE", "langsmith") if sandbox_type == "langsmith": - token = github_proxy_token or await get_github_app_installation_token() + token, expires_at = await _resolve_proxy_token(github_proxy_token) if not token: msg = "Cannot configure proxy: GitHub App installation token is unavailable" logger.error(msg) raise ValueError(msg) await _start_langsmith_sandbox_if_needed(sandbox_backend) await asyncio.to_thread(_configure_github_proxy, sandbox_backend.id, token) + record_proxy_token_expiry(thread_id, expires_at, repositories=github_proxy_repositories) return sandbox_backend @@ -184,12 +205,15 @@ async def _create_sandbox_with_proxy( async def _refresh_github_proxy( sandbox_backend: SandboxBackendProtocol, github_proxy_token: str | None = None, + *, + thread_id: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> None: """Refresh GitHub proxy credentials for reused LangSmith sandboxes.""" if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith": return - token = github_proxy_token or await get_github_app_installation_token() + token, expires_at = await _resolve_proxy_token(github_proxy_token) if not token: logger.warning( "Skipping GitHub proxy refresh for sandbox %s: installation token unavailable", @@ -200,16 +224,23 @@ async def _refresh_github_proxy( current_backend = unwrap_sandbox_backend(sandbox_backend) await _start_langsmith_sandbox_if_needed(current_backend) await asyncio.to_thread(_configure_github_proxy, current_backend.id, token) + record_proxy_token_expiry(thread_id, expires_at, repositories=github_proxy_repositories) async def _refresh_github_proxy_or_recreate( sandbox_backend: SandboxBackendProtocol, thread_id: str, github_proxy_token: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> SandboxBackendProtocol: """Refresh proxy credentials, recreating stale LangSmith sandboxes on failure.""" try: - await _refresh_github_proxy(sandbox_backend, github_proxy_token) + await _refresh_github_proxy( + sandbox_backend, + github_proxy_token, + thread_id=thread_id, + github_proxy_repositories=github_proxy_repositories, + ) except Exception: # noqa: BLE001 logger.warning( "Failed to refresh GitHub proxy for sandbox %s on thread %s, recreating sandbox", @@ -217,7 +248,11 @@ async def _refresh_github_proxy_or_recreate( thread_id, exc_info=True, ) - return await _recreate_sandbox(thread_id, github_proxy_token=github_proxy_token) + return await _recreate_sandbox( + thread_id, + github_proxy_token=github_proxy_token, + github_proxy_repositories=github_proxy_repositories, + ) return sandbox_backend @@ -233,6 +268,7 @@ async def _recreate_sandbox( thread_id: str, *, github_proxy_token: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> SandboxBackendProtocol: """Recreate a sandbox after a connection failure. @@ -244,7 +280,11 @@ async def _recreate_sandbox( try: sandbox_backend = set_sandbox_backend( thread_id, - await _create_sandbox_with_proxy(github_proxy_token), + await _create_sandbox_with_proxy( + github_proxy_token, + thread_id=thread_id, + github_proxy_repositories=github_proxy_repositories, + ), ) except Exception: logger.exception("Failed to recreate sandbox after connection failure") @@ -257,6 +297,7 @@ async def check_or_recreate_sandbox( sandbox_backend: SandboxBackendProtocol, thread_id: str, github_proxy_token: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> SandboxBackendProtocol: """Check if a cached sandbox is reachable; recreate it if not. @@ -273,7 +314,11 @@ async def check_or_recreate_sandbox( "Cached sandbox is no longer reachable for thread %s, recreating", thread_id, ) - sandbox_backend = await _recreate_sandbox(thread_id, github_proxy_token=github_proxy_token) + sandbox_backend = await _recreate_sandbox( + thread_id, + github_proxy_token=github_proxy_token, + github_proxy_repositories=github_proxy_repositories, + ) return sandbox_backend @@ -328,6 +373,7 @@ async def ensure_sandbox_for_thread( thread_id: str, *, github_proxy_token: str | None = None, + github_proxy_repositories: Sequence[str] | None = None, ) -> SandboxBackendProtocol: """Get-or-create a healthy sandbox bound to ``thread_id``. @@ -354,17 +400,21 @@ async def ensure_sandbox_for_thread( logger.info("Using cached sandbox backend for thread %s", thread_id) original_sandbox_id = sandbox_backend.id sandbox_backend = await check_or_recreate_sandbox( - sandbox_backend, thread_id, github_proxy_token + sandbox_backend, thread_id, github_proxy_token, github_proxy_repositories ) if sandbox_backend.id == original_sandbox_id: sandbox_backend = await _refresh_github_proxy_or_recreate( - sandbox_backend, thread_id, github_proxy_token + sandbox_backend, thread_id, github_proxy_token, github_proxy_repositories ) elif sandbox_id is None: logger.info("Creating new sandbox for thread %s", thread_id) await client.threads.update(thread_id=thread_id, metadata=_creating_metadata()) try: - sandbox_backend = await _create_sandbox_with_proxy(github_proxy_token) + sandbox_backend = await _create_sandbox_with_proxy( + github_proxy_token, + thread_id=thread_id, + github_proxy_repositories=github_proxy_repositories, + ) logger.info("Sandbox created: %s", sandbox_backend.id) except Exception: logger.exception("Failed to create sandbox") @@ -382,7 +432,11 @@ async def ensure_sandbox_for_thread( 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()) try: - sandbox_backend = await _create_sandbox_with_proxy(github_proxy_token) + sandbox_backend = await _create_sandbox_with_proxy( + github_proxy_token, + thread_id=thread_id, + github_proxy_repositories=github_proxy_repositories, + ) created_replacement_sandbox = True except Exception: logger.exception("Failed to create replacement sandbox") @@ -391,11 +445,11 @@ async def ensure_sandbox_for_thread( if not created_replacement_sandbox: original_sandbox_id = sandbox_backend.id sandbox_backend = await check_or_recreate_sandbox( - sandbox_backend, thread_id, github_proxy_token + sandbox_backend, thread_id, github_proxy_token, github_proxy_repositories ) if sandbox_backend.id == original_sandbox_id: sandbox_backend = await _refresh_github_proxy_or_recreate( - sandbox_backend, thread_id, github_proxy_token + sandbox_backend, thread_id, github_proxy_token, github_proxy_repositories ) sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend) @@ -662,6 +716,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"), ToolErrorMiddleware(), ToolArtifactMiddleware(), + refresh_github_proxy_before_model, check_message_queue_before_model, SlackAssistantStatusMiddleware(), ensure_no_empty_msg, diff --git a/agent/utils/github_proxy.py b/agent/utils/github_proxy.py new file mode 100644 index 00000000..af288e5f --- /dev/null +++ b/agent/utils/github_proxy.py @@ -0,0 +1,129 @@ +"""Track and refresh the GitHub App token baked into a sandbox's proxy. + +The LangSmith sandbox proxy is configured once at run start with a GitHub App +installation token. Those tokens expire after exactly one hour, so any agent +run longer than ~1h would start seeing 401s on every ``gh``/``git`` call in the +sandbox. This module records when each thread's proxy token expires and lets a +before-model middleware re-configure the proxy before it goes stale. +""" + +from __future__ import annotations + +import asyncio +import logging +import os +from collections.abc import Sequence +from datetime import UTC, datetime, timedelta +from typing import Any + +from .github_app import get_github_app_installation_token_with_expiry +from .sandbox_state import SANDBOX_BACKENDS, unwrap_sandbox_backend + +logger = logging.getLogger(__name__) + +# Refresh the proxy token once it is within this window of expiring. +PROXY_TOKEN_REFRESH_WINDOW = timedelta(minutes=5) +# Used only when the token's own expiry is unknown: refresh after this age. +PROXY_TOKEN_FALLBACK_TTL = timedelta(minutes=50) + +# thread_id -> (token_expires_at | None, recorded_at, repositories scope | None) +_PROXY_TOKEN_EXPIRY: dict[str, tuple[datetime | None, datetime, tuple[str, ...] | None]] = {} + + +def _parse_expiry(expires_at: Any) -> datetime | None: + """Best-effort parse of a GitHub ``expires_at`` value to an aware datetime.""" + if expires_at is None: + return None + if isinstance(expires_at, datetime): + return expires_at if expires_at.tzinfo else expires_at.replace(tzinfo=UTC) + if isinstance(expires_at, int | float): + try: + return datetime.fromtimestamp(float(expires_at), tz=UTC) + except (OverflowError, OSError, ValueError): + return None + if isinstance(expires_at, str): + raw = expires_at.strip() + if not raw: + return None + if raw.endswith("Z"): + raw = raw[:-1] + "+00:00" + try: + parsed = datetime.fromisoformat(raw) + except ValueError: + return None + return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC) + return None + + +def record_proxy_token_expiry( + thread_id: str | None, + expires_at: Any, + *, + repositories: Sequence[str] | None = None, +) -> None: + """Record when ``thread_id``'s proxy token expires and the repo scope it was minted with. + + ``repositories`` preserves the original token scope (reviewer runs mint a + repo-scoped installation token) so a later refresh doesn't broaden it to an + installation-wide token. + """ + if not thread_id: + return + scope = tuple(repositories) if repositories else None + _PROXY_TOKEN_EXPIRY[thread_id] = (_parse_expiry(expires_at), datetime.now(UTC), scope) + + +def clear_proxy_token_expiry(thread_id: str | None) -> None: + if thread_id: + _PROXY_TOKEN_EXPIRY.pop(thread_id, None) + + +def proxy_token_needs_refresh(thread_id: str | None, *, now: datetime | None = None) -> bool: + """Whether the recorded proxy token is at/near expiry and should be refreshed.""" + if not thread_id: + return False + record = _PROXY_TOKEN_EXPIRY.get(thread_id) + if record is None: + return False + expires_at, recorded_at, _scope = record + current = (now or datetime.now(UTC)).astimezone(UTC) + if expires_at is not None: + return (expires_at - current) <= PROXY_TOKEN_REFRESH_WINDOW + return (current - recorded_at) >= PROXY_TOKEN_FALLBACK_TTL + + +async def maybe_refresh_proxy_token(thread_id: str | None, *, now: datetime | None = None) -> bool: + """Re-configure the sandbox proxy with a fresh token when near expiry. + + Returns True when a refresh was performed. Only applies to LangSmith + sandboxes; other providers don't use the proxy. + """ + if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith": + return False + if not thread_id or not proxy_token_needs_refresh(thread_id, now=now): + return False + + sandbox_backend = SANDBOX_BACKENDS.get(thread_id) + if sandbox_backend is None: + return False + + # Preserve the original token scope: reviewer runs mint a repo-scoped token, + # so refreshing must not broaden it to an installation-wide token. + _expires, _recorded, repositories = _PROXY_TOKEN_EXPIRY.get(thread_id, (None, None, None)) + token, expires_at = await get_github_app_installation_token_with_expiry( + repositories=list(repositories) if repositories else None + ) + if not token: + logger.warning( + "Proxy token for thread %s is near expiry but no installation token is available", + thread_id, + ) + return False + + from ..integrations.langsmith import _configure_github_proxy + + current_backend = unwrap_sandbox_backend(sandbox_backend) + await asyncio.to_thread(_configure_github_proxy, current_backend.id, token) + record_proxy_token_expiry(thread_id, expires_at, repositories=repositories) + logger.info("Refreshed GitHub proxy token for thread %s before expiry", thread_id) + return True diff --git a/tests/test_github_proxy_refresh.py b/tests/test_github_proxy_refresh.py new file mode 100644 index 00000000..8d630b92 --- /dev/null +++ b/tests/test_github_proxy_refresh.py @@ -0,0 +1,210 @@ +"""Tests for mid-run GitHub proxy token refresh.""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from agent.utils import github_proxy +from agent.utils.github_proxy import ( + PROXY_TOKEN_FALLBACK_TTL, + clear_proxy_token_expiry, + maybe_refresh_proxy_token, + proxy_token_needs_refresh, + record_proxy_token_expiry, +) + + +@pytest.fixture(autouse=True) +def _clear_state() -> None: + github_proxy._PROXY_TOKEN_EXPIRY.clear() + yield + github_proxy._PROXY_TOKEN_EXPIRY.clear() + + +class TestProxyTokenNeedsRefresh: + def test_false_when_no_record(self) -> None: + assert proxy_token_needs_refresh("thread-1") is False + + def test_false_when_thread_id_missing(self) -> None: + assert proxy_token_needs_refresh(None) is False + + def test_true_when_near_expiry(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=2)) + assert proxy_token_needs_refresh("thread-1", now=now) is True + + def test_false_when_far_from_expiry(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=55)) + assert proxy_token_needs_refresh("thread-1", now=now) is False + + def test_parses_iso_z_suffix(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", "2025-01-01T12:03:00Z") + assert proxy_token_needs_refresh("thread-1", now=now) is True + + def test_fallback_ttl_when_expiry_unknown(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", None) + github_proxy._PROXY_TOKEN_EXPIRY["thread-1"] = (None, now, None) + assert proxy_token_needs_refresh("thread-1", now=now) is False + later = now + PROXY_TOKEN_FALLBACK_TTL + assert proxy_token_needs_refresh("thread-1", now=later) is True + + def test_clear_removes_record(self) -> None: + record_proxy_token_expiry("thread-1", datetime.now(UTC)) + clear_proxy_token_expiry("thread-1") + assert proxy_token_needs_refresh("thread-1") is False + + +class TestMaybeRefreshProxyToken: + @pytest.mark.asyncio + async def test_skips_when_not_langsmith(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=1)) + with patch.dict("os.environ", {"SANDBOX_TYPE": "local"}): + assert await maybe_refresh_proxy_token("thread-1", now=now) is False + + @pytest.mark.asyncio + async def test_skips_when_not_near_expiry(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=55)) + with patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}): + assert await maybe_refresh_proxy_token("thread-1", now=now) is False + + @pytest.mark.asyncio + async def test_skips_when_no_sandbox(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=1)) + with ( + patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), + patch.dict(github_proxy.SANDBOX_BACKENDS, {}, clear=True), + ): + assert await maybe_refresh_proxy_token("thread-1", now=now) is False + + @pytest.mark.asyncio + async def test_refreshes_and_records_new_expiry(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=1)) + backend = MagicMock(id="sb-1") + new_expiry = "2025-01-01T13:00:00Z" + + with ( + patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), + patch.dict(github_proxy.SANDBOX_BACKENDS, {"thread-1": backend}, clear=True), + patch( + "agent.utils.github_proxy.get_github_app_installation_token_with_expiry", + new=AsyncMock(return_value=("ghs_new", new_expiry)), + ), + patch("agent.integrations.langsmith._configure_github_proxy") as mock_configure, + ): + result = await maybe_refresh_proxy_token("thread-1", now=now) + + assert result is True + mock_configure.assert_called_once_with("sb-1", "ghs_new") + expires_at, _recorded, _scope = github_proxy._PROXY_TOKEN_EXPIRY["thread-1"] + assert expires_at == datetime(2025, 1, 1, 13, 0, 0, tzinfo=UTC) + + @pytest.mark.asyncio + async def test_preserves_repo_scope_on_refresh(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=1), repositories=["open-swe"]) + backend = MagicMock(id="sb-1") + token_mock = AsyncMock(return_value=("ghs_new", "2025-01-01T13:00:00Z")) + + with ( + patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), + patch.dict(github_proxy.SANDBOX_BACKENDS, {"thread-1": backend}, clear=True), + patch( + "agent.utils.github_proxy.get_github_app_installation_token_with_expiry", + new=token_mock, + ), + patch("agent.integrations.langsmith._configure_github_proxy"), + ): + result = await maybe_refresh_proxy_token("thread-1", now=now) + + assert result is True + token_mock.assert_awaited_once_with(repositories=["open-swe"]) + _expires, _recorded, scope = github_proxy._PROXY_TOKEN_EXPIRY["thread-1"] + assert scope == ("open-swe",) + + @pytest.mark.asyncio + async def test_no_refresh_when_token_unavailable(self) -> None: + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + record_proxy_token_expiry("thread-1", now + timedelta(minutes=1)) + backend = MagicMock(id="sb-1") + + with ( + patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), + patch.dict(github_proxy.SANDBOX_BACKENDS, {"thread-1": backend}, clear=True), + patch( + "agent.utils.github_proxy.get_github_app_installation_token_with_expiry", + new=AsyncMock(return_value=(None, None)), + ), + patch("agent.integrations.langsmith._configure_github_proxy") as mock_configure, + ): + result = await maybe_refresh_proxy_token("thread-1", now=now) + + assert result is False + mock_configure.assert_not_called() + + +class TestRefreshGithubProxyMiddleware: + @pytest.mark.asyncio + async def test_calls_refresh_with_thread_id(self) -> None: + from agent.middleware.refresh_github_proxy import refresh_github_proxy_before_model + + with ( + patch( + "agent.middleware.refresh_github_proxy.get_config", + return_value={"configurable": {"thread_id": "thread-9"}}, + ), + patch( + "agent.middleware.refresh_github_proxy.maybe_refresh_proxy_token", + new=AsyncMock(return_value=True), + ) as mock_refresh, + ): + result = await refresh_github_proxy_before_model.abefore_model({}, MagicMock()) + + assert result is None + mock_refresh.assert_awaited_once_with("thread-9") + + @pytest.mark.asyncio + async def test_no_thread_id_is_noop(self) -> None: + from agent.middleware.refresh_github_proxy import refresh_github_proxy_before_model + + with ( + patch( + "agent.middleware.refresh_github_proxy.get_config", + return_value={"configurable": {}}, + ), + patch( + "agent.middleware.refresh_github_proxy.maybe_refresh_proxy_token", + new=AsyncMock(), + ) as mock_refresh, + ): + result = await refresh_github_proxy_before_model.abefore_model({}, MagicMock()) + + assert result is None + mock_refresh.assert_not_called() + + @pytest.mark.asyncio + async def test_swallows_refresh_errors(self) -> None: + from agent.middleware.refresh_github_proxy import refresh_github_proxy_before_model + + with ( + patch( + "agent.middleware.refresh_github_proxy.get_config", + return_value={"configurable": {"thread_id": "thread-9"}}, + ), + patch( + "agent.middleware.refresh_github_proxy.maybe_refresh_proxy_token", + new=AsyncMock(side_effect=RuntimeError("boom")), + ), + ): + result = await refresh_github_proxy_before_model.abefore_model({}, MagicMock()) + + assert result is None diff --git a/tests/test_proxy_auth.py b/tests/test_proxy_auth.py index 5bbaf33c..a7eb2596 100644 --- a/tests/test_proxy_auth.py +++ b/tests/test_proxy_auth.py @@ -184,9 +184,9 @@ class TestCreateSandboxWithProxy: """Installation token should be used for proxy auth on langsmith sandboxes.""" with ( patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_install", + return_value=("ghs_install", None), ), patch("agent.server.create_sandbox") as mock_create, patch("agent.server._configure_github_proxy") as mock_proxy, @@ -224,9 +224,9 @@ class TestCreateSandboxWithProxy: with ( patch("agent.server.create_sandbox") as mock_create, patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value=None, + return_value=(None, None), ), patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), ): @@ -275,9 +275,9 @@ class TestRefreshProxyOnSandboxReuse: return_value="sandbox-cached", ), patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_fresh", + return_value=("ghs_fresh", None), ), patch("agent.server._configure_github_proxy") as mock_proxy, patch( @@ -325,9 +325,9 @@ class TestRefreshProxyOnSandboxReuse: ), patch("agent.server.create_sandbox", return_value=mock_sandbox) as mock_create, patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_fresh", + return_value=("ghs_fresh", None), ), patch("agent.server._configure_github_proxy") as mock_proxy, patch( @@ -360,9 +360,9 @@ class TestRefreshProxyOnSandboxReuse: with ( patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_fresh", + return_value=("ghs_fresh", None), ), patch( "agent.server._configure_github_proxy", @@ -385,7 +385,9 @@ class TestRefreshProxyOnSandboxReuse: assert sandbox is replacement_sandbox mock_proxy.assert_called_once_with("sandbox-stale", "ghs_fresh") - mock_recreate.assert_awaited_once_with("thread-123", github_proxy_token=None) + mock_recreate.assert_awaited_once_with( + "thread-123", github_proxy_token=None, github_proxy_repositories=None + ) @pytest.mark.asyncio async def test_starts_stopped_langsmith_sandbox_before_proxy_refresh(self) -> None: @@ -396,9 +398,9 @@ class TestRefreshProxyOnSandboxReuse: with ( patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_fresh", + return_value=("ghs_fresh", None), ), patch("agent.server._configure_github_proxy") as mock_proxy, patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), @@ -423,9 +425,9 @@ class TestRefreshProxyOnSandboxReuse: with ( patch( - "agent.server.get_github_app_installation_token", + "agent.server.get_github_app_installation_token_with_expiry", new_callable=AsyncMock, - return_value="ghs_fresh", + return_value=("ghs_fresh", None), ), patch("agent.server._configure_github_proxy") as mock_proxy, patch.dict("os.environ", {"SANDBOX_TYPE": "langsmith"}), diff --git a/tests/test_reviewer.py b/tests/test_reviewer.py index 3ff89ebd..5d92be9e 100644 --- a/tests/test_reviewer.py +++ b/tests/test_reviewer.py @@ -202,6 +202,7 @@ async def test_reviewer_reuses_app_token_for_sandbox_proxy() -> None: mock_sandbox.assert_awaited_once_with( "reviewer-thread-id", github_proxy_token="app-token", + github_proxy_repositories=["repo"], ) diff --git a/tests/test_stale_sandbox_creating.py b/tests/test_stale_sandbox_creating.py index 1088457e..bb477f24 100644 --- a/tests/test_stale_sandbox_creating.py +++ b/tests/test_stale_sandbox_creating.py @@ -67,7 +67,9 @@ async def test_fresh_sandbox_creating_waits_for_other_worker() -> None: {"metadata": {"sandbox_id": "sandbox-existing", "sandbox_creating_at": fresh_at}}, ] - async def passthrough(sb, _thread_id, _github_proxy_token=None): + async def passthrough( + sb, _thread_id, _github_proxy_token=None, _github_proxy_repositories=None + ): return sb with (