refactor(agent): capture subagent runtime before spawn
This commit is contained in:
@@ -340,11 +340,8 @@ class AgentLoop:
|
||||
self._file_state_store = FileStateStore()
|
||||
self.runner = AgentRunner()
|
||||
self.subagents = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=workspace,
|
||||
bus=bus,
|
||||
model=self.model,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
tools_config=_tc,
|
||||
max_tool_result_chars=self.max_tool_result_chars,
|
||||
restrict_to_workspace=restrict_to_workspace,
|
||||
@@ -495,7 +492,6 @@ class AgentLoop:
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
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._sync_replay_max_messages()
|
||||
self._provider_signature = snapshot.signature
|
||||
|
||||
+20
-26
@@ -12,14 +12,18 @@ from loguru import logger
|
||||
|
||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||
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.loader import ToolLoader
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.config.schema import AgentDefaults, ToolsConfig
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.security.workspace_access import (
|
||||
WorkspaceScope,
|
||||
bind_workspace_scope,
|
||||
@@ -77,12 +81,9 @@ class SubagentManager:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: LLMProvider,
|
||||
workspace: Path,
|
||||
bus: MessageBus,
|
||||
max_tool_result_chars: int,
|
||||
model: str | None = None,
|
||||
context_window_tokens: int | None = None,
|
||||
tools_config: ToolsConfig | None = None,
|
||||
restrict_to_workspace: bool = False,
|
||||
disabled_skills: list[str] | None = None,
|
||||
@@ -92,11 +93,8 @@ class SubagentManager:
|
||||
llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None,
|
||||
):
|
||||
defaults = AgentDefaults()
|
||||
self.provider = provider
|
||||
self.workspace = workspace
|
||||
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.max_tool_result_chars = max_tool_result_chars
|
||||
self.restrict_to_workspace = restrict_to_workspace
|
||||
@@ -152,17 +150,6 @@ class SubagentManager:
|
||||
ToolLoader().load(ctx, registry, scope="subagent")
|
||||
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(
|
||||
self,
|
||||
task: str,
|
||||
@@ -173,8 +160,12 @@ class SubagentManager:
|
||||
origin_message_id: str | None = None,
|
||||
temperature: float | None = None,
|
||||
workspace_scope: WorkspaceScope | None = None,
|
||||
*,
|
||||
runtime: LLMRuntime,
|
||||
) -> str:
|
||||
"""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]
|
||||
display_label = label or task[:30] + ("..." if len(task) > 30 else "")
|
||||
origin = {"channel": origin_channel, "chat_id": origin_chat_id, "session_key": session_key}
|
||||
@@ -194,8 +185,8 @@ class SubagentManager:
|
||||
display_label,
|
||||
origin,
|
||||
status,
|
||||
runtime,
|
||||
origin_message_id,
|
||||
temperature,
|
||||
workspace_scope,
|
||||
)
|
||||
)
|
||||
@@ -223,8 +214,8 @@ class SubagentManager:
|
||||
label: str,
|
||||
origin: dict[str, str],
|
||||
status: SubagentStatus,
|
||||
runtime: LLMRuntime,
|
||||
origin_message_id: str | None = None,
|
||||
temperature: float | None = None,
|
||||
workspace_scope: WorkspaceScope | None = None,
|
||||
) -> None:
|
||||
"""Execute the subagent task and announce the result."""
|
||||
@@ -253,13 +244,15 @@ class SubagentManager:
|
||||
if self._llm_wall_timeout_for_session
|
||||
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
|
||||
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(
|
||||
initial_messages=messages,
|
||||
tools=tools,
|
||||
@@ -279,6 +272,7 @@ class SubagentManager:
|
||||
finally:
|
||||
if token is not None:
|
||||
reset_workspace_scope(token)
|
||||
reset_request_context(request_token)
|
||||
status.phase = "done"
|
||||
status.stop_reason = result.stop_reason
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
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.schema import NumberSchema, StringSchema, tool_parameters_schema
|
||||
from nanobot.security.workspace_access import current_workspace_scope
|
||||
@@ -70,20 +70,19 @@ class SpawnTool(Tool):
|
||||
f"to complete before spawning a new one."
|
||||
)
|
||||
request_ctx = current_request_context()
|
||||
origin_channel = request_ctx.channel if request_ctx is not None else "cli"
|
||||
origin_chat_id = request_ctx.chat_id if request_ctx is not None else "direct"
|
||||
session_key = (
|
||||
request_ctx.session_key or f"{origin_channel}:{origin_chat_id}"
|
||||
if request_ctx is not None
|
||||
else "cli:direct"
|
||||
)
|
||||
if request_ctx is None or request_ctx.runtime is None:
|
||||
return ToolResult.error("Error: spawn requires an active model runtime")
|
||||
origin_channel = request_ctx.channel
|
||||
origin_chat_id = request_ctx.chat_id
|
||||
session_key = request_ctx.session_key or f"{origin_channel}:{origin_chat_id}"
|
||||
return await self._manager.spawn(
|
||||
task=task,
|
||||
runtime=request_ctx.runtime,
|
||||
label=label,
|
||||
origin_channel=origin_channel,
|
||||
origin_chat_id=origin_chat_id,
|
||||
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,
|
||||
workspace_scope=current_workspace_scope(),
|
||||
)
|
||||
|
||||
@@ -8,14 +8,16 @@ import pytest
|
||||
from nanobot.agent.runner import AgentRunResult
|
||||
from nanobot.agent.subagent import SubagentManager, SubagentStatus
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import GenerationSettings
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_forwards_resolver_to_agent_run_spec(tmp_path: Path) -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "m"
|
||||
provider.generation = GenerationSettings()
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
max_tool_result_chars=64,
|
||||
@@ -39,6 +41,7 @@ async def test_subagent_forwards_resolver_to_agent_run_spec(tmp_path: Path) -> N
|
||||
"lbl",
|
||||
{"channel": "cli", "chat_id": "direct", "session_key": "cli:direct"},
|
||||
status,
|
||||
LLMRuntime.capture(provider, "m", context_window_tokens=128_000),
|
||||
)
|
||||
mgr.runner.run.assert_called_once()
|
||||
spec = mgr.runner.run.call_args[0][0]
|
||||
|
||||
@@ -9,7 +9,8 @@ import pytest
|
||||
|
||||
from nanobot.bus.outbound_events import StreamedResponseEvent
|
||||
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
|
||||
|
||||
@@ -21,6 +22,7 @@ def _make_loop(tmp_path):
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
provider.generation = GenerationSettings()
|
||||
|
||||
with patch("nanobot.agent.loop.ContextBuilder"), \
|
||||
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": "."})],
|
||||
))
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
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)
|
||||
|
||||
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()
|
||||
args = mgr._announce_result.await_args.args
|
||||
|
||||
@@ -1098,11 +1098,13 @@ async def test_request_context_uses_effective_key_for_spawn_tool(tmp_path: Path)
|
||||
spawn_tool = loop.tools.get("spawn")
|
||||
assert spawn_tool is not None
|
||||
spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined]
|
||||
runtime = loop.llm_runtime()
|
||||
|
||||
with request_context(RequestContext(
|
||||
channel="discord",
|
||||
chat_id="thread-777",
|
||||
session_key="discord:parent-456:thread:thread-777",
|
||||
runtime=runtime,
|
||||
)):
|
||||
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_chat_id"] == "thread-777"
|
||||
assert call["session_key"] == "discord:parent-456:thread:thread-777"
|
||||
assert call["runtime"] is runtime
|
||||
|
||||
|
||||
@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")
|
||||
assert spawn_tool is not None
|
||||
spawn_tool._manager.spawn = AsyncMock(return_value="started") # type: ignore[attr-defined]
|
||||
runtime = loop.llm_runtime()
|
||||
|
||||
with request_context(RequestContext(
|
||||
channel="slack",
|
||||
@@ -1450,12 +1454,14 @@ async def test_request_context_passes_thread_session_key_to_spawn(tmp_path: Path
|
||||
message_id="msg-123",
|
||||
metadata={"slack": {"thread_ts": "1700.42", "channel_type": "channel"}},
|
||||
session_key="slack:C123:1700.42",
|
||||
runtime=runtime,
|
||||
)):
|
||||
await spawn_tool.execute(task="inspect thread")
|
||||
|
||||
call = spawn_tool._manager.spawn.await_args.kwargs # type: ignore[attr-defined]
|
||||
assert call["session_key"] == "slack:C123:1700.42"
|
||||
assert call["origin_message_id"] == "msg-123"
|
||||
assert call["runtime"] is runtime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -41,8 +41,8 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
|
||||
assert loop.model == "new-model"
|
||||
assert loop.context_window_tokens == 2000
|
||||
assert not hasattr(loop.runner, "provider")
|
||||
assert loop.subagents.provider is new_provider
|
||||
assert loop.subagents.model == "new-model"
|
||||
assert not hasattr(loop.subagents, "provider")
|
||||
assert not hasattr(loop.subagents, "model")
|
||||
assert not hasattr(loop.subagents.runner, "provider")
|
||||
assert loop.consolidator.provider is new_provider
|
||||
assert loop.consolidator.model == "new-model"
|
||||
|
||||
@@ -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.max_tokens == 4096
|
||||
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.context_window_tokens == 32_768
|
||||
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 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 loop.consolidator.provider is new_provider
|
||||
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 == "base-model"
|
||||
assert loop.subagents.model == "base-model"
|
||||
assert not hasattr(loop.subagents, "model")
|
||||
assert loop.consolidator.model == "base-model"
|
||||
assert loop.context_window_tokens == 1000
|
||||
assert loop.consolidator.max_completion_tokens == 123
|
||||
|
||||
@@ -10,7 +10,13 @@ from nanobot.agent.subagent import SubagentManager, SubagentStatus
|
||||
from nanobot.agent.tools.filesystem import FileToolsConfig
|
||||
from nanobot.bus.queue import MessageBus
|
||||
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
|
||||
@@ -19,10 +25,8 @@ async def test_subagent_uses_tool_loader():
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.get_default_model.return_value = "test"
|
||||
sm = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=Path("/tmp"),
|
||||
bus=MessageBus(),
|
||||
model="test",
|
||||
max_tool_result_chars=16_000,
|
||||
)
|
||||
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.get_default_model.return_value = "test"
|
||||
sm = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
model="test",
|
||||
max_tool_result_chars=16_000,
|
||||
)
|
||||
|
||||
@@ -60,10 +62,8 @@ def test_subagent_respects_file_tool_toggle(tmp_path):
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.get_default_model.return_value = "test"
|
||||
sm = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
model="test",
|
||||
max_tool_result_chars=16_000,
|
||||
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.get_default_model.return_value = "test"
|
||||
sm = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
model="test",
|
||||
max_tool_result_chars=16_000,
|
||||
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,
|
||||
)
|
||||
|
||||
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]
|
||||
assert spec.fail_on_tool_error is False
|
||||
|
||||
@@ -14,27 +14,35 @@ from nanobot.agent.subagent import (
|
||||
SubagentStatus,
|
||||
_SubagentHook,
|
||||
)
|
||||
from nanobot.agent.tools.context import current_request_context
|
||||
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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _manager(tmp_path: Path, **kw) -> SubagentManager:
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
defaults = dict(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
model="test-model",
|
||||
max_tool_result_chars=16_000,
|
||||
)
|
||||
defaults.update(kw)
|
||||
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:
|
||||
defaults = dict(
|
||||
iteration=1,
|
||||
@@ -77,17 +85,16 @@ class TestSubagentStatus:
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# set_provider
|
||||
# Runtime ownership
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSetProvider:
|
||||
def test_updates_provider_model_runner(self, tmp_path):
|
||||
class TestRuntimeOwnership:
|
||||
def test_manager_has_no_provider_model_mirrors(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
new_provider = MagicMock(spec=LLMProvider)
|
||||
sm.set_provider(new_provider, "new-model")
|
||||
assert sm.provider is new_provider
|
||||
assert sm.model == "new-model"
|
||||
assert not hasattr(sm, "provider")
|
||||
assert not hasattr(sm, "model")
|
||||
assert not hasattr(sm, "context_window_tokens")
|
||||
assert not hasattr(sm.runner, "provider")
|
||||
|
||||
|
||||
@@ -103,7 +110,7 @@ class TestSpawn:
|
||||
sm.runner.run = AsyncMock(return_value=AgentRunResult(
|
||||
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 "id:" in result
|
||||
|
||||
@@ -116,7 +123,7 @@ class TestSpawn:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
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
|
||||
|
||||
block.set()
|
||||
@@ -129,7 +136,7 @@ class TestSpawn:
|
||||
sm.runner.run = AsyncMock(return_value=AgentRunResult(
|
||||
final_content="done", messages=[], stop_reason="completed",
|
||||
))
|
||||
await sm.spawn("my task")
|
||||
await sm.spawn("my task", runtime=_runtime())
|
||||
await _drain_subagent_tasks(sm)
|
||||
# Status cleaned up after task completes
|
||||
assert len(sm._task_statuses) == 0
|
||||
@@ -143,7 +150,7 @@ class TestSpawn:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
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 len(sm._session_tasks["s1"]) == 1
|
||||
|
||||
@@ -160,7 +167,7 @@ class TestSpawn:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
sm.runner.run = _slow_run
|
||||
|
||||
await sm.spawn("task")
|
||||
await sm.spawn("task", runtime=_runtime())
|
||||
assert len(sm._session_tasks) == 0
|
||||
|
||||
block.set()
|
||||
@@ -176,7 +183,7 @@ class TestSpawn:
|
||||
sm.runner.run = _slow_run
|
||||
|
||||
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()))
|
||||
assert status.label == long_task[:30] + "..."
|
||||
|
||||
@@ -192,7 +199,9 @@ class TestSpawn:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
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()))
|
||||
assert status.label == "Custom Label"
|
||||
|
||||
@@ -205,12 +214,47 @@ class TestSpawn:
|
||||
sm.runner.run = AsyncMock(return_value=AgentRunResult(
|
||||
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)
|
||||
assert len(sm._running_tasks) == 0
|
||||
assert len(sm._task_statuses) == 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
|
||||
@@ -229,6 +273,7 @@ class TestRunSubagent:
|
||||
"t1", "do task", "label",
|
||||
{"channel": "cli", "chat_id": "direct"},
|
||||
SubagentStatus(task_id="t1", label="label", task_description="do task", started_at=time.monotonic()),
|
||||
_runtime(),
|
||||
)
|
||||
mock_announce.assert_called_once()
|
||||
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:
|
||||
await sm._run_subagent(
|
||||
"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"
|
||||
|
||||
@@ -256,7 +301,7 @@ class TestRunSubagent:
|
||||
with patch.object(sm, "_announce_result", new_callable=AsyncMock) as mock_announce:
|
||||
await sm._run_subagent(
|
||||
"t1", "do task", "label",
|
||||
{"channel": "cli", "chat_id": "direct"}, status,
|
||||
{"channel": "cli", "chat_id": "direct"}, status, _runtime(),
|
||||
)
|
||||
assert status.phase == "error"
|
||||
assert "LLM down" in status.error
|
||||
@@ -272,7 +317,7 @@ class TestRunSubagent:
|
||||
with patch.object(sm, "_announce_result", new_callable=AsyncMock):
|
||||
await sm._run_subagent(
|
||||
"t1", "do task", "label",
|
||||
{"channel": "cli", "chat_id": "direct"}, status,
|
||||
{"channel": "cli", "chat_id": "direct"}, status, _runtime(),
|
||||
)
|
||||
assert status.phase == "done"
|
||||
assert status.stop_reason == "completed"
|
||||
@@ -451,8 +496,9 @@ class TestCancelBySession:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
sm.runner.run = _slow_run
|
||||
|
||||
await sm.spawn("task1", session_key="s1")
|
||||
await sm.spawn("task2", session_key="s1")
|
||||
runtime = _runtime()
|
||||
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
|
||||
|
||||
count = await sm.cancel_by_session("s1")
|
||||
@@ -472,7 +518,7 @@ class TestCancelBySession:
|
||||
sm.runner.run = AsyncMock(return_value=AgentRunResult(
|
||||
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)
|
||||
|
||||
count = await sm.cancel_by_session("s1")
|
||||
@@ -499,8 +545,9 @@ class TestRunningCounts:
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
sm.runner.run = _slow_run
|
||||
|
||||
await sm.spawn("t1", session_key="s1")
|
||||
await sm.spawn("t2", session_key="s1")
|
||||
runtime = _runtime()
|
||||
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_by_session("s1") == 2
|
||||
|
||||
|
||||
@@ -10,11 +10,19 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
from nanobot.providers.base import GenerationSettings
|
||||
from nanobot.session.keys import UNIFIED_SESSION_KEY
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
_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):
|
||||
"""Create a minimal AgentLoop with mocked dependencies."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
@@ -201,10 +209,7 @@ class TestSubagentCancellation:
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=MagicMock(),
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -234,10 +239,7 @@ class TestSubagentCancellation:
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=MagicMock(),
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -271,7 +273,6 @@ class TestSubagentCancellation:
|
||||
return LLMResponse(content="done", tool_calls=[])
|
||||
provider.chat_with_retry = scripted_chat_with_retry
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -284,7 +285,14 @@ class TestSubagentCancellation:
|
||||
|
||||
from nanobot.agent.subagent import SubagentStatus
|
||||
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 = [
|
||||
msg for msg in captured_second_call
|
||||
@@ -305,7 +313,6 @@ class TestSubagentCancellation:
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -326,7 +333,14 @@ class TestSubagentCancellation:
|
||||
|
||||
from nanobot.agent.subagent import SubagentStatus
|
||||
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._announce_result.assert_awaited_once()
|
||||
@@ -345,7 +359,6 @@ class TestSubagentCancellation:
|
||||
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
|
||||
))
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -364,7 +377,14 @@ class TestSubagentCancellation:
|
||||
|
||||
from nanobot.agent.subagent import SubagentStatus
|
||||
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()
|
||||
args = mgr._announce_result.await_args.args
|
||||
@@ -388,7 +408,6 @@ class TestSubagentCancellation:
|
||||
tool_calls=[ToolCallRequest(id="call_1", name="list_dir", arguments={"path": "."})],
|
||||
))
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -412,6 +431,7 @@ class TestSubagentCancellation:
|
||||
mgr._run_subagent(
|
||||
"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()),
|
||||
_runtime(provider),
|
||||
)
|
||||
)
|
||||
mgr._running_tasks["sub-1"] = task
|
||||
@@ -436,10 +456,7 @@ class TestSubagentAnnounceSessionKey:
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=MagicMock(),
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
@@ -509,6 +526,7 @@ class TestSubagentAnnounceSessionKey:
|
||||
"sub-4", "task", "label",
|
||||
{"channel": "telegram", "chat_id": "444", "session_key": UNIFIED_SESSION_KEY},
|
||||
status,
|
||||
_runtime(),
|
||||
)
|
||||
|
||||
msg = await bus.consume_inbound()
|
||||
|
||||
@@ -2,10 +2,12 @@ import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
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.image_generation import ImageGenerationError, ImageGenerationTool
|
||||
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]
|
||||
token = bind_workspace_scope(scope)
|
||||
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:
|
||||
reset_workspace_scope(token)
|
||||
|
||||
|
||||
@@ -8,10 +8,17 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
|
||||
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
|
||||
|
||||
|
||||
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
|
||||
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."""
|
||||
@@ -24,7 +31,6 @@ async def test_subagent_exec_tool_receives_allowed_env_keys(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
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()
|
||||
)
|
||||
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()
|
||||
@@ -65,7 +76,6 @@ async def test_subagent_uses_configured_max_iterations(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
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()
|
||||
)
|
||||
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()
|
||||
@@ -104,27 +119,30 @@ async def test_spawn_forwards_temperature_to_run_spec(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
)
|
||||
mgr._announce_result = AsyncMock()
|
||||
|
||||
parent_runtime = _runtime(provider)
|
||||
seen = {}
|
||||
|
||||
async def fake_run(spec):
|
||||
seen["temperature"] = spec.runtime.generation.temperature
|
||||
seen["runtime"] = spec.runtime
|
||||
return SimpleNamespace(
|
||||
stop_reason="done", final_content="done", error=None, tool_events=[],
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
assert seen["temperature"] == 0.9
|
||||
assert seen["runtime"] is not parent_runtime
|
||||
assert parent_runtime.generation.temperature == 0.1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -138,7 +156,6 @@ async def test_spawn_tool_rejects_when_at_concurrency_limit(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
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
|
||||
|
||||
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
|
||||
result = await tool.execute(task="first task")
|
||||
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
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
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
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
|
||||
@@ -85,7 +85,7 @@ async def test_model_command_switches_preset(tmp_path) -> None:
|
||||
assert "Model: `openai/gpt-4.1`" in out.content
|
||||
assert loop.model_preset == "fast"
|
||||
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"
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -9,7 +10,15 @@ from nanobot.agent.tools.cron import CronTool
|
||||
from nanobot.agent.tools.message import MessageTool
|
||||
from nanobot.agent.tools.spawn import SpawnTool
|
||||
from nanobot.cron.service import CronService
|
||||
from nanobot.providers.base import GenerationSettings, LLMProvider
|
||||
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
|
||||
@@ -60,6 +69,7 @@ async def test_spawn_tool_keeps_task_local_context() -> None:
|
||||
self,
|
||||
*,
|
||||
task: str,
|
||||
runtime: LLMRuntime,
|
||||
label: str | None,
|
||||
origin_channel: str,
|
||||
origin_chat_id: str,
|
||||
@@ -74,14 +84,22 @@ async def test_spawn_tool_keeps_task_local_context() -> None:
|
||||
tool = SpawnTool(_Manager())
|
||||
|
||||
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()
|
||||
await release.wait()
|
||||
return await tool.execute(task="one")
|
||||
|
||||
async def task_two() -> str:
|
||||
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()
|
||||
return await tool.execute(task="two")
|
||||
|
||||
@@ -182,6 +200,7 @@ async def test_spawn_tool_basic_request_context_and_execute() -> None:
|
||||
self,
|
||||
*,
|
||||
task,
|
||||
runtime,
|
||||
label,
|
||||
origin_channel,
|
||||
origin_chat_id,
|
||||
@@ -194,15 +213,19 @@ async def test_spawn_tool_basic_request_context_and_execute() -> None:
|
||||
return f"ok: {task}"
|
||||
|
||||
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")
|
||||
assert result == "ok: do something"
|
||||
assert seen == [("feishu", "chat-abc", "feishu:chat-abc")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spawn_tool_default_values_without_request_context() -> None:
|
||||
"""Without a request context, default cli:direct should be used."""
|
||||
async def test_spawn_tool_rejects_missing_request_runtime() -> None:
|
||||
"""Spawning cannot reconstruct a model runtime outside turn admission."""
|
||||
seen: list[tuple[str, str, str]] = []
|
||||
|
||||
class _Manager:
|
||||
@@ -215,6 +238,7 @@ async def test_spawn_tool_default_values_without_request_context() -> None:
|
||||
self,
|
||||
*,
|
||||
task,
|
||||
runtime,
|
||||
label,
|
||||
origin_channel,
|
||||
origin_chat_id,
|
||||
@@ -228,8 +252,10 @@ async def test_spawn_tool_default_values_without_request_context() -> None:
|
||||
|
||||
tool = SpawnTool(_Manager())
|
||||
|
||||
await tool.execute(task="test")
|
||||
assert seen == [("cli", "direct", "cli:direct")]
|
||||
result = await tool.execute(task="test")
|
||||
assert result == "Error: spawn requires an active model runtime"
|
||||
assert result.is_error
|
||||
assert seen == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -16,6 +16,8 @@ from nanobot.agent.tools.search import FindFilesTool, GrepTool
|
||||
from nanobot.agent.tools.web import WebSearchTool
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.config.schema import WebSearchConfig
|
||||
from nanobot.providers.base import GenerationSettings
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -320,8 +322,8 @@ async def test_subagent_registers_grep(tmp_path: Path) -> None:
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
provider.generation = GenerationSettings()
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=4096,
|
||||
@@ -341,7 +343,14 @@ async def test_subagent_registers_grep(tmp_path: Path) -> None:
|
||||
mgr._announce_result = AsyncMock()
|
||||
|
||||
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 "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:
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
skills_dir = tmp_path / "skills"
|
||||
(skills_dir / "alpha").mkdir(parents=True)
|
||||
(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")
|
||||
|
||||
mgr = SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=bus,
|
||||
max_tool_result_chars=4096,
|
||||
|
||||
Reference in New Issue
Block a user