refactor(agent): capture subagent runtime before spawn

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent 198fd9f869
commit af85c356b8
16 changed files with 261 additions and 132 deletions
-4
View File
@@ -340,11 +340,8 @@ class AgentLoop:
self._file_state_store = FileStateStore() self._file_state_store = FileStateStore()
self.runner = AgentRunner() self.runner = AgentRunner()
self.subagents = SubagentManager( self.subagents = SubagentManager(
provider=provider,
workspace=workspace, workspace=workspace,
bus=bus, bus=bus,
model=self.model,
context_window_tokens=self.context_window_tokens,
tools_config=_tc, tools_config=_tc,
max_tool_result_chars=self.max_tool_result_chars, max_tool_result_chars=self.max_tool_result_chars,
restrict_to_workspace=restrict_to_workspace, restrict_to_workspace=restrict_to_workspace,
@@ -495,7 +492,6 @@ class AgentLoop:
self.provider = provider self.provider = provider
self.model = model self.model = model
self.context_window_tokens = context_window_tokens self.context_window_tokens = context_window_tokens
self.subagents.set_provider(provider, model, context_window_tokens)
self.consolidator.set_provider(provider, model, context_window_tokens) self.consolidator.set_provider(provider, model, context_window_tokens)
self._sync_replay_max_messages() self._sync_replay_max_messages()
self._provider_signature = snapshot.signature self._provider_signature = snapshot.signature
+20 -26
View File
@@ -12,14 +12,18 @@ from loguru import logger
from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.agent.hook import AgentHook, AgentHookContext
from nanobot.agent.runner import AgentRunner, AgentRunSpec from nanobot.agent.runner import AgentRunner, AgentRunSpec
from nanobot.agent.tools.context import ToolContext from nanobot.agent.tools.context import (
RequestContext,
ToolContext,
bind_request_context,
reset_request_context,
)
from nanobot.agent.tools.file_state import FileStates from nanobot.agent.tools.file_state import FileStates
from nanobot.agent.tools.loader import ToolLoader from nanobot.agent.tools.loader import ToolLoader
from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import AgentDefaults, ToolsConfig from nanobot.config.schema import AgentDefaults, ToolsConfig
from nanobot.providers.base import LLMProvider
from nanobot.security.workspace_access import ( from nanobot.security.workspace_access import (
WorkspaceScope, WorkspaceScope,
bind_workspace_scope, bind_workspace_scope,
@@ -77,12 +81,9 @@ class SubagentManager:
def __init__( def __init__(
self, self,
provider: LLMProvider,
workspace: Path, workspace: Path,
bus: MessageBus, bus: MessageBus,
max_tool_result_chars: int, max_tool_result_chars: int,
model: str | None = None,
context_window_tokens: int | None = None,
tools_config: ToolsConfig | None = None, tools_config: ToolsConfig | None = None,
restrict_to_workspace: bool = False, restrict_to_workspace: bool = False,
disabled_skills: list[str] | None = None, disabled_skills: list[str] | None = None,
@@ -92,11 +93,8 @@ class SubagentManager:
llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None, llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None,
): ):
defaults = AgentDefaults() defaults = AgentDefaults()
self.provider = provider
self.workspace = workspace self.workspace = workspace
self.bus = bus self.bus = bus
self.model = model or provider.get_default_model()
self.context_window_tokens = context_window_tokens or defaults.context_window_tokens
self.tools_config = tools_config or ToolsConfig() self.tools_config = tools_config or ToolsConfig()
self.max_tool_result_chars = max_tool_result_chars self.max_tool_result_chars = max_tool_result_chars
self.restrict_to_workspace = restrict_to_workspace self.restrict_to_workspace = restrict_to_workspace
@@ -152,17 +150,6 @@ class SubagentManager:
ToolLoader().load(ctx, registry, scope="subagent") ToolLoader().load(ctx, registry, scope="subagent")
return registry return registry
def set_provider(
self,
provider: LLMProvider,
model: str,
context_window_tokens: int | None = None,
) -> None:
self.provider = provider
self.model = model
if context_window_tokens is not None:
self.context_window_tokens = context_window_tokens
async def spawn( async def spawn(
self, self,
task: str, task: str,
@@ -173,8 +160,12 @@ class SubagentManager:
origin_message_id: str | None = None, origin_message_id: str | None = None,
temperature: float | None = None, temperature: float | None = None,
workspace_scope: WorkspaceScope | None = None, workspace_scope: WorkspaceScope | None = None,
*,
runtime: LLMRuntime,
) -> str: ) -> str:
"""Spawn a subagent to execute a task in the background.""" """Spawn a subagent to execute a task in the background."""
if temperature is not None:
runtime = runtime.with_generation_overrides(temperature=temperature)
task_id = str(uuid.uuid4())[:8] task_id = str(uuid.uuid4())[:8]
display_label = label or task[:30] + ("..." if len(task) > 30 else "") display_label = label or task[:30] + ("..." if len(task) > 30 else "")
origin = {"channel": origin_channel, "chat_id": origin_chat_id, "session_key": session_key} origin = {"channel": origin_channel, "chat_id": origin_chat_id, "session_key": session_key}
@@ -194,8 +185,8 @@ class SubagentManager:
display_label, display_label,
origin, origin,
status, status,
runtime,
origin_message_id, origin_message_id,
temperature,
workspace_scope, workspace_scope,
) )
) )
@@ -223,8 +214,8 @@ class SubagentManager:
label: str, label: str,
origin: dict[str, str], origin: dict[str, str],
status: SubagentStatus, status: SubagentStatus,
runtime: LLMRuntime,
origin_message_id: str | None = None, origin_message_id: str | None = None,
temperature: float | None = None,
workspace_scope: WorkspaceScope | None = None, workspace_scope: WorkspaceScope | None = None,
) -> None: ) -> None:
"""Execute the subagent task and announce the result.""" """Execute the subagent task and announce the result."""
@@ -253,13 +244,15 @@ class SubagentManager:
if self._llm_wall_timeout_for_session if self._llm_wall_timeout_for_session
else None else None
) )
request_token = bind_request_context(RequestContext(
channel=origin["channel"],
chat_id=origin["chat_id"],
message_id=origin_message_id,
session_key=sess_key,
runtime=runtime,
))
token = bind_workspace_scope(workspace_scope) if workspace_scope is not None else None token = bind_workspace_scope(workspace_scope) if workspace_scope is not None else None
try: try:
runtime = LLMRuntime.capture(
self.provider,
self.model,
context_window_tokens=self.context_window_tokens,
).with_generation_overrides(temperature=temperature)
result = await self.runner.run(AgentRunSpec( result = await self.runner.run(AgentRunSpec(
initial_messages=messages, initial_messages=messages,
tools=tools, tools=tools,
@@ -279,6 +272,7 @@ class SubagentManager:
finally: finally:
if token is not None: if token is not None:
reset_workspace_scope(token) reset_workspace_scope(token)
reset_request_context(request_token)
status.phase = "done" status.phase = "done"
status.stop_reason = result.stop_reason status.stop_reason = result.stop_reason
+8 -9
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from nanobot.agent.tools.base import Tool, tool_parameters from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
from nanobot.agent.tools.context import current_request_context from nanobot.agent.tools.context import current_request_context
from nanobot.agent.tools.schema import NumberSchema, StringSchema, tool_parameters_schema from nanobot.agent.tools.schema import NumberSchema, StringSchema, tool_parameters_schema
from nanobot.security.workspace_access import current_workspace_scope from nanobot.security.workspace_access import current_workspace_scope
@@ -70,20 +70,19 @@ class SpawnTool(Tool):
f"to complete before spawning a new one." f"to complete before spawning a new one."
) )
request_ctx = current_request_context() request_ctx = current_request_context()
origin_channel = request_ctx.channel if request_ctx is not None else "cli" if request_ctx is None or request_ctx.runtime is None:
origin_chat_id = request_ctx.chat_id if request_ctx is not None else "direct" return ToolResult.error("Error: spawn requires an active model runtime")
session_key = ( origin_channel = request_ctx.channel
request_ctx.session_key or f"{origin_channel}:{origin_chat_id}" origin_chat_id = request_ctx.chat_id
if request_ctx is not None session_key = request_ctx.session_key or f"{origin_channel}:{origin_chat_id}"
else "cli:direct"
)
return await self._manager.spawn( return await self._manager.spawn(
task=task, task=task,
runtime=request_ctx.runtime,
label=label, label=label,
origin_channel=origin_channel, origin_channel=origin_channel,
origin_chat_id=origin_chat_id, origin_chat_id=origin_chat_id,
session_key=session_key, session_key=session_key,
origin_message_id=request_ctx.message_id if request_ctx is not None else None, origin_message_id=request_ctx.message_id,
temperature=temperature, temperature=temperature,
workspace_scope=current_workspace_scope(), workspace_scope=current_workspace_scope(),
) )
+4 -1
View File
@@ -8,14 +8,16 @@ import pytest
from nanobot.agent.runner import AgentRunResult from nanobot.agent.runner import AgentRunResult
from nanobot.agent.subagent import SubagentManager, SubagentStatus from nanobot.agent.subagent import SubagentManager, SubagentStatus
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.providers.base import GenerationSettings
from nanobot.utils.llm_runtime import LLMRuntime
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_subagent_forwards_resolver_to_agent_run_spec(tmp_path: Path) -> None: async def test_subagent_forwards_resolver_to_agent_run_spec(tmp_path: Path) -> None:
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "m" provider.get_default_model.return_value = "m"
provider.generation = GenerationSettings()
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=MessageBus(), bus=MessageBus(),
max_tool_result_chars=64, max_tool_result_chars=64,
@@ -39,6 +41,7 @@ async def test_subagent_forwards_resolver_to_agent_run_spec(tmp_path: Path) -> N
"lbl", "lbl",
{"channel": "cli", "chat_id": "direct", "session_key": "cli:direct"}, {"channel": "cli", "chat_id": "direct", "session_key": "cli:direct"},
status, status,
LLMRuntime.capture(provider, "m", context_window_tokens=128_000),
) )
mgr.runner.run.assert_called_once() mgr.runner.run.assert_called_once()
spec = mgr.runner.run.call_args[0][0] spec = mgr.runner.run.call_args[0][0]
+11 -3
View File
@@ -9,7 +9,8 @@ import pytest
from nanobot.bus.outbound_events import StreamedResponseEvent from nanobot.bus.outbound_events import StreamedResponseEvent
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMResponse, ToolCallRequest from nanobot.providers.base import GenerationSettings, LLMResponse, ToolCallRequest
from nanobot.utils.llm_runtime import LLMRuntime
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
@@ -21,6 +22,7 @@ def _make_loop(tmp_path):
bus = MessageBus() bus = MessageBus()
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
provider.generation = GenerationSettings()
with patch("nanobot.agent.loop.ContextBuilder"), \ with patch("nanobot.agent.loop.ContextBuilder"), \
patch("nanobot.agent.loop.SessionManager"), \ patch("nanobot.agent.loop.SessionManager"), \
@@ -316,7 +318,6 @@ async def test_subagent_max_iterations_announces_existing_fallback(tmp_path, mon
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
)) ))
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -329,7 +330,14 @@ async def test_subagent_max_iterations_announces_existing_fallback(tmp_path, mon
monkeypatch.setattr("nanobot.agent.tools.filesystem.ListDirTool.execute", fake_execute) monkeypatch.setattr("nanobot.agent.tools.filesystem.ListDirTool.execute", fake_execute)
status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()) status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic())
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status) await mgr._run_subagent(
"sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
LLMRuntime.capture(provider, "test-model", context_window_tokens=128_000),
)
mgr._announce_result.assert_awaited_once() mgr._announce_result.assert_awaited_once()
args = mgr._announce_result.await_args.args args = mgr._announce_result.await_args.args
+6
View File
@@ -1098,11 +1098,13 @@ async def test_request_context_uses_effective_key_for_spawn_tool(tmp_path: Path)
spawn_tool = loop.tools.get("spawn") spawn_tool = loop.tools.get("spawn")
assert spawn_tool is not None assert spawn_tool is not None
spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined] spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined]
runtime = loop.llm_runtime()
with request_context(RequestContext( with request_context(RequestContext(
channel="discord", channel="discord",
chat_id="thread-777", chat_id="thread-777",
session_key="discord:parent-456:thread:thread-777", session_key="discord:parent-456:thread:thread-777",
runtime=runtime,
)): )):
await spawn_tool.execute(task="inspect context") await spawn_tool.execute(task="inspect context")
@@ -1110,6 +1112,7 @@ async def test_request_context_uses_effective_key_for_spawn_tool(tmp_path: Path)
assert call["origin_channel"] == "discord" assert call["origin_channel"] == "discord"
assert call["origin_chat_id"] == "thread-777" assert call["origin_chat_id"] == "thread-777"
assert call["session_key"] == "discord:parent-456:thread:thread-777" assert call["session_key"] == "discord:parent-456:thread:thread-777"
assert call["runtime"] is runtime
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -1443,6 +1446,7 @@ async def test_request_context_passes_thread_session_key_to_spawn(tmp_path: Path
spawn_tool = loop.tools.get("spawn") spawn_tool = loop.tools.get("spawn")
assert spawn_tool is not None assert spawn_tool is not None
spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined] spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined]
runtime = loop.llm_runtime()
with request_context(RequestContext( with request_context(RequestContext(
channel="slack", channel="slack",
@@ -1450,12 +1454,14 @@ async def test_request_context_passes_thread_session_key_to_spawn(tmp_path: Path
message_id="msg-123", message_id="msg-123",
metadata={"slack": {"thread_ts": "1700.42", "channel_type": "channel"}}, metadata={"slack": {"thread_ts": "1700.42", "channel_type": "channel"}},
session_key="slack:C123:1700.42", session_key="slack:C123:1700.42",
runtime=runtime,
)): )):
await spawn_tool.execute(task="inspect thread") await spawn_tool.execute(task="inspect thread")
call = spawn_tool._manager.spawn.await_args.kwargs # type: ignore[attr-defined] call = spawn_tool._manager.spawn.await_args.kwargs # type: ignore[attr-defined]
assert call["session_key"] == "slack:C123:1700.42" assert call["session_key"] == "slack:C123:1700.42"
assert call["origin_message_id"] == "msg-123" assert call["origin_message_id"] == "msg-123"
assert call["runtime"] is runtime
@pytest.mark.asyncio @pytest.mark.asyncio
+2 -2
View File
@@ -41,8 +41,8 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
assert loop.model == "new-model" assert loop.model == "new-model"
assert loop.context_window_tokens == 2000 assert loop.context_window_tokens == 2000
assert not hasattr(loop.runner, "provider") assert not hasattr(loop.runner, "provider")
assert loop.subagents.provider is new_provider assert not hasattr(loop.subagents, "provider")
assert loop.subagents.model == "new-model" assert not hasattr(loop.subagents, "model")
assert not hasattr(loop.subagents.runner, "provider") assert not hasattr(loop.subagents.runner, "provider")
assert loop.consolidator.provider is new_provider assert loop.consolidator.provider is new_provider
assert loop.consolidator.model == "new-model" assert loop.consolidator.model == "new-model"
+3 -3
View File
@@ -57,7 +57,7 @@ def test_model_preset_setter_updates_state(tmp_path) -> None:
assert loop.provider.generation.temperature == 0.5 assert loop.provider.generation.temperature == 0.5
assert loop.provider.generation.max_tokens == 4096 assert loop.provider.generation.max_tokens == 4096
assert loop.provider.generation.reasoning_effort == "low" assert loop.provider.generation.reasoning_effort == "low"
assert loop.subagents.model == "openai/gpt-4.1" assert not hasattr(loop.subagents, "model")
assert loop.consolidator.model == "openai/gpt-4.1" assert loop.consolidator.model == "openai/gpt-4.1"
assert loop.consolidator.context_window_tokens == 32_768 assert loop.consolidator.context_window_tokens == 32_768
assert loop.consolidator.max_completion_tokens == 4096 assert loop.consolidator.max_completion_tokens == 4096
@@ -108,7 +108,7 @@ def test_model_preset_setter_replaces_provider_from_snapshot(tmp_path) -> None:
assert loop.provider is new_provider assert loop.provider is new_provider
assert not hasattr(loop.runner, "provider") assert not hasattr(loop.runner, "provider")
assert loop.subagents.provider is new_provider assert not hasattr(loop.subagents, "provider")
assert not hasattr(loop.subagents.runner, "provider") assert not hasattr(loop.subagents.runner, "provider")
assert loop.consolidator.provider is new_provider assert loop.consolidator.provider is new_provider
assert loop.model == "anthropic/claude-opus-4-5" assert loop.model == "anthropic/claude-opus-4-5"
@@ -135,7 +135,7 @@ def test_model_preset_setter_failure_leaves_old_state(tmp_path) -> None:
assert loop.model_preset is None assert loop.model_preset is None
assert loop.model == "base-model" assert loop.model == "base-model"
assert loop.subagents.model == "base-model" assert not hasattr(loop.subagents, "model")
assert loop.consolidator.model == "base-model" assert loop.consolidator.model == "base-model"
assert loop.context_window_tokens == 1000 assert loop.context_window_tokens == 1000
assert loop.consolidator.max_completion_tokens == 123 assert loop.consolidator.max_completion_tokens == 123
+15 -10
View File
@@ -10,7 +10,13 @@ from nanobot.agent.subagent import SubagentManager, SubagentStatus
from nanobot.agent.tools.filesystem import FileToolsConfig from nanobot.agent.tools.filesystem import FileToolsConfig
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import ToolsConfig from nanobot.config.schema import ToolsConfig
from nanobot.providers.base import LLMProvider from nanobot.providers.base import GenerationSettings, LLMProvider
from nanobot.utils.llm_runtime import LLMRuntime
def _runtime(provider: LLMProvider) -> LLMRuntime:
provider.generation = GenerationSettings()
return LLMRuntime.capture(provider, "test", context_window_tokens=128_000)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -19,10 +25,8 @@ async def test_subagent_uses_tool_loader():
provider = MagicMock(spec=LLMProvider) provider = MagicMock(spec=LLMProvider)
provider.get_default_model.return_value = "test" provider.get_default_model.return_value = "test"
sm = SubagentManager( sm = SubagentManager(
provider=provider,
workspace=Path("/tmp"), workspace=Path("/tmp"),
bus=MessageBus(), bus=MessageBus(),
model="test",
max_tool_result_chars=16_000, max_tool_result_chars=16_000,
) )
tools = sm._build_tools() tools = sm._build_tools()
@@ -39,10 +43,8 @@ async def test_subagent_build_tools_isolates_file_read_state(tmp_path):
provider = MagicMock(spec=LLMProvider) provider = MagicMock(spec=LLMProvider)
provider.get_default_model.return_value = "test" provider.get_default_model.return_value = "test"
sm = SubagentManager( sm = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=MessageBus(), bus=MessageBus(),
model="test",
max_tool_result_chars=16_000, max_tool_result_chars=16_000,
) )
@@ -60,10 +62,8 @@ def test_subagent_respects_file_tool_toggle(tmp_path):
provider = MagicMock(spec=LLMProvider) provider = MagicMock(spec=LLMProvider)
provider.get_default_model.return_value = "test" provider.get_default_model.return_value = "test"
sm = SubagentManager( sm = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=MessageBus(), bus=MessageBus(),
model="test",
max_tool_result_chars=16_000, max_tool_result_chars=16_000,
tools_config=ToolsConfig(file=FileToolsConfig(enable=False)), tools_config=ToolsConfig(file=FileToolsConfig(enable=False)),
) )
@@ -87,10 +87,8 @@ async def test_subagent_forwards_fail_on_tool_error_to_runner(tmp_path):
provider = MagicMock(spec=LLMProvider) provider = MagicMock(spec=LLMProvider)
provider.get_default_model.return_value = "test" provider.get_default_model.return_value = "test"
sm = SubagentManager( sm = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=MessageBus(), bus=MessageBus(),
model="test",
max_tool_result_chars=16_000, max_tool_result_chars=16_000,
fail_on_tool_error=False, fail_on_tool_error=False,
) )
@@ -106,7 +104,14 @@ async def test_subagent_forwards_fail_on_tool_error_to_runner(tmp_path):
started_at=0.0, started_at=0.0,
) )
await sm._run_subagent("t1", "task", "label", {"channel": "cli", "chat_id": "direct"}, status) await sm._run_subagent(
"t1",
"task",
"label",
{"channel": "cli", "chat_id": "direct"},
status,
_runtime(provider),
)
spec = sm.runner.run.call_args.args[0] spec = sm.runner.run.call_args.args[0]
assert spec.fail_on_tool_error is False assert spec.fail_on_tool_error is False
+75 -28
View File
@@ -14,27 +14,35 @@ from nanobot.agent.subagent import (
SubagentStatus, SubagentStatus,
_SubagentHook, _SubagentHook,
) )
from nanobot.agent.tools.context import current_request_context
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMProvider from nanobot.providers.base import GenerationSettings, LLMProvider
from nanobot.utils.llm_runtime import LLMRuntime
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Helpers # Helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _manager(tmp_path: Path, **kw) -> SubagentManager: def _manager(tmp_path: Path, **kw) -> SubagentManager:
provider = MagicMock(spec=LLMProvider)
provider.get_default_model.return_value = "test-model"
defaults = dict( defaults = dict(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=MessageBus(), bus=MessageBus(),
model="test-model",
max_tool_result_chars=16_000, max_tool_result_chars=16_000,
) )
defaults.update(kw) defaults.update(kw)
return SubagentManager(**defaults) return SubagentManager(**defaults)
def _runtime(*, model: str = "test-model", temperature: float = 0.1) -> LLMRuntime:
provider = MagicMock(spec=LLMProvider)
provider.generation = GenerationSettings(temperature=temperature, max_tokens=4096)
return LLMRuntime.capture(
provider,
model,
context_window_tokens=128_000,
)
def _make_hook_context(**overrides) -> AgentHookContext: def _make_hook_context(**overrides) -> AgentHookContext:
defaults = dict( defaults = dict(
iteration=1, iteration=1,
@@ -77,17 +85,16 @@ class TestSubagentStatus:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# set_provider # Runtime ownership
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestSetProvider: class TestRuntimeOwnership:
def test_updates_provider_model_runner(self, tmp_path): def test_manager_has_no_provider_model_mirrors(self, tmp_path):
sm = _manager(tmp_path) sm = _manager(tmp_path)
new_provider = MagicMock(spec=LLMProvider) assert not hasattr(sm, "provider")
sm.set_provider(new_provider, "new-model") assert not hasattr(sm, "model")
assert sm.provider is new_provider assert not hasattr(sm, "context_window_tokens")
assert sm.model == "new-model"
assert not hasattr(sm.runner, "provider") assert not hasattr(sm.runner, "provider")
@@ -103,7 +110,7 @@ class TestSpawn:
sm.runner.run = AsyncMock(return_value=AgentRunResult( sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed", final_content="done", messages=[], stop_reason="completed",
)) ))
result = await sm.spawn("do something") result = await sm.spawn("do something", runtime=_runtime())
assert "started" in result assert "started" in result
assert "id:" in result assert "id:" in result
@@ -116,7 +123,7 @@ class TestSpawn:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("task", session_key="s1") await sm.spawn("task", runtime=_runtime(), session_key="s1")
assert len(sm._running_tasks) == 1 assert len(sm._running_tasks) == 1
block.set() block.set()
@@ -129,7 +136,7 @@ class TestSpawn:
sm.runner.run = AsyncMock(return_value=AgentRunResult( sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed", final_content="done", messages=[], stop_reason="completed",
)) ))
await sm.spawn("my task") await sm.spawn("my task", runtime=_runtime())
await _drain_subagent_tasks(sm) await _drain_subagent_tasks(sm)
# Status cleaned up after task completes # Status cleaned up after task completes
assert len(sm._task_statuses) == 0 assert len(sm._task_statuses) == 0
@@ -143,7 +150,7 @@ class TestSpawn:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("task", session_key="s1") await sm.spawn("task", runtime=_runtime(), session_key="s1")
assert "s1" in sm._session_tasks assert "s1" in sm._session_tasks
assert len(sm._session_tasks["s1"]) == 1 assert len(sm._session_tasks["s1"]) == 1
@@ -160,7 +167,7 @@ class TestSpawn:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("task") await sm.spawn("task", runtime=_runtime())
assert len(sm._session_tasks) == 0 assert len(sm._session_tasks) == 0
block.set() block.set()
@@ -176,7 +183,7 @@ class TestSpawn:
sm.runner.run = _slow_run sm.runner.run = _slow_run
long_task = "A" * 50 long_task = "A" * 50
await sm.spawn(long_task, session_key="s1") await sm.spawn(long_task, runtime=_runtime(), session_key="s1")
status = next(iter(sm._task_statuses.values())) status = next(iter(sm._task_statuses.values()))
assert status.label == long_task[:30] + "..." assert status.label == long_task[:30] + "..."
@@ -192,7 +199,9 @@ class TestSpawn:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("task", label="Custom Label", session_key="s1") await sm.spawn(
"task", runtime=_runtime(), label="Custom Label", session_key="s1"
)
status = next(iter(sm._task_statuses.values())) status = next(iter(sm._task_statuses.values()))
assert status.label == "Custom Label" assert status.label == "Custom Label"
@@ -205,12 +214,47 @@ class TestSpawn:
sm.runner.run = AsyncMock(return_value=AgentRunResult( sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed", final_content="done", messages=[], stop_reason="completed",
)) ))
await sm.spawn("task", session_key="s1") await sm.spawn("task", runtime=_runtime(), session_key="s1")
await _drain_subagent_tasks(sm) await _drain_subagent_tasks(sm)
assert len(sm._running_tasks) == 0 assert len(sm._running_tasks) == 0
assert len(sm._task_statuses) == 0 assert len(sm._task_statuses) == 0
assert len(sm._session_tasks) == 0 assert len(sm._session_tasks) == 0
@pytest.mark.asyncio
async def test_runtime_is_captured_before_background_task_starts(self, tmp_path):
sm = _manager(tmp_path)
runtime = _runtime(temperature=0.2)
entered = asyncio.Event()
release = asyncio.Event()
seen: dict[str, object] = {}
async def observe(spec):
seen["spec_runtime"] = spec.runtime
request_ctx = current_request_context()
seen["context_runtime"] = request_ctx.runtime if request_ctx else None
entered.set()
await release.wait()
return AgentRunResult(
final_content="done",
messages=[],
stop_reason="completed",
)
sm.runner.run = observe
await sm.spawn("task", runtime=runtime, session_key="s1")
runtime.provider.generation = GenerationSettings(
temperature=0.9,
max_tokens=128,
)
await asyncio.wait_for(entered.wait(), timeout=1)
assert seen["spec_runtime"] is runtime
assert seen["context_runtime"] is runtime
assert runtime.generation.temperature == 0.2
release.set()
await _drain_subagent_tasks(sm)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# _run_subagent # _run_subagent
@@ -229,6 +273,7 @@ class TestRunSubagent:
"t1", "do task", "label", "t1", "do task", "label",
{"channel": "cli", "chat_id": "direct"}, {"channel": "cli", "chat_id": "direct"},
SubagentStatus(task_id="t1", label="label", task_description="do task", started_at=time.monotonic()), SubagentStatus(task_id="t1", label="label", task_description="do task", started_at=time.monotonic()),
_runtime(),
) )
mock_announce.assert_called_once() mock_announce.assert_called_once()
assert mock_announce.call_args.args[-2] == "ok" assert mock_announce.call_args.args[-2] == "ok"
@@ -244,7 +289,7 @@ class TestRunSubagent:
with patch.object(sm, "_announce_result", new_callable=AsyncMock) as mock_announce: with patch.object(sm, "_announce_result", new_callable=AsyncMock) as mock_announce:
await sm._run_subagent( await sm._run_subagent(
"t1", "do task", "label", "t1", "do task", "label",
{"channel": "cli", "chat_id": "direct"}, status, {"channel": "cli", "chat_id": "direct"}, status, _runtime(),
) )
assert mock_announce.call_args.args[-2] == "error" assert mock_announce.call_args.args[-2] == "error"
@@ -256,7 +301,7 @@ class TestRunSubagent:
with patch.object(sm, "_announce_result", new_callable=AsyncMock) as mock_announce: with patch.object(sm, "_announce_result", new_callable=AsyncMock) as mock_announce:
await sm._run_subagent( await sm._run_subagent(
"t1", "do task", "label", "t1", "do task", "label",
{"channel": "cli", "chat_id": "direct"}, status, {"channel": "cli", "chat_id": "direct"}, status, _runtime(),
) )
assert status.phase == "error" assert status.phase == "error"
assert "LLM down" in status.error assert "LLM down" in status.error
@@ -272,7 +317,7 @@ class TestRunSubagent:
with patch.object(sm, "_announce_result", new_callable=AsyncMock): with patch.object(sm, "_announce_result", new_callable=AsyncMock):
await sm._run_subagent( await sm._run_subagent(
"t1", "do task", "label", "t1", "do task", "label",
{"channel": "cli", "chat_id": "direct"}, status, {"channel": "cli", "chat_id": "direct"}, status, _runtime(),
) )
assert status.phase == "done" assert status.phase == "done"
assert status.stop_reason == "completed" assert status.stop_reason == "completed"
@@ -451,8 +496,9 @@ class TestCancelBySession:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("task1", session_key="s1") runtime = _runtime()
await sm.spawn("task2", session_key="s1") await sm.spawn("task1", runtime=runtime, session_key="s1")
await sm.spawn("task2", runtime=runtime, session_key="s1")
assert len(sm._session_tasks.get("s1", set())) == 2 assert len(sm._session_tasks.get("s1", set())) == 2
count = await sm.cancel_by_session("s1") count = await sm.cancel_by_session("s1")
@@ -472,7 +518,7 @@ class TestCancelBySession:
sm.runner.run = AsyncMock(return_value=AgentRunResult( sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed", final_content="done", messages=[], stop_reason="completed",
)) ))
await sm.spawn("task1", session_key="s1") await sm.spawn("task1", runtime=_runtime(), session_key="s1")
await _drain_subagent_tasks(sm) await _drain_subagent_tasks(sm)
count = await sm.cancel_by_session("s1") count = await sm.cancel_by_session("s1")
@@ -499,8 +545,9 @@ class TestRunningCounts:
return AgentRunResult(final_content="done", messages=[], stop_reason="completed") return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
sm.runner.run = _slow_run sm.runner.run = _slow_run
await sm.spawn("t1", session_key="s1") runtime = _runtime()
await sm.spawn("t2", session_key="s1") await sm.spawn("t1", runtime=runtime, session_key="s1")
await sm.spawn("t2", runtime=runtime, session_key="s1")
assert sm.get_running_count() == 2 assert sm.get_running_count() == 2
assert sm.get_running_count_by_session("s1") == 2 assert sm.get_running_count_by_session("s1") == 2
+34 -16
View File
@@ -10,11 +10,19 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import GenerationSettings
from nanobot.session.keys import UNIFIED_SESSION_KEY from nanobot.session.keys import UNIFIED_SESSION_KEY
from nanobot.utils.llm_runtime import LLMRuntime
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
def _runtime(provider: MagicMock | None = None) -> LLMRuntime:
provider = provider or MagicMock()
provider.generation = GenerationSettings()
return LLMRuntime.capture(provider, "test-model", context_window_tokens=128_000)
def _make_loop(*, tools_config=None): def _make_loop(*, tools_config=None):
"""Create a minimal AgentLoop with mocked dependencies.""" """Create a minimal AgentLoop with mocked dependencies."""
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
@@ -201,10 +209,7 @@ class TestSubagentCancellation:
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=MagicMock(), workspace=MagicMock(),
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -234,10 +239,7 @@ class TestSubagentCancellation:
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=MagicMock(), workspace=MagicMock(),
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -271,7 +273,6 @@ class TestSubagentCancellation:
return LLMResponse(content="done", tool_calls=[]) return LLMResponse(content="done", tool_calls=[])
provider.chat_with_retry = scripted_chat_with_retry provider.chat_with_retry = scripted_chat_with_retry
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -284,7 +285,14 @@ class TestSubagentCancellation:
from nanobot.agent.subagent import SubagentStatus from nanobot.agent.subagent import SubagentStatus
status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()) status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic())
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status) await mgr._run_subagent(
"sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
_runtime(provider),
)
assistant_messages = [ assistant_messages = [
msg for msg in captured_second_call msg for msg in captured_second_call
@@ -305,7 +313,6 @@ class TestSubagentCancellation:
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -326,7 +333,14 @@ class TestSubagentCancellation:
from nanobot.agent.subagent import SubagentStatus from nanobot.agent.subagent import SubagentStatus
status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()) status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic())
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status) await mgr._run_subagent(
"sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
_runtime(provider),
)
mgr.runner.run.assert_awaited_once() mgr.runner.run.assert_awaited_once()
mgr._announce_result.assert_awaited_once() mgr._announce_result.assert_awaited_once()
@@ -345,7 +359,6 @@ class TestSubagentCancellation:
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
)) ))
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -364,7 +377,14 @@ class TestSubagentCancellation:
from nanobot.agent.subagent import SubagentStatus from nanobot.agent.subagent import SubagentStatus
status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()) status = SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic())
await mgr._run_subagent("sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status) await mgr._run_subagent(
"sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
_runtime(provider),
)
mgr._announce_result.assert_awaited_once() mgr._announce_result.assert_awaited_once()
args = mgr._announce_result.await_args.args args = mgr._announce_result.await_args.args
@@ -388,7 +408,6 @@ class TestSubagentCancellation:
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})], tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
)) ))
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -412,6 +431,7 @@ class TestSubagentCancellation:
mgr._run_subagent( mgr._run_subagent(
"sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, "sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"},
SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()), SubagentStatus(task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()),
_runtime(provider),
) )
) )
mgr._running_tasks["sub-1"] = task mgr._running_tasks["sub-1"] = task
@@ -436,10 +456,7 @@ class TestSubagentAnnounceSessionKey:
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=MagicMock(), workspace=MagicMock(),
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -509,6 +526,7 @@ class TestSubagentAnnounceSessionKey:
"sub-4", "task", "label", "sub-4", "task", "label",
{"channel": "telegram", "chat_id": "444", "session_key": UNIFIED_SESSION_KEY}, {"channel": "telegram", "chat_id": "444", "session_key": UNIFIED_SESSION_KEY},
status, status,
_runtime(),
) )
msg = await bus.consume_inbound() msg = await bus.consume_inbound()
+8 -1
View File
@@ -2,10 +2,12 @@ import json
import time import time
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest import pytest
from nanobot.agent.tools.cli_apps import CliAppsTool from nanobot.agent.tools.cli_apps import CliAppsTool
from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.agent.tools.filesystem import ReadFileTool, WriteFileTool from nanobot.agent.tools.filesystem import ReadFileTool, WriteFileTool
from nanobot.agent.tools.image_generation import ImageGenerationError, ImageGenerationTool from nanobot.agent.tools.image_generation import ImageGenerationError, ImageGenerationTool
from nanobot.agent.tools.message import MessageTool from nanobot.agent.tools.message import MessageTool
@@ -362,7 +364,12 @@ async def test_spawn_tool_forwards_current_workspace_scope(tmp_path: Path) -> No
tool = SpawnTool(manager) # type: ignore[arg-type] tool = SpawnTool(manager) # type: ignore[arg-type]
token = bind_workspace_scope(scope) token = bind_workspace_scope(scope)
try: try:
result = await tool.execute(task="inspect") with request_context(RequestContext(
channel="test",
chat_id="chat",
runtime=MagicMock(),
)):
result = await tool.execute(task="inspect")
finally: finally:
reset_workspace_scope(token) reset_workspace_scope(token)
+30 -16
View File
@@ -8,10 +8,17 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import GenerationSettings
from nanobot.utils.llm_runtime import LLMRuntime
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
def _runtime(provider: MagicMock, model: str = "test-model") -> LLMRuntime:
provider.generation = GenerationSettings(temperature=0.1, max_tokens=4096)
return LLMRuntime.capture(provider, model, context_window_tokens=128_000)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_subagent_exec_tool_receives_allowed_env_keys(tmp_path): async def test_subagent_exec_tool_receives_allowed_env_keys(tmp_path):
"""allowed_env_keys from ExecToolConfig must be forwarded to the subagent's ExecTool.""" """allowed_env_keys from ExecToolConfig must be forwarded to the subagent's ExecTool."""
@@ -24,7 +31,6 @@ async def test_subagent_exec_tool_receives_allowed_env_keys(tmp_path):
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -49,7 +55,12 @@ async def test_subagent_exec_tool_receives_allowed_env_keys(tmp_path):
task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic() task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()
) )
await mgr._run_subagent( await mgr._run_subagent(
"sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status "sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
_runtime(provider),
) )
mgr.runner.run.assert_awaited_once() mgr.runner.run.assert_awaited_once()
@@ -65,7 +76,6 @@ async def test_subagent_uses_configured_max_iterations(tmp_path):
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -88,7 +98,12 @@ async def test_subagent_uses_configured_max_iterations(tmp_path):
task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic() task_id="sub-1", label="label", task_description="do task", started_at=time.monotonic()
) )
await mgr._run_subagent( await mgr._run_subagent(
"sub-1", "do task", "label", {"channel": "test", "chat_id": "c1"}, status "sub-1",
"do task",
"label",
{"channel": "test", "chat_id": "c1"},
status,
_runtime(provider),
) )
mgr.runner.run.assert_awaited_once() mgr.runner.run.assert_awaited_once()
@@ -104,27 +119,30 @@ async def test_spawn_forwards_temperature_to_run_spec(tmp_path):
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
) )
mgr._announce_result = AsyncMock() mgr._announce_result = AsyncMock()
parent_runtime = _runtime(provider)
seen = {} seen = {}
async def fake_run(spec): async def fake_run(spec):
seen["temperature"] = spec.runtime.generation.temperature seen["temperature"] = spec.runtime.generation.temperature
seen["runtime"] = spec.runtime
return SimpleNamespace( return SimpleNamespace(
stop_reason="done", final_content="done", error=None, tool_events=[], stop_reason="done", final_content="done", error=None, tool_events=[],
) )
mgr.runner.run = AsyncMock(side_effect=fake_run) mgr.runner.run = AsyncMock(side_effect=fake_run)
await mgr.spawn(task="do task", temperature=0.9) await mgr.spawn(task="do task", runtime=parent_runtime, temperature=0.9)
await asyncio.gather(*mgr._running_tasks.values(), return_exceptions=True) await asyncio.gather(*mgr._running_tasks.values(), return_exceptions=True)
assert seen["temperature"] == 0.9 assert seen["temperature"] == 0.9
assert seen["runtime"] is not parent_runtime
assert parent_runtime.generation.temperature == 0.1
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -138,7 +156,6 @@ async def test_spawn_tool_rejects_when_at_concurrency_limit(tmp_path):
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -162,7 +179,12 @@ async def test_spawn_tool_rejects_when_at_concurrency_limit(tmp_path):
from nanobot.agent.tools.context import RequestContext, request_context from nanobot.agent.tools.context import RequestContext, request_context
tool = SpawnTool(mgr) tool = SpawnTool(mgr)
with request_context(RequestContext(channel="test", chat_id="c1", session_key="test:c1")): with request_context(RequestContext(
channel="test",
chat_id="c1",
session_key="test:c1",
runtime=_runtime(provider),
)):
# First spawn succeeds # First spawn succeeds
result = await tool.execute(task="first task") result = await tool.execute(task="first task")
assert "started" in result assert "started" in result
@@ -184,11 +206,7 @@ def test_subagent_default_max_concurrent_matches_agent_defaults(tmp_path):
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
@@ -203,11 +221,7 @@ def test_subagent_default_max_iterations_matches_agent_defaults(tmp_path):
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
+1 -1
View File
@@ -85,7 +85,7 @@ async def test_model_command_switches_preset(tmp_path) -> None:
assert "Model: `openai/gpt-4.1`" in out.content assert "Model: `openai/gpt-4.1`" in out.content
assert loop.model_preset == "fast" assert loop.model_preset == "fast"
assert loop.model == "openai/gpt-4.1" assert loop.model == "openai/gpt-4.1"
assert loop.subagents.model == "openai/gpt-4.1" assert not hasattr(loop.subagents, "model")
assert loop.consolidator.model == "openai/gpt-4.1" assert loop.consolidator.model == "openai/gpt-4.1"
+33 -7
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
from unittest.mock import MagicMock
import pytest import pytest
@@ -9,7 +10,15 @@ from nanobot.agent.tools.cron import CronTool
from nanobot.agent.tools.message import MessageTool from nanobot.agent.tools.message import MessageTool
from nanobot.agent.tools.spawn import SpawnTool from nanobot.agent.tools.spawn import SpawnTool
from nanobot.cron.service import CronService from nanobot.cron.service import CronService
from nanobot.providers.base import GenerationSettings, LLMProvider
from nanobot.session.keys import UNIFIED_SESSION_KEY from nanobot.session.keys import UNIFIED_SESSION_KEY
from nanobot.utils.llm_runtime import LLMRuntime
def _runtime(model: str = "test-model") -> LLMRuntime:
provider = MagicMock(spec=LLMProvider)
provider.generation = GenerationSettings()
return LLMRuntime.capture(provider, model, context_window_tokens=128_000)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -60,6 +69,7 @@ async def test_spawn_tool_keeps_task_local_context() -> None:
self, self,
*, *,
task: str, task: str,
runtime: LLMRuntime,
label: str | None, label: str | None,
origin_channel: str, origin_channel: str,
origin_chat_id: str, origin_chat_id: str,
@@ -74,14 +84,22 @@ async def test_spawn_tool_keeps_task_local_context() -> None:
tool = SpawnTool(_Manager()) tool = SpawnTool(_Manager())
async def task_one() -> str: async def task_one() -> str:
with request_context(RequestContext(channel="whatsapp", chat_id="chat-a")): with request_context(RequestContext(
channel="whatsapp",
chat_id="chat-a",
runtime=_runtime("model-a"),
)):
entered.set() entered.set()
await release.wait() await release.wait()
return await tool.execute(task="one") return await tool.execute(task="one")
async def task_two() -> str: async def task_two() -> str:
await entered.wait() await entered.wait()
with request_context(RequestContext(channel="telegram", chat_id="chat-b")): with request_context(RequestContext(
channel="telegram",
chat_id="chat-b",
runtime=_runtime("model-b"),
)):
release.set() release.set()
return await tool.execute(task="two") return await tool.execute(task="two")
@@ -182,6 +200,7 @@ async def test_spawn_tool_basic_request_context_and_execute() -> None:
self, self,
*, *,
task, task,
runtime,
label, label,
origin_channel, origin_channel,
origin_chat_id, origin_chat_id,
@@ -194,15 +213,19 @@ async def test_spawn_tool_basic_request_context_and_execute() -> None:
return f"ok: {task}" return f"ok: {task}"
tool = SpawnTool(_Manager()) tool = SpawnTool(_Manager())
with request_context(RequestContext(channel="feishu", chat_id="chat-abc")): with request_context(RequestContext(
channel="feishu",
chat_id="chat-abc",
runtime=_runtime(),
)):
result = await tool.execute(task="do something") result = await tool.execute(task="do something")
assert result == "ok: do something" assert result == "ok: do something"
assert seen == [("feishu", "chat-abc", "feishu:chat-abc")] assert seen == [("feishu", "chat-abc", "feishu:chat-abc")]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_spawn_tool_default_values_without_request_context() -> None: async def test_spawn_tool_rejects_missing_request_runtime() -> None:
"""Without a request context, default cli:direct should be used.""" """Spawning cannot reconstruct a model runtime outside turn admission."""
seen: list[tuple[str, str, str]] = [] seen: list[tuple[str, str, str]] = []
class _Manager: class _Manager:
@@ -215,6 +238,7 @@ async def test_spawn_tool_default_values_without_request_context() -> None:
self, self,
*, *,
task, task,
runtime,
label, label,
origin_channel, origin_channel,
origin_chat_id, origin_chat_id,
@@ -228,8 +252,10 @@ async def test_spawn_tool_default_values_without_request_context() -> None:
tool = SpawnTool(_Manager()) tool = SpawnTool(_Manager())
await tool.execute(task="test") result = await tool.execute(task="test")
assert seen == [("cli", "direct", "cli:direct")] assert result == "Error: spawn requires an active model runtime"
assert result.is_error
assert seen == []
@pytest.mark.asyncio @pytest.mark.asyncio
+11 -5
View File
@@ -16,6 +16,8 @@ from nanobot.agent.tools.search import FindFilesTool, GrepTool
from nanobot.agent.tools.web import WebSearchTool from nanobot.agent.tools.web import WebSearchTool
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import WebSearchConfig from nanobot.config.schema import WebSearchConfig
from nanobot.providers.base import GenerationSettings
from nanobot.utils.llm_runtime import LLMRuntime
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -320,8 +322,8 @@ async def test_subagent_registers_grep(tmp_path: Path) -> None:
bus = MessageBus() bus = MessageBus()
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
provider.generation = GenerationSettings()
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=4096, max_tool_result_chars=4096,
@@ -341,7 +343,14 @@ async def test_subagent_registers_grep(tmp_path: Path) -> None:
mgr._announce_result = AsyncMock() mgr._announce_result = AsyncMock()
status = SubagentStatus(task_id="sub-1", label="label", task_description="search task", started_at=time.monotonic()) status = SubagentStatus(task_id="sub-1", label="label", task_description="search task", started_at=time.monotonic())
await mgr._run_subagent("sub-1", "search task", "label", {"channel": "cli", "chat_id": "direct"}, status) await mgr._run_subagent(
"sub-1",
"search task",
"label",
{"channel": "cli", "chat_id": "direct"},
status,
LLMRuntime.capture(provider, "test-model", context_window_tokens=128_000),
)
assert "find_files" in captured["tool_names"] assert "find_files" in captured["tool_names"]
assert "grep" in captured["tool_names"] assert "grep" in captured["tool_names"]
@@ -349,8 +358,6 @@ async def test_subagent_registers_grep(tmp_path: Path) -> None:
def test_subagent_prompt_respects_disabled_skills(tmp_path: Path) -> None: def test_subagent_prompt_respects_disabled_skills(tmp_path: Path) -> None:
bus = MessageBus() bus = MessageBus()
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
skills_dir = tmp_path / "skills" skills_dir = tmp_path / "skills"
(skills_dir / "alpha").mkdir(parents=True) (skills_dir / "alpha").mkdir(parents=True)
(skills_dir / "alpha" / "SKILL.md").write_text("# Alpha\n\nhidden\n", encoding="utf-8") (skills_dir / "alpha" / "SKILL.md").write_text("# Alpha\n\nhidden\n", encoding="utf-8")
@@ -358,7 +365,6 @@ def test_subagent_prompt_respects_disabled_skills(tmp_path: Path) -> None:
(skills_dir / "beta" / "SKILL.md").write_text("# Beta\n\nshown\n", encoding="utf-8") (skills_dir / "beta" / "SKILL.md").write_text("# Beta\n\nshown\n", encoding="utf-8")
mgr = SubagentManager( mgr = SubagentManager(
provider=provider,
workspace=tmp_path, workspace=tmp_path,
bus=bus, bus=bus,
max_tool_result_chars=4096, max_tool_result_chars=4096,