mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 08:03:15 +00:00
fix: recover from mid-run sandbox death (#1274)
* fix: recover from mid-run sandbox death Recreate dead sandboxes during tool execution and stop repeated unrecoverable timeout loops with a user-facing notification. * fix: count repeated sandbox recreations Treat consecutive sandbox recreations as an unrecovered failure streak so outages cannot loop until the model-call limit.
This commit is contained in:
parent
dc9a0b98da
commit
743b2b9ba4
5 changed files with 611 additions and 19 deletions
|
|
@ -3,6 +3,7 @@ from .ensure_no_empty_msg import ensure_no_empty_msg
|
|||
from .exclude_tools import ExcludeToolsMiddleware
|
||||
from .notify_step_limit import notify_step_limit_reached
|
||||
from .refresh_slack_status import SlackAssistantStatusMiddleware
|
||||
from .sandbox_circuit_breaker import SandboxCircuitBreakerMiddleware
|
||||
from .sanitize_tool_inputs import SanitizeToolInputsMiddleware
|
||||
from .tool_error_handler import ToolErrorMiddleware
|
||||
|
||||
|
|
@ -10,6 +11,7 @@ __all__ = [
|
|||
"ExcludeToolsMiddleware",
|
||||
"SanitizeToolInputsMiddleware",
|
||||
"ToolErrorMiddleware",
|
||||
"SandboxCircuitBreakerMiddleware",
|
||||
"SlackAssistantStatusMiddleware",
|
||||
"check_message_queue_before_model",
|
||||
"ensure_no_empty_msg",
|
||||
|
|
|
|||
261
agent/middleware/sandbox_circuit_breaker.py
Normal file
261
agent/middleware/sandbox_circuit_breaker.py
Normal file
|
|
@ -0,0 +1,261 @@
|
|||
"""Circuit breaker for repeated unrecoverable sandbox failures."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
from langchain.agents.middleware import AgentMiddleware, AgentState, hook_config
|
||||
from langchain_core.messages import AIMessage, BaseMessage, ToolMessage
|
||||
from langgraph.config import get_config
|
||||
from langgraph.runtime import Runtime
|
||||
|
||||
from ..utils.github_app import get_github_app_installation_token
|
||||
from ..utils.github_comments import post_github_comment
|
||||
from ..utils.github_token import get_github_token
|
||||
from ..utils.linear import comment_on_linear_issue
|
||||
from ..utils.slack import post_slack_thread_reply
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SANDBOX_CIRCUIT_BREAKER_THRESHOLD = 2
|
||||
SANDBOX_UNRECOVERABLE_MESSAGE = "Sandbox became unrecoverable mid-task. Please retrigger."
|
||||
|
||||
_CIRCUIT_BREAKER_MARKER = "Sandbox circuit breaker triggered"
|
||||
_SANDBOX_RECREATED_AFTER_CLIENT_ERROR = "sandbox_recreated_after_client_error"
|
||||
_SANDBOX_ID_RE = re.compile(r"\bsb-[A-Za-z0-9-]+\b")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SandboxErrorStreak:
|
||||
reason: Literal["client_error", "recreated"]
|
||||
sandbox_id: str | None
|
||||
count: int
|
||||
|
||||
|
||||
def _content_to_text(content: object) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if not isinstance(content, list):
|
||||
return str(content)
|
||||
|
||||
parts: list[str] = []
|
||||
for block in content:
|
||||
if isinstance(block, Mapping):
|
||||
text = block.get("text", "")
|
||||
parts.append(text if isinstance(text, str) else str(text))
|
||||
else:
|
||||
parts.append(str(block))
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _extract_sandbox_id(text: str) -> str | None:
|
||||
match = _SANDBOX_ID_RE.search(text)
|
||||
return match.group(0) if match else None
|
||||
|
||||
|
||||
def _last_message_has_circuit_breaker_marker(messages: Sequence[BaseMessage]) -> bool:
|
||||
if not messages:
|
||||
return False
|
||||
content = _content_to_text(getattr(messages[-1], "content", "") or "")
|
||||
return _CIRCUIT_BREAKER_MARKER in content
|
||||
|
||||
|
||||
def _sandbox_error_streak(messages: Sequence[BaseMessage]) -> SandboxErrorStreak | None:
|
||||
sandbox_id: str | None = None
|
||||
reason: Literal["client_error", "recreated"] | None = None
|
||||
count = 0
|
||||
|
||||
for message in reversed(messages):
|
||||
if isinstance(message, ToolMessage):
|
||||
text = _content_to_text(message.content)
|
||||
if _SANDBOX_RECREATED_AFTER_CLIENT_ERROR in text:
|
||||
if reason is None:
|
||||
reason = "recreated"
|
||||
elif reason != "recreated":
|
||||
break
|
||||
count += 1
|
||||
continue
|
||||
|
||||
message_sandbox_id = _extract_sandbox_id(text)
|
||||
if "SandboxClientError" not in text or message_sandbox_id is None:
|
||||
break
|
||||
if reason is None:
|
||||
reason = "client_error"
|
||||
sandbox_id = message_sandbox_id
|
||||
elif reason != "client_error" or message_sandbox_id != sandbox_id:
|
||||
break
|
||||
count += 1
|
||||
continue
|
||||
|
||||
text = _content_to_text(getattr(message, "content", "") or "")
|
||||
if _CIRCUIT_BREAKER_MARKER in text:
|
||||
return None
|
||||
if getattr(message, "type", "") in {"human", "system"}:
|
||||
break
|
||||
|
||||
if reason is None:
|
||||
return None
|
||||
return SandboxErrorStreak(reason=reason, sandbox_id=sandbox_id, count=count)
|
||||
|
||||
|
||||
def _get_slack_target(configurable: Mapping[str, Any]) -> tuple[str, str] | None:
|
||||
slack_thread = configurable.get("slack_thread")
|
||||
if not isinstance(slack_thread, Mapping):
|
||||
return None
|
||||
channel_id = slack_thread.get("channel_id")
|
||||
thread_ts = slack_thread.get("thread_ts")
|
||||
if not isinstance(channel_id, str) or not isinstance(thread_ts, str):
|
||||
return None
|
||||
if not channel_id or not thread_ts:
|
||||
return None
|
||||
return channel_id, thread_ts
|
||||
|
||||
|
||||
def _get_linear_issue_id(configurable: Mapping[str, Any]) -> str | None:
|
||||
linear_issue = configurable.get("linear_issue")
|
||||
if not isinstance(linear_issue, Mapping):
|
||||
return None
|
||||
issue_id = linear_issue.get("id")
|
||||
return issue_id if isinstance(issue_id, str) and issue_id else None
|
||||
|
||||
|
||||
def _coerce_issue_number(value: object) -> int | None:
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, str) and value.isdigit():
|
||||
return int(value)
|
||||
return None
|
||||
|
||||
|
||||
def _get_github_target(configurable: Mapping[str, Any]) -> tuple[dict[str, str], int] | None:
|
||||
repo_config = configurable.get("repo")
|
||||
if not isinstance(repo_config, Mapping):
|
||||
return None
|
||||
owner = repo_config.get("owner")
|
||||
name = repo_config.get("name")
|
||||
if not isinstance(owner, str) or not isinstance(name, str) or not owner or not name:
|
||||
return None
|
||||
repo = {"owner": owner, "name": name}
|
||||
|
||||
github_pr_or_issue = configurable.get("github_pr_or_issue")
|
||||
if isinstance(github_pr_or_issue, Mapping):
|
||||
number = _coerce_issue_number(github_pr_or_issue.get("number"))
|
||||
target_repo = github_pr_or_issue.get("repo")
|
||||
if isinstance(target_repo, Mapping):
|
||||
target_owner = target_repo.get("owner")
|
||||
target_name = target_repo.get("name")
|
||||
if isinstance(target_owner, str) and isinstance(target_name, str):
|
||||
repo = {"owner": target_owner, "name": target_name}
|
||||
if number is not None:
|
||||
return repo, number
|
||||
|
||||
github_issue = configurable.get("github_issue")
|
||||
if isinstance(github_issue, Mapping):
|
||||
number = _coerce_issue_number(github_issue.get("number"))
|
||||
if number is not None:
|
||||
return repo, number
|
||||
|
||||
pr_number = _coerce_issue_number(configurable.get("pr_number"))
|
||||
if pr_number is not None:
|
||||
return repo, pr_number
|
||||
return None
|
||||
|
||||
|
||||
async def _post_unrecoverable_notification(config: Mapping[str, Any]) -> None:
|
||||
configurable = config.get("configurable", {})
|
||||
if not isinstance(configurable, Mapping):
|
||||
logger.info("No runtime configurable found for sandbox circuit breaker notification")
|
||||
return
|
||||
|
||||
slack_target = _get_slack_target(configurable)
|
||||
if slack_target is not None:
|
||||
channel_id, thread_ts = slack_target
|
||||
await post_slack_thread_reply(channel_id, thread_ts, SANDBOX_UNRECOVERABLE_MESSAGE)
|
||||
logger.info("Sent sandbox circuit breaker notification to Slack thread %s", thread_ts)
|
||||
return
|
||||
|
||||
linear_issue_id = _get_linear_issue_id(configurable)
|
||||
if linear_issue_id is not None:
|
||||
await comment_on_linear_issue(linear_issue_id, SANDBOX_UNRECOVERABLE_MESSAGE)
|
||||
logger.info("Sent sandbox circuit breaker notification to Linear issue %s", linear_issue_id)
|
||||
return
|
||||
|
||||
github_target = _get_github_target(configurable)
|
||||
if github_target is not None:
|
||||
token = get_github_token(config) or await get_github_app_installation_token()
|
||||
if not token:
|
||||
logger.info("No GitHub token available for sandbox circuit breaker notification")
|
||||
return
|
||||
repo, issue_number = github_target
|
||||
await post_github_comment(
|
||||
repo,
|
||||
issue_number,
|
||||
SANDBOX_UNRECOVERABLE_MESSAGE,
|
||||
token=token,
|
||||
)
|
||||
logger.info("Sent sandbox circuit breaker notification to GitHub item #%s", issue_number)
|
||||
return
|
||||
|
||||
logger.info("No user-facing target found for sandbox circuit breaker notification")
|
||||
|
||||
|
||||
class SandboxCircuitBreakerMiddleware(AgentMiddleware[AgentState, Any]):
|
||||
"""Stop runs that repeatedly hit the same dead sandbox."""
|
||||
|
||||
state_schema = AgentState
|
||||
|
||||
def __init__(self, *, threshold: int = SANDBOX_CIRCUIT_BREAKER_THRESHOLD) -> None:
|
||||
self.threshold = threshold
|
||||
|
||||
@hook_config(can_jump_to=["end"])
|
||||
def before_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: # noqa: ARG002
|
||||
messages = state.get("messages", [])
|
||||
if _last_message_has_circuit_breaker_marker(messages):
|
||||
return None
|
||||
|
||||
streak = _sandbox_error_streak(messages)
|
||||
if streak is None or streak.count <= self.threshold:
|
||||
return None
|
||||
|
||||
if streak.reason == "recreated":
|
||||
detail = (
|
||||
f"{streak.count} consecutive sandbox recreations did not recover tool execution"
|
||||
)
|
||||
else:
|
||||
detail = f"{streak.count} consecutive sandbox tool failures against {streak.sandbox_id}"
|
||||
content = f"{_CIRCUIT_BREAKER_MARKER}: {detail}. {SANDBOX_UNRECOVERABLE_MESSAGE}"
|
||||
return {"jump_to": "end", "messages": [AIMessage(content=content)]}
|
||||
|
||||
@hook_config(can_jump_to=["end"])
|
||||
async def abefore_model(
|
||||
self,
|
||||
state: AgentState,
|
||||
runtime: Runtime,
|
||||
) -> dict[str, Any] | None:
|
||||
return self.before_model(state, runtime)
|
||||
|
||||
async def aafter_agent(
|
||||
self,
|
||||
state: AgentState,
|
||||
runtime: Runtime, # noqa: ARG002
|
||||
) -> dict[str, Any] | None:
|
||||
messages = state.get("messages", [])
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
last_msg = messages[-1]
|
||||
content = _content_to_text(getattr(last_msg, "content", "") or "")
|
||||
if _CIRCUIT_BREAKER_MARKER not in content:
|
||||
return None
|
||||
|
||||
try:
|
||||
config = get_config()
|
||||
await _post_unrecoverable_notification(config)
|
||||
except Exception:
|
||||
logger.exception("Failed to send sandbox circuit breaker notification")
|
||||
|
||||
return None
|
||||
|
|
@ -6,20 +6,26 @@ returned as error ToolMessages instead of crashing the agent run.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
AgentState,
|
||||
)
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.config import get_config
|
||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||
from langgraph.types import Command
|
||||
from langsmith.sandbox import SandboxClientError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SANDBOX_RECREATED_AFTER_CLIENT_ERROR = "sandbox_recreated_after_client_error"
|
||||
|
||||
|
||||
def _get_name(candidate: object) -> str | None:
|
||||
if not candidate:
|
||||
|
|
@ -55,12 +61,101 @@ def _to_error_payload(e: Exception, request: ToolCallRequest | None = None) -> d
|
|||
return data
|
||||
|
||||
|
||||
def _to_sandbox_recreated_payload(
|
||||
e: SandboxClientError,
|
||||
sandbox_id: str,
|
||||
request: ToolCallRequest | None = None,
|
||||
) -> dict[str, str]:
|
||||
data: dict[str, str] = {
|
||||
"status": "error",
|
||||
"error_type": e.__class__.__name__,
|
||||
"previous_error": str(e),
|
||||
"recovery": SANDBOX_RECREATED_AFTER_CLIENT_ERROR,
|
||||
"sandbox_id": sandbox_id,
|
||||
"error": (
|
||||
"The previous sandbox became unreachable mid-run. A fresh sandbox "
|
||||
f"({sandbox_id}) has been created and cached for this thread. "
|
||||
"Retry the last tool call; if repository files are missing, re-clone or "
|
||||
"reinitialize the workspace first."
|
||||
),
|
||||
}
|
||||
tool_name = _extract_tool_name(request)
|
||||
if tool_name:
|
||||
data["name"] = tool_name
|
||||
return data
|
||||
|
||||
|
||||
def _get_tool_call_id(request: ToolCallRequest) -> str | None:
|
||||
if isinstance(request.tool_call, dict):
|
||||
return request.tool_call.get("id")
|
||||
return None
|
||||
|
||||
|
||||
def _get_thread_id(request: ToolCallRequest) -> str | None:
|
||||
runtime_config = getattr(getattr(request, "runtime", None), "config", None)
|
||||
config: Mapping[str, Any] | None = (
|
||||
runtime_config if isinstance(runtime_config, Mapping) else None
|
||||
)
|
||||
if config is None:
|
||||
try:
|
||||
maybe_config = get_config()
|
||||
except Exception:
|
||||
logger.exception("Failed to read runnable config while handling sandbox error")
|
||||
return None
|
||||
config = maybe_config if isinstance(maybe_config, Mapping) else None
|
||||
if config is None:
|
||||
return None
|
||||
|
||||
configurable = config.get("configurable", {})
|
||||
if not isinstance(configurable, Mapping):
|
||||
return None
|
||||
thread_id = configurable.get("thread_id")
|
||||
return thread_id if isinstance(thread_id, str) and thread_id else 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
|
||||
|
||||
sandbox_backend = await _recreate_sandbox(thread_id)
|
||||
SANDBOX_BACKENDS[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
|
||||
|
||||
|
||||
def _recreate_sandbox_for_thread_sync(thread_id: str) -> str:
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(_recreate_sandbox_for_thread(thread_id))
|
||||
raise RuntimeError(
|
||||
"Cannot recreate sandbox from a sync tool call while an event loop is running"
|
||||
)
|
||||
|
||||
|
||||
def _sandbox_recreated_tool_message(
|
||||
e: SandboxClientError,
|
||||
sandbox_id: str,
|
||||
request: ToolCallRequest,
|
||||
) -> ToolMessage:
|
||||
data = _to_sandbox_recreated_payload(e, sandbox_id, request)
|
||||
return ToolMessage(
|
||||
content=json.dumps(data),
|
||||
tool_call_id=_get_tool_call_id(request),
|
||||
status="error",
|
||||
)
|
||||
|
||||
|
||||
def _generic_error_tool_message(e: Exception, request: ToolCallRequest) -> ToolMessage:
|
||||
data = _to_error_payload(e, request)
|
||||
return ToolMessage(
|
||||
content=json.dumps(data),
|
||||
tool_call_id=_get_tool_call_id(request),
|
||||
status="error",
|
||||
)
|
||||
|
||||
|
||||
class ToolErrorMiddleware(AgentMiddleware):
|
||||
"""Normalize tool execution errors into predictable payloads.
|
||||
|
||||
|
|
@ -78,14 +173,19 @@ class ToolErrorMiddleware(AgentMiddleware):
|
|||
) -> ToolMessage | Command:
|
||||
try:
|
||||
return handler(request)
|
||||
except SandboxClientError as e:
|
||||
logger.exception("Sandbox error during tool call handling; request=%r", request)
|
||||
thread_id = _get_thread_id(request)
|
||||
if thread_id:
|
||||
try:
|
||||
sandbox_id = _recreate_sandbox_for_thread_sync(thread_id)
|
||||
return _sandbox_recreated_tool_message(e, sandbox_id, request)
|
||||
except Exception:
|
||||
logger.exception("Failed to recreate sandbox for thread %s", thread_id)
|
||||
return _generic_error_tool_message(e, request)
|
||||
except Exception as e:
|
||||
logger.exception("Error during tool call handling; request=%r", request)
|
||||
data = _to_error_payload(e, request)
|
||||
return ToolMessage(
|
||||
content=json.dumps(data),
|
||||
tool_call_id=_get_tool_call_id(request),
|
||||
status="error",
|
||||
)
|
||||
return _generic_error_tool_message(e, request)
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
|
|
@ -94,11 +194,16 @@ class ToolErrorMiddleware(AgentMiddleware):
|
|||
) -> ToolMessage | Command:
|
||||
try:
|
||||
return await handler(request)
|
||||
except SandboxClientError as e:
|
||||
logger.exception("Sandbox error during tool call handling; request=%r", request)
|
||||
thread_id = _get_thread_id(request)
|
||||
if thread_id:
|
||||
try:
|
||||
sandbox_id = await _recreate_sandbox_for_thread(thread_id)
|
||||
return _sandbox_recreated_tool_message(e, sandbox_id, request)
|
||||
except Exception:
|
||||
logger.exception("Failed to recreate sandbox for thread %s", thread_id)
|
||||
return _generic_error_tool_message(e, request)
|
||||
except Exception as e:
|
||||
logger.exception("Error during tool call handling; request=%r", request)
|
||||
data = _to_error_payload(e, request)
|
||||
return ToolMessage(
|
||||
content=json.dumps(data),
|
||||
tool_call_id=_get_tool_call_id(request),
|
||||
status="error",
|
||||
)
|
||||
return _generic_error_tool_message(e, request)
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from langsmith.sandbox import SandboxClientError
|
|||
|
||||
from .integrations.langsmith import _configure_github_proxy
|
||||
from .middleware import (
|
||||
SandboxCircuitBreakerMiddleware,
|
||||
SanitizeToolInputsMiddleware,
|
||||
SlackAssistantStatusMiddleware,
|
||||
ToolErrorMiddleware,
|
||||
|
|
@ -149,6 +150,14 @@ async def _refresh_github_proxy_or_recreate(
|
|||
return sandbox_backend
|
||||
|
||||
|
||||
async def _configure_git_identity(sandbox_backend: SandboxBackendProtocol) -> None:
|
||||
await asyncio.to_thread(
|
||||
sandbox_backend.execute,
|
||||
"git config --global user.name 'open-swe[bot]' && "
|
||||
"git config --global user.email 'open-swe@users.noreply.github.com'",
|
||||
)
|
||||
|
||||
|
||||
async def _recreate_sandbox(thread_id: str) -> SandboxBackendProtocol:
|
||||
"""Recreate a sandbox after a connection failure.
|
||||
|
||||
|
|
@ -298,11 +307,7 @@ async def ensure_sandbox_for_thread(thread_id: str) -> SandboxBackendProtocol:
|
|||
# lost their `--global` config (or had it overwritten), and Vercel preview
|
||||
# deploys reject commits whose author email can't be resolved to a GitHub
|
||||
# account.
|
||||
await asyncio.to_thread(
|
||||
sandbox_backend.execute,
|
||||
"git config --global user.name 'open-swe[bot]' && "
|
||||
"git config --global user.email 'open-swe@users.noreply.github.com'",
|
||||
)
|
||||
await _configure_git_identity(sandbox_backend)
|
||||
|
||||
return sandbox_backend
|
||||
|
||||
|
|
@ -314,6 +319,13 @@ DEFAULT_RECURSION_LIMIT = 9_999
|
|||
MODEL_CALL_RECURSION_LIMIT = 5_000 # ~half the recursion limit to account for tool calls
|
||||
|
||||
|
||||
def _get_cached_sandbox_backend(thread_id: str) -> SandboxBackendProtocol:
|
||||
sandbox_backend = SANDBOX_BACKENDS.get(thread_id)
|
||||
if sandbox_backend is None:
|
||||
raise RuntimeError(f"No sandbox backend cached for thread {thread_id}")
|
||||
return sandbox_backend
|
||||
|
||||
|
||||
async def get_agent(config: RunnableConfig) -> Pregel:
|
||||
"""Get or create an agent with a sandbox for the given thread."""
|
||||
thread_id = config["configurable"].get("thread_id", None)
|
||||
|
|
@ -342,6 +354,9 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
|||
|
||||
work_dir = await aresolve_sandbox_work_dir(sandbox_backend)
|
||||
|
||||
def backend_factory(_runtime: object, _thread_id: str = thread_id) -> SandboxBackendProtocol:
|
||||
return _get_cached_sandbox_backend(_thread_id)
|
||||
|
||||
model_id = os.environ.get("LLM_MODEL_ID", DEFAULT_LLM_MODEL_ID)
|
||||
model_kwargs: ModelKwargs = {"max_tokens": DEFAULT_LLM_MAX_TOKENS}
|
||||
if model_id == DEFAULT_LLM_MODEL_ID:
|
||||
|
|
@ -371,7 +386,7 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
|||
slack_read_thread_messages,
|
||||
slack_thread_reply,
|
||||
],
|
||||
backend=sandbox_backend,
|
||||
backend=backend_factory,
|
||||
middleware=[
|
||||
SanitizeToolInputsMiddleware(),
|
||||
ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"),
|
||||
|
|
@ -380,5 +395,6 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
|||
SlackAssistantStatusMiddleware(),
|
||||
ensure_no_empty_msg,
|
||||
notify_step_limit_reached,
|
||||
SandboxCircuitBreakerMiddleware(),
|
||||
],
|
||||
).with_config(config)
|
||||
|
|
|
|||
208
tests/middleware/test_sandbox_recovery.py
Normal file
208
tests/middleware/test_sandbox_recovery.py
Normal file
|
|
@ -0,0 +1,208 @@
|
|||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||
from langsmith.sandbox import SandboxClientError
|
||||
|
||||
from agent.middleware.sandbox_circuit_breaker import (
|
||||
SANDBOX_UNRECOVERABLE_MESSAGE,
|
||||
SandboxCircuitBreakerMiddleware,
|
||||
)
|
||||
from agent.middleware.tool_error_handler import ToolErrorMiddleware
|
||||
from agent.utils.sandbox_state import SANDBOX_BACKENDS
|
||||
|
||||
|
||||
class FakeSandboxBackend:
|
||||
id = "sb-new"
|
||||
|
||||
def execute(self, _command: str) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _tool_request(thread_id: str = "thread-1") -> ToolCallRequest:
|
||||
runtime = MagicMock(config={"configurable": {"thread_id": thread_id}})
|
||||
return ToolCallRequest(
|
||||
tool_call={"name": "ls", "args": {"path": "/"}, "id": "tc1"},
|
||||
tool=MagicMock(),
|
||||
state={},
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
|
||||
def _sandbox_error_message(tool_call_id: str, sandbox_id: str = "sb-dead") -> ToolMessage:
|
||||
return ToolMessage(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": f"Sandbox request timed out: {sandbox_id}",
|
||||
"error_type": "SandboxClientError",
|
||||
"status": "error",
|
||||
}
|
||||
),
|
||||
tool_call_id=tool_call_id,
|
||||
status="error",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sandbox_client_error_recreates_sandbox() -> None:
|
||||
middleware = ToolErrorMiddleware()
|
||||
request = _tool_request()
|
||||
backend = FakeSandboxBackend()
|
||||
|
||||
async def handler(_request: ToolCallRequest) -> ToolMessage:
|
||||
raise SandboxClientError("Sandbox request timed out: sb-dead")
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("agent.server._recreate_sandbox", new_callable=AsyncMock) as mock_recreate,
|
||||
patch("agent.server.client") as mock_client,
|
||||
):
|
||||
mock_recreate.return_value = backend
|
||||
mock_client.threads.update = AsyncMock()
|
||||
|
||||
result = await middleware.awrap_tool_call(request, handler)
|
||||
|
||||
assert isinstance(result, ToolMessage)
|
||||
mock_recreate.assert_awaited_once_with("thread-1")
|
||||
mock_client.threads.update.assert_awaited_once_with(
|
||||
thread_id="thread-1",
|
||||
metadata={"sandbox_id": "sb-new"},
|
||||
)
|
||||
assert SANDBOX_BACKENDS["thread-1"] is backend
|
||||
|
||||
payload = json.loads(result.content)
|
||||
assert payload["status"] == "error"
|
||||
assert payload["error_type"] == "SandboxClientError"
|
||||
assert payload["recovery"] == "sandbox_recreated_after_client_error"
|
||||
assert payload["previous_error"] == "Sandbox request timed out: sb-dead"
|
||||
assert "sb-new" in payload["error"]
|
||||
finally:
|
||||
SANDBOX_BACKENDS.pop("thread-1", None)
|
||||
|
||||
|
||||
def test_repeated_sandbox_errors_trigger_circuit_breaker_once() -> None:
|
||||
middleware = SandboxCircuitBreakerMiddleware(threshold=2)
|
||||
messages = [
|
||||
HumanMessage(content="please fix this"),
|
||||
AIMessage(content="", tool_calls=[{"name": "ls", "args": {}, "id": "tc1"}]),
|
||||
_sandbox_error_message("tc1"),
|
||||
AIMessage(content="", tool_calls=[{"name": "grep", "args": {}, "id": "tc2"}]),
|
||||
_sandbox_error_message("tc2"),
|
||||
AIMessage(content="", tool_calls=[{"name": "execute", "args": {}, "id": "tc3"}]),
|
||||
_sandbox_error_message("tc3"),
|
||||
]
|
||||
|
||||
result = middleware.before_model({"messages": messages}, MagicMock())
|
||||
|
||||
assert result is not None
|
||||
assert result["jump_to"] == "end"
|
||||
assert len(result["messages"]) == 1
|
||||
assert "Sandbox circuit breaker triggered" in result["messages"][0].content
|
||||
|
||||
repeated = middleware.before_model(
|
||||
{"messages": [*messages, *result["messages"]]},
|
||||
MagicMock(),
|
||||
)
|
||||
assert repeated is None
|
||||
|
||||
|
||||
def test_repeated_sandbox_recreations_trigger_circuit_breaker() -> None:
|
||||
middleware = SandboxCircuitBreakerMiddleware(threshold=2)
|
||||
messages = [
|
||||
HumanMessage(content="please fix this"),
|
||||
AIMessage(content="", tool_calls=[{"name": "ls", "args": {}, "id": "tc1"}]),
|
||||
ToolMessage(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error_type": "SandboxClientError",
|
||||
"previous_error": "Sandbox request timed out: sb-old-1",
|
||||
"recovery": "sandbox_recreated_after_client_error",
|
||||
"sandbox_id": "sb-new-1",
|
||||
"status": "error",
|
||||
}
|
||||
),
|
||||
tool_call_id="tc1",
|
||||
status="error",
|
||||
),
|
||||
AIMessage(content="", tool_calls=[{"name": "grep", "args": {}, "id": "tc2"}]),
|
||||
ToolMessage(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error_type": "SandboxClientError",
|
||||
"previous_error": "Sandbox request timed out: sb-new-1",
|
||||
"recovery": "sandbox_recreated_after_client_error",
|
||||
"sandbox_id": "sb-new-2",
|
||||
"status": "error",
|
||||
}
|
||||
),
|
||||
tool_call_id="tc2",
|
||||
status="error",
|
||||
),
|
||||
AIMessage(content="", tool_calls=[{"name": "execute", "args": {}, "id": "tc3"}]),
|
||||
ToolMessage(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error_type": "SandboxClientError",
|
||||
"previous_error": "Sandbox request timed out: sb-new-2",
|
||||
"recovery": "sandbox_recreated_after_client_error",
|
||||
"sandbox_id": "sb-new-3",
|
||||
"status": "error",
|
||||
}
|
||||
),
|
||||
tool_call_id="tc3",
|
||||
status="error",
|
||||
),
|
||||
]
|
||||
|
||||
result = middleware.before_model({"messages": messages}, MagicMock())
|
||||
|
||||
assert result is not None
|
||||
assert result["jump_to"] == "end"
|
||||
assert "consecutive sandbox recreations" in result["messages"][0].content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_circuit_breaker_posts_one_user_notification() -> None:
|
||||
middleware = SandboxCircuitBreakerMiddleware(threshold=2)
|
||||
state = {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
content=(
|
||||
"Sandbox circuit breaker triggered: 3 consecutive sandbox tool failures "
|
||||
"against sb-dead."
|
||||
)
|
||||
)
|
||||
]
|
||||
}
|
||||
config = {
|
||||
"configurable": {
|
||||
"slack_thread": {"channel_id": "C123", "thread_ts": "171.123"},
|
||||
"linear_issue": {"id": "lin-1"},
|
||||
"repo": {"owner": "langchain-ai", "name": "open-swe"},
|
||||
"pr_number": 7,
|
||||
}
|
||||
}
|
||||
|
||||
with (
|
||||
patch("agent.middleware.sandbox_circuit_breaker.get_config", return_value=config),
|
||||
patch(
|
||||
"agent.middleware.sandbox_circuit_breaker.post_slack_thread_reply",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_slack,
|
||||
patch(
|
||||
"agent.middleware.sandbox_circuit_breaker.comment_on_linear_issue",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_linear,
|
||||
patch(
|
||||
"agent.middleware.sandbox_circuit_breaker.post_github_comment",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_github,
|
||||
):
|
||||
result = await middleware.aafter_agent(state, MagicMock())
|
||||
|
||||
assert result is None
|
||||
mock_slack.assert_awaited_once_with("C123", "171.123", SANDBOX_UNRECOVERABLE_MESSAGE)
|
||||
mock_linear.assert_not_called()
|
||||
mock_github.assert_not_called()
|
||||
Loading…
Add table
Reference in a new issue