mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
feat: enforce timeouts on execute tool commands (#955)
* feat: enforce timeouts on execute tool commands * refactor: dedupe execute-timeout middleware logic * cr * move info to debug --------- Co-authored-by: bracesproul <braceasproul@gmail.com>
This commit is contained in:
parent
8297c61fe1
commit
e7ed860114
4 changed files with 89 additions and 0 deletions
|
|
@ -1,9 +1,11 @@
|
|||
from .check_message_queue import check_message_queue_before_model
|
||||
from .open_pr import open_pr_if_needed
|
||||
from .post_to_linear import post_to_linear_after_model
|
||||
from .timeout_execute_tool import TimeoutExecuteToolMiddleware
|
||||
from .tool_error_handler import ToolErrorMiddleware
|
||||
|
||||
__all__ = [
|
||||
"TimeoutExecuteToolMiddleware",
|
||||
"ToolErrorMiddleware",
|
||||
"check_message_queue_before_model",
|
||||
"open_pr_if_needed",
|
||||
|
|
|
|||
83
apps/agent/agent/middleware/timeout_execute_tool.py
Normal file
83
apps/agent/agent/middleware/timeout_execute_tool.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
"""Tool middleware that wraps execute commands with a timeout."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import shlex
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from langchain.agents.middleware.types import AgentMiddleware, AgentState
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||
from langgraph.types import Command
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT_SECONDS = 300
|
||||
TIMEOUT_REGEX = re.compile(r"\btimeout\s+\d+(?:\.\d+)?\s*[smhd]?\b", re.IGNORECASE)
|
||||
|
||||
|
||||
def _get_tool_name(request: ToolCallRequest) -> str | None:
|
||||
tool_call = request.tool_call
|
||||
if isinstance(tool_call, dict):
|
||||
return tool_call.get("name")
|
||||
return None
|
||||
|
||||
|
||||
def _get_command_arg(request: ToolCallRequest) -> str | None:
|
||||
tool_call = request.tool_call
|
||||
if not isinstance(tool_call, dict):
|
||||
return None
|
||||
args = tool_call.get("args")
|
||||
if not isinstance(args, dict):
|
||||
return None
|
||||
command = args.get("command")
|
||||
return command if isinstance(command, str) else None
|
||||
|
||||
|
||||
def _wrap_command(command: str) -> str:
|
||||
if TIMEOUT_REGEX.search(command):
|
||||
return command
|
||||
quoted = shlex.quote(command)
|
||||
return f"timeout {DEFAULT_TIMEOUT_SECONDS}s sh -c {quoted}"
|
||||
|
||||
|
||||
def _overwrite_request_if_needed(request: ToolCallRequest) -> ToolCallRequest:
|
||||
if _get_tool_name(request) != "execute":
|
||||
return request
|
||||
|
||||
command = _get_command_arg(request)
|
||||
if not command:
|
||||
return request
|
||||
|
||||
wrapped = _wrap_command(command)
|
||||
if wrapped == command:
|
||||
return request
|
||||
|
||||
tool_call = dict(request.tool_call)
|
||||
args = dict(tool_call.get("args", {}))
|
||||
args["command"] = wrapped
|
||||
tool_call["args"] = args
|
||||
logger.debug("Wrapped execute command with timeout")
|
||||
return request.override(tool_call=tool_call)
|
||||
|
||||
|
||||
class TimeoutExecuteToolMiddleware(AgentMiddleware):
|
||||
"""Ensure execute tool calls are wrapped with a timeout."""
|
||||
|
||||
state_schema = AgentState
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||
) -> ToolMessage | Command:
|
||||
return handler(_overwrite_request_if_needed(request))
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
|
||||
) -> ToolMessage | Command:
|
||||
return await handler(_overwrite_request_if_needed(request))
|
||||
|
|
@ -8,6 +8,8 @@ All code execution and file operations happen in this sandbox environment.
|
|||
|
||||
**Important:**
|
||||
- Use `{working_dir}` as your working directory for all operations
|
||||
- The `execute` tool enforces a 5-minute timeout by default (`timeout 300s`)
|
||||
- If a command times out and needs longer, rerun it by explicitly appending `timeout Ns`
|
||||
|
||||
|
||||
---
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from .middleware import (
|
|||
check_message_queue_before_model,
|
||||
open_pr_if_needed,
|
||||
post_to_linear_after_model,
|
||||
TimeoutExecuteToolMiddleware,
|
||||
)
|
||||
from .prompt import construct_system_prompt
|
||||
from .tools import commit_and_open_pr, fetch_url, http_request
|
||||
|
|
@ -345,6 +346,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915
|
|||
tools=[http_request, fetch_url, commit_and_open_pr],
|
||||
backend=sandbox_backend,
|
||||
middleware=[
|
||||
TimeoutExecuteToolMiddleware(),
|
||||
ToolErrorMiddleware(),
|
||||
check_message_queue_before_model,
|
||||
post_to_linear_after_model,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue