feat(agents): let callers choose the top-level agent's finish tool (#1496)
Some checks are pending
CI / python (push) Waiting to run
CI / tui (push) Waiting to run
CI / viewer (push) Waiting to run
CI / package (push) Waiting to run
CI / workflows (push) Waiting to run
CI / ci-passed (push) Blocked by required conditions

This commit is contained in:
ian-at-strix 2026-10-08 19:11:50 -04:00 • committed by GitHub
parent f5a900416b
commit 62b496430d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 88 additions and 12 deletions

View file

@ -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)

View file

@ -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

View file

@ -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(

View file

@ -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:

View file

@ -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:

View file

@ -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,

View file

@ -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"))

View file

@ -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,

View file

@ -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: