Support chat-compatible sandbox patch tool

This commit is contained in:
0xallam
2026-04-26 16:54:34 -07:00
parent 88fc7be8c1
commit 756457f108
4 changed files with 152 additions and 3 deletions
+134 -2
View File
@@ -17,6 +17,8 @@ runtime skill-loading tool.
from __future__ import annotations from __future__ import annotations
import copy
import inspect
import json import json
import logging import logging
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -24,7 +26,7 @@ from typing import TYPE_CHECKING, Any
from agents.agent import ToolsToFinalOutputResult from agents.agent import ToolsToFinalOutputResult
from agents.sandbox import SandboxAgent from agents.sandbox import SandboxAgent
from agents.sandbox.capabilities import Filesystem, Shell from agents.sandbox.capabilities import Filesystem, Shell
from agents.tool import Tool from agents.tool import CustomTool, FunctionTool, Tool
from strix.agents.prompt import render_system_prompt from strix.agents.prompt import render_system_prompt
from strix.tools.agents_graph.tools import ( from strix.tools.agents_graph.tools import (
@@ -65,6 +67,8 @@ from strix.tools.web_search.tool import web_search
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from agents import RunContextWrapper from agents import RunContextWrapper
from agents.tool import FunctionToolResult from agents.tool import FunctionToolResult
@@ -72,6 +76,118 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_CUSTOM_TOOL_INPUT_FIELD_BY_NAME = {
"apply_patch": "patch",
}
_DEFAULT_CUSTOM_TOOL_INPUT_FIELD = "input"
def _custom_tool_input_field(tool: CustomTool) -> str:
return _CUSTOM_TOOL_INPUT_FIELD_BY_NAME.get(tool.name, _DEFAULT_CUSTOM_TOOL_INPUT_FIELD)
def _raw_input_schema(tool: CustomTool) -> dict[str, Any]:
input_field = _custom_tool_input_field(tool)
return {
"type": "object",
"properties": {
input_field: {
"type": "string",
"description": (
f"Complete `{tool.name}` payload. Follow the tool description exactly."
),
},
},
"required": [input_field],
"additionalProperties": False,
}
def _extract_custom_input(tool: CustomTool, raw_input: str | dict[str, Any]) -> str:
if isinstance(raw_input, str):
try:
parsed = json.loads(raw_input)
except json.JSONDecodeError:
return ""
else:
parsed = raw_input
value = parsed.get(_custom_tool_input_field(tool))
return value if isinstance(value, str) else ""
def _format_tool_error(exc: Exception) -> str:
return str(exc) or exc.__class__.__name__
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
safe_tool = copy.copy(tool)
invoke_tool = safe_tool.on_invoke_tool
async def invoke(ctx: Any, raw_input: str) -> Any:
try:
return await invoke_tool(ctx, raw_input)
except Exception as exc: # noqa: BLE001 - tool errors should be model-visible results.
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
return _format_tool_error(exc)
safe_tool.on_invoke_tool = invoke
return safe_tool
def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
"""Expose an SDK raw-input custom tool through Chat-Completions function calling."""
async def invoke(ctx: Any, raw_input: str) -> Any:
custom_input = _extract_custom_input(tool, raw_input)
if not custom_input:
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
try:
return await tool.on_invoke_tool(ctx, custom_input)
except Exception as exc: # noqa: BLE001 - matches SDK CustomTool error-as-result behavior.
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
return _format_tool_error(exc)
needs_approval = tool.runtime_needs_approval()
function_needs_approval: bool | Callable[[Any, dict[str, Any], str], Awaitable[bool]]
if callable(needs_approval):
async def approve(ctx: Any, args: dict[str, Any], call_id: str) -> bool:
result = needs_approval(ctx, _extract_custom_input(tool, args), call_id)
if inspect.isawaitable(result):
result = await result
return bool(result)
function_needs_approval = approve
else:
function_needs_approval = needs_approval
return FunctionTool(
name=tool.name,
description=(
f"{tool.description}\n\n"
f"Pass the complete `{tool.name}` payload in `{_custom_tool_input_field(tool)}`."
),
params_json_schema=_raw_input_schema(tool),
on_invoke_tool=invoke,
strict_json_schema=False,
needs_approval=function_needs_approval,
)
def _configure_chat_completions_filesystem_tools(toolset: Any) -> None:
for name, tool in vars(toolset).items():
if isinstance(tool, CustomTool):
setattr(toolset, name, _custom_tool_as_function_tool(tool))
elif isinstance(tool, FunctionTool):
setattr(toolset, name, _function_tool_with_error_result(tool))
def _configure_chat_completions_shell_tools(toolset: Any) -> None:
for name, tool in vars(toolset).items():
if isinstance(tool, FunctionTool):
setattr(toolset, name, _function_tool_with_error_result(tool))
def _lifecycle_tool_completed(tool_name: str, output: Any) -> bool: def _lifecycle_tool_completed(tool_name: str, output: Any) -> bool:
if tool_name == "agent_finish": if tool_name == "agent_finish":
completion_key = "agent_completed" completion_key = "agent_completed"
@@ -177,6 +293,7 @@ def build_strix_agent(
scan_mode: str = "deep", scan_mode: str = "deep",
is_whitebox: bool = False, is_whitebox: bool = False,
interactive: bool = False, interactive: bool = False,
chat_completions_tools: bool = False,
system_prompt_context: dict[str, Any] | None = None, system_prompt_context: dict[str, Any] | None = None,
) -> SandboxAgent[Any]: ) -> SandboxAgent[Any]:
"""Build a ``SandboxAgent`` configured for either root or child use. """Build a ``SandboxAgent`` configured for either root or child use.
@@ -203,6 +320,8 @@ def build_strix_agent(
the create_agent / wiki integration. the create_agent / wiki integration.
interactive: Renders the interactive-mode communication block interactive: Renders the interactive-mode communication block
in the system prompt. in the system prompt.
chat_completions_tools: Wrap SDK custom tools as function tools
when the selected backend cannot accept Responses custom tools.
system_prompt_context: Free-form dict the prompt template system_prompt_context: Free-form dict the prompt template
renders into the ``system_prompt_context`` variable — renders into the ``system_prompt_context`` variable —
today carries the scan scope / authorization block. today carries the scan scope / authorization block.
@@ -244,7 +363,18 @@ def build_strix_agent(
# model=None so ``RunConfig.model`` drives provider selection # model=None so ``RunConfig.model`` drives provider selection
# through the SDK's default MultiProvider. # through the SDK's default MultiProvider.
model=None, model=None,
capabilities=[Filesystem(), Shell()], capabilities=[
Filesystem(
configure_tools=(
_configure_chat_completions_filesystem_tools if chat_completions_tools else None
),
),
Shell(
configure_tools=(
_configure_chat_completions_shell_tools if chat_completions_tools else None
),
),
],
) )
@@ -253,6 +383,7 @@ def make_child_factory(
scan_mode: str = "deep", scan_mode: str = "deep",
is_whitebox: bool = False, is_whitebox: bool = False,
interactive: bool = False, interactive: bool = False,
chat_completions_tools: bool = False,
system_prompt_context: dict[str, Any] | None = None, system_prompt_context: dict[str, Any] | None = None,
) -> Any: ) -> Any:
"""Return the runner-owned builder used by ``spawn_child_agent``. """Return the runner-owned builder used by ``spawn_child_agent``.
@@ -270,6 +401,7 @@ def make_child_factory(
scan_mode=scan_mode, scan_mode=scan_mode,
is_whitebox=is_whitebox, is_whitebox=is_whitebox,
interactive=interactive, interactive=interactive,
chat_completions_tools=chat_completions_tools,
system_prompt_context=system_prompt_context, system_prompt_context=system_prompt_context,
) )
+8
View File
@@ -96,3 +96,11 @@ def normalize_model_name(model_name: str) -> str:
return f"litellm/gemini/{model}" return f"litellm/gemini/{model}"
return model return model
def uses_chat_completions_tool_schema(model_name: str, settings: Settings) -> bool:
"""Return whether the resolved SDK route can only receive JSON function tools."""
model = model_name.strip().lower()
if model.startswith(("litellm/", "any-llm/")):
return True
return bool(settings.llm.api_base)
+2
View File
@@ -370,6 +370,8 @@ async def _run_cycle(
logger.exception("agent run failed for %s; parking as %s", agent_id, status) logger.exception("agent run failed for %s; parking as %s", agent_id, status)
await coordinator.set_status(agent_id, status) await coordinator.set_status(agent_id, status)
await _notify_parent_on_crash(coordinator, agent_id, status) await _notify_parent_on_crash(coordinator, agent_id, status)
if context.get("parent_id") is None and status in {"failed", "crashed"}:
raise
return None return None
else: else:
await _settle_run_result(coordinator, agent_id, interactive) await _settle_run_result(coordinator, agent_id, interactive)
+8 -1
View File
@@ -19,7 +19,11 @@ from agents.sandbox import SandboxRunConfig
from strix.agents.factory import build_strix_agent, make_child_factory from strix.agents.factory import build_strix_agent, make_child_factory
from strix.config import load_settings from strix.config import load_settings
from strix.config.models import configure_sdk_model_defaults, normalize_model_name from strix.config.models import (
configure_sdk_model_defaults,
normalize_model_name,
uses_chat_completions_tool_schema,
)
from strix.core.agents import AgentCoordinator from strix.core.agents import AgentCoordinator
from strix.core.execution import ( from strix.core.execution import (
respawn_subagents, respawn_subagents,
@@ -99,6 +103,7 @@ async def run_strix_scan(
"No LLM model configured. Set STRIX_LLM env or pass model= to run_strix_scan().", "No LLM model configured. Set STRIX_LLM env or pass model= to run_strix_scan().",
) )
logger.info("LLM model resolved: %s", resolved_model) logger.info("LLM model resolved: %s", resolved_model)
chat_completions_tools = uses_chat_completions_tool_schema(resolved_model, settings)
# Caller may pre-create the coordinator so it can route stop/chat # Caller may pre-create the coordinator so it can route stop/chat
# commands while the scan loop runs in another thread. # commands while the scan loop runs in another thread.
@@ -178,6 +183,7 @@ async def run_strix_scan(
scan_mode=scan_mode, scan_mode=scan_mode,
is_whitebox=is_whitebox, is_whitebox=is_whitebox,
interactive=interactive, interactive=interactive,
chat_completions_tools=chat_completions_tools,
system_prompt_context=scope_context, system_prompt_context=scope_context,
) )
@@ -194,6 +200,7 @@ async def run_strix_scan(
scan_mode=scan_mode, scan_mode=scan_mode,
is_whitebox=is_whitebox, is_whitebox=is_whitebox,
interactive=interactive, interactive=interactive,
chat_completions_tools=chat_completions_tools,
system_prompt_context=scope_context, system_prompt_context=scope_context,
) )