diff --git a/apps/agent/agent/middleware/open_pr.py b/apps/agent/agent/middleware/open_pr.py index 158b7671..af19bf70 100644 --- a/apps/agent/agent/middleware/open_pr.py +++ b/apps/agent/agent/middleware/open_pr.py @@ -32,7 +32,7 @@ from ..utils.github import ( git_push, ) from ..utils.linear import comment_on_linear_issue -from ..utils.sandbox_state import SANDBOX_BACKENDS +from ..utils.sandbox_state import get_sandbox_backend logger = logging.getLogger(__name__) @@ -147,7 +147,7 @@ I've {action} pull request to address this issue: repo_owner = repo_config.get("owner") repo_name = repo_config.get("name") - sandbox_backend = SANDBOX_BACKENDS.get(thread_id) + sandbox_backend = await get_sandbox_backend(thread_id) if thread_id else None repo_dir = f"/workspace/{repo_name}" diff --git a/apps/agent/agent/server.py b/apps/agent/agent/server.py index c44c76b4..b3d92cb7 100644 --- a/apps/agent/agent/server.py +++ b/apps/agent/agent/server.py @@ -223,13 +223,29 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915 tools=[], ).with_config(config) + sandbox_backend = SANDBOX_BACKENDS.get(thread_id) sandbox_id = await _get_sandbox_id_from_metadata(thread_id) - if sandbox_id == SANDBOX_CREATING: + if sandbox_id == SANDBOX_CREATING and not sandbox_backend: logger.info("Sandbox creation in progress, waiting...") sandbox_id = await _wait_for_sandbox_id(thread_id) - if sandbox_id is None: + if sandbox_backend: + logger.info("Using cached sandbox backend for thread %s", thread_id) + thread = await client.threads.get(thread_id=thread_id) + repo_dir = thread.get("metadata", {}).get("repo_dir") + + if repo_owner and repo_name: + logger.info("Pulling latest changes for repo %s/%s", repo_owner, repo_name) + try: + repo_dir = await _clone_or_pull_repo_in_sandbox( + sandbox_backend, repo_owner, repo_name, github_token + ) + except Exception: + logger.exception("Failed to pull repo in cached sandbox") + raise + + elif sandbox_id is None: logger.info("Creating new sandbox for thread %s", thread_id) await client.threads.update(thread_id=thread_id, metadata={"sandbox_id": SANDBOX_CREATING}) diff --git a/apps/agent/agent/tools/commit_and_open_pr.py b/apps/agent/agent/tools/commit_and_open_pr.py index 6a40ac73..2afc3615 100644 --- a/apps/agent/agent/tools/commit_and_open_pr.py +++ b/apps/agent/agent/tools/commit_and_open_pr.py @@ -5,7 +5,6 @@ from typing import Any from langgraph.config import get_config from ..encryption import decrypt_token -from ..integrations.langsmith import _create_langsmith_sandbox from ..utils.github import ( create_github_pr, get_github_default_branch, @@ -19,7 +18,7 @@ from ..utils.github import ( git_has_unpushed_commits, git_push, ) -from ..utils.sandbox_state import SANDBOX_BACKENDS +from ..utils.sandbox_state import get_sandbox_backend_sync logger = logging.getLogger(__name__) @@ -128,15 +127,9 @@ def commit_and_open_pr( "pr_url": None, } - sandbox_backend = SANDBOX_BACKENDS.get(thread_id) + sandbox_backend = get_sandbox_backend_sync(thread_id) if not sandbox_backend: - sandbox_id = configurable.get("sandbox_id") - - if not sandbox_id: - return {"success": False, "error": "No sandbox found for thread", "pr_url": None} - - sandbox_backend = _create_langsmith_sandbox(sandbox_id) - SANDBOX_BACKENDS[thread_id] = sandbox_backend + return {"success": False, "error": "No sandbox found for thread", "pr_url": None} repo_dir = f"/workspace/{repo_name}" diff --git a/apps/agent/agent/utils/sandbox_state.py b/apps/agent/agent/utils/sandbox_state.py index 5d8d2b9c..46adc855 100644 --- a/apps/agent/agent/utils/sandbox_state.py +++ b/apps/agent/agent/utils/sandbox_state.py @@ -2,7 +2,46 @@ from __future__ import annotations +import asyncio +import logging from typing import Any +from langgraph_sdk import get_client + +from ..integrations.langsmith import _create_langsmith_sandbox + +logger = logging.getLogger(__name__) +client = get_client() + # Thread ID -> SandboxBackend mapping, shared between server.py and middleware SANDBOX_BACKENDS: dict[str, Any] = {} + + +async def _get_sandbox_id_from_metadata(thread_id: str) -> str | None: + """Fetch sandbox_id from thread metadata.""" + try: + thread = await client.threads.get(thread_id=thread_id) + except Exception: + logger.exception("Failed to fetch thread metadata for sandbox") + return None + return thread.get("metadata", {}).get("sandbox_id") + + +async def get_sandbox_backend(thread_id: str) -> Any | None: + """Get sandbox backend from cache, or connect using thread metadata.""" + sandbox_backend = SANDBOX_BACKENDS.get(thread_id) + if sandbox_backend: + return sandbox_backend + + sandbox_id = await _get_sandbox_id_from_metadata(thread_id) + if not sandbox_id: + return None + + sandbox_backend = await asyncio.to_thread(_create_langsmith_sandbox, sandbox_id) + SANDBOX_BACKENDS[thread_id] = sandbox_backend + return sandbox_backend + + +def get_sandbox_backend_sync(thread_id: str) -> Any | None: + """Sync wrapper for get_sandbox_backend.""" + return asyncio.run(get_sandbox_backend(thread_id))