Compare commits

...
Author SHA1 Message Date
chengyongruandchengyongru 3c25f826ea fix(telegram): detect unsupported rich stream edits 2026-08-31 18:39:35 +08:00
Nolanandchengyongru 195e4c281d fix(telegram): preserve final-edit retry contract in rich stream upgrade
Address review feedback on the rich edit error classification:

- Transport, rate-limit, and unexpected errors now propagate instead of
  returning False. Returning False made send_delta fall through to an
  immediate legacy edit_message_text, which under connection-pool
  exhaustion doubled demand and discarded the buffered retry state that
  ChannelManager relies on.
- 'Message is not modified' is treated as success: when the rich edit is
  applied server-side but its response times out, the retry inside
  _call_with_retry reports the edit as already applied, and the previous
  BadRequest branch would have let the legacy edit overwrite the
  successful rich result.
- False is now returned only for capability errors (pre-10.1 servers,
  which still trip the rich latch) and content-shaped rejections, where
  the legacy HTML path is the intended fallback.

Adds regression coverage for a direct NetworkError (propagates, buffer
kept for manager retry) and TimedOut followed by Message-is-not-modified
(treated as success, no legacy overwrite).
2026-08-31 18:39:35 +08:00
Nolanandchengyongru 37663ac947 fix(telegram): upgrade streaming preview to rich in place at stream end
The rich branch in send_delta(stream_end=True) was unreachable: it was
guarded by 'not buf.message_id' after an earlier return had already
ensured buf.message_id is set, so sendRichMessage never fired with
streaming enabled and the final message always went through the legacy
HTML editMessageText path.

Bot API 10.1 added a rich_message parameter to editMessageText, which
upgrades an existing message to rich in place. Use it at stream end via
do_api_request so the streaming preview keeps its identity — no
delete-and-resend, so none of the flickering or dropped line breaks
that made rich-at-stream-end fail in #4470.

Capability errors (server older than 10.1) trip the existing
_rich_send_disabled latch and fall back to the legacy HTML path, which
is unchanged.

Fixes #5516
2026-08-31 18:39:35 +08:00
chengyongruandGitHub e111b83af6 refactor(agent): unify runner request fitting (#5612)
* refactor(agent): unify runner request fitting

* fix(agent): fit every runner model request

* fix(agent): count resumed state during request fitting

* refactor(agent): consolidate request fitting state
2026-08-31 18:08:21 +08:00
chengyongruandchengyongru 6d6d58d329 style(tui): simplify runtime header controls 2026-08-31 17:03:24 +08:00
chengyongruandchengyongru e69159cdae style(tui): replace header dots with spacing 2026-08-31 17:03:24 +08:00
feb33e1f99 docs(tools): clarify edit_file selector exclusivity (#5598)
* docs(tools): clarify edit selector exclusivity

* docs(tools): streamline edit_file guidance

---------

Co-authored-by: chengyongru <chengyongru.ai@gmail.com>
2026-08-31 14:28:41 +08:00
chengyongruandGitHub bb34b58f47 refactor(agent): make memory summaries cumulative (#5610)
* refactor(agent): make memory summaries cumulative

Treat the latest session summary as a replacement checkpoint, preserve it through bounded raw fallbacks, and reserve history.jsonl for Dream ingestion.

* fix(agent): preserve cumulative checkpoint context

* fix(agent): preserve memory archive prompt cache

* refactor(agent): state archive prompt positively

* refactor(agent): remove checkpoint version migration

* refactor(agent): summarize full archive context

* test(agent): align cumulative archive prompt assertion

* refactor(agent): clarify memory checkpoint contract
2026-08-31 13:37:20 +08:00
chengyongruandGitHub 6cd7063682 refactor(agent): defer transcript assembly to runner (#5608)
* refactor(agent): defer transcript assembly to runner

Keep persisted history and the fresh turn as explicit inputs until the Runner assembles the provider transcript. Preserve ContextBuilder and direct AgentRunner compatibility while making the save boundary structural.

Refs NAN-81.

* fix(providers): preserve mixed adjacent user content
2026-08-31 00:06:15 +08:00
Xubin Ren d019658501 fix(dingtalk): stop late inbound task creation 2026-08-30 17:56:34 +08:00
yu-xin-candXubin Ren e8385d9257 fix(dingtalk): drain inbound background tasks 2026-08-30 17:56:34 +08:00
Xubin Ren 5c71ef6e49 test(email): assert header-first fetch contract 2026-08-30 17:38:44 +08:00
Till AdamandXubin Ren f573ecfe56 perf(email): fetch headers before body, use UID SEARCH to skip re-fetch
The IMAP poll loop previously downloaded the entire message body for
every UNSEEN message before running any filter (self-sent, SPF/DKIM,
allow-list), and only learned the UID by parsing it back out of that
fetch response. Rejected messages stay unseen and get re-fetched in
full on every subsequent poll.

Switch to UID SEARCH (UIDs come back directly, no per-message fetch
needed to learn them) so already-processed UIDs are skipped before any
network fetch, then fetch headers only to evaluate every filter — the
full body is downloaded only for messages that pass every check and
are actually delivered. No behavior change: filter outcomes and
\Seen semantics are unchanged.
2026-08-30 17:38:44 +08:00
Xubin Ren 1ac1b35c84 fix(cron): sanitize migrated legacy metadata 2026-08-30 17:24:05 +08:00
Oxygen56andXubin Ren 679a07460e fix(cron): sanitize persisted origin metadata 2026-08-30 17:24:05 +08:00
Oxygen56andXubin Ren 919e3d341e fix(agent): add retry hint to tool exceptions 2026-08-30 17:03:51 +08:00
yu-xin-candXubin Ren 2c55934198 fix(agent): bound session message rate-limit state 2026-08-30 16:46:36 +08:00
Xubin Ren 5afdffff51 fix(agent): settle reasoning close before cancellation 2026-08-30 16:32:08 +08:00
KDBandXubin Ren bfe041def7 fix(agent): close reasoning stream on cancellation 2026-08-30 16:32:08 +08:00
Kail Tianandchengyongru 1c1b13a3a9 fix(tui): preserve cursor position on Windows exit 2026-08-30 12:17:40 +08:00
Xubin Ren 2c87143f77 fix(cli): preserve gateway log stream boundaries 2026-08-29 23:03:23 +08:00
Xubin Ren d7df2726de fix(cli): stream gateway logs in WebUI launcher 2026-08-29 23:03:23 +08:00
61 changed files with 2791 additions and 1446 deletions
+1 -1
View File
@@ -139,7 +139,7 @@ Interactive mode exits with `exit`, `quit`, `/exit`, `/quit`, `:q`, or `Ctrl+D`.
| Command | Description | | Command | Description |
|---|---| |---|---|
| `nanobot webui` | Create config/workspace if needed, enable the local WebUI channel after confirmation, start the gateway, and open `http://127.0.0.1:8765` | | `nanobot webui` | Create config/workspace if needed, enable the local WebUI channel after confirmation, start the gateway, open `http://127.0.0.1:8765`, and follow new gateway logs |
| `nanobot webui --background` | Deprecated; prints the equivalent explicit `nanobot gateway --background` command and exits | | `nanobot webui --background` | Deprecated; prints the equivalent explicit `nanobot gateway --background` command and exits |
| `nanobot webui --dev` | Start the gateway and Vite together at `http://127.0.0.1:5173`, with live frontend updates | | `nanobot webui --dev` | Start the gateway and Vite together at `http://127.0.0.1:5173`, with live frontend updates |
| `nanobot webui --no-open` | Prepare and start the WebUI without opening a browser | | `nanobot webui --no-open` | Prepare and start the WebUI without opening a browser |
+4 -8
View File
@@ -2268,16 +2268,12 @@ When a user is idle for longer than a configured threshold, nanobot **proactivel
How it works: How it works:
1. **Idle detection**: On each idle tick (~1 s), checks whether an idle-session scan is due. By default, the full scan runs at most once per minute. 1. **Idle detection**: On each idle tick (~1 s), checks whether an idle-session scan is due. By default, the full scan runs at most once per minute.
2. **Background compaction**: Idle sessions summarize the older live prefix via LLM and keep the most recent legal suffix (currently 8 messages). 2. **Background compaction**: Older context is summarized while the most recent messages remain available.
3. **Summary injection**: When the user returns, the summary is injected as runtime context (one-shot, not persisted) alongside the retained recent suffix. 3. **Session preservation**: The complete session history remains stored for later inspection and reuse.
4. **Restart-safe resume**: The summary is also mirrored into session metadata so it can still be recovered after a process restart. 4. **Restart-safe resume**: The compacted context remains available after a process restart.
> [!NOTE] > [!NOTE]
> Mental model: "summarize older context, keep the freshest live turns, **and overwrite the session file with the compact form.**" It is not a full `session.clear()`, but it is a write — not a soft cursor move. > Auto compact shortens the context sent to the model without deleting the session's structured message history.
>
> Concretely, auto compact rewrites `sessions/<key>.jsonl` in place: older messages (including their structured `tool_calls` / `tool_call_id` / `reasoning_content`) are replaced by just the retained recent suffix (currently 8 messages), while the archived prefix is preserved only as a plain-text summary appended to `memory/history.jsonl` (or a `[RAW] ...` flattened dump if LLM summarization fails). The original structured JSON of those turns is no longer recoverable from the session file.
>
> This differs from the **token-driven soft consolidation** that fires when a prompt exceeds the context budget: that path only advances an internal `last_consolidated` cursor and leaves the session file untouched, so the raw tool-call trail stays on disk and can still be replayed or audited. If you rely on that trail for debugging or auditing, set `idleCompactAfterMinutes` to `0` and let only the token-driven path run.
## Timezone ## Timezone
+3 -1
View File
@@ -23,7 +23,9 @@ one is missing, starts or joins the same on-demand gateway used by the native
TUI, and opens the browser. With a fresh config, TUI, and opens the browser. With a fresh config,
it can open before a model is configured so you can finish setup in **Settings it can open before a model is configured so you can finish setup in **Settings
→ Models**. The first-run path binds the WebUI to `127.0.0.1` by default, so → Models**. The first-run path binds the WebUI to `127.0.0.1` by default, so
it is not available from other devices on your LAN. it is not available from other devices on your LAN. While the launcher remains
attached, it mirrors new log output from that exact gateway instance in the
terminal without replaying older logs.
After model setup, explicitly promote the shared gateway when you do not want to keep a client open: After model setup, explicitly promote the shared gateway when you do not want to keep a client open:
+67 -79
View File
@@ -30,11 +30,7 @@ from nanobot.security.workspace_access import WorkspaceScopeResolver
from nanobot.session.keys import last_channel_from_metadata from nanobot.session.keys import last_channel_from_metadata
from nanobot.session.manager import Session from nanobot.session.manager import Session
from nanobot.session.summary import SessionSummary from nanobot.session.summary import SessionSummary
from nanobot.utils.helpers import ( from nanobot.utils.helpers import detect_image_mime, load_bundled_template
detect_image_mime,
load_bundled_template,
truncate_text_to_tokens,
)
from nanobot.utils.prompt_templates import render_template from nanobot.utils.prompt_templates import render_template
@@ -75,14 +71,29 @@ 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."""
BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md"] BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md"]
_SKIPPABLE_DEFAULTS = {"AGENTS.md", "USER.md"} _SKIPPABLE_DEFAULTS = {"AGENTS.md", "USER.md"}
_RUNTIME_CONTEXT_TAG = RUNTIME_CONTEXT_TAG _RUNTIME_CONTEXT_TAG = RUNTIME_CONTEXT_TAG
_MAX_RECENT_HISTORY = 50
_MAX_HISTORY_TOKENS = 8_000 # hard cap on recent history section size (tokens)
_RUNTIME_CONTEXT_END = RUNTIME_CONTEXT_END _RUNTIME_CONTEXT_END = RUNTIME_CONTEXT_END
def __init__(self, workspace: Path, timezone: str | None = None, disabled_skills: list[str] | None = None): def __init__(self, workspace: Path, timezone: str | None = None, disabled_skills: list[str] | None = None):
@@ -98,9 +109,6 @@ class ContextBuilder:
session_summary: SessionSummary | None = None, session_summary: SessionSummary | None = None,
workspace: Path | None = None, workspace: Path | None = None,
include_memory: bool = True, include_memory: bool = True,
include_memory_recent_history: bool = True,
session_key: str | None = None,
unified_session: bool = False,
) -> str: ) -> str:
"""Build the system prompt from identity, bootstrap files, memory, and skills.""" """Build the system prompt from identity, bootstrap files, memory, and skills."""
root = workspace or self.workspace root = workspace or self.workspace
@@ -138,29 +146,6 @@ class ContextBuilder:
if skills_summary: if skills_summary:
parts.append(render_template("agent/skills_section.md", skills_summary=skills_summary)) parts.append(render_template("agent/skills_section.md", skills_summary=skills_summary))
if include_memory_recent_history:
entries = self.memory.read_recent_history_for_prompt(
since_cursor=self.memory.get_last_dream_cursor(),
session_key=session_key,
unified_session=unified_session,
)
if entries:
capped = entries[-self._MAX_RECENT_HISTORY:]
capped = self._without_duplicate_session_summary(
capped,
session_key=session_key,
session_summary=session_summary,
)
if capped:
history_text = "\n".join(
f"- [{e['timestamp']}] {e['content']}" for e in capped
)
history_text = truncate_text_to_tokens(
history_text,
self._MAX_HISTORY_TOKENS,
)
parts.append("# Recent History\n\n" + history_text)
if session_summary: if session_summary:
parts.append( parts.append(
"[Archived Context Summary]\n\n" "[Archived Context Summary]\n\n"
@@ -170,25 +155,6 @@ class ContextBuilder:
return "\n\n---\n\n".join(parts) return "\n\n---\n\n".join(parts)
@staticmethod
def _without_duplicate_session_summary(
entries: list[dict[str, Any]],
*,
session_key: str | None,
session_summary: SessionSummary | None,
) -> list[dict[str, Any]]:
"""Drop the history entry already represented by the session summary."""
if not session_summary:
return entries
for index in range(len(entries) - 1, -1, -1):
entry = entries[index]
if (
entry.get("session_key") == session_key
and entry.get("content") == session_summary["text"]
):
return [*entries[:index], *entries[index + 1:]]
return entries
def _get_identity(self, channel: str | None = None, workspace: Path | None = None) -> str: def _get_identity(self, channel: str | None = None, workspace: Path | None = None) -> str:
"""Get the core identity section.""" """Get the core identity section."""
root = workspace or self.workspace root = workspace or self.workspace
@@ -278,46 +244,68 @@ class ContextBuilder:
runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None, runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None,
workspace: Path | None = None, workspace: Path | None = None,
include_memory: bool = True, include_memory: bool = True,
include_memory_recent_history: bool = True,
session_key: str | None = None,
unified_session: bool = False,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""Build the complete message list for an LLM call.""" """Compatibility wrapper for callers that need merged adjacent roles."""
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,
)
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,
) -> 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=session_summary, session_summary=transcript.session_summary,
workspace=root, workspace=root,
include_memory=include_memory, include_memory=include_memory,
include_memory_recent_history=include_memory_recent_history,
session_key=session_key,
unified_session=unified_session,
), ),
}, },
*history, *transcript.history,
] ]
current = self.build_current_message( if transcript.current_message is None:
current_message,
media=media,
current_role=current_role,
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 return messages
current = self.build_current_message(
transcript.current_message,
media=list(transcript.media) if transcript.media else None,
current_role=transcript.current_role,
runtime_context_blocks=transcript.runtime_context_blocks,
)
messages.append(current) messages.append(current)
return messages return messages
+101 -134
View File
@@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, cast
from loguru import logger from loguru import logger
from nanobot.providers.base import LLMUsage
from nanobot.utils.helpers import ( from nanobot.utils.helpers import (
estimate_message_tokens, estimate_message_tokens,
estimate_prompt_tokens_chain, estimate_prompt_tokens_chain,
@@ -27,12 +28,6 @@ if TYPE_CHECKING:
from nanobot.providers.base import LLMProvider from nanobot.providers.base import LLMProvider
SNIP_SAFETY_BUFFER = 1024 SNIP_SAFETY_BUFFER = 1024
MICROCOMPACT_MIN_CHARS = 500
INFLIGHT_COMPACT_TARGET_RATIO = 0.85
COMPACTABLE_TOOLS = frozenset({
"read_file", "exec", "grep", "find_files",
"web_search", "web_fetch", "list_dir", "list_exec_sessions",
})
# read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops. # read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops.
TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"}) TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"})
BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]" BACKFILL_CONTENT = "[Tool result unavailable — call was interrupted or lost]"
@@ -41,6 +36,27 @@ PLACEHOLDER_TEXTS = frozenset({
}) })
class ContextWindowExceededError(RuntimeError):
"""Raised before a locally fitted request that still exceeds its budget."""
def __init__(
self,
*,
session_key: str | None,
estimated_tokens: int,
input_budget: int,
source: str,
) -> None:
self.session_key = session_key
self.estimated_tokens = estimated_tokens
self.input_budget = input_budget
self.source = source
super().__init__(
"Model input still exceeds the local context budget after request fitting "
f"for {session_key or 'default'}: {estimated_tokens}/{input_budget} via {source}"
)
def _tool_call_name_is_valid(tool_call: Any) -> bool: def _tool_call_name_is_valid(tool_call: Any) -> bool:
"""Whether a persisted OpenAI-style tool_call carries a usable name. """Whether a persisted OpenAI-style tool_call carries a usable name.
@@ -67,7 +83,6 @@ class ContextGovernanceConfig:
context_window_tokens: int | None = None context_window_tokens: int | None = None
context_block_limit: int | None = None context_block_limit: int | None = None
max_tokens: int | None = None max_tokens: int | None = None
inflight_start_index: int = 0
class ContextGovernor: class ContextGovernor:
@@ -77,17 +92,85 @@ class ContextGovernor:
self, self,
config: ContextGovernanceConfig, config: ContextGovernanceConfig,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
updated = self.strip_placeholder_assistant_messages(messages) updated = self.strip_placeholder_assistant_messages(messages)
updated = self.strip_malformed_tool_calls(updated) updated = self.strip_malformed_tool_calls(updated)
updated = self.drop_orphan_tool_results(updated) updated = self.drop_orphan_tool_results(updated)
updated = self.backfill_missing_tool_results(updated) updated = self.backfill_missing_tool_results(updated)
updated = self.apply_tool_result_budget(config, updated) return self.apply_tool_result_budget(config, updated)
updated = self.compact_inflight_overflow(config, updated, compacted_tool_call_ids)
updated = self.snip_history(config, updated) def fit_to_budget(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Fit a model-facing copy while keeping the source transcript intact."""
updated = self.snip_history(
config,
messages,
tool_definitions=tool_definitions,
force=True,
)
updated = self.drop_orphan_tool_results(updated) updated = self.drop_orphan_tool_results(updated)
return self.backfill_missing_tool_results(updated) updated = self.backfill_missing_tool_results(updated)
if not config.context_window_tokens:
return updated
budget = self.input_budget(config)
estimated, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tool_definitions,
)
if budget > 0 and estimated <= budget:
return updated
raise ContextWindowExceededError(
session_key=config.session_key,
estimated_tokens=estimated,
input_budget=budget,
source=source,
)
def fit_request(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
usage: LLMUsage | None,
*,
usage_matches_messages: bool,
tool_definitions: list[dict[str, Any]] | None,
request_context_tokens: int | None = None,
) -> tuple[list[dict[str, Any]], bool]:
"""Fit the request when its measured or estimated input is pressured."""
if not config.context_window_tokens:
return messages, False
budget = self.input_budget(config)
if (
request_context_tokens is None
and usage_matches_messages
and usage is not None
and usage.context_tokens is not None
):
pressured = budget <= 0 or usage.context_tokens >= budget
else:
estimated, _ = estimate_prompt_tokens_chain(
config.provider,
config.model,
messages,
tool_definitions,
)
if request_context_tokens is not None:
estimated = max(estimated, request_context_tokens)
pressured = budget <= 0 or estimated >= budget
if not pressured:
return messages, False
return self.fit_to_budget(
config,
messages,
tool_definitions=tool_definitions,
), True
@staticmethod @staticmethod
def input_budget(config: ContextGovernanceConfig) -> int: def input_budget(config: ContextGovernanceConfig) -> int:
@@ -326,71 +409,13 @@ class ContextGovernor:
updated[idx]["content"] = normalized updated[idx]["content"] = normalized
return updated return updated
def compact_inflight_overflow(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]:
"""Compact in-flight tool results only when the request would overflow."""
budget = self.input_budget(config)
if budget <= 0:
return messages
tools = config.tools.get_definitions()
updated = self._apply_recorded_compactions(messages, compacted_tool_call_ids)
estimate, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tools,
)
if estimate <= budget:
return updated
target = int(budget * INFLIGHT_COMPACT_TARGET_RATIO)
candidates = self._inflight_compaction_candidates(
config,
updated,
compacted_tool_call_ids,
)
if not candidates:
return updated
for candidate_idx, (idx, tool_call_id) in enumerate(candidates):
is_newest_candidate = candidate_idx == len(candidates) - 1
if is_newest_candidate and estimate <= budget:
break
if tool_call_id in compacted_tool_call_ids:
continue
if updated is messages:
updated = [dict(m) for m in messages]
compacted_tool_call_ids.add(tool_call_id)
self._compact_tool_result_at(updated, idx)
estimate, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
updated,
tools,
)
if estimate <= target:
break
logger.debug(
"In-flight context compaction for {}: prompt={} budget={} target={} via {}, ids={}",
config.session_key or "default",
estimate,
budget,
target,
source,
len(compacted_tool_call_ids),
)
return updated
def snip_history( def snip_history(
self, self,
config: ContextGovernanceConfig, config: ContextGovernanceConfig,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
force: bool = False,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
if not messages or not config.context_window_tokens: if not messages or not config.context_window_tokens:
return messages return messages
@@ -399,14 +424,13 @@ class ContextGovernor:
if budget <= 0: if budget <= 0:
return messages return messages
tools = config.tools.get_definitions()
estimate, _ = estimate_prompt_tokens_chain( estimate, _ = estimate_prompt_tokens_chain(
config.provider, config.provider,
config.model, config.model,
messages, messages,
tools, tool_definitions,
) )
if estimate <= budget: if not force and estimate <= budget:
return messages return messages
system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"] system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"]
@@ -419,7 +443,7 @@ class ContextGovernor:
config.provider, config.provider,
config.model, config.model,
system_messages, system_messages,
tools, tool_definitions,
) )
remaining_budget = max(0, budget - max(system_tokens, fixed_tokens)) remaining_budget = max(0, budget - max(system_tokens, fixed_tokens))
kept: list[dict[str, Any]] = [] kept: list[dict[str, Any]] = []
@@ -434,16 +458,6 @@ class ContextGovernor:
return system_messages + self._legal_history_tail(kept, non_system) return system_messages + self._legal_history_tail(kept, non_system)
@staticmethod
def _tool_result_compaction_message(message: dict[str, Any]) -> str:
name = message.get("name", "tool")
return (
f"Error: The previous {name} result was compacted to fit context because it was too "
"large. Do not repeat the same call unchanged. Retry with a narrower path, query, "
"range, or result limit, use another tool, or tell the user the task cannot fit in "
"the available context."
)
def _legal_history_tail( def _legal_history_tail(
self, self,
kept: list[dict[str, Any]], kept: list[dict[str, Any]],
@@ -462,50 +476,3 @@ class ContextGovernor:
if messages[idx].get("role") == "user": if messages[idx].get("role") == "user":
return messages[idx:] return messages[idx:]
return [] return []
def _apply_recorded_compactions(
self,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[dict[str, Any]]:
if not compacted_tool_call_ids:
return messages
updated = messages
for idx, msg in enumerate(messages):
if msg.get("role") != "tool":
continue
tool_call_id = msg.get("tool_call_id")
if not tool_call_id or str(tool_call_id) not in compacted_tool_call_ids:
continue
compaction_message = self._tool_result_compaction_message(msg)
if msg.get("content") == compaction_message:
continue
if updated is messages:
updated = [dict(m) for m in messages]
updated[idx]["content"] = compaction_message
return updated
def _inflight_compaction_candidates(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
compacted_tool_call_ids: set[str],
) -> list[tuple[int, str]]:
compactable: list[tuple[int, str]] = []
for idx, msg in enumerate(messages):
if idx < config.inflight_start_index:
continue
if msg.get("role") != "tool" or msg.get("name") not in COMPACTABLE_TOOLS:
continue
tool_call_id = msg.get("tool_call_id")
if not tool_call_id or str(tool_call_id) in compacted_tool_call_ids:
continue
content = msg.get("content")
if not isinstance(content, str) or len(content) < MICROCOMPACT_MIN_CHARS:
continue
compactable.append((idx, str(tool_call_id)))
return compactable
def _compact_tool_result_at(self, messages: list[dict[str, Any]], idx: int) -> None:
messages[idx]["content"] = self._tool_result_compaction_message(messages[idx])
+26 -17
View File
@@ -14,6 +14,7 @@ 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
@@ -23,7 +24,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 from nanobot.agent.context import ContextBuilder, PersistedPromptContextResolver, TranscriptInput
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
@@ -135,7 +136,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)
initial_messages: list[dict[str, Any]] = field(default_factory=list) transcript_input: TranscriptInput | None = None
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)
@@ -443,7 +444,6 @@ class AgentLoop:
workspace_scopes=self.workspace_scopes, workspace_scopes=self.workspace_scopes,
unified_session=unified_session, unified_session=unified_session,
), ),
unified_session=unified_session,
) )
self.auto_compact = AutoCompact( self.auto_compact = AutoCompact(
sessions=self.sessions, sessions=self.sessions,
@@ -723,22 +723,15 @@ class AgentLoop:
return True return True
return False return False
def _build_initial_messages(self, ctx: TurnContext) -> list[dict[str, Any]]: def _build_transcript_input(self, ctx: TurnContext) -> TranscriptInput:
"""Build the initial message list for the LLM turn.""" """Capture the persisted history and fresh input as separate transcript parts."""
assert ctx.session is not None assert ctx.session is not None
scope = self.workspace_scopes.for_message(ctx.msg, ctx.session.metadata) return TranscriptInput(
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:
@@ -929,7 +922,7 @@ class AgentLoop:
async def _run_agent_loop( async def _run_agent_loop(
self, self,
initial_messages: list[dict[str, Any]], transcript_input: TranscriptInput,
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,
@@ -1110,6 +1103,12 @@ 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,
)
if request_context is None: if request_context is None:
request_ctx = dataclasses.replace( request_ctx = dataclasses.replace(
request_ctx, request_ctx,
@@ -1156,11 +1155,13 @@ 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=initial_messages, initial_messages=None,
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,
@@ -1878,6 +1879,13 @@ class AgentLoop:
session, session,
runtime=runtime, runtime=runtime,
) )
# Token consolidation may have committed a replacement checkpoint
# after the compact stage captured its summary for this request.
ctx.session, ctx.pending_summary = self.auto_compact.prepare_session(
session,
ctx.session_key,
)
session = ctx.require_session()
is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent" is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent"
_hist_kwargs: dict[str, Any] = { _hist_kwargs: dict[str, Any] = {
@@ -1968,7 +1976,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.initial_messages = self._build_initial_messages(ctx) ctx.transcript_input = self._build_transcript_input(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()
@@ -1980,9 +1988,10 @@ 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.initial_messages, ctx.transcript_input,
runtime=runtime, runtime=runtime,
on_progress=ctx.on_progress, on_progress=ctx.on_progress,
on_stream=ctx.on_stream, on_stream=ctx.on_stream,
+142 -146
View File
@@ -35,6 +35,7 @@ from nanobot.utils.helpers import (
estimate_prompt_tokens_chain, estimate_prompt_tokens_chain,
strip_think, strip_think,
truncate_text, truncate_text,
truncate_text_to_tokens,
) )
from nanobot.utils.prompt_templates import render_template from nanobot.utils.prompt_templates import render_template
from nanobot.utils.workspace_prompts import ( from nanobot.utils.workspace_prompts import (
@@ -65,8 +66,6 @@ class MemoryStore:
# durable files are tiny in practice (~5 KB total), but a runaway file must # durable files are tiny in practice (~5 KB total), but a runaway file must
# not unbounded the prompt. # not unbounded the prompt.
_DREAM_FILE_EMBED_CAP = 8000 _DREAM_FILE_EMBED_CAP = 8000
_INTERNAL_HISTORY_SESSION_PREFIXES = ("cron:", "dream:")
_INTERNAL_HISTORY_SESSION_KEYS = {"heartbeat"}
_LEGACY_ENTRY_START_RE = re.compile(r"^\[(\d{4}-\d{2}-\d{2}[^\]]*)\]\s*") _LEGACY_ENTRY_START_RE = re.compile(r"^\[(\d{4}-\d{2}-\d{2}[^\]]*)\]\s*")
_LEGACY_TIMESTAMP_RE = re.compile(r"^\[(\d{4}-\d{2}-\d{2} \d{2}:\d{2})\]\s*") _LEGACY_TIMESTAMP_RE = re.compile(r"^\[(\d{4}-\d{2}-\d{2} \d{2}:\d{2})\]\s*")
_LEGACY_RAW_MESSAGE_RE = re.compile( _LEGACY_RAW_MESSAGE_RE = re.compile(
@@ -260,6 +259,29 @@ class MemoryStore:
# -- history.jsonl — append-only, JSONL format --------------------------- # -- history.jsonl — append-only, JSONL format ---------------------------
def _normalize_history_entry(
self,
entry: str,
*,
max_chars: int | None = None,
) -> str:
"""Return the exact bounded, model-safe text accepted by the journal."""
limit = max_chars if max_chars is not None else _HISTORY_ENTRY_HARD_CAP
raw = entry.rstrip()
content = strip_think(raw)
if len(content) > limit:
if not self._oversize_logged:
self._oversize_logged = True
logger.warning(
"history entry exceeds {} chars ({}); truncating. "
"Usually means a caller forgot its own cap; "
"further occurrences suppressed.",
limit,
len(content),
)
content = truncate_text(content, limit)
return content
def append_history( def append_history(
self, self,
entry: str, entry: str,
@@ -274,27 +296,16 @@ class MemoryStore:
persisted. If the cleaned content is empty but the raw entry wasn't, persisted. If the cleaned content is empty but the raw entry wasn't,
the record is persisted with an empty string rather than falling back the record is persisted with an empty string rather than falling back
to the raw leak — otherwise `strip_think`'s guarantees would be to the raw leak — otherwise `strip_think`'s guarantees would be
undone by history replay / consolidation downstream. undone when Dream consumes the journal entry.
A defensive cap (*max_chars*, default ``_HISTORY_ENTRY_HARD_CAP``) is A defensive cap (*max_chars*, default ``_HISTORY_ENTRY_HARD_CAP``) is
applied as a final safety net: individual callers should cap their own applied as a final safety net: individual callers should cap their own
content more tightly; this default only exists to catch unintentional content more tightly; this default only exists to catch unintentional
large writes (e.g. an LLM echoing its input back as a "summary"). large writes (e.g. an LLM echoing its input back as a "summary").
""" """
limit = max_chars if max_chars is not None else _HISTORY_ENTRY_HARD_CAP
ts = datetime.now().strftime("%Y-%m-%d %H:%M") ts = datetime.now().strftime("%Y-%m-%d %H:%M")
raw = entry.rstrip() raw = entry.rstrip()
if len(raw) > limit: content = self._normalize_history_entry(entry, max_chars=max_chars)
if not self._oversize_logged:
self._oversize_logged = True
logger.warning(
"history entry exceeds {} chars ({}); truncating. "
"Usually means a caller forgot its own cap; "
"further occurrences suppressed.",
limit, len(raw),
)
raw = truncate_text(raw, limit)
content = strip_think(raw)
# Cursor allocation and the append must be atomic: concurrent writers # Cursor allocation and the append must be atomic: concurrent writers
# could otherwise read the same current cursor and emit duplicates. # could otherwise read the same current cursor and emit duplicates.
with self._append_lock: with self._append_lock:
@@ -302,7 +313,7 @@ class MemoryStore:
if raw and not content: if raw and not content:
logger.debug( logger.debug(
"history entry {} stripped to empty (likely template leak); " "history entry {} stripped to empty (likely template leak); "
"persisting empty content to avoid re-polluting context", "persisting empty content to avoid re-polluting Dream input",
cursor, cursor,
) )
record = {"cursor": cursor, "timestamp": ts, "content": content} record = {"cursor": cursor, "timestamp": ts, "content": content}
@@ -392,36 +403,6 @@ class MemoryStore:
"""Return history entries with a valid cursor > *since_cursor*.""" """Return history entries with a valid cursor > *since_cursor*."""
return [e for e, c in self._iter_valid_entries() if c > since_cursor] return [e for e, c in self._iter_valid_entries() if c > since_cursor]
@classmethod
def _is_internal_history_session(cls, session_key: str | None) -> bool:
if not session_key:
return False
return (
session_key in cls._INTERNAL_HISTORY_SESSION_KEYS
or session_key.startswith(cls._INTERNAL_HISTORY_SESSION_PREFIXES)
)
def read_recent_history_for_prompt(
self,
since_cursor: int,
*,
session_key: str | None,
unified_session: bool = False,
) -> list[dict[str, Any]]:
"""Return unprocessed history entries safe to inject into a turn prompt."""
entries = self.read_unprocessed_history(since_cursor=since_cursor)
if session_key is None:
return entries
if not unified_session:
return [e for e in entries if e.get("session_key") == session_key]
return [
entry
for entry in entries
if (entry_session := entry.get("session_key")) == session_key
or not self._is_internal_history_session(entry_session)
]
def compact_history(self) -> None: def compact_history(self) -> None:
"""Drop oldest processed entries without discarding pending Dream input.""" """Drop oldest processed entries without discarding pending Dream input."""
if self.max_history_entries <= 0: if self.max_history_entries <= 0:
@@ -718,21 +699,28 @@ class MemoryStore:
*, *,
max_chars: int | None = None, max_chars: int | None = None,
session_key: str | None = None, session_key: str | None = None,
) -> None: ) -> str:
"""Fallback: dump raw messages to history.jsonl without LLM summarization.""" """Persist and return a bounded raw checkpoint when summarization degrades."""
limit = max_chars if max_chars is not None else _RAW_ARCHIVE_MAX_CHARS checkpoint = self._build_raw_checkpoint(messages, max_chars=max_chars)
formatted = truncate_text( self.append_history(checkpoint, session_key=session_key)
self._format_messages(public_history_messages(messages)),
limit,
)
self.append_history(
f"[RAW] {len(messages)} messages\n"
f"{formatted}",
session_key=session_key,
)
logger.warning( logger.warning(
"Memory consolidation degraded: raw-archived {} messages", len(messages) "Memory consolidation degraded: raw-archived {} messages", len(messages)
) )
return checkpoint
def _build_raw_checkpoint(
self,
messages: list[dict[str, Any]],
*,
max_chars: int | None = None,
) -> str:
"""Build the same bounded checkpoint as :meth:`raw_archive` without writing it."""
limit = max_chars if max_chars is not None else _RAW_ARCHIVE_MAX_CHARS
checkpoint = (
f"[RAW] {len(messages)} messages\n"
f"{self._format_messages(public_history_messages(messages))}"
)
return self._normalize_history_entry(checkpoint, max_chars=limit)
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Dream helpers # Dream helpers
@@ -787,12 +775,11 @@ class MemoryStore:
# Memory ingestion and legacy context-pressure coordination # Memory ingestion and legacy context-pressure coordination
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Individual history.jsonl writers cap their own payloads tightly; the # Raw fallbacks use a tighter cap. Completed model summaries may scale with the
# _HISTORY_ENTRY_HARD_CAP at append_history() is a belt-and-suspenders default # configured generation budget, while append_history() still enforces the
# that catches any new caller that forgot to set its own cap. # emergency hard cap against pathological provider output.
_RAW_ARCHIVE_MAX_CHARS = 16_000 # fallback dump (LLM failed) _RAW_ARCHIVE_MAX_CHARS = 16_000 # fallback dump (LLM failed)
_ARCHIVE_SUMMARY_MAX_CHARS = 8_000 # LLM-produced consolidation summary _HISTORY_ENTRY_HARD_CAP = 64_000 # emergency cap in append_history
_HISTORY_ENTRY_HARD_CAP = 64_000 # emergency cap in append_history
class MemoryArchiver: class MemoryArchiver:
@@ -809,13 +796,45 @@ class MemoryArchiver:
build_messages: Callable[..., list[dict[str, Any]]], build_messages: Callable[..., list[dict[str, Any]]],
get_tool_definitions: Callable[[], list[dict[str, Any]]], get_tool_definitions: Callable[[], list[dict[str, Any]]],
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None, resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
unified_session: bool = False,
) -> None: ) -> None:
self.store = store self.store = store
self._build_messages = build_messages self._build_messages = build_messages
self._get_tool_definitions = get_tool_definitions self._get_tool_definitions = get_tool_definitions
self._resolve_prompt_context = resolve_prompt_context self._resolve_prompt_context = resolve_prompt_context
self.unified_session = unified_session
def _raw_checkpoint(
self,
messages: list[dict[str, Any]],
*,
session_key: str,
previous_summary: str | None,
max_tokens: int,
) -> str:
"""Persist the failed chunk and return a bounded replacement checkpoint."""
raw = self.store.raw_archive(messages, session_key=session_key)
token_limit = max(1, max_tokens)
if not previous_summary:
return truncate_text_to_tokens(raw, token_limit)
combined = (
"[Previous archived context]\n"
f"{previous_summary}\n\n"
"[Newly archived raw context]\n"
f"{raw}"
)
bounded = truncate_text_to_tokens(combined, token_limit)
if bounded == combined:
return combined
# Keep evidence from both sides when their full concatenation cannot fit.
section_limit = max(1, (token_limit - 32) // 2)
return truncate_text_to_tokens(
"[Previous archived context]\n"
f"{truncate_text_to_tokens(previous_summary, section_limit)}\n\n"
"[Newly archived raw context]\n"
f"{truncate_text_to_tokens(raw, section_limit)}",
token_limit,
)
async def archive( async def archive(
self, self,
@@ -825,48 +844,53 @@ class MemoryArchiver:
session_key: str, session_key: str,
request_messages: list[dict[str, Any]], request_messages: list[dict[str, Any]],
request_tools: list[dict[str, Any]], request_tools: list[dict[str, Any]],
previous_summary: str | None = None,
) -> str | None: ) -> str | None:
"""Execute a prepared archive request and persist its result.""" """Execute a prepared archive request and persist its result."""
if not messages: if not messages:
return None return None
def raw_fallback() -> str:
return self._raw_checkpoint(
messages,
session_key=session_key,
previous_summary=previous_summary,
max_tokens=runtime.generation.max_tokens,
)
try: try:
with llm_usage_source("dream"): with llm_usage_source("dream"):
response = await runtime.provider.chat_with_retry( response = await runtime.provider.chat_with_retry(
model=runtime.model, model=runtime.model,
messages=request_messages, messages=request_messages,
tools=request_tools, tools=request_tools,
tool_choice="none",
temperature=runtime.generation.temperature, temperature=runtime.generation.temperature,
max_tokens=runtime.generation.max_tokens, max_tokens=runtime.generation.max_tokens,
reasoning_effort=runtime.generation.reasoning_effort, reasoning_effort=runtime.generation.reasoning_effort,
) )
except Exception: except Exception:
logger.warning("Memory archive provider call failed, raw-dumping to history") logger.warning("Memory archive provider call failed, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key) return raw_fallback()
return None
if response.finish_reason in {"error", "length"}: if response.finish_reason in {"error", "length"}:
logger.warning( logger.warning(
"Memory archive provider did not complete ({}), raw-dumping to history", "Memory archive provider did not complete ({}), raw-dumping to history",
response.finish_reason, response.finish_reason,
) )
self.store.raw_archive(messages, session_key=session_key) return raw_fallback()
return None
if response.has_tool_calls is True: if response.has_tool_calls is True:
logger.warning("Memory archive provider returned tool calls, raw-dumping to history") logger.warning("Memory archive provider returned tool calls, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key) return raw_fallback()
return None
summary = response.content summary = response.content
if not summary or not summary.strip(): if not summary or not summary.strip():
logger.warning("Memory archive provider returned no summary, raw-dumping to history") logger.warning("Memory archive provider returned no summary, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key) return raw_fallback()
return None summary = self.store._normalize_history_entry(summary)
if summary.strip() == "(nothing)": if not summary:
logger.warning("Memory archive provider summary was not safe to replay, raw-dumping")
return raw_fallback()
if summary == "(nothing)":
return "(nothing)" return "(nothing)"
self.store.append_history( self.store.append_history(summary, session_key=session_key)
summary,
max_chars=_ARCHIVE_SUMMARY_MAX_CHARS,
session_key=session_key,
)
return summary return summary
async def archive_session( async def archive_session(
@@ -881,13 +905,26 @@ class MemoryArchiver:
messages = list(session.messages[session.last_archived:archive_end]) messages = list(session.messages[session.last_archived:archive_end])
if not messages: if not messages:
return None return None
session_summary = session_summary_from_metadata(
session.metadata,
fallback_last_active=session.updated_at,
)
previous_summary = session_summary["text"] if session_summary else None
def raw_fallback() -> str:
return self._raw_checkpoint(
messages,
session_key=session.key,
previous_summary=previous_summary,
max_tokens=runtime.generation.max_tokens,
)
if input_token_budget <= 0: if input_token_budget <= 0:
logger.debug( logger.debug(
"Memory archive has no safe input budget for {}; raw-dumping", "Memory archive has no safe input budget for {}; raw-dumping",
session.key, session.key,
) )
self.store.raw_archive(messages, session_key=session.key) return raw_fallback()
return None
prefix = Session( prefix = Session(
key=session.key, key=session.key,
messages=list(session.messages[:archive_end]), messages=list(session.messages[:archive_end]),
@@ -903,13 +940,8 @@ class MemoryArchiver:
"Memory archive cannot replay the full chunk for {}; raw-dumping", "Memory archive cannot replay the full chunk for {}; raw-dumping",
session.key, session.key,
) )
self.store.raw_archive(messages, session_key=session.key) return raw_fallback()
return None prompt = render_template("agent/consolidator_archive.md", strip=True)
prompt = render_template(
"agent/consolidator_archive.md",
strip=True,
archive_count=len(archive_history),
)
channel = session.key.split(":", 1)[0] if ":" in session.key else None channel = session.key.split(":", 1)[0] if ":" in session.key else None
workspace: Path | None = None workspace: Path | None = None
if self._resolve_prompt_context is not None: if self._resolve_prompt_context is not None:
@@ -918,13 +950,8 @@ class MemoryArchiver:
history=history, history=history,
current_message=prompt, current_message=prompt,
channel=channel, channel=channel,
session_summary=session_summary_from_metadata( session_summary=session_summary,
session.metadata,
fallback_last_active=session.updated_at,
),
workspace=workspace, workspace=workspace,
session_key=session.key,
unified_session=self.unified_session,
) )
tools = self._get_tool_definitions() tools = self._get_tool_definitions()
estimated, source = estimate_prompt_tokens_chain( estimated, source = estimate_prompt_tokens_chain(
@@ -941,14 +968,14 @@ class MemoryArchiver:
input_token_budget, input_token_budget,
source, source,
) )
self.store.raw_archive(messages, session_key=session.key) return raw_fallback()
return None
return await self.archive( return await self.archive(
messages, messages,
runtime=runtime, runtime=runtime,
session_key=session.key, session_key=session.key,
request_messages=request_messages, request_messages=request_messages,
request_tools=tools, request_tools=tools,
previous_summary=previous_summary,
) )
@@ -964,20 +991,16 @@ class Consolidator:
build_messages: Callable[..., list[dict[str, Any]]], build_messages: Callable[..., list[dict[str, Any]]],
get_tool_definitions: Callable[[], list[dict[str, Any]]], get_tool_definitions: Callable[[], list[dict[str, Any]]],
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None, resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
unified_session: bool = False,
): ):
self.store = store self.store = store
self.sessions = sessions self.sessions = sessions
self.unified_session = unified_session
self._build_messages = build_messages self._build_messages = build_messages
self._get_tool_definitions = get_tool_definitions self._get_tool_definitions = get_tool_definitions
self._resolve_prompt_context = resolve_prompt_context
self.archiver = MemoryArchiver( self.archiver = MemoryArchiver(
store=store, store=store,
build_messages=build_messages, build_messages=build_messages,
get_tool_definitions=get_tool_definitions, get_tool_definitions=get_tool_definitions,
resolve_prompt_context=resolve_prompt_context, resolve_prompt_context=resolve_prompt_context,
unified_session=unified_session,
) )
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = ( self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
weakref.WeakValueDictionary() weakref.WeakValueDictionary()
@@ -1013,13 +1036,18 @@ class Consolidator:
return [] return []
return session.get_history() return session.get_history()
def _persist_last_summary(self, session: Session, summary: str | None) -> None: @staticmethod
if summary and summary != "(nothing)": def _set_last_summary(
session: Session,
summary: str,
*,
last_active: datetime | None = None,
) -> None:
if summary != "(nothing)":
session.metadata["_last_summary"] = { session.metadata["_last_summary"] = {
"text": summary, "text": summary,
"last_active": session.updated_at.isoformat(), "last_active": (last_active or session.updated_at).isoformat(),
} }
self.sessions.save(session)
def estimate_session_prompt_tokens( def estimate_session_prompt_tokens(
self, self,
@@ -1039,8 +1067,6 @@ class Consolidator:
current_message="[token-probe]", current_message="[token-probe]",
channel=channel, channel=channel,
session_summary=summary, session_summary=summary,
session_key=session.key,
unified_session=self.unified_session,
) )
return estimate_prompt_tokens_chain( return estimate_prompt_tokens_chain(
runtime.provider, runtime.provider,
@@ -1057,24 +1083,6 @@ class Consolidator:
- self._SAFETY_BUFFER - self._SAFETY_BUFFER
) )
async def archive(
self,
messages: list[dict[str, Any]],
*,
runtime: LLMRuntime,
session_key: str,
request_messages: list[dict[str, Any]],
request_tools: list[dict[str, Any]],
) -> str | None:
"""Compatibility wrapper for the extracted MemoryArchiver."""
return await self.archiver.archive(
messages,
runtime=runtime,
session_key=session_key,
request_messages=request_messages,
request_tools=request_tools,
)
async def archive_session( async def archive_session(
self, self,
session: Session, session: Session,
@@ -1101,26 +1109,23 @@ class Consolidator:
The budget reserves space for completion tokens and a safety buffer The budget reserves space for completion tokens and a safety buffer
so the LLM request never exceeds the context window. so the LLM request never exceeds the context window.
""" """
if runtime.context_window_tokens <= 0:
return
lock = self.get_lock(session.key) lock = self.get_lock(session.key)
async with lock: async with lock:
# Refresh session reference: AutoCompact may have replaced it. # Refresh session reference: AutoCompact may have replaced it.
fresh = self.sessions.get_or_create(session.key) fresh = self.sessions.get_or_create(session.key)
if fresh is not session: if fresh is not session:
session = fresh session = fresh
if runtime.context_window_tokens <= 0:
return
if not session.messages: if not session.messages:
return return
budget = self._input_token_budget(runtime) budget = self._input_token_budget(runtime)
last_summary: str | None = None
estimated, source = self.estimate_session_prompt_tokens( estimated, source = self.estimate_session_prompt_tokens(
session, session,
runtime=runtime, runtime=runtime,
) )
if estimated <= 0: if estimated <= 0:
self._persist_last_summary(session, last_summary)
return return
if estimated < budget: if estimated < budget:
unarchived_count = len(session.messages) - session.last_archived unarchived_count = len(session.messages) - session.last_archived
@@ -1132,7 +1137,6 @@ class Consolidator:
source, source,
unarchived_count, unarchived_count,
) )
self._persist_last_summary(session, last_summary)
return return
end_idx = self.pick_consolidation_boundary(session) end_idx = self.pick_consolidation_boundary(session)
@@ -1160,18 +1164,12 @@ class Consolidator:
archive_end=end_idx, archive_end=end_idx,
runtime=runtime, runtime=runtime,
) )
# Advance either way: archive_session raw-archives on degradation, if summary is None:
# and replaying the same chunk would duplicate Memory material. return
if summary: self._set_last_summary(session, summary)
last_summary = summary
session.last_archived = end_idx session.last_archived = end_idx
self.sessions.save(session) self.sessions.save(session)
# Persist the last summary to session metadata so it can be injected
# into the runtime context on the next prepare_session() call, aligning
# the summary injection strategy with AutoCompact._archive().
self._persist_last_summary(session, last_summary)
async def compact_idle_session( async def compact_idle_session(
self, self,
session_key: str, session_key: str,
@@ -1209,12 +1207,10 @@ class Consolidator:
archive_end=archive_end, archive_end=archive_end,
runtime=runtime, runtime=runtime,
) )
if summary is None:
return None
if summary and summary != "(nothing)": self._set_last_summary(session, summary, last_active=last_active)
session.metadata["_last_summary"] = {
"text": summary,
"last_active": last_active.isoformat(),
}
# A turn can append while the provider call is in flight. Advance only # A turn can append while the provider call is in flight. Advance only
# through the captured batch so new messages remain eligible next time. # through the captured batch so new messages remain eligible next time.
+183 -65
View File
@@ -14,6 +14,7 @@ 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,
@@ -66,6 +67,7 @@ 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 = (
@@ -94,11 +96,13 @@ 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]] initial_messages: list[dict[str, Any]] | None
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
@@ -135,6 +139,17 @@ class AgentRunResult:
provider_state: ProviderConversationState | None = field(default=None, repr=False) provider_state: ProviderConversationState | None = field(default=None, repr=False)
@dataclass(slots=True)
class _ModelRequestState:
"""Per-run state used to govern the next provider request."""
config: ContextGovernanceConfig
conversation: ProviderConversationStateController
usage: LLMUsage | None = None
messages: list[dict[str, Any]] | None = None
tool_definitions: list[dict[str, Any]] | None = None
class AgentRunner: class AgentRunner:
"""Run a tool-capable LLM loop without product-layer concerns.""" """Run a tool-capable LLM loop without product-layer concerns."""
@@ -410,7 +425,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 = list(spec.initial_messages) messages = self._initial_transcript(spec)
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)
@@ -462,6 +477,19 @@ 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,
@@ -483,7 +511,6 @@ class AgentRunner:
length_recovery_parts: list[str] = [] length_recovery_parts: list[str] = []
had_injections = False had_injections = False
injection_cycles = 0 injection_cycles = 0
compacted_tool_call_ids: set[str] = set()
pending_stream_content: str | None = None pending_stream_content: str | None = None
conversation_state = ProviderConversationStateController( conversation_state = ProviderConversationStateController(
provider=spec.runtime.provider, provider=spec.runtime.provider,
@@ -502,39 +529,29 @@ 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(spec.initial_messages), )
request_state = _ModelRequestState(
config=governance_config,
conversation=conversation_state,
) )
for iteration in range(spec.max_iterations): for iteration in range(spec.max_iterations):
# Keep the persisted conversation untouched. Context governance
# may repair or compact historical messages for the model, but
# those synthetic edits must not shift the append boundary used
# later when the caller saves only the new turn. A governance
# failure must stop the run instead of sending an ungoverned copy.
messages_for_model = self.context_governor.prepare_for_model(
governance_config,
messages,
compacted_tool_call_ids,
)
context = AgentHookContext( context = AgentHookContext(
iteration=iteration, iteration=iteration,
messages=messages, messages=messages,
session_key=spec.session_key, session_key=spec.session_key,
) )
await hook.before_iteration(context) await hook.before_iteration(context)
provider_context = conversation_state.prepare_request(
messages,
context_window_tokens=spec.runtime.context_window_tokens,
model_messages=messages_for_model,
)
response = await self._request_model( response = await self._request_model(
spec, spec,
messages_for_model, messages,
hook, hook,
context, context,
conversation_state=conversation_state, request_state=request_state,
provider_context=provider_context, transcript=messages,
) )
assert request_state.messages is not None
messages_for_model = request_state.messages
conversation_state.observe_response(response, messages) conversation_state.observe_response(response, messages)
context.response = response context.response = response
context.tool_calls = list(response.tool_calls) context.tool_calls = list(response.tool_calls)
@@ -546,7 +563,7 @@ class AgentRunner:
response.content, response.content,
) )
response.content = cleaned_content response.content = cleaned_content
raw_usage = self._usage_or_estimate(spec, messages_for_model, response) raw_usage = self._record_request_usage(spec, request_state, response)
context.usage = raw_usage context.usage = raw_usage
usage = self._merge_usage(usage, raw_usage) usage = self._merge_usage(usage, raw_usage)
if reasoning_text and not context.streamed_reasoning: if reasoning_text and not context.streamed_reasoning:
@@ -620,7 +637,6 @@ class AgentRunner:
self.context_governor.prepare_for_model( self.context_governor.prepare_for_model(
governance_config, governance_config,
messages, messages,
compacted_tool_call_ids,
) )
if response.provider_state is not None if response.provider_state is not None
else None else None
@@ -686,14 +702,13 @@ class AgentRunner:
) )
if hook.wants_streaming(): if hook.wants_streaming():
await hook.on_stream_end(context, resuming=False) await hook.on_stream_end(context, resuming=False)
retry_messages = self._finalization_retry_messages(messages_for_model)
response = await self._request_finalization_retry( response = await self._request_finalization_retry(
spec, spec,
messages_for_model, messages_for_model,
request_state=request_state,
transcript=messages, transcript=messages,
conversation_state=conversation_state,
) )
retry_usage = self._usage_or_estimate(spec, retry_messages, response) retry_usage = self._record_request_usage(spec, request_state, response)
usage = self._merge_usage(usage, retry_usage) usage = self._merge_usage(usage, retry_usage)
raw_usage = self._merge_usage(raw_usage, retry_usage) raw_usage = self._merge_usage(raw_usage, retry_usage)
context.response = response context.response = response
@@ -880,7 +895,7 @@ class AgentRunner:
hook, hook,
messages, messages,
usage, usage,
conversation_state, request_state=request_state,
) )
if terminal_content is None: if terminal_content is None:
terminal_content = self._max_iterations_fallback(spec) terminal_content = self._max_iterations_fallback(spec)
@@ -927,6 +942,60 @@ class AgentRunner:
kwargs["reasoning_effort"] = generation.reasoning_effort kwargs["reasoning_effort"] = generation.reasoning_effort
return kwargs return kwargs
def _prepare_model_request(
self,
state: _ModelRequestState,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
transcript: list[dict[str, Any]] | None = None,
) -> tuple[list[dict[str, Any]], ProviderCallContext | None]:
"""Prepare, fit, and record the exact payload sent to a provider."""
prepared = self.context_governor.prepare_for_model(state.config, messages)
supplemental_messages = (
[prepared[-1]] if transcript is not None and tool_definitions is None else None
)
model_messages = None if supplemental_messages is not None else prepared
request_context_tokens = (
state.conversation.estimate_request_context_tokens(
transcript,
model_messages=model_messages,
supplemental_messages=supplemental_messages,
tool_definitions=tool_definitions,
)
if transcript is not None
else None
)
usage_matches_messages = (
state.messages is not None
and prepared == state.messages
and tool_definitions == state.tool_definitions
)
prepared, fitted = self.context_governor.fit_request(
state.config,
prepared,
state.usage,
usage_matches_messages=usage_matches_messages,
tool_definitions=tool_definitions,
request_context_tokens=request_context_tokens,
)
provider_context = (
state.conversation.prepare_request(
transcript,
context_window_tokens=state.config.context_window_tokens,
model_messages=model_messages,
supplemental_messages=supplemental_messages,
resume_state=not fitted,
)
if transcript is not None
else state.conversation.independent_request_context(
context_window_tokens=state.config.context_window_tokens,
)
)
state.messages = deepcopy(prepared)
state.tool_definitions = deepcopy(tool_definitions)
return prepared, provider_context
async def _request_model( async def _request_model(
self, self,
spec: AgentRunSpec, spec: AgentRunSpec,
@@ -934,21 +1003,29 @@ class AgentRunner:
hook: AgentHook, hook: AgentHook,
context: AgentHookContext, context: AgentHookContext,
*, *,
request_state: _ModelRequestState,
malformed_retry: bool = False, malformed_retry: bool = False,
conversation_state: ProviderConversationStateController, transcript: list[dict[str, Any]] | None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse: ) -> LLMResponse:
timeout_s = self._resolve_llm_timeout_s(spec) timeout_s = self._resolve_llm_timeout_s(spec)
tool_definitions = spec.tools.get_definitions()
messages, provider_context = self._prepare_model_request(
request_state,
messages,
tool_definitions=tool_definitions,
transcript=transcript,
)
kwargs = self._build_request_kwargs( kwargs = self._build_request_kwargs(
spec, spec,
messages, messages,
tools=spec.tools.get_definitions(), tools=tool_definitions,
) )
wants_streaming = hook.wants_streaming() wants_streaming = hook.wants_streaming()
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
@@ -972,11 +1049,29 @@ 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 nonlocal native_reasoning_open, native_reasoning_close_task
if not native_reasoning_open: if native_reasoning_close_task is None:
return if not native_reasoning_open:
native_reasoning_open = False return
await hook.emit_reasoning_end() native_reasoning_open = False
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":
@@ -1051,6 +1146,10 @@ 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(
@@ -1098,11 +1197,9 @@ class AgentRunner:
) )
return await self._request_model( return await self._request_model(
spec, retry_messages, hook, context, spec, retry_messages, hook, context,
request_state=request_state,
malformed_retry=True, malformed_retry=True,
conversation_state=conversation_state, transcript=None,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
) )
if ( if (
all_dropped all_dropped
@@ -1118,9 +1215,7 @@ class AgentRunner:
return await self._request_no_tools( return await self._request_no_tools(
spec, spec,
fallback_messages, fallback_messages,
provider_context=conversation_state.independent_request_context( request_state=request_state,
context_window_tokens=spec.runtime.context_window_tokens,
),
) )
return response return response
@@ -1188,21 +1283,17 @@ class AgentRunner:
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
*, *,
request_state: _ModelRequestState,
transcript: list[dict[str, Any]], transcript: list[dict[str, Any]],
conversation_state: ProviderConversationStateController,
) -> LLMResponse: ) -> LLMResponse:
retry_messages = self._finalization_retry_messages(messages) retry_messages = self._finalization_retry_messages(messages)
provider_context = conversation_state.prepare_request(
transcript,
context_window_tokens=spec.runtime.context_window_tokens,
supplemental_messages=[retry_messages[-1]],
)
response = await self._request_no_tools( response = await self._request_no_tools(
spec, spec,
retry_messages, retry_messages,
provider_context=provider_context, request_state=request_state,
transcript=transcript,
) )
conversation_state.observe_response( request_state.conversation.observe_response(
response, response,
transcript, transcript,
adopt_candidate_state=False, adopt_candidate_state=False,
@@ -1221,16 +1312,15 @@ class AgentRunner:
hook: AgentHook, hook: AgentHook,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
usage: LLMUsage | None, usage: LLMUsage | None,
conversation_state: ProviderConversationStateController, *,
request_state: _ModelRequestState,
) -> tuple[str | None, LLMUsage | None]: ) -> tuple[str | None, LLMUsage | None]:
retry_messages = self._budget_exhausted_finalization_messages(messages) retry_messages = self._budget_exhausted_finalization_messages(messages)
try: try:
response = await self._request_no_tools( response = await self._request_no_tools(
spec, spec,
retry_messages, retry_messages,
provider_context=conversation_state.independent_request_context( request_state=request_state,
context_window_tokens=spec.runtime.context_window_tokens,
),
) )
except Exception: except Exception:
logger.exception( logger.exception(
@@ -1239,7 +1329,7 @@ class AgentRunner:
) )
return None, usage return None, usage
raw_usage = self._usage_or_estimate(spec, retry_messages, response) raw_usage = self._record_request_usage(spec, request_state, response)
usage = self._merge_usage(usage, raw_usage) usage = self._merge_usage(usage, raw_usage)
if response.finish_reason == "error" or response.has_tool_calls: if response.finish_reason == "error" or response.has_tool_calls:
logger.warning( logger.warning(
@@ -1268,8 +1358,15 @@ class AgentRunner:
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
*, *,
provider_context: ProviderCallContext | None = None, request_state: _ModelRequestState,
transcript: list[dict[str, Any]] | None = None,
) -> LLMResponse: ) -> LLMResponse:
messages, provider_context = self._prepare_model_request(
request_state,
messages,
tool_definitions=None,
transcript=transcript,
)
kwargs = self._build_request_kwargs( kwargs = self._build_request_kwargs(
spec, spec,
messages, messages,
@@ -1281,17 +1378,18 @@ class AgentRunner:
) )
timeout_s = self._resolve_llm_timeout_s(spec) timeout_s = self._resolve_llm_timeout_s(spec)
try: try:
return ( response = (
await coro await coro
if timeout_s is None if timeout_s is None
else await asyncio.wait_for(coro, timeout=timeout_s) else await asyncio.wait_for(coro, timeout=timeout_s)
) )
except asyncio.TimeoutError: except asyncio.TimeoutError:
return LLMResponse( response = LLMResponse(
content=f"Error calling LLM: timed out after {timeout_s:g}s", content=f"Error calling LLM: timed out after {timeout_s:g}s",
finish_reason="error", finish_reason="error",
error_kind="timeout", error_kind="timeout",
) )
return response
@staticmethod @staticmethod
def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None: def _resolve_llm_timeout_s(spec: AgentRunSpec) -> float | None:
@@ -1333,33 +1431,53 @@ class AgentRunner:
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
response: LLMResponse, response: LLMResponse,
*,
tool_definitions: list[dict[str, Any]] | None,
) -> LLMUsage | None: ) -> LLMUsage | None:
usage = response.usage usage = response.usage
if response.finish_reason == "error": if response.finish_reason == "error":
if usage is None or usage.total_tokens == 0: if usage is None or usage.total_tokens == 0:
usage = LLMUsage.empty_request() usage = LLMUsage.empty_request()
elif usage is None or usage.total_tokens == 0: elif usage is None or usage.total_tokens == 0:
usage = self._estimate_response_usage(spec, messages, response) usage = self._estimate_response_usage(
spec,
messages,
response,
tool_definitions=tool_definitions,
)
return usage.with_timing( return usage.with_timing(
generation_ms=response.generation_ms, generation_ms=response.generation_ms,
ttft_ms=response.ttft_ms, ttft_ms=response.ttft_ms,
) )
def _record_request_usage(
self,
spec: AgentRunSpec,
state: _ModelRequestState,
response: LLMResponse,
) -> LLMUsage | None:
assert state.messages is not None
state.usage = self._usage_or_estimate(
spec,
state.messages,
response,
tool_definitions=state.tool_definitions,
)
return state.usage
def _estimate_response_usage( def _estimate_response_usage(
self, self,
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
response: LLMResponse, response: LLMResponse,
*,
tool_definitions: list[dict[str, Any]] | None,
) -> LLMUsage: ) -> LLMUsage:
try:
tools = spec.tools.get_definitions()
except Exception:
tools = None
prompt_tokens, _ = estimate_prompt_tokens_chain( prompt_tokens, _ = estimate_prompt_tokens_chain(
spec.runtime.provider, spec.runtime.provider,
spec.runtime.model, spec.runtime.model,
messages, messages,
tools, tool_definitions,
) )
assistant_message = build_assistant_message( assistant_message = build_assistant_message(
response.content or "", response.content or "",
+15 -7
View File
@@ -43,6 +43,13 @@ _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],
@@ -105,7 +112,7 @@ async def _execute_tool_call(
"status": "error", "status": "error",
"detail": "repeated external lookup blocked", "detail": "repeated external lookup blocked",
} }
return lookup_error + _RETRY_HINT, event return _with_retry_hint(lookup_error), event
prepare_call = cast( prepare_call = cast(
Callable[[str, Any], object] | None, Callable[[str, Any], object] | None,
@@ -119,6 +126,7 @@ 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",
@@ -126,14 +134,14 @@ async def _execute_tool_call(
} }
handled = _classify_violation( handled = _classify_violation(
raw_text=prep_error, raw_text=prep_error,
soft_payload=prep_error + _RETRY_HINT, soft_payload=payload,
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 prep_error + _RETRY_HINT, event return payload, event
await hook.before_execute_tool(context, tool_call, tool, params) await hook.before_execute_tool(context, tool_call, tool, params)
try: try:
@@ -150,10 +158,9 @@ async def _execute_tool_call(
"status": "error", "status": "error",
"detail": str(exc), "detail": str(exc),
} }
payload = f"Error: {type(exc).__name__}: {exc}" payload = _with_retry_hint(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,
@@ -165,6 +172,7 @@ 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",
@@ -172,14 +180,14 @@ async def _execute_tool_call(
} }
handled = _classify_violation( handled = _classify_violation(
raw_text=result, raw_text=result,
soft_payload=result + _RETRY_HINT, soft_payload=payload,
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 result + _RETRY_HINT, event return payload, event
await hook.after_execute_tool(context, tool_call, tool, params, result) await hook.after_execute_tool(context, tool_call, tool, params, result)
+7 -11
View File
@@ -861,8 +861,10 @@ def _best_window(old_text: str, content: str) -> tuple[float, int, list[str], li
@tool_parameters( @tool_parameters(
tool_parameters_schema( tool_parameters_schema(
path=StringSchema("The file path to edit"), path=StringSchema("The file path to edit"),
old_text=StringSchema("The text to find and replace"), old_text=StringSchema("The text to find and replace; copy it from read_file."),
new_text=StringSchema("The text to replace with"), new_text=StringSchema(
"The replacement text; must differ from old_text for an existing file."
),
replace_all=BooleanSchema(description="Replace all occurrences (default false)"), replace_all=BooleanSchema(description="Replace all occurrences (default false)"),
occurrence=IntegerSchema( occurrence=IntegerSchema(
description="Optional 1-based occurrence to replace when old_text appears multiple times.", description="Optional 1-based occurrence to replace when old_text appears multiple times.",
@@ -899,15 +901,9 @@ class EditFileTool(_FsTool):
@property @property
def description(self) -> str: def description(self) -> str:
return ( return (
"Perform a small, exact replacement in one file by replacing " "Perform a small, exact replacement in one file. "
"old_text with new_text. When replacing text in an existing file, " "Prefer apply_patch for multi-file, structural, or generated edits. "
"old_text and new_text must be different. Use this for narrow text substitutions " "occurrence, line_hint, and replace_all=true are mutually exclusive."
"with old_text copied from read_file. For multi-file, structural, "
"or generated code edits, prefer apply_patch. If old_text matches "
"multiple times, provide more context or set occurrence, line_hint, "
"replace_all, and expected_replacements. When editing from numbered "
"read_file output, set line_hint to the exact target line. "
"Shows closest-match diagnostics on failure."
) )
@staticmethod @staticmethod
+16 -3
View File
@@ -7,7 +7,7 @@ from __future__ import annotations
import asyncio import asyncio
import json import json
import time import time
from collections import deque from collections import OrderedDict, 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: dict[str, deque[float]] = {} self._sent_at: OrderedDict[str, deque[float]] = OrderedDict()
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,8 +240,11 @@ 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:
@@ -259,6 +262,8 @@ 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)
@@ -271,6 +276,14 @@ 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,
+24 -2
View File
@@ -182,6 +182,12 @@ 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).
@@ -196,7 +202,7 @@ class NanobotDingTalkHandler(_CallbackHandlerBase):
) )
) )
self.channel._background_tasks.add(task) self.channel._background_tasks.add(task)
task.add_done_callback(self.channel._background_tasks.discard) task.add_done_callback(self.channel._on_background_task_done)
return AckMessage.STATUS_OK, "OK" return AckMessage.STATUS_OK, "OK"
@@ -256,6 +262,17 @@ 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."""
@@ -272,6 +289,7 @@ 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)
@@ -309,6 +327,7 @@ 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
@@ -326,8 +345,11 @@ 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
for task in self._background_tasks: background_tasks = tuple(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 from unittest.mock import AsyncMock, MagicMock
import httpx import httpx
import pytest import pytest
@@ -402,6 +402,61 @@ 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."""
@@ -451,6 +506,72 @@ 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
@@ -650,6 +771,41 @@ 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)."""
+58 -43
View File
@@ -430,7 +430,13 @@ 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)
@@ -438,29 +444,30 @@ class EmailChannel(BaseChannel):
return messages return messages
try: try:
status, data = client.search(None, *search_criteria) status, data = client.uid("SEARCH", None, *search_criteria)
if status != "OK" or not data: if status != "OK" or not data or not data[0]:
return messages return messages
ids = data[0].split() uids = [raw.decode("ascii", errors="ignore") for raw in data[0].split()]
if limit > 0 and len(ids) > limit: if limit > 0 and len(uids) > limit:
ids = ids[-limit:] uids = uids[-limit:]
for imap_id in ids:
status, fetched = client.fetch(imap_id, "(BODY.PEEK[] UID)") features: _ServerFeatures | None = None
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)
raw_bytes = self._extract_message_bytes(fetched) if header_bytes is None:
if raw_bytes is None:
continue continue
uid = self._extract_uid(fetched) parsed = BytesParser(policy=policy.default).parsebytes(header_bytes)
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
@@ -468,9 +475,8 @@ 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:
client.store(imap_id, "+FLAGS", "\\Seen") features = self._mark_seen_uid(client, uid, features)
if uid: skipped_uids.add(uid)
skipped_uids.add(uid)
continue continue
# --- Anti-spoofing: verify Authentication-Results --- # --- Anti-spoofing: verify Authentication-Results ---
@@ -482,8 +488,7 @@ class EmailChannel(BaseChannel):
sender, sender,
) )
self._remember_processed_uid(uid, dedupe, cycle_uids) self._remember_processed_uid(uid, dedupe, cycle_uids)
if uid: skipped_uids.add(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(
@@ -492,18 +497,26 @@ class EmailChannel(BaseChannel):
sender, sender,
) )
self._remember_processed_uid(uid, dedupe, cycle_uids) self._remember_processed_uid(uid, dedupe, cycle_uids)
if uid: skipped_uids.add(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:
client.store(imap_id, "+FLAGS", "\\Seen") features = self._mark_seen_uid(client, uid, features)
if uid: skipped_uids.add(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()
@@ -556,10 +569,19 @@ 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:
client.store(imap_id, "+FLAGS", "\\Seen") features = self._mark_seen_uid(client, uid, features)
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)
@@ -714,11 +736,14 @@ 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", "(\\Deleted)") status, _ = client.uid("STORE", uid, "+FLAGS", f"({flag})")
if status == "OK": if status == "OK":
features.uid_store = True features.uid_store = True
return True return True
@@ -728,12 +753,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("Post-action skipped: UID {} not found", uid) self.logger.warning("Could not locate UID {} to set flag {}", uid, flag)
return False return False
status, _ = client.store(imap_id, "+FLAGS", "\\Deleted") status, _ = client.store(imap_id, "+FLAGS", flag)
if status != "OK": if status != "OK":
self.logger.warning("Post-action failed: could not mark UID {} as deleted", uid) self.logger.warning("Failed to set flag {} on UID {}", flag, uid)
return False return False
return True return True
@@ -773,16 +798,6 @@ 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,30 +53,7 @@ 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")
class FakeIMAP: fake = _make_fake_imap(raw, uid=b"123")
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())
@@ -86,38 +63,25 @@ 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 fake.store_calls == [(b"1", "+FLAGS", "\\Seen")] assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
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")
class FakeIMAP: fake = _make_fake_imap(raw, uid=b"123")
def login(self, _user: str, _pw: str): monkeypatch.setattr("nanobot.channels.email.runtime.imaplib.IMAP4_SSL", lambda _h, _p: fake)
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()
@@ -130,26 +94,10 @@ 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")
class FakeIMAP: monkeypatch.setattr(
def login(self, _user: str, _pw: str): "nanobot.channels.email.runtime.imaplib.IMAP4_SSL",
return "OK", [b"logged in"] lambda _h, _p: _make_fake_imap(raw, uid=b"123"),
)
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),
@@ -545,30 +493,7 @@ 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")
class FakeIMAP: fake = _make_fake_imap(raw, uid=b"123")
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())
@@ -576,7 +501,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 fake.store_calls == [(b"1", "+FLAGS", "\\Seen")] assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
# 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()
@@ -614,37 +539,14 @@ 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")
class FakeIMAP: fake = _make_fake_imap(raw, uid=b"123")
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 fake.store_calls == [(b"1", "+FLAGS", "\\Seen")] assert ("STORE", "123", "+FLAGS", "(\\Seen)") in fake.uid_calls
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:
@@ -662,15 +564,16 @@ 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 search(self, *_args): def uid(self, command: str, *args):
self.search_calls += 1 if command == "SEARCH":
if fail_once["pending"]: self.search_calls += 1
fail_once["pending"] = False if fail_once["pending"]:
raise imaplib.IMAP4.abort("socket error") fail_once["pending"] = False
return "OK", [b"1"] raise imaplib.IMAP4.abort("socket error")
return "OK", [b"123"]
def fetch(self, _imap_id: bytes, _parts: str): if command == "FETCH":
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))
@@ -700,10 +603,7 @@ 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 = { mailbox_state = {"123": raw_first, "124": raw_second}
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:
@@ -713,20 +613,18 @@ 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 search(self, *_args): def uid(self, command: str, *args):
unseen_ids = [imap_id for imap_id, item in mailbox_state.items() if not item["seen"]] if command == "SEARCH":
return "OK", [b" ".join(unseen_ids)] keys = " ".join(sorted(mailbox_state.keys(), key=int))
return "OK", [keys.encode()]
def fetch(self, imap_id: bytes, _parts: str): if command == "FETCH":
if imap_id == b"2" and fail_once["pending"]: uid = args[0]
fail_once["pending"] = False if uid == "124" and fail_once["pending"]:
raise imaplib.IMAP4.abort("socket error") fail_once["pending"] = False
item = mailbox_state[imap_id] raise imaplib.IMAP4.abort("socket error")
header = b"%s (UID %s BODY[] {200})" % (imap_id, item["uid"]) raw = mailbox_state[uid]
return "OK", [(header, item["raw"]), b")"] header = f"{uid} (UID {uid} BODY[] {{200}})".encode()
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):
@@ -1044,12 +942,13 @@ 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 search(self, *_args): def uid(self, command: str, *args):
self.search_args = _args if command == "SEARCH":
return "OK", [b"5"] self.search_args = args
return "OK", [b"999"]
def fetch(self, _imap_id: bytes, _parts: str): if command == "FETCH":
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))
@@ -1070,7 +969,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"
# search(None, "SINCE", "06-Feb-2026", "BEFORE", "07-Feb-2026") # uid("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 == []
@@ -1080,11 +979,12 @@ 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): def _make_fake_imap(raw: bytes, uid: bytes = b"500"):
"""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"]
@@ -1092,11 +992,16 @@ def _make_fake_imap(raw: bytes):
def select(self, _mailbox: str): def select(self, _mailbox: str):
return "OK", [b"1"] return "OK", [b"1"]
def search(self, *_args): def capability(self):
return "OK", [b"1"] return "OK", [b"IMAP4rev1"]
def fetch(self, _imap_id: bytes, _parts: str): def uid(self, command: str, *args):
return "OK", [(b"1 (UID 500 BODY[] {200})", raw), b")"] self.uid_calls.append((command, *args))
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))
@@ -1292,7 +1197,10 @@ 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 fake.store_calls == [(b"1", "+FLAGS", "\\Seen")] assert [call for call in fake.uid_calls if call[0] == "FETCH"] == [
("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:
+71 -19
View File
@@ -897,6 +897,68 @@ class TelegramChannel(BaseChannel):
self.logger.debug("sendRichMessage failed: {}", exc) self.logger.debug("sendRichMessage failed: {}", exc)
return False return False
async def _try_edit_rich(self, chat_id: int, message_id: int, content: str) -> bool:
"""Upgrade an existing message to rich in place via editMessageText (Bot API 10.1).
Editing in place keeps the message identity, so the streaming preview is
upgraded without the delete-and-resend pattern that caused flickering and
dropped line breaks (issue #4470).
Returns True when the rich edit is in place (including the ambiguous
"message is not modified" retry outcome after a response timeout).
Returns False only when the legacy HTML path should take over:
capability errors (server older than Bot API 10.1, which also trip the
rich latch) and content-shaped BadRequest rejections. Transport,
rate-limit, and unexpected errors propagate so the final-edit retry
contract is preserved ChannelManager retries the buffered send
instead of an immediate legacy edit doubling connection demand.
"""
if not self._app:
return False
payload: dict[str, Any] = {
"chat_id": chat_id,
"message_id": message_id,
"rich_message": {
"markdown": content,
},
}
try:
await self._call_with_retry(
self._app.bot.do_api_request,
"editMessageText",
api_kwargs=payload,
)
return True
except BadRequest as exc:
if self._is_not_modified_error(exc):
# Ambiguous success: the rich edit was applied server-side but
# its response timed out, so the retry hit "message is not
# modified". Treat it as done rather than letting the legacy
# edit overwrite the already-successful rich result.
self.logger.debug("Rich stream edit already applied for {}", chat_id)
return True
# Before Bot API 10.1, editMessageText ignores rich_message and
# reports the absent text argument instead.
pre_rich_edit_server = (
bool(content)
and str(exc).strip().lower() == "message text is empty"
)
if self._is_rich_capability_error(exc) or pre_rich_edit_server:
self.logger.debug("editMessageText rich_message not available, disabling")
self._rich_send_disabled = True
return False
# Content-shaped rejections (invalid markdown, unsupported media in
# the rich payload, …) fall back to the legacy HTML edit.
self.logger.debug("editMessageText rich_message rejected: {}", exc)
return False
except Exception:
# Transport, rate-limit, and unexpected errors propagate so the
# final-edit retry contract stays intact: ChannelManager retries
# the buffered send instead of this handler doubling connection
# demand with an immediate legacy edit.
raise
async def send(self, msg: OutboundMessage) -> None: async def send(self, msg: OutboundMessage) -> None:
"""Send a message through Telegram.""" """Send a message through Telegram."""
app = await self._wait_for_app() app = await self._wait_for_app()
@@ -1136,26 +1198,16 @@ class TelegramChannel(BaseChannel):
thread_kwargs["message_thread_id"] = message_thread_id thread_kwargs["message_thread_id"] = message_thread_id
raw_text = buf.text raw_text = buf.text
# Try sendRichMessage for final output (Bot API 10.1). # Try upgrading the streaming preview to rich in place (Bot API 10.1:
# Skip when a streaming preview already exists to avoid the # editMessageText gained a rich_message parameter). Editing in place
# delete-and-resend pattern that causes flickering and drops # keeps the message identity, so there is no delete-and-resend and
# line breaks (issue #4470). # none of the flickering / dropped line breaks from issue #4470.
if not buf.message_id and self.config.rich_messages and not getattr(self, "_rich_send_disabled", False): # The previous branch here was unreachable: it was guarded by
reply_params = None # ``not buf.message_id`` after an early return had already ensured
if reply_to_message_id := meta.get("message_id"): # ``buf.message_id`` is set (issue #5516).
reply_params = {"message_id": int(reply_to_message_id), "allow_sending_without_reply": True} if self.config.rich_messages and not getattr(self, "_rich_send_disabled", False):
rich_ok = await self._try_send_rich( rich_ok = await self._try_edit_rich(int_chat_id, buf.message_id, raw_text)
int_chat_id, raw_text, reply_params, thread_kwargs, None,
)
if rich_ok: if rich_ok:
# Delete the streaming preview message
try:
await self._call_with_retry(
app.bot.delete_message,
chat_id=int_chat_id, message_id=buf.message_id,
)
except Exception:
pass # Preview stays if delete fails
self._stream_bufs.pop(chat_id, None) self._stream_bufs.pop(chat_id, None)
return return
@@ -2735,3 +2735,130 @@ def test_markdown_to_html_code_block_same_line_no_newline() -> None:
stripped = _strip_md_block(text) stripped = _strip_md_block(text)
assert stripped == "Use <tag> here" assert stripped == "Use <tag> here"
@pytest.mark.asyncio
async def test_send_delta_stream_end_upgrades_preview_to_rich_in_place() -> None:
"""Rich messages finally work with streaming: the preview is upgraded via
editMessageText rich_message (in place), not delete-and-resend (issue #5516)."""
from telegram.error import BadRequest
channel = TelegramChannel(
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"], rich_messages=True),
MessageBus(),
)
_install_ready_app(channel)
channel._app.bot.do_api_request = AsyncMock()
channel._app.bot.edit_message_text = AsyncMock(side_effect=BadRequest("should not be reached"))
channel._stream_bufs["123"] = _StreamBuf(text="**hello**", message_id=7, last_edit=0.0)
await channel.send_delta("123", "", stream_end=True)
# editMessageText with rich_message payload, in place (same message_id)
channel._app.bot.do_api_request.assert_awaited_once()
args, kwargs = channel._app.bot.do_api_request.await_args
assert args[0] == "editMessageText"
assert kwargs["api_kwargs"]["chat_id"] == 123
assert kwargs["api_kwargs"]["message_id"] == 7
assert kwargs["api_kwargs"]["rich_message"] == {"markdown": "**hello**"}
# No delete-and-resend, no legacy HTML edit
channel._app.bot.edit_message_text.assert_not_awaited()
assert "123" not in channel._stream_bufs
@pytest.mark.asyncio
async def test_send_delta_stream_end_rich_capability_error_latches_and_falls_back() -> None:
"""On a pre-10.1 Bot API server the rich edit fails, the latch trips, and the
legacy HTML edit handles the final output."""
from telegram.error import BadRequest
channel = TelegramChannel(
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"], rich_messages=True),
MessageBus(),
)
_install_ready_app(channel)
# Before Bot API 10.1, editMessageText ignores rich_message and requires text.
channel._app.bot.do_api_request = AsyncMock(
side_effect=BadRequest("Message text is empty")
)
channel._app.bot.edit_message_text = AsyncMock()
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0)
await channel.send_delta("123", "", stream_end=True)
channel._app.bot.do_api_request.assert_awaited_once()
# Latch tripped: subsequent sends skip the rich path entirely
assert channel._rich_send_disabled is True
# Legacy HTML edit handled the final message
channel._app.bot.edit_message_text.assert_awaited_once()
assert "123" not in channel._stream_bufs
@pytest.mark.asyncio
async def test_send_delta_stream_end_rich_disabled_uses_legacy_html() -> None:
"""rich_messages=False (the default) keeps the legacy HTML path untouched."""
channel = TelegramChannel(
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
MessageBus(),
)
_install_ready_app(channel)
channel._app.bot.do_api_request = AsyncMock()
channel._app.bot.edit_message_text = AsyncMock()
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0)
await channel.send_delta("123", "", stream_end=True)
channel._app.bot.do_api_request.assert_not_called()
channel._app.bot.edit_message_text.assert_awaited_once()
assert "123" not in channel._stream_bufs
@pytest.mark.asyncio
async def test_send_delta_stream_end_rich_network_error_propagates_for_retry() -> None:
"""A transport failure on the rich edit must propagate so ChannelManager
retries the buffered send not fall through to an immediate legacy edit
that doubles connection demand during pool exhaustion."""
from telegram.error import NetworkError
channel = TelegramChannel(
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"], rich_messages=True),
MessageBus(),
)
_install_ready_app(channel)
channel._app.bot.do_api_request = AsyncMock(side_effect=NetworkError("pool exhausted"))
channel._app.bot.edit_message_text = AsyncMock()
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0)
with pytest.raises(NetworkError):
await channel.send_delta("123", "", stream_end=True)
# No legacy fallback edit: the buffered state stays for the manager retry.
channel._app.bot.edit_message_text.assert_not_awaited()
assert "123" in channel._stream_bufs
@pytest.mark.asyncio
async def test_send_delta_stream_end_rich_not_modified_after_timeout_is_success() -> None:
"""Ambiguous success: the rich edit applied server-side but its response
timed out, so the retry hit "message is not modified". That is a completed
rich upgrade the legacy edit must not overwrite it."""
from telegram.error import BadRequest, TimedOut
channel = TelegramChannel(
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"], rich_messages=True),
MessageBus(),
)
_install_ready_app(channel)
# First attempt (inside _call_with_retry) times out, retry reports the
# edit as already applied.
channel._app.bot.do_api_request = AsyncMock(
side_effect=[TimedOut(), BadRequest("Message is not modified")]
)
channel._app.bot.edit_message_text = AsyncMock(side_effect=AssertionError("must not overwrite rich result"))
channel._stream_bufs["123"] = _StreamBuf(text="hello", message_id=7, last_edit=0.0)
await channel.send_delta("123", "", stream_end=True)
assert channel._app.bot.do_api_request.await_count == 2
channel._app.bot.edit_message_text.assert_not_awaited()
assert "123" not in channel._stream_bufs
+83 -4
View File
@@ -1,12 +1,14 @@
"""Shared WebUI setup, URL, health, and browser helpers.""" """Shared WebUI setup, URL, health, and browser helpers."""
import os
import subprocess import subprocess
import sys import sys
import time import time
import webbrowser import webbrowser
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, BinaryIO
import typer import typer
from pydantic import ValidationError from pydantic import ValidationError
@@ -457,27 +459,104 @@ def _print_webui_foreground_lifecycle(*, attached: bool) -> None:
console.print("[green]WebUI is attached to the shared gateway.[/green]") console.print("[green]WebUI is attached to the shared gateway.[/green]")
console.print("[dim]Closing the browser does not stop channels or automations.[/dim]") console.print("[dim]Closing the browser does not stop channels or automations.[/dim]")
console.print( console.print(
"[dim]Press Ctrl+C to detach; the gateway stops only when the last local client exits.[/dim]" "[dim]Following live gateway logs. Press Ctrl+C to detach; the gateway stops "
"only when the last local client exits.[/dim]"
) )
_LOG_ANCHOR_BYTES = 64
@dataclass
class _GatewayLogCursor:
offset: int = 0
identity: tuple[int, int] | None = None
anchor: bytes = b""
pending: bytes = b""
def _log_anchor(handle: BinaryIO, offset: int) -> bytes:
size = min(offset, _LOG_ANCHOR_BYTES)
handle.seek(offset - size)
return handle.read(size)
def _start_gateway_log_cursor(log_path: Path) -> _GatewayLogCursor:
"""Start following at the current end of *log_path*."""
try:
with log_path.open("rb") as handle:
stat = os.fstat(handle.fileno())
offset = stat.st_size
return _GatewayLogCursor(
offset=offset,
identity=(stat.st_dev, stat.st_ino),
anchor=_log_anchor(handle, offset),
)
except OSError:
return _GatewayLogCursor()
def _read_new_gateway_logs(
log_path: Path,
cursor: _GatewayLogCursor,
*,
flush: bool = False,
) -> list[str]:
"""Read complete gateway log lines appended after *cursor*."""
try:
with log_path.open("rb") as handle:
stat = os.fstat(handle.fileno())
identity = (stat.st_dev, stat.st_ino)
reset = cursor.identity != identity or stat.st_size < cursor.offset
if not reset and cursor.offset:
reset = _log_anchor(handle, cursor.offset) != cursor.anchor
if reset:
cursor.offset = 0
cursor.pending = b""
handle.seek(cursor.offset)
chunk = handle.read()
cursor.offset = handle.tell()
cursor.identity = identity
cursor.anchor = _log_anchor(handle, cursor.offset)
except OSError:
return []
parts = (cursor.pending + chunk).split(b"\n")
cursor.pending = parts.pop()
if flush and cursor.pending:
parts.append(cursor.pending)
cursor.pending = b""
return [part.removesuffix(b"\r").decode("utf-8", errors="replace") for part in parts]
def _attach_to_background_gateway( def _attach_to_background_gateway(
runtime: "GatewayRuntime", runtime: "GatewayRuntime",
*, *,
poll_hook: Callable[[], None] | None = None, poll_hook: Callable[[], None] | None = None,
sleep: Callable[[float], None] = time.sleep, sleep: Callable[[float], None] = time.sleep,
) -> None: ) -> None:
"""Keep a WebUI launcher attached without taking ownership of the gateway.""" """Keep the launcher attached and mirror this gateway's new log output."""
status = runtime.status()
log_path = status.log_path
cursor = _start_gateway_log_cursor(log_path)
_print_webui_foreground_lifecycle(attached=True) _print_webui_foreground_lifecycle(attached=True)
try: try:
while runtime.status().running: while status.running:
for line in _read_new_gateway_logs(log_path, cursor):
console.print(line, markup=False, highlight=False)
if poll_hook is not None: if poll_hook is not None:
poll_hook() poll_hook()
sleep(0.5) sleep(0.5)
status = runtime.status()
except KeyboardInterrupt: except KeyboardInterrupt:
for line in _read_new_gateway_logs(log_path, cursor, flush=True):
console.print(line, markup=False, highlight=False)
console.print("\n[yellow]WebUI launcher detached.[/yellow]") console.print("\n[yellow]WebUI launcher detached.[/yellow]")
return return
for line in _read_new_gateway_logs(log_path, cursor, flush=True):
console.print(line, markup=False, highlight=False)
console.print("[yellow]Gateway stopped.[/yellow]") console.print("[yellow]Gateway stopped.[/yellow]")
+21 -3
View File
@@ -25,6 +25,7 @@ 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,
) )
@@ -115,8 +116,21 @@ 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:
"""Migrate legacy user cron payloads into session-bound payloads. """Make routing metadata persistable and migrate legacy user cron 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
@@ -124,8 +138,12 @@ 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 False return changed
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)
@@ -135,7 +153,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 = dict(payload.channel_meta or {}) payload.origin_metadata = _persistable_origin_metadata(payload.channel_meta or {})
payload.deliver = False payload.deliver = False
payload.channel = None payload.channel = None
+21
View File
@@ -1029,6 +1029,20 @@ 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.
@@ -1063,6 +1077,13 @@ 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:
+42 -1
View File
@@ -11,6 +11,7 @@ from nanobot.providers.base import (
ProviderCallContext, ProviderCallContext,
ProviderConversationState, ProviderConversationState,
) )
from nanobot.utils.helpers import estimate_prompt_tokens_chain
_PROVIDER_STATE_OUTPUT_META = "provider_state_output" _PROVIDER_STATE_OUTPUT_META = "provider_state_output"
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary" _PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
@@ -69,6 +70,37 @@ class ProviderConversationStateController:
session_id=self._session_id, session_id=self._session_id,
) )
def estimate_request_context_tokens(
self,
messages: list[dict[str, Any]],
*,
model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None,
tool_definitions: list[dict[str, Any]] | None = None,
) -> int | None:
"""Estimate resumed state plus the pending delta for the next request."""
state = self.checkpoint(messages, model_messages=model_messages)
if state is None:
return None
context_tokens = state.payload.get("context_tokens")
if (
isinstance(context_tokens, bool)
or not isinstance(context_tokens, int)
or context_tokens < 0
):
return None
pending_messages = [
*state.pending_messages,
*(supplemental_messages or []),
]
delta_tokens, _ = estimate_prompt_tokens_chain(
self._provider,
self._model,
pending_messages,
tool_definitions,
)
return context_tokens + max(0, delta_tokens)
def prepare_request( def prepare_request(
self, self,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
@@ -76,11 +108,20 @@ class ProviderConversationStateController:
context_window_tokens: int | None, context_window_tokens: int | None,
model_messages: list[dict[str, Any]] | None = None, model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None, supplemental_messages: list[dict[str, Any]] | None = None,
resume_state: bool = True,
) -> ProviderCallContext | None: ) -> ProviderCallContext | None:
"""Build typed context for the next request and remember its durable delta.""" """Build context for the next request and remember its durable delta.
``resume_state=False`` abandons opaque history when local request
fitting has produced a new independent model-facing context.
"""
independent_context = self.independent_request_context( independent_context = self.independent_request_context(
context_window_tokens=context_window_tokens, context_window_tokens=context_window_tokens,
) )
if not resume_state:
self._state = None
self._request_messages = []
return independent_context
if self._state is None: if self._state is None:
self._request_messages = [] self._request_messages = []
return independent_context return independent_context
+3 -7
View File
@@ -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=len(ctx.initial_messages), initial_message_count=ctx.transcript_input.message_count,
history_count=len(ctx.history),
input_persisted_early=ctx.input_persisted_early, input_persisted_early=ctx.input_persisted_early,
) )
@@ -185,7 +185,6 @@ 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."""
@@ -193,10 +192,7 @@ 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
# build_messages may merge the current message into a same-role history tail. if not input_persisted_early:
# 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
+34 -19
View File
@@ -1,27 +1,42 @@
Create a memory overview for only the final {{ archive_count }} conversation messages immediately before this instruction. Earlier messages are context for resolving references; do not summarize them again. Create a compact replacement checkpoint for this session.
Use [skip] unless a fact meets all SNIP criteria: When `[Archived Context Summary]` appears in the system prompt, update that previous checkpoint to reflect the current conversation state.
- Signal: would the user need to repeat this if forgotten?
- Novel: not just a restatement of another fact in this same conversation chunk
- Important: prevents rework or captures preferences / rules
- Persistent: still relevant after 2 weeks
Also preserve a compact working-state handoff even when it is not Persistent: the active objective, current status, completed steps, unresolved blockers, next action, and exact identifiers needed to continue without rework. Mark these facts [ephemeral]. ## Merge rules
Format each fact as: - Use the latest correction or decision as the current version of a fact, and merge duplicates.
- [mark] fact content - Preserve exact names, identifiers, paths, commands, decisions, results, and unresolved blockers when they are needed to continue the session.
- Retain a fact already present in long-term memory when it is needed for session continuity.
Marks (choose the best match): ## What to retain
- [permanent] Core preferences, personal traits, habits — never becomes stale
- [durable] Technical discoveries, project knowledge, config details — valid for months
- [ephemeral] Active task state, temporary decisions — may change in weeks
- [correction] Correction to a previous memory — state what changed
- [skip] Conversational filler, code/source facts derivable from the repo, or audit-only breadcrumbs
Priority: user corrections and preferences > solutions > decisions > events > environment facts. Always retain a compact working-state handoff:
- active objective
- current status
- completed results that constrain later work
- unresolved blockers
- next action
- exact identifiers needed for that action
Do not output facts already present in the system prompt's Recent History. Mark working-state facts `[ephemeral]`.
Do not mark something [skip] merely because it might already exist in long-term memory. For other facts, retain a candidate only when it meets all four SNIP criteria:
- Signal: remembering it saves the user from repeating it
- Novel: it adds a distinct fact to this checkpoint
- Important: losing it would cause rework or discard a preference or rule
- Persistent: it is expected to remain useful for at least two weeks
Return only formatted fact lines, or `(nothing)` if nothing noteworthy happened. Assign each retained fact its best current mark:
- `[permanent]` for core preferences, personal traits, and habits that remain relevant indefinitely
- `[durable]` for technical discoveries, project knowledge, and configuration that remains valid for months
- `[ephemeral]` for active task state and temporary decisions that may change within weeks
- `[correction]` for the current fact that supersedes conflicting earlier long-term memory
When space is limited, prioritize user corrections and preferences, then solutions, decisions, events, and environment facts.
## Output
Return one concise retained fact per line in this form:
- [mark] fact
Use `(nothing)` when no fact qualifies and there is no active working state.
+1 -1
View File
@@ -40,7 +40,7 @@
result with its original consumer or checker when one is available. result with its original consumer or checker when one is available.
- Use `apply_patch` as the default code editing tool, especially for multi-file changes, structural edits, generated code, moves, adds, or deletes. - Use `apply_patch` as the default code editing tool, especially for multi-file changes, structural edits, generated code, moves, adds, or deletes.
- Use `apply_patch dry_run=true` when the patch is uncertain and you want validation plus a change summary before writing. - Use `apply_patch dry_run=true` when the patch is uncertain and you want validation plus a change summary before writing.
- Use `edit_file` only for small exact replacements in one file, with `old_text` copied from `read_file`; when editing a specific numbered line, pass that exact line as `line_hint`; add `occurrence` or `expected_replacements` when ambiguity matters. - Use `edit_file` only for small exact replacements in one file, with `old_text` copied from `read_file`.
- Use `write_file` for new files or intentional full-file rewrites, not routine partial edits. - Use `write_file` for new files or intentional full-file rewrites, not routine partial edits.
- If `apply_patch` or `edit_file` fails, re-read with `force=true`, narrow the context, and try a smaller patch rather than switching to shell `sed` or `echo`. - If `apply_patch` or `edit_file` fails, re-read with `force=true`, narrow the context, and try a smaller patch rather than switching to shell `sed` or `echo`.
+5 -1
View File
@@ -5,6 +5,7 @@ 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
@@ -148,7 +149,10 @@ 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(
[{"role": "user", "content": "hello"}], TranscriptInput(
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,
+3 -3
View File
@@ -1302,9 +1302,9 @@ class TestSummaryPersistence:
assert "_last_summary" in reloaded.metadata assert "_last_summary" in reloaded.metadata
# Simulate /new command # Simulate /new command
session.clear() reloaded.clear()
loop.sessions.save(session) loop.sessions.save(reloaded)
loop.sessions.invalidate(session.key) loop.sessions.invalidate(reloaded.key)
# After /new, metadata should no longer contain _last_summary # After /new, metadata should no longer contain _last_summary
fresh = loop.sessions.get_or_create("cli:test") fresh = loop.sessions.get_or_create("cli:test")
+242 -43
View File
@@ -1,4 +1,4 @@
"""Tests for the lightweight Consolidator — append-only to HISTORY.md.""" """Tests for Memory checkpoint consolidation and history journaling."""
from dataclasses import replace from dataclasses import replace
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
@@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from nanobot.agent.memory import ( from nanobot.agent.memory import (
_ARCHIVE_SUMMARY_MAX_CHARS, _HISTORY_ENTRY_HARD_CAP,
Consolidator, Consolidator,
MemoryStore, MemoryStore,
) )
@@ -26,6 +26,8 @@ from nanobot.session.manager import Session
from nanobot.utils.llm_runtime import LLMRuntime from nanobot.utils.llm_runtime import LLMRuntime
from nanobot.utils.prompt_templates import render_template from nanobot.utils.prompt_templates import render_template
_ARCHIVE_PROMPT = render_template("agent/consolidator_archive.md", strip=True)
@pytest.fixture @pytest.fixture
def store(tmp_path): def store(tmp_path):
@@ -98,8 +100,15 @@ def _build_test_messages(**kwargs):
] ]
async def _archive(consolidator, messages, runtime, *, session_key="test:session"): async def _archive(
return await consolidator.archive( consolidator,
messages,
runtime,
*,
session_key="test:session",
previous_summary=None,
):
return await consolidator.archiver.archive(
messages, messages,
runtime=runtime, runtime=runtime,
session_key=session_key, session_key=session_key,
@@ -108,6 +117,7 @@ async def _archive(consolidator, messages, runtime, *, session_key="test:session
current_message="consolidate", current_message="consolidate",
), ),
request_tools=[], request_tools=[],
previous_summary=previous_summary,
) )
@@ -201,7 +211,9 @@ class TestConsolidatorSummarize:
mock_provider.chat_with_retry.side_effect = Exception("API error") mock_provider.chat_with_retry.side_effect = Exception("API error")
messages = [{"role": "user", "content": "hello"}] messages = [{"role": "user", "content": "hello"}]
result = await _archive(consolidator, messages, runtime) result = await _archive(consolidator, messages, runtime)
assert result is None # no summary on raw dump fallback assert result is not None
assert "[RAW]" in result
assert "hello" in result
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 1 assert len(entries) == 1
assert "[RAW]" in entries[0]["content"] assert "[RAW]" in entries[0]["content"]
@@ -226,23 +238,51 @@ class TestConsolidatorSummarize:
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert entries[0]["session_key"] == "slack:chat-2" assert entries[0]["session_key"] == "slack:chat-2"
async def test_raw_fallback_represents_previous_checkpoint_and_new_chunk(
self,
consolidator,
mock_provider,
runtime,
):
runtime = replace(runtime, generation=GenerationSettings(max_tokens=96))
mock_provider.chat_with_retry.side_effect = RuntimeError("API error")
result = await _archive(
consolidator,
[{"role": "user", "content": "NEW_MARKER " + "new " * 200}],
runtime,
previous_summary="OLD_MARKER " + "old " * 200,
)
assert result is not None
assert "[Previous archived context]" in result
assert "OLD_MARKER" in result
assert "[Newly archived raw context]" in result
assert "NEW_MARKER" in result
assert "... (truncated)" in result
async def test_summarize_skips_empty_messages(self, consolidator, runtime): async def test_summarize_skips_empty_messages(self, consolidator, runtime):
result = await _archive(consolidator, [], runtime) result = await _archive(consolidator, [], runtime)
assert result is None assert result is None
class TestConsolidatorPromptContract: class TestConsolidatorPromptContract:
def test_archive_prompt_preserves_working_state_with_memory_facts(self): def test_archive_prompt_requests_a_cumulative_replacement_checkpoint(self):
prompt = render_template("agent/consolidator_archive.md", strip=True, archive_count=4) prompt = _ARCHIVE_PROMPT
for section in ("## Merge rules", "## What to retain", "## Output"):
assert section in prompt
assert "replacement checkpoint" in prompt
assert "[Archived Context Summary]" in prompt
assert "current conversation state" in prompt
assert "SNIP" in prompt assert "SNIP" in prompt
assert "final 4 conversation messages" in prompt for mark in ("[permanent]", "[durable]", "[ephemeral]", "[correction]"):
for mark in ("[permanent]", "[durable]", "[ephemeral]", "[correction]", "[skip]"):
assert mark in prompt assert mark in prompt
assert "working-state handoff" in prompt assert "working-state handoff" in prompt
assert "exact identifiers needed to continue without rework" in prompt assert "- [mark] fact" in prompt
assert "Do not output facts already present in the system prompt's Recent History" in prompt assert "[skip]" not in prompt
assert "Do not mark something [skip] merely because it might already exist" in prompt assert "(nothing)" in prompt
assert "history.jsonl" not in prompt
class TestConsolidatorArchiveErrorHandling: class TestConsolidatorArchiveErrorHandling:
@@ -272,7 +312,8 @@ class TestConsolidatorArchiveErrorHandling:
{"role": "assistant", "content": "Done, fixed the race condition."}, {"role": "assistant", "content": "Done, fixed the race condition."},
] ]
result = await _archive(consolidator, messages, runtime) result = await _archive(consolidator, messages, runtime)
assert result is None assert result is not None
assert "[RAW]" in result
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 1 assert len(entries) == 1
assert "[RAW]" in entries[0]["content"] assert "[RAW]" in entries[0]["content"]
@@ -436,9 +477,9 @@ class TestConsolidatorTokenBudget:
assert [message["content"] for message in request["messages"][1:-1]] == [ assert [message["content"] for message in request["messages"][1:-1]] == [
f"m{i}" for i in range(50) f"m{i}" for i in range(50)
] ]
assert "final 50 conversation messages" in request["messages"][-1]["content"] assert request["messages"][-1]["content"] == _ARCHIVE_PROMPT
assert request["tools"] == [] assert request["tools"] == []
assert request["tool_choice"] == "none" assert "tool_choice" not in request
assert session.last_archived == 50 assert session.last_archived == 50
assert session.provider_state == _provider_state() assert session.provider_state == _provider_state()
@@ -460,8 +501,7 @@ class TestConsolidatorTokenBudget:
consolidator.estimate_session_prompt_tokens = MagicMock( consolidator.estimate_session_prompt_tokens = MagicMock(
side_effect=[(1200, "tiktoken"), (400, "tiktoken")] side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
) )
# LLM consolidation fails after raw_archive fires. consolidator.archive_session = AsyncMock(return_value="[RAW] checkpoint")
consolidator.archive_session = AsyncMock(return_value=None)
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime) await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
@@ -491,7 +531,7 @@ class TestConsolidatorTokenBudget:
consolidator.estimate_session_prompt_tokens = MagicMock( consolidator.estimate_session_prompt_tokens = MagicMock(
return_value=(1200, "tiktoken") return_value=(1200, "tiktoken")
) )
consolidator.archive_session = AsyncMock(return_value=None) consolidator.archive_session = AsyncMock(return_value="[RAW] checkpoint")
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime) await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
@@ -613,27 +653,62 @@ class TestCompactIdleSession:
assert reloaded.last_archived == 2 assert reloaded.last_archived == 2
assert [message["content"] for message in reloaded.get_history()] == ["hello", "hi"] assert [message["content"] for message in reloaded.get_history()] == ["hello", "hi"]
@pytest.mark.asyncio
async def test_idle_compaction_with_no_new_messages_is_noop(
self, real_consolidator, mock_provider, store, runtime
):
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:archived-idle")
session.add_message("user", "already archived")
session.add_message("assistant", "old answer")
session.last_archived = 2
sessions.save(session)
sessions.invalidate("cli:archived-idle")
result = await real_consolidator.compact_idle_session(
"cli:archived-idle",
runtime=runtime,
)
assert result == ""
mock_provider.chat_with_retry.assert_not_awaited()
reloaded = sessions.get_or_create("cli:archived-idle")
assert reloaded.last_archived == 2
assert "_last_summary" not in reloaded.metadata
assert store.read_unprocessed_history(since_cursor=0) == []
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_new_messages_advance_existing_archive_progress( async def test_new_messages_advance_existing_archive_progress(
self, real_consolidator, mock_provider, runtime self, real_consolidator, mock_provider, runtime
): ):
mock_provider.chat_with_retry.return_value = MagicMock( mock_provider.chat_with_retry.side_effect = [
content="Summary.", finish_reason="stop" MagicMock(content="First replacement checkpoint.", finish_reason="stop"),
) MagicMock(content="Second replacement checkpoint.", finish_reason="stop"),
]
sessions = real_consolidator.sessions sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:incremental") session = sessions.get_or_create("cli:incremental")
session.add_message("user", "first user") session.add_message("user", "first user")
session.add_message("assistant", "first assistant") session.add_message("assistant", "first assistant")
sessions.save(session) sessions.save(session)
await real_consolidator.compact_idle_session("cli:incremental", runtime=runtime) first = await real_consolidator.compact_idle_session(
"cli:incremental",
runtime=runtime,
)
current = sessions.get_or_create("cli:incremental") current = sessions.get_or_create("cli:incremental")
current.add_message("user", "second user") current.add_message("user", "second user")
current.add_message("assistant", "second assistant") current.add_message("assistant", "second assistant")
sessions.save(current) sessions.save(current)
await real_consolidator.compact_idle_session("cli:incremental", runtime=runtime) second = await real_consolidator.compact_idle_session(
"cli:incremental",
runtime=runtime,
)
assert first == "First replacement checkpoint."
assert second == "Second replacement checkpoint."
assert mock_provider.chat_with_retry.await_count == 2 assert mock_provider.chat_with_retry.await_count == 2
latest_build = real_consolidator.archiver._build_messages.call_args_list[-1].kwargs
assert latest_build["session_summary"]["text"] == "First replacement checkpoint."
latest_messages = mock_provider.chat_with_retry.await_args_list[-1].kwargs["messages"] latest_messages = mock_provider.chat_with_retry.await_args_list[-1].kwargs["messages"]
assert [message["content"] for message in latest_messages[1:5]] == [ assert [message["content"] for message in latest_messages[1:5]] == [
"first user", "first user",
@@ -641,8 +716,91 @@ class TestCompactIdleSession:
"second user", "second user",
"second assistant", "second assistant",
] ]
assert "final 2 conversation messages" in latest_messages[-1]["content"] assert latest_messages[-1]["content"] == _ARCHIVE_PROMPT
assert sessions.get_or_create("cli:incremental").last_archived == 4 sessions.invalidate("cli:incremental")
reloaded = sessions.get_or_create("cli:incremental")
assert reloaded.last_archived == 4
assert reloaded.metadata["_last_summary"]["text"] == second
@pytest.mark.asyncio
async def test_raw_fallback_preserves_previous_checkpoint_and_new_chunk(
self,
real_consolidator,
mock_provider,
store,
runtime,
):
mock_provider.chat_with_retry.side_effect = [
LLMResponse(content="Earlier durable checkpoint.", finish_reason="stop"),
RuntimeError("LLM unavailable"),
]
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:cumulative-fallback")
session.add_message("user", "first user")
session.add_message("assistant", "first answer")
sessions.save(session)
await real_consolidator.compact_idle_session(
"cli:cumulative-fallback",
runtime=runtime,
)
current = sessions.get_or_create("cli:cumulative-fallback")
current.add_message("user", "second user")
current.add_message("assistant", "newest working state")
sessions.save(current)
fallback = await real_consolidator.compact_idle_session(
"cli:cumulative-fallback",
runtime=runtime,
)
assert fallback is not None
assert "[Previous archived context]" in fallback
assert "Earlier durable checkpoint." in fallback
assert "[Newly archived raw context]" in fallback
assert "newest working state" in fallback
entries = store.read_unprocessed_history(0)
assert entries[0]["content"] == "Earlier durable checkpoint."
assert entries[1]["content"].startswith("[RAW] 2 messages")
sessions.invalidate("cli:cumulative-fallback")
reloaded = sessions.get_or_create("cli:cumulative-fallback")
assert reloaded.metadata["_last_summary"]["text"] == fallback
@pytest.mark.asyncio
async def test_nothing_keeps_previous_replacement_checkpoint(
self,
real_consolidator,
mock_provider,
runtime,
):
mock_provider.chat_with_retry.side_effect = [
LLMResponse(content="Existing checkpoint.", finish_reason="stop"),
LLMResponse(content="(nothing)", finish_reason="stop"),
]
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:nothing-after-summary")
session.add_message("user", "important first turn")
session.add_message("assistant", "important result")
sessions.save(session)
await real_consolidator.compact_idle_session(
"cli:nothing-after-summary",
runtime=runtime,
)
current = sessions.get_or_create("cli:nothing-after-summary")
current.add_message("user", "thanks")
current.add_message("assistant", "you're welcome")
sessions.save(current)
result = await real_consolidator.compact_idle_session(
"cli:nothing-after-summary",
runtime=runtime,
)
assert result == "(nothing)"
sessions.invalidate("cli:nothing-after-summary")
reloaded = sessions.get_or_create("cli:nothing-after-summary")
assert reloaded.last_archived == 4
assert reloaded.metadata["_last_summary"]["text"] == "Existing checkpoint."
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_concurrent_append_remains_unarchived( async def test_concurrent_append_remains_unarchived(
@@ -792,11 +950,16 @@ class TestCompactIdleSession:
result = await real_consolidator.compact_idle_session( result = await real_consolidator.compact_idle_session(
"cli:nothing", runtime=runtime, max_suffix=4 "cli:nothing", runtime=runtime, max_suffix=4
) )
second = await real_consolidator.compact_idle_session(
"cli:nothing", runtime=runtime, max_suffix=4
)
assert result == "(nothing)" assert result == "(nothing)"
assert second == ""
reloaded = sessions.get_or_create("cli:nothing") reloaded = sessions.get_or_create("cli:nothing")
assert "_last_summary" not in reloaded.metadata assert "_last_summary" not in reloaded.metadata
assert real_consolidator.store.read_unprocessed_history(0) == [] assert real_consolidator.store.read_unprocessed_history(0) == []
mock_provider.chat_with_retry.assert_awaited_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_llm_failure_preserves_history_but_advances_replay_boundary( async def test_llm_failure_preserves_history_but_advances_replay_boundary(
@@ -813,7 +976,8 @@ class TestCompactIdleSession:
result = await real_consolidator.compact_idle_session( result = await real_consolidator.compact_idle_session(
"cli:fail", runtime=runtime, max_suffix=4 "cli:fail", runtime=runtime, max_suffix=4
) )
assert result is None assert result is not None
assert "[RAW]" in result
# raw_archive should have been called (history.jsonl gets an entry) # raw_archive should have been called (history.jsonl gets an entry)
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
@@ -823,6 +987,7 @@ class TestCompactIdleSession:
assert len(reloaded.messages) == 20 assert len(reloaded.messages) == 20
assert reloaded.messages[0]["content"] == "u0" assert reloaded.messages[0]["content"] == "u0"
assert reloaded.last_archived == 20 assert reloaded.last_archived == 20
assert reloaded.metadata["_last_summary"]["text"] == result
assert [m["content"] for m in reloaded.get_history(max_messages=20)] == [ assert [m["content"] for m in reloaded.get_history(max_messages=20)] == [
"u6", "u6",
"a6", "a6",
@@ -863,11 +1028,10 @@ class TestCompactIdleSession:
archived_call = mock_provider.chat_with_retry.call_args archived_call = mock_provider.chat_with_retry.call_args
sent_messages = archived_call.kwargs["messages"] sent_messages = archived_call.kwargs["messages"]
sent_content = [message.get("content") for message in sent_messages] sent_content = [message.get("content") for message in sent_messages]
# The ordinary replay prefix contributes recent context, while the # The replacement overview covers all model-visible conversation context.
# temporary instruction limits the new overview to the unarchived tail.
assert "u0" not in sent_content assert "u0" not in sent_content
assert "u26" in sent_content assert "u26" in sent_content
assert "final 10 conversation messages" in sent_messages[-1]["content"] assert sent_messages[-1]["content"] == _ARCHIVE_PROMPT
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_full_archive_keeps_extended_legal_replay_suffix( async def test_full_archive_keeps_extended_legal_replay_suffix(
@@ -956,9 +1120,9 @@ class TestCompactIdleSession:
"user", "user",
] ]
assert sent_messages[2]["tool_calls"][0]["id"] == "call-1" assert sent_messages[2]["tool_calls"][0]["id"] == "call-1"
assert "final 4 conversation messages" in sent_messages[-1]["content"] assert sent_messages[-1]["content"] == _ARCHIVE_PROMPT
assert call["tools"] == tools assert call["tools"] == tools
assert call["tool_choice"] == "none" assert "tool_choice" not in call
reloaded = sessions.get_or_create("cli:tool-history") reloaded = sessions.get_or_create("cli:tool-history")
assert len(reloaded.messages) == 4 assert len(reloaded.messages) == 4
@@ -996,7 +1160,8 @@ class TestCompactIdleSession:
runtime=runtime, runtime=runtime,
) )
assert result is None assert result is not None
assert "[RAW]" in result
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 1 assert len(entries) == 1
assert entries[0]["content"].startswith("[RAW] ") assert entries[0]["content"].startswith("[RAW] ")
@@ -1026,7 +1191,8 @@ class TestCompactIdleSession:
runtime=runtime, runtime=runtime,
) )
assert result is None assert result is not None
assert "[RAW]" in result
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 1 assert len(entries) == 1
assert entries[0]["content"].startswith("[RAW] ") assert entries[0]["content"].startswith("[RAW] ")
@@ -1052,7 +1218,8 @@ class TestCompactIdleSession:
runtime=runtime, runtime=runtime,
) )
assert result is None assert result is not None
assert "[RAW]" in result
mock_provider.chat_with_retry.assert_not_awaited() mock_provider.chat_with_retry.assert_not_awaited()
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 1 assert len(entries) == 1
@@ -1060,7 +1227,7 @@ class TestCompactIdleSession:
assert sessions.get_or_create("sdk:oversized").last_archived == 1 assert sessions.get_or_create("sdk:oversized").last_archived == 1
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_incremental_scope_counts_only_model_visible_messages( async def test_archive_context_contains_only_model_visible_messages(
self, self,
real_consolidator, real_consolidator,
mock_provider, mock_provider,
@@ -1093,7 +1260,7 @@ class TestCompactIdleSession:
"new user", "new user",
"new answer", "new answer",
] ]
assert "final 2 conversation messages" in sent[-1]["content"] assert sent[-1]["content"] == _ARCHIVE_PROMPT
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_reuses_real_prefix_for_unified_session_workspace( async def test_reuses_real_prefix_for_unified_session_workspace(
@@ -1126,8 +1293,6 @@ class TestCompactIdleSession:
current_message="next project question", current_message="next project question",
channel="websocket", channel="websocket",
workspace=project, workspace=project,
session_key=session.key,
unified_session=True,
) )
await loop.consolidator.compact_idle_session( await loop.consolidator.compact_idle_session(
@@ -1137,7 +1302,7 @@ class TestCompactIdleSession:
sent_messages = runtime.provider.chat_with_retry.call_args.kwargs["messages"] sent_messages = runtime.provider.chat_with_retry.call_args.kwargs["messages"]
assert sent_messages[:-1] == ordinary_messages[:-1] assert sent_messages[:-1] == ordinary_messages[:-1]
assert "final 2 conversation messages" in sent_messages[-1]["content"] assert sent_messages[-1]["content"] == _ARCHIVE_PROMPT
system = sent_messages[0]["content"] system = sent_messages[0]["content"]
assert "PROJECT_WORKSPACE_MARKER" in system assert "PROJECT_WORKSPACE_MARKER" in system
assert "GLOBAL_WORKSPACE_MARKER" not in system assert "GLOBAL_WORKSPACE_MARKER" not in system
@@ -1307,6 +1472,21 @@ class TestRawArchiveTruncation:
assert len(entries) == 1 assert len(entries) == 1
assert "hello" in entries[0]["content"] assert "hello" in entries[0]["content"]
def test_raw_archive_returns_the_sanitized_persisted_checkpoint(self, store):
messages = [
{
"role": "user",
"content": "<think>PRIVATE_REASONING</think>visible result",
}
]
checkpoint = store.raw_archive(messages, session_key="cli:test")
persisted = store.read_unprocessed_history(since_cursor=0)[0]["content"]
assert checkpoint == persisted
assert "PRIVATE_REASONING" not in checkpoint
assert "visible result" in checkpoint
def test_raw_archive_excludes_model_only_runtime_context(self, store): def test_raw_archive_excludes_model_only_runtime_context(self, store):
content, marker = append_runtime_context( content, marker = append_runtime_context(
"ship the feature", "ship the feature",
@@ -1338,21 +1518,40 @@ class TestRawArchiveTruncation:
class TestArchivePersistence: class TestArchivePersistence:
async def test_oversized_summary_is_capped_before_append( async def test_archive_returns_the_sanitized_persisted_summary(
self, consolidator, mock_provider, store, runtime
):
mock_provider.chat_with_retry.return_value = MagicMock(
content="<think>PRIVATE_REASONING</think>safe summary",
finish_reason="stop",
has_tool_calls=False,
)
summary = await _archive(
consolidator,
[{"role": "user", "content": "hi"}],
runtime,
)
persisted = store.read_unprocessed_history(since_cursor=0)[0]["content"]
assert summary == persisted == "safe summary"
async def test_oversized_summary_uses_history_emergency_cap(
self, consolidator, mock_provider, store, runtime self, consolidator, mock_provider, store, runtime
): ):
"""A pathologically large LLM summary must not land full-length in """A pathologically large LLM summary must not land full-length in
history.jsonl that would re-open the #3412 bloat vector from the history.jsonl that would re-open the #3412 bloat vector from the
*success* path instead of the fallback path.""" *success* path instead of the fallback path."""
mock_provider.chat_with_retry.return_value = MagicMock( mock_provider.chat_with_retry.return_value = MagicMock(
content="S" * (_ARCHIVE_SUMMARY_MAX_CHARS * 10), content="S" * (_HISTORY_ENTRY_HARD_CAP * 2),
finish_reason="stop", finish_reason="stop",
) )
await _archive( summary = await _archive(
consolidator, consolidator,
[{"role": "user", "content": "hi"}], [{"role": "user", "content": "hi"}],
runtime, runtime,
) )
entry = store.read_unprocessed_history(since_cursor=0)[0] entry = store.read_unprocessed_history(since_cursor=0)[0]
assert len(entry["content"]) <= _ARCHIVE_SUMMARY_MAX_CHARS + 50 assert len(entry["content"]) <= _HISTORY_ENTRY_HARD_CAP + 50
assert summary == entry["content"]
+26 -9
View File
@@ -4,7 +4,7 @@ from pathlib import Path
import pytest import pytest
from nanobot.agent.context import ContextBuilder from nanobot.agent.context import ContextBuilder, TranscriptInput
from nanobot.runtime_context import RuntimeContextBlock from nanobot.runtime_context import RuntimeContextBlock
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -133,10 +133,7 @@ class TestLoadBootstrapFiles:
(project / "SOUL.md").write_text("project soul collision", encoding="utf-8") (project / "SOUL.md").write_text("project soul collision", encoding="utf-8")
(project / "USER.md").write_text("project user collision", encoding="utf-8") (project / "USER.md").write_text("project user collision", encoding="utf-8")
result = ContextBuilder(agent_home).build_system_prompt( result = ContextBuilder(agent_home).build_system_prompt(workspace=project)
workspace=project,
include_memory_recent_history=False,
)
assert "selected project rules" in result assert "selected project rules" in result
assert "global project rules" not in result assert "global project rules" not in result
@@ -152,10 +149,7 @@ class TestLoadBootstrapFiles:
project.mkdir() project.mkdir()
(agent_home / "AGENTS.md").write_text("default workspace rules", encoding="utf-8") (agent_home / "AGENTS.md").write_text("default workspace rules", encoding="utf-8")
result = ContextBuilder(agent_home).build_system_prompt( result = ContextBuilder(agent_home).build_system_prompt(workspace=project)
workspace=project,
include_memory_recent_history=False,
)
assert "default workspace rules" not in result assert "default workspace rules" not in result
@@ -403,6 +397,15 @@ 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)
@@ -472,6 +475,20 @@ 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(
-168
View File
@@ -3,7 +3,6 @@
from __future__ import annotations from __future__ import annotations
import datetime as datetime_module import datetime as datetime_module
import re
from datetime import datetime as real_datetime from datetime import datetime as real_datetime
from importlib.resources import files as pkg_files from importlib.resources import files as pkg_files
from pathlib import Path from pathlib import Path
@@ -104,173 +103,6 @@ def test_provider_context_appended_after_user_content(tmp_path) -> None:
assert user_pos < context_pos, "user content must precede provider context" assert user_pos < context_pos, "user content must precede provider context"
def test_unprocessed_history_injected_into_system_prompt(tmp_path) -> None:
"""Entries in history.jsonl not yet consumed by Dream appear with timestamps."""
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
builder.memory.append_history("User asked about weather in Tokyo")
builder.memory.append_history("Agent fetched forecast via web_search")
prompt = builder.build_system_prompt()
assert "# Recent History" in prompt
assert "User asked about weather in Tokyo" in prompt
assert "Agent fetched forecast via web_search" in prompt
assert re.search(r"\[\d{4}-\d{2}-\d{2} \d{2}:\d{2}\]", prompt)
def test_recent_history_injection_is_session_scoped(tmp_path) -> None:
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
builder.memory.append_history("legacy entry without session")
builder.memory.append_history("telegram history", session_key="telegram:chat-1")
builder.memory.append_history("slack history", session_key="slack:chat-2")
prompt = builder.build_system_prompt(session_key="telegram:chat-1")
assert "# Recent History" in prompt
assert "telegram history" in prompt
assert "slack history" not in prompt
assert "legacy entry without session" not in prompt
def test_session_summary_replaces_interleaved_recent_history_entry(tmp_path) -> None:
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
session_key = "unified:default"
overview = "CURRENT_SESSION_OVERVIEW_MARKER"
builder.memory.append_history("another session event", session_key=session_key)
builder.memory.append_history(overview, session_key=session_key)
latest_cursor = builder.memory.append_history(
"later telegram event",
session_key="telegram:chat-1",
)
summary = {"text": overview, "last_active": "2026-08-19T10:00:00"}
prompt = builder.build_system_prompt(
session_key=session_key,
session_summary=summary,
unified_session=True,
)
assert "# Recent History" in prompt
assert "another session event" in prompt
assert "later telegram event" in prompt
assert "[Archived Context Summary]" in prompt
assert prompt.count(overview) == 1
builder.memory.set_last_dream_cursor(latest_cursor)
processed_prompt = builder.build_system_prompt(
session_key=session_key,
session_summary=summary,
unified_session=True,
)
assert "# Recent History" not in processed_prompt
assert processed_prompt.count(overview) == 1
def test_recent_history_injection_unified_excludes_cron_internals(tmp_path) -> None:
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
builder.memory.append_history("unified user history", session_key="unified:default")
builder.memory.append_history("channel user history", session_key="telegram:chat-1")
builder.memory.append_history("cron internal history", session_key="cron:job-1")
prompt = builder.build_system_prompt(
session_key="unified:default",
unified_session=True,
)
assert "unified user history" in prompt
assert "channel user history" in prompt
assert "cron internal history" not in prompt
def test_cron_recent_history_can_see_own_history_and_unified_context(tmp_path) -> None:
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
builder.memory.append_history("unified user history", session_key="unified:default")
builder.memory.append_history("own cron history", session_key="cron:job-1")
builder.memory.append_history("other cron history", session_key="cron:job-2")
prompt = builder.build_system_prompt(
session_key="cron:job-1",
unified_session=True,
)
assert "unified user history" in prompt
assert "own cron history" in prompt
assert "other cron history" not in prompt
def test_recent_history_capped_at_max(tmp_path) -> None:
"""Only the most recent _MAX_RECENT_HISTORY entries are injected."""
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
for i in range(builder._MAX_RECENT_HISTORY + 20):
builder.memory.append_history(f"entry-{i}")
prompt = builder.build_system_prompt()
assert "entry-0" not in prompt
assert "entry-19" not in prompt
assert f"entry-{builder._MAX_RECENT_HISTORY + 19}" in prompt
def test_recent_history_truncated_at_max_tokens(tmp_path) -> None:
"""Recent History section must be truncated to _MAX_HISTORY_TOKENS."""
import tiktoken
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
big_entry = "word " * (builder._MAX_HISTORY_TOKENS + 5_000)
builder.memory.append_history(big_entry)
prompt = builder.build_system_prompt()
history_section = prompt.split("# Recent History\n\n", 1)
assert len(history_section) == 2
enc = tiktoken.get_encoding("cl100k_base")
assert len(enc.encode(history_section[1])) <= builder._MAX_HISTORY_TOKENS
def test_no_recent_history_when_dream_has_processed_all(tmp_path) -> None:
"""If Dream has consumed everything, no Recent History section should appear."""
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
cursor = builder.memory.append_history("already processed entry")
builder.memory.set_last_dream_cursor(cursor)
prompt = builder.build_system_prompt()
assert "# Recent History" not in prompt
def test_partial_dream_processing_shows_only_remainder(tmp_path) -> None:
"""When Dream has processed some entries, only the unprocessed ones appear."""
workspace = _make_workspace(tmp_path)
builder = ContextBuilder(workspace)
builder.memory.append_history("old conversation about Python")
c2 = builder.memory.append_history("old conversation about Rust")
builder.memory.append_history("recent question about Docker")
builder.memory.append_history("recent question about K8s")
builder.memory.set_last_dream_cursor(c2)
prompt = builder.build_system_prompt()
assert "# Recent History" in prompt
assert "old conversation about Python" not in prompt
assert "old conversation about Rust" not in prompt
assert "recent question about Docker" in prompt
assert "recent question about K8s" in prompt
def test_execution_rules_in_system_prompt(tmp_path) -> None: def test_execution_rules_in_system_prompt(tmp_path) -> None:
"""Execution rules should appear in the system prompt via the default templates.""" """Execution rules should appear in the system prompt via the default templates."""
from nanobot.utils.helpers import sync_workspace_templates from nanobot.utils.helpers import sync_workspace_templates
+3 -3
View File
@@ -426,7 +426,7 @@ class TestEphemeralDirect:
bus=bus, bus=bus,
provider=provider, provider=provider,
workspace=tmp_path, workspace=tmp_path,
context_window_tokens=8000, context_window_tokens=32_000,
) )
return loop, store return loop, store
@@ -606,7 +606,7 @@ class TestEphemeralDirect:
bus=MessageBus(), bus=MessageBus(),
provider=provider, provider=provider,
workspace=tmp_path, workspace=tmp_path,
context_window_tokens=8000, context_window_tokens=32_000,
) )
await loop.process_direct( await loop.process_direct(
@@ -666,7 +666,7 @@ class TestEphemeralHooks:
bus=bus, bus=bus,
provider=provider, provider=provider,
workspace=tmp_path, workspace=tmp_path,
context_window_tokens=8000, context_window_tokens=32_000,
hooks=[spy], hooks=[spy],
) )
+9 -5
View File
@@ -6,6 +6,7 @@ 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,
@@ -459,7 +460,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(
[{"role": "user", "content": "hi"}], TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
) )
@@ -504,7 +505,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(
[{"role": "user", "content": "hi"}], TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
runtime=runtime, runtime=runtime,
on_progress=on_progress, on_progress=on_progress,
request_context=RequestContext( request_context=RequestContext(
@@ -551,7 +552,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(
[{"role": "user", "content": "hi"}], TranscriptInput(history=[{"role": "user", "content": "hi"}], current_message=None),
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
) )
@@ -577,7 +578,9 @@ 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(
[], runtime=loop.llm_runtime(), on_progress=bad_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=bad_progress,
) )
@@ -596,7 +599,8 @@ 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(
[], runtime=loop.llm_runtime() TranscriptInput(history=[], current_message=None),
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) "
+39 -2
View File
@@ -7,11 +7,17 @@ from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse from nanobot.providers.base import LLMResponse
def _make_loop(tmp_path, *, estimated_tokens: int, context_window_tokens: int) -> AgentLoop: def _make_loop(
tmp_path,
*,
estimated_tokens: int,
context_window_tokens: int,
max_tokens: int = 0,
) -> AgentLoop:
from nanobot.providers.base import GenerationSettings from nanobot.providers.base import GenerationSettings
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
provider.generation = GenerationSettings(max_tokens=0) provider.generation = GenerationSettings(max_tokens=max_tokens)
provider.estimate_prompt_tokens.return_value = (estimated_tokens, "test-counter") provider.estimate_prompt_tokens.return_value = (estimated_tokens, "test-counter")
_response = LLMResponse(content="ok", tool_calls=[]) _response = LLMResponse(content="ok", tool_calls=[])
provider.chat_with_retry = AsyncMock(return_value=_response) provider.chat_with_retry = AsyncMock(return_value=_response)
@@ -23,6 +29,9 @@ def _make_loop(tmp_path, *, estimated_tokens: int, context_window_tokens: int) -
workspace=tmp_path, workspace=tmp_path,
model="test-model", model="test-model",
context_window_tokens=context_window_tokens, context_window_tokens=context_window_tokens,
# These tests isolate Memory consolidation; Runner request fitting is
# covered separately with realistic context windows.
context_block_limit=10_000,
) )
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator._SAFETY_BUFFER = 0 loop.consolidator._SAFETY_BUFFER = 0
@@ -56,6 +65,34 @@ async def test_prompt_above_threshold_triggers_consolidation(tmp_path) -> None:
assert loop.consolidator.archive_session.await_count >= 1 assert loop.consolidator.archive_session.await_count >= 1
@pytest.mark.asyncio
async def test_token_consolidation_refreshes_summary_for_current_request(tmp_path) -> None:
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
loop.consolidator.archive_session = AsyncMock( # type: ignore[method-assign]
return_value="FRESH_CHECKPOINT"
)
loop.consolidator.estimate_session_prompt_tokens = MagicMock( # type: ignore[method-assign]
return_value=(1000, "test")
)
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test")
session.messages = [
{"role": role, "content": f"{role[0]}{turn}"}
for turn in range(10)
for role in ("user", "assistant")
]
loop.sessions.save(session)
await loop.process_direct("hello", session_key="cli:test")
request_messages = loop.provider.chat_with_retry.await_args.kwargs["messages"]
system_prompt = request_messages[0]["content"]
assert "FRESH_CHECKPOINT" in system_prompt
assert all(message.get("content") != "u0" for message in request_messages)
assert loop.sessions.get_or_create("cli:test").last_archived == 12
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None: async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None:
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200) loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
+14 -5
View File
@@ -6,6 +6,7 @@ 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
@@ -84,7 +85,9 @@ 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(
[], runtime=loop.llm_runtime(), on_progress=on_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=on_progress,
) )
assert result.final_content == "Done" assert result.final_content == "Done"
@@ -155,7 +158,9 @@ 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(
[], runtime=loop.llm_runtime(), on_progress=on_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=on_progress,
) )
assert result.final_content == "Done" assert result.final_content == "Done"
@@ -225,7 +230,9 @@ class TestToolEventProgress:
) )
result = await loop._run_agent_loop( result = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=on_progress,
) )
assert result.final_content == "Done" assert result.final_content == "Done"
@@ -263,7 +270,9 @@ class TestToolEventProgress:
file_events.extend(file_edit_events) file_events.extend(file_edit_events)
await loop._run_agent_loop( await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=on_progress,
) )
assert file_events == [] assert file_events == []
@@ -1019,7 +1028,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,
+14 -7
View File
@@ -7,6 +7,7 @@ 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
@@ -55,7 +56,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)],
@@ -340,7 +341,8 @@ 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(
[], runtime=loop.llm_runtime() TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
) )
assert result.final_content == ( assert result.final_content == (
@@ -362,7 +364,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",
@@ -401,7 +403,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,
@@ -428,7 +430,9 @@ 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(
[], runtime=loop.llm_runtime(), on_stream=on_stream TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_stream=on_stream,
) )
assert result.final_content == "Hello World" assert result.final_content == "Hello World"
@@ -451,7 +455,9 @@ 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(
[], runtime=loop.llm_runtime(), on_stream=on_stream TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_stream=on_stream,
) )
assert result.final_content == "Hello World" assert result.final_content == "Hello World"
@@ -472,7 +478,8 @@ 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(
[], runtime=loop.llm_runtime() TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
) )
assert result.final_content == "Recovered answer" assert result.final_content == "Recovered answer"
+46 -23
View File
@@ -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 from nanobot.agent.context import ContextBuilder, TranscriptInput
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,6 +79,13 @@ 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
@@ -930,10 +937,13 @@ 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(
{"role": "system", "content": "system"}, history=[
{"role": "user", "content": "question"}, {"role": "system", "content": "system"},
], {"role": "user", "content": "question"},
],
current_message=None,
),
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
session=session, session=session,
) )
@@ -1008,7 +1018,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._build_initial_messages = MagicMock( # type: ignore[method-assign] loop.context.build_system_prompt = 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")
@@ -1041,8 +1051,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_initial_messages = loop._build_initial_messages build_system_prompt = loop.context.build_system_prompt
loop._build_initial_messages = MagicMock( # type: ignore[method-assign] loop.context.build_system_prompt = 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")
@@ -1066,7 +1076,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._build_initial_messages = build_initial_messages # type: ignore[method-assign] loop.context.build_system_prompt = build_system_prompt # 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"),
) )
@@ -1319,7 +1329,8 @@ async def test_internal_continuation_queues_turn_without_fake_user_history(
calls: list[dict] = [] calls: list[dict] = []
async def fake_run_agent_loop(initial_messages, *, metadata=None, **_kwargs): async def fake_run_agent_loop(transcript_input, *, 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(
@@ -1387,8 +1398,9 @@ async def test_internal_continuation_preserves_streaming_route_metadata(
calls = 0 calls = 0
async def fake_run_agent_loop(initial_messages, *, on_stream=None, on_stream_end=None, **_kwargs): async def fake_run_agent_loop(transcript_input, *, 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(
@@ -1460,8 +1472,9 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
calls = 0 calls = 0
async def fake_run_agent_loop(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_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(
@@ -1623,7 +1636,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(
@@ -1753,7 +1766,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(_initial_messages, *, session=None, **_kwargs): async def interrupted_run_agent_loop(_transcript_input, *, session=None, **_kwargs):
assert session is not None assert session is not None
loop._set_runtime_checkpoint( loop._set_runtime_checkpoint(
session, session,
@@ -1813,7 +1826,8 @@ 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(initial_messages, **_kwargs): async def resumed_run_agent_loop(transcript_input, **_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"}],
@@ -1864,7 +1878,8 @@ 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(initial_messages, **kwargs): async def fake_run_agent_loop(transcript_input, **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"]
@@ -1940,7 +1955,8 @@ 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(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_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"}],
@@ -1966,7 +1982,8 @@ 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(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_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"}],
@@ -2022,7 +2039,8 @@ 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(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_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"}],
@@ -2065,7 +2083,8 @@ 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(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_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"}],
@@ -2196,7 +2215,8 @@ 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(initial_messages, **kwargs): async def fake_run_agent_loop(transcript_input, **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(
@@ -2252,8 +2272,11 @@ 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(initial_messages, **_kwargs): async def fake_run_agent_loop(transcript_input, **_kwargs):
assert [m["role"] for m in initial_messages] == ["system", "user"] initial_messages = _assembled_messages(loop.context, transcript_input)
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",
[ [
+3 -2
View File
@@ -5,6 +5,7 @@ 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,
@@ -133,7 +134,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",
@@ -234,7 +235,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",
-48
View File
@@ -113,54 +113,6 @@ class TestHistoryWithCursor:
entries = store.read_unprocessed_history(since_cursor=0) entries = store.read_unprocessed_history(since_cursor=0)
assert len(entries) == 2 assert len(entries) == 2
def test_prompt_history_filters_to_current_session(self, store):
store.append_history("legacy entry without session")
store.append_history("telegram entry", session_key="telegram:chat-1")
store.append_history("slack entry", session_key="slack:chat-2")
entries = store.read_recent_history_for_prompt(
since_cursor=0,
session_key="telegram:chat-1",
)
assert [e["content"] for e in entries] == ["telegram entry"]
assert [e["content"] for e in store.read_unprocessed_history(0)] == [
"legacy entry without session",
"telegram entry",
"slack entry",
]
def test_unified_prompt_history_excludes_internal_cron_sessions(self, store):
store.append_history("legacy entry without session")
store.append_history("unified entry", session_key="unified:default")
store.append_history("telegram entry", session_key="telegram:chat-1")
store.append_history("cron internal entry", session_key="cron:job-1")
entries = store.read_recent_history_for_prompt(
since_cursor=0,
session_key="unified:default",
unified_session=True,
)
assert [e["content"] for e in entries] == [
"legacy entry without session",
"unified entry",
"telegram entry",
]
def test_unified_cron_prompt_history_includes_own_cron_entry(self, store):
store.append_history("unified entry", session_key="unified:default")
store.append_history("other cron entry", session_key="cron:job-2")
store.append_history("own cron entry", session_key="cron:job-1")
entries = store.read_recent_history_for_prompt(
since_cursor=0,
session_key="cron:job-1",
unified_session=True,
)
assert [e["content"] for e in entries] == ["unified entry", "own cron entry"]
def test_read_unprocessed_skips_entries_without_cursor(self, store): def test_read_unprocessed_skips_entries_without_cursor(self, store):
"""Regression: entries missing the cursor key should be silently skipped.""" """Regression: entries missing the cursor key should be silently skipped."""
store.history_file.write_text( store.history_file.write_text(
+5 -1
View File
@@ -8,6 +8,10 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from nanobot.utils.prompt_templates import render_template
_ARCHIVE_PROMPT = render_template("agent/consolidator_archive.md", strip=True)
class TestNewCommandArchival: class TestNewCommandArchival:
"""Test /new archival behavior with the structured archive flow.""" """Test /new archival behavior with the structured archive flow."""
@@ -117,7 +121,7 @@ class TestNewCommandArchival:
await loop.aclose() await loop.aclose()
sent = loop.provider.chat_with_retry.call_args.kwargs["messages"] sent = loop.provider.chat_with_retry.call_args.kwargs["messages"]
assert sent[1:-1] == ordinary_history assert sent[1:-1] == ordinary_history
assert "final 2 conversation messages" in sent[-1]["content"] assert sent[-1]["content"] == _ARCHIVE_PROMPT
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_new_clears_session_and_responds(self, tmp_path: Path) -> None: async def test_new_clears_session_and_responds(self, tmp_path: Path) -> None:
+60 -29
View File
@@ -10,6 +10,8 @@ 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.agent.context_governance import ContextWindowExceededError
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import ( from nanobot.providers.base import (
LLMProvider, LLMProvider,
@@ -34,6 +36,35 @@ 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
@@ -56,6 +87,7 @@ def test_usage_or_estimate_replaces_reported_zero_for_content(monkeypatch) -> No
_make_usage_spec(provider, tools), _make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
response, response,
tool_definitions=tools.get_definitions(),
) )
assert usage == LLMUsage.estimated(input_tokens=12, output_tokens=7).with_timing( assert usage == LLMUsage.estimated(input_tokens=12, output_tokens=7).with_timing(
@@ -100,6 +132,7 @@ def test_usage_or_estimate_counts_tool_call_output_for_reported_zero(monkeypatch
_make_usage_spec(provider, tools), _make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
response, response,
tool_definitions=tools.get_definitions(),
) )
assert usage == LLMUsage.estimated(input_tokens=13, output_tokens=9) assert usage == LLMUsage.estimated(input_tokens=13, output_tokens=9)
@@ -132,6 +165,7 @@ def test_usage_or_estimate_counts_error_without_estimating_tokens(
_make_usage_spec(provider, tools), _make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
response, response,
tool_definitions=tools.get_definitions(),
) )
assert usage is not None assert usage is not None
@@ -167,6 +201,7 @@ def test_usage_or_estimate_trusts_positive_reported_total(monkeypatch) -> None:
_make_usage_spec(provider, tools), _make_usage_spec(provider, tools),
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
response, response,
tool_definitions=tools.get_definitions(),
) )
assert usage is not None assert usage is not None
@@ -336,14 +371,12 @@ async def test_runner_replays_provider_state_without_chat_projection_duplicates(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runner_governs_tool_result_before_adding_it_to_provider_state(): async def test_runner_preserves_tool_result_before_rejecting_unfit_followup():
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider) provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
calls = 0 calls = 0
captured_context: ProviderCallContext | None = None
checkpoints: list[dict] = [] checkpoints: list[dict] = []
state = ProviderConversationState( state = ProviderConversationState(
kind="openai_responses", kind="openai_responses",
@@ -354,7 +387,7 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
) )
async def chat_with_retry(**kwargs): async def chat_with_retry(**kwargs):
nonlocal calls, captured_context nonlocal calls
calls += 1 calls += 1
if calls == 1: if calls == 1:
return LLMResponse( return LLMResponse(
@@ -368,7 +401,6 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
], ],
provider_state=state, provider_state=state,
) )
captured_context = kwargs["provider_context"]
return LLMResponse(content="done") return LLMResponse(content="done")
provider.chat_with_retry = chat_with_retry provider.chat_with_retry = chat_with_retry
@@ -379,37 +411,36 @@ async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
async def checkpoint(payload: dict) -> None: async def checkpoint(payload: dict) -> None:
checkpoints.append(payload) checkpoints.append(payload)
await AgentRunner().run(make_run_spec( with pytest.raises(ContextWindowExceededError):
provider, await AgentRunner().run(make_run_spec(
initial_messages=[ provider,
{"role": "system", "content": "system"}, initial_messages=[
{"role": "user", "content": "read the file"}, {"role": "system", "content": "system"},
], {"role": "user", "content": "read the file"},
tools=tools, ],
model="gpt-5.6", tools=tools,
context_window_tokens=3_000, model="gpt-5.6",
context_block_limit=200, context_window_tokens=3_000,
max_tokens=1_000, context_block_limit=200,
max_iterations=3, max_tokens=1_000,
max_tool_result_chars=10_000, max_iterations=3,
checkpoint_callback=checkpoint, max_tool_result_chars=10_000,
)) checkpoint_callback=checkpoint,
))
assert captured_context is not None assert calls == 1
assert captured_context.conversation_state is not None
pending = captured_context.conversation_state.pending_messages
assert len(pending) == 1
assert pending[0]["role"] == "tool"
assert "compacted to fit context" in pending[0]["content"]
assert pending[0]["content"] != "x" * 5_000
completed_checkpoint = next( completed_checkpoint = next(
checkpoint checkpoint
for checkpoint in checkpoints for checkpoint in checkpoints
if checkpoint["phase"] == "tools_completed" if checkpoint["phase"] == "tools_completed"
) )
checkpoint_pending = completed_checkpoint["provider_state"].pending_messages checkpoint_pending = completed_checkpoint["provider_state"].pending_messages
assert "compacted to fit context" in checkpoint_pending[0]["content"] assert checkpoint_pending == [{
assert checkpoint_pending[0]["content"] != "x" * 5_000 "role": "tool",
"tool_call_id": "call_1",
"name": "read_file",
"content": "x" * 5_000,
}]
@pytest.mark.asyncio @pytest.mark.asyncio
+24
View File
@@ -52,7 +52,31 @@ 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
+526 -264
View File
@@ -1,8 +1,7 @@
"""Tests for AgentRunner context governance: backfill, orphan cleanup, microcompact, snip_history.""" """Tests for AgentRunner context governance: repair and request fitting."""
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
@@ -12,11 +11,14 @@ from nanobot.agent.context_governance import (
BACKFILL_CONTENT, BACKFILL_CONTENT,
ContextGovernanceConfig, ContextGovernanceConfig,
ContextGovernor, ContextGovernor,
ContextWindowExceededError,
) )
from nanobot.agent.runner import AgentRunSpec from nanobot.agent.runner import AgentRunSpec
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import ( from nanobot.providers.base import (
LLMProvider,
LLMResponse, LLMResponse,
LLMUsage,
ProviderConversationState, ProviderConversationState,
ToolCallRequest, ToolCallRequest,
) )
@@ -28,8 +30,6 @@ def _governance_config(
provider, provider,
tools, tools,
spec: AgentRunSpec, spec: AgentRunSpec,
*,
inflight_start_index: int = 0,
) -> ContextGovernanceConfig: ) -> ContextGovernanceConfig:
return ContextGovernanceConfig( return ContextGovernanceConfig(
provider=provider, provider=provider,
@@ -41,7 +41,6 @@ def _governance_config(
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=inflight_start_index,
) )
@@ -89,6 +88,508 @@ async def test_runner_propagates_context_governance_failure():
provider.chat_with_retry.assert_not_awaited() provider.chat_with_retry.assert_not_awaited()
@pytest.mark.asyncio
async def test_runner_locally_fits_oversized_initial_transcript(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="done"))
tools = MagicMock()
tools.get_definitions.return_value = []
old_content = "x" * 20_000
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda _provider, _model, messages, _tools: (
(600, "test-counter")
if any(message.get("content") == old_content for message in messages)
else (100, "test-counter")
),
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
{"role": "user", "content": "old question"},
{"role": "assistant", "content": old_content},
{"role": "user", "content": "continue"},
],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert provider.chat_with_retry.await_args.kwargs["messages"] == [
{"role": "system", "content": "system"},
{"role": "user", "content": "continue"},
]
assert any(message.get("content") == old_content for message in result.messages)
@pytest.mark.asyncio
async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatch):
from nanobot.agent.hook import AgentHook, AgentHookContext
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="unexpected"))
tools = MagicMock()
tools.get_definitions.return_value = []
oversized = "hook-added-oversized-message"
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda _provider, _model, messages, _tools: (
(2_000, "test-counter")
if any(message.get("content") == oversized for message in messages)
else (100, "test-counter")
),
)
class MutatingHook(AgentHook):
async def before_iteration(self, context: AgentHookContext) -> None:
context.messages.append({"role": "user", "content": oversized})
with pytest.raises(ContextWindowExceededError):
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "hello"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
hook=MutatingHook(),
))
provider.chat_with_retry.assert_not_awaited()
@pytest.mark.asyncio
async def test_runner_drops_resumable_provider_state_when_request_is_fitted(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
captured_contexts = []
old_content = "old-oversized-history"
candidate = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="local-model",
version=1,
payload={"items": [{"type": "message", "content": "fresh state"}]},
)
async def chat_with_retry(*, provider_context=None, **_kwargs):
captured_contexts.append(provider_context)
return LLMResponse(
content="done",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
provider_state=candidate,
)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda _provider, _model, messages, _tools: (
(600, "test-counter")
if any(message.get("content") == old_content for message in messages)
else (100, "test-counter")
),
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_message_tokens",
lambda message: 450 if message.get("content") == old_content else 50,
)
saved_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="local-model",
version=1,
payload={"items": [{"type": "message", "content": "stale state"}]},
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "assistant", "content": old_content},
{"role": "user", "content": "continue"},
],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=saved_state,
))
assert captured_contexts[0].conversation_state is None
assert result.provider_state is not None
assert result.provider_state.payload == candidate.payload
@pytest.mark.asyncio
async def test_runner_fits_each_malformed_retry_with_its_actual_tools(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
calls: list[dict] = []
estimated_tools: list[object] = []
definitions = [{"type": "function", "function": {"name": "read_file"}}]
async def chat_with_retry(*, messages, tools=None, **_kwargs):
calls.append({"messages": [dict(message) for message in messages], "tools": tools})
if len(calls) < 3:
return LLMResponse(
content="bad tool request",
tool_calls=[ToolCallRequest(id=f"bad_{len(calls)}", name=None, arguments={})],
finish_reason="tool_calls",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
return LLMResponse(
content="recovered",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
def estimate(_provider, _model, messages, _tools):
estimated_tools.append(_tools)
user_count = sum(message.get("role") == "user" for message in messages)
return (600 if user_count > 1 else 100), "test-counter"
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = definitions
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_message_tokens",
lambda _message: 300,
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "use a tool"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert [call["tools"] for call in calls] == [definitions, definitions, None]
assert definitions in estimated_tools
assert None in estimated_tools
assert [len(call["messages"]) for call in calls] == [1, 1, 1]
assert result.final_content == "recovered"
assert result.messages == [
{"role": "user", "content": "use a tool"},
{"role": "assistant", "content": "recovered"},
]
@pytest.mark.asyncio
async def test_runner_fits_empty_response_finalization_before_dispatch(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
calls: list[dict] = []
async def chat_with_retry(*, messages, tools=None, **_kwargs):
calls.append({"messages": [dict(message) for message in messages], "tools": tools})
if len(calls) < 3:
return LLMResponse(
content=None,
usage=LLMUsage.reported(input_tokens=100, output_tokens=1),
)
return LLMResponse(
content="finalized",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
def estimate(_provider, _model, messages, _tools):
contents = [str(message.get("content") or "") for message in messages]
has_original = "do task" in contents
has_finalization = any("conversation above" in content for content in contents)
return (600 if has_original and has_finalization else 100), "test-counter"
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_message_tokens",
lambda _message: 300,
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert len(calls) == 3
assert calls[-1]["tools"] is None
assert all(message.get("content") != "do task" for message in calls[-1]["messages"])
assert result.final_content == "finalized"
@pytest.mark.asyncio
async def test_runner_fits_max_iteration_finalization_before_dispatch(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
calls: list[dict] = []
oversized_result = "oversized-current-tool-result"
async def chat_with_retry(*, messages, tools=None, **_kwargs):
calls.append({"messages": [dict(message) for message in messages], "tools": tools})
if len(calls) == 1:
return LLMResponse(
content="working",
tool_calls=[ToolCallRequest(id="call_1", name="read_file", arguments={})],
finish_reason="tool_calls",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
return LLMResponse(
content="safe summary",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
def estimate(_provider, _model, messages, _tools):
has_oversized = any(
message.get("content") == oversized_result for message in messages
)
return (600 if has_oversized else 100), "test-counter"
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value=oversized_result)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_message_tokens",
lambda message: 600 if message.get("content") == oversized_result else 50,
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "inspect"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert len(calls) == 2
assert calls[-1]["tools"] is None
assert all(
message.get("content") != oversized_result
for message in calls[-1]["messages"]
)
assert any(message.get("content") == oversized_result for message in result.messages)
assert result.final_content == "safe summary"
@pytest.mark.parametrize(
("input_tokens", "expected_fitted"),
[(500, True), (100, False)],
)
def test_matching_reported_provider_usage_avoids_local_estimate(
monkeypatch,
input_tokens,
expected_fitted,
):
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
tools.get_definitions.return_value = []
spec = make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "hello"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (_ for _ in ()).throw(
AssertionError("matching provider usage must be authoritative")
),
)
governor = ContextGovernor()
monkeypatch.setattr(governor, "fit_to_budget", lambda *_args, **_kwargs: [])
_messages, fitted = governor.fit_request(
_governance_config(provider, tools, spec),
spec.initial_messages,
LLMUsage.reported(input_tokens=input_tokens, output_tokens=10),
usage_matches_messages=True,
tool_definitions=tools.get_definitions(),
)
assert fitted is expected_fitted
def test_changed_messages_use_local_estimate_after_reported_usage(monkeypatch):
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
tools.get_definitions.return_value = []
spec = make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "new tool output"}],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
)
estimate = MagicMock(return_value=(600, "test-counter"))
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
governor = ContextGovernor()
monkeypatch.setattr(governor, "fit_to_budget", lambda *_args, **_kwargs: [])
_messages, fitted = governor.fit_request(
_governance_config(provider, tools, spec),
spec.initial_messages,
LLMUsage.reported(input_tokens=900, output_tokens=10),
usage_matches_messages=False,
tool_definitions=tools.get_definitions(),
)
assert fitted is True
estimate.assert_called_once()
@pytest.mark.asyncio
async def test_runner_counts_resumed_provider_state_before_dispatch(monkeypatch):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
captured_contexts = []
async def chat_with_retry(*, provider_context=None, **_kwargs):
captured_contexts.append(provider_context)
return LLMResponse(
content="done",
usage=LLMUsage.reported(input_tokens=100, output_tokens=10),
)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
current_message = {"role": "user", "content": "new delta"}
saved_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="local-model",
version=1,
payload={
"items": [{"type": "reasoning", "encrypted_content": "opaque"}],
"context_tokens": 450,
},
pending_messages=[current_message],
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (100, "test-counter"),
)
monkeypatch.setattr(
"nanobot.providers.conversation_state.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (100, "test-counter"),
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[current_message],
tools=tools,
model="local-model",
context_window_tokens=2_000,
context_block_limit=500,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=saved_state,
))
assert captured_contexts[0].conversation_state is None
assert result.messages == [
current_message,
{"role": "assistant", "content": "done"},
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("context_block_limit", "expected_budget"),
[(500, 500), (None, 0)],
)
async def test_runner_refuses_locally_fitted_request_that_still_cannot_fit(
monkeypatch,
context_block_limit,
expected_budget,
):
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="unexpected"))
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (2_000, "test-counter"),
)
with pytest.raises(ContextWindowExceededError) as exc_info:
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "oversized system"},
{"role": "user", "content": "oversized user"},
],
tools=tools,
model="local-model",
context_window_tokens=1_000,
context_block_limit=context_block_limit,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert exc_info.value.estimated_tokens == 2_000
assert exc_info.value.input_budget == expected_budget
provider.chat_with_retry.assert_not_awaited()
def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch): def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch):
provider = MagicMock() provider = MagicMock()
tools = MagicMock() tools = MagicMock()
@@ -130,7 +631,11 @@ def test_snip_history_drops_orphaned_tool_results_from_trimmed_slice(monkeypatch
lambda msg: token_sizes.get(str(msg.get("content")), 40), lambda msg: token_sizes.get(str(msg.get("content")), 40),
) )
trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) trimmed = ContextGovernor().snip_history(
_governance_config(provider, tools, spec),
messages,
tool_definitions=tools.get_definitions(),
)
# After the fix, the user message is recovered so the sequence is valid # After the fix, the user message is recovered so the sequence is valid
# for providers that require system → user (e.g. GLM error 1214). # for providers that require system → user (e.g. GLM error 1214).
@@ -182,7 +687,11 @@ def test_snip_history_reserves_budget_for_tool_definitions(monkeypatch):
lambda msg: token_sizes.get(str(msg.get("content")), 40), lambda msg: token_sizes.get(str(msg.get("content")), 40),
) )
trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) trimmed = ContextGovernor().snip_history(
_governance_config(provider, tools, spec),
messages,
tool_definitions=tools.get_definitions(),
)
contents = [message.get("content") for message in trimmed] contents = [message.get("content") for message in trimmed]
assert contents == ["system", "recent two"] assert contents == ["system", "recent two"]
@@ -465,260 +974,6 @@ async def test_runner_backfill_only_mutates_model_context_not_returned_messages(
] ]
# ---------------------------------------------------------------------------
# Microcompact (stale tool result compaction)
# ---------------------------------------------------------------------------
def _microcompact_messages(*, total: int, tool_name: str, content: str) -> list[dict]:
messages: list[dict] = [{"role": "system", "content": "sys"}]
for i in range(total):
messages.append({
"role": "assistant",
"content": "",
"tool_calls": [{
"id": f"c{i}",
"type": "function",
"function": {"name": tool_name, "arguments": "{}"},
}],
})
messages.append({
"role": "tool",
"tool_call_id": f"c{i}",
"name": tool_name,
"content": content,
})
return messages
def test_microcompact_skips_when_prompt_under_hard_budget(monkeypatch):
"""Cache-friendly path: in-flight tool results stay stable while prompt fits."""
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
total = 15
long_content = "x" * 600
messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content)
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=20_000,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (1000, "test"),
)
result = ContextGovernor().compact_inflight_overflow(
_governance_config(provider, tools, spec),
messages,
set(),
)
assert result is messages
def test_microcompact_overflow_compacts_to_low_watermark(monkeypatch):
"""Overflow path: compact in-flight stale results with headroom for later calls."""
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
total = 18
long_content = "x" * 600
messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content)
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=2224, # input budget 1200, low target 1020
)
def estimate(_provider, _model, msgs, _tools):
return sum(
100 if (content := msg.get("content")) == long_content
else 1 if isinstance(content, str) and "compacted to fit context" in content
else 0
for msg in msgs
if msg.get("role") == "tool"
), "test"
monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate)
result = ContextGovernor().compact_inflight_overflow(
_governance_config(provider, tools, spec),
messages,
set(),
)
tool_msgs = [m for m in result if m.get("role") == "tool"]
compacted = [m for m in tool_msgs if "compacted to fit context" in str(m.get("content", ""))]
preserved = [m for m in tool_msgs if m.get("content") == long_content]
assert len(compacted) == 8
assert len(preserved) == total - 8
assert [m["tool_call_id"] for m in compacted] == [f"c{i}" for i in range(8)]
def test_microcompact_compacts_newest_when_it_alone_overflows(monkeypatch):
"""An unfit newest result tells the model to retry narrowly or report the limit."""
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
long_content = "x" * 600
messages = _microcompact_messages(total=1, tool_name="read_file", content=long_content)
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=2000,
context_block_limit=500,
)
def estimate(_provider, _model, msgs, _tools):
return sum(
1000 if msg.get("content") == long_content else 1
for msg in msgs
if msg.get("role") == "tool"
), "test"
monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate)
compacted_tool_call_ids: set[str] = set()
result = ContextGovernor().compact_inflight_overflow(
_governance_config(provider, tools, spec),
messages,
compacted_tool_call_ids,
)
tool_msg = next(m for m in result if m.get("role") == "tool")
assert "compacted to fit context" in tool_msg["content"]
assert "Do not repeat the same call unchanged" in tool_msg["content"]
assert "Retry with a narrower path, query, range, or result limit" in tool_msg["content"]
assert "tell the user the task cannot fit" in tool_msg["content"]
assert compacted_tool_call_ids == {"c0"}
def test_context_governor_keeps_compaction_boundary_stable(monkeypatch):
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
total = 18
long_content = "x" * 600
messages = _microcompact_messages(total=total, tool_name="read_file", content=long_content)
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=2224,
)
def estimate(_provider, _model, msgs, _tools):
return sum(
100 if msg.get("content") == long_content else 1
for msg in msgs
if msg.get("role") == "tool"
), "test"
monkeypatch.setattr("nanobot.agent.context_governance.estimate_prompt_tokens_chain", estimate)
governor = ContextGovernor()
compacted_tool_call_ids: set[str] = set()
config = _governance_config(provider, tools, spec, inflight_start_index=0)
first = governor.compact_inflight_overflow(config, messages, compacted_tool_call_ids)
first_ids = set(compacted_tool_call_ids)
second = governor.compact_inflight_overflow(config, messages, compacted_tool_call_ids)
assert compacted_tool_call_ids == first_ids
assert [m.get("content") for m in second] == [m.get("content") for m in first]
def test_microcompact_preserves_short_results(monkeypatch):
"""Short tool results below the compaction threshold should not be replaced."""
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
total = 15
messages = _microcompact_messages(total=total, tool_name="exec", content="short")
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=2024,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (2000, "test"),
)
result = ContextGovernor().compact_inflight_overflow(
_governance_config(provider, tools, spec),
messages,
set(),
)
assert result is messages # no copy needed — all stale results are short
def test_microcompact_skips_non_compactable_tools(monkeypatch):
"""Non-compactable tools (e.g. 'message') should never be replaced."""
provider = MagicMock()
provider.generation = SimpleNamespace(max_tokens=0)
tools = MagicMock()
tools.get_definitions.return_value = []
total = 15
long_content = "y" * 1000
messages = _microcompact_messages(total=total, tool_name="message", content=long_content)
spec = make_run_spec(provider,
initial_messages=messages,
tools=tools,
model="test-model",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
max_tokens=0,
context_window_tokens=2024,
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (2000, "test"),
)
result = ContextGovernor().compact_inflight_overflow(
_governance_config(provider, tools, spec),
messages,
set(),
)
assert result is messages # no compactable tools found
def test_governance_repairs_orphans_after_snip(): def test_governance_repairs_orphans_after_snip():
"""After snipping clips an assistant+tool_calls, orphan repair cleans up the tail.""" """After snipping clips an assistant+tool_calls, orphan repair cleans up the tail."""
# Simulate snipping that keeps only the tail: drop the assistant with # Simulate snipping that keeps only the tail: drop the assistant with
@@ -818,7 +1073,11 @@ def test_snip_history_preserves_user_message_after_truncation(monkeypatch):
lambda msg: token_sizes.get(str(msg.get("content")), 100), lambda msg: token_sizes.get(str(msg.get("content")), 100),
) )
trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) trimmed = ContextGovernor().snip_history(
_governance_config(provider, tools, spec),
messages,
tool_definitions=tools.get_definitions(),
)
# The first non-system message MUST be user (not assistant). # The first non-system message MUST be user (not assistant).
non_system = [m for m in trimmed if m.get("role") != "system"] non_system = [m for m in trimmed if m.get("role") != "system"]
@@ -863,7 +1122,11 @@ def test_snip_history_no_user_at_all_falls_back_gracefully(monkeypatch):
lambda msg: 100, lambda msg: 100,
) )
trimmed = ContextGovernor().snip_history(_governance_config(provider, tools, spec), messages) trimmed = ContextGovernor().snip_history(
_governance_config(provider, tools, spec),
messages,
tool_definitions=tools.get_definitions(),
)
# Should not crash. The result should still be a valid list. # Should not crash. The result should still be a valid list.
assert isinstance(trimmed, list) assert isinstance(trimmed, list)
@@ -871,7 +1134,6 @@ def test_snip_history_no_user_at_all_falls_back_gracefully(monkeypatch):
assert any(m.get("role") == "system" for m in trimmed) assert any(m.get("role") == "system" for m in trimmed)
# The _enforce_role_alternation safety net must be able to fix whatever # The _enforce_role_alternation safety net must be able to fix whatever
# _snip_history returns here — verify it produces a valid sequence. # _snip_history returns here — verify it produces a valid sequence.
from nanobot.providers.base import LLMProvider
fixed = LLMProvider._enforce_role_alternation(trimmed) fixed = LLMProvider._enforce_role_alternation(trimmed)
non_system = [m for m in fixed if m["role"] != "system"] non_system = [m for m in fixed if m["role"] != "system"]
if non_system: if non_system:
+8 -4
View File
@@ -10,6 +10,7 @@ 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
@@ -617,7 +618,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(
[{"role": "user", "content": "hello"}], TranscriptInput(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,
@@ -711,7 +712,10 @@ 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(
[{"role": "user", "content": "initial message from user A"}], TranscriptInput(
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(
@@ -812,7 +816,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(
[{"role": "user", "content": "hello"}], TranscriptInput(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,
@@ -1476,7 +1480,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(
[{"role": "user", "content": "hello"}], TranscriptInput(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,
+93
View File
@@ -9,6 +9,7 @@ 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
@@ -82,6 +83,18 @@ 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
@@ -554,6 +567,86 @@ 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
+2 -2
View File
@@ -114,7 +114,7 @@ async def test_removed_session_model_preset_falls_back_and_clears_metadata(tmp_p
provider=base, provider=base,
workspace=tmp_path, workspace=tmp_path,
model="base-model", model="base-model",
context_window_tokens=8_000, context_window_tokens=16_000,
) )
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign] loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
session_key = "sdk:removed-preset" session_key = "sdk:removed-preset"
@@ -196,7 +196,7 @@ async def test_sdk_custom_model_preset_metadata_does_not_select_runtime(
provider=base, provider=base,
workspace=tmp_path, workspace=tmp_path,
model="base-model", model="base-model",
context_window_tokens=8_000, context_window_tokens=16_000,
) )
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign] loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
bot = Nanobot(loop) bot = Nanobot(loop)
+8 -4
View File
@@ -7,6 +7,7 @@ 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
@@ -568,7 +569,10 @@ 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([], runtime=loop.llm_runtime()) await loop._run_agent_loop(
TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
)
loop.runner.run.assert_awaited_once() loop.runner.run.assert_awaited_once()
@@ -609,7 +613,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(
[{"role": "user", "content": "test"}], TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
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),
@@ -668,7 +672,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(
[{"role": "user", "content": "test"}], TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
runtime=runtime, runtime=runtime,
session=session, session=session,
request_context=RequestContext( request_context=RequestContext(
@@ -742,7 +746,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(
[{"role": "user", "content": "test"}], TranscriptInput(history=[{"role": "user", "content": "test"}], current_message=None),
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
session=session, session=session,
pending_queue=pending_queue, pending_queue=pending_queue,
+84 -4
View File
@@ -2654,12 +2654,14 @@ def test_webui_foreground_attaches_to_existing_managed_gateway(monkeypatch, tmp_
assert seen["lease_release_wait_for_stop"] is False assert seen["lease_release_wait_for_stop"] is False
def test_attach_to_background_gateway_detaches_on_ctrl_c(capsys) -> None: def test_attach_to_background_gateway_detaches_on_ctrl_c(capsys, tmp_path: Path) -> None:
stopped = False stopped = False
log_path = tmp_path / "gateway.log"
log_path.touch()
class _FakeRuntime: class _FakeRuntime:
def status(self): def status(self):
return SimpleNamespace(running=True) return SimpleNamespace(running=True, log_path=log_path)
def stop(self): def stop(self):
nonlocal stopped nonlocal stopped
@@ -2679,10 +2681,88 @@ def test_attach_to_background_gateway_detaches_on_ctrl_c(capsys) -> None:
assert "WebUI launcher detached" in rendered assert "WebUI launcher detached" in rendered
def test_attach_to_background_gateway_checks_owned_sidecar() -> None: def test_attach_to_background_gateway_follows_only_new_logs(capsys, tmp_path: Path) -> None:
log_path = tmp_path / "gateway.log"
log_path.write_text("historical log\n", encoding="utf-8")
polls = 0
class _FakeRuntime: class _FakeRuntime:
def status(self): def status(self):
return SimpleNamespace(running=True) return SimpleNamespace(running=True, log_path=log_path)
def _append_then_interrupt(_seconds: float) -> None:
nonlocal polls
if polls == 0:
with log_path.open("a", encoding="utf-8") as handle:
handle.write("[websocket] live log\n")
polls += 1
return
raise KeyboardInterrupt
cli_webui_support._attach_to_background_gateway(
_FakeRuntime(),
sleep=_append_then_interrupt,
)
output = capsys.readouterr().out
assert "[websocket] live log" in output
assert "historical log" not in output
def test_read_new_gateway_logs_recovers_after_truncation(tmp_path: Path) -> None:
log_path = tmp_path / "gateway.log"
log_path.write_text("a much longer historical log line\n", encoding="utf-8")
cursor = cli_webui_support._start_gateway_log_cursor(log_path)
log_path.write_text("fresh log\n", encoding="utf-8")
lines = cli_webui_support._read_new_gateway_logs(log_path, cursor)
assert lines == ["fresh log"]
assert cursor.offset == log_path.stat().st_size
def test_read_new_gateway_logs_detects_fast_rewrite_past_offset(tmp_path: Path) -> None:
log_path = tmp_path / "gateway.log"
log_path.write_text("historical log\n", encoding="utf-8")
cursor = cli_webui_support._start_gateway_log_cursor(log_path)
log_path.write_text("first fresh log\nsecond fresh log\n", encoding="utf-8")
lines = cli_webui_support._read_new_gateway_logs(log_path, cursor)
assert lines == ["first fresh log", "second fresh log"]
def test_read_new_gateway_logs_waits_for_complete_utf8_line(tmp_path: Path) -> None:
log_path = tmp_path / "gateway.log"
log_path.touch()
cursor = cli_webui_support._start_gateway_log_cursor(log_path)
encoded = "模型 ready\n".encode()
log_path.write_bytes(encoded[:2])
assert cli_webui_support._read_new_gateway_logs(log_path, cursor) == []
with log_path.open("ab") as handle:
handle.write(encoded[2:])
assert cli_webui_support._read_new_gateway_logs(log_path, cursor) == ["模型 ready"]
def test_read_new_gateway_logs_tolerates_missing_file(tmp_path: Path) -> None:
log_path = tmp_path / "missing.log"
cursor = cli_webui_support._start_gateway_log_cursor(log_path)
lines = cli_webui_support._read_new_gateway_logs(log_path, cursor)
assert lines == []
assert cursor.offset == 0
def test_attach_to_background_gateway_checks_owned_sidecar(tmp_path: Path) -> None:
log_path = tmp_path / "gateway.log"
log_path.touch()
class _FakeRuntime:
def status(self):
return SimpleNamespace(running=True, log_path=log_path)
def sidecar_exited() -> None: def sidecar_exited() -> None:
raise WebUIDevError("WebUI development server exited unexpectedly (code 23)") raise WebUIDevError("WebUI development server exited unexpectedly (code 23)")
+9 -2
View File
@@ -11,6 +11,7 @@ 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
@@ -311,10 +312,16 @@ class TestRestartCommand:
LLMResponse(content="second", usage=None), LLMResponse(content="second", usage=None),
]) ])
first = await loop._run_agent_loop([], runtime=loop.llm_runtime()) first = await loop._run_agent_loop(
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([], runtime=loop.llm_runtime()) second = await loop._run_agent_loop(
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
+40 -1
View File
@@ -7,6 +7,7 @@ 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:
@@ -292,7 +293,12 @@ 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": {"message_thread_id": 42}, "channelMeta": {
"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": {},
@@ -411,6 +417,39 @@ 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"
@@ -146,6 +146,51 @@ def test_controller_uses_governed_messages_for_provider_state_delta() -> None:
assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result" assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result"
def test_controller_estimates_active_state_plus_pending_delta(monkeypatch) -> None:
provider = _provider()
current_message = {"role": "user", "content": "new delta"}
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={
"items": [{"type": "reasoning", "encrypted_content": "opaque"}],
"context_tokens": 450,
},
pending_messages=[current_message],
)
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=[current_message],
state=state,
)
seen = {}
def estimate(_provider, _model, messages, tools):
seen["messages"] = messages
seen["tools"] = tools
return 100, "test-counter"
monkeypatch.setattr(
"nanobot.providers.conversation_state.estimate_prompt_tokens_chain",
estimate,
)
tokens = controller.estimate_request_context_tokens(
[current_message],
model_messages=[current_message],
tool_definitions=[{"type": "web_search"}],
)
assert tokens == 550
assert seen == {
"messages": [current_message],
"tools": [{"type": "web_search"}],
}
def test_transient_response_preserves_only_durable_request_messages() -> None: def test_transient_response_preserves_only_durable_request_messages() -> None:
provider = _provider() provider = _provider()
current_message = {"role": "user", "content": "continue"} current_message = {"role": "user", "content": "continue"}
@@ -112,14 +112,49 @@ 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_non_string_content_uses_latest(self): def test_consecutive_user_messages_preserve_text_before_multimodal_content(self):
image = {
"type": "image_url",
"image_url": {"url": "data:image/png;base64,aW1hZ2U="},
}
msgs = [ msgs = [
{"role": "user", "content": [{"type": "text", "text": "A"}]}, {"role": "user", "content": "Earlier unanswered question"},
{"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 len(result) == 1 assert result == [{
assert result[0]["content"] == "B" "role": "user",
"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 = [
+3 -6
View File
@@ -141,14 +141,13 @@ def test_internal_continuation_requires_budget_boundary_and_queue():
) )
def test_save_skip_matches_prefix_when_current_message_merged(): def test_save_skip_matches_prefix_when_current_message_was_persisted():
skip = _save_skip_for_turn( skip = _save_skip_for_turn(
message_metadata=None, message_metadata=None,
initial_message_count=2, # [system, merged user] initial_message_count=3, # [system, history user, current user]
history_count=1,
input_persisted_early=True, input_persisted_early=True,
) )
assert skip == 2 assert skip == 3
def test_save_skip_unchanged_for_standalone_current_message(): def test_save_skip_unchanged_for_standalone_current_message():
@@ -156,12 +155,10 @@ 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
+37
View File
@@ -1,6 +1,7 @@
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
@@ -11,6 +12,7 @@ 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
@@ -299,6 +301,41 @@ 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."""
+4 -1
View File
@@ -6,6 +6,7 @@ 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
@@ -178,7 +179,9 @@ class TestMessageToolSuppressLogic:
progress.append((content, tool_hint)) progress.append((content, tool_hint))
result = await loop._run_agent_loop( result = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress TranscriptInput(history=[], current_message=None),
runtime=loop.llm_runtime(),
on_progress=on_progress,
) )
assert result.final_content == "Done" assert result.final_content == "Done"
+60
View File
@@ -183,6 +183,66 @@ 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,
+6 -2
View File
@@ -9,7 +9,9 @@ from nanobot.agent.tools.shell import ExecTool
def test_coding_tool_descriptions_steer_editing_priority() -> None: def test_coding_tool_descriptions_steer_editing_priority() -> None:
apply_patch = ApplyPatchTool().description.lower() apply_patch = ApplyPatchTool().description.lower()
edit_file = EditFileTool().description.lower() edit_tool = EditFileTool()
edit_file = edit_tool.description.lower()
edit_parameters = edit_tool.parameters["properties"]
write_file = WriteFileTool().description.lower() write_file = WriteFileTool().description.lower()
assert "default tool for code edits" in apply_patch assert "default tool for code edits" in apply_patch
@@ -18,8 +20,10 @@ def test_coding_tool_descriptions_steer_editing_priority() -> None:
assert "edit_file only for small exact replacements" in apply_patch assert "edit_file only for small exact replacements" in apply_patch
assert "small, exact replacement" in edit_file assert "small, exact replacement" in edit_file
assert "copied from read_file" in edit_file
assert "prefer apply_patch" in edit_file assert "prefer apply_patch" in edit_file
assert "occurrence, line_hint, and replace_all=true are mutually exclusive" in edit_file
assert "copy it from read_file" in edit_parameters["old_text"]["description"].lower()
assert "must differ from old_text" in edit_parameters["new_text"]["description"].lower()
assert "replace an entire file" in write_file assert "replace an entire file" in write_file
assert "prefer apply_patch" in write_file assert "prefer apply_patch" in write_file
+19 -28
View File
@@ -205,7 +205,7 @@ describe("NanobotTui layout", () => {
expect(occurrences(frame, "Ask nanobot anything")).toBe(1) expect(occurrences(frame, "Ask nanobot anything")).toBe(1)
expect(occurrences(frame, "Ready")).toBe(0) expect(occurrences(frame, "Ready")).toBe(0)
expect(occurrences(frame, "Getting ready…")).toBe(1) expect(occurrences(frame, "Getting ready…")).toBe(1)
expect(occurrences(frame, "nanobot · test/model")).toBe(1) expect(occurrences(frame, "default ▾")).toBe(1)
} }
app.accept({ event: "attached", chat_id: "chat" }) app.accept({ event: "attached", chat_id: "chat" })
@@ -985,7 +985,6 @@ describe("NanobotTui layout", () => {
const ui = app as unknown as { const ui = app as unknown as {
composer: TextareaRenderable composer: TextareaRenderable
sessionMenu: { visible: boolean } sessionMenu: { visible: boolean }
titleText: { plainText: string }
runtimeControls: { modelText: { plainText: string } } runtimeControls: { modelText: { plainText: string } }
} }
@@ -999,8 +998,7 @@ describe("NanobotTui layout", () => {
ui.composer.submit() ui.composer.submit()
await waitUntil(() => attached.length === 1) await waitUntil(() => attached.length === 1)
expect(attached).toEqual(["other"]) expect(attached).toEqual(["other"])
expect(ui.titleText.plainText).toContain("Release checklist") expect(ui.runtimeControls.modelText.plainText).toBe("Deep Research ▾")
expect(ui.runtimeControls.modelText.plainText).toContain("Deep Research")
expect(ui.runtimeControls.modelText.plainText).not.toContain("test/model") expect(ui.runtimeControls.modelText.plainText).not.toContain("test/model")
app.accept({ event: "attached", chat_id: "other" }) app.accept({ event: "attached", chat_id: "other" })
@@ -1009,8 +1007,7 @@ describe("NanobotTui layout", () => {
ui.composer.submit() ui.composer.submit()
await waitUntil(() => newChats.length === 1) await waitUntil(() => newChats.length === 1)
expect(newChats).toEqual(["new"]) expect(newChats).toEqual(["new"])
expect(ui.titleText.plainText).toContain("New chat") expect(ui.runtimeControls.modelText.plainText).toBe("default ▾")
expect(ui.runtimeControls.modelText.plainText).toContain("test/model")
} finally { } finally {
globalThis.fetch = original globalThis.fetch = original
} }
@@ -1182,7 +1179,8 @@ describe("NanobotTui layout", () => {
model_preset: "Codex", model_preset: "Codex",
}) })
await setup.flush() await setup.flush()
expect(ui.runtimeControls.modelText.plainText).toContain("Codex · openai/gpt-5.6") expect(ui.runtimeControls.modelText.plainText).toBe("Codex ")
expect(ui.runtimeControls.modelText.plainText).not.toContain("openai/gpt-5.6")
app.accept({ app.accept({
event: "runtime_model_updated", event: "runtime_model_updated",
@@ -1190,7 +1188,7 @@ describe("NanobotTui layout", () => {
model_preset: "DeepSeek", model_preset: "DeepSeek",
}) })
await setup.flush() await setup.flush()
expect(ui.runtimeControls.modelText.plainText).toContain("Codex · openai/gpt-5.6") expect(ui.runtimeControls.modelText.plainText).toBe("Codex ")
expect(ui.runtimeControls.modelText.plainText).not.toContain("DeepSeek") expect(ui.runtimeControls.modelText.plainText).not.toContain("DeepSeek")
}) })
@@ -1212,8 +1210,8 @@ describe("NanobotTui layout", () => {
}) })
await setup.flush() await setup.flush()
expect(ui.runtimeControls.modelText.plainText).toContain("deepseek/deepseek-chat") expect(ui.runtimeControls.modelText.plainText).toBe("default ▾")
expect(ui.runtimeControls.modelText.plainText).not.toContain("Codex") expect(ui.runtimeControls.modelText.plainText).not.toContain("deepseek/deepseek-chat")
}) })
test("refreshes the canonical preset after the model command completes", async () => { test("refreshes the canonical preset after the model command completes", async () => {
@@ -1309,7 +1307,6 @@ describe("NanobotTui layout", () => {
menuRoot: { getChildren(): unknown[] } menuRoot: { getChildren(): unknown[] }
} }
composer: TextareaRenderable composer: TextareaRenderable
titleText: TextRenderable
status: TextRenderable status: TextRenderable
meta: TextRenderable meta: TextRenderable
} }
@@ -1324,7 +1321,6 @@ describe("NanobotTui layout", () => {
expect(ui.runtimeControls.modelText.selectable).toBe(false) expect(ui.runtimeControls.modelText.selectable).toBe(false)
expect(ui.runtimeControls.accessText.selectable).toBe(false) expect(ui.runtimeControls.accessText.selectable).toBe(false)
expect(ui.runtimeControls.contextText.selectable).toBe(false) expect(ui.runtimeControls.contextText.selectable).toBe(false)
expect(ui.titleText.selectable).toBe(false)
expect(ui.status.selectable).toBe(false) expect(ui.status.selectable).toBe(false)
expect(ui.meta.selectable).toBe(false) expect(ui.meta.selectable).toBe(false)
app.accept({ event: "goal_status", chat_id: "chat", status: "running" }) app.accept({ event: "goal_status", chat_id: "chat", status: "running" })
@@ -1394,7 +1390,7 @@ describe("NanobotTui layout", () => {
} }
}) })
test("opens and switches sessions from the clickable title", async () => { test("switches sessions only through the sessions command", async () => {
const original = globalThis.fetch const original = globalThis.fetch
globalThis.fetch = ((input: string | URL | Request) => { globalThis.fetch = ((input: string | URL | Request) => {
const url = String(input) const url = String(input)
@@ -1420,14 +1416,18 @@ describe("NanobotTui layout", () => {
const ui = app as unknown as { const ui = app as unknown as {
composer: TextareaRenderable composer: TextareaRenderable
sessionMenu: { visible: boolean; root: { getChildren(): unknown[] } } sessionMenu: { visible: boolean; root: { getChildren(): unknown[] } }
titleText: TextRenderable title: { getChildren(): unknown[] }
status: TextRenderable
} }
try { try {
await waitUntil(() => (app as unknown as { ready: boolean }).ready) await waitUntil(() => (app as unknown as { ready: boolean }).ready)
await setup.renderOnce() await setup.renderOnce()
await setup.mockMouse.click(ui.titleText.x + 2, ui.titleText.y) const titleItems = ui.title.getChildren() as TextRenderable[]
expect(titleItems.some((item) => item.id === "nanobot-tui-title-text")).toBe(false)
expect(ui.sessionMenu.visible).toBe(false)
ui.composer.setText("/sessions")
ui.composer.submit()
await waitUntil(() => ui.sessionMenu.visible) await waitUntil(() => ui.sessionMenu.visible)
await setup.flush() await setup.flush()
expect(ui.composer.placeholder).toBe("Search sessions") expect(ui.composer.placeholder).toBe("Search sessions")
@@ -1441,15 +1441,6 @@ describe("NanobotTui layout", () => {
expect(attached).toEqual(["other"]) expect(attached).toEqual(["other"])
expect(ui.sessionMenu.visible).toBe(false) expect(ui.sessionMenu.visible).toBe(false)
expect(ui.composer.focused).toBe(true) expect(ui.composer.focused).toBe(true)
expect(ui.titleText.plainText).toContain("Release checklist")
app.accept({ event: "attached", chat_id: "other" })
await setup.mockMouse.click(ui.titleText.x + 2, ui.titleText.y)
await waitUntil(() => ui.sessionMenu.visible)
ui.composer.blur()
await setup.mockMouse.click(ui.status.x, ui.status.y)
expect(ui.sessionMenu.visible).toBe(false)
expect(ui.composer.focused).toBe(true)
} finally { } finally {
globalThis.fetch = original globalThis.fetch = original
} }
@@ -1715,7 +1706,7 @@ describe("NanobotTui layout", () => {
await setup.flush() await setup.flush()
const frame = setup.captureCharFrame() const frame = setup.captureCharFrame()
expect(frame).toContain("Release checklist") expect(frame).toContain("Release checklist")
expect(occurrences(frame, "Current chat")).toBe(1) expect(occurrences(frame, "Current chat")).toBe(0)
} finally { } finally {
globalThis.fetch = original globalThis.fetch = original
} }
@@ -1959,7 +1950,7 @@ describe("NanobotTui layout", () => {
} else if (width >= 28 && height >= 9) { } else if (width >= 28 && height >= 9) {
expect(occurrences(frame, "Enter now · Tab next")).toBe(1) expect(occurrences(frame, "Enter now · Tab next")).toBe(1)
} }
expect(occurrences(frame, "nanobot · test/model")).toBe(height >= 14 ? 1 : 0) expect(occurrences(frame, "default ▾")).toBe(height >= 14 ? 1 : 0)
} }
}) })
@@ -3122,7 +3113,7 @@ describe("NanobotTui with a Herdr pane title reporter", () => {
await setup.flush() await setup.flush()
const activeFrame = setup.captureCharFrame() const activeFrame = setup.captureCharFrame()
expect(activeFrame).toContain(">_ nanobot") expect(activeFrame).toContain(">_ nanobot")
expect(activeFrame).toContain("test/model") expect(activeFrame).toContain("default ▾")
expect(occurrences(activeFrame, " Ship the Herdr integration")).toBe(1) expect(occurrences(activeFrame, " Ship the Herdr integration")).toBe(1)
expect(occurrences(activeFrame, "app.ts")).toBe(1) expect(occurrences(activeFrame, "app.ts")).toBe(1)
expect(ui.composer.placeholder).toBe("Enter send now · Tab send next") expect(ui.composer.placeholder).toBe("Enter send now · Tab send next")
+3 -40
View File
@@ -94,7 +94,7 @@ import {
type FooterMode, type FooterMode,
type FooterHintTheme, type FooterHintTheme,
} from "./footer-hints" } from "./footer-hints"
import { createTuiHost, type TuiHost } from "./host" import { configureOpenTuiEnvironment, createTuiHost, type TuiHost } from "./host"
interface AppOptions { interface AppOptions {
wsUrl?: string wsUrl?: string
@@ -442,7 +442,6 @@ export class NanobotTui {
private readonly client: ChatClient private readonly client: ChatClient
private readonly shell: BoxRenderable private readonly shell: BoxRenderable
private readonly title: BoxRenderable private readonly title: BoxRenderable
private readonly titleText: TextRenderable
private readonly composerFrame: BoxRenderable private readonly composerFrame: BoxRenderable
private readonly composer: TextareaRenderable private readonly composer: TextareaRenderable
private composerSyntax: SyntaxStyle private composerSyntax: SyntaxStyle
@@ -649,28 +648,6 @@ export class NanobotTui {
alignItems: "center", alignItems: "center",
backgroundColor: RGBA.defaultBackground(), backgroundColor: RGBA.defaultBackground(),
}) })
this.titleText = new TextRenderable(renderer, {
id: "nanobot-tui-title-text",
content: "nanobot",
height: 1,
flexShrink: 0,
truncate: true,
fg: this.palette.muted,
selectable: false,
onMouseOver: () => { this.titleText.fg = this.palette.accent },
onMouseOut: () => this.renderTitleColor(),
onMouseDown: (event) => {
if (event.button !== 0) return
event.preventDefault()
event.stopPropagation()
this.renderer.clearSelection()
if (this.sessionLoading || this.sessionMenu.visible) {
this.closeSessions()
return
}
void this.openSessions()
},
})
this.runtimeControls = new RuntimeControls( this.runtimeControls = new RuntimeControls(
renderer, renderer,
runtimeControlsTheme(this.palette), runtimeControlsTheme(this.palette),
@@ -703,7 +680,6 @@ export class NanobotTui {
}, },
}, },
) )
this.title.add(this.titleText)
this.title.add(this.runtimeControls.modelText) this.title.add(this.runtimeControls.modelText)
this.title.add(this.runtimeControls.accessText) this.title.add(this.runtimeControls.accessText)
this.title.add(this.runtimeControls.contextText) this.title.add(this.runtimeControls.contextText)
@@ -820,6 +796,7 @@ 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,
@@ -1876,7 +1853,6 @@ export class NanobotTui {
this.composer.syntaxStyle = this.composerSyntax this.composer.syntaxStyle = this.composerSyntax
this.syncComposerImageHighlights(this.composer.plainText) this.syncComposerImageHighlights(this.composer.plainText)
void this.renderer.idle().catch(() => {}).finally(() => previousComposerSyntax.destroy()) void this.renderer.idle().catch(() => {}).finally(() => previousComposerSyntax.destroy())
this.renderTitleColor()
this.status.fg = this.palette.muted this.status.fg = this.palette.muted
this.meta.fg = this.palette.faint this.meta.fg = this.palette.faint
this.updateMeta() this.updateMeta()
@@ -1956,24 +1932,15 @@ export class NanobotTui {
} }
private updateTitle(): void { private updateTitle(): void {
const identity = this.sessionTitle.trim() || "nanobot"
this.titleText.maxWidth = Math.max(8, Math.floor(this.renderer.width * 0.38))
this.titleText.content = identity
const context = this.contextTokens === null const context = this.contextTokens === null
? "" ? ""
: ` · ~${formatTokenCount(this.contextTokens)}${this.contextWindowTokens : ` ~${formatTokenCount(this.contextTokens)}${this.contextWindowTokens
? `/${formatTokenCount(this.contextWindowTokens)}` ? `/${formatTokenCount(this.contextWindowTokens)}`
: ""} ctx` : ""} ctx`
this.runtimeControls.updateModel(this.modelName, this.modelPreset) this.runtimeControls.updateModel(this.modelName, this.modelPreset)
this.runtimeControls.updateContext(context) this.runtimeControls.updateContext(context)
} }
private renderTitleColor(): void {
this.titleText.fg = this.sessionLoading || this.sessionMenu.visible
? this.palette.accent
: this.palette.muted
}
private resizeComposer(): void { private resizeComposer(): void {
const verticalPadding = this.renderer.height >= 12 ? 1 : 0 const verticalPadding = this.renderer.height >= 12 ? 1 : 0
const maxContentHeight = Math.max(1, Math.min(12, Math.floor(this.renderer.height / 3))) const maxContentHeight = Math.max(1, Math.min(12, Math.floor(this.renderer.height / 3)))
@@ -2384,7 +2351,6 @@ export class NanobotTui {
this.contextPanel.hide() this.contextPanel.hide()
this.clearComposer() this.clearComposer()
this.sessionLoading = true this.sessionLoading = true
this.renderTitleColor()
const loadId = ++this.sessionLoadId const loadId = ++this.sessionLoadId
this.status.content = "Loading sessions…" this.status.content = "Loading sessions…"
try { try {
@@ -2411,7 +2377,6 @@ export class NanobotTui {
this.defaultModelPreset, this.defaultModelPreset,
) )
this.startSessionRefresh() this.startSessionRefresh()
this.renderTitleColor()
this.sessionMenu.update(this.composer.plainText, limit) this.sessionMenu.update(this.composer.plainText, limit)
this.syncComposerPlaceholder() this.syncComposerPlaceholder()
this.updateMeta() this.updateMeta()
@@ -2419,7 +2384,6 @@ export class NanobotTui {
} catch (error) { } catch (error) {
if (loadId !== this.sessionLoadId) return if (loadId !== this.sessionLoadId) return
this.sessionLoading = false this.sessionLoading = false
this.renderTitleColor()
this.status.content = error instanceof Error ? error.message : String(error) this.status.content = error instanceof Error ? error.message : String(error)
} }
} }
@@ -2578,7 +2542,6 @@ export class NanobotTui {
this.sessionLoadId += 1 this.sessionLoadId += 1
this.sessionLoading = false this.sessionLoading = false
this.hideSessionMenu() this.hideSessionMenu()
this.renderTitleColor()
this.clearComposer() this.clearComposer()
this.syncComposerPlaceholder() this.syncComposerPlaceholder()
this.composer.focus() this.composer.focus()
+25 -1
View File
@@ -1,6 +1,9 @@
import { describe, expect, test } from "bun:test" import { describe, expect, test } from "bun:test"
import { createTuiHost } from "./host" import {
configureOpenTuiEnvironment,
createTuiHost,
} from "./host"
async function settle(): Promise<void> { async function settle(): Promise<void> {
await Bun.sleep(0) await Bun.sleep(0)
@@ -8,6 +11,27 @@ 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]) })
+13
View File
@@ -8,6 +8,19 @@ 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 {}
+2 -5
View File
@@ -309,12 +309,9 @@ export class RuntimeControls {
} }
private render(): void { private render(): void {
const runtime = this.modelPreset !== "default" this.modelText.content = `${this.modelPreset}`
? [this.modelPreset, this.model].filter(Boolean).join(" · ")
: this.model
this.modelText.content = ` · ${runtime}`
const access = this.scope.access_mode === "full" ? "full access" : "workspace access" const access = this.scope.access_mode === "full" ? "full access" : "workspace access"
this.accessText.content = ` · ${access}` this.accessText.content = ` ${access}`
this.renderColors() this.renderColors()
} }
+1 -1
View File
@@ -211,7 +211,7 @@ export class Transcript {
const title = this.createText(`>_ nanobot v${options.version}`, "text", true) const title = this.createText(`>_ nanobot v${options.version}`, "text", true)
const context = this.createText([ const context = this.createText([
"", "",
`${options.model} · ${options.access}`, `${options.model} ${options.access}`,
options.workspace, options.workspace,
].join("\n"), "muted") ].join("\n"), "muted")
row.add(title) row.add(title)