mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Support chat-compatible sandbox patch tool
This commit is contained in:
+134
-2
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user