mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 08:13:11 +03:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
234da592f5 | ||
|
|
2d138b92fc |
+20
-71
@@ -75,23 +75,6 @@ class PersistedPromptContextResolver:
|
|||||||
return channel, scope.project_path
|
return channel, scope.project_path
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
|
||||||
class TranscriptInput:
|
|
||||||
"""Raw turn inputs from which ``ContextBuilder`` assembles a transcript."""
|
|
||||||
|
|
||||||
history: list[dict[str, Any]]
|
|
||||||
current_message: str | None
|
|
||||||
media: Sequence[str] | None = None
|
|
||||||
current_role: str = "user"
|
|
||||||
session_summary: SessionSummary | None = None
|
|
||||||
runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def message_count(self) -> int:
|
|
||||||
"""Number of boundary-preserving messages in the assembled transcript."""
|
|
||||||
return 1 + len(self.history) + (self.current_message is not None)
|
|
||||||
|
|
||||||
|
|
||||||
class ContextBuilder:
|
class ContextBuilder:
|
||||||
"""Builds the context (system prompt + messages) for the agent."""
|
"""Builds the context (system prompt + messages) for the agent."""
|
||||||
|
|
||||||
@@ -299,58 +282,14 @@ class ContextBuilder:
|
|||||||
session_key: str | None = None,
|
session_key: str | None = None,
|
||||||
unified_session: bool = False,
|
unified_session: bool = False,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Compatibility wrapper for callers that need merged adjacent roles."""
|
"""Build the complete message list for an LLM call."""
|
||||||
messages = self.build_transcript(
|
|
||||||
TranscriptInput(
|
|
||||||
history=history,
|
|
||||||
current_message=current_message,
|
|
||||||
media=media,
|
|
||||||
current_role=current_role,
|
|
||||||
session_summary=session_summary,
|
|
||||||
runtime_context_blocks=runtime_context_blocks,
|
|
||||||
),
|
|
||||||
channel=channel,
|
|
||||||
workspace=workspace,
|
|
||||||
include_memory=include_memory,
|
|
||||||
include_memory_recent_history=include_memory_recent_history,
|
|
||||||
session_key=session_key,
|
|
||||||
unified_session=unified_session,
|
|
||||||
)
|
|
||||||
current = messages[-1]
|
|
||||||
if len(messages) < 2 or messages[-2].get("role") != current.get("role"):
|
|
||||||
return messages
|
|
||||||
|
|
||||||
merged = dict(messages[-2])
|
|
||||||
merged["content"] = self._merge_message_content(
|
|
||||||
merged.get("content"),
|
|
||||||
current.get("content"),
|
|
||||||
)
|
|
||||||
current_meta = current.get("_meta")
|
|
||||||
if current.get("role") == "user" and isinstance(current_meta, dict):
|
|
||||||
internal_meta = dict(merged.get("_meta") or {})
|
|
||||||
internal_meta.update(cast(dict[str, Any], current_meta))
|
|
||||||
merged["_meta"] = internal_meta
|
|
||||||
return [*messages[:-2], merged]
|
|
||||||
|
|
||||||
def build_transcript(
|
|
||||||
self,
|
|
||||||
transcript: TranscriptInput,
|
|
||||||
*,
|
|
||||||
channel: str | None = None,
|
|
||||||
workspace: Path | None = None,
|
|
||||||
include_memory: bool = True,
|
|
||||||
include_memory_recent_history: bool = True,
|
|
||||||
session_key: str | None = None,
|
|
||||||
unified_session: bool = False,
|
|
||||||
) -> list[dict[str, Any]]:
|
|
||||||
"""Build a model transcript while preserving the fresh-turn boundary."""
|
|
||||||
root = workspace or self.workspace
|
root = workspace or self.workspace
|
||||||
messages: list[dict[str, Any]] = [
|
messages: list[dict[str, Any]] = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": self.build_system_prompt(
|
"content": self.build_system_prompt(
|
||||||
channel=channel,
|
channel=channel,
|
||||||
session_summary=transcript.session_summary,
|
session_summary=session_summary,
|
||||||
workspace=root,
|
workspace=root,
|
||||||
include_memory=include_memory,
|
include_memory=include_memory,
|
||||||
include_memory_recent_history=include_memory_recent_history,
|
include_memory_recent_history=include_memory_recent_history,
|
||||||
@@ -358,17 +297,27 @@ class ContextBuilder:
|
|||||||
unified_session=unified_session,
|
unified_session=unified_session,
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
*transcript.history,
|
*history,
|
||||||
]
|
]
|
||||||
if transcript.current_message is None:
|
|
||||||
return messages
|
|
||||||
|
|
||||||
current = self.build_current_message(
|
current = self.build_current_message(
|
||||||
transcript.current_message,
|
current_message,
|
||||||
media=list(transcript.media) if transcript.media else None,
|
media=media,
|
||||||
current_role=transcript.current_role,
|
current_role=current_role,
|
||||||
runtime_context_blocks=transcript.runtime_context_blocks,
|
runtime_context_blocks=runtime_context_blocks,
|
||||||
)
|
)
|
||||||
|
if messages[-1].get("role") == current_role:
|
||||||
|
last = dict(messages[-1])
|
||||||
|
last["content"] = self._merge_message_content(
|
||||||
|
last.get("content"),
|
||||||
|
current.get("content"),
|
||||||
|
)
|
||||||
|
current_meta = current.get("_meta")
|
||||||
|
if current_role == "user" and isinstance(current_meta, dict):
|
||||||
|
internal_meta = dict(last.get("_meta") or {})
|
||||||
|
internal_meta.update(cast(dict[str, Any], current_meta))
|
||||||
|
last["_meta"] = internal_meta
|
||||||
|
messages[-1] = last
|
||||||
|
return messages
|
||||||
messages.append(current)
|
messages.append(current)
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
|
|||||||
+16
-22
@@ -14,7 +14,6 @@ from collections.abc import Coroutine, Iterable, Mapping
|
|||||||
from contextlib import AbstractContextManager, ExitStack, nullcontext, suppress
|
from contextlib import AbstractContextManager, ExitStack, nullcontext, suppress
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from functools import partial
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypeVar, cast
|
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypeVar, cast
|
||||||
|
|
||||||
@@ -24,7 +23,7 @@ from nanobot.agent import context as agent_context
|
|||||||
from nanobot.agent import model_presets as preset_helpers
|
from nanobot.agent import model_presets as preset_helpers
|
||||||
from nanobot.agent.autocompact import AutoCompact
|
from nanobot.agent.autocompact import AutoCompact
|
||||||
from nanobot.agent.automation_turns import publish_next_deferred_turn
|
from nanobot.agent.automation_turns import publish_next_deferred_turn
|
||||||
from nanobot.agent.context import ContextBuilder, PersistedPromptContextResolver, TranscriptInput
|
from nanobot.agent.context import ContextBuilder, PersistedPromptContextResolver
|
||||||
from nanobot.agent.cron_turns import CronTurnCoordinator
|
from nanobot.agent.cron_turns import CronTurnCoordinator
|
||||||
from nanobot.agent.hook import AgentHook, AgentTurnHookFactory
|
from nanobot.agent.hook import AgentHook, AgentTurnHookFactory
|
||||||
from nanobot.agent.memory import Consolidator
|
from nanobot.agent.memory import Consolidator
|
||||||
@@ -136,7 +135,7 @@ class TurnContext:
|
|||||||
session: Session | None = None
|
session: Session | None = None
|
||||||
|
|
||||||
history: list[dict[str, Any]] = field(default_factory=list)
|
history: list[dict[str, Any]] = field(default_factory=list)
|
||||||
transcript_input: TranscriptInput | None = None
|
initial_messages: list[dict[str, Any]] = field(default_factory=list)
|
||||||
provider_state: ProviderConversationState | None = field(default=None, repr=False)
|
provider_state: ProviderConversationState | None = field(default=None, repr=False)
|
||||||
request_context: RequestContext | None = None
|
request_context: RequestContext | None = None
|
||||||
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
|
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
|
||||||
@@ -724,15 +723,22 @@ class AgentLoop:
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _build_transcript_input(self, ctx: TurnContext) -> TranscriptInput:
|
def _build_initial_messages(self, ctx: TurnContext) -> list[dict[str, Any]]:
|
||||||
"""Capture the persisted history and fresh input as separate transcript parts."""
|
"""Build the initial message list for the LLM turn."""
|
||||||
assert ctx.session is not None
|
assert ctx.session is not None
|
||||||
return TranscriptInput(
|
scope = self.workspace_scopes.for_message(ctx.msg, ctx.session.metadata)
|
||||||
|
return self.context.build_messages(
|
||||||
history=ctx.history,
|
history=ctx.history,
|
||||||
current_message=ctx.msg.content,
|
current_message=ctx.msg.content,
|
||||||
media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
|
media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
|
||||||
|
channel=ctx.delivery.route.channel,
|
||||||
session_summary=ctx.pending_summary,
|
session_summary=ctx.pending_summary,
|
||||||
|
workspace=scope.project_path,
|
||||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||||
|
include_memory=ctx.session.policy.persist,
|
||||||
|
include_memory_recent_history=not ctx.ephemeral,
|
||||||
|
session_key=ctx.session.key,
|
||||||
|
unified_session=self._unified_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _request_context_for_turn(self, ctx: TurnContext) -> RequestContext:
|
def _request_context_for_turn(self, ctx: TurnContext) -> RequestContext:
|
||||||
@@ -923,7 +929,7 @@ class AgentLoop:
|
|||||||
|
|
||||||
async def _run_agent_loop(
|
async def _run_agent_loop(
|
||||||
self,
|
self,
|
||||||
transcript_input: TranscriptInput,
|
initial_messages: list[dict[str, Any]],
|
||||||
on_progress: Callable[..., Awaitable[None]] | None = None,
|
on_progress: Callable[..., Awaitable[None]] | None = None,
|
||||||
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||||
@@ -1104,15 +1110,6 @@ class AgentLoop:
|
|||||||
message_metadata=request_metadata,
|
message_metadata=request_metadata,
|
||||||
session_metadata=session.metadata if session is not None else None,
|
session_metadata=session.metadata if session is not None else None,
|
||||||
)
|
)
|
||||||
transcript_builder = partial(
|
|
||||||
self.context.build_transcript,
|
|
||||||
channel=request_ctx.channel,
|
|
||||||
workspace=effective_scope.project_path,
|
|
||||||
include_memory=session.policy.persist if session is not None else True,
|
|
||||||
include_memory_recent_history=not ephemeral,
|
|
||||||
session_key=session.key if session is not None else request_ctx.session_key,
|
|
||||||
unified_session=self._unified_session,
|
|
||||||
)
|
|
||||||
if request_context is None:
|
if request_context is None:
|
||||||
request_ctx = dataclasses.replace(
|
request_ctx = dataclasses.replace(
|
||||||
request_ctx,
|
request_ctx,
|
||||||
@@ -1159,13 +1156,11 @@ class AgentLoop:
|
|||||||
run_extra_hooks_for_ephemeral=run_extra_hooks_for_ephemeral,
|
run_extra_hooks_for_ephemeral=run_extra_hooks_for_ephemeral,
|
||||||
))
|
))
|
||||||
result = await self.runner.run(AgentRunSpec(
|
result = await self.runner.run(AgentRunSpec(
|
||||||
initial_messages=None,
|
initial_messages=initial_messages,
|
||||||
tools=effective_tools,
|
tools=effective_tools,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
max_iterations=self.max_iterations,
|
max_iterations=self.max_iterations,
|
||||||
max_tool_result_chars=self.max_tool_result_chars,
|
max_tool_result_chars=self.max_tool_result_chars,
|
||||||
transcript_input=transcript_input,
|
|
||||||
transcript_builder=transcript_builder,
|
|
||||||
hook=hook,
|
hook=hook,
|
||||||
concurrent_tools=True,
|
concurrent_tools=True,
|
||||||
workspace=effective_scope.project_path,
|
workspace=effective_scope.project_path,
|
||||||
@@ -1973,7 +1968,7 @@ class AgentLoop:
|
|||||||
# Upgrade the replay-safe baseline to the resumable state before
|
# Upgrade the replay-safe baseline to the resumable state before
|
||||||
# prompt assembly and the first model checkpoint.
|
# prompt assembly and the first model checkpoint.
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
ctx.transcript_input = self._build_transcript_input(ctx)
|
ctx.initial_messages = self._build_initial_messages(ctx)
|
||||||
|
|
||||||
if ctx.on_progress is None:
|
if ctx.on_progress is None:
|
||||||
ctx.on_progress = ctx.delivery.progress_callback()
|
ctx.on_progress = ctx.delivery.progress_callback()
|
||||||
@@ -1985,10 +1980,9 @@ class AgentLoop:
|
|||||||
if ctx.visible_run_started_at is None:
|
if ctx.visible_run_started_at is None:
|
||||||
ctx.visible_run_started_at = time.time()
|
ctx.visible_run_started_at = time.time()
|
||||||
await ctx.delivery.running(started_at=ctx.visible_run_started_at)
|
await ctx.delivery.running(started_at=ctx.visible_run_started_at)
|
||||||
assert ctx.transcript_input is not None
|
|
||||||
with capture_message_deliveries() as message_sends:
|
with capture_message_deliveries() as message_sends:
|
||||||
result = await self._run_agent_loop(
|
result = await self._run_agent_loop(
|
||||||
ctx.transcript_input,
|
ctx.initial_messages,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
on_progress=ctx.on_progress,
|
on_progress=ctx.on_progress,
|
||||||
on_stream=ctx.on_stream,
|
on_stream=ctx.on_stream,
|
||||||
|
|||||||
+8
-48
@@ -14,7 +14,6 @@ from typing import Any, cast
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.context_governance import (
|
from nanobot.agent.context_governance import (
|
||||||
ContextGovernanceConfig,
|
ContextGovernanceConfig,
|
||||||
ContextGovernor,
|
ContextGovernor,
|
||||||
@@ -67,7 +66,6 @@ ContinuationCallback = Callable[[], str | None]
|
|||||||
RetryWaitCallback = Callable[[str], Awaitable[None]]
|
RetryWaitCallback = Callable[[str], Awaitable[None]]
|
||||||
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
|
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
|
||||||
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
|
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
|
||||||
TranscriptBuilder = Callable[[TranscriptInput], list[dict[str, Any]]]
|
|
||||||
|
|
||||||
_DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model."
|
_DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model."
|
||||||
_ARREARAGE_ERROR_MESSAGE = (
|
_ARREARAGE_ERROR_MESSAGE = (
|
||||||
@@ -96,13 +94,11 @@ def _restore_outer_whitespace(content: str, original: str | None) -> str:
|
|||||||
class AgentRunSpec:
|
class AgentRunSpec:
|
||||||
"""Configuration for a single agent execution."""
|
"""Configuration for a single agent execution."""
|
||||||
|
|
||||||
initial_messages: list[dict[str, Any]] | None
|
initial_messages: list[dict[str, Any]]
|
||||||
tools: ToolRegistry
|
tools: ToolRegistry
|
||||||
runtime: LLMRuntime
|
runtime: LLMRuntime
|
||||||
max_iterations: int
|
max_iterations: int
|
||||||
max_tool_result_chars: int
|
max_tool_result_chars: int
|
||||||
transcript_input: TranscriptInput | None = None
|
|
||||||
transcript_builder: TranscriptBuilder | None = None
|
|
||||||
hook: AgentHook | None = None
|
hook: AgentHook | None = None
|
||||||
error_message: str | None = _DEFAULT_ERROR_MESSAGE
|
error_message: str | None = _DEFAULT_ERROR_MESSAGE
|
||||||
max_iterations_message: str | None = None
|
max_iterations_message: str | None = None
|
||||||
@@ -414,7 +410,7 @@ class AgentRunner:
|
|||||||
|
|
||||||
async def run(self, spec: AgentRunSpec) -> AgentRunResult:
|
async def run(self, spec: AgentRunSpec) -> AgentRunResult:
|
||||||
hook = spec.hook or AgentHook()
|
hook = spec.hook or AgentHook()
|
||||||
messages = self._initial_transcript(spec)
|
messages = list(spec.initial_messages)
|
||||||
context = AgentRunHookContext(messages=deepcopy(messages))
|
context = AgentRunHookContext(messages=deepcopy(messages))
|
||||||
llm_usage_source_token = bind_llm_usage_source(
|
llm_usage_source_token = bind_llm_usage_source(
|
||||||
spec.llm_usage_source or source_from_session_key(spec.session_key)
|
spec.llm_usage_source or source_from_session_key(spec.session_key)
|
||||||
@@ -466,19 +462,6 @@ class AgentRunner:
|
|||||||
finally:
|
finally:
|
||||||
reset_llm_usage_source(llm_usage_source_token)
|
reset_llm_usage_source(llm_usage_source_token)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _initial_transcript(spec: AgentRunSpec) -> list[dict[str, Any]]:
|
|
||||||
"""Resolve exactly one supported source for the initial model transcript."""
|
|
||||||
if spec.transcript_input is not None:
|
|
||||||
if spec.initial_messages is not None:
|
|
||||||
raise ValueError("provide either transcript_input or initial_messages, not both")
|
|
||||||
if spec.transcript_builder is None:
|
|
||||||
raise ValueError("transcript_builder is required with transcript_input")
|
|
||||||
return list(spec.transcript_builder(spec.transcript_input))
|
|
||||||
if spec.initial_messages is None:
|
|
||||||
raise ValueError("initial_messages is required without transcript_input")
|
|
||||||
return list(spec.initial_messages)
|
|
||||||
|
|
||||||
async def _run_core(
|
async def _run_core(
|
||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
@@ -519,7 +502,7 @@ class AgentRunner:
|
|||||||
context_window_tokens=spec.runtime.context_window_tokens,
|
context_window_tokens=spec.runtime.context_window_tokens,
|
||||||
context_block_limit=spec.context_block_limit,
|
context_block_limit=spec.context_block_limit,
|
||||||
max_tokens=spec.runtime.generation.max_tokens,
|
max_tokens=spec.runtime.generation.max_tokens,
|
||||||
inflight_start_index=len(messages),
|
inflight_start_index=len(spec.initial_messages),
|
||||||
)
|
)
|
||||||
|
|
||||||
for iteration in range(spec.max_iterations):
|
for iteration in range(spec.max_iterations):
|
||||||
@@ -966,7 +949,6 @@ class AgentRunner:
|
|||||||
|
|
||||||
active_hosted_tools: dict[str, dict[str, Any]] = {}
|
active_hosted_tools: dict[str, dict[str, Any]] = {}
|
||||||
native_reasoning_open = False
|
native_reasoning_open = False
|
||||||
native_reasoning_close_task: asyncio.Task[None] | None = None
|
|
||||||
request_started_at = 0.0
|
request_started_at = 0.0
|
||||||
first_output_at: float | None = None
|
first_output_at: float | None = None
|
||||||
generation_started_at: float | None = None
|
generation_started_at: float | None = None
|
||||||
@@ -990,29 +972,11 @@ class AgentRunner:
|
|||||||
generation_started_at = None
|
generation_started_at = None
|
||||||
|
|
||||||
async def _close_native_reasoning() -> None:
|
async def _close_native_reasoning() -> None:
|
||||||
nonlocal native_reasoning_open, native_reasoning_close_task
|
nonlocal native_reasoning_open
|
||||||
if native_reasoning_close_task is None:
|
if not native_reasoning_open:
|
||||||
if not native_reasoning_open:
|
return
|
||||||
return
|
native_reasoning_open = False
|
||||||
native_reasoning_open = False
|
await hook.emit_reasoning_end()
|
||||||
native_reasoning_close_task = asyncio.create_task(
|
|
||||||
hook.emit_reasoning_end()
|
|
||||||
)
|
|
||||||
|
|
||||||
close_task = native_reasoning_close_task
|
|
||||||
cancellation: asyncio.CancelledError | None = None
|
|
||||||
while not close_task.done():
|
|
||||||
try:
|
|
||||||
await asyncio.shield(close_task)
|
|
||||||
except asyncio.CancelledError as exc:
|
|
||||||
cancellation = cancellation or exc
|
|
||||||
try:
|
|
||||||
close_task.result()
|
|
||||||
finally:
|
|
||||||
if native_reasoning_close_task is close_task:
|
|
||||||
native_reasoning_close_task = None
|
|
||||||
if cancellation is not None:
|
|
||||||
raise cancellation
|
|
||||||
|
|
||||||
async def _provider_tool_event(event: dict[str, Any]) -> None:
|
async def _provider_tool_event(event: dict[str, Any]) -> None:
|
||||||
if event.get("kind") != "hosted_tool":
|
if event.get("kind") != "hosted_tool":
|
||||||
@@ -1087,10 +1051,6 @@ class AgentRunner:
|
|||||||
await coro if outer_timeout_s is None
|
await coro if outer_timeout_s is None
|
||||||
else await asyncio.wait_for(coro, timeout=outer_timeout_s)
|
else await asyncio.wait_for(coro, timeout=outer_timeout_s)
|
||||||
)
|
)
|
||||||
except asyncio.CancelledError:
|
|
||||||
_pause_generation()
|
|
||||||
await _close_native_reasoning()
|
|
||||||
raise
|
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
if outer_timeout_s is None:
|
if outer_timeout_s is None:
|
||||||
response = LLMResponse(
|
response = LLMResponse(
|
||||||
|
|||||||
@@ -43,13 +43,6 @@ _WORKSPACE_VIOLATION_MARKERS: tuple[str, ...] = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _with_retry_hint(payload: str) -> str:
|
|
||||||
"""Append the recovery hint exactly once."""
|
|
||||||
if payload.endswith(_RETRY_HINT):
|
|
||||||
return payload
|
|
||||||
return payload + _RETRY_HINT
|
|
||||||
|
|
||||||
|
|
||||||
async def execute_tool_calls(
|
async def execute_tool_calls(
|
||||||
tools: ToolRegistry,
|
tools: ToolRegistry,
|
||||||
tool_calls: list[ToolCallRequest],
|
tool_calls: list[ToolCallRequest],
|
||||||
@@ -112,7 +105,7 @@ async def _execute_tool_call(
|
|||||||
"status": "error",
|
"status": "error",
|
||||||
"detail": "repeated external lookup blocked",
|
"detail": "repeated external lookup blocked",
|
||||||
}
|
}
|
||||||
return _with_retry_hint(lookup_error), event
|
return lookup_error + _RETRY_HINT, event
|
||||||
|
|
||||||
prepare_call = cast(
|
prepare_call = cast(
|
||||||
Callable[[str, Any], object] | None,
|
Callable[[str, Any], object] | None,
|
||||||
@@ -126,7 +119,6 @@ async def _execute_tool_call(
|
|||||||
if len(prepared_tuple) == 3:
|
if len(prepared_tuple) == 3:
|
||||||
tool, params, prep_error = cast(tuple[Any, Any, str | None], prepared_tuple)
|
tool, params, prep_error = cast(tuple[Any, Any, str | None], prepared_tuple)
|
||||||
if prep_error:
|
if prep_error:
|
||||||
payload = _with_retry_hint(prep_error)
|
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
"status": "error",
|
"status": "error",
|
||||||
@@ -134,14 +126,14 @@ async def _execute_tool_call(
|
|||||||
}
|
}
|
||||||
handled = _classify_violation(
|
handled = _classify_violation(
|
||||||
raw_text=prep_error,
|
raw_text=prep_error,
|
||||||
soft_payload=payload,
|
soft_payload=prep_error + _RETRY_HINT,
|
||||||
event=event,
|
event=event,
|
||||||
tool_call=tool_call,
|
tool_call=tool_call,
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
workspace_violation_counts=workspace_violation_counts,
|
||||||
)
|
)
|
||||||
if handled is not None:
|
if handled is not None:
|
||||||
return handled
|
return handled
|
||||||
return payload, event
|
return prep_error + _RETRY_HINT, event
|
||||||
|
|
||||||
await hook.before_execute_tool(context, tool_call, tool, params)
|
await hook.before_execute_tool(context, tool_call, tool, params)
|
||||||
try:
|
try:
|
||||||
@@ -158,9 +150,10 @@ async def _execute_tool_call(
|
|||||||
"status": "error",
|
"status": "error",
|
||||||
"detail": str(exc),
|
"detail": str(exc),
|
||||||
}
|
}
|
||||||
payload = _with_retry_hint(f"Error: {type(exc).__name__}: {exc}")
|
payload = f"Error: {type(exc).__name__}: {exc}"
|
||||||
handled = _classify_violation(
|
handled = _classify_violation(
|
||||||
raw_text=str(exc),
|
raw_text=str(exc),
|
||||||
|
# Preserve legacy exception payloads without the retry hint.
|
||||||
soft_payload=payload,
|
soft_payload=payload,
|
||||||
event=event,
|
event=event,
|
||||||
tool_call=tool_call,
|
tool_call=tool_call,
|
||||||
@@ -172,7 +165,6 @@ async def _execute_tool_call(
|
|||||||
|
|
||||||
if is_tool_error_result(result):
|
if is_tool_error_result(result):
|
||||||
await hook.on_execute_tool_error(context, tool_call, tool, params, result)
|
await hook.on_execute_tool_error(context, tool_call, tool, params, result)
|
||||||
payload = _with_retry_hint(result)
|
|
||||||
event = {
|
event = {
|
||||||
"name": tool_call.name,
|
"name": tool_call.name,
|
||||||
"status": "error",
|
"status": "error",
|
||||||
@@ -180,14 +172,14 @@ async def _execute_tool_call(
|
|||||||
}
|
}
|
||||||
handled = _classify_violation(
|
handled = _classify_violation(
|
||||||
raw_text=result,
|
raw_text=result,
|
||||||
soft_payload=payload,
|
soft_payload=result + _RETRY_HINT,
|
||||||
event=event,
|
event=event,
|
||||||
tool_call=tool_call,
|
tool_call=tool_call,
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
workspace_violation_counts=workspace_violation_counts,
|
||||||
)
|
)
|
||||||
if handled is not None:
|
if handled is not None:
|
||||||
return handled
|
return handled
|
||||||
return payload, event
|
return result + _RETRY_HINT, event
|
||||||
|
|
||||||
await hook.after_execute_tool(context, tool_call, tool, params, result)
|
await hook.after_execute_tool(context, tool_call, tool, params, result)
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
from collections import OrderedDict, deque
|
from collections import deque
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Protocol
|
from typing import Any, Protocol
|
||||||
@@ -127,7 +127,7 @@ class SendSessionMessageTool(Tool):
|
|||||||
self._max_messages_per_minute = max_messages_per_minute
|
self._max_messages_per_minute = max_messages_per_minute
|
||||||
self._schedule_later = schedule_later
|
self._schedule_later = schedule_later
|
||||||
self._clock = clock or time.monotonic
|
self._clock = clock or time.monotonic
|
||||||
self._sent_at: OrderedDict[str, deque[float]] = OrderedDict()
|
self._sent_at: dict[str, deque[float]] = {}
|
||||||
self._pending_replies: dict[tuple[str, str], _PendingReply] = {}
|
self._pending_replies: dict[tuple[str, str], _PendingReply] = {}
|
||||||
self._expiry_tasks: set[asyncio.Task[None]] = set()
|
self._expiry_tasks: set[asyncio.Task[None]] = set()
|
||||||
self._send_lock = asyncio.Lock()
|
self._send_lock = asyncio.Lock()
|
||||||
@@ -240,11 +240,8 @@ class SendSessionMessageTool(Tool):
|
|||||||
|
|
||||||
async with self._send_lock:
|
async with self._send_lock:
|
||||||
now = self._clock()
|
now = self._clock()
|
||||||
|
sent_at = self._sent_at.setdefault(source.session_key, deque())
|
||||||
cutoff = now - _RATE_LIMIT_WINDOW_SECONDS
|
cutoff = now - _RATE_LIMIT_WINDOW_SECONDS
|
||||||
self._prune_expired_rate_limits(cutoff)
|
|
||||||
sent_at = self._sent_at.get(source.session_key)
|
|
||||||
if sent_at is None:
|
|
||||||
sent_at = deque[float]()
|
|
||||||
while sent_at and sent_at[0] <= cutoff:
|
while sent_at and sent_at[0] <= cutoff:
|
||||||
sent_at.popleft()
|
sent_at.popleft()
|
||||||
if len(sent_at) >= self._max_messages_per_minute:
|
if len(sent_at) >= self._max_messages_per_minute:
|
||||||
@@ -262,8 +259,6 @@ class SendSessionMessageTool(Tool):
|
|||||||
input_role="user",
|
input_role="user",
|
||||||
))
|
))
|
||||||
sent_at.append(now)
|
sent_at.append(now)
|
||||||
self._sent_at[source.session_key] = sent_at
|
|
||||||
self._sent_at.move_to_end(source.session_key)
|
|
||||||
self._cancel_pending_reply(reverse_wait_key)
|
self._cancel_pending_reply(reverse_wait_key)
|
||||||
if timeout_seconds is not None:
|
if timeout_seconds is not None:
|
||||||
self._cancel_pending_reply(wait_key)
|
self._cancel_pending_reply(wait_key)
|
||||||
@@ -276,14 +271,6 @@ class SendSessionMessageTool(Tool):
|
|||||||
|
|
||||||
return f"@{target.name}"
|
return f"@{target.name}"
|
||||||
|
|
||||||
def _prune_expired_rate_limits(self, cutoff: float) -> None:
|
|
||||||
"""Drop sources ordered by their most recent successful send."""
|
|
||||||
while self._sent_at:
|
|
||||||
_, sent_at = next(iter(self._sent_at.items()))
|
|
||||||
if sent_at[-1] > cutoff:
|
|
||||||
return
|
|
||||||
self._sent_at.popitem(last=False)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _validate_reply_timeout(
|
def _validate_reply_timeout(
|
||||||
expect_reply: bool,
|
expect_reply: bool,
|
||||||
|
|||||||
@@ -182,12 +182,6 @@ class NanobotDingTalkHandler(_CallbackHandlerBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.channel._accepting_inbound_tasks:
|
|
||||||
self.channel.logger.debug(
|
|
||||||
"Skipping DingTalk inbound dispatch during channel shutdown"
|
|
||||||
)
|
|
||||||
return AckMessage.STATUS_OK, "OK"
|
|
||||||
|
|
||||||
self.channel.logger.info("Received message from {} ({}): {}", sender_name, sender_id, content)
|
self.channel.logger.info("Received message from {} ({}): {}", sender_name, sender_id, content)
|
||||||
|
|
||||||
# Forward to Nanobot via _on_message (non-blocking).
|
# Forward to Nanobot via _on_message (non-blocking).
|
||||||
@@ -202,7 +196,7 @@ class NanobotDingTalkHandler(_CallbackHandlerBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.channel._background_tasks.add(task)
|
self.channel._background_tasks.add(task)
|
||||||
task.add_done_callback(self.channel._on_background_task_done)
|
task.add_done_callback(self.channel._background_tasks.discard)
|
||||||
|
|
||||||
return AckMessage.STATUS_OK, "OK"
|
return AckMessage.STATUS_OK, "OK"
|
||||||
|
|
||||||
@@ -262,17 +256,6 @@ class DingTalkChannel(BaseChannel):
|
|||||||
|
|
||||||
# Hold references to background tasks to prevent GC
|
# Hold references to background tasks to prevent GC
|
||||||
self._background_tasks: set[asyncio.Task[None]] = set()
|
self._background_tasks: set[asyncio.Task[None]] = set()
|
||||||
self._accepting_inbound_tasks = True
|
|
||||||
|
|
||||||
def _on_background_task_done(self, task: asyncio.Task[None]) -> None:
|
|
||||||
self._background_tasks.discard(task)
|
|
||||||
if task.cancelled():
|
|
||||||
return
|
|
||||||
exception = task.exception()
|
|
||||||
if exception is not None:
|
|
||||||
self.logger.opt(exception=exception).error(
|
|
||||||
"DingTalk inbound message task failed"
|
|
||||||
)
|
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start the DingTalk bot with Stream Mode."""
|
"""Start the DingTalk bot with Stream Mode."""
|
||||||
@@ -289,7 +272,6 @@ class DingTalkChannel(BaseChannel):
|
|||||||
self.logger.error("client_id and client_secret not configured")
|
self.logger.error("client_id and client_secret not configured")
|
||||||
return
|
return
|
||||||
|
|
||||||
self._accepting_inbound_tasks = True
|
|
||||||
self._running = True
|
self._running = True
|
||||||
self._http = httpx.AsyncClient(
|
self._http = httpx.AsyncClient(
|
||||||
timeout=httpx.Timeout(10.0, connect=10.0, read=30.0, write=30.0, pool=10.0)
|
timeout=httpx.Timeout(10.0, connect=10.0, read=30.0, write=30.0, pool=10.0)
|
||||||
@@ -327,7 +309,6 @@ class DingTalkChannel(BaseChannel):
|
|||||||
|
|
||||||
async def stop(self) -> None:
|
async def stop(self) -> None:
|
||||||
"""Stop the DingTalk bot."""
|
"""Stop the DingTalk bot."""
|
||||||
self._accepting_inbound_tasks = False
|
|
||||||
self._running = False
|
self._running = False
|
||||||
await self._close_stream_client()
|
await self._close_stream_client()
|
||||||
start_task = self._start_task
|
start_task = self._start_task
|
||||||
@@ -345,11 +326,8 @@ class DingTalkChannel(BaseChannel):
|
|||||||
await self._http.aclose()
|
await self._http.aclose()
|
||||||
self._http = None
|
self._http = None
|
||||||
# Cancel outstanding background tasks
|
# Cancel outstanding background tasks
|
||||||
background_tasks = tuple(self._background_tasks)
|
for task in self._background_tasks:
|
||||||
for task in background_tasks:
|
|
||||||
task.cancel()
|
task.cancel()
|
||||||
if background_tasks:
|
|
||||||
await asyncio.gather(*background_tasks, return_exceptions=True)
|
|
||||||
self._background_tasks.clear()
|
self._background_tasks.clear()
|
||||||
|
|
||||||
async def _close_stream_client(self) -> None:
|
async def _close_stream_client(self) -> None:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import json
|
|||||||
import zipfile
|
import zipfile
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
@@ -402,61 +402,6 @@ async def test_handler_uses_voice_recognition_text_when_text_is_empty(monkeypatc
|
|||||||
assert msg.chat_id == "group:conv123"
|
assert msg.chat_id == "group:conv123"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_handler_retrieves_background_message_failure(monkeypatch) -> None:
|
|
||||||
bus = MessageBus()
|
|
||||||
channel = DingTalkChannel(
|
|
||||||
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["user1"]),
|
|
||||||
bus,
|
|
||||||
)
|
|
||||||
handler = NanobotDingTalkHandler(channel)
|
|
||||||
failure = RuntimeError("inbound dispatch failed")
|
|
||||||
mock_logger = MagicMock()
|
|
||||||
channel.logger = mock_logger
|
|
||||||
|
|
||||||
class _FakeChatbotMessage:
|
|
||||||
text = SimpleNamespace(content="hello")
|
|
||||||
extensions = {}
|
|
||||||
sender_staff_id = "user1"
|
|
||||||
sender_id = "fallback-user"
|
|
||||||
sender_nick = "Alice"
|
|
||||||
message_type = "text"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def from_dict(_data):
|
|
||||||
return _FakeChatbotMessage()
|
|
||||||
|
|
||||||
async def fail(*_args) -> None:
|
|
||||||
raise failure
|
|
||||||
|
|
||||||
monkeypatch.setattr(dingtalk_module, "ChatbotMessage", _FakeChatbotMessage)
|
|
||||||
monkeypatch.setattr(dingtalk_module, "AckMessage", SimpleNamespace(STATUS_OK="OK"))
|
|
||||||
monkeypatch.setattr(channel, "_on_message", fail)
|
|
||||||
event_loop = asyncio.get_running_loop()
|
|
||||||
previous_handler = event_loop.get_exception_handler()
|
|
||||||
loop_errors: list[dict[str, object]] = []
|
|
||||||
event_loop.set_exception_handler(lambda _loop, context: loop_errors.append(context))
|
|
||||||
|
|
||||||
try:
|
|
||||||
status, body = await handler.process(
|
|
||||||
SimpleNamespace(data={"conversationType": "1", "text": {"content": "hello"}})
|
|
||||||
)
|
|
||||||
for _ in range(10):
|
|
||||||
await asyncio.sleep(0)
|
|
||||||
if not channel._background_tasks:
|
|
||||||
break
|
|
||||||
finally:
|
|
||||||
event_loop.set_exception_handler(previous_handler)
|
|
||||||
|
|
||||||
assert (status, body) == ("OK", "OK")
|
|
||||||
assert not channel._background_tasks
|
|
||||||
assert not loop_errors
|
|
||||||
mock_logger.opt.assert_called_once_with(exception=failure)
|
|
||||||
mock_logger.opt.return_value.error.assert_called_once_with(
|
|
||||||
"DingTalk inbound message task failed"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_handler_processes_file_message(monkeypatch) -> None:
|
async def test_handler_processes_file_message(monkeypatch) -> None:
|
||||||
"""Test that file messages are handled and forwarded with downloaded path."""
|
"""Test that file messages are handled and forwarded with downloaded path."""
|
||||||
@@ -506,72 +451,6 @@ async def test_handler_processes_file_message(monkeypatch) -> None:
|
|||||||
assert "/tmp/nanobot_dingtalk/user1/report.xlsx" in msg.content
|
assert "/tmp/nanobot_dingtalk/user1/report.xlsx" in msg.content
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_handler_does_not_spawn_message_task_after_stop_during_download(
|
|
||||||
monkeypatch,
|
|
||||||
) -> None:
|
|
||||||
channel = DingTalkChannel(
|
|
||||||
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["user1"]),
|
|
||||||
MessageBus(),
|
|
||||||
)
|
|
||||||
handler = NanobotDingTalkHandler(channel)
|
|
||||||
download_started = asyncio.Event()
|
|
||||||
release_download = asyncio.Event()
|
|
||||||
message_task_started = asyncio.Event()
|
|
||||||
|
|
||||||
class _FakeFileChatbotMessage:
|
|
||||||
text = None
|
|
||||||
extensions = {}
|
|
||||||
image_content = None
|
|
||||||
rich_text_content = None
|
|
||||||
sender_staff_id = "user1"
|
|
||||||
sender_id = "fallback-user"
|
|
||||||
sender_nick = "Alice"
|
|
||||||
message_type = "file"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def from_dict(_data):
|
|
||||||
return _FakeFileChatbotMessage()
|
|
||||||
|
|
||||||
async def delayed_download(*_args):
|
|
||||||
download_started.set()
|
|
||||||
await release_download.wait()
|
|
||||||
return "/tmp/nanobot_dingtalk/user1/report.xlsx"
|
|
||||||
|
|
||||||
async def block_message(*_args) -> None:
|
|
||||||
message_task_started.set()
|
|
||||||
await asyncio.Future()
|
|
||||||
|
|
||||||
monkeypatch.setattr(dingtalk_module, "ChatbotMessage", _FakeFileChatbotMessage)
|
|
||||||
monkeypatch.setattr(dingtalk_module, "AckMessage", SimpleNamespace(STATUS_OK="OK"))
|
|
||||||
monkeypatch.setattr(channel, "_download_dingtalk_file", delayed_download)
|
|
||||||
monkeypatch.setattr(channel, "_on_message", block_message)
|
|
||||||
|
|
||||||
process_task = asyncio.create_task(handler.process(SimpleNamespace(data={
|
|
||||||
"conversationType": "1",
|
|
||||||
"content": {"downloadCode": "abc123", "fileName": "report.xlsx"},
|
|
||||||
"text": {"content": ""},
|
|
||||||
})))
|
|
||||||
await download_started.wait()
|
|
||||||
|
|
||||||
try:
|
|
||||||
await channel.stop()
|
|
||||||
release_download.set()
|
|
||||||
assert await process_task == ("OK", "OK")
|
|
||||||
await asyncio.sleep(0)
|
|
||||||
|
|
||||||
assert not message_task_started.is_set()
|
|
||||||
assert not channel._background_tasks
|
|
||||||
finally:
|
|
||||||
release_download.set()
|
|
||||||
if not process_task.done():
|
|
||||||
process_task.cancel()
|
|
||||||
pending = tuple(channel._background_tasks)
|
|
||||||
for task in pending:
|
|
||||||
task.cancel()
|
|
||||||
await asyncio.gather(process_task, *pending, return_exceptions=True)
|
|
||||||
|
|
||||||
|
|
||||||
def _rich_text_message(rich_text_list):
|
def _rich_text_message(rich_text_list):
|
||||||
class _FakeRichTextChatbotMessage:
|
class _FakeRichTextChatbotMessage:
|
||||||
text = None
|
text = None
|
||||||
@@ -771,41 +650,6 @@ async def test_stop_cancels_stream_client_after_sdk_swallows_first_cancel(monkey
|
|||||||
assert start_task.cancelled()
|
assert start_task.cancelled()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_stop_waits_for_background_message_tasks() -> None:
|
|
||||||
channel = DingTalkChannel(
|
|
||||||
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["*"]),
|
|
||||||
MessageBus(),
|
|
||||||
)
|
|
||||||
mock_logger = MagicMock()
|
|
||||||
channel.logger = mock_logger
|
|
||||||
started = asyncio.Event()
|
|
||||||
cancelled = asyncio.Event()
|
|
||||||
|
|
||||||
async def wait_forever() -> None:
|
|
||||||
started.set()
|
|
||||||
try:
|
|
||||||
await asyncio.Future()
|
|
||||||
finally:
|
|
||||||
cancelled.set()
|
|
||||||
|
|
||||||
task = asyncio.create_task(wait_forever())
|
|
||||||
channel._background_tasks.add(task)
|
|
||||||
task.add_done_callback(channel._on_background_task_done)
|
|
||||||
await started.wait()
|
|
||||||
|
|
||||||
try:
|
|
||||||
await channel.stop()
|
|
||||||
assert task.done()
|
|
||||||
assert cancelled.is_set()
|
|
||||||
assert not channel._background_tasks
|
|
||||||
mock_logger.opt.assert_not_called()
|
|
||||||
finally:
|
|
||||||
if not task.done():
|
|
||||||
task.cancel()
|
|
||||||
await asyncio.gather(task, return_exceptions=True)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_download_dingtalk_file(tmp_path, monkeypatch) -> None:
|
async def test_download_dingtalk_file(tmp_path, monkeypatch) -> None:
|
||||||
"""Test the two-step file download flow (get URL then download content)."""
|
"""Test the two-step file download flow (get URL then download content)."""
|
||||||
|
|||||||
@@ -430,13 +430,7 @@ class EmailChannel(BaseChannel):
|
|||||||
skipped_uids: set[str],
|
skipped_uids: set[str],
|
||||||
cycle_uids: set[str],
|
cycle_uids: set[str],
|
||||||
) -> list[dict[str, Any]] | None:
|
) -> list[dict[str, Any]] | None:
|
||||||
"""Fetch messages by arbitrary IMAP search criteria.
|
"""Fetch messages by arbitrary IMAP search criteria."""
|
||||||
|
|
||||||
Uses UID SEARCH so already-processed UIDs are recognized before any
|
|
||||||
FETCH at all, then fetches headers only to evaluate every filter — the
|
|
||||||
full body (and any attachments) is downloaded only for messages that
|
|
||||||
pass every check and are actually going to be delivered.
|
|
||||||
"""
|
|
||||||
mailbox = self.config.imap_mailbox or "INBOX"
|
mailbox = self.config.imap_mailbox or "INBOX"
|
||||||
|
|
||||||
client = self._open_imap_client(mailbox=mailbox, missing_mailbox_ok=True)
|
client = self._open_imap_client(mailbox=mailbox, missing_mailbox_ok=True)
|
||||||
@@ -444,30 +438,29 @@ class EmailChannel(BaseChannel):
|
|||||||
return messages
|
return messages
|
||||||
|
|
||||||
try:
|
try:
|
||||||
status, data = client.uid("SEARCH", None, *search_criteria)
|
status, data = client.search(None, *search_criteria)
|
||||||
if status != "OK" or not data or not data[0]:
|
if status != "OK" or not data:
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
uids = [raw.decode("ascii", errors="ignore") for raw in data[0].split()]
|
ids = data[0].split()
|
||||||
if limit > 0 and len(uids) > limit:
|
if limit > 0 and len(ids) > limit:
|
||||||
uids = uids[-limit:]
|
ids = ids[-limit:]
|
||||||
|
for imap_id in ids:
|
||||||
features: _ServerFeatures | None = None
|
status, fetched = client.fetch(imap_id, "(BODY.PEEK[] UID)")
|
||||||
|
|
||||||
for uid in uids:
|
|
||||||
if not uid or uid in cycle_uids:
|
|
||||||
continue
|
|
||||||
if dedupe and uid in self._processed_uids:
|
|
||||||
continue
|
|
||||||
|
|
||||||
status, fetched = client.uid("FETCH", uid, "(BODY.PEEK[HEADER])")
|
|
||||||
if status != "OK" or not fetched:
|
if status != "OK" or not fetched:
|
||||||
continue
|
continue
|
||||||
header_bytes = self._extract_message_bytes(fetched)
|
|
||||||
if header_bytes is None:
|
raw_bytes = self._extract_message_bytes(fetched)
|
||||||
|
if raw_bytes is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
parsed = BytesParser(policy=policy.default).parsebytes(header_bytes)
|
uid = self._extract_uid(fetched)
|
||||||
|
if uid and uid in cycle_uids:
|
||||||
|
continue
|
||||||
|
if dedupe and uid and uid in self._processed_uids:
|
||||||
|
continue
|
||||||
|
|
||||||
|
parsed = BytesParser(policy=policy.default).parsebytes(raw_bytes)
|
||||||
sender = parseaddr(parsed.get("From", ""))[1].strip().lower()
|
sender = parseaddr(parsed.get("From", ""))[1].strip().lower()
|
||||||
if not sender:
|
if not sender:
|
||||||
continue
|
continue
|
||||||
@@ -475,8 +468,9 @@ class EmailChannel(BaseChannel):
|
|||||||
self.logger.info("From {} ignored: matches bot-owned address", sender)
|
self.logger.info("From {} ignored: matches bot-owned address", sender)
|
||||||
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
||||||
if mark_seen:
|
if mark_seen:
|
||||||
features = self._mark_seen_uid(client, uid, features)
|
client.store(imap_id, "+FLAGS", "\\Seen")
|
||||||
skipped_uids.add(uid)
|
if uid:
|
||||||
|
skipped_uids.add(uid)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# --- Anti-spoofing: verify Authentication-Results ---
|
# --- Anti-spoofing: verify Authentication-Results ---
|
||||||
@@ -488,7 +482,8 @@ class EmailChannel(BaseChannel):
|
|||||||
sender,
|
sender,
|
||||||
)
|
)
|
||||||
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
||||||
skipped_uids.add(uid)
|
if uid:
|
||||||
|
skipped_uids.add(uid)
|
||||||
continue
|
continue
|
||||||
if self.config.verify_dkim and not dkim_pass:
|
if self.config.verify_dkim and not dkim_pass:
|
||||||
self.logger.warning(
|
self.logger.warning(
|
||||||
@@ -497,26 +492,18 @@ class EmailChannel(BaseChannel):
|
|||||||
sender,
|
sender,
|
||||||
)
|
)
|
||||||
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
||||||
skipped_uids.add(uid)
|
if uid:
|
||||||
|
skipped_uids.add(uid)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not self.is_allowed(sender):
|
if not self.is_allowed(sender):
|
||||||
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
||||||
if mark_seen:
|
if mark_seen:
|
||||||
features = self._mark_seen_uid(client, uid, features)
|
client.store(imap_id, "+FLAGS", "\\Seen")
|
||||||
skipped_uids.add(uid)
|
if uid:
|
||||||
|
skipped_uids.add(uid)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Passed every filter — only now fetch the full message body
|
|
||||||
# (and any attachments) for the message we're actually delivering.
|
|
||||||
status, full_fetched = client.uid("FETCH", uid, "(BODY.PEEK[])")
|
|
||||||
if status != "OK" or not full_fetched:
|
|
||||||
continue
|
|
||||||
raw_bytes = self._extract_message_bytes(full_fetched)
|
|
||||||
if raw_bytes is None:
|
|
||||||
continue
|
|
||||||
parsed = BytesParser(policy=policy.default).parsebytes(raw_bytes)
|
|
||||||
|
|
||||||
subject = self._decode_header_value(parsed.get("Subject", ""))
|
subject = self._decode_header_value(parsed.get("Subject", ""))
|
||||||
date_value = parsed.get("Date", "")
|
date_value = parsed.get("Date", "")
|
||||||
message_id = parsed.get("Message-ID", "").strip()
|
message_id = parsed.get("Message-ID", "").strip()
|
||||||
@@ -569,19 +556,10 @@ class EmailChannel(BaseChannel):
|
|||||||
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
self._remember_processed_uid(uid, dedupe, cycle_uids)
|
||||||
|
|
||||||
if mark_seen:
|
if mark_seen:
|
||||||
features = self._mark_seen_uid(client, uid, features)
|
client.store(imap_id, "+FLAGS", "\\Seen")
|
||||||
finally:
|
finally:
|
||||||
self._close_imap_client(client)
|
self._close_imap_client(client)
|
||||||
|
|
||||||
def _mark_seen_uid(
|
|
||||||
self, client: Any, uid: str, features: _ServerFeatures | None
|
|
||||||
) -> _ServerFeatures:
|
|
||||||
"""Mark a single UID \\Seen, reusing session-learned STORE support."""
|
|
||||||
if features is None:
|
|
||||||
features = self._server_features(client)
|
|
||||||
self._uid_store_flag(client, uid, "\\Seen", features)
|
|
||||||
return features
|
|
||||||
|
|
||||||
def _open_imap_client(self, mailbox: str, *, missing_mailbox_ok: bool = False) -> Any | None:
|
def _open_imap_client(self, mailbox: str, *, missing_mailbox_ok: bool = False) -> Any | None:
|
||||||
if self.config.imap_use_ssl:
|
if self.config.imap_use_ssl:
|
||||||
client: Any = imaplib.IMAP4_SSL(self.config.imap_host, self.config.imap_port)
|
client: Any = imaplib.IMAP4_SSL(self.config.imap_host, self.config.imap_port)
|
||||||
@@ -736,14 +714,11 @@ class EmailChannel(BaseChannel):
|
|||||||
return data[0].split()[0]
|
return data[0].split()[0]
|
||||||
|
|
||||||
def _uid_store_deleted(self, client: Any, uid: str, features: _ServerFeatures) -> bool:
|
def _uid_store_deleted(self, client: Any, uid: str, features: _ServerFeatures) -> bool:
|
||||||
return self._uid_store_flag(client, uid, "\\Deleted", features)
|
|
||||||
|
|
||||||
def _uid_store_flag(self, client: Any, uid: str, flag: str, features: _ServerFeatures) -> bool:
|
|
||||||
# Optimistic path: try UID STORE first because UID is stable and avoids
|
# Optimistic path: try UID STORE first because UID is stable and avoids
|
||||||
# sequence-number lookup. If this fails once for the session, remember it
|
# sequence-number lookup. If this fails once for the session, remember it
|
||||||
# and use the sequence STORE fallback directly for remaining UIDs.
|
# and use the sequence STORE fallback directly for remaining UIDs.
|
||||||
if features.uid_store is not False:
|
if features.uid_store is not False:
|
||||||
status, _ = client.uid("STORE", uid, "+FLAGS", f"({flag})")
|
status, _ = client.uid("STORE", uid, "+FLAGS", "(\\Deleted)")
|
||||||
if status == "OK":
|
if status == "OK":
|
||||||
features.uid_store = True
|
features.uid_store = True
|
||||||
return True
|
return True
|
||||||
@@ -753,12 +728,12 @@ class EmailChannel(BaseChannel):
|
|||||||
# unreliable: resolve the current sequence number from UID and use STORE.
|
# unreliable: resolve the current sequence number from UID and use STORE.
|
||||||
imap_id = self._lookup_imap_id_by_uid(client, uid)
|
imap_id = self._lookup_imap_id_by_uid(client, uid)
|
||||||
if not imap_id:
|
if not imap_id:
|
||||||
self.logger.warning("Could not locate UID {} to set flag {}", uid, flag)
|
self.logger.warning("Post-action skipped: UID {} not found", uid)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
status, _ = client.store(imap_id, "+FLAGS", flag)
|
status, _ = client.store(imap_id, "+FLAGS", "\\Deleted")
|
||||||
if status != "OK":
|
if status != "OK":
|
||||||
self.logger.warning("Failed to set flag {} on UID {}", flag, uid)
|
self.logger.warning("Post-action failed: could not mark UID {} as deleted", uid)
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -798,6 +773,16 @@ class EmailChannel(BaseChannel):
|
|||||||
return bytes(fetched_item[1])
|
return bytes(fetched_item[1])
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_uid(fetched: list[Any]) -> str:
|
||||||
|
for item in fetched:
|
||||||
|
if isinstance(item, tuple) and item and isinstance(item[0], (bytes, bytearray)):
|
||||||
|
head = bytes(item[0]).decode("utf-8", errors="ignore")
|
||||||
|
m = re.search(r"UID\s+(\d+)", head)
|
||||||
|
if m:
|
||||||
|
return m.group(1)
|
||||||
|
return ""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _decode_header_value(value: str) -> str:
|
def _decode_header_value(value: str) -> str:
|
||||||
if not value:
|
if not value:
|
||||||
|
|||||||
@@ -53,7 +53,30 @@ def _make_raw_email(
|
|||||||
def test_fetch_new_messages_parses_unseen_and_marks_seen(monkeypatch) -> None:
|
def test_fetch_new_messages_parses_unseen_and_marks_seen(monkeypatch) -> None:
|
||||||
raw = _make_raw_email(subject="Invoice", body="Please pay")
|
raw = _make_raw_email(subject="Invoice", body="Please pay")
|
||||||
|
|
||||||
fake = _make_fake_imap(raw, uid=b"123")
|
class FakeIMAP:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.store_calls: list[tuple[bytes, str, str]] = []
|
||||||
|
|
||||||
|
def login(self, _user: str, _pw: str):
|
||||||
|
return "OK", [b"logged in"]
|
||||||
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
|
self.store_calls.append((imap_id, op, flags))
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
fake = FakeIMAP()
|
||||||
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
channel = EmailChannel(_make_config(), MessageBus())
|
channel = EmailChannel(_make_config(), MessageBus())
|
||||||
@@ -63,25 +86,38 @@ def test_fetch_new_messages_parses_unseen_and_marks_seen(monkeypatch) -> None:
|
|||||||
assert items[0]["sender"] == "alice@example.com"
|
assert items[0]["sender"] == "alice@example.com"
|
||||||
assert items[0]["subject"] == "Invoice"
|
assert items[0]["subject"] == "Invoice"
|
||||||
assert "Please pay" in items[0]["content"]
|
assert "Please pay" in items[0]["content"]
|
||||||
assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
|
assert fake.store_calls == [(b"1", "+FLAGS", "\\Seen")]
|
||||||
assert [call for call in fake.uid_calls if call[0] == "FETCH"] == [
|
|
||||||
("FETCH", "123", "(BODY.PEEK[HEADER])"),
|
|
||||||
("FETCH", "123", "(BODY.PEEK[])"),
|
|
||||||
]
|
|
||||||
assert skipped_uids == set()
|
assert skipped_uids == set()
|
||||||
|
|
||||||
# Same UID should be deduped in-process.
|
# Same UID should be deduped in-process.
|
||||||
items_again, skipped_again = channel._fetch_new_messages()
|
items_again, skipped_again = channel._fetch_new_messages()
|
||||||
assert items_again == []
|
assert items_again == []
|
||||||
assert skipped_again == set()
|
assert skipped_again == set()
|
||||||
assert len([call for call in fake.uid_calls if call[0] == "FETCH"]) == 2
|
|
||||||
|
|
||||||
|
|
||||||
def test_fetch_new_messages_returns_accepted_and_skipped_uids(monkeypatch) -> None:
|
def test_fetch_new_messages_returns_accepted_and_skipped_uids(monkeypatch) -> None:
|
||||||
raw = _make_raw_email(subject="Invoice", body="Please pay")
|
raw = _make_raw_email(subject="Invoice", body="Please pay")
|
||||||
|
|
||||||
fake = _make_fake_imap(raw, uid=b"123")
|
class FakeIMAP:
|
||||||
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
def login(self, _user: str, _pw: str):
|
||||||
|
return "OK", [b"logged in"]
|
||||||
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, _imap_id: bytes, _op: str, _flags: str):
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: FakeIMAP())
|
||||||
|
|
||||||
channel = EmailChannel(_make_config(post_action="delete"), MessageBus())
|
channel = EmailChannel(_make_config(post_action="delete"), MessageBus())
|
||||||
items, skipped_uids = channel._fetch_new_messages()
|
items, skipped_uids = channel._fetch_new_messages()
|
||||||
@@ -94,10 +130,26 @@ def test_fetch_new_messages_returns_accepted_and_skipped_uids(monkeypatch) -> No
|
|||||||
def test_fetch_new_messages_rejected_returns_skipped_uid(monkeypatch) -> None:
|
def test_fetch_new_messages_rejected_returns_skipped_uid(monkeypatch) -> None:
|
||||||
raw = _make_raw_email(from_addr="Nanobot <bot@example.com>", subject="Loop test")
|
raw = _make_raw_email(from_addr="Nanobot <bot@example.com>", subject="Loop test")
|
||||||
|
|
||||||
monkeypatch.setattr(
|
class FakeIMAP:
|
||||||
"nanobot.channels.email.runtime.imaplib.IMAP4_SSL",
|
def login(self, _user: str, _pw: str):
|
||||||
lambda _h, _p: _make_fake_imap(raw, uid=b"123"),
|
return "OK", [b"logged in"]
|
||||||
)
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, _imap_id: bytes, _op: str, _flags: str):
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: FakeIMAP())
|
||||||
|
|
||||||
channel_skip = EmailChannel(
|
channel_skip = EmailChannel(
|
||||||
_make_config(from_address="bot@example.com", post_action="delete", post_action_ignore_skipped=True),
|
_make_config(from_address="bot@example.com", post_action="delete", post_action_ignore_skipped=True),
|
||||||
@@ -493,7 +545,30 @@ async def test_start_keeps_post_actions_for_successful_emails_when_later_deliver
|
|||||||
def test_fetch_new_messages_skips_self_sent_email_and_marks_seen(monkeypatch) -> None:
|
def test_fetch_new_messages_skips_self_sent_email_and_marks_seen(monkeypatch) -> None:
|
||||||
raw = _make_raw_email(from_addr="Nanobot <bot@example.com>", subject="Loop test")
|
raw = _make_raw_email(from_addr="Nanobot <bot@example.com>", subject="Loop test")
|
||||||
|
|
||||||
fake = _make_fake_imap(raw, uid=b"123")
|
class FakeIMAP:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.store_calls: list[tuple[bytes, str, str]] = []
|
||||||
|
|
||||||
|
def login(self, _user: str, _pw: str):
|
||||||
|
return "OK", [b"logged in"]
|
||||||
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
|
self.store_calls.append((imap_id, op, flags))
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
fake = FakeIMAP()
|
||||||
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
channel = EmailChannel(_make_config(from_address="bot@example.com"), MessageBus())
|
channel = EmailChannel(_make_config(from_address="bot@example.com"), MessageBus())
|
||||||
@@ -501,7 +576,7 @@ def test_fetch_new_messages_skips_self_sent_email_and_marks_seen(monkeypatch) ->
|
|||||||
|
|
||||||
assert items == []
|
assert items == []
|
||||||
assert skipped_uids == {"123"}
|
assert skipped_uids == {"123"}
|
||||||
assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
|
assert fake.store_calls == [(b"1", "+FLAGS", "\\Seen")]
|
||||||
|
|
||||||
# Same UID should still be deduped after being ignored.
|
# Same UID should still be deduped after being ignored.
|
||||||
items_again, skipped_again = channel._fetch_new_messages()
|
items_again, skipped_again = channel._fetch_new_messages()
|
||||||
@@ -539,14 +614,37 @@ def test_fetch_new_messages_skips_self_sent_across_identity_sources(
|
|||||||
imap_username matches, and must be case-insensitive."""
|
imap_username matches, and must be case-insensitive."""
|
||||||
raw = _make_raw_email(from_addr=from_header, subject="Loop test")
|
raw = _make_raw_email(from_addr=from_header, subject="Loop test")
|
||||||
|
|
||||||
fake = _make_fake_imap(raw, uid=b"123")
|
class FakeIMAP:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.store_calls: list[tuple[bytes, str, str]] = []
|
||||||
|
|
||||||
|
def login(self, _user: str, _pw: str):
|
||||||
|
return "OK", [b"logged in"]
|
||||||
|
|
||||||
|
def select(self, _mailbox: str):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def search(self, *_args):
|
||||||
|
return "OK", [b"1"]
|
||||||
|
|
||||||
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
|
|
||||||
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
|
self.store_calls.append((imap_id, op, flags))
|
||||||
|
return "OK", [b""]
|
||||||
|
|
||||||
|
def logout(self):
|
||||||
|
return "BYE", [b""]
|
||||||
|
|
||||||
|
fake = FakeIMAP()
|
||||||
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
|
||||||
|
|
||||||
channel = EmailChannel(_make_config(**config_override), MessageBus())
|
channel = EmailChannel(_make_config(**config_override), MessageBus())
|
||||||
items, _ = channel._fetch_new_messages()
|
items, _ = channel._fetch_new_messages()
|
||||||
|
|
||||||
assert items == []
|
assert items == []
|
||||||
assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
|
assert fake.store_calls == [(b"1", "+FLAGS", "\\Seen")]
|
||||||
|
|
||||||
|
|
||||||
def test_fetch_new_messages_retries_once_when_imap_connection_goes_stale(monkeypatch) -> None:
|
def test_fetch_new_messages_retries_once_when_imap_connection_goes_stale(monkeypatch) -> None:
|
||||||
@@ -564,16 +662,15 @@ def test_fetch_new_messages_retries_once_when_imap_connection_goes_stale(monkeyp
|
|||||||
def select(self, _mailbox: str):
|
def select(self, _mailbox: str):
|
||||||
return "OK", [b"1"]
|
return "OK", [b"1"]
|
||||||
|
|
||||||
def uid(self, command: str, *args):
|
def search(self, *_args):
|
||||||
if command == "SEARCH":
|
self.search_calls += 1
|
||||||
self.search_calls += 1
|
if fail_once["pending"]:
|
||||||
if fail_once["pending"]:
|
fail_once["pending"] = False
|
||||||
fail_once["pending"] = False
|
raise imaplib.IMAP4.abort("socket error")
|
||||||
raise imaplib.IMAP4.abort("socket error")
|
return "OK", [b"1"]
|
||||||
return "OK", [b"123"]
|
|
||||||
if command == "FETCH":
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
return "OK", [(b"1 (UID 123 BODY[] {200})", raw), b")"]
|
||||||
return "OK", [b""]
|
|
||||||
|
|
||||||
def store(self, imap_id: bytes, op: str, flags: str):
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
self.store_calls.append((imap_id, op, flags))
|
self.store_calls.append((imap_id, op, flags))
|
||||||
@@ -603,7 +700,10 @@ def test_fetch_new_messages_retries_once_when_imap_connection_goes_stale(monkeyp
|
|||||||
def test_fetch_new_messages_keeps_messages_collected_before_stale_retry(monkeypatch) -> None:
|
def test_fetch_new_messages_keeps_messages_collected_before_stale_retry(monkeypatch) -> None:
|
||||||
raw_first = _make_raw_email(subject="First", body="First body")
|
raw_first = _make_raw_email(subject="First", body="First body")
|
||||||
raw_second = _make_raw_email(subject="Second", body="Second body")
|
raw_second = _make_raw_email(subject="Second", body="Second body")
|
||||||
mailbox_state = {"123": raw_first, "124": raw_second}
|
mailbox_state = {
|
||||||
|
b"1": {"uid": b"123", "raw": raw_first, "seen": False},
|
||||||
|
b"2": {"uid": b"124", "raw": raw_second, "seen": False},
|
||||||
|
}
|
||||||
fail_once = {"pending": True}
|
fail_once = {"pending": True}
|
||||||
|
|
||||||
class FlakyIMAP:
|
class FlakyIMAP:
|
||||||
@@ -613,18 +713,20 @@ def test_fetch_new_messages_keeps_messages_collected_before_stale_retry(monkeypa
|
|||||||
def select(self, _mailbox: str):
|
def select(self, _mailbox: str):
|
||||||
return "OK", [b"2"]
|
return "OK", [b"2"]
|
||||||
|
|
||||||
def uid(self, command: str, *args):
|
def search(self, *_args):
|
||||||
if command == "SEARCH":
|
unseen_ids = [imap_id for imap_id, item in mailbox_state.items() if not item["seen"]]
|
||||||
keys = " ".join(sorted(mailbox_state.keys(), key=int))
|
return "OK", [b" ".join(unseen_ids)]
|
||||||
return "OK", [keys.encode()]
|
|
||||||
if command == "FETCH":
|
def fetch(self, imap_id: bytes, _parts: str):
|
||||||
uid = args[0]
|
if imap_id == b"2" and fail_once["pending"]:
|
||||||
if uid == "124" and fail_once["pending"]:
|
fail_once["pending"] = False
|
||||||
fail_once["pending"] = False
|
raise imaplib.IMAP4.abort("socket error")
|
||||||
raise imaplib.IMAP4.abort("socket error")
|
item = mailbox_state[imap_id]
|
||||||
raw = mailbox_state[uid]
|
header = b"%s (UID %s BODY[] {200})" % (imap_id, item["uid"])
|
||||||
header = f"{uid} (UID {uid} BODY[] {{200}})".encode()
|
return "OK", [(header, item["raw"]), b")"]
|
||||||
return "OK", [(header, raw), b")"]
|
|
||||||
|
def store(self, imap_id: bytes, _op: str, _flags: str):
|
||||||
|
mailbox_state[imap_id]["seen"] = True
|
||||||
return "OK", [b""]
|
return "OK", [b""]
|
||||||
|
|
||||||
def logout(self):
|
def logout(self):
|
||||||
@@ -942,13 +1044,12 @@ def test_fetch_messages_between_dates_uses_imap_since_before_without_mark_seen(m
|
|||||||
def select(self, _mailbox: str):
|
def select(self, _mailbox: str):
|
||||||
return "OK", [b"1"]
|
return "OK", [b"1"]
|
||||||
|
|
||||||
def uid(self, command: str, *args):
|
def search(self, *_args):
|
||||||
if command == "SEARCH":
|
self.search_args = _args
|
||||||
self.search_args = args
|
return "OK", [b"5"]
|
||||||
return "OK", [b"999"]
|
|
||||||
if command == "FETCH":
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
return "OK", [(b"5 (UID 999 BODY[] {200})", raw), b")"]
|
return "OK", [(b"5 (UID 999 BODY[] {200})", raw), b")"]
|
||||||
return "OK", [b""]
|
|
||||||
|
|
||||||
def store(self, imap_id: bytes, op: str, flags: str):
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
self.store_calls.append((imap_id, op, flags))
|
self.store_calls.append((imap_id, op, flags))
|
||||||
@@ -969,7 +1070,7 @@ def test_fetch_messages_between_dates_uses_imap_since_before_without_mark_seen(m
|
|||||||
|
|
||||||
assert len(items) == 1
|
assert len(items) == 1
|
||||||
assert items[0]["subject"] == "Status"
|
assert items[0]["subject"] == "Status"
|
||||||
# uid("SEARCH", None, "SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
# search(None, "SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
||||||
assert fake.search_args is not None
|
assert fake.search_args is not None
|
||||||
assert fake.search_args[1:] == ("SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
assert fake.search_args[1:] == ("SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026")
|
||||||
assert fake.store_calls == []
|
assert fake.store_calls == []
|
||||||
@@ -979,12 +1080,11 @@ def test_fetch_messages_between_dates_uses_imap_since_before_without_mark_seen(m
|
|||||||
# Security: Anti-spoofing tests for Authentication-Results verification
|
# Security: Anti-spoofing tests for Authentication-Results verification
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
def _make_fake_imap(raw: bytes, uid: bytes = b"500"):
|
def _make_fake_imap(raw: bytes):
|
||||||
"""Return a FakeIMAP class pre-loaded with the given raw email."""
|
"""Return a FakeIMAP class pre-loaded with the given raw email."""
|
||||||
class FakeIMAP:
|
class FakeIMAP:
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.store_calls: list[tuple[bytes, str, str]] = []
|
self.store_calls: list[tuple[bytes, str, str]] = []
|
||||||
self.uid_calls: list[tuple] = []
|
|
||||||
|
|
||||||
def login(self, _user: str, _pw: str):
|
def login(self, _user: str, _pw: str):
|
||||||
return "OK", [b"logged in"]
|
return "OK", [b"logged in"]
|
||||||
@@ -992,16 +1092,11 @@ def _make_fake_imap(raw: bytes, uid: bytes = b"500"):
|
|||||||
def select(self, _mailbox: str):
|
def select(self, _mailbox: str):
|
||||||
return "OK", [b"1"]
|
return "OK", [b"1"]
|
||||||
|
|
||||||
def capability(self):
|
def search(self, *_args):
|
||||||
return "OK", [b"IMAP4rev1"]
|
return "OK", [b"1"]
|
||||||
|
|
||||||
def uid(self, command: str, *args):
|
def fetch(self, _imap_id: bytes, _parts: str):
|
||||||
self.uid_calls.append((command, *args))
|
return "OK", [(b"1 (UID 500 BODY[] {200})", raw), b")"]
|
||||||
if command == "SEARCH":
|
|
||||||
return "OK", [uid]
|
|
||||||
if command == "FETCH":
|
|
||||||
return "OK", [(b"1 (UID " + uid + b" BODY[] {200})", raw), b")"]
|
|
||||||
return "OK", [b""]
|
|
||||||
|
|
||||||
def store(self, imap_id: bytes, op: str, flags: str):
|
def store(self, imap_id: bytes, op: str, flags: str):
|
||||||
self.store_calls.append((imap_id, op, flags))
|
self.store_calls.append((imap_id, op, flags))
|
||||||
@@ -1197,10 +1292,7 @@ def test_fetch_new_messages_ignores_unauthorized_sender_before_attachments(monke
|
|||||||
|
|
||||||
assert channel._fetch_new_messages() == ([], {"500"})
|
assert channel._fetch_new_messages() == ([], {"500"})
|
||||||
assert called["attachments"] is False
|
assert called["attachments"] is False
|
||||||
assert [call for call in fake.uid_calls if call[0] == "FETCH"] == [
|
assert fake.store_calls == [(b"1", "+FLAGS", "\\Seen")]
|
||||||
("FETCH", "500", "(BODY.PEEK[HEADER])")
|
|
||||||
]
|
|
||||||
assert ("STORE", "500", "+FLAGS", "(\\Seen)") in fake.uid_calls
|
|
||||||
|
|
||||||
|
|
||||||
def test_extract_attachments_saves_pdf(tmp_path, monkeypatch) -> None:
|
def test_extract_attachments_saves_pdf(tmp_path, monkeypatch) -> None:
|
||||||
|
|||||||
+3
-21
@@ -25,7 +25,6 @@ from nanobot.cron.types import (
|
|||||||
CronSchedule,
|
CronSchedule,
|
||||||
CronStore,
|
CronStore,
|
||||||
)
|
)
|
||||||
from nanobot.runtime_context import RUNTIME_CONTEXT_INPUT_META
|
|
||||||
from nanobot.utils.run_records import (
|
from nanobot.utils.run_records import (
|
||||||
write_run_record as write_automation_run_record,
|
write_run_record as write_automation_run_record,
|
||||||
)
|
)
|
||||||
@@ -116,21 +115,8 @@ def _disable_malformed_legacy_job(job: CronJob) -> None:
|
|||||||
logger.warning("Cron: disabled malformed legacy job '{}' ({}): {}", job.name, job.id, reason)
|
logger.warning("Cron: disabled malformed legacy job '{}' ({}): {}", job.name, job.id, reason)
|
||||||
|
|
||||||
|
|
||||||
def _persistable_origin_metadata(metadata: dict[str, Any]) -> dict[str, Any]:
|
|
||||||
"""Return a detached JSON-safe routing snapshot for a cron payload."""
|
|
||||||
snapshot: dict[str, Any] = {}
|
|
||||||
for key, value in metadata.items():
|
|
||||||
if key == RUNTIME_CONTEXT_INPUT_META:
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
snapshot[key] = json.loads(json.dumps(value, ensure_ascii=False, allow_nan=False))
|
|
||||||
except (TypeError, ValueError, RecursionError):
|
|
||||||
continue
|
|
||||||
return snapshot
|
|
||||||
|
|
||||||
|
|
||||||
def _normalize_agent_turn_job(job: CronJob) -> bool:
|
def _normalize_agent_turn_job(job: CronJob) -> bool:
|
||||||
"""Make routing metadata persistable and migrate legacy user cron payloads.
|
"""Migrate legacy user cron payloads into session-bound payloads.
|
||||||
|
|
||||||
Pre-bound user cron jobs stored their delivery target in ``channel``/``to``.
|
Pre-bound user cron jobs stored their delivery target in ``channel``/``to``.
|
||||||
Normal user-created legacy jobs always have those fields; if they are
|
Normal user-created legacy jobs always have those fields; if they are
|
||||||
@@ -138,12 +124,8 @@ def _normalize_agent_turn_job(job: CronJob) -> bool:
|
|||||||
a runtime legacy execution path.
|
a runtime legacy execution path.
|
||||||
"""
|
"""
|
||||||
payload = job.payload
|
payload = job.payload
|
||||||
origin_metadata = _persistable_origin_metadata(payload.origin_metadata)
|
|
||||||
changed = origin_metadata != payload.origin_metadata
|
|
||||||
payload.origin_metadata = origin_metadata
|
|
||||||
|
|
||||||
if payload.kind != "agent_turn" or not _has_legacy_delivery_context(payload):
|
if payload.kind != "agent_turn" or not _has_legacy_delivery_context(payload):
|
||||||
return changed
|
return False
|
||||||
|
|
||||||
if not payload.channel or not payload.to:
|
if not payload.channel or not payload.to:
|
||||||
_disable_malformed_legacy_job(job)
|
_disable_malformed_legacy_job(job)
|
||||||
@@ -153,7 +135,7 @@ def _normalize_agent_turn_job(job: CronJob) -> bool:
|
|||||||
payload.origin_channel = payload.origin_channel or payload.channel
|
payload.origin_channel = payload.origin_channel or payload.channel
|
||||||
payload.origin_chat_id = payload.origin_chat_id or payload.to
|
payload.origin_chat_id = payload.origin_chat_id or payload.to
|
||||||
if not payload.origin_metadata:
|
if not payload.origin_metadata:
|
||||||
payload.origin_metadata = _persistable_origin_metadata(payload.channel_meta or {})
|
payload.origin_metadata = dict(payload.channel_meta or {})
|
||||||
|
|
||||||
payload.deliver = False
|
payload.deliver = False
|
||||||
payload.channel = None
|
payload.channel = None
|
||||||
|
|||||||
@@ -1029,20 +1029,6 @@ class LLMProvider(ABC):
|
|||||||
# Unknown 429 defaults to WAIT+retry.
|
# Unknown 429 defaults to WAIT+retry.
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _content_as_blocks(content: Any) -> list[dict[str, Any]]:
|
|
||||||
"""Convert message content to blocks so mixed user content can be merged."""
|
|
||||||
if isinstance(content, list):
|
|
||||||
return [
|
|
||||||
dict(cast(dict[str, Any], item))
|
|
||||||
if isinstance(item, dict)
|
|
||||||
else {"type": "text", "text": str(item)}
|
|
||||||
for item in cast(list[object], content)
|
|
||||||
]
|
|
||||||
if content is None:
|
|
||||||
return []
|
|
||||||
return [{"type": "text", "text": str(content)}]
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _enforce_role_alternation(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
def _enforce_role_alternation(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||||
"""Merge consecutive same-role messages and drop trailing assistant messages.
|
"""Merge consecutive same-role messages and drop trailing assistant messages.
|
||||||
@@ -1077,13 +1063,6 @@ class LLMProvider(ABC):
|
|||||||
curr_content = msg.get("content") or ""
|
curr_content = msg.get("content") or ""
|
||||||
if isinstance(prev_content, str) and isinstance(curr_content, str):
|
if isinstance(prev_content, str) and isinstance(curr_content, str):
|
||||||
prev["content"] = (prev_content + "\n\n" + curr_content).strip()
|
prev["content"] = (prev_content + "\n\n" + curr_content).strip()
|
||||||
elif role == "user":
|
|
||||||
combined = dict(msg)
|
|
||||||
combined["content"] = [
|
|
||||||
*LLMProvider._content_as_blocks(prev_content),
|
|
||||||
*LLMProvider._content_as_blocks(curr_content),
|
|
||||||
]
|
|
||||||
merged[-1] = combined
|
|
||||||
else:
|
else:
|
||||||
merged[-1] = dict(msg)
|
merged[-1] = dict(msg)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -147,10 +147,10 @@ def prepare_save_boundary(ctx: TurnContext) -> None:
|
|||||||
if ctx.session is not None:
|
if ctx.session is not None:
|
||||||
clear_internal_continuation_state(ctx.session.metadata)
|
clear_internal_continuation_state(ctx.session.metadata)
|
||||||
|
|
||||||
assert ctx.transcript_input is not None
|
|
||||||
ctx.save_skip = _save_skip_for_turn(
|
ctx.save_skip = _save_skip_for_turn(
|
||||||
message_metadata=ctx.msg.metadata,
|
message_metadata=ctx.msg.metadata,
|
||||||
initial_message_count=ctx.transcript_input.message_count,
|
initial_message_count=len(ctx.initial_messages),
|
||||||
|
history_count=len(ctx.history),
|
||||||
input_persisted_early=ctx.input_persisted_early,
|
input_persisted_early=ctx.input_persisted_early,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -185,6 +185,7 @@ def _save_skip_for_turn(
|
|||||||
*,
|
*,
|
||||||
message_metadata: Mapping[str, Any] | None,
|
message_metadata: Mapping[str, Any] | None,
|
||||||
initial_message_count: int,
|
initial_message_count: int,
|
||||||
|
history_count: int,
|
||||||
input_persisted_early: bool,
|
input_persisted_early: bool,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Return the persisted-message append boundary for this turn."""
|
"""Return the persisted-message append boundary for this turn."""
|
||||||
@@ -192,7 +193,10 @@ def _save_skip_for_turn(
|
|||||||
return initial_message_count
|
return initial_message_count
|
||||||
if internal_continuation_inbound(message_metadata):
|
if internal_continuation_inbound(message_metadata):
|
||||||
return initial_message_count
|
return initial_message_count
|
||||||
if not input_persisted_early:
|
# build_messages may merge the current message into a same-role history tail.
|
||||||
|
# Runner-appended messages start at initial_message_count in either shape.
|
||||||
|
has_standalone_current = initial_message_count > 1 + history_count
|
||||||
|
if has_standalone_current and not input_persisted_early:
|
||||||
return initial_message_count - 1
|
return initial_message_count - 1
|
||||||
return initial_message_count
|
return initial_message_count
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.loop import AgentLoop, TurnContext, TurnKind
|
from nanobot.agent.loop import AgentLoop, TurnContext, TurnKind
|
||||||
from nanobot.agent.tools.context import RequestContext
|
from nanobot.agent.tools.context import RequestContext
|
||||||
from nanobot.agent.tools.filesystem import ReadFileTool
|
from nanobot.agent.tools.filesystem import ReadFileTool
|
||||||
@@ -149,10 +148,7 @@ async def test_pending_document_attachment_keeps_body_out_of_prompt(
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(
|
[{"role": "user", "content": "hello"}],
|
||||||
history=[{"role": "user", "content": "hello"}],
|
|
||||||
current_message=None,
|
|
||||||
),
|
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import ContextBuilder, TranscriptInput
|
from nanobot.agent.context import ContextBuilder
|
||||||
from nanobot.runtime_context import RuntimeContextBlock
|
from nanobot.runtime_context import RuntimeContextBlock
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -403,15 +403,6 @@ class TestBuildMessages:
|
|||||||
assert "user-only runtime context" not in messages[-1]["content"]
|
assert "user-only runtime context" not in messages[-1]["content"]
|
||||||
assert "_meta" not in messages[-1]
|
assert "_meta" not in messages[-1]
|
||||||
|
|
||||||
def test_compatibility_builder_merges_system_role_without_history(self, tmp_path):
|
|
||||||
builder = _builder(tmp_path)
|
|
||||||
|
|
||||||
messages = builder.build_messages([], "system event", current_role="system")
|
|
||||||
|
|
||||||
assert len(messages) == 1
|
|
||||||
assert messages[0]["role"] == "system"
|
|
||||||
assert str(messages[0]["content"]).endswith("system event")
|
|
||||||
|
|
||||||
def test_explicit_skill_reference_loads_full_instructions_for_this_turn(self, tmp_path):
|
def test_explicit_skill_reference_loads_full_instructions_for_this_turn(self, tmp_path):
|
||||||
skill_dir = tmp_path / "skills" / "review"
|
skill_dir = tmp_path / "skills" / "review"
|
||||||
skill_dir.mkdir(parents=True)
|
skill_dir.mkdir(parents=True)
|
||||||
@@ -481,20 +472,6 @@ class TestBuildMessages:
|
|||||||
assert "previous user message" in str(messages[1]["content"])
|
assert "previous user message" in str(messages[1]["content"])
|
||||||
assert "new message" in str(messages[1]["content"])
|
assert "new message" in str(messages[1]["content"])
|
||||||
|
|
||||||
def test_structured_transcript_preserves_fresh_turn_boundary(self, tmp_path):
|
|
||||||
builder = _builder(tmp_path)
|
|
||||||
transcript = TranscriptInput(
|
|
||||||
history=[{"role": "user", "content": "previous user message"}],
|
|
||||||
current_message="new message",
|
|
||||||
)
|
|
||||||
|
|
||||||
messages = builder.build_transcript(transcript)
|
|
||||||
|
|
||||||
assert [message["role"] for message in messages] == ["system", "user", "user"]
|
|
||||||
assert messages[-2]["content"] == "previous user message"
|
|
||||||
assert messages[-1]["content"] == "new message"
|
|
||||||
assert transcript.message_count == 3
|
|
||||||
|
|
||||||
def test_current_message_can_be_built_without_history_merge(self, tmp_path):
|
def test_current_message_can_be_built_without_history_merge(self, tmp_path):
|
||||||
builder = _builder(tmp_path)
|
builder = _builder(tmp_path)
|
||||||
current = builder.build_current_message(
|
current = builder.build_current_message(
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.hook import (
|
from nanobot.agent.hook import (
|
||||||
AgentHook,
|
AgentHook,
|
||||||
AgentHookContext,
|
AgentHookContext,
|
||||||
@@ -460,7 +459,7 @@ async def test_agent_loop_extra_hook_receives_calls(tmp_path):
|
|||||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
|
[{"role": "user", "content": "hi"}],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -505,7 +504,7 @@ async def test_agent_loop_turn_hook_factories_receive_context(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
|
[{"role": "user", "content": "hi"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
@@ -552,7 +551,7 @@ async def test_agent_loop_extra_hook_error_isolation(tmp_path):
|
|||||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
|
[{"role": "user", "content": "hi"}],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -578,9 +577,7 @@ async def test_agent_loop_extra_hooks_do_not_swallow_loop_hook_errors(tmp_path):
|
|||||||
|
|
||||||
with pytest.raises(RuntimeError, match="progress failed"):
|
with pytest.raises(RuntimeError, match="progress failed"):
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=bad_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=bad_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -599,8 +596,7 @@ async def test_agent_loop_no_hooks_backward_compat(tmp_path):
|
|||||||
loop.max_iterations = 2
|
loop.max_iterations = 2
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime()
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
)
|
||||||
assert result.final_content == (
|
assert result.final_content == (
|
||||||
"I reached the maximum number of tool call iterations (2) "
|
"I reached the maximum number of tool call iterations (2) "
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.hooks import create_file_edit_activity_hook
|
from nanobot.agent.hooks import create_file_edit_activity_hook
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.tools.context import current_request_context
|
from nanobot.agent.tools.context import current_request_context
|
||||||
@@ -85,9 +84,7 @@ class TestToolEventProgress:
|
|||||||
progress.append((content, tool_hint, tool_events))
|
progress.append((content, tool_hint, tool_events))
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=on_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=on_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Done"
|
assert result.final_content == "Done"
|
||||||
@@ -158,9 +155,7 @@ class TestToolEventProgress:
|
|||||||
file_events.extend(file_edit_events)
|
file_events.extend(file_edit_events)
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=on_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=on_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Done"
|
assert result.final_content == "Done"
|
||||||
@@ -230,9 +225,7 @@ class TestToolEventProgress:
|
|||||||
)
|
)
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=on_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=on_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Done"
|
assert result.final_content == "Done"
|
||||||
@@ -270,9 +263,7 @@ class TestToolEventProgress:
|
|||||||
file_events.extend(file_edit_events)
|
file_events.extend(file_edit_events)
|
||||||
|
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=on_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=on_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert file_events == []
|
assert file_events == []
|
||||||
@@ -1028,7 +1019,7 @@ class TestToolEventProgress:
|
|||||||
progress.append((content, tool_hint, tool_events))
|
progress.append((content, tool_hint, tool_events))
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
on_stream=on_stream,
|
on_stream=on_stream,
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.goal_permission import goal_mutation_allowed, goal_mutation_permission
|
from nanobot.agent.goal_permission import goal_mutation_allowed, goal_mutation_permission
|
||||||
from nanobot.agent.tools.context import RequestContext
|
from nanobot.agent.tools.context import RequestContext
|
||||||
from nanobot.bus.outbound_events import StreamedResponseEvent
|
from nanobot.bus.outbound_events import StreamedResponseEvent
|
||||||
@@ -56,7 +55,7 @@ async def test_ephemeral_runner_enters_and_restores_turn_scopes(tmp_path):
|
|||||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
ephemeral=True,
|
ephemeral=True,
|
||||||
turn_scopes=[goal_mutation_permission(True)],
|
turn_scopes=[goal_mutation_permission(True)],
|
||||||
@@ -341,8 +340,7 @@ async def test_loop_max_iterations_message_stays_stable(tmp_path):
|
|||||||
loop.max_iterations = 2
|
loop.max_iterations = 2
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime()
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == (
|
assert result.final_content == (
|
||||||
@@ -364,7 +362,7 @@ async def test_loop_goal_turn_uses_standard_iteration_budget(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
channel="cli",
|
channel="cli",
|
||||||
@@ -403,7 +401,7 @@ async def test_loop_stream_filter_handles_think_only_prefix_without_crashing(tmp
|
|||||||
endings.append(resuming)
|
endings.append(resuming)
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
on_stream=on_stream,
|
on_stream=on_stream,
|
||||||
on_stream_end=on_stream_end,
|
on_stream_end=on_stream_end,
|
||||||
@@ -430,9 +428,7 @@ async def test_loop_stream_filter_hides_partial_trailing_think_prefix(tmp_path):
|
|||||||
deltas.append(delta)
|
deltas.append(delta)
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_stream=on_stream
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_stream=on_stream,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Hello World"
|
assert result.final_content == "Hello World"
|
||||||
@@ -455,9 +451,7 @@ async def test_loop_stream_filter_hides_complete_trailing_think_tag(tmp_path):
|
|||||||
deltas.append(delta)
|
deltas.append(delta)
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_stream=on_stream
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_stream=on_stream,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Hello World"
|
assert result.final_content == "Hello World"
|
||||||
@@ -478,8 +472,7 @@ async def test_loop_retries_think_only_final_response(tmp_path):
|
|||||||
loop.provider.chat_with_retry = chat_with_retry
|
loop.provider.chat_with_retry = chat_with_retry
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime()
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Recovered answer"
|
assert result.final_content == "Recovered answer"
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.agent.context import ContextBuilder, TranscriptInput
|
from nanobot.agent.context import ContextBuilder
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.runner import AgentRunResult
|
from nanobot.agent.runner import AgentRunResult
|
||||||
from nanobot.agent.tools.context import RequestContext, request_context
|
from nanobot.agent.tools.context import RequestContext, request_context
|
||||||
@@ -79,13 +79,6 @@ def _agent_run_result(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _assembled_messages(
|
|
||||||
builder: ContextBuilder,
|
|
||||||
transcript_input: TranscriptInput,
|
|
||||||
) -> list[dict]:
|
|
||||||
return builder.build_transcript(transcript_input, include_memory=False)
|
|
||||||
|
|
||||||
|
|
||||||
def _mk_loop() -> AgentLoop:
|
def _mk_loop() -> AgentLoop:
|
||||||
loop = AgentLoop.__new__(AgentLoop)
|
loop = AgentLoop.__new__(AgentLoop)
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
@@ -937,13 +930,10 @@ async def test_runtime_checkpoint_keeps_provider_state_out_of_public_metadata(
|
|||||||
session = loop.sessions.get_or_create("cli:private-checkpoint")
|
session = loop.sessions.get_or_create("cli:private-checkpoint")
|
||||||
|
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(
|
[
|
||||||
history=[
|
{"role": "system", "content": "system"},
|
||||||
{"role": "system", "content": "system"},
|
{"role": "user", "content": "question"},
|
||||||
{"role": "user", "content": "question"},
|
],
|
||||||
],
|
|
||||||
current_message=None,
|
|
||||||
),
|
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
session=session,
|
session=session,
|
||||||
)
|
)
|
||||||
@@ -1018,7 +1008,7 @@ async def test_subagent_followup_state_is_durable_before_prompt_assembly(
|
|||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
loop.provider.can_resume_conversation_state.return_value = True
|
loop.provider.can_resume_conversation_state.return_value = True
|
||||||
loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign]
|
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
|
||||||
side_effect=RuntimeError("prompt boom"),
|
side_effect=RuntimeError("prompt boom"),
|
||||||
)
|
)
|
||||||
session = loop.sessions.get_or_create("cli:subagent-prompt-crash")
|
session = loop.sessions.get_or_create("cli:subagent-prompt-crash")
|
||||||
@@ -1051,8 +1041,8 @@ async def test_subagent_redelivery_does_not_duplicate_staged_provider_input(
|
|||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
loop.provider.can_resume_conversation_state.return_value = True
|
loop.provider.can_resume_conversation_state.return_value = True
|
||||||
build_system_prompt = loop.context.build_system_prompt
|
build_initial_messages = loop._build_initial_messages
|
||||||
loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign]
|
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
|
||||||
side_effect=RuntimeError("prompt boom"),
|
side_effect=RuntimeError("prompt boom"),
|
||||||
)
|
)
|
||||||
session = loop.sessions.get_or_create("cli:subagent-redelivery")
|
session = loop.sessions.get_or_create("cli:subagent-redelivery")
|
||||||
@@ -1076,7 +1066,7 @@ async def test_subagent_redelivery_does_not_duplicate_staged_provider_input(
|
|||||||
message.get("content")
|
message.get("content")
|
||||||
for message in persisted.provider_state.pending_messages
|
for message in persisted.provider_state.pending_messages
|
||||||
].count("subagent result") == 1
|
].count("subagent result") == 1
|
||||||
loop.context.build_system_prompt = build_system_prompt # type: ignore[method-assign]
|
loop._build_initial_messages = build_initial_messages # type: ignore[method-assign]
|
||||||
loop._run_agent_loop = AsyncMock( # type: ignore[method-assign]
|
loop._run_agent_loop = AsyncMock( # type: ignore[method-assign]
|
||||||
side_effect=RuntimeError("provider boom"),
|
side_effect=RuntimeError("provider boom"),
|
||||||
)
|
)
|
||||||
@@ -1329,8 +1319,7 @@ async def test_internal_continuation_queues_turn_without_fake_user_history(
|
|||||||
|
|
||||||
calls: list[dict] = []
|
calls: list[dict] = []
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, *, metadata=None, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, *, metadata=None, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
calls.append({"initial_messages": initial_messages, "metadata": metadata})
|
calls.append({"initial_messages": initial_messages, "metadata": metadata})
|
||||||
if len(calls) == 1:
|
if len(calls) == 1:
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
@@ -1398,9 +1387,8 @@ async def test_internal_continuation_preserves_streaming_route_metadata(
|
|||||||
|
|
||||||
calls = 0
|
calls = 0
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, *, on_stream=None, on_stream_end=None, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, *, on_stream=None, on_stream_end=None, **_kwargs):
|
||||||
nonlocal calls
|
nonlocal calls
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
calls += 1
|
calls += 1
|
||||||
if calls == 1:
|
if calls == 1:
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
@@ -1472,9 +1460,8 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
|
|||||||
|
|
||||||
calls = 0
|
calls = 0
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
nonlocal calls
|
nonlocal calls
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
calls += 1
|
calls += 1
|
||||||
if calls == 1:
|
if calls == 1:
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
@@ -1636,7 +1623,7 @@ async def test_run_agent_loop_continuation_reads_latest_goal_metadata(
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
session=session,
|
session=session,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
@@ -1766,7 +1753,7 @@ async def test_stop_preserves_runtime_checkpoint_for_next_turn(tmp_path: Path) -
|
|||||||
|
|
||||||
checkpoint_saved = asyncio.Event()
|
checkpoint_saved = asyncio.Event()
|
||||||
|
|
||||||
async def interrupted_run_agent_loop(_transcript_input, *, session=None, **_kwargs):
|
async def interrupted_run_agent_loop(_initial_messages, *, session=None, **_kwargs):
|
||||||
assert session is not None
|
assert session is not None
|
||||||
loop._set_runtime_checkpoint(
|
loop._set_runtime_checkpoint(
|
||||||
session,
|
session,
|
||||||
@@ -1826,8 +1813,7 @@ async def test_stop_preserves_runtime_checkpoint_for_next_turn(tmp_path: Path) -
|
|||||||
assert interrupted.metadata.get(AgentLoop._PENDING_USER_TURN_KEY) is True
|
assert interrupted.metadata.get(AgentLoop._PENDING_USER_TURN_KEY) is True
|
||||||
assert interrupted.metadata.get(AgentLoop._RUNTIME_CHECKPOINT_KEY) is not None
|
assert interrupted.metadata.get(AgentLoop._RUNTIME_CHECKPOINT_KEY) is not None
|
||||||
|
|
||||||
async def resumed_run_agent_loop(transcript_input, **_kwargs):
|
async def resumed_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"next answer",
|
"next answer",
|
||||||
[*initial_messages, {"role": "assistant", "content": "next answer"}],
|
[*initial_messages, {"role": "assistant", "content": "next answer"}],
|
||||||
@@ -1878,8 +1864,7 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
|
|||||||
record_runtime = MagicMock(wraps=loop.runtime_event_publisher.record_turn_runtime)
|
record_runtime = MagicMock(wraps=loop.runtime_event_publisher.record_turn_runtime)
|
||||||
loop.runtime_event_publisher.record_turn_runtime = record_runtime
|
loop.runtime_event_publisher.record_turn_runtime = record_runtime
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **kwargs):
|
async def fake_run_agent_loop(initial_messages, **kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
seen["initial_messages"] = initial_messages
|
seen["initial_messages"] = initial_messages
|
||||||
seen["runtime"] = kwargs["runtime"]
|
seen["runtime"] = kwargs["runtime"]
|
||||||
seen["request_context"] = kwargs["request_context"]
|
seen["request_context"] = kwargs["request_context"]
|
||||||
@@ -1955,8 +1940,7 @@ async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) ->
|
|||||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
turn_usage = LLMUsage.reported(input_tokens=64, output_tokens=9)
|
turn_usage = LLMUsage.reported(input_tokens=64, output_tokens=9)
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"done",
|
"done",
|
||||||
[*initial_messages, {"role": "assistant", "content": "done"}],
|
[*initial_messages, {"role": "assistant", "content": "done"}],
|
||||||
@@ -1982,8 +1966,7 @@ async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) ->
|
|||||||
return_value=False
|
return_value=False
|
||||||
)
|
)
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"done",
|
"done",
|
||||||
[*initial_messages, {"role": "assistant", "content": "done"}],
|
[*initial_messages, {"role": "assistant", "content": "done"}],
|
||||||
@@ -2039,8 +2022,7 @@ async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Pat
|
|||||||
|
|
||||||
setattr(loop, name, record)
|
setattr(loop, name, record)
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"done",
|
"done",
|
||||||
[*initial_messages, {"role": "assistant", "content": "done"}],
|
[*initial_messages, {"role": "assistant", "content": "done"}],
|
||||||
@@ -2083,8 +2065,7 @@ async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp
|
|||||||
loop = _make_full_loop(tmp_path)
|
loop = _make_full_loop(tmp_path)
|
||||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"ack",
|
"ack",
|
||||||
[*initial_messages, {"role": "assistant", "content": "ack"}],
|
[*initial_messages, {"role": "assistant", "content": "ack"}],
|
||||||
@@ -2215,8 +2196,7 @@ async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(t
|
|||||||
|
|
||||||
seen: dict[str, object] = {}
|
seen: dict[str, object] = {}
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **kwargs):
|
async def fake_run_agent_loop(initial_messages, **kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
|
||||||
seen["initial_messages"] = initial_messages
|
seen["initial_messages"] = initial_messages
|
||||||
seen["request_context"] = kwargs["request_context"]
|
seen["request_context"] = kwargs["request_context"]
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
@@ -2272,11 +2252,8 @@ async def test_turn_after_unanswered_user_keeps_tool_call_pairing(tmp_path: Path
|
|||||||
session.add_message("user", "earlier question that never got an answer")
|
session.add_message("user", "earlier question that never got an answer")
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
async def fake_run_agent_loop(transcript_input, **_kwargs):
|
async def fake_run_agent_loop(initial_messages, **_kwargs):
|
||||||
initial_messages = _assembled_messages(loop.context, transcript_input)
|
assert [m["role"] for m in initial_messages] == ["system", "user"]
|
||||||
assert [m["role"] for m in initial_messages] == ["system", "user", "user"]
|
|
||||||
assert initial_messages[-2]["content"] == "earlier question that never got an answer"
|
|
||||||
assert initial_messages[-1]["content"] == "and another thing"
|
|
||||||
return _agent_run_result(
|
return _agent_run_result(
|
||||||
"done",
|
"done",
|
||||||
[
|
[
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.tools.context import (
|
from nanobot.agent.tools.context import (
|
||||||
RequestContext,
|
RequestContext,
|
||||||
@@ -134,7 +133,7 @@ async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) ->
|
|||||||
metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}}
|
metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}}
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
channel="slack",
|
channel="slack",
|
||||||
@@ -235,7 +234,7 @@ async def test_agent_loop_restores_outer_request_context_after_runner_exception(
|
|||||||
try:
|
try:
|
||||||
with pytest.raises(RuntimeError, match="runner failed"):
|
with pytest.raises(RuntimeError, match="runner failed"):
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
channel="slack",
|
channel="slack",
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import (
|
from nanobot.providers.base import (
|
||||||
LLMProvider,
|
LLMProvider,
|
||||||
@@ -35,35 +34,6 @@ def _make_usage_spec(provider, tools):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_initial_transcript_is_built_from_structured_turn_input() -> None:
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock(spec=LLMProvider)
|
|
||||||
transcript_input = TranscriptInput(
|
|
||||||
history=[{"role": "user", "content": "earlier"}],
|
|
||||||
current_message="fresh",
|
|
||||||
)
|
|
||||||
expected = [
|
|
||||||
{"role": "system", "content": "system"},
|
|
||||||
{"role": "user", "content": "earlier"},
|
|
||||||
{"role": "user", "content": "fresh"},
|
|
||||||
]
|
|
||||||
transcript_builder = MagicMock(return_value=expected)
|
|
||||||
spec = make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=None,
|
|
||||||
transcript_input=transcript_input,
|
|
||||||
transcript_builder=transcript_builder,
|
|
||||||
tools=MagicMock(),
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=1,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert AgentRunner._initial_transcript(spec) == expected
|
|
||||||
transcript_builder.assert_called_once_with(transcript_input)
|
|
||||||
|
|
||||||
|
|
||||||
def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> None:
|
def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> None:
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
|
|||||||
@@ -52,31 +52,7 @@ async def test_runner_returns_tool_exception_to_model_for_recovery():
|
|||||||
{"name": "list_dir", "status": "error", "detail": "boom"}
|
{"name": "list_dir", "status": "error", "detail": "boom"}
|
||||||
]
|
]
|
||||||
tool_message = next(message for message in result.messages if message.get("role") == "tool")
|
tool_message = next(message for message in result.messages if message.get("role") == "tool")
|
||||||
retry_hint = "[Analyze the error above and try a different approach.]"
|
|
||||||
assert "Error: RuntimeError: boom" in tool_message["content"]
|
assert "Error: RuntimeError: boom" in tool_message["content"]
|
||||||
assert tool_message["content"].count(retry_hint) == 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_tool_execution_does_not_duplicate_existing_retry_hint():
|
|
||||||
retry_hint = "\n\n[Analyze the error above and try a different approach.]"
|
|
||||||
tools = SimpleNamespace(
|
|
||||||
execute=AsyncMock(return_value=ToolResult.error("Error: boom" + retry_hint)),
|
|
||||||
)
|
|
||||||
|
|
||||||
results, events = await execute_tool_calls(
|
|
||||||
tools,
|
|
||||||
[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
|
||||||
concurrent=False,
|
|
||||||
external_lookup_counts={},
|
|
||||||
workspace_violation_counts={},
|
|
||||||
hook=AgentHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
|
||||||
)
|
|
||||||
|
|
||||||
assert results == ["Error: boom" + retry_hint]
|
|
||||||
assert results[0].count(retry_hint) == 1
|
|
||||||
assert events[0]["status"] == "error"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import pytest
|
|||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.automation_turns import publish_next_deferred_turn
|
from nanobot.agent.automation_turns import publish_next_deferred_turn
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.tools.context import RequestContext
|
from nanobot.agent.tools.context import RequestContext
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
@@ -618,7 +617,7 @@ async def test_loop_injected_followup_preserves_image_media(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hello"}], current_message=None),
|
[{"role": "user", "content": "hello"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
@@ -712,10 +711,7 @@ async def test_pending_injection_resolves_its_own_runtime_context(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(
|
[{"role": "user", "content": "initial message from user A"}],
|
||||||
history=[{"role": "user", "content": "initial message from user A"}],
|
|
||||||
current_message=None,
|
|
||||||
),
|
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
session=session,
|
session=session,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
@@ -816,7 +812,7 @@ async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_p
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hello"}], current_message=None),
|
[{"role": "user", "content": "hello"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
@@ -1480,7 +1476,7 @@ async def test_pending_queue_preserves_overflow_for_next_injection_cycle(tmp_pat
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "hello"}], current_message=None),
|
[{"role": "user", "content": "hello"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
request_context=RequestContext(channel="cli", chat_id="c", runtime=runtime),
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ channels, gated by ``context.streamed_reasoning`` rather than
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
@@ -83,18 +82,6 @@ class _LifecycleRecordingHook(AgentHook):
|
|||||||
self.events.append(f"hosted_tool:{event.get('phase')}")
|
self.events.append(f"hosted_tool:{event.get('phase')}")
|
||||||
|
|
||||||
|
|
||||||
class _BlockingReasoningEndHook(_LifecycleRecordingHook):
|
|
||||||
def __init__(self) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.reasoning_end_started = asyncio.Event()
|
|
||||||
self.release_reasoning_end = asyncio.Event()
|
|
||||||
|
|
||||||
async def emit_reasoning_end(self) -> None:
|
|
||||||
self.reasoning_end_started.set()
|
|
||||||
await self.release_reasoning_end.wait()
|
|
||||||
await super().emit_reasoning_end()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_preserves_reasoning_fields_in_assistant_history():
|
async def test_runner_preserves_reasoning_fields_in_assistant_history():
|
||||||
"""Reasoning fields ride along on the persisted assistant message so
|
"""Reasoning fields ride along on the persisted assistant message so
|
||||||
@@ -567,86 +554,6 @@ async def test_runner_closes_native_reasoning_before_hosted_tool_event():
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_closes_native_reasoning_when_stream_is_cancelled():
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock()
|
|
||||||
reasoning_started = asyncio.Event()
|
|
||||||
release_provider = asyncio.Event()
|
|
||||||
|
|
||||||
async def chat_stream_with_retry(
|
|
||||||
*, on_thinking_delta=None, **kwargs
|
|
||||||
):
|
|
||||||
if on_thinking_delta:
|
|
||||||
await on_thinking_delta("inspect")
|
|
||||||
reasoning_started.set()
|
|
||||||
await release_provider.wait()
|
|
||||||
raise AssertionError("the cancelled provider call should not complete")
|
|
||||||
|
|
||||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
|
||||||
tools = MagicMock()
|
|
||||||
tools.get_definitions.return_value = []
|
|
||||||
hook = _LifecycleRecordingHook()
|
|
||||||
|
|
||||||
task = asyncio.create_task(AgentRunner().run(make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=[{"role": "user", "content": "inspect"}],
|
|
||||||
tools=tools,
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=1,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
hook=hook,
|
|
||||||
)))
|
|
||||||
await reasoning_started.wait()
|
|
||||||
|
|
||||||
task.cancel()
|
|
||||||
with pytest.raises(asyncio.CancelledError):
|
|
||||||
await task
|
|
||||||
|
|
||||||
assert hook.events == ["reasoning:inspect", "reasoning_end"]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_settles_native_reasoning_end_before_propagating_cancellation():
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock()
|
|
||||||
|
|
||||||
async def chat_stream_with_retry(
|
|
||||||
*, on_content_delta=None, on_thinking_delta=None, **kwargs
|
|
||||||
):
|
|
||||||
if on_thinking_delta:
|
|
||||||
await on_thinking_delta("inspect")
|
|
||||||
if on_content_delta:
|
|
||||||
await on_content_delta("done")
|
|
||||||
raise AssertionError("the cancelled provider call should not complete")
|
|
||||||
|
|
||||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
|
||||||
tools = MagicMock()
|
|
||||||
tools.get_definitions.return_value = []
|
|
||||||
hook = _BlockingReasoningEndHook()
|
|
||||||
|
|
||||||
task = asyncio.create_task(AgentRunner().run(make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=[{"role": "user", "content": "inspect"}],
|
|
||||||
tools=tools,
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=1,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
hook=hook,
|
|
||||||
)))
|
|
||||||
await hook.reasoning_end_started.wait()
|
|
||||||
|
|
||||||
task.cancel()
|
|
||||||
await asyncio.sleep(0)
|
|
||||||
hook.release_reasoning_end.set()
|
|
||||||
with pytest.raises(asyncio.CancelledError):
|
|
||||||
await task
|
|
||||||
|
|
||||||
assert hook.events == ["reasoning:inspect", "reasoning_end"]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_strips_thinking_tags_from_native_thinking_deltas():
|
async def test_runner_strips_thinking_tags_from_native_thinking_deltas():
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.tools.context import RequestContext
|
from nanobot.agent.tools.context import RequestContext
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import GenerationSettings
|
from nanobot.providers.base import GenerationSettings
|
||||||
@@ -569,10 +568,7 @@ async def test_agent_loop_syncs_updated_max_iterations_before_run(tmp_path):
|
|||||||
loop.runner.run = AsyncMock(side_effect=fake_run)
|
loop.runner.run = AsyncMock(side_effect=fake_run)
|
||||||
loop.max_iterations = 55
|
loop.max_iterations = 55
|
||||||
|
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop([], runtime=loop.llm_runtime())
|
||||||
TranscriptInput(history=[], current_message=None),
|
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
|
||||||
|
|
||||||
loop.runner.run.assert_awaited_once()
|
loop.runner.run.assert_awaited_once()
|
||||||
|
|
||||||
@@ -613,7 +609,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
|
[{"role": "user", "content": "test"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
session=None,
|
session=None,
|
||||||
request_context=RequestContext(channel="test", chat_id="c1", runtime=runtime),
|
request_context=RequestContext(channel="test", chat_id="c1", runtime=runtime),
|
||||||
@@ -672,7 +668,7 @@ async def test_terminal_drain_timeout(tmp_path):
|
|||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
|
[{"role": "user", "content": "test"}],
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
session=session,
|
session=session,
|
||||||
request_context=RequestContext(
|
request_context=RequestContext(
|
||||||
@@ -746,7 +742,7 @@ async def test_terminal_drain_reuses_one_timeout_budget(tmp_path):
|
|||||||
loop.subagents._running_tasks["sub-deadline-1"] = hang_task
|
loop.subagents._running_tasks["sub-deadline-1"] = hang_task
|
||||||
|
|
||||||
await loop._run_agent_loop(
|
await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
|
[{"role": "user", "content": "test"}],
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
session=session,
|
session=session,
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.bus.events import InboundMessage
|
from nanobot.bus.events import InboundMessage
|
||||||
from nanobot.providers.base import LLMResponse, LLMUsage
|
from nanobot.providers.base import LLMResponse, LLMUsage
|
||||||
|
|
||||||
@@ -312,16 +311,10 @@ class TestRestartCommand:
|
|||||||
LLMResponse(content="second", usage=None),
|
LLMResponse(content="second", usage=None),
|
||||||
])
|
])
|
||||||
|
|
||||||
first = await loop._run_agent_loop(
|
first = await loop._run_agent_loop([], runtime=loop.llm_runtime())
|
||||||
TranscriptInput(history=[], current_message=None),
|
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
|
||||||
assert first.usage == LLMUsage.reported(input_tokens=9, output_tokens=4)
|
assert first.usage == LLMUsage.reported(input_tokens=9, output_tokens=4)
|
||||||
|
|
||||||
second = await loop._run_agent_loop(
|
second = await loop._run_agent_loop([], runtime=loop.llm_runtime())
|
||||||
TranscriptInput(history=[], current_message=None),
|
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
)
|
|
||||||
assert second.usage == LLMUsage.estimated(input_tokens=123, output_tokens=7)
|
assert second.usage == LLMUsage.estimated(input_tokens=123, output_tokens=7)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.cron.service import CronJobSkippedError, CronService
|
from nanobot.cron.service import CronJobSkippedError, CronService
|
||||||
from nanobot.cron.types import CronJob, CronPayload, CronSchedule
|
from nanobot.cron.types import CronJob, CronPayload, CronSchedule
|
||||||
from nanobot.runtime_context import RUNTIME_CONTEXT_INPUT_META
|
|
||||||
|
|
||||||
|
|
||||||
async def _wait_until(predicate, *, timeout: float = 1.0, interval: float = 0.01) -> None:
|
async def _wait_until(predicate, *, timeout: float = 1.0, interval: float = 0.01) -> None:
|
||||||
@@ -293,12 +292,7 @@ def test_load_store_migrates_legacy_delivery_context(tmp_path) -> None:
|
|||||||
"deliver": True,
|
"deliver": True,
|
||||||
"channel": "telegram",
|
"channel": "telegram",
|
||||||
"to": "user-1",
|
"to": "user-1",
|
||||||
"channelMeta": {
|
"channelMeta": {"message_thread_id": 42},
|
||||||
"message_thread_id": 42,
|
|
||||||
RUNTIME_CONTEXT_INPUT_META: [
|
|
||||||
{"source": "webui_quote", "content": "stale quote"}
|
|
||||||
],
|
|
||||||
},
|
|
||||||
"sessionKey": "telegram:user-1:topic:42",
|
"sessionKey": "telegram:user-1:topic:42",
|
||||||
},
|
},
|
||||||
"state": {},
|
"state": {},
|
||||||
@@ -417,39 +411,6 @@ def test_add_job_preserves_origin_delivery_context(tmp_path) -> None:
|
|||||||
assert reloaded.payload.origin_metadata == metadata
|
assert reloaded.payload.origin_metadata == metadata
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_start_heals_runtime_context_from_pending_external_add(tmp_path) -> None:
|
|
||||||
"""Flattened runtime blocks from older action files must not be replayed."""
|
|
||||||
store_path = tmp_path / "cron" / "jobs.json"
|
|
||||||
external = CronService(store_path)
|
|
||||||
job = external.add_job(
|
|
||||||
name="quoted reminder",
|
|
||||||
schedule=CronSchedule(kind="every", every_ms=60_000),
|
|
||||||
message="remember this",
|
|
||||||
origin_metadata={"webui": True},
|
|
||||||
**_bound_chat("quoted"),
|
|
||||||
)
|
|
||||||
|
|
||||||
action_path = tmp_path / "cron" / "action.jsonl"
|
|
||||||
action = json.loads(action_path.read_text(encoding="utf-8"))
|
|
||||||
action["params"]["payload"]["origin_metadata"][RUNTIME_CONTEXT_INPUT_META] = [
|
|
||||||
{"source": "webui_quote", "content": "quoted reply"}
|
|
||||||
]
|
|
||||||
action_path.write_text(json.dumps(action), encoding="utf-8")
|
|
||||||
|
|
||||||
owner = CronService(store_path)
|
|
||||||
await owner.start()
|
|
||||||
try:
|
|
||||||
loaded = owner.get_job(job.id)
|
|
||||||
assert loaded is not None
|
|
||||||
assert loaded.payload.origin_metadata == {"webui": True}
|
|
||||||
|
|
||||||
raw = json.loads(store_path.read_text(encoding="utf-8"))
|
|
||||||
assert raw["jobs"][0]["payload"]["originMetadata"] == {"webui": True}
|
|
||||||
finally:
|
|
||||||
owner.stop()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_channel_meta_and_session_key_survive_store_reload(tmp_path) -> None:
|
async def test_channel_meta_and_session_key_survive_store_reload(tmp_path) -> None:
|
||||||
store_path = tmp_path / "cron" / "jobs.json"
|
store_path = tmp_path / "cron" / "jobs.json"
|
||||||
|
|||||||
@@ -112,49 +112,14 @@ class TestEnforceRoleAlternation:
|
|||||||
assert result[1]["content"] is None
|
assert result[1]["content"] is None
|
||||||
assert result[2]["role"] == "tool"
|
assert result[2]["role"] == "tool"
|
||||||
|
|
||||||
def test_consecutive_user_messages_preserve_text_before_multimodal_content(self):
|
def test_non_string_content_uses_latest(self):
|
||||||
image = {
|
|
||||||
"type": "image_url",
|
|
||||||
"image_url": {"url": "data:image/png;base64,aW1hZ2U="},
|
|
||||||
}
|
|
||||||
msgs = [
|
msgs = [
|
||||||
{"role": "user", "content": "Earlier unanswered question"},
|
{"role": "user", "content": [{"type": "text", "text": "A"}]},
|
||||||
{
|
{"role": "user", "content": "B"},
|
||||||
"role": "user",
|
|
||||||
"content": [image, {"type": "text", "text": "The error is here"}],
|
|
||||||
},
|
|
||||||
]
|
]
|
||||||
result = LLMProvider._enforce_role_alternation(msgs)
|
result = LLMProvider._enforce_role_alternation(msgs)
|
||||||
assert result == [{
|
assert len(result) == 1
|
||||||
"role": "user",
|
assert result[0]["content"] == "B"
|
||||||
"content": [
|
|
||||||
{"type": "text", "text": "Earlier unanswered question"},
|
|
||||||
image,
|
|
||||||
{"type": "text", "text": "The error is here"},
|
|
||||||
],
|
|
||||||
}]
|
|
||||||
|
|
||||||
def test_consecutive_user_messages_preserve_multimodal_content_before_text(self):
|
|
||||||
image = {
|
|
||||||
"type": "image_url",
|
|
||||||
"image_url": {"url": "data:image/png;base64,aW1hZ2U="},
|
|
||||||
}
|
|
||||||
msgs = [
|
|
||||||
{
|
|
||||||
"role": "user",
|
|
||||||
"content": [image, {"type": "text", "text": "First question"}],
|
|
||||||
},
|
|
||||||
{"role": "user", "content": "Follow-up detail"},
|
|
||||||
]
|
|
||||||
result = LLMProvider._enforce_role_alternation(msgs)
|
|
||||||
assert result == [{
|
|
||||||
"role": "user",
|
|
||||||
"content": [
|
|
||||||
image,
|
|
||||||
{"type": "text", "text": "First question"},
|
|
||||||
{"type": "text", "text": "Follow-up detail"},
|
|
||||||
],
|
|
||||||
}]
|
|
||||||
|
|
||||||
def test_original_messages_not_mutated(self):
|
def test_original_messages_not_mutated(self):
|
||||||
msgs = [
|
msgs = [
|
||||||
|
|||||||
@@ -141,13 +141,14 @@ def test_internal_continuation_requires_budget_boundary_and_queue():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_save_skip_matches_prefix_when_current_message_was_persisted():
|
def test_save_skip_matches_prefix_when_current_message_merged():
|
||||||
skip = _save_skip_for_turn(
|
skip = _save_skip_for_turn(
|
||||||
message_metadata=None,
|
message_metadata=None,
|
||||||
initial_message_count=3, # [system, history user, current user]
|
initial_message_count=2, # [system, merged user]
|
||||||
|
history_count=1,
|
||||||
input_persisted_early=True,
|
input_persisted_early=True,
|
||||||
)
|
)
|
||||||
assert skip == 3
|
assert skip == 2
|
||||||
|
|
||||||
|
|
||||||
def test_save_skip_unchanged_for_standalone_current_message():
|
def test_save_skip_unchanged_for_standalone_current_message():
|
||||||
@@ -155,10 +156,12 @@ def test_save_skip_unchanged_for_standalone_current_message():
|
|||||||
assert _save_skip_for_turn(
|
assert _save_skip_for_turn(
|
||||||
message_metadata=None,
|
message_metadata=None,
|
||||||
initial_message_count=3,
|
initial_message_count=3,
|
||||||
|
history_count=1,
|
||||||
input_persisted_early=True,
|
input_persisted_early=True,
|
||||||
) == 3
|
) == 3
|
||||||
assert _save_skip_for_turn(
|
assert _save_skip_for_turn(
|
||||||
message_metadata=None,
|
message_metadata=None,
|
||||||
initial_message_count=3,
|
initial_message_count=3,
|
||||||
|
history_count=1,
|
||||||
input_persisted_early=False,
|
input_persisted_early=False,
|
||||||
) == 2
|
) == 2
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -12,7 +11,6 @@ from nanobot.agent.tools.message import MessageTool
|
|||||||
from nanobot.agent.tools.spawn import SpawnTool
|
from nanobot.agent.tools.spawn import SpawnTool
|
||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
from nanobot.providers.base import GenerationSettings, LLMProvider
|
from nanobot.providers.base import GenerationSettings, LLMProvider
|
||||||
from nanobot.runtime_context import RUNTIME_CONTEXT_INPUT_META, RuntimeContextBlock
|
|
||||||
from nanobot.session.keys import UNIFIED_SESSION_KEY
|
from nanobot.session.keys import UNIFIED_SESSION_KEY
|
||||||
from nanobot.utils.llm_runtime import LLMRuntime
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
|
|
||||||
@@ -301,41 +299,6 @@ async def test_webui_cron_tool_uses_origin_session_when_unified_enabled(tmp_path
|
|||||||
assert jobs[0].payload.origin_metadata == {"webui": True}
|
assert jobs[0].payload.origin_metadata == {"webui": True}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_cron_tool_snapshots_only_persistable_request_metadata(tmp_path) -> None:
|
|
||||||
"""Live runtime context must not poison a persisted WebUI cron job."""
|
|
||||||
store_path = tmp_path / "jobs.json"
|
|
||||||
service = CronService(store_path)
|
|
||||||
tool = CronTool(service)
|
|
||||||
await service.start()
|
|
||||||
try:
|
|
||||||
with request_context(
|
|
||||||
RequestContext(
|
|
||||||
channel="websocket",
|
|
||||||
chat_id="chat-123",
|
|
||||||
metadata={
|
|
||||||
"webui": True,
|
|
||||||
RUNTIME_CONTEXT_INPUT_META: [
|
|
||||||
RuntimeContextBlock(source="webui_quote", content="quoted reply")
|
|
||||||
],
|
|
||||||
"opaque": object(),
|
|
||||||
},
|
|
||||||
session_key=UNIFIED_SESSION_KEY,
|
|
||||||
)
|
|
||||||
):
|
|
||||||
result = await tool.execute(action="add", message="standup", every_seconds=300)
|
|
||||||
|
|
||||||
assert result.startswith("Created job")
|
|
||||||
jobs = service.list_jobs()
|
|
||||||
assert len(jobs) == 1
|
|
||||||
assert jobs[0].payload.origin_metadata == {"webui": True}
|
|
||||||
|
|
||||||
raw = json.loads(store_path.read_text(encoding="utf-8"))
|
|
||||||
assert raw["jobs"][0]["payload"]["originMetadata"] == {"webui": True}
|
|
||||||
finally:
|
|
||||||
service.stop()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cron_tool_preserves_thread_scoped_session_key(tmp_path) -> None:
|
async def test_cron_tool_preserves_thread_scoped_session_key(tmp_path) -> None:
|
||||||
"""Channel-provided thread session keys should remain the cron owner."""
|
"""Channel-provided thread session keys should remain the cron owner."""
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.context import TranscriptInput
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.tools.message import MessageTool
|
from nanobot.agent.tools.message import MessageTool
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
@@ -179,9 +178,7 @@ class TestMessageToolSuppressLogic:
|
|||||||
progress.append((content, tool_hint))
|
progress.append((content, tool_hint))
|
||||||
|
|
||||||
result = await loop._run_agent_loop(
|
result = await loop._run_agent_loop(
|
||||||
TranscriptInput(history=[], current_message=None),
|
[], runtime=loop.llm_runtime(), on_progress=on_progress
|
||||||
runtime=loop.llm_runtime(),
|
|
||||||
on_progress=on_progress,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result.final_content == "Done"
|
assert result.final_content == "Done"
|
||||||
|
|||||||
@@ -183,66 +183,6 @@ async def test_rate_limit_is_per_source_session_and_uses_a_rolling_minute(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_rate_limit_releases_expired_source_state_and_keeps_recent_sources(
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
sessions = SessionManager(tmp_path)
|
|
||||||
_persist(
|
|
||||||
sessions,
|
|
||||||
"websocket:a",
|
|
||||||
"websocket:b",
|
|
||||||
"websocket:c",
|
|
||||||
"websocket:target",
|
|
||||||
)
|
|
||||||
now = 0.0
|
|
||||||
tool = SendSessionMessageTool(
|
|
||||||
sessions=sessions,
|
|
||||||
bus=MessageBus(),
|
|
||||||
max_messages_per_minute=2,
|
|
||||||
clock=lambda: now,
|
|
||||||
)
|
|
||||||
target = _handle(sessions, "websocket:target").name
|
|
||||||
|
|
||||||
for source in ("websocket:a", "websocket:b"):
|
|
||||||
await tool.enqueue(
|
|
||||||
source_session_key=source,
|
|
||||||
target_handle=target,
|
|
||||||
content="initial",
|
|
||||||
expect_reply=False,
|
|
||||||
)
|
|
||||||
now = 30.0
|
|
||||||
await tool.enqueue(
|
|
||||||
source_session_key="websocket:a",
|
|
||||||
target_handle=target,
|
|
||||||
content="recent",
|
|
||||||
expect_reply=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
now = 61.0
|
|
||||||
await tool.enqueue(
|
|
||||||
source_session_key="websocket:c",
|
|
||||||
target_handle=target,
|
|
||||||
content="trigger cleanup",
|
|
||||||
expect_reply=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert set(tool._sent_at) == {"websocket:a", "websocket:c"}
|
|
||||||
await tool.enqueue(
|
|
||||||
source_session_key="websocket:a",
|
|
||||||
target_handle=target,
|
|
||||||
content="within rolling window",
|
|
||||||
expect_reply=False,
|
|
||||||
)
|
|
||||||
with pytest.raises(SessionMessageError, match="rate limit"):
|
|
||||||
await tool.enqueue(
|
|
||||||
source_session_key="websocket:a",
|
|
||||||
target_handle=target,
|
|
||||||
content="over limit",
|
|
||||||
expect_reply=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_reply_timeout_injects_a_user_input_back_into_the_source(
|
async def test_reply_timeout_injects_a_user_input_back_into_the_source(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
|
|||||||
+1
-2
@@ -94,7 +94,7 @@ import {
|
|||||||
type FooterMode,
|
type FooterMode,
|
||||||
type FooterHintTheme,
|
type FooterHintTheme,
|
||||||
} from "./footer-hints"
|
} from "./footer-hints"
|
||||||
import { configureOpenTuiEnvironment, createTuiHost, type TuiHost } from "./host"
|
import { createTuiHost, type TuiHost } from "./host"
|
||||||
|
|
||||||
interface AppOptions {
|
interface AppOptions {
|
||||||
wsUrl?: string
|
wsUrl?: string
|
||||||
@@ -820,7 +820,6 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
|
|
||||||
static async create(options: AppOptions): Promise<NanobotTui> {
|
static async create(options: AppOptions): Promise<NanobotTui> {
|
||||||
configureOpenTuiEnvironment()
|
|
||||||
const host = createTuiHost()
|
const host = createTuiHost()
|
||||||
const renderer = await createCliRenderer({
|
const renderer = await createCliRenderer({
|
||||||
targetFps: 30,
|
targetFps: 30,
|
||||||
|
|||||||
+1
-25
@@ -1,9 +1,6 @@
|
|||||||
import { describe, expect, test } from "bun:test"
|
import { describe, expect, test } from "bun:test"
|
||||||
|
|
||||||
import {
|
import { createTuiHost } from "./host"
|
||||||
configureOpenTuiEnvironment,
|
|
||||||
createTuiHost,
|
|
||||||
} from "./host"
|
|
||||||
|
|
||||||
async function settle(): Promise<void> {
|
async function settle(): Promise<void> {
|
||||||
await Bun.sleep(0)
|
await Bun.sleep(0)
|
||||||
@@ -11,27 +8,6 @@ async function settle(): Promise<void> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
describe("TUI host integration", () => {
|
describe("TUI host integration", () => {
|
||||||
test("disables the explicit-width probe on Windows", () => {
|
|
||||||
const environment: Record<string, string | undefined> = {}
|
|
||||||
|
|
||||||
configureOpenTuiEnvironment(environment, "win32")
|
|
||||||
|
|
||||||
expect(environment.OPENTUI_FORCE_EXPLICIT_WIDTH).toBe("false")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("preserves explicit probe choices and leaves other platforms unchanged", () => {
|
|
||||||
const overridden = {
|
|
||||||
OPENTUI_FORCE_EXPLICIT_WIDTH: "true",
|
|
||||||
}
|
|
||||||
const nonWindows: Record<string, string | undefined> = {}
|
|
||||||
|
|
||||||
configureOpenTuiEnvironment(overridden, "win32")
|
|
||||||
configureOpenTuiEnvironment(nonWindows, "linux")
|
|
||||||
|
|
||||||
expect(overridden.OPENTUI_FORCE_EXPLICIT_WIDTH).toBe("true")
|
|
||||||
expect(nonWindows.OPENTUI_FORCE_EXPLICIT_WIDTH).toBeUndefined()
|
|
||||||
})
|
|
||||||
|
|
||||||
test("standalone terminals remain a no-op", async () => {
|
test("standalone terminals remain a no-op", async () => {
|
||||||
const commands: string[][] = []
|
const commands: string[][] = []
|
||||||
const host = createTuiHost({}, async (command) => { commands.push([...command]) })
|
const host = createTuiHost({}, async (command) => { commands.push([...command]) })
|
||||||
|
|||||||
@@ -8,19 +8,6 @@ type CommandRunner = (command: readonly string[]) => Promise<void>
|
|||||||
|
|
||||||
const METADATA_SOURCE = "nanobot:tui:metadata"
|
const METADATA_SOURCE = "nanobot:tui:metadata"
|
||||||
|
|
||||||
export function configureOpenTuiEnvironment(
|
|
||||||
environment: Environment = process.env,
|
|
||||||
platform = process.platform,
|
|
||||||
): void {
|
|
||||||
if (platform !== "win32") return
|
|
||||||
|
|
||||||
// OpenTUI probes OSC 66 support on the main screen before its renderer is
|
|
||||||
// active. Some Windows terminal hosts do not restore the cursor around that
|
|
||||||
// probe, so shutdown resumes in terminal history instead of below the TUI.
|
|
||||||
// Keep an explicit user choice, but use the safe default on Windows.
|
|
||||||
environment.OPENTUI_FORCE_EXPLICIT_WIDTH ??= "false"
|
|
||||||
}
|
|
||||||
|
|
||||||
class StandaloneHost implements TuiHost {
|
class StandaloneHost implements TuiHost {
|
||||||
reportTitle(): void {}
|
reportTitle(): void {}
|
||||||
release(): void {}
|
release(): void {}
|
||||||
|
|||||||
Reference in New Issue
Block a user