mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 13:42:16 +00:00
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:
parent
0c51513862
commit
5c7c78406c
5 changed files with 181 additions and 28 deletions
|
|
@ -115,10 +115,10 @@ def _get_thread_id(request: ToolCallRequest) -> str | None:
|
||||||
|
|
||||||
async def _recreate_sandbox_for_thread(thread_id: str) -> str:
|
async def _recreate_sandbox_for_thread(thread_id: str) -> str:
|
||||||
from agent.server import _configure_git_identity, _recreate_sandbox, client
|
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_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 client.threads.update(thread_id=thread_id, metadata={"sandbox_id": sandbox_backend.id})
|
||||||
await _configure_git_identity(sandbox_backend)
|
await _configure_git_identity(sandbox_backend)
|
||||||
return sandbox_backend.id
|
return sandbox_backend.id
|
||||||
|
|
|
||||||
|
|
@ -68,17 +68,23 @@ SANDBOX_CREATING = "__creating__"
|
||||||
SANDBOX_CREATION_TIMEOUT = 180
|
SANDBOX_CREATION_TIMEOUT = 180
|
||||||
SANDBOX_POLL_INTERVAL = 1.0
|
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:
|
async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProtocol) -> None:
|
||||||
"""Start a LangSmith sandbox before operations that require it to be running."""
|
"""Start a LangSmith sandbox before operations that require it to be running."""
|
||||||
if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith":
|
if os.getenv("SANDBOX_TYPE", "langsmith") != "langsmith":
|
||||||
return
|
return
|
||||||
if not isinstance(sandbox_backend, LangSmithSandbox):
|
current_backend = unwrap_sandbox_backend(sandbox_backend)
|
||||||
|
if not isinstance(current_backend, LangSmithSandbox):
|
||||||
return
|
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 = await asyncio.to_thread(sandbox._client.get_sandbox_status, sandbox.name) # noqa: SLF001
|
||||||
status_name = getattr(status, "status", status)
|
status_name = getattr(status, "status", status)
|
||||||
status_name = getattr(status_name, "value", status_name)
|
status_name = getattr(status_name, "value", status_name)
|
||||||
|
|
@ -88,7 +94,7 @@ async def _start_langsmith_sandbox_if_needed(sandbox_backend: SandboxBackendProt
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Starting LangSmith sandbox %s before proxy refresh (status=%s)",
|
"Starting LangSmith sandbox %s before proxy refresh (status=%s)",
|
||||||
sandbox_backend.id,
|
current_backend.id,
|
||||||
status_text or "unknown",
|
status_text or "unknown",
|
||||||
)
|
)
|
||||||
await asyncio.to_thread(sandbox.start)
|
await asyncio.to_thread(sandbox.start)
|
||||||
|
|
@ -130,8 +136,9 @@ async def _refresh_github_proxy(
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
await _start_langsmith_sandbox_if_needed(sandbox_backend)
|
current_backend = unwrap_sandbox_backend(sandbox_backend)
|
||||||
await asyncio.to_thread(_configure_github_proxy, sandbox_backend.id, installation_token)
|
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(
|
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:
|
async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol:
|
||||||
"""Recreate a sandbox after a connection failure.
|
"""Recreate a sandbox after a connection failure.
|
||||||
|
|
||||||
Clears the stale cache entry, sets the SANDBOX_CREATING sentinel,
|
Sets the SANDBOX_CREATING sentinel and creates a fresh sandbox
|
||||||
and creates a fresh sandbox (with proxy auth configured).
|
(with proxy auth configured), swapping the per-thread proxy target.
|
||||||
The agent is responsible for cloning repos via tools.
|
The agent is responsible for cloning repos via tools.
|
||||||
"""
|
"""
|
||||||
SANDBOX_BACKENDS.pop(thread_id, None)
|
|
||||||
await client.threads.update(
|
await client.threads.update(
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
metadata={"sandbox_id": SANDBOX_CREATING},
|
metadata={"sandbox_id": SANDBOX_CREATING},
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
sandbox_backend = await _create_sandbox_with_proxy()
|
sandbox_backend = set_sandbox_backend(thread_id, await _create_sandbox_with_proxy())
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to recreate sandbox after connection failure")
|
logger.exception("Failed to recreate sandbox after connection failure")
|
||||||
await client.threads.update(thread_id=thread_id, metadata={"sandbox_id": None})
|
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_backend, thread_id
|
||||||
)
|
)
|
||||||
|
|
||||||
SANDBOX_BACKENDS[thread_id] = sandbox_backend
|
sandbox_backend = set_sandbox_backend(thread_id, sandbox_backend)
|
||||||
|
|
||||||
if sandbox_id != sandbox_backend.id:
|
if sandbox_id != sandbox_backend.id:
|
||||||
await client.threads.update(
|
await client.threads.update(
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,7 @@
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
from deepagents.backends.protocol import SandboxBackendProtocol
|
||||||
|
|
||||||
from agent.integrations.daytona import create_daytona_sandbox
|
from agent.integrations.daytona import create_daytona_sandbox
|
||||||
from agent.integrations.langsmith import create_langsmith_sandbox
|
from agent.integrations.langsmith import create_langsmith_sandbox
|
||||||
from agent.integrations.local import create_local_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.
|
"""Create or reconnect to a sandbox using the configured provider.
|
||||||
|
|
||||||
The provider is selected via the SANDBOX_TYPE environment variable.
|
The provider is selected via the SANDBOX_TYPE environment variable.
|
||||||
|
|
|
||||||
|
|
@ -4,16 +4,150 @@ from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
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 langgraph.config import get_config
|
||||||
|
|
||||||
from .sandbox import create_sandbox
|
from .sandbox import create_sandbox
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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:
|
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:
|
except Exception:
|
||||||
logger.exception("Failed to read thread metadata for sandbox")
|
logger.exception("Failed to read thread metadata for sandbox")
|
||||||
return None
|
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."""
|
"""Get sandbox backend from cache, or connect using thread metadata."""
|
||||||
sandbox_backend = SANDBOX_BACKENDS.get(thread_id)
|
sandbox_backend = SANDBOX_BACKENDS.get(thread_id)
|
||||||
if sandbox_backend:
|
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}")
|
raise ValueError(f"Missing sandbox_id in thread metadata for {thread_id}")
|
||||||
|
|
||||||
sandbox_backend = await asyncio.to_thread(create_sandbox, sandbox_id)
|
sandbox_backend = await asyncio.to_thread(create_sandbox, sandbox_id)
|
||||||
SANDBOX_BACKENDS[thread_id] = sandbox_backend
|
return set_sandbox_backend(thread_id, sandbox_backend)
|
||||||
return 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."""
|
"""Sync wrapper for get_sandbox_backend."""
|
||||||
return asyncio.run(get_sandbox_backend(thread_id))
|
return asyncio.run(get_sandbox_backend(thread_id))
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ import json
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from deepagents.backends.protocol import ExecuteResponse, SandboxBackendProtocol
|
||||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||||
from langsmith.sandbox import SandboxClientError
|
from langsmith.sandbox import SandboxClientError
|
||||||
|
|
@ -11,14 +12,14 @@ from agent.middleware.sandbox_circuit_breaker import (
|
||||||
SandboxCircuitBreakerMiddleware,
|
SandboxCircuitBreakerMiddleware,
|
||||||
)
|
)
|
||||||
from agent.middleware.tool_error_handler import ToolErrorMiddleware
|
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"
|
id = "sb-new"
|
||||||
|
|
||||||
def execute(self, _command: str) -> None:
|
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||||
return None
|
return ExecuteResponse(output=f"{self.id}: {command}: {timeout}", exit_code=0)
|
||||||
|
|
||||||
|
|
||||||
def _tool_request(thread_id: str = "thread-1") -> ToolCallRequest:
|
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:
|
async def test_sandbox_client_error_recreates_sandbox() -> None:
|
||||||
middleware = ToolErrorMiddleware()
|
middleware = ToolErrorMiddleware()
|
||||||
request = _tool_request()
|
request = _tool_request()
|
||||||
|
old_backend = FakeSandboxBackend()
|
||||||
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:
|
async def handler(_request: ToolCallRequest) -> ToolMessage:
|
||||||
raise SandboxClientError("Sandbox request timed out: sb-dead")
|
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",
|
thread_id="thread-1",
|
||||||
metadata={"sandbox_id": "sb-new"},
|
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)
|
payload = json.loads(result.content)
|
||||||
assert payload["status"] == "error"
|
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 payload["previous_error"] == "Sandbox request timed out: sb-dead"
|
||||||
assert "sb-new" in payload["error"]
|
assert "sb-new" in payload["error"]
|
||||||
finally:
|
finally:
|
||||||
SANDBOX_BACKENDS.pop("thread-1", None)
|
clear_sandbox_backend("thread-1")
|
||||||
|
|
||||||
|
|
||||||
def test_repeated_sandbox_errors_trigger_circuit_breaker_once() -> None:
|
def test_repeated_sandbox_errors_trigger_circuit_breaker_once() -> None:
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue