mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 13:28:43 +03:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
43d592f8e4 | ||
|
|
4a6853f0ff | ||
|
|
3b03cc2079 |
@@ -29,7 +29,7 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
|
||||
|
||||
def runtime_lines(state: Any, msg: Any, workspace: Path, *, skip: bool = False) -> list[str]:
|
||||
"""Return model-visible runtime annotations for turn-attached capabilities."""
|
||||
return [
|
||||
lines = [
|
||||
*cli_app_utils.runtime_lines(msg, workspace, skip=skip),
|
||||
*mcp_tools.runtime_lines(
|
||||
msg,
|
||||
@@ -38,6 +38,11 @@ def runtime_lines(state: Any, msg: Any, workspace: Path, *, skip: bool = False)
|
||||
skip=skip,
|
||||
),
|
||||
]
|
||||
if not skip and getattr(state, "subagents", None) is not None:
|
||||
session_key = getattr(msg, "session_key", None)
|
||||
if session_key:
|
||||
lines.extend(state.subagents.runtime_status_lines(session_key))
|
||||
return lines
|
||||
|
||||
|
||||
async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
|
||||
|
||||
+35
-27
@@ -25,6 +25,10 @@ from nanobot.agent.memory import Consolidator
|
||||
from nanobot.agent.progress_hook import AgentProgressHook
|
||||
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.agent.subagent_delivery import (
|
||||
build_subagent_result_continuation,
|
||||
materialize_subagent_result_continuation,
|
||||
)
|
||||
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
|
||||
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
|
||||
from nanobot.agent.tools.message import MessageTool
|
||||
@@ -287,6 +291,7 @@ class AgentLoop:
|
||||
max_iterations=self.max_iterations,
|
||||
max_concurrent_subagents=max_concurrent_subagents,
|
||||
llm_wall_timeout_for_session=lambda sk: runner_wall_llm_timeout_s(self.sessions, sk),
|
||||
on_result_ready=self._on_subagent_result_ready,
|
||||
)
|
||||
self._unified_session = unified_session
|
||||
self._max_messages = max_messages if max_messages > 0 else 120
|
||||
@@ -548,6 +553,21 @@ class AgentLoop:
|
||||
"""Build a progress callback that publishes to the message bus."""
|
||||
return build_bus_progress_callback(self.bus, msg)
|
||||
|
||||
async def _on_subagent_result_ready(self, result: Any) -> None:
|
||||
"""Wake the owning session when a subagent result becomes ready."""
|
||||
msg = build_subagent_result_continuation(result)
|
||||
queue = self._pending_queues.get(result.session_key)
|
||||
if queue is not None:
|
||||
try:
|
||||
queue.put_nowait(msg)
|
||||
return
|
||||
except asyncio.QueueFull:
|
||||
logger.warning(
|
||||
"Pending queue full for subagent result in session {}; queueing fresh turn",
|
||||
result.session_key,
|
||||
)
|
||||
await self.bus.publish_inbound(msg)
|
||||
|
||||
async def _build_retry_wait_callback(
|
||||
self, msg: InboundMessage
|
||||
) -> Callable[[str], Awaitable[None]]:
|
||||
@@ -731,11 +751,9 @@ class AgentLoop:
|
||||
async def _drain_pending(*, limit: int = _MAX_INJECTIONS_PER_TURN) -> list[dict[str, Any]]:
|
||||
"""Drain follow-up messages from the pending queue.
|
||||
|
||||
When no messages are immediately available but sub-agents
|
||||
spawned in this dispatch are still running, blocks until at
|
||||
least one result arrives (or timeout). This keeps the runner
|
||||
loop alive so subsequent sub-agent completions are consumed
|
||||
in-order rather than dispatched separately.
|
||||
This path is only for real same-session user follow-up messages.
|
||||
Worker results are read explicitly through the subagent mailbox
|
||||
tools instead of being injected as ordinary inbound messages.
|
||||
"""
|
||||
if pending_queue is None:
|
||||
return []
|
||||
@@ -752,30 +770,15 @@ class AgentLoop:
|
||||
items: list[dict[str, Any]] = []
|
||||
while len(items) < limit:
|
||||
try:
|
||||
items.append(_to_user_message(pending_queue.get_nowait()))
|
||||
pending_msg = pending_queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
# Block if nothing drained but sub-agents spawned in this dispatch
|
||||
# are still running. Keeps the runner loop alive so subsequent
|
||||
# completions are injected in-order rather than dispatched separately.
|
||||
if (not items
|
||||
and session is not None
|
||||
and self.subagents.get_running_count_by_session(session.key) > 0):
|
||||
try:
|
||||
msg = await asyncio.wait_for(pending_queue.get(), timeout=300)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
"Timeout waiting for sub-agent completion in session {}",
|
||||
session.key,
|
||||
)
|
||||
return items
|
||||
items.append(_to_user_message(msg))
|
||||
while len(items) < limit:
|
||||
try:
|
||||
items.append(_to_user_message(pending_queue.get_nowait()))
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
pending_msg = await materialize_subagent_result_continuation(
|
||||
pending_msg,
|
||||
session_key=active_session_key or pending_msg.session_key,
|
||||
subagents=self.subagents,
|
||||
)
|
||||
items.append(_to_user_message(pending_msg))
|
||||
|
||||
return items
|
||||
|
||||
@@ -1432,6 +1435,11 @@ class AgentLoop:
|
||||
ctx.session,
|
||||
replay_max_messages=self._max_messages,
|
||||
)
|
||||
ctx.msg = await materialize_subagent_result_continuation(
|
||||
ctx.msg,
|
||||
session_key=ctx.session_key,
|
||||
subagents=self.subagents,
|
||||
)
|
||||
self._set_tool_context(
|
||||
ctx.msg.channel,
|
||||
ctx.msg.chat_id,
|
||||
|
||||
@@ -0,0 +1,415 @@
|
||||
"""Durable mailbox primitives for manager-worker task coordination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from nanobot.utils.helpers import ensure_dir, safe_filename
|
||||
|
||||
TaskState = str # running | completed | failed | cancelled
|
||||
MailboxReadState = str # ready | running | not_found | consumed | timeout
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskRequest:
|
||||
"""Task request recorded when the manager dispatches a worker."""
|
||||
|
||||
task_id: str
|
||||
session_key: str
|
||||
label: str
|
||||
task: str
|
||||
origin: dict[str, Any] = field(default_factory=dict)
|
||||
created_at: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskResult:
|
||||
"""Worker result written to the manager mailbox."""
|
||||
|
||||
task_id: str
|
||||
session_key: str
|
||||
label: str
|
||||
task: str
|
||||
status: str
|
||||
content: str
|
||||
sender: str = "subagent"
|
||||
completed_at: float = field(default_factory=time.time)
|
||||
dedupe_key: str | None = None
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskSnapshot:
|
||||
"""Read-only view of a task in the mailbox."""
|
||||
|
||||
task_id: str
|
||||
session_key: str
|
||||
label: str
|
||||
task: str
|
||||
state: TaskState
|
||||
created_at: float
|
||||
completed_at: float | None = None
|
||||
consumed_at: float | None = None
|
||||
result_status: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class MailboxRead:
|
||||
"""Result of a mailbox wait/consume operation."""
|
||||
|
||||
state: MailboxReadState
|
||||
task: TaskSnapshot | None = None
|
||||
result: TaskResult | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _TaskRecord:
|
||||
request: TaskRequest
|
||||
state: TaskState = "running"
|
||||
result: TaskResult | None = None
|
||||
consumed_at: float | None = None
|
||||
completed_at: float | None = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class MailboxStore:
|
||||
"""Durable task mailbox for local subagent coordination.
|
||||
|
||||
JSON files are the source of truth. The condition variable only wakes
|
||||
waiters inside this process; persisted records remain readable after a
|
||||
manager restart.
|
||||
"""
|
||||
|
||||
def __init__(self, workspace: str | Path, *, root: str | Path | None = None) -> None:
|
||||
base = Path(root).expanduser() if root is not None else Path(workspace) / "tasks" / "subagents"
|
||||
self.root = ensure_dir(base)
|
||||
self._changed = asyncio.Condition()
|
||||
|
||||
async def dispatch(self, request: TaskRequest) -> None:
|
||||
"""Record that a task was dispatched."""
|
||||
async with self._changed:
|
||||
path, record = self._load_by_task_id(request.task_id, session_key=request.session_key)
|
||||
if record is not None:
|
||||
return
|
||||
path = self._record_path(request.session_key, request.task_id)
|
||||
self._write_record(path, _TaskRecord(request=request))
|
||||
self._changed.notify_all()
|
||||
|
||||
async def record_result(self, result: TaskResult) -> bool:
|
||||
"""Record a worker result.
|
||||
|
||||
Returns ``True`` when this call writes a new terminal result and
|
||||
``False`` when the task was already finalized.
|
||||
"""
|
||||
async with self._changed:
|
||||
path, record = self._load_by_task_id(result.task_id, session_key=result.session_key)
|
||||
if record is None:
|
||||
request = TaskRequest(
|
||||
task_id=result.task_id,
|
||||
session_key=result.session_key,
|
||||
label=result.label,
|
||||
task=result.task,
|
||||
origin=dict(result.metadata),
|
||||
created_at=result.completed_at,
|
||||
)
|
||||
record = _TaskRecord(request=request)
|
||||
path = self._record_path(result.session_key, result.task_id)
|
||||
elif record.result is not None or record.state != "running":
|
||||
return False
|
||||
|
||||
record.result = result
|
||||
record.completed_at = result.completed_at
|
||||
record.state = self._state_for_result(result.status)
|
||||
record.error = result.content if result.status in {"error", "cancelled"} else None
|
||||
self._write_record(path, record)
|
||||
self._changed.notify_all()
|
||||
return True
|
||||
|
||||
async def mark_cancelled(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
session_key: str | None = None,
|
||||
reason: str = "Cancelled.",
|
||||
) -> bool:
|
||||
"""Mark a task cancelled and make the cancellation consumable once."""
|
||||
async with self._changed:
|
||||
path, record = self._load_by_task_id(task_id, session_key=session_key)
|
||||
if record is None or record.result is not None or record.state != "running":
|
||||
return False
|
||||
result = TaskResult(
|
||||
task_id=task_id,
|
||||
session_key=record.request.session_key,
|
||||
label=record.request.label,
|
||||
task=record.request.task,
|
||||
status="cancelled",
|
||||
content=reason,
|
||||
dedupe_key=task_id,
|
||||
)
|
||||
record.result = result
|
||||
record.completed_at = result.completed_at
|
||||
record.state = "cancelled"
|
||||
record.error = reason
|
||||
self._write_record(path, record)
|
||||
self._changed.notify_all()
|
||||
return True
|
||||
|
||||
async def poll(
|
||||
self,
|
||||
session_key: str,
|
||||
*,
|
||||
task_id: str | None = None,
|
||||
) -> list[TaskSnapshot]:
|
||||
"""Return snapshots for one task or all tasks in a session."""
|
||||
async with self._changed:
|
||||
return self.snapshot_sync(session_key, task_id=task_id)
|
||||
|
||||
def snapshot_sync(
|
||||
self,
|
||||
session_key: str,
|
||||
*,
|
||||
task_id: str | None = None,
|
||||
) -> list[TaskSnapshot]:
|
||||
"""Synchronous snapshot used while building runtime context."""
|
||||
if task_id is not None:
|
||||
_, record = self._load_by_task_id(task_id, session_key=session_key)
|
||||
if record is None:
|
||||
return []
|
||||
return [self._snapshot(record)]
|
||||
|
||||
records = self._load_session_records(session_key)
|
||||
snapshots = [self._snapshot(record) for record in records]
|
||||
snapshots.sort(key=lambda item: (item.completed_at is None, item.created_at, item.task_id))
|
||||
return snapshots
|
||||
|
||||
async def wait_for_result(
|
||||
self,
|
||||
session_key: str,
|
||||
*,
|
||||
task_id: str | None = None,
|
||||
timeout_seconds: float = 30.0,
|
||||
) -> MailboxRead:
|
||||
"""Wait for and consume a result once."""
|
||||
deadline = time.monotonic() + max(0.0, timeout_seconds)
|
||||
async with self._changed:
|
||||
while True:
|
||||
read = self._consume_ready_locked(session_key, task_id)
|
||||
if read.state != "running":
|
||||
return read
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
return MailboxRead("timeout", task=read.task)
|
||||
try:
|
||||
await asyncio.wait_for(self._changed.wait(), timeout=remaining)
|
||||
except asyncio.TimeoutError:
|
||||
return MailboxRead("timeout", task=read.task)
|
||||
|
||||
def _consume_ready_locked(
|
||||
self,
|
||||
session_key: str,
|
||||
task_id: str | None,
|
||||
) -> MailboxRead:
|
||||
if task_id is not None:
|
||||
path, record = self._load_by_task_id(task_id, session_key=session_key)
|
||||
if record is None:
|
||||
return MailboxRead("not_found")
|
||||
snapshot = self._snapshot(record)
|
||||
if record.result is None:
|
||||
return MailboxRead("running", task=snapshot)
|
||||
if record.consumed_at is not None:
|
||||
return MailboxRead("consumed", task=snapshot, result=record.result)
|
||||
record.consumed_at = time.time()
|
||||
self._write_record(path, record)
|
||||
return MailboxRead("ready", task=self._snapshot(record), result=record.result)
|
||||
|
||||
records_with_paths = self._load_session_records_with_paths(session_key)
|
||||
ready = [
|
||||
(path, record)
|
||||
for path, record in records_with_paths
|
||||
if record.result is not None and record.consumed_at is None
|
||||
]
|
||||
if ready:
|
||||
ready.sort(key=lambda item: (
|
||||
item[1].completed_at or item[1].request.created_at,
|
||||
item[1].request.task_id,
|
||||
))
|
||||
path, record = ready[0]
|
||||
record.consumed_at = time.time()
|
||||
self._write_record(path, record)
|
||||
return MailboxRead("ready", task=self._snapshot(record), result=record.result)
|
||||
|
||||
running = [record for _, record in records_with_paths if record.result is None]
|
||||
if running:
|
||||
running.sort(key=lambda record: (record.request.created_at, record.request.task_id))
|
||||
return MailboxRead("running", task=self._snapshot(running[0]))
|
||||
if records_with_paths:
|
||||
records = [record for _, record in records_with_paths]
|
||||
records.sort(key=lambda record: (
|
||||
record.completed_at or record.request.created_at,
|
||||
record.request.task_id,
|
||||
))
|
||||
return MailboxRead("consumed", task=self._snapshot(records[-1]))
|
||||
return MailboxRead("not_found")
|
||||
|
||||
def _session_dir(self, session_key: str) -> Path:
|
||||
return self.root / safe_filename(session_key)
|
||||
|
||||
def _record_path(self, session_key: str, task_id: str) -> Path:
|
||||
return ensure_dir(self._session_dir(session_key)) / f"{safe_filename(task_id)}.json"
|
||||
|
||||
def _load_by_task_id(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
session_key: str | None = None,
|
||||
) -> tuple[Path, _TaskRecord | None]:
|
||||
if session_key is not None:
|
||||
path = self._record_path(session_key, task_id)
|
||||
return path, self._read_record(path)
|
||||
|
||||
filename = f"{safe_filename(task_id)}.json"
|
||||
for path in self.root.glob(f"*/{filename}"):
|
||||
record = self._read_record(path)
|
||||
if record is not None:
|
||||
return path, record
|
||||
return self.root / "_missing" / filename, None
|
||||
|
||||
def _load_session_records(self, session_key: str) -> list[_TaskRecord]:
|
||||
return [record for _, record in self._load_session_records_with_paths(session_key)]
|
||||
|
||||
def _load_session_records_with_paths(self, session_key: str) -> list[tuple[Path, _TaskRecord]]:
|
||||
directory = self._session_dir(session_key)
|
||||
if not directory.exists():
|
||||
return []
|
||||
records: list[tuple[Path, _TaskRecord]] = []
|
||||
for path in directory.glob("*.json"):
|
||||
record = self._read_record(path)
|
||||
if record is not None:
|
||||
records.append((path, record))
|
||||
return records
|
||||
|
||||
def _read_record(self, path: Path) -> _TaskRecord | None:
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
return self._record_from_json(data)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _write_record(self, path: Path, record: _TaskRecord) -> None:
|
||||
ensure_dir(path.parent)
|
||||
payload = json.dumps(self._record_to_json(record), ensure_ascii=False, indent=2)
|
||||
tmp = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
|
||||
try:
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
f.write(payload)
|
||||
f.write("\n")
|
||||
with suppress(OSError):
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
with suppress(OSError):
|
||||
fd = os.open(str(path.parent), os.O_RDONLY)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
finally:
|
||||
os.close(fd)
|
||||
finally:
|
||||
tmp.unlink(missing_ok=True)
|
||||
|
||||
@staticmethod
|
||||
def _record_to_json(record: _TaskRecord) -> dict[str, Any]:
|
||||
result = record.result
|
||||
return {
|
||||
"version": 1,
|
||||
"task_id": record.request.task_id,
|
||||
"session_key": record.request.session_key,
|
||||
"label": record.request.label,
|
||||
"task": record.request.task,
|
||||
"origin": record.request.origin,
|
||||
"state": record.state,
|
||||
"result": None if result is None else {
|
||||
"task_id": result.task_id,
|
||||
"session_key": result.session_key,
|
||||
"label": result.label,
|
||||
"task": result.task,
|
||||
"status": result.status,
|
||||
"content": result.content,
|
||||
"sender": result.sender,
|
||||
"completed_at": result.completed_at,
|
||||
"dedupe_key": result.dedupe_key,
|
||||
"metadata": result.metadata,
|
||||
},
|
||||
"consumed_at": record.consumed_at,
|
||||
"created_at": record.request.created_at,
|
||||
"completed_at": record.completed_at,
|
||||
"updated_at": time.time(),
|
||||
"error": record.error,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _record_from_json(data: dict[str, Any]) -> _TaskRecord:
|
||||
request = TaskRequest(
|
||||
task_id=str(data["task_id"]),
|
||||
session_key=str(data["session_key"]),
|
||||
label=str(data.get("label") or data["task_id"]),
|
||||
task=str(data.get("task") or ""),
|
||||
origin=dict(data.get("origin") or {}),
|
||||
created_at=float(data.get("created_at") or time.time()),
|
||||
)
|
||||
raw_result = data.get("result")
|
||||
result = None
|
||||
if isinstance(raw_result, dict):
|
||||
result = TaskResult(
|
||||
task_id=str(raw_result.get("task_id") or request.task_id),
|
||||
session_key=str(raw_result.get("session_key") or request.session_key),
|
||||
label=str(raw_result.get("label") or request.label),
|
||||
task=str(raw_result.get("task") or request.task),
|
||||
status=str(raw_result.get("status") or "error"),
|
||||
content=str(raw_result.get("content") or ""),
|
||||
sender=str(raw_result.get("sender") or "subagent"),
|
||||
completed_at=float(raw_result.get("completed_at") or time.time()),
|
||||
dedupe_key=raw_result.get("dedupe_key"),
|
||||
metadata=dict(raw_result.get("metadata") or {}),
|
||||
)
|
||||
return _TaskRecord(
|
||||
request=request,
|
||||
state=str(data.get("state") or "running"),
|
||||
result=result,
|
||||
consumed_at=data.get("consumed_at"),
|
||||
completed_at=data.get("completed_at"),
|
||||
error=data.get("error"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _state_for_result(status: str) -> TaskState:
|
||||
if status == "ok":
|
||||
return "completed"
|
||||
if status == "cancelled":
|
||||
return "cancelled"
|
||||
return "failed"
|
||||
|
||||
@staticmethod
|
||||
def _snapshot(record: _TaskRecord) -> TaskSnapshot:
|
||||
result = record.result
|
||||
return TaskSnapshot(
|
||||
task_id=record.request.task_id,
|
||||
session_key=record.request.session_key,
|
||||
label=record.request.label,
|
||||
task=record.request.task,
|
||||
state=record.state,
|
||||
created_at=record.request.created_at,
|
||||
completed_at=record.completed_at,
|
||||
consumed_at=record.consumed_at,
|
||||
result_status=result.status if result is not None else None,
|
||||
error=record.error,
|
||||
)
|
||||
+141
-30
@@ -4,19 +4,20 @@ import asyncio
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
from typing import Any, Awaitable, Callable
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||
from nanobot.agent.mailbox import MailboxRead, MailboxStore, TaskRequest, TaskResult, TaskSnapshot
|
||||
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
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
|
||||
@@ -87,6 +88,8 @@ class SubagentManager:
|
||||
max_iterations: int | None = None,
|
||||
max_concurrent_subagents: int | None = None,
|
||||
llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None,
|
||||
mailbox: MailboxStore | None = None,
|
||||
on_result_ready: Callable[[TaskResult], Awaitable[None]] | None = None,
|
||||
):
|
||||
defaults = AgentDefaults()
|
||||
self.provider = provider
|
||||
@@ -109,6 +112,8 @@ class SubagentManager:
|
||||
)
|
||||
self.runner = AgentRunner(provider)
|
||||
self._llm_wall_timeout_for_session = llm_wall_timeout_for_session
|
||||
self.mailbox = mailbox or MailboxStore(workspace)
|
||||
self._on_result_ready = on_result_ready
|
||||
self._running_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._task_statuses: dict[str, SubagentStatus] = {}
|
||||
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
||||
@@ -162,6 +167,7 @@ class SubagentManager:
|
||||
"""Spawn a subagent to execute a task in the background."""
|
||||
task_id = str(uuid.uuid4())[:8]
|
||||
display_label = label or task[:30] + ("..." if len(task) > 30 else "")
|
||||
mailbox_session_key = session_key or f"{origin_channel}:{origin_chat_id}"
|
||||
origin = {"channel": origin_channel, "chat_id": origin_chat_id, "session_key": session_key}
|
||||
|
||||
status = SubagentStatus(
|
||||
@@ -171,6 +177,18 @@ class SubagentManager:
|
||||
started_at=time.monotonic(),
|
||||
)
|
||||
self._task_statuses[task_id] = status
|
||||
await self.mailbox.dispatch(TaskRequest(
|
||||
task_id=task_id,
|
||||
session_key=mailbox_session_key,
|
||||
label=display_label,
|
||||
task=task,
|
||||
origin={
|
||||
"channel": origin_channel,
|
||||
"chat_id": origin_chat_id,
|
||||
"session_key": session_key,
|
||||
"origin_message_id": origin_message_id,
|
||||
},
|
||||
))
|
||||
|
||||
bg_task = asyncio.create_task(
|
||||
self._run_subagent(
|
||||
@@ -199,14 +217,17 @@ class SubagentManager:
|
||||
bg_task.add_done_callback(_cleanup)
|
||||
|
||||
logger.info("Spawned subagent [{}]: {}", task_id, display_label)
|
||||
return f"Subagent [{display_label}] started (id: {task_id}). I'll notify you when it completes."
|
||||
return (
|
||||
f"Subagent [{display_label}] started (id: {task_id}). "
|
||||
f"Use poll_subagents or wait_subagents with id {task_id} to get the result."
|
||||
)
|
||||
|
||||
async def _run_subagent(
|
||||
self,
|
||||
task_id: str,
|
||||
task: str,
|
||||
label: str,
|
||||
origin: dict[str, str],
|
||||
origin: dict[str, Any],
|
||||
status: SubagentStatus,
|
||||
origin_message_id: str | None = None,
|
||||
temperature: float | None = None,
|
||||
@@ -281,6 +302,12 @@ class SubagentManager:
|
||||
logger.info("Subagent [{}] completed successfully", task_id)
|
||||
await self._announce_result(task_id, label, task, final_result, origin, "ok", origin_message_id)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
status.phase = "cancelled"
|
||||
status.stop_reason = "cancelled"
|
||||
await self.mailbox.mark_cancelled(task_id, reason="Cancelled.")
|
||||
logger.info("Subagent [{}] cancelled", task_id)
|
||||
raise
|
||||
except Exception as e:
|
||||
status.phase = "error"
|
||||
status.error = str(e)
|
||||
@@ -293,44 +320,45 @@ class SubagentManager:
|
||||
label: str,
|
||||
task: str,
|
||||
result: str,
|
||||
origin: dict[str, str],
|
||||
origin: dict[str, Any],
|
||||
status: str,
|
||||
origin_message_id: str | None = None,
|
||||
) -> None:
|
||||
"""Announce the subagent result to the main agent via the message bus."""
|
||||
status_text = "completed successfully" if status == "ok" else "failed"
|
||||
|
||||
announce_content = render_template(
|
||||
"agent/subagent_announce.md",
|
||||
label=label,
|
||||
status_text=status_text,
|
||||
task=task,
|
||||
result=result,
|
||||
)
|
||||
|
||||
# Inject as system message to trigger main agent.
|
||||
# Use session_key_override to align with the main agent's effective
|
||||
# session key (which accounts for unified sessions) so the result is
|
||||
# routed to the correct pending queue (mid-turn injection) instead of
|
||||
# being dispatched as a competing independent task.
|
||||
"""Record the subagent result in the mailbox for explicit manager polling."""
|
||||
override = origin.get("session_key") or f"{origin['channel']}:{origin['chat_id']}"
|
||||
metadata: dict[str, Any] = {
|
||||
"injected_event": "subagent_result",
|
||||
"subagent_task_id": task_id,
|
||||
"origin_channel": origin.get("channel"),
|
||||
"origin_chat_id": origin.get("chat_id"),
|
||||
}
|
||||
if origin_message_id:
|
||||
metadata["origin_message_id"] = origin_message_id
|
||||
msg = InboundMessage(
|
||||
channel="system",
|
||||
sender_id="subagent",
|
||||
chat_id=f"{origin['channel']}:{origin['chat_id']}",
|
||||
content=announce_content,
|
||||
session_key_override=override,
|
||||
|
||||
task_result = TaskResult(
|
||||
task_id=task_id,
|
||||
session_key=override,
|
||||
label=label,
|
||||
task=task,
|
||||
status=status,
|
||||
content=result,
|
||||
dedupe_key=task_id,
|
||||
metadata=metadata,
|
||||
)
|
||||
written = await self.mailbox.record_result(task_result)
|
||||
|
||||
await self.bus.publish_inbound(msg)
|
||||
logger.debug("Subagent [{}] announced result to {}:{}", task_id, origin['channel'], origin['chat_id'])
|
||||
if written:
|
||||
logger.debug(
|
||||
"Subagent [{}] wrote result to mailbox for session {}",
|
||||
task_id,
|
||||
override,
|
||||
)
|
||||
if self._on_result_ready is not None:
|
||||
try:
|
||||
await self._on_result_ready(task_result)
|
||||
except Exception:
|
||||
logger.exception("Subagent result-ready callback failed")
|
||||
else:
|
||||
logger.debug("Subagent [{}] result already recorded", task_id)
|
||||
|
||||
@staticmethod
|
||||
def _format_partial_progress(result) -> str:
|
||||
@@ -375,12 +403,95 @@ class SubagentManager:
|
||||
"""Cancel all subagents for the given session. Returns count cancelled."""
|
||||
tasks = [self._running_tasks[tid] for tid in self._session_tasks.get(session_key, [])
|
||||
if tid in self._running_tasks and not self._running_tasks[tid].done()]
|
||||
for tid in list(self._session_tasks.get(session_key, [])):
|
||||
if tid in self._running_tasks and not self._running_tasks[tid].done():
|
||||
await self.mailbox.mark_cancelled(
|
||||
tid,
|
||||
session_key=session_key,
|
||||
reason="Cancelled by /stop.",
|
||||
)
|
||||
for t in tasks:
|
||||
t.cancel()
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
return len(tasks)
|
||||
|
||||
async def cancel_task(self, task_id: str, session_key: str | None = None) -> str:
|
||||
"""Cancel one running subagent task and record a cancelled mailbox state."""
|
||||
snapshots = await self.mailbox.poll(session_key, task_id=task_id) if session_key else []
|
||||
if session_key and not snapshots:
|
||||
return "not_found"
|
||||
task = self._running_tasks.get(task_id)
|
||||
if task is None or task.done():
|
||||
if snapshots:
|
||||
return snapshots[0].state
|
||||
return "not_found"
|
||||
await self.mailbox.mark_cancelled(
|
||||
task_id,
|
||||
session_key=session_key,
|
||||
reason="Cancelled by manager.",
|
||||
)
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError, Exception):
|
||||
await task
|
||||
return "cancelled"
|
||||
|
||||
async def poll(
|
||||
self,
|
||||
session_key: str,
|
||||
task_id: str | None = None,
|
||||
) -> list[TaskSnapshot]:
|
||||
"""Return mailbox task status snapshots for a session."""
|
||||
return await self.mailbox.poll(session_key, task_id=task_id)
|
||||
|
||||
async def wait_for_result(
|
||||
self,
|
||||
session_key: str,
|
||||
task_id: str | None = None,
|
||||
timeout_seconds: float = 30.0,
|
||||
) -> MailboxRead:
|
||||
"""Wait for and consume a mailbox result for a session."""
|
||||
return await self.mailbox.wait_for_result(
|
||||
session_key,
|
||||
task_id=task_id,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
def runtime_status_lines(self, session_key: str, *, limit: int = 8) -> list[str]:
|
||||
"""Return compact model-visible task status lines for runtime context."""
|
||||
snapshots = self.mailbox.snapshot_sync(session_key)
|
||||
if not snapshots:
|
||||
return []
|
||||
|
||||
now = time.time()
|
||||
ordered = sorted(
|
||||
snapshots,
|
||||
key=lambda item: (
|
||||
item.consumed_at is not None,
|
||||
item.completed_at is None,
|
||||
item.created_at,
|
||||
item.task_id,
|
||||
),
|
||||
)
|
||||
lines = ["Subagent tasks:"]
|
||||
for snapshot in ordered[: max(0, limit)]:
|
||||
state = snapshot.state
|
||||
if snapshot.result_status and snapshot.consumed_at is None:
|
||||
state = f"{state}, result ready"
|
||||
elif snapshot.consumed_at is not None:
|
||||
state = f"{state}, result consumed"
|
||||
elapsed = max(0, int((snapshot.completed_at or now) - snapshot.created_at))
|
||||
label = " ".join(snapshot.label.split())
|
||||
if len(label) > 48:
|
||||
label = label[:45] + "..."
|
||||
lines.append(
|
||||
f"- {snapshot.task_id}: {state}, label=\"{label}\", elapsed={elapsed}s"
|
||||
)
|
||||
remaining = len(ordered) - limit
|
||||
if remaining > 0:
|
||||
lines.append(f"- ... {remaining} more subagent task(s)")
|
||||
return lines
|
||||
|
||||
def get_running_count(self) -> int:
|
||||
"""Return the number of currently running subagents."""
|
||||
return len(self._running_tasks)
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Runtime delivery helpers for completed subagent task results."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from typing import Any
|
||||
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.session import turn_continuation
|
||||
|
||||
_FORWARDED_METADATA_KEYS = frozenset({
|
||||
"message_id",
|
||||
"origin_message_id",
|
||||
"_wants_stream",
|
||||
"webui",
|
||||
"slack",
|
||||
})
|
||||
|
||||
|
||||
def build_subagent_result_continuation(result: Any) -> InboundMessage:
|
||||
"""Build an internal inbound wake-up for a ready subagent result."""
|
||||
metadata = dict(result.metadata or {})
|
||||
channel = str(metadata.get("origin_channel") or "")
|
||||
chat_id = str(metadata.get("origin_chat_id") or "")
|
||||
if not channel or not chat_id:
|
||||
channel, chat_id = _channel_chat_from_session_key(result.session_key)
|
||||
|
||||
wake_meta = turn_continuation.subagent_result_continuation_metadata(
|
||||
{key: value for key, value in metadata.items() if key in _FORWARDED_METADATA_KEYS},
|
||||
task_id=result.task_id,
|
||||
)
|
||||
return InboundMessage(
|
||||
channel=channel,
|
||||
sender_id="system:continuation",
|
||||
chat_id=chat_id,
|
||||
content=(
|
||||
"A subagent task result is ready. The runtime will attach the "
|
||||
"result to this continuation turn."
|
||||
),
|
||||
metadata=wake_meta,
|
||||
session_key_override=result.session_key,
|
||||
)
|
||||
|
||||
|
||||
async def materialize_subagent_result_continuation(
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
session_key: str,
|
||||
subagents: Any,
|
||||
) -> InboundMessage:
|
||||
"""Replace a subagent-result continuation placeholder with the mailbox result."""
|
||||
task_id = turn_continuation.subagent_result_continuation_task_id(msg.metadata)
|
||||
if not task_id:
|
||||
return msg
|
||||
read = await subagents.wait_for_result(
|
||||
session_key,
|
||||
task_id=task_id,
|
||||
timeout_seconds=0,
|
||||
)
|
||||
return dataclasses.replace(msg, content=_subagent_result_continuation_content(read, task_id))
|
||||
|
||||
|
||||
def _channel_chat_from_session_key(session_key: str) -> tuple[str, str]:
|
||||
channel, _, chat_id = session_key.partition(":")
|
||||
return channel or "cli", chat_id or "direct"
|
||||
|
||||
|
||||
def _subagent_result_continuation_content(read: Any, requested_task_id: str) -> str:
|
||||
if read.state == "ready" and read.result is not None:
|
||||
status_text = {
|
||||
"ok": "completed",
|
||||
"error": "failed",
|
||||
"cancelled": "cancelled",
|
||||
}.get(read.result.status, read.result.status)
|
||||
return (
|
||||
"A subagent result was delivered by the runtime. Use this result "
|
||||
"as authoritative context for the next answer; do not mention the "
|
||||
"internal continuation boundary.\n\n"
|
||||
f"Subagent [{read.result.label}] "
|
||||
f"(id: {read.result.task_id}, status: {status_text})\n\n"
|
||||
f"Task:\n{read.result.task}\n\n"
|
||||
f"Result:\n{read.result.content}"
|
||||
)
|
||||
if read.state == "consumed":
|
||||
return (
|
||||
f"Subagent task {requested_task_id} already has a consumed result. "
|
||||
"Check poll_subagents if you need its current status."
|
||||
)
|
||||
if read.state == "running":
|
||||
return (
|
||||
f"Subagent task {requested_task_id} is still running. "
|
||||
"Use poll_subagents or wait_subagents if you need to block."
|
||||
)
|
||||
return f"Subagent task {requested_task_id} result is not available ({read.state})."
|
||||
@@ -63,7 +63,8 @@ class SpawnTool(Tool, ContextAware):
|
||||
return (
|
||||
"Spawn a subagent to handle a task in the background. "
|
||||
"Use this for complex or time-consuming tasks that can run independently. "
|
||||
"The subagent will complete the task and report back when done. "
|
||||
"The subagent writes its result to a mailbox; use poll_subagents "
|
||||
"or wait_subagents to retrieve it explicitly. "
|
||||
"For deliverables or existing projects, inspect the workspace first "
|
||||
"and use a dedicated subdirectory when helpful."
|
||||
)
|
||||
@@ -81,8 +82,8 @@ class SpawnTool(Tool, ContextAware):
|
||||
if running >= limit:
|
||||
return (
|
||||
f"Cannot spawn subagent: concurrency limit reached "
|
||||
f"({running}/{limit} running). Wait for a running subagent "
|
||||
f"to complete before spawning a new one."
|
||||
f"({running}/{limit} running). Use wait_subagents or cancel_subagent "
|
||||
f"before spawning a new one."
|
||||
)
|
||||
return await self._manager.spawn(
|
||||
task=task,
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Explicit mailbox tools for subagent coordination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextvars import ContextVar
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from nanobot.agent.mailbox import MailboxRead, TaskSnapshot
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.context import ContextAware, RequestContext
|
||||
from nanobot.agent.tools.schema import NumberSchema, StringSchema, tool_parameters_schema
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
|
||||
|
||||
def _normalize_task_id(task_id: str | None) -> str | None:
|
||||
if task_id is None:
|
||||
return None
|
||||
task_id = task_id.strip()
|
||||
return task_id or None
|
||||
|
||||
|
||||
def _truncate(text: str, limit: int = 120) -> str:
|
||||
text = " ".join(text.split())
|
||||
return text if len(text) <= limit else text[: limit - 3] + "..."
|
||||
|
||||
|
||||
class _SubagentMailboxTool(Tool, ContextAware):
|
||||
"""Shared context plumbing for subagent mailbox tools."""
|
||||
|
||||
def __init__(self, manager: "SubagentManager"):
|
||||
self._manager = manager
|
||||
self._session_key: ContextVar[str] = ContextVar(
|
||||
f"{self.__class__.__name__}_session_key",
|
||||
default="cli:direct",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
return getattr(ctx, "subagent_manager", None) is not None
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
return cls(manager=ctx.subagent_manager)
|
||||
|
||||
def set_context(self, ctx: RequestContext) -> None:
|
||||
self._session_key.set(ctx.session_key or f"{ctx.channel}:{ctx.chat_id}")
|
||||
|
||||
|
||||
@tool_parameters(
|
||||
tool_parameters_schema(
|
||||
task_id=StringSchema(
|
||||
"Optional subagent task id. Omit to list all subagent tasks for this session.",
|
||||
nullable=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
class PollSubagentsTool(_SubagentMailboxTool):
|
||||
"""Non-blocking task status check."""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "poll_subagents"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return (
|
||||
"Check subagent task status without blocking. Use this to see whether a "
|
||||
"spawned subagent is still running or has a result ready to consume."
|
||||
)
|
||||
|
||||
@property
|
||||
def read_only(self) -> bool:
|
||||
return True
|
||||
|
||||
async def execute(self, task_id: str | None = None, **_: Any) -> str:
|
||||
task_id = _normalize_task_id(task_id)
|
||||
session_key = self._session_key.get()
|
||||
snapshots = await self._manager.poll(session_key, task_id=task_id)
|
||||
if not snapshots:
|
||||
if task_id:
|
||||
return f"Subagent task {task_id} not found for this session."
|
||||
return "No subagent tasks found for this session."
|
||||
return self._format_snapshots(snapshots)
|
||||
|
||||
@staticmethod
|
||||
def _format_snapshots(snapshots: list[TaskSnapshot]) -> str:
|
||||
lines = ["Subagent task status:"]
|
||||
for snapshot in snapshots:
|
||||
state = snapshot.state
|
||||
if snapshot.result_status and snapshot.consumed_at is None:
|
||||
state = f"{state}, result ready"
|
||||
elif snapshot.consumed_at is not None:
|
||||
state = f"{state}, result consumed"
|
||||
lines.append(
|
||||
f"- id: {snapshot.task_id} | label: {snapshot.label} | "
|
||||
f"status: {state} | task: {_truncate(snapshot.task)}"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
@tool_parameters(
|
||||
tool_parameters_schema(
|
||||
task_id=StringSchema(
|
||||
"Optional subagent task id. Omit to consume the next ready result.",
|
||||
nullable=True,
|
||||
),
|
||||
timeout_seconds=NumberSchema(
|
||||
description="How long to wait for a result before returning. Defaults to 30 seconds.",
|
||||
minimum=0.0,
|
||||
maximum=300.0,
|
||||
),
|
||||
)
|
||||
)
|
||||
class WaitSubagentsTool(_SubagentMailboxTool):
|
||||
"""Wait for and consume one task result."""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "wait_subagents"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return (
|
||||
"Wait for a subagent result and consume it once. Use this after spawn "
|
||||
"when you need the worker's result before continuing."
|
||||
)
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
task_id: str | None = None,
|
||||
timeout_seconds: float = 30.0,
|
||||
**_: Any,
|
||||
) -> str:
|
||||
task_id = _normalize_task_id(task_id)
|
||||
read = await self._manager.wait_for_result(
|
||||
self._session_key.get(),
|
||||
task_id=task_id,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
return self._format_read(read, task_id)
|
||||
|
||||
@staticmethod
|
||||
def _format_read(read: MailboxRead, requested_task_id: str | None) -> str:
|
||||
if read.state == "not_found":
|
||||
target = f" {requested_task_id}" if requested_task_id else ""
|
||||
return f"Subagent task{target} not found for this session."
|
||||
if read.state == "timeout":
|
||||
target = f" {read.task.task_id}" if read.task is not None else ""
|
||||
return f"Timed out waiting for subagent task{target}."
|
||||
if read.state == "consumed":
|
||||
target = f" {read.task.task_id}" if read.task is not None else ""
|
||||
return f"Subagent result for task{target} was already consumed."
|
||||
if read.result is None or read.task is None:
|
||||
return "No subagent result is ready."
|
||||
|
||||
status_text = {
|
||||
"ok": "completed",
|
||||
"error": "failed",
|
||||
"cancelled": "cancelled",
|
||||
}.get(read.result.status, read.result.status)
|
||||
return (
|
||||
f"Subagent result for [{read.result.label}] "
|
||||
f"(id: {read.result.task_id}, status: {status_text}).\n\n"
|
||||
f"Task: {read.result.task}\n\n"
|
||||
f"Result:\n{read.result.content}"
|
||||
)
|
||||
|
||||
|
||||
@tool_parameters(
|
||||
tool_parameters_schema(
|
||||
task_id=StringSchema("Subagent task id to cancel"),
|
||||
required=["task_id"],
|
||||
)
|
||||
)
|
||||
class CancelSubagentTool(_SubagentMailboxTool):
|
||||
"""Cancel one running task."""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "cancel_subagent"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return (
|
||||
"Cancel a running subagent task and record a cancelled mailbox state. "
|
||||
"Use this only when the delegated task is no longer needed."
|
||||
)
|
||||
|
||||
async def execute(self, task_id: str, **_: Any) -> str:
|
||||
task_id = _normalize_task_id(task_id)
|
||||
if task_id is None:
|
||||
return "Error: task_id is required."
|
||||
state = await self._manager.cancel_task(task_id, session_key=self._session_key.get())
|
||||
if state == "cancelled":
|
||||
return f"Cancelled subagent task {task_id}."
|
||||
if state == "not_found":
|
||||
return f"Subagent task {task_id} not found for this session."
|
||||
if state in {"completed", "failed"}:
|
||||
return (
|
||||
f"Subagent task {task_id} already {state}; "
|
||||
"use wait_subagents to consume its result if needed."
|
||||
)
|
||||
if state == "cancelled":
|
||||
return f"Subagent task {task_id} is already cancelled."
|
||||
return f"Subagent task {task_id} is {state}."
|
||||
@@ -25,6 +25,8 @@ INTERNAL_CONTINUATION_RUN_STARTED_AT_META = "_internal_continuation_run_started_
|
||||
SKIP_USER_PERSIST_META = "_skip_user_persist"
|
||||
|
||||
_GOAL_CONTINUATION_KIND = "sustained_goal"
|
||||
SUBAGENT_RESULT_CONTINUATION_KIND = "subagent_result"
|
||||
SUBAGENT_RESULT_TASK_ID_META = "_subagent_result_task_id"
|
||||
_GOAL_CONTINUATION_SENDER = "system:continuation"
|
||||
_GOAL_CONTINUATION_ROUNDS_KEY = "_sustained_goal_continuation_rounds"
|
||||
_MAX_GOAL_CONTINUATION_ROUNDS = 12
|
||||
@@ -58,6 +60,38 @@ def internal_continuation_run_started_at(metadata: Mapping[str, Any] | None) ->
|
||||
return started_at if started_at > 0 else None
|
||||
|
||||
|
||||
def subagent_result_continuation_inbound(metadata: Mapping[str, Any] | None) -> bool:
|
||||
"""True for an internal continuation caused by a ready subagent result."""
|
||||
return bool(
|
||||
internal_continuation_inbound(metadata)
|
||||
and metadata.get(INTERNAL_CONTINUATION_KIND_META) == SUBAGENT_RESULT_CONTINUATION_KIND
|
||||
)
|
||||
|
||||
|
||||
def subagent_result_continuation_task_id(metadata: Mapping[str, Any] | None) -> str | None:
|
||||
"""Return the ready subagent task id carried by a continuation message."""
|
||||
if not subagent_result_continuation_inbound(metadata):
|
||||
return None
|
||||
value = metadata.get(SUBAGENT_RESULT_TASK_ID_META) if metadata else None
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def subagent_result_continuation_metadata(
|
||||
message_metadata: Mapping[str, Any] | None,
|
||||
*,
|
||||
task_id: str,
|
||||
run_started_at: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build sanitized metadata for a subagent-result continuation turn."""
|
||||
metadata = _internal_continuation_metadata(
|
||||
message_metadata,
|
||||
kind=SUBAGENT_RESULT_CONTINUATION_KIND,
|
||||
run_started_at=run_started_at,
|
||||
)
|
||||
metadata[SUBAGENT_RESULT_TASK_ID_META] = task_id
|
||||
return metadata
|
||||
|
||||
|
||||
def should_persist_user_message(metadata: Mapping[str, Any] | None) -> bool:
|
||||
"""Return whether this inbound message should be persisted as user input."""
|
||||
if metadata and metadata.get(SKIP_USER_PERSIST_META) is True:
|
||||
@@ -223,11 +257,12 @@ def _increment_goal_continuation_round(session_metadata: MutableMapping[str, Any
|
||||
def _internal_continuation_metadata(
|
||||
message_metadata: Mapping[str, Any] | None,
|
||||
*,
|
||||
kind: str = _GOAL_CONTINUATION_KIND,
|
||||
run_started_at: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
metadata = dict(message_metadata or {})
|
||||
metadata[INTERNAL_CONTINUATION_META] = True
|
||||
metadata[INTERNAL_CONTINUATION_KIND_META] = _GOAL_CONTINUATION_KIND
|
||||
metadata[INTERNAL_CONTINUATION_KIND_META] = kind
|
||||
if run_started_at is not None:
|
||||
metadata[INTERNAL_CONTINUATION_RUN_STARTED_AT_META] = float(run_started_at)
|
||||
for key in _STRIPPED_INBOUND_META_KEYS:
|
||||
|
||||
@@ -14,8 +14,10 @@ from nanobot.providers.base import LLMResponse
|
||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
from nanobot.session.turn_continuation import (
|
||||
INTERNAL_CONTINUATION_KIND_META,
|
||||
INTERNAL_CONTINUATION_META,
|
||||
INTERNAL_CONTINUATION_RUN_STARTED_AT_META,
|
||||
SUBAGENT_RESULT_CONTINUATION_KIND,
|
||||
)
|
||||
from nanobot.session.webui_turns import (
|
||||
TITLE_GENERATION_MAX_TOKENS,
|
||||
@@ -864,6 +866,100 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
|
||||
assert isinstance(turn_end[0].metadata.get("latency_ms"), int)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_context_lists_ready_subagent_result(tmp_path: Path) -> None:
|
||||
loop = _make_full_loop(tmp_path)
|
||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||
|
||||
await loop.subagents._announce_result(
|
||||
"sub-ready",
|
||||
"research",
|
||||
"look up the answer",
|
||||
"worker answer",
|
||||
{"channel": "cli", "chat_id": "test", "session_key": "cli:test"},
|
||||
"ok",
|
||||
)
|
||||
|
||||
seen: dict[str, list[dict]] = {}
|
||||
|
||||
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||
seen["initial_messages"] = initial_messages
|
||||
return (
|
||||
"done",
|
||||
[],
|
||||
[*initial_messages, {"role": "assistant", "content": "done"}],
|
||||
"completed",
|
||||
False,
|
||||
)
|
||||
|
||||
loop._run_agent_loop = fake_run_agent_loop # type: ignore[method-assign]
|
||||
|
||||
await loop._process_message(
|
||||
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="continue")
|
||||
)
|
||||
|
||||
rendered = "\n".join(str(msg.get("content", "")) for msg in seen["initial_messages"])
|
||||
assert "Subagent tasks:" in rendered
|
||||
assert "sub-ready: completed, result ready" in rendered
|
||||
assert "worker answer" not in rendered
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_result_continuation_delivers_result_without_user_history(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
loop = _make_full_loop(tmp_path)
|
||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||
|
||||
await loop.subagents._announce_result(
|
||||
"sub-deliver",
|
||||
"worker",
|
||||
"calculate the answer",
|
||||
"the worker result",
|
||||
{"channel": "cli", "chat_id": "test", "session_key": "cli:test"},
|
||||
"ok",
|
||||
)
|
||||
queued = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=0.5)
|
||||
|
||||
assert queued.metadata[INTERNAL_CONTINUATION_META] is True
|
||||
assert queued.metadata[INTERNAL_CONTINUATION_KIND_META] == SUBAGENT_RESULT_CONTINUATION_KIND
|
||||
assert "the worker result" not in queued.content
|
||||
|
||||
seen: dict[str, list[dict]] = {}
|
||||
|
||||
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||
seen["initial_messages"] = initial_messages
|
||||
return (
|
||||
"reported",
|
||||
[],
|
||||
[*initial_messages, {"role": "assistant", "content": "reported"}],
|
||||
"completed",
|
||||
False,
|
||||
)
|
||||
|
||||
loop._run_agent_loop = fake_run_agent_loop # type: ignore[method-assign]
|
||||
|
||||
response = await loop._process_message(queued, pending_queue=asyncio.Queue())
|
||||
|
||||
assert response is not None
|
||||
assert response.content == "reported"
|
||||
rendered = "\n".join(str(msg.get("content", "")) for msg in seen["initial_messages"])
|
||||
assert "the worker result" in rendered
|
||||
|
||||
read = await loop.subagents.wait_for_result(
|
||||
"cli:test",
|
||||
task_id="sub-deliver",
|
||||
timeout_seconds=0,
|
||||
)
|
||||
assert read.state == "consumed"
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
assert [
|
||||
{k: v for k, v in m.items() if k in {"role", "content"}}
|
||||
for m in session.messages
|
||||
] == [{"role": "assistant", "content": "reported"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_message_uses_context_chat_id_for_runtime_prompt(tmp_path: Path) -> None:
|
||||
loop = _make_full_loop(tmp_path)
|
||||
|
||||
@@ -285,80 +285,76 @@ class TestRunSubagent:
|
||||
|
||||
class TestAnnounceResult:
|
||||
@pytest.mark.asyncio
|
||||
async def test_publishes_inbound_message(self, tmp_path):
|
||||
async def test_records_mailbox_result_without_publishing_inbound(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
sm.bus.publish_inbound = AsyncMock()
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "result text",
|
||||
{"channel": "cli", "chat_id": "direct"}, "ok",
|
||||
)
|
||||
|
||||
assert len(published) == 1
|
||||
msg = published[0]
|
||||
assert msg.channel == "system"
|
||||
assert msg.sender_id == "subagent"
|
||||
assert msg.metadata["injected_event"] == "subagent_result"
|
||||
assert msg.metadata["subagent_task_id"] == "t1"
|
||||
sm.bus.publish_inbound.assert_not_awaited()
|
||||
snapshots = await sm.mailbox.poll("cli:direct", task_id="t1")
|
||||
assert snapshots[0].state == "completed"
|
||||
read = await sm.mailbox.wait_for_result("cli:direct", task_id="t1", timeout_seconds=0)
|
||||
assert read.state == "ready"
|
||||
assert read.result is not None
|
||||
assert read.result.content == "result text"
|
||||
assert read.result.metadata["subagent_task_id"] == "t1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_key_override(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "result",
|
||||
{"channel": "telegram", "chat_id": "123", "session_key": "s1"}, "ok",
|
||||
)
|
||||
|
||||
assert published[0].session_key_override == "s1"
|
||||
assert await sm.mailbox.poll("s1", task_id="t1")
|
||||
assert await sm.mailbox.poll("telegram:123", task_id="t1") == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_key_override_fallback(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "result",
|
||||
{"channel": "telegram", "chat_id": "123"}, "ok",
|
||||
)
|
||||
|
||||
assert published[0].session_key_override == "telegram:123"
|
||||
snapshots = await sm.mailbox.poll("telegram:123", task_id="t1")
|
||||
assert snapshots[0].session_key == "telegram:123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ok_status_text(self, tmp_path):
|
||||
async def test_ok_status_records_completed_state(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "result",
|
||||
{"channel": "cli", "chat_id": "direct"}, "ok",
|
||||
)
|
||||
|
||||
assert "completed successfully" in published[0].content
|
||||
snapshots = await sm.mailbox.poll("cli:direct", task_id="t1")
|
||||
assert snapshots[0].state == "completed"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_status_text(self, tmp_path):
|
||||
async def test_error_status_records_failed_state(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "error details",
|
||||
{"channel": "cli", "chat_id": "direct"}, "error",
|
||||
)
|
||||
|
||||
assert "failed" in published[0].content
|
||||
snapshots = await sm.mailbox.poll("cli:direct", task_id="t1")
|
||||
assert snapshots[0].state == "failed"
|
||||
assert snapshots[0].error == "error details"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_origin_message_id_in_metadata(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
published = []
|
||||
sm.bus.publish_inbound = AsyncMock(side_effect=lambda msg: published.append(msg))
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "result",
|
||||
@@ -366,7 +362,29 @@ class TestAnnounceResult:
|
||||
origin_message_id="msg-123",
|
||||
)
|
||||
|
||||
assert published[0].metadata["origin_message_id"] == "msg-123"
|
||||
read = await sm.mailbox.wait_for_result("cli:direct", task_id="t1", timeout_seconds=0)
|
||||
assert read.result is not None
|
||||
assert read.result.metadata["origin_message_id"] == "msg-123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_results_are_not_consumed_twice(self, tmp_path):
|
||||
sm = _manager(tmp_path)
|
||||
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "first",
|
||||
{"channel": "cli", "chat_id": "direct"}, "ok",
|
||||
)
|
||||
await sm._announce_result(
|
||||
"t1", "label", "task", "second",
|
||||
{"channel": "cli", "chat_id": "direct"}, "ok",
|
||||
)
|
||||
|
||||
first = await sm.mailbox.wait_for_result("cli:direct", task_id="t1", timeout_seconds=0)
|
||||
second = await sm.mailbox.wait_for_result("cli:direct", task_id="t1", timeout_seconds=0)
|
||||
assert first.state == "ready"
|
||||
assert first.result is not None
|
||||
assert first.result.content == "first"
|
||||
assert second.state == "consumed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -427,7 +427,7 @@ class TestSubagentCancellation:
|
||||
|
||||
|
||||
class TestSubagentAnnounceSessionKey:
|
||||
"""Verify _announce_result uses the effective session key for mid-turn routing."""
|
||||
"""Verify _announce_result stores results under the effective session key."""
|
||||
|
||||
def _make_mgr(self):
|
||||
"""Create a SubagentManager with mocked deps and its bus."""
|
||||
@@ -448,27 +448,27 @@ class TestSubagentAnnounceSessionKey:
|
||||
@pytest.mark.asyncio
|
||||
async def test_announce_uses_effective_key_in_unified_mode(self):
|
||||
"""In unified session mode, session_key_override must be 'unified:default'
|
||||
so the result matches the pending queue key."""
|
||||
so the result matches the manager mailbox session key."""
|
||||
mgr, bus = self._make_mgr()
|
||||
|
||||
origin = {"channel": "telegram", "chat_id": "111", "session_key": UNIFIED_SESSION_KEY}
|
||||
await mgr._announce_result("sub-1", "label", "task", "result", origin, "ok")
|
||||
|
||||
msg = await bus.consume_inbound()
|
||||
assert msg.session_key_override == UNIFIED_SESSION_KEY
|
||||
assert msg.session_key == UNIFIED_SESSION_KEY
|
||||
assert bus.inbound.empty()
|
||||
snapshots = await mgr.mailbox.poll("unified:default", task_id="sub-1")
|
||||
assert snapshots[0].session_key == "unified:default"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_announce_uses_raw_key_in_normal_mode(self):
|
||||
"""Without unified sessions, session_key_override is the raw channel:chat_id."""
|
||||
"""Without unified sessions, the mailbox session is the raw channel:chat_id."""
|
||||
mgr, bus = self._make_mgr()
|
||||
|
||||
origin = {"channel": "telegram", "chat_id": "222", "session_key": "telegram:222"}
|
||||
await mgr._announce_result("sub-2", "label", "task", "result", origin, "ok")
|
||||
|
||||
msg = await bus.consume_inbound()
|
||||
assert msg.session_key_override == "telegram:222"
|
||||
assert msg.session_key == "telegram:222"
|
||||
assert bus.inbound.empty()
|
||||
snapshots = await mgr.mailbox.poll("telegram:222", task_id="sub-2")
|
||||
assert snapshots[0].session_key == "telegram:222"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_announce_falls_back_to_origin_when_no_session_key(self):
|
||||
@@ -478,10 +478,9 @@ class TestSubagentAnnounceSessionKey:
|
||||
origin = {"channel": "discord", "chat_id": "333", "session_key": None}
|
||||
await mgr._announce_result("sub-3", "label", "task", "result", origin, "ok")
|
||||
|
||||
msg = await bus.consume_inbound()
|
||||
assert msg.session_key_override == "discord:333"
|
||||
assert msg.channel == "system"
|
||||
assert msg.chat_id == "discord:333"
|
||||
assert bus.inbound.empty()
|
||||
snapshots = await mgr.mailbox.poll("discord:333", task_id="sub-3")
|
||||
assert snapshots[0].session_key == "discord:333"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_key_flows_through_run_subagent(self):
|
||||
@@ -510,5 +509,6 @@ class TestSubagentAnnounceSessionKey:
|
||||
status,
|
||||
)
|
||||
|
||||
msg = await bus.consume_inbound()
|
||||
assert msg.session_key_override == UNIFIED_SESSION_KEY
|
||||
assert bus.inbound.empty()
|
||||
snapshots = await mgr.mailbox.poll("unified:default", task_id="sub-4")
|
||||
assert snapshots[0].session_key == "unified:default"
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Tests for explicit subagent mailbox tools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.runner import AgentRunResult
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.agent.tools.context import RequestContext
|
||||
from nanobot.agent.tools.subagent_mailbox import (
|
||||
CancelSubagentTool,
|
||||
PollSubagentsTool,
|
||||
WaitSubagentsTool,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
|
||||
|
||||
def _manager(tmp_path: Path) -> SubagentManager:
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
return SubagentManager(
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
bus=MessageBus(),
|
||||
max_tool_result_chars=AgentDefaults().max_tool_result_chars,
|
||||
)
|
||||
|
||||
|
||||
def _bind(tool, session_key: str = "cli:test") -> None:
|
||||
tool.set_context(RequestContext(channel="cli", chat_id="test", session_key=session_key))
|
||||
|
||||
|
||||
async def _drain(mgr: SubagentManager) -> None:
|
||||
tasks = list(mgr._running_tasks.values())
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_subagents_returns_result_once(tmp_path: Path) -> None:
|
||||
mgr = _manager(tmp_path)
|
||||
mgr.runner.run = AsyncMock(
|
||||
return_value=AgentRunResult(final_content="worker result", messages=[], stop_reason="completed")
|
||||
)
|
||||
|
||||
await mgr.spawn("do work", label="worker", session_key="cli:test")
|
||||
task_id = next(iter(mgr._running_tasks))
|
||||
await _drain(mgr)
|
||||
|
||||
wait_tool = WaitSubagentsTool(mgr)
|
||||
_bind(wait_tool)
|
||||
|
||||
first = await wait_tool.execute(task_id=task_id, timeout_seconds=0)
|
||||
second = await wait_tool.execute(task_id=task_id, timeout_seconds=0)
|
||||
|
||||
assert "worker result" in first
|
||||
assert f"id: {task_id}" in first
|
||||
assert "already consumed" in second
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_subagents_reads_result_after_manager_recreation(tmp_path: Path) -> None:
|
||||
mgr = _manager(tmp_path)
|
||||
mgr.runner.run = AsyncMock(
|
||||
return_value=AgentRunResult(final_content="durable worker result", messages=[], stop_reason="completed")
|
||||
)
|
||||
|
||||
await mgr.spawn("do durable work", label="worker", session_key="cli:test")
|
||||
task_id = next(iter(mgr._running_tasks))
|
||||
await _drain(mgr)
|
||||
|
||||
recreated = _manager(tmp_path)
|
||||
wait_tool = WaitSubagentsTool(recreated)
|
||||
poll_tool = PollSubagentsTool(recreated)
|
||||
_bind(wait_tool)
|
||||
_bind(poll_tool)
|
||||
|
||||
first = await wait_tool.execute(task_id=task_id, timeout_seconds=0)
|
||||
after = await poll_tool.execute(task_id=task_id)
|
||||
|
||||
assert "durable worker result" in first
|
||||
assert "result consumed" in after
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poll_subagents_reports_running_completed_and_not_found(tmp_path: Path) -> None:
|
||||
mgr = _manager(tmp_path)
|
||||
release = asyncio.Event()
|
||||
|
||||
async def _run(_spec):
|
||||
await release.wait()
|
||||
return AgentRunResult(final_content="done", messages=[], stop_reason="completed")
|
||||
|
||||
mgr.runner.run = AsyncMock(side_effect=_run)
|
||||
await mgr.spawn("slow work", label="slow", session_key="cli:test")
|
||||
task_id = next(iter(mgr._running_tasks))
|
||||
|
||||
poll_tool = PollSubagentsTool(mgr)
|
||||
_bind(poll_tool)
|
||||
|
||||
running = await poll_tool.execute(task_id=task_id)
|
||||
missing = await poll_tool.execute(task_id="missing")
|
||||
release.set()
|
||||
await _drain(mgr)
|
||||
completed = await poll_tool.execute(task_id=task_id)
|
||||
|
||||
assert "status: running" in running
|
||||
assert "not found" in missing
|
||||
assert "completed, result ready" in completed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_subagent_marks_cancelled_result(tmp_path: Path) -> None:
|
||||
mgr = _manager(tmp_path)
|
||||
started = asyncio.Event()
|
||||
|
||||
async def _run(_spec):
|
||||
started.set()
|
||||
await asyncio.Event().wait()
|
||||
|
||||
mgr.runner.run = AsyncMock(side_effect=_run)
|
||||
await mgr.spawn("slow work", label="slow", session_key="cli:test")
|
||||
task_id = next(iter(mgr._running_tasks))
|
||||
await asyncio.wait_for(started.wait(), timeout=1.0)
|
||||
|
||||
cancel_tool = CancelSubagentTool(mgr)
|
||||
wait_tool = WaitSubagentsTool(mgr)
|
||||
_bind(cancel_tool)
|
||||
_bind(wait_tool)
|
||||
|
||||
cancelled = await cancel_tool.execute(task_id=task_id)
|
||||
result = await wait_tool.execute(task_id=task_id, timeout_seconds=0)
|
||||
|
||||
assert cancelled == f"Cancelled subagent task {task_id}."
|
||||
assert "status: cancelled" in result
|
||||
assert "Cancelled by manager." in result
|
||||
@@ -279,8 +279,8 @@ async def test_agent_loop_syncs_updated_max_iterations_before_run(tmp_path):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_drain_pending_blocks_while_subagents_running(tmp_path):
|
||||
"""_drain_pending should block when no messages are available but sub-agents are still running."""
|
||||
async def test_drain_pending_does_not_block_while_subagents_running(tmp_path):
|
||||
"""_drain_pending should ignore running workers unless user messages are queued."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
@@ -336,31 +336,24 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path):
|
||||
|
||||
assert injection_callback is not None
|
||||
|
||||
# Now test the callback directly
|
||||
# With sub-agents running and an empty queue, it should block
|
||||
drain_task = asyncio.create_task(injection_callback())
|
||||
# Running subagents alone must not keep the current turn alive.
|
||||
results = await asyncio.wait_for(injection_callback(), timeout=1.0)
|
||||
assert results == []
|
||||
|
||||
# Let the task enter the blocking queue wait.
|
||||
await asyncio.sleep(0)
|
||||
|
||||
# Should still be running (blocked on pending_queue.get())
|
||||
assert not drain_task.done(), "drain should block while sub-agents are running"
|
||||
|
||||
# Now put a message in the queue (simulating sub-agent completion)
|
||||
# Real follow-up messages still use the ordinary pending queue path.
|
||||
await pending_queue.put(InboundMessage(
|
||||
sender_id="subagent",
|
||||
sender_id="user",
|
||||
channel="test",
|
||||
chat_id="c1",
|
||||
content="Sub-agent result",
|
||||
content="User follow-up",
|
||||
media=None,
|
||||
metadata={},
|
||||
))
|
||||
|
||||
# Should unblock and return results
|
||||
results = await asyncio.wait_for(drain_task, timeout=2.0)
|
||||
results = await asyncio.wait_for(injection_callback(), timeout=1.0)
|
||||
assert len(results) >= 1
|
||||
assert results[0]["role"] == "user"
|
||||
assert "Sub-agent result" in str(results[0]["content"])
|
||||
assert "User follow-up" in str(results[0]["content"])
|
||||
|
||||
# Cleanup
|
||||
hang_task.cancel()
|
||||
@@ -417,8 +410,8 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_drain_pending_timeout(tmp_path):
|
||||
"""_drain_pending should return empty after timeout when sub-agents hang."""
|
||||
async def test_drain_pending_does_not_wait_for_hung_subagents(tmp_path):
|
||||
"""_drain_pending should not call asyncio.wait_for for hung subagents."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.session.manager import Session
|
||||
@@ -467,14 +460,10 @@ async def test_drain_pending_timeout(tmp_path):
|
||||
|
||||
assert injection_callback is not None
|
||||
|
||||
# Patch the timeout path without leaking the queue.get() coroutine.
|
||||
async def _timeout(awaitable, timeout):
|
||||
awaitable.close()
|
||||
raise asyncio.TimeoutError
|
||||
|
||||
with patch("nanobot.agent.loop.asyncio.wait_for", side_effect=_timeout):
|
||||
with patch("nanobot.agent.loop.asyncio.wait_for") as wait_for:
|
||||
results = await injection_callback()
|
||||
assert results == []
|
||||
wait_for.assert_not_called()
|
||||
|
||||
# Cleanup
|
||||
hang_task.cancel()
|
||||
|
||||
@@ -14,12 +14,16 @@ from nanobot.session.turn_continuation import (
|
||||
INTERNAL_CONTINUATION_META,
|
||||
INTERNAL_CONTINUATION_PENDING_META,
|
||||
INTERNAL_CONTINUATION_RUN_STARTED_AT_META,
|
||||
SUBAGENT_RESULT_CONTINUATION_KIND,
|
||||
_save_skip_for_turn,
|
||||
internal_continuation_pending,
|
||||
internal_continuation_run_started_at,
|
||||
maybe_continue_turn,
|
||||
should_finalize_on_max_iterations,
|
||||
should_stream_budget_response,
|
||||
subagent_result_continuation_inbound,
|
||||
subagent_result_continuation_metadata,
|
||||
subagent_result_continuation_task_id,
|
||||
)
|
||||
|
||||
|
||||
@@ -165,3 +169,23 @@ def test_save_skip_unchanged_for_standalone_current_message():
|
||||
history_count=1,
|
||||
user_persisted_early=False,
|
||||
) == 2
|
||||
|
||||
|
||||
def test_subagent_result_continuation_metadata():
|
||||
meta = subagent_result_continuation_metadata(
|
||||
{
|
||||
"message_id": "msg-1",
|
||||
"_stream_id": "old-stream",
|
||||
"_stream_delta": True,
|
||||
},
|
||||
task_id="sub-1",
|
||||
run_started_at=42.0,
|
||||
)
|
||||
|
||||
assert meta[INTERNAL_CONTINUATION_META] is True
|
||||
assert meta[INTERNAL_CONTINUATION_KIND_META] == SUBAGENT_RESULT_CONTINUATION_KIND
|
||||
assert meta[INTERNAL_CONTINUATION_RUN_STARTED_AT_META] == 42.0
|
||||
assert subagent_result_continuation_inbound(meta)
|
||||
assert subagent_result_continuation_task_id(meta) == "sub-1"
|
||||
assert "_stream_id" not in meta
|
||||
assert "_stream_delta" not in meta
|
||||
|
||||
Reference in New Issue
Block a user