fix: keep sandbox backend stable across recovery (#1294)w

Use a per-thread proxy so in-flight tools continue through the latest recreated sandbox instead of holding a stale backend reference.
This commit is contained in:
Johannes du Plessis 2026-05-11 16:03:38 -07:00 • committed by GitHub
parent 0c51513862
commit 5c7c78406c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 181 additions and 28 deletions

View file

@ -115,10 +115,10 @@ def _get_thread_id(request: ToolCallRequest) -> str | None:
async def _recreate_sandbox_for_thread(thread_id: str) -> str:
from agent.server import _configure_git_identity, _recreate_sandbox, client
from agent.utils.sandbox_state import SANDBOX_BACKENDS
from agent.utils.sandbox_state import set_sandbox_backend
sandbox_backend = await _recreate_sandbox(thread_id)
SANDBOX_BACKENDS[thread_id] = sandbox_backend
sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend)
await client.threads.update(thread_id=thread_id, metadata={"sandbox_id": sandbox_backend.id})
await _configure_git_identity(sandbox_backend)
return sandbox_backend.id

View file

@ -68,17 +68,23 @@ SANDBOX_CREATING = "__creating__"
SANDBOX_CREATION_TIMEOUT = 180
SANDBOX_POLL_INTERVAL = 1.0
from .utils.sandbox_state import SANDBOX_BACKENDS, get_sandbox_id_from_metadata
from .utils.sandbox_state import (
SANDBOX_BACKENDS,
get_sandbox_id_from_metadata,
set_sandbox_backend,
unwrap_sandbox_backend,
)
async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProtocol) -> None:
"""Start a LangSmith sandbox before operations that require it to be running."""
if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith":
return
if not isinstance(sandbox_backend, LangSmithSandbox):
current_backend = unwrap_sandbox_backend(sandbox_backend)
if not isinstance(current_backend, LangSmithSandbox):
return
sandbox = sandbox_backend._sandbox # noqa: SLF001
sandbox = current_backend._sandbox # noqa: SLF001
status = await asyncio.to_thread(sandbox._client.get_sandbox_status, sandbox.name) # noqa: SLF001
status_name = getattr(status, "status", status)
status_name = getattr(status_name, "value", status_name)
@ -88,7 +94,7 @@ async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProt
logger.info(
"Starting LangSmith sandbox %s before proxy refresh (status=%s)",
sandbox_backend.id,
current_backend.id,
status_text or "unknown",
)
await asyncio.to_thread(sandbox.start)
@ -130,8 +136,9 @@ async def _refresh_github_proxy(
)
return
await _start_langsmith_sandbox_if_needed(sandbox_backend)
await asyncio.to_thread(_configure_github_proxy, sandbox_backend.id, installation_token)
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, installation_token)
async def _refresh_github_proxy_or_recreate(
@ -163,17 +170,16 @@ async def _configure_git_identity(sandbox_backend: SandboxBackendProtocol) -> No
async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol:
"""Recreate a sandbox after a connection failure.
Clears the stale cache entry, sets the SANDBOX_CREATING sentinel,
and creates a fresh sandbox (with proxy auth configured).
Sets the SANDBOX_CREATING sentinel and creates a fresh sandbox
(with proxy auth configured), swapping the per-thread proxy target.
The agent is responsible for cloning repos via tools.
"""
SANDBOX_BACKENDS.pop(thread_id, None)
await client.threads.update(
thread_id=thread_id,
metadata={"sandbox_id": SANDBOX_CREATING},
)
try:
sandbox_backend = await _create_sandbox_with_proxy()
sandbox_backend = set_sandbox_backend(thread_id, await _create_sandbox_with_proxy())
except Exception:
logger.exception("Failed to recreate sandbox after connection failure")
await client.threads.update(thread_id=thread_id, metadata={"sandbox_id": None})
@ -298,7 +304,7 @@ async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol:
sandbox_backend, thread_id
)
SANDBOX_BACKENDS[thread_id] = sandbox_backend
sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend)
if sandbox_id != sandbox_backend.id:
await client.threads.update(

View file

@ -1,5 +1,7 @@
import os
from deepagents.backends.protocol import SandboxBackendProtocol
from agent.integrations.daytona import create_daytona_sandbox
from agent.integrations.langsmith import create_langsmith_sandbox
from agent.integrations.local import create_local_sandbox
@ -15,7 +17,7 @@ SANDBOX_FACTORIES = {
}
def create_sandbox(sandbox_id: str | None = None):
def create_sandbox(sandbox_id: str | None = None) -> SandboxBackendProtocol:
"""Create or reconnect to a sandbox using the configured provider.
The provider is selected via the SANDBOX_TYPE environment variable.

View file

@ -4,16 +4,150 @@ from __future__ import annotations
import asyncio
import logging
from typing import Any
from deepagents.backends.protocol import (
EditResult,
ExecuteResponse,
FileDownloadResponse,
FileUploadResponse,
GlobResult,
GrepResult,
LsResult,
ReadResult,
SandboxBackendProtocol,
WriteResult,
)
from langgraph.config import get_config
from .sandbox import create_sandbox
logger = logging.getLogger(__name__)
# Thread ID -> SandboxBackend mapping, shared between server.py and middleware
SANDBOX_BACKENDS: dict[str, Any] = {}
class SandboxBackendProxy(SandboxBackendProtocol):
"""Stable per-thread backend handle whose target can be replaced."""
def __init__(self, backend: SandboxBackendProtocol) -> None:
self._backend = backend
@property
def current(self) -> SandboxBackendProtocol:
return self._backend
@property
def id(self) -> str:
return self._backend.id
def replace_backend(self, backend: SandboxBackendProtocol) -> None:
self._backend = backend
def ls(self, path: str) -> LsResult:
return self._backend.ls(path)
async def als(self, path: str) -> LsResult:
return await self._backend.als(path)
def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult:
return self._backend.read(file_path, offset, limit)
async def aread(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult:
return await self._backend.aread(file_path, offset, limit)
def grep(
self,
pattern: str,
path: str | None = None,
glob: str | None = None,
) -> GrepResult:
return self._backend.grep(pattern, path, glob)
async def agrep(
self,
pattern: str,
path: str | None = None,
glob: str | None = None,
) -> GrepResult:
return await self._backend.agrep(pattern, path, glob)
def glob(self, pattern: str, path: str = "/") -> GlobResult:
return self._backend.glob(pattern, path)
async def aglob(self, pattern: str, path: str = "/") -> GlobResult:
return await self._backend.aglob(pattern, path)
def write(self, file_path: str, content: str) -> WriteResult:
return self._backend.write(file_path, content)
async def awrite(self, file_path: str, content: str) -> WriteResult:
return await self._backend.awrite(file_path, content)
def edit(
self,
file_path: str,
old_string: str,
new_string: str,
replace_all: bool = False,
) -> EditResult:
return self._backend.edit(file_path, old_string, new_string, replace_all)
async def aedit(
self,
file_path: str,
old_string: str,
new_string: str,
replace_all: bool = False,
) -> EditResult:
return await self._backend.aedit(file_path, old_string, new_string, replace_all)
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
return self._backend.upload_files(files)
async def aupload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
return await self._backend.aupload_files(files)
def download_files(self, paths: list[str]) -> list[FileDownloadResponse]:
return self._backend.download_files(paths)
async def adownload_files(self, paths: list[str]) -> list[FileDownloadResponse]:
return await self._backend.adownload_files(paths)
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
return self._backend.execute(command, timeout=timeout)
async def aexecute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
return await self._backend.aexecute(command, timeout=timeout)
# Thread ID -> stable SandboxBackendProxy, shared between server.py and middleware.
SANDBOX_BACKENDS: dict[str, SandboxBackendProxy] = {}
def unwrap_sandbox_backend(sandbox_backend: SandboxBackendProtocol) -> SandboxBackendProtocol:
if isinstance(sandbox_backend, SandboxBackendProxy):
return sandbox_backend.current
return sandbox_backend
def set_sandbox_backend(
thread_id: str,
sandbox_backend: SandboxBackendProtocol,
) -> SandboxBackendProxy:
if isinstance(sandbox_backend, SandboxBackendProxy):
SANDBOX_BACKENDS[thread_id] = sandbox_backend
return sandbox_backend
existing = SANDBOX_BACKENDS.get(thread_id)
if isinstance(existing, SandboxBackendProxy):
existing.replace_backend(sandbox_backend)
return existing
proxy = SandboxBackendProxy(sandbox_backend)
SANDBOX_BACKENDS[thread_id] = proxy
return proxy
def clear_sandbox_backend(thread_id: str) -> None:
SANDBOX_BACKENDS.pop(thread_id, None)
async def get_sandbox_id_from_metadata(thread_id: str) -> str | None:
@ -23,10 +157,14 @@ async def get_sandbox_id_from_metadata(thread_id: str) -> str | None:
except Exception:
logger.exception("Failed to read thread metadata for sandbox")
return None
return config.get("metadata", {}).get("sandbox_id")
metadata = config.get("metadata", {})
if not isinstance(metadata, dict):
return None
sandbox_id = metadata.get("sandbox_id")
return sandbox_id if isinstance(sandbox_id, str) else None
async def get_sandbox_backend(thread_id: str) -> Any | None:
async def get_sandbox_backend(thread_id: str) -> SandboxBackendProxy:
"""Get sandbox backend from cache, or connect using thread metadata."""
sandbox_backend = SANDBOX_BACKENDS.get(thread_id)
if sandbox_backend:
@ -37,10 +175,9 @@ async def get_sandbox_backend(thread_id: str) -> Any | None:
raise ValueError(f"Missing sandbox_id in thread metadata for {thread_id}")
sandbox_backend = await asyncio.to_thread(create_sandbox, sandbox_id)
SANDBOX_BACKENDS[thread_id] = sandbox_backend
return sandbox_backend
return set_sandbox_backend(thread_id, sandbox_backend)
def get_sandbox_backend_sync(thread_id: str) -> Any | None:
def get_sandbox_backend_sync(thread_id: str) -> SandboxBackendProxy:
"""Sync wrapper for get_sandbox_backend."""
return asyncio.run(get_sandbox_backend(thread_id))

View file

@ -2,6 +2,7 @@ import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from deepagents.backends.protocol import ExecuteResponse, SandboxBackendProtocol
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langgraph.prebuilt.tool_node import ToolCallRequest
from langsmith.sandbox import SandboxClientError
@ -11,14 +12,14 @@ from agent.middleware.sandbox_circuit_breaker import (
SandboxCircuitBreakerMiddleware,
)
from agent.middleware.tool_error_handler import ToolErrorMiddleware
from agent.utils.sandbox_state import SANDBOX_BACKENDS
from agent.utils.sandbox_state import SANDBOX_BACKENDS, clear_sandbox_backend, set_sandbox_backend
class FakeSandboxBackend:
class FakeSandboxBackend(SandboxBackendProtocol):
id = "sb-new"
def execute(self, _command: str) -> None:
return None
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
return ExecuteResponse(output=f"{self.id}: {command}: {timeout}", exit_code=0)
def _tool_request(thread_id: str = "thread-1") -> ToolCallRequest:
@ -49,7 +50,11 @@ def _sandbox_error_message(tool_call_id: str, sandbox_id: str = "sb-dead") -> To
async def test_sandbox_client_error_recreates_sandbox() -> None:
middleware = ToolErrorMiddleware()
request = _tool_request()
old_backend = FakeSandboxBackend()
backend = FakeSandboxBackend()
old_backend.id = "sb-old"
backend.id = "sb-new"
proxy = set_sandbox_backend("thread-1", old_backend)
async def handler(_request: ToolCallRequest) -> ToolMessage:
raise SandboxClientError("Sandbox request timed out: sb-dead")
@ -70,7 +75,10 @@ async def test_sandbox_client_error_recreates_sandbox() -> None:
thread_id="thread-1",
metadata={"sandbox_id": "sb-new"},
)
assert SANDBOX_BACKENDS["thread-1"] is backend
assert SANDBOX_BACKENDS["thread-1"] is proxy
assert proxy.current is backend
assert proxy.id == "sb-new"
assert proxy.execute("echo ok").output == "sb-new: echo ok: None"
payload = json.loads(result.content)
assert payload["status"] == "error"
@ -79,7 +87,7 @@ async def test_sandbox_client_error_recreates_sandbox() -> None:
assert payload["previous_error"] == "Sandbox request timed out: sb-dead"
assert "sb-new" in payload["error"]
finally:
SANDBOX_BACKENDS.pop("thread-1", None)
clear_sandbox_backend("thread-1")
def test_repeated_sandbox_errors_trigger_circuit_breaker_once() -> None: