mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-01 00:31:51 +03:00
* refactor(agent): defer transcript assembly to runner Keep persisted history and the fresh turn as explicit inputs until the Runner assembles the provider transcript. Preserve ContextBuilder and direct AgentRunner compatibility while making the save boundary structural. Refs NAN-81. * fix(providers): preserve mixed adjacent user content
298 lines
9.0 KiB
Python
298 lines
9.0 KiB
Python
import asyncio
|
|
import inspect
|
|
from pathlib import Path
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
from nanobot.agent.context import TranscriptInput
|
|
from nanobot.agent.loop import AgentLoop
|
|
from nanobot.agent.tools.context import (
|
|
RequestContext,
|
|
bind_request_context,
|
|
current_request_context,
|
|
reset_request_context,
|
|
)
|
|
from nanobot.agent.tools.registry import ToolRegistry
|
|
from nanobot.bus.events import InboundMessage
|
|
from nanobot.bus.queue import MessageBus
|
|
from nanobot.config.schema import Config
|
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
|
from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META
|
|
|
|
|
|
class _ContextRecordingTool:
|
|
name = "cron"
|
|
concurrency_safe = False
|
|
|
|
def __init__(self) -> None:
|
|
self.contexts: list[dict] = []
|
|
self.runtimes: list[object] = []
|
|
|
|
async def execute(self, **_kwargs) -> str:
|
|
ctx = current_request_context()
|
|
assert ctx is not None
|
|
self.runtimes.append(ctx.runtime)
|
|
self.contexts.append({
|
|
"channel": ctx.channel,
|
|
"chat_id": ctx.chat_id,
|
|
"metadata": ctx.metadata,
|
|
"session_key": ctx.session_key,
|
|
})
|
|
return "created"
|
|
|
|
|
|
class _Tools:
|
|
def __init__(self, tool: _ContextRecordingTool) -> None:
|
|
self.tool = tool
|
|
|
|
@property
|
|
def tool_names(self) -> list[str]:
|
|
return ["cron"]
|
|
|
|
def get(self, name: str):
|
|
return self.tool if name == "cron" else None
|
|
|
|
def get_definitions(self) -> list:
|
|
return []
|
|
|
|
def prepare_call(self, name: str, arguments: dict):
|
|
return (self.tool, arguments, None) if name == "cron" else (None, arguments, None)
|
|
|
|
|
|
def test_loop_registers_default_tools_in_injected_registry(tmp_path: Path) -> None:
|
|
provider = MagicMock()
|
|
provider.get_default_model.return_value = "test-model"
|
|
registry = ToolRegistry()
|
|
|
|
loop = AgentLoop(
|
|
bus=MessageBus(),
|
|
provider=provider,
|
|
workspace=tmp_path,
|
|
tool_registry=registry,
|
|
)
|
|
|
|
assert loop.tools is registry
|
|
assert registry.has("read_file")
|
|
|
|
|
|
def _config_for_loop(tmp_path: Path) -> Config:
|
|
return Config.model_validate({"agents": {"defaults": {"workspace": str(tmp_path)}}})
|
|
|
|
|
|
def _provider_for_loop() -> MagicMock:
|
|
provider = MagicMock()
|
|
provider.get_default_model.return_value = "test-model"
|
|
return provider
|
|
|
|
|
|
def test_loop_from_config_requires_caller_owned_registry(tmp_path: Path) -> None:
|
|
signature = inspect.signature(AgentLoop.from_config)
|
|
|
|
with pytest.raises(TypeError, match="tool_registry"):
|
|
signature.bind(_config_for_loop(tmp_path))
|
|
|
|
|
|
def test_loop_from_config_uses_caller_owned_registry(tmp_path: Path) -> None:
|
|
registry = ToolRegistry()
|
|
loop = AgentLoop.from_config(
|
|
_config_for_loop(tmp_path),
|
|
tool_registry=registry,
|
|
provider=_provider_for_loop(),
|
|
)
|
|
|
|
assert loop.tools is registry
|
|
assert loop.tools.has("read_file")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) -> None:
|
|
provider = MagicMock()
|
|
calls = {"n": 0}
|
|
|
|
async def chat_with_retry(**_kwargs):
|
|
calls["n"] += 1
|
|
if calls["n"] == 1:
|
|
return LLMResponse(
|
|
content=None,
|
|
tool_calls=[ToolCallRequest(id="call_1", name="cron", arguments={"action": "add"})],
|
|
)
|
|
return LLMResponse(content="done", tool_calls=[])
|
|
|
|
provider.chat_with_retry = chat_with_retry
|
|
provider.get_default_model.return_value = "test-model"
|
|
|
|
loop = AgentLoop(
|
|
bus=MessageBus(),
|
|
provider=provider,
|
|
workspace=tmp_path,
|
|
model="test-model",
|
|
)
|
|
cron = _ContextRecordingTool()
|
|
loop.tools = _Tools(cron)
|
|
|
|
metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}}
|
|
runtime = loop.llm_runtime()
|
|
await loop._run_agent_loop(
|
|
TranscriptInput(history=[], current_message=None),
|
|
runtime=runtime,
|
|
request_context=RequestContext(
|
|
channel="slack",
|
|
chat_id="C123",
|
|
session_key="slack:C123:111.222",
|
|
runtime=runtime,
|
|
metadata=metadata,
|
|
),
|
|
)
|
|
|
|
assert cron.contexts[-1] == {
|
|
"channel": "slack",
|
|
"chat_id": "C123",
|
|
"metadata": metadata,
|
|
"session_key": "slack:C123:111.222",
|
|
}
|
|
assert cron.runtimes[-1] is runtime
|
|
|
|
|
|
def test_request_context_nested_bind_restores_outer_context() -> None:
|
|
outer = RequestContext(channel="slack", chat_id="outer", session_key="slack:outer")
|
|
inner = RequestContext(channel="email", chat_id="inner", session_key="email:inner")
|
|
|
|
outer_token = bind_request_context(outer)
|
|
try:
|
|
assert current_request_context() is outer
|
|
inner_token = bind_request_context(inner)
|
|
try:
|
|
assert current_request_context() is inner
|
|
finally:
|
|
reset_request_context(inner_token)
|
|
assert current_request_context() is outer
|
|
finally:
|
|
reset_request_context(outer_token)
|
|
|
|
assert current_request_context() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_request_context_bindings_are_isolated_between_concurrent_tasks() -> None:
|
|
entered = asyncio.Event()
|
|
release = asyncio.Event()
|
|
|
|
async def observe(ctx: RequestContext, *, wait_first: bool) -> RequestContext | None:
|
|
token = bind_request_context(ctx)
|
|
try:
|
|
if wait_first:
|
|
entered.set()
|
|
await release.wait()
|
|
else:
|
|
await entered.wait()
|
|
release.set()
|
|
await asyncio.sleep(0)
|
|
return current_request_context()
|
|
finally:
|
|
reset_request_context(token)
|
|
|
|
first = RequestContext(channel="feishu", chat_id="first", session_key="feishu:first")
|
|
second = RequestContext(channel="telegram", chat_id="second", session_key="telegram:second")
|
|
|
|
observed = await asyncio.gather(
|
|
observe(first, wait_first=True),
|
|
observe(second, wait_first=False),
|
|
)
|
|
|
|
assert observed == [first, second]
|
|
assert current_request_context() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_agent_loop_restores_outer_request_context_after_runner_exception(
|
|
tmp_path: Path,
|
|
) -> None:
|
|
provider = MagicMock()
|
|
provider.get_default_model.return_value = "test-model"
|
|
loop = AgentLoop(
|
|
bus=MessageBus(),
|
|
provider=provider,
|
|
workspace=tmp_path,
|
|
model="test-model",
|
|
)
|
|
outer = RequestContext(channel="test", chat_id="outer", session_key="test:outer")
|
|
runtime = loop.llm_runtime()
|
|
|
|
async def fail_run(spec):
|
|
current = current_request_context()
|
|
assert current is not None
|
|
assert spec.runtime is runtime
|
|
assert current.runtime is runtime
|
|
assert current.channel == "slack"
|
|
assert current.chat_id == "C123"
|
|
assert current.session_key == "slack:C123:111.222"
|
|
assert current.original_user_text == " unchanged user text "
|
|
raise RuntimeError("runner failed")
|
|
|
|
loop.runner.run = AsyncMock(side_effect=fail_run)
|
|
outer_token = bind_request_context(outer)
|
|
try:
|
|
with pytest.raises(RuntimeError, match="runner failed"):
|
|
await loop._run_agent_loop(
|
|
TranscriptInput(history=[], current_message=None),
|
|
runtime=runtime,
|
|
request_context=RequestContext(
|
|
channel="slack",
|
|
chat_id="C123",
|
|
session_key="slack:C123:111.222",
|
|
original_user_text=" unchanged user text ",
|
|
runtime=runtime,
|
|
),
|
|
)
|
|
assert current_request_context() is outer
|
|
finally:
|
|
reset_request_context(outer_token)
|
|
|
|
assert current_request_context() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
("metadata", "expected"),
|
|
[
|
|
({}, " original user text "),
|
|
({INTERNAL_CONTINUATION_META: True}, None),
|
|
],
|
|
)
|
|
async def test_process_message_captures_original_text_before_restore(
|
|
tmp_path: Path,
|
|
metadata: dict,
|
|
expected: str | None,
|
|
) -> None:
|
|
provider = MagicMock()
|
|
provider.get_default_model.return_value = "test-model"
|
|
loop = AgentLoop(
|
|
bus=MessageBus(),
|
|
provider=provider,
|
|
workspace=tmp_path,
|
|
model="test-model",
|
|
)
|
|
runtime = loop.llm_runtime()
|
|
seen: list[tuple[str | None, object]] = []
|
|
|
|
async def stop_after_capture(ctx) -> str:
|
|
seen.append((ctx.original_user_text, ctx.runtime))
|
|
raise RuntimeError("captured before restore")
|
|
|
|
loop._restore_turn = stop_after_capture # type: ignore[method-assign]
|
|
|
|
with pytest.raises(RuntimeError, match="captured before restore"):
|
|
await loop._process_message(
|
|
InboundMessage(
|
|
channel="slack",
|
|
sender_id="user",
|
|
chat_id="C123",
|
|
content=" original user text ",
|
|
metadata=metadata,
|
|
),
|
|
runtime=runtime,
|
|
)
|
|
|
|
assert seen == [(expected, runtime)]
|