mirror of
https://github.com/usestrix/strix.git
synced 2026-10-09 03:18:31 +00:00
feat(agents): let callers choose the top-level agent's finish tool (#1496)
This commit is contained in:
parent
f5a900416b
commit
62b496430d
9 changed files with 88 additions and 12 deletions
|
|
@ -675,6 +675,7 @@ def build_strix_agent(
|
|||
extra_tools: Sequence[Tool] | None = None,
|
||||
instructions_override: str | None = None,
|
||||
supports_images: bool = True,
|
||||
finish_tool: Tool = finish_scan,
|
||||
) -> SandboxAgent[Any]:
|
||||
"""Build a SandboxAgent for either root or child use.
|
||||
|
||||
|
|
@ -687,6 +688,7 @@ def build_strix_agent(
|
|||
registered via ``register_agent_tools``.
|
||||
instructions_override: Use this verbatim as the system prompt instead
|
||||
of rendering the built-in scan prompt.
|
||||
finish_tool: The tool that ends the run, given to the root agent.
|
||||
"""
|
||||
if instructions_override is not None:
|
||||
instructions = instructions_override
|
||||
|
|
@ -707,7 +709,7 @@ def build_strix_agent(
|
|||
# Yielding to the user is only meaningful when one is attached.
|
||||
agent_tools.append(wait_for_user)
|
||||
if is_root:
|
||||
tools: list[Tool] = [*_BASE_TOOLS, *agent_tools, finish_scan]
|
||||
tools: list[Tool] = [*_BASE_TOOLS, *agent_tools, finish_tool]
|
||||
else:
|
||||
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
|
||||
_ensure_unique_tool_names(tools)
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ class AgentCoordinator:
|
|||
self._lock = asyncio.Lock()
|
||||
self._snapshot_path: Path | None = None
|
||||
self.is_shutting_down = False
|
||||
self.root_finish_tool = "finish_scan"
|
||||
self._budget_stopped = False
|
||||
self._reserve_stopped = False
|
||||
self._budget_paused = False
|
||||
|
|
|
|||
|
|
@ -302,7 +302,7 @@ async def _run_agent_loop(
|
|||
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
|
||||
|
||||
if reserve_stopped and start_parked and interactive and context.get("parent_id") is None:
|
||||
await coordinator.send(agent_id, _reserve_notice())
|
||||
await coordinator.send(agent_id, _reserve_notice(coordinator.root_finish_tool))
|
||||
|
||||
if not (start_parked and interactive):
|
||||
with contextlib.suppress(BudgetPausedError):
|
||||
|
|
@ -642,6 +642,7 @@ async def _run_until_lifecycle(
|
|||
input_data = await _append_tool_required_message(
|
||||
session=session,
|
||||
context=context,
|
||||
root_finish_tool=coordinator.root_finish_tool,
|
||||
attempt=recoveries,
|
||||
limit=recovery_limit,
|
||||
interactive=interactive,
|
||||
|
|
@ -988,8 +989,9 @@ async def _append_tool_required_message(
|
|||
limit: int,
|
||||
interactive: bool,
|
||||
silent_yield: bool = False,
|
||||
root_finish_tool: str = "finish_scan",
|
||||
) -> list[dict[str, str]]:
|
||||
finish_tool = "finish_scan" if context.get("parent_id") is None else "agent_finish"
|
||||
finish_tool = root_finish_tool if context.get("parent_id") is None else "agent_finish"
|
||||
if silent_yield:
|
||||
message = (
|
||||
"You called wait_for_user without having written anything to the user since "
|
||||
|
|
@ -1107,7 +1109,7 @@ async def notify_parent_on_terminal(
|
|||
)
|
||||
|
||||
|
||||
def _reserve_notice() -> dict[str, Any]:
|
||||
def _reserve_notice(finish_tool: str) -> dict[str, Any]:
|
||||
return {
|
||||
"from": "system",
|
||||
"type": "budget_reserve_stop",
|
||||
|
|
@ -1117,7 +1119,7 @@ def _reserve_notice() -> dict[str, Any]:
|
|||
"sub-agent is being force-stopped as soon as its in-flight turn completes, and "
|
||||
"none will send a completion report. Their confirmed vulnerabilities are "
|
||||
"already filed as they were found. Do not wait on any sub-agents and do not "
|
||||
"spawn new ones — wrap up now and call finish_scan."
|
||||
f"spawn new ones — wrap up now and call {finish_tool}."
|
||||
),
|
||||
}
|
||||
|
||||
|
|
@ -1126,7 +1128,7 @@ async def _notify_root_on_budget_reserve(coordinator: AgentCoordinator) -> None:
|
|||
root = await coordinator.claim_reserve_notification()
|
||||
if root is None:
|
||||
return
|
||||
await coordinator.send(root, _reserve_notice())
|
||||
await coordinator.send(root, _reserve_notice(coordinator.root_finish_tool))
|
||||
|
||||
|
||||
async def _notify_parent_on_exit(
|
||||
|
|
|
|||
|
|
@ -104,16 +104,16 @@ _ROOT_DIRECTIVES: tuple[str, ...] = (
|
|||
(
|
||||
"As the root agent, begin planning your wind-down of the whole scan: avoid "
|
||||
"starting large new lines of investigation, and keep your required objectives on "
|
||||
"track so you can call finish_scan comfortably before the limit."
|
||||
"track so you can call {finish_tool} comfortably before the limit."
|
||||
),
|
||||
(
|
||||
"As the root agent, prioritize wrapping up the whole scan now: stop opening new "
|
||||
"lines of investigation, close out only what is essential, and move toward calling "
|
||||
"finish_scan to compile and deliver the final report."
|
||||
"{finish_tool} to compile and deliver the final report."
|
||||
),
|
||||
(
|
||||
"As the root agent, STOP all other work on the whole scan and finish immediately: "
|
||||
"secure your findings and call finish_scan now — anything left unfinished when the "
|
||||
"secure your findings and call {finish_tool} now — anything left unfinished when the "
|
||||
"limit is hit is discarded."
|
||||
),
|
||||
)
|
||||
|
|
@ -138,8 +138,11 @@ _SUBAGENT_DIRECTIVES: tuple[str, ...] = (
|
|||
|
||||
def _wrapup_directive(context: RunContextWrapper[dict[str, Any]], stage: int) -> str:
|
||||
is_root = context.context.get("parent_id") is None
|
||||
directives = _ROOT_DIRECTIVES if is_root else _SUBAGENT_DIRECTIVES
|
||||
return directives[stage]
|
||||
if not is_root:
|
||||
return _SUBAGENT_DIRECTIVES[stage]
|
||||
coordinator = coordinator_from_context(context.context)
|
||||
finish_tool = coordinator.root_finish_tool if coordinator is not None else "finish_scan"
|
||||
return _ROOT_DIRECTIVES[stage].format(finish_tool=finish_tool)
|
||||
|
||||
|
||||
def _urgency(stage: int) -> str:
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ from strix.report.state import get_global_report_state
|
|||
from strix.runtime import session_manager
|
||||
from strix.telemetry import set_scan_phase
|
||||
from strix.telemetry.logging import set_scan_id, setup_scan_logging
|
||||
from strix.tools.finish.tool import finish_scan
|
||||
from strix.tools.output_store import (
|
||||
WORKSPACE_SPILL_DIR,
|
||||
configure_spill_writer,
|
||||
|
|
@ -58,6 +59,7 @@ from strix.tools.output_store import (
|
|||
if TYPE_CHECKING:
|
||||
from agents.memory import SQLiteSession
|
||||
from agents.result import RunResultBase
|
||||
from agents.tool import Tool
|
||||
|
||||
from strix.runtime.status import StatusSink
|
||||
from strix.tools.mcp import (
|
||||
|
|
@ -215,6 +217,7 @@ async def run_strix_scan(
|
|||
status_sink: StatusSink | None = None,
|
||||
mcp_connection_requests: list[McpConnectionRequest] | None = None,
|
||||
mcp_status_sink: McpStatusSink | None = None,
|
||||
root_finish_tool: Tool = finish_scan,
|
||||
) -> RunResultBase | None:
|
||||
"""Run or resume one Strix scan against a sandbox.
|
||||
|
||||
|
|
@ -237,6 +240,7 @@ async def run_strix_scan(
|
|||
command-line default) it reads ``~/.strix/mcp-servers.json`` itself. Either
|
||||
way the engine does the connecting, so the caller passes inert configs plus
|
||||
metadata and never live sessions.
|
||||
``root_finish_tool`` is the tool the root agent ends the run with.
|
||||
"""
|
||||
|
||||
def report(phase: str) -> None:
|
||||
|
|
@ -293,6 +297,7 @@ async def run_strix_scan(
|
|||
coordinator = AgentCoordinator()
|
||||
coordinator.set_snapshot_path(agents_path)
|
||||
coordinator.set_budget_policy(budget_policy)
|
||||
coordinator.root_finish_tool = root_finish_tool.name
|
||||
|
||||
from strix.tools.coverage.tools import hydrate_coverage_from_disk
|
||||
from strix.tools.notes.tools import hydrate_notes_from_disk
|
||||
|
|
@ -503,6 +508,7 @@ async def run_strix_scan(
|
|||
system_prompt_context=root_context,
|
||||
instructions_override=root_instructions,
|
||||
supports_images=supports_images,
|
||||
finish_tool=root_finish_tool,
|
||||
)
|
||||
|
||||
if not is_resume:
|
||||
|
|
|
|||
|
|
@ -683,7 +683,8 @@ async def agent_finish(
|
|||
{
|
||||
"success": False,
|
||||
"error": (
|
||||
"agent_finish is for subagents. Root/main agents must call finish_scan instead"
|
||||
"agent_finish is for subagents. Root/main agents must call "
|
||||
f"{coordinator.root_finish_tool} instead"
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
|
|
|
|||
|
|
@ -59,6 +59,15 @@ def test_registered_tools_appear_before_lifecycle_tool() -> None:
|
|||
assert child_names[-2:] == ["extra", "agent_finish"]
|
||||
|
||||
|
||||
def test_root_ends_with_the_given_finish_tool() -> None:
|
||||
root = factory.build_strix_agent(is_root=True, finish_tool=_tool("finish_pr_review"))
|
||||
|
||||
root_names = [t.name for t in root.tools]
|
||||
|
||||
assert root_names[-1] == "finish_pr_review"
|
||||
assert "finish_scan" not in root_names
|
||||
|
||||
|
||||
def test_per_call_extra_tools_stack_with_registry() -> None:
|
||||
factory.register_agent_tools(_tool("registered"))
|
||||
|
||||
|
|
|
|||
|
|
@ -132,6 +132,27 @@ async def test_reserve_stop_notifies_root_once(monkeypatch: pytest.MonkeyPatch)
|
|||
assert "finish_scan" in str(message["content"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_notice_names_the_root_finish_tool(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.root_finish_tool = "finish_pr_review"
|
||||
await coordinator.register("root", "strix", parent_id=None)
|
||||
await coordinator.register("child-a", "recon", parent_id="root")
|
||||
|
||||
sent: list[dict[str, Any]] = []
|
||||
|
||||
async def _record(_target_agent_id: str, message: dict[str, Any]) -> bool:
|
||||
sent.append(message)
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(coordinator, "send", _record)
|
||||
|
||||
await _notify_root_on_budget_reserve(coordinator)
|
||||
|
||||
assert "call finish_pr_review" in str(sent[0]["content"])
|
||||
assert "finish_scan" not in str(sent[0]["content"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_reserve_claims_yield_single_root() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
|
|
@ -1295,6 +1316,21 @@ async def test_interactive_nudge_offers_waiting_without_repeating() -> None:
|
|||
assert "do not repeat it" in items[0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nudge_names_the_root_finish_tool() -> None:
|
||||
items = await execution._append_tool_required_message(
|
||||
session=None,
|
||||
context={"parent_id": None},
|
||||
attempt=1,
|
||||
limit=3,
|
||||
interactive=False,
|
||||
root_finish_tool="finish_pr_review",
|
||||
)
|
||||
|
||||
assert "call finish_pr_review" in items[0]["content"]
|
||||
assert "finish_scan" not in items[0]["content"]
|
||||
|
||||
|
||||
def _cycle_with_items(
|
||||
coordinator: AgentCoordinator,
|
||||
agent_id: str,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
|
|||
|
||||
import pytest
|
||||
|
||||
from strix.core.agents import AgentCoordinator
|
||||
from strix.core.hooks import (
|
||||
BudgetExceededError,
|
||||
BudgetPausedError,
|
||||
|
|
@ -335,6 +336,21 @@ async def test_budget_warning_root_directive_distinct_from_subagent() -> None:
|
|||
assert "confirmed" in sub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_root_directive_names_the_coordinators_finish_tool() -> None:
|
||||
hooks = ReportUsageHooks(model="test-model", max_turns=100)
|
||||
coordinator = AgentCoordinator()
|
||||
coordinator.root_finish_tool = "finish_pr_review"
|
||||
ctx = _make_warn_context(requests=85, parent_id=None)
|
||||
ctx.context["coordinator"] = coordinator
|
||||
|
||||
items: list[Any] = []
|
||||
await hooks.on_llm_start(ctx, MagicMock(), None, items)
|
||||
|
||||
assert "finish_pr_review" in items[0]["content"]
|
||||
assert "finish_scan" not in items[0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("parent_id", [None, "root-1"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_turn_warning_directive_escalates_per_stage(parent_id: str | None) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue