diff --git a/navi/core/agent.py b/navi/core/agent.py index 4943595..7317456 100644 --- a/navi/core/agent.py +++ b/navi/core/agent.py @@ -60,6 +60,7 @@ AIHelperTokensUsed, CompressionStarted, ModelInfo, + ProfileSwitched, StreamEnd, StreamStopped, SubagentComplete, @@ -784,6 +785,19 @@ ) -> list[Tool]: return build_tool_list(scope.native, scope.mcp, self._tools, self._mcp_manager) + def _tool_list_for_profile(self, profile_id: str) -> "list[Tool] | None": + """Resolve another profile's tools by id; None when it is unknown. + + Used mid-batch right after a switch_profile call, where the running + turn's own ``profile`` binding is still the one it started with. + """ + try: + profile = self._profiles.get(profile_id) + except Exception: + log.warning("agent.switch_profile_unknown", profile_id=profile_id) + return None + return self._tool_list(profile.get_agent_tools()) + def _get_backend(self, backend_key: str) -> LLMBackend: return self._backends.get(backend_key) @@ -944,6 +958,29 @@ The mode is decided per-turn by ``turn_ctx.parallel_tool_calls`` (profile override ? global setting, resolved at tool_ctx build). """ + # A profile switch replaces the toolset mid-turn, so it must not share a + # batch: every neighbour would still be dispatched against the old + # tool_map and die with "tool 'X' not found" (seen with reload_tools after + # leaving tool_developer). Run it alone first, then re-resolve the tools + # it switched to for the rest of the batch. The turn's own ``profile`` + # binding stays old on purpose — the end-of-iteration reload in + # run_stream() is what rebinds profile/llm/schemas. + switching = next((tc for tc in turn_tool_calls if tc.name == "switch_profile"), None) + if switching is not None and len(turn_tool_calls) > 1: + rest = [tc for tc in turn_tool_calls if tc is not switching] + switched_to = None + async for ev in self._execute_tools_sequential( + [switching], tools, turn_ctx, session, stop_event, tool_ctx + ): + if isinstance(ev, ProfileSwitched): + switched_to = ev.profile_id + yield ev + if switched_to: + new_tools = self._tool_list_for_profile(switched_to) + if new_tools is not None: + tools = new_tools + turn_tool_calls = rest + if turn_tool_calls and turn_ctx.parallel_tool_calls and len(turn_tool_calls) > 1: async for ev in self._execute_tools_parallel( turn_tool_calls, tools, turn_ctx, session, stop_event, tool_ctx diff --git a/tests/unit/core/test_parallel_tools.py b/tests/unit/core/test_parallel_tools.py index 98e23e3..ba39747 100644 --- a/tests/unit/core/test_parallel_tools.py +++ b/tests/unit/core/test_parallel_tools.py @@ -251,6 +251,82 @@ assert kinds == ["ToolStarted", "ToolEvent"] +class CountingTool: + """Records how many times it was dispatched.""" + + def __init__(self, name): + self.name = name + self.used = 0 + + async def execute(self, arguments, ctx=None): + self.used += 1 + return ToolResult(success=True, output=f"done {self.name}") + + +class SwitchingTool: + """switch_profile stand-in: emits the event the agent loop keys off.""" + + name = "switch_profile" + + async def execute(self, arguments, ctx=None): + from navi.core.events import ProfileSwitched + + await current_event_sink.get().put( + ProfileSwitched(profile_id="new", profile_name="New") + ) + return ToolResult(success=True, output="switched") + + +class TestSwitchProfileBatch: + """A profile switch swaps the toolset — the batch must not race it.""" + + def _agent_with_new_toolset(self, new_tools): + agent, _ = make_agent() + agent._profiles = SimpleNamespace( + get=lambda pid: SimpleNamespace(get_agent_tools=lambda: None) + ) + agent._tool_list = lambda scope: new_tools + return agent + + async def test_switch_runs_first_and_the_rest_uses_the_new_toolset(self): + old_probe, new_probe = CountingTool("probe"), CountingTool("probe") + switch = SwitchingTool() + agent = self._agent_with_new_toolset([switch, new_probe]) + session = SimpleNamespace(messages=[], context=[]) + + # The switch is second in call order: it must still run first. + events = [ev async for ev in agent._execute_tools_with_sink( + make_tcs("probe", "switch_profile"), [old_probe, switch], + make_turn_ctx(parallel=True), session, None, None)] + + kinds = [type(e).__name__ for e in events] + assert kinds == ["ToolStarted", "ProfileSwitched", "ToolEvent", + "ToolStarted", "ToolEvent"] + # ...and the switch's result lands first, ahead of the call it was batched with + assert [e.tool_name for e in events if isinstance(e, ToolEvent)] == \ + ["switch_profile", "probe"] + # the late call was dispatched against the NEW profile's tool object + assert new_probe.used == 1 + assert old_probe.used == 0 + + async def test_unknown_target_keeps_the_old_toolset(self): + """Tools are re-resolved only when the profile resolves — an unknown id + logs and leaves the rest of the batch on the old list.""" + probe = CountingTool("probe") + agent, _ = make_agent() + agent._profiles = SimpleNamespace(get=lambda pid: (_ for _ in ()).throw(KeyError(pid))) + agent._tool_list = lambda scope: [] # must never be reached + session = SimpleNamespace(messages=[], context=[]) + + events = [ev async for ev in agent._execute_tools_with_sink( + make_tcs("switch_profile", "probe"), [SwitchingTool(), probe], + make_turn_ctx(parallel=True), session, None, None)] + + assert [e.tool_name for e in events if isinstance(e, ToolEvent)] == \ + ["switch_profile", "probe"] + assert probe.used == 1 + + class TestRepair: def test_repair_dangling_tool_calls(self): from navi.core.pg_session_store import repair_dangling_tool_calls