mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 11:22:10 +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 .check_message_queue import check_message_queue_before_model
|
||||||
from .open_pr import open_pr_if_needed
|
from .open_pr import open_pr_if_needed
|
||||||
from .post_to_linear import post_to_linear_after_model
|
from .post_to_linear import post_to_linear_after_model
|
||||||
|
from .timeout_execute_tool import TimeoutExecuteToolMiddleware
|
||||||
from .tool_error_handler import ToolErrorMiddleware
|
from .tool_error_handler import ToolErrorMiddleware
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"TimeoutExecuteToolMiddleware",
|
||||||
"ToolErrorMiddleware",
|
"ToolErrorMiddleware",
|
||||||
"check_message_queue_before_model",
|
"check_message_queue_before_model",
|
||||||
"open_pr_if_needed",
|
"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:**
|
**Important:**
|
||||||
- Use `{working_dir}` as your working directory for all operations
|
- 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,
|
check_message_queue_before_model,
|
||||||
open_pr_if_needed,
|
open_pr_if_needed,
|
||||||
post_to_linear_after_model,
|
post_to_linear_after_model,
|
||||||
|
TimeoutExecuteToolMiddleware,
|
||||||
)
|
)
|
||||||
from .prompt import construct_system_prompt
|
from .prompt import construct_system_prompt
|
||||||
from .tools import commit_and_open_pr, fetch_url, http_request
|
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],
|
tools=[http_request, fetch_url, commit_and_open_pr],
|
||||||
backend=sandbox_backend,
|
backend=sandbox_backend,
|
||||||
middleware=[
|
middleware=[
|
||||||
|
TimeoutExecuteToolMiddleware(),
|
||||||
ToolErrorMiddleware(),
|
ToolErrorMiddleware(),
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
post_to_linear_after_model,
|
post_to_linear_after_model,
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue