diff --git a/strix/agents/factory.py b/strix/agents/factory.py index 3b9dd616..4cd9e037 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -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) diff --git a/strix/core/agents.py b/strix/core/agents.py index efd7cb60..7e738be1 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -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 diff --git a/strix/core/execution.py b/strix/core/execution.py index 5997e3db..87a2bf02 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -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( diff --git a/strix/core/hooks.py b/strix/core/hooks.py index eae4d7c7..06650be9 100644 --- a/strix/core/hooks.py +++ b/strix/core/hooks.py @@ -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: diff --git a/strix/core/runner.py b/strix/core/runner.py index 58124be2..f9bb4940 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -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: diff --git a/strix/tools/agents_graph/tools.py b/strix/tools/agents_graph/tools.py index 60c88525..a26c1dc4 100644 --- a/strix/tools/agents_graph/tools.py +++ b/strix/tools/agents_graph/tools.py @@ -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, diff --git a/tests/test_agent_tool_registration.py b/tests/test_agent_tool_registration.py index 489fea70..e459670b 100644 --- a/tests/test_agent_tool_registration.py +++ b/tests/test_agent_tool_registration.py @@ -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")) diff --git a/tests/test_execution.py b/tests/test_execution.py index c3e1776d..0daf0da2 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -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, diff --git a/tests/test_hooks.py b/tests/test_hooks.py index 16fce426..975a99b8 100644 --- a/tests/test_hooks.py +++ b/tests/test_hooks.py @@ -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: