refactor(agent): let runner own context compaction (#5568)

* refactor(agent): consolidate accepted history under pressure

* fix(agent): align provider and session compaction

* refactor(agent): simplify runner context compaction

* refactor(agent): remove background token consolidation

* fix(agent): keep injected transcript messages distinct

* refactor(agent): unify native compaction summaries

* fix(agent): preserve native compaction boundary

* fix(agent): unify context compaction paths

* fix(agent): preserve exact compaction request boundaries
This commit is contained in:
chengyongru
2026-09-02 18:05:54 +08:00
committed by GitHub
parent da96c5c6eb
commit d81aa5a4ab
44 changed files with 1914 additions and 1666 deletions
+3 -1
View File
@@ -235,7 +235,7 @@ class ContextBuilder:
def build_messages( def build_messages(
self, self,
history: list[dict[str, Any]], history: list[dict[str, Any]],
current_message: str, current_message: str | None,
*, *,
media: list[str] | None = None, media: list[str] | None = None,
channel: str | None = None, channel: str | None = None,
@@ -259,6 +259,8 @@ class ContextBuilder:
workspace=workspace, workspace=workspace,
include_memory=include_memory, include_memory=include_memory,
) )
if current_message is None:
return messages
current = messages[-1] current = messages[-1]
if len(messages) < 2 or messages[-2].get("role") != current.get("role"): if len(messages) < 2 or messages[-2].get("role") != current.get("role"):
return messages return messages
+493 -30
View File
@@ -1,19 +1,43 @@
"""Model-message governance for agent runner requests. """Model-message governance and compaction for agent runner requests.
This module owns model-facing message shaping and tool-result content normalization. This module owns model-facing message shaping, request pressure, H/delta
It may return copied messages or persisted-result placeholders, but it must not compaction state, and tool-result content normalization. It may return copied
mutate an existing session history list in place. messages or persisted-result placeholders, but it must not mutate an existing
session history list in place.
""" """
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from collections.abc import Awaitable, Callable
from copy import deepcopy
from dataclasses import dataclass, replace
from datetime import datetime
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
from loguru import logger from loguru import logger
from nanobot.providers.base import LLMUsage from nanobot.agent.context import TranscriptInput
from nanobot.providers.base import (
LLMResponse,
LLMUsage,
ProviderCallContext,
ProviderConversationState,
)
from nanobot.providers.conversation_state import (
ProviderConversationStateController,
allows_conversation_message_merge,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_MESSAGE_META,
detach_runtime_context,
reattach_runtime_context,
)
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.summary import (
SUMMARY_CONTINUATION_TEXT,
SessionSummaryCheckpoint,
)
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,6 +51,16 @@ if TYPE_CHECKING:
from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.registry import ToolRegistry
from nanobot.providers.base import LLMProvider from nanobot.providers.base import LLMProvider
TranscriptBuilder = Callable[[TranscriptInput], list[dict[str, Any]]]
HistoryConsolidator = Callable[
[list[dict[str, Any]], str | None],
Awaitable[str | None],
]
ProviderCompactionConsolidator = Callable[
[ProviderConversationState, list[dict[str, Any]], str | None],
Awaitable[str | None],
]
SNIP_SAFETY_BUFFER = 1024 SNIP_SAFETY_BUFFER = 1024
# 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"})
@@ -85,8 +119,204 @@ class ContextGovernanceConfig:
max_tokens: int | None = None max_tokens: int | None = None
@dataclass(slots=True)
class ContextCompactionState:
"""Track accepted provider input H separately from the unsent delta."""
raw_messages: list[dict[str, Any]]
accepted_messages: list[dict[str, Any]]
raw_accepted_boundary: int
active_summary: str | None
transcript_input: TranscriptInput
transcript_builder: TranscriptBuilder
consolidate_history: HistoryConsolidator
consolidate_provider_compaction: ProviderCompactionConsolidator | None
summary_checkpoint: SessionSummaryCheckpoint | None = None
@classmethod
def from_transcript(
cls,
transcript_input: TranscriptInput,
transcript_builder: TranscriptBuilder,
consolidate_history: HistoryConsolidator | None,
consolidate_provider_compaction: ProviderCompactionConsolidator | None,
) -> tuple[list[dict[str, Any]], ContextCompactionState | None]:
"""Build the raw transcript and its initial H/delta boundary."""
messages = list(transcript_builder(transcript_input))
if consolidate_history is None:
return messages, None
accepted_history_boundary = 1 + len(transcript_input.history)
return messages, cls(
raw_messages=messages,
accepted_messages=deepcopy(messages[:accepted_history_boundary]),
raw_accepted_boundary=accepted_history_boundary,
active_summary=(
transcript_input.session_summary["text"]
if transcript_input.session_summary is not None
else None
),
transcript_input=transcript_input,
transcript_builder=transcript_builder,
consolidate_history=consolidate_history,
consolidate_provider_compaction=consolidate_provider_compaction,
)
def request_messages(
self,
raw_messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
return [
*deepcopy(self.accepted_messages),
*deepcopy(raw_messages[self.raw_accepted_boundary:]),
]
def delta_after_accepted(
self,
request_messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
return deepcopy(request_messages[len(self.accepted_messages):])
def accept_request(
self,
model_messages: list[dict[str, Any]],
*,
raw_boundary: int,
) -> None:
"""Advance H after the provider has received one request."""
self.accepted_messages = deepcopy(model_messages)
self.raw_accepted_boundary = raw_boundary
@dataclass(slots=True)
class ModelRequestState:
"""Context state shared by every provider request in one runner turn."""
config: ContextGovernanceConfig
conversation: ProviderConversationStateController
usage: LLMUsage | None = None
messages: list[dict[str, Any]] | None = None
tool_definitions: list[dict[str, Any]] | None = None
compaction: ContextCompactionState | None = None
provider_compaction_applied: bool = False
class ContextGovernor: class ContextGovernor:
"""Prepare model-copy messages while preserving persisted history.""" """Own model-request context while preserving persisted history."""
@staticmethod
def _merge_message_content(left: Any, right: Any) -> str | list[dict[str, Any]]:
if isinstance(left, str) and isinstance(right, str):
return f"{left}\n\n{right}" if left else right
def _to_blocks(value: Any) -> list[dict[str, Any]]:
if isinstance(value, list):
return [
cast(dict[str, Any], item)
if isinstance(item, dict)
else {"type": "text", "text": str(item)}
for item in cast(list[Any], value)
]
if value is None:
return []
return [{"type": "text", "text": str(value)}]
return _to_blocks(left) + _to_blocks(right)
@classmethod
def _merge_adjacent_user_messages_for_model(
cls,
messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Merge adjacent visible user messages only in the model-facing copy."""
prepared: list[dict[str, Any]] = []
for source in messages:
injection = deepcopy(source)
if (
prepared
and injection.get("role") == "user"
and prepared[-1].get("role") == "user"
and injection.get("content") != SUMMARY_CONTINUATION_TEXT
and prepared[-1].get("content") != SUMMARY_CONTINUATION_TEXT
and not is_hidden_history_message(injection)
and not is_hidden_history_message(prepared[-1])
and allows_conversation_message_merge(injection)
and allows_conversation_message_merge(prepared[-1])
):
merged = dict(prepared[-1])
left_meta = merged.get("_meta")
right_meta = injection.get("_meta")
left_meta_dict = (
cast(dict[str, Any], left_meta) if isinstance(left_meta, dict) else None
)
right_meta_dict = (
cast(dict[str, Any], right_meta) if isinstance(right_meta, dict) else None
)
left_marker = (
left_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
if left_meta_dict is not None
else None
)
right_marker = (
right_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
if right_meta_dict is not None
else None
)
left_marker_dict = (
cast(dict[str, Any], left_marker) if isinstance(left_marker, dict) else None
)
right_marker_dict = (
cast(dict[str, Any], right_marker) if isinstance(right_marker, dict) else None
)
empty_sources: list[str] = []
empty_blocks: list[dict[str, Any]] = []
detached_left = (
detach_runtime_context(merged.get("content"), left_marker_dict)
if left_marker_dict is not None
else (merged.get("content"), empty_sources, empty_blocks)
)
detached_right = (
detach_runtime_context(injection.get("content"), right_marker_dict)
if right_marker_dict is not None
else (injection.get("content"), empty_sources, empty_blocks)
)
if detached_left is not None and detached_right is not None:
left_content, left_sources, left_blocks = detached_left
right_content, right_sources, right_blocks = detached_right
merged_content = cls._merge_message_content(left_content, right_content)
context_blocks = [*left_blocks, *right_blocks]
if context_blocks:
merged_content, marker = reattach_runtime_context(
merged_content,
[*left_sources, *right_sources],
context_blocks,
)
internal_meta = (
dict(left_meta_dict) if left_meta_dict is not None else {}
)
if right_meta_dict is not None:
for key, value in right_meta_dict.items():
internal_meta.setdefault(key, value)
internal_meta[RUNTIME_CONTEXT_MESSAGE_META] = marker
merged["_meta"] = internal_meta
merged["content"] = merged_content
else:
merged["content"] = cls._merge_message_content(
merged.get("content"),
injection.get("content"),
)
prepared[-1] = merged
continue
prepared.append(injection)
return prepared
def prepare_messages_for_model(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Build the normalized model-facing copy of a raw transcript."""
governed = self.prepare_for_model(config, messages)
return self._merge_adjacent_user_messages_for_model(governed)
def prepare_for_model( def prepare_for_model(
self, self,
@@ -115,17 +345,31 @@ class ContextGovernor:
) )
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)
return self.ensure_request_fits(
config,
updated,
tool_definitions=tool_definitions,
)
def ensure_request_fits(
self,
config: ContextGovernanceConfig,
messages: list[dict[str, Any]],
*,
tool_definitions: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Validate an exact model request without dropping any messages."""
if not config.context_window_tokens: if not config.context_window_tokens:
return updated return messages
budget = self.input_budget(config) budget = self.input_budget(config)
estimated, source = estimate_prompt_tokens_chain( estimated, source = estimate_prompt_tokens_chain(
config.provider, config.provider,
config.model, config.model,
updated, messages,
tool_definitions, tool_definitions,
) )
if budget > 0 and estimated <= budget: if budget > 0 and estimated <= budget:
return updated return messages
raise ContextWindowExceededError( raise ContextWindowExceededError(
session_key=config.session_key, session_key=config.session_key,
estimated_tokens=estimated, estimated_tokens=estimated,
@@ -133,6 +377,41 @@ class ContextGovernor:
source=source, source=source,
) )
def request_pressure(
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[int, str] | None:
"""Return the authoritative measurement when a request is pressured."""
if not config.context_window_tokens:
return None
budget = self.input_budget(config)
if request_context_tokens is not None:
measured = request_context_tokens
source = "resumed provider state plus pending messages"
elif (
usage_matches_messages
and usage is not None
and usage.context_tokens is not None
):
measured = usage.context_tokens
source = "matching provider usage"
else:
measured, source = estimate_prompt_tokens_chain(
config.provider,
config.model,
messages,
tool_definitions,
)
if budget > 0 and measured < budget:
return None
return measured, source
def fit_request( def fit_request(
self, self,
config: ContextGovernanceConfig, config: ContextGovernanceConfig,
@@ -144,27 +423,15 @@ class ContextGovernor:
request_context_tokens: int | None = None, request_context_tokens: int | None = None,
) -> tuple[list[dict[str, Any]], bool]: ) -> tuple[list[dict[str, Any]], bool]:
"""Fit the request when its measured or estimated input is pressured.""" """Fit the request when its measured or estimated input is pressured."""
if not config.context_window_tokens: pressure = self.request_pressure(
return messages, False config,
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, messages,
tool_definitions, usage,
usage_matches_messages=usage_matches_messages,
tool_definitions=tool_definitions,
request_context_tokens=request_context_tokens,
) )
if request_context_tokens is not None: if pressure is None:
estimated = max(estimated, request_context_tokens)
pressured = budget <= 0 or estimated >= budget
if not pressured:
return messages, False return messages, False
return self.fit_to_budget( return self.fit_to_budget(
config, config,
@@ -172,6 +439,201 @@ class ContextGovernor:
tool_definitions=tool_definitions, tool_definitions=tool_definitions,
), True ), True
@staticmethod
def _summary_transcript(
compaction: ContextCompactionState,
summary: str,
) -> list[dict[str, Any]]:
"""Rebuild only the stable system prefix around a replacement summary."""
return compaction.transcript_builder(
replace(
compaction.transcript_input,
history=[],
current_message=None,
media=None,
session_summary={
"text": summary,
"last_active": datetime.now().astimezone().isoformat(),
},
runtime_context_blocks=None,
)
)
async def summarize_provider_compaction(
self,
state: ModelRequestState,
response: LLMResponse,
*,
current_request_boundary: int | None,
) -> None:
"""Materialize the exact input replaced by provider-native compaction."""
compaction = state.compaction
if (
not response.provider_compaction_applied
or response.provider_compaction_state is None
or compaction is None
or compaction.consolidate_provider_compaction is None
):
return
if response.provider_compaction_scope == "prior_context":
accepted_messages = compaction.accepted_messages
transcript_boundary = compaction.raw_accepted_boundary
elif (
response.provider_compaction_scope == "current_request"
and state.messages is not None
and current_request_boundary is not None
):
accepted_messages = state.messages
transcript_boundary = current_request_boundary
else:
logger.warning(
"Ignoring provider compaction with missing request-boundary scope for {}",
state.config.session_key or "default",
)
return
summary = await compaction.consolidate_provider_compaction(
response.provider_compaction_state,
deepcopy(accepted_messages),
compaction.active_summary,
)
if not summary:
return
compaction.active_summary = summary
compaction.summary_checkpoint = SessionSummaryCheckpoint(
summary=summary,
transcript_boundary=transcript_boundary,
)
async def _compact_request_history(
self,
state: ModelRequestState,
compaction: ContextCompactionState,
messages: list[dict[str, Any]],
pressure: tuple[int, str],
*,
tool_definitions: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Replace accepted history H with a checkpoint while preserving delta."""
delta_messages = compaction.delta_after_accepted(messages)
consolidation_prefix = self.prepare_messages_for_model(
state.config,
compaction.accepted_messages,
)
summary = await compaction.consolidate_history(
deepcopy(consolidation_prefix),
compaction.active_summary,
)
if not summary:
measured, source = pressure
raise ContextWindowExceededError(
session_key=state.config.session_key,
estimated_tokens=measured,
input_budget=self.input_budget(state.config),
source=source,
)
compaction.active_summary = summary
prepared = self.prepare_messages_for_model(
state.config,
[
*self._summary_transcript(compaction, summary),
{"role": "user", "content": SUMMARY_CONTINUATION_TEXT},
*delta_messages,
],
)
# Responses-style state is append-only. Replacing H with a
# checkpoint requires a fresh request; a successful response may
# establish a new provider-owned state at the rewritten boundary.
state.conversation.replace_transcript(compaction.raw_messages)
state.usage = None
prepared = self.ensure_request_fits(
state.config,
prepared,
tool_definitions=tool_definitions,
)
compaction.summary_checkpoint = SessionSummaryCheckpoint(
summary=summary,
transcript_boundary=compaction.raw_accepted_boundary,
)
return prepared
async def prepare_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, compact or fit, and record the exact provider payload."""
prepared = self.prepare_messages_for_model(state.config, messages)
model_messages: list[dict[str, Any]] | None = prepared
supplemental_messages: list[dict[str, Any]] | None = None
request_context_tokens = None
if transcript is not None:
if tool_definitions is None:
model_messages = None
supplemental_messages = [prepared[-1]]
request_context_tokens = state.conversation.estimate_request_context_tokens(
transcript,
model_messages=model_messages,
supplemental_messages=supplemental_messages,
tool_definitions=tool_definitions,
)
usage_matches_messages = (
state.messages is not None
and prepared == state.messages
and tool_definitions == state.tool_definitions
)
request_was_fitted = False
compaction = state.compaction
if compaction is None:
prepared, request_was_fitted = self.fit_request(
state.config,
prepared,
state.usage,
usage_matches_messages=usage_matches_messages,
tool_definitions=tool_definitions,
request_context_tokens=request_context_tokens,
)
else:
pressure = self.request_pressure(
state.config,
prepared,
state.usage,
usage_matches_messages=usage_matches_messages,
tool_definitions=tool_definitions,
request_context_tokens=request_context_tokens,
)
if pressure is not None:
prepared = await self._compact_request_history(
state,
compaction,
messages,
pressure,
tool_definitions=tool_definitions,
)
model_messages = prepared
supplemental_messages = None
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 request_was_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
@staticmethod @staticmethod
def input_budget(config: ContextGovernanceConfig) -> int: def input_budget(config: ContextGovernanceConfig) -> int:
if not config.context_window_tokens: if not config.context_window_tokens:
@@ -424,13 +886,14 @@ class ContextGovernor:
if budget <= 0: if budget <= 0:
return messages return messages
if not force:
estimate, _ = estimate_prompt_tokens_chain( estimate, _ = estimate_prompt_tokens_chain(
config.provider, config.provider,
config.model, config.model,
messages, messages,
tool_definitions, tool_definitions,
) )
if not force and estimate <= budget: if 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"]
+119 -38
View File
@@ -13,6 +13,7 @@ import weakref
from collections.abc import Coroutine, Iterable, Mapping 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 datetime import datetime
from enum import Enum, auto from enum import Enum, auto
from functools import partial from functools import partial
from pathlib import Path from pathlib import Path
@@ -93,7 +94,11 @@ from nanobot.session.recovery import (
restore_pending_interruption, restore_pending_interruption,
restore_runtime_checkpoint, restore_runtime_checkpoint,
) )
from nanobot.session.summary import SessionSummary from nanobot.session.summary import (
SUMMARY_CONTINUATION_TEXT,
SessionSummary,
SessionSummaryCheckpoint,
)
from nanobot.triggers.local_turns import LocalTriggerTurnCoordinator from nanobot.triggers.local_turns import LocalTriggerTurnCoordinator
from nanobot.utils.cancellation import task_is_cancelling from nanobot.utils.cancellation import task_is_cancelling
from nanobot.utils.document import reference_non_image_attachments from nanobot.utils.document import reference_non_image_attachments
@@ -161,6 +166,8 @@ class TurnContext:
pending_queue: asyncio.Queue[InboundMessage] | None = None pending_queue: asyncio.Queue[InboundMessage] | None = None
pending_summary: SessionSummary | None = None pending_summary: SessionSummary | None = None
summary_checkpoint: SessionSummaryCheckpoint | None = None
provider_compaction_applied: bool = False
ephemeral: bool = False ephemeral: bool = False
run_extra_hooks_for_ephemeral: bool = False run_extra_hooks_for_ephemeral: bool = False
@@ -923,19 +930,6 @@ class AgentLoop:
return return
remember_last_channel(session.metadata, msg.channel, msg.chat_id) remember_last_channel(session.metadata, msg.channel, msg.chat_id)
@staticmethod
def _replay_token_budget(runtime: LLMRuntime) -> int:
"""Derive a token budget for session history replay from the context window."""
if runtime.context_window_tokens <= 0:
return 0
max_output = runtime.generation.max_tokens
try:
reserved_output = int(max_output)
except (TypeError, ValueError):
reserved_output = 4096
budget = runtime.context_window_tokens - max(1, reserved_output) - 1024
return budget if budget > 0 else max(128, runtime.context_window_tokens // 2)
async def _run_agent_loop( async def _run_agent_loop(
self, self,
transcript_input: TranscriptInput, transcript_input: TranscriptInput,
@@ -1186,6 +1180,26 @@ class AgentLoop:
provider_retry_mode=self.provider_retry_mode, provider_retry_mode=self.provider_retry_mode,
retry_wait_callback=on_retry_wait, retry_wait_callback=on_retry_wait,
checkpoint_callback=_checkpoint, checkpoint_callback=_checkpoint,
consolidate_history=(
partial(
self.consolidator.summarize_transcript,
runtime=runtime,
session_key=session.key,
tools=effective_tools.get_definitions(),
)
if session is not None and not ephemeral
else None
),
consolidate_provider_compaction=(
partial(
self.consolidator.summarize_provider_compaction,
runtime=runtime,
session_key=session.key,
tools=effective_tools.get_definitions(),
)
if session is not None and not ephemeral
else None
),
injection_callback=_drain_pending, injection_callback=_drain_pending,
terminal_injection_callback=_wait_for_pending, terminal_injection_callback=_wait_for_pending,
# Sustained goals may legitimately exceed NANOBOT_LLM_TIMEOUT_S; idle stall # Sustained goals may legitimately exceed NANOBOT_LLM_TIMEOUT_S; idle stall
@@ -1886,12 +1900,6 @@ class AgentLoop:
if ctx.on_runtime_admitted is not None: if ctx.on_runtime_admitted is not None:
await ctx.on_runtime_admitted(runtime) await ctx.on_runtime_admitted(runtime)
if not ctx.ephemeral: if not ctx.ephemeral:
await self.consolidator.maybe_consolidate_by_tokens(
session,
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( ctx.session, ctx.pending_summary = self.auto_compact.prepare_session(
session, session,
ctx.session_key, ctx.session_key,
@@ -1899,11 +1907,7 @@ class AgentLoop:
session = ctx.require_session() 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] = { ctx.history = session.get_history(extend_to_user=is_subagent)
"max_tokens": self._replay_token_budget(runtime),
"extend_to_user": is_subagent,
}
ctx.history = session.get_history(**_hist_kwargs)
stored_state = session.provider_state stored_state = session.provider_state
subagent_followup_persisted = False subagent_followup_persisted = False
if is_subagent: if is_subagent:
@@ -2021,6 +2025,8 @@ class AgentLoop:
) )
ctx.final_content = result.final_content ctx.final_content = result.final_content
ctx.all_messages = result.messages ctx.all_messages = result.messages
ctx.summary_checkpoint = result.summary_checkpoint
ctx.provider_compaction_applied = result.provider_compaction_applied
ctx.stop_reason = result.stop_reason ctx.stop_reason = result.stop_reason
if ( if (
ctx.kind is TurnKind.USER ctx.kind is TurnKind.USER
@@ -2034,7 +2040,6 @@ class AgentLoop:
await turn_continuation.maybe_continue_turn(ctx) await turn_continuation.maybe_continue_turn(ctx)
async def _persist_turn(self, ctx: TurnContext) -> None: async def _persist_turn(self, ctx: TurnContext) -> None:
runtime = ctx.require_runtime()
session = ctx.require_session() session = ctx.require_session()
turn_continuation.prepare_save_boundary(ctx) turn_continuation.prepare_save_boundary(ctx)
@@ -2060,15 +2065,18 @@ class AgentLoop:
self._save_turn( self._save_turn(
session, ctx.all_messages, ctx.save_skip, session, ctx.all_messages, ctx.save_skip,
turn_latency_ms=ctx.turn_latency_ms, turn_latency_ms=ctx.turn_latency_ms,
summary_checkpoint=ctx.summary_checkpoint,
input_persisted_early=ctx.input_persisted_early,
) )
if (
not ctx.ephemeral
and ctx.provider_compaction_applied
and ctx.summary_checkpoint is not None
):
# The next request must rebuild from the portable checkpoint;
# the opaque continuation predates that transcript rewrite.
session.provider_state = None
ctx.delivery.record_latency(ctx.turn_latency_ms) ctx.delivery.record_latency(ctx.turn_latency_ms)
if not ctx.ephemeral:
self.schedule_background(
self.consolidator.maybe_consolidate_by_tokens(
session,
runtime=runtime,
)
)
self._clear_pending_user_turn(session) self._clear_pending_user_turn(session)
self._clear_runtime_checkpoint(session) self._clear_runtime_checkpoint(session)
self.sessions.save(session) self.sessions.save(session)
@@ -2142,6 +2150,55 @@ class AgentLoop:
return filtered return filtered
@staticmethod
def _insert_summary_checkpoint(
session: Session,
checkpoint: SessionSummaryCheckpoint,
*,
insert_at: int | None = None,
) -> None:
"""Commit a replacement summary and its hidden transcript boundary."""
hint = {
"role": "user",
"content": SUMMARY_CONTINUATION_TEXT,
HIDDEN_HISTORY_META: True,
"timestamp": datetime.now().isoformat(),
}
if insert_at is None:
session.messages.append(hint)
checkpoint_session_index = len(session.messages) - 1
else:
session.messages.insert(insert_at, hint)
checkpoint_session_index = insert_at
session.metadata["_last_summary"] = {
"text": checkpoint.summary,
"last_active": session.updated_at.isoformat(),
}
session.last_archived = checkpoint_session_index
@staticmethod
def _validated_checkpoint_boundary(
checkpoint: SessionSummaryCheckpoint | None,
*,
skip: int,
message_count: int,
session_key: str,
) -> int | None:
"""Return a checkpoint boundary only when it belongs to this turn."""
if checkpoint is None:
return None
boundary = checkpoint.transcript_boundary
if skip - 1 <= boundary <= message_count:
return boundary
logger.warning(
"Ignoring invalid summary boundary {} outside [{}, {}] for {}",
boundary,
skip - 1,
message_count,
session_key,
)
return None
def _save_turn( def _save_turn(
self, self,
session: Session, session: Session,
@@ -2149,10 +2206,10 @@ class AgentLoop:
skip: int, skip: int,
*, *,
turn_latency_ms: int | None = None, turn_latency_ms: int | None = None,
summary_checkpoint: SessionSummaryCheckpoint | None = None,
input_persisted_early: bool = False,
) -> None: ) -> None:
"""Save new-turn messages into session, truncating large tool results.""" """Commit new-turn messages and an optional summary boundary."""
from datetime import datetime
declared_tool_call_ids = { declared_tool_call_ids = {
str(tc["id"]) str(tc["id"])
for m in session.messages for m in session.messages
@@ -2169,8 +2226,30 @@ class AgentLoop:
} }
last_assistant_idx: int | None = None last_assistant_idx: int | None = None
saved_followup_ids: set[str] = set() saved_followup_ids: set[str] = set()
for m in messages[skip:]: checkpoint_boundary = self._validated_checkpoint_boundary(
entry = dict(m) summary_checkpoint,
skip=skip,
message_count=len(messages),
session_key=session.key,
)
# The trigger input may already be the session tail while still being
# the first message after the replacement checkpoint.
if summary_checkpoint is not None and checkpoint_boundary == skip - 1:
insert_at = len(session.messages) - (1 if input_persisted_early else 0)
self._insert_summary_checkpoint(
session,
summary_checkpoint,
insert_at=insert_at,
)
for message_index, message in enumerate(messages[skip:], start=skip):
# Insert against the raw transcript index before filtering the
# message so persistence cleanup cannot shift the H/Δ boundary.
if summary_checkpoint is not None and checkpoint_boundary == message_index:
self._insert_summary_checkpoint(session, summary_checkpoint)
entry = dict(message)
followup_id_value = cast(object, entry.pop(PENDING_FOLLOWUP_ID_KEY, None)) followup_id_value = cast(object, entry.pop(PENDING_FOLLOWUP_ID_KEY, None))
followup_ids = ( followup_ids = (
[followup_id_value] [followup_id_value]
@@ -2249,6 +2328,8 @@ class AgentLoop:
for tc in (cast(dict[str, Any], tc_value),) for tc in (cast(dict[str, Any], tc_value),)
if tc.get("id") if tc.get("id")
) )
if summary_checkpoint is not None and checkpoint_boundary == len(messages):
self._insert_summary_checkpoint(session, summary_checkpoint)
if turn_latency_ms is not None and last_assistant_idx is not None: if turn_latency_ms is not None and last_assistant_idx is not None:
session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms) session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms)
if saved_followup_ids: if saved_followup_ids:
+164 -124
View File
@@ -1,4 +1,4 @@
"""Memory storage, transcript archiving, and legacy consolidation coordination.""" """Memory storage, transcript archiving, and session checkpoint consolidation."""
# Tool schemas are installed by the ``@tool_parameters`` class decorator at # Tool schemas are installed by the ``@tool_parameters`` class decorator at
# runtime; static analyzers cannot observe that it clears ``parameters`` from # runtime; static analyzers cannot observe that it clears ``parameters`` from
@@ -21,6 +21,7 @@ from typing import TYPE_CHECKING, Any, Callable, Iterator, cast
from loguru import logger from loguru import logger
from nanobot.llm_usage.context import llm_usage_source from nanobot.llm_usage.context import llm_usage_source
from nanobot.providers.base import ProviderCallContext, ProviderConversationState
from nanobot.runtime_context import public_history_messages from nanobot.runtime_context import public_history_messages
from nanobot.session.manager import ( from nanobot.session.manager import (
MIN_COMPACTED_REPLAY_MESSAGES, MIN_COMPACTED_REPLAY_MESSAGES,
@@ -740,7 +741,7 @@ class MemoryStore:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Memory ingestion and legacy context-pressure coordination # Memory ingestion and context-pressure coordination
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Raw fallbacks use a tighter cap. Completed model summaries may scale with the # Raw fallbacks use a tighter cap. Completed model summaries may scale with the
@@ -780,6 +781,20 @@ class MemoryArchiver:
) -> str: ) -> str:
"""Persist the failed chunk and return a bounded replacement checkpoint.""" """Persist the failed chunk and return a bounded replacement checkpoint."""
raw = self.store.raw_archive(messages, session_key=session_key) raw = self.store.raw_archive(messages, session_key=session_key)
return self._combine_raw_checkpoint(
raw,
previous_summary=previous_summary,
max_tokens=max_tokens,
)
@staticmethod
def _combine_raw_checkpoint(
raw: str,
*,
previous_summary: str | None,
max_tokens: int,
) -> str:
"""Return a bounded checkpoint that preserves prior and newly archived context."""
token_limit = max(1, max_tokens) token_limit = max(1, max_tokens)
if not previous_summary: if not previous_summary:
return truncate_text_to_tokens(raw, token_limit) return truncate_text_to_tokens(raw, token_limit)
@@ -806,35 +821,94 @@ class MemoryArchiver:
async def archive( async def archive(
self, self,
messages: list[dict[str, Any]], source_messages: list[dict[str, Any]],
*, *,
runtime: LLMRuntime, runtime: LLMRuntime,
session_key: str, session_key: str,
request_messages: list[dict[str, Any]], history: list[dict[str, Any]],
request_tools: list[dict[str, Any]], request_tools: list[dict[str, Any]],
previous_summary: str | None = None, previous_summary: str | None = None,
input_token_budget: int | None = None,
fallback_max_tokens: int | None = None,
provider_state: ProviderConversationState | None = None,
) -> str | None: ) -> str | None:
"""Execute a prepared archive request and persist its result.""" """Append the archive prompt to H and persist its summary."""
if not messages: if not source_messages:
return None return None
def raw_fallback() -> str: def raw_fallback() -> str:
return self._raw_checkpoint( return self._raw_checkpoint(
messages, source_messages,
session_key=session_key, session_key=session_key,
previous_summary=previous_summary, previous_summary=previous_summary,
max_tokens=runtime.generation.max_tokens, max_tokens=(
fallback_max_tokens
if fallback_max_tokens is not None
else runtime.generation.max_tokens
),
) )
prompt = render_template(
"agent/consolidator_archive.md",
strip=True,
archive_count=len(source_messages),
)
prompt_message = {"role": "user", "content": prompt}
provider_context = None
call_tools = request_tools
if provider_state is not None:
if not runtime.provider.can_resume_conversation_state(
provider_state,
runtime.model,
):
return raw_fallback()
instruction_messages: list[dict[str, Any]] = []
for message in history:
if message.get("role") not in {"system", "developer"}:
break
instruction_messages.append(dict(message))
request_messages = [*instruction_messages, prompt_message]
provider_context = ProviderCallContext(
conversation_state=provider_state.with_pending_messages([
*provider_state.pending_messages,
prompt_message,
]),
context_window_tokens=runtime.context_window_tokens,
session_id=session_key,
)
call_tools = []
else:
request_messages = [
*[dict(message) for message in history],
prompt_message,
]
if input_token_budget is not None and provider_context is None:
estimated, source = estimate_prompt_tokens_chain(
runtime.provider,
runtime.model,
request_messages,
call_tools,
)
if input_token_budget <= 0 or estimated > input_token_budget:
logger.debug(
"Memory archive input does not fit for {}: {}/{} via {}; raw-dumping",
session_key,
estimated,
input_token_budget,
source,
)
return raw_fallback()
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=call_tools,
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,
provider_context=provider_context,
) )
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")
@@ -879,20 +953,17 @@ class MemoryArchiver:
) )
previous_summary = session_summary["text"] if session_summary else None previous_summary = session_summary["text"] if session_summary else None
def raw_fallback() -> str: if input_token_budget <= 0:
logger.debug(
"Memory archive has no safe input budget for {}; raw-dumping",
session.key,
)
return self._raw_checkpoint( return self._raw_checkpoint(
messages, messages,
session_key=session.key, session_key=session.key,
previous_summary=previous_summary, previous_summary=previous_summary,
max_tokens=runtime.generation.max_tokens, max_tokens=runtime.generation.max_tokens,
) )
if input_token_budget <= 0:
logger.debug(
"Memory archive has no safe input budget for {}; raw-dumping",
session.key,
)
return raw_fallback()
prefix = Session( prefix = Session(
key=session.key, key=session.key,
messages=list(session.messages[:archive_end]), messages=list(session.messages[:archive_end]),
@@ -908,47 +979,37 @@ 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,
) )
return raw_fallback() return self._raw_checkpoint(
prompt = render_template("agent/consolidator_archive.md", strip=True) messages,
session_key=session.key,
previous_summary=previous_summary,
max_tokens=runtime.generation.max_tokens,
)
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:
channel, workspace = self._resolve_prompt_context(session) channel, workspace = self._resolve_prompt_context(session)
request_messages = self._build_messages( history_messages = self._build_messages(
history=history, history=history,
current_message=prompt, current_message=None,
channel=channel, channel=channel,
session_summary=session_summary, session_summary=session_summary,
workspace=workspace, workspace=workspace,
) )
tools = self._get_tool_definitions() tools = self._get_tool_definitions()
estimated, source = estimate_prompt_tokens_chain(
runtime.provider,
runtime.model,
request_messages,
tools,
)
if estimated > input_token_budget:
logger.debug(
"Memory archive prefix exceeds budget for {}; raw-dumping: {}/{} via {}",
session.key,
estimated,
input_token_budget,
source,
)
return raw_fallback()
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, history=history_messages,
request_tools=tools, request_tools=tools,
previous_summary=previous_summary, previous_summary=previous_summary,
input_token_budget=input_token_budget,
) )
class Consolidator: class Consolidator:
"""Legacy context-pressure coordinator backed by a MemoryArchiver.""" """Coordinate session Memory checkpoints through ``MemoryArchiver``."""
_SAFETY_BUFFER = 1024 # extra headroom for tokenizer estimation drift _SAFETY_BUFFER = 1024 # extra headroom for tokenizer estimation drift
@@ -978,22 +1039,73 @@ class Consolidator:
"""Return the shared consolidation lock for one session.""" """Return the shared consolidation lock for one session."""
return self._locks.setdefault(session_key, asyncio.Lock()) return self._locks.setdefault(session_key, asyncio.Lock())
def pick_consolidation_boundary( async def summarize_transcript(
self, self,
session: Session, accepted_messages: list[dict[str, Any]],
) -> int | None: previous_summary: str | None,
"""Return the fixed user-led boundary before the recent replay tail.""" *,
if not session.messages: runtime: LLMRuntime,
session_key: str,
tools: list[dict[str, Any]],
provider_state: ProviderConversationState | None = None,
) -> str | None:
"""Summarize the exact transcript prefix already accepted by the model."""
source_messages = [
dict(message)
for message in accepted_messages
if message.get("role") != "system"
]
if not source_messages:
return None return None
boundary = max(0, len(session.messages) - MIN_COMPACTED_REPLAY_MESSAGES)
while boundary > 0 and session.messages[boundary].get("role") != "user": max_output_tokens = max(0, runtime.generation.max_tokens)
boundary -= 1 input_token_budget = runtime.context_window_tokens - max_output_tokens
if ( checkpoint_tokens = min(
boundary <= session.last_archived max_output_tokens,
or session.messages[boundary].get("role") != "user" max(1, (input_token_budget - self._SAFETY_BUFFER) // 2),
): )
summary = await self.archiver.archive(
source_messages,
runtime=runtime,
session_key=session_key,
history=accepted_messages,
request_tools=tools,
previous_summary=previous_summary,
input_token_budget=input_token_budget,
fallback_max_tokens=max(1, checkpoint_tokens),
provider_state=provider_state,
)
if summary == "(nothing)":
summary = self.archiver._raw_checkpoint(
source_messages,
session_key=session_key,
previous_summary=previous_summary,
max_tokens=max_output_tokens,
)
if summary is None:
return None return None
return boundary return truncate_text_to_tokens(summary, max(1, max_output_tokens))
async def summarize_provider_compaction(
self,
state: ProviderConversationState,
fallback_messages: list[dict[str, Any]],
previous_summary: str | None,
*,
runtime: LLMRuntime,
session_key: str,
tools: list[dict[str, Any]],
) -> str | None:
"""Prompt a native compacted state without replaying its raw history."""
return await self.summarize_transcript(
fallback_messages,
previous_summary,
runtime=runtime,
session_key=session_key,
tools=tools,
provider_state=state,
)
@staticmethod @staticmethod
def _full_replay_history( def _full_replay_history(
@@ -1058,7 +1170,7 @@ class Consolidator:
archive_end: int, archive_end: int,
runtime: LLMRuntime, runtime: LLMRuntime,
) -> str | None: ) -> str | None:
"""Compatibility wrapper for the extracted MemoryArchiver.""" """Archive one captured session range through the shared Memory path."""
return await self.archiver.archive_session( return await self.archiver.archive_session(
session, session,
archive_end=archive_end, archive_end=archive_end,
@@ -1066,78 +1178,6 @@ class Consolidator:
input_token_budget=self._input_token_budget(runtime), input_token_budget=self._input_token_budget(runtime),
) )
async def maybe_consolidate_by_tokens(
self,
session: Session,
*,
runtime: LLMRuntime,
) -> None:
"""Archive one fixed old prefix when the prompt exceeds the safe budget.
The budget reserves space for completion tokens and a safety buffer
so the LLM request never exceeds the context window.
"""
lock = self.get_lock(session.key)
async with lock:
# Refresh session reference: AutoCompact may have replaced it.
fresh = self.sessions.get_or_create(session.key)
if fresh is not session:
session = fresh
if runtime.context_window_tokens <= 0:
return
if not session.messages:
return
budget = self._input_token_budget(runtime)
estimated, source = self.estimate_session_prompt_tokens(
session,
runtime=runtime,
)
if estimated <= 0:
return
if estimated < budget:
unarchived_count = len(session.messages) - session.last_archived
logger.debug(
"Token consolidation idle {}: {}/{} via {}, msgs={}",
session.key,
estimated,
runtime.context_window_tokens,
source,
unarchived_count,
)
return
end_idx = self.pick_consolidation_boundary(session)
if end_idx is None:
logger.debug(
"Token consolidation: no safe fixed boundary for {}",
session.key,
)
return
chunk = session.messages[session.last_archived:end_idx]
if not chunk:
return
logger.info(
"Token consolidation for {}: {}/{} via {}, chunk={} msgs",
session.key,
estimated,
runtime.context_window_tokens,
source,
len(chunk),
)
summary = await self.archive_session(
session,
archive_end=end_idx,
runtime=runtime,
)
if summary is None:
return
self._set_last_summary(session, summary)
session.last_archived = end_idx
self.sessions.save(session)
async def compact_idle_session( async def compact_idle_session(
self, self,
session_key: str, session_key: str,
+81 -205
View File
@@ -16,8 +16,13 @@ from loguru import logger
from nanobot.agent.context import TranscriptInput from nanobot.agent.context import TranscriptInput
from nanobot.agent.context_governance import ( from nanobot.agent.context_governance import (
ContextCompactionState,
ContextGovernanceConfig, ContextGovernanceConfig,
ContextGovernor, ContextGovernor,
HistoryConsolidator,
ModelRequestState,
ProviderCompactionConsolidator,
TranscriptBuilder,
) )
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
from nanobot.agent.tools.execution import execute_tool_calls from nanobot.agent.tools.execution import execute_tool_calls
@@ -32,20 +37,10 @@ from nanobot.providers.base import (
LLMProvider, LLMProvider,
LLMResponse, LLMResponse,
LLMUsage, LLMUsage,
ProviderCallContext,
ProviderConversationState, ProviderConversationState,
) )
from nanobot.providers.conversation_state import ( from nanobot.providers.conversation_state import ProviderConversationStateController
ProviderConversationStateController, from nanobot.session.summary import SessionSummaryCheckpoint
allows_conversation_message_merge,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_MESSAGE_META,
detach_runtime_context,
reattach_runtime_context,
)
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY
from nanobot.utils.helpers import ( from nanobot.utils.helpers import (
build_assistant_message, build_assistant_message,
estimate_message_tokens, estimate_message_tokens,
@@ -67,7 +62,6 @@ ContinuationCallback = Callable[[], str | None]
RetryWaitCallback = Callable[[str], Awaitable[None]] RetryWaitCallback = Callable[[str], Awaitable[None]]
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]] CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]] InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
TranscriptBuilder = Callable[[TranscriptInput], list[dict[str, Any]]]
_DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model." _DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model."
_ARREARAGE_ERROR_MESSAGE = ( _ARREARAGE_ERROR_MESSAGE = (
@@ -113,6 +107,8 @@ class AgentRunSpec:
provider_retry_mode: str = "standard" provider_retry_mode: str = "standard"
retry_wait_callback: RetryWaitCallback | None = None retry_wait_callback: RetryWaitCallback | None = None
checkpoint_callback: CheckpointCallback | None = None checkpoint_callback: CheckpointCallback | None = None
consolidate_history: HistoryConsolidator | None = None
consolidate_provider_compaction: ProviderCompactionConsolidator | None = None
injection_callback: InjectionCallback | None = None injection_callback: InjectionCallback | None = None
terminal_injection_callback: InjectionCallback | None = None terminal_injection_callback: InjectionCallback | None = None
llm_timeout_s: float | None = None llm_timeout_s: float | None = None
@@ -137,17 +133,8 @@ class AgentRunResult:
# Terminal tail to emit when the preceding final-content prefix was already streamed. # Terminal tail to emit when the preceding final-content prefix was already streamed.
pending_stream_content: str | None = None pending_stream_content: str | None = None
provider_state: ProviderConversationState | None = field(default=None, repr=False) provider_state: ProviderConversationState | None = field(default=None, repr=False)
summary_checkpoint: SessionSummaryCheckpoint | None = field(default=None, repr=False)
provider_compaction_applied: bool = field(default=False, 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:
@@ -157,118 +144,12 @@ class AgentRunner:
self.context_governor = ContextGovernor() self.context_governor = ContextGovernor()
@staticmethod @staticmethod
def _merge_message_content(left: Any, right: Any) -> str | list[dict[str, Any]]:
if isinstance(left, str) and isinstance(right, str):
return f"{left}\n\n{right}" if left else right
def _to_blocks(value: Any) -> list[dict[str, Any]]:
if isinstance(value, list):
return [
cast(dict[str, Any], item)
if isinstance(item, dict)
else {"type": "text", "text": str(item)}
for item in cast(list[Any], value)
]
if value is None:
return []
return [{"type": "text", "text": str(value)}]
return _to_blocks(left) + _to_blocks(right)
@classmethod
def _append_injected_messages( def _append_injected_messages(
cls,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
injections: list[dict[str, Any]], injections: list[dict[str, Any]],
) -> None: ) -> None:
"""Append injected user messages while preserving role alternation.""" """Append injected messages without rewriting the raw transcript."""
for injection in injections: messages.extend(injections)
if (
messages
and injection.get("role") == "user"
and messages[-1].get("role") == "user"
and not is_hidden_history_message(injection)
and not is_hidden_history_message(messages[-1])
and allows_conversation_message_merge(messages[-1])
):
merged = dict(messages[-1])
left_meta = merged.get("_meta")
right_meta = injection.get("_meta")
left_meta_dict = cast(dict[str, Any], left_meta) if isinstance(left_meta, dict) else None
right_meta_dict = (
cast(dict[str, Any], right_meta) if isinstance(right_meta, dict) else None
)
left_marker = (
left_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
if left_meta_dict is not None
else None
)
right_marker = (
right_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
if right_meta_dict is not None
else None
)
left_marker_dict = (
cast(dict[str, Any], left_marker) if isinstance(left_marker, dict) else None
)
right_marker_dict = (
cast(dict[str, Any], right_marker) if isinstance(right_marker, dict) else None
)
empty_sources: list[str] = []
empty_blocks: list[dict[str, Any]] = []
detached_left = (
detach_runtime_context(merged.get("content"), left_marker_dict)
if left_marker_dict is not None
else (merged.get("content"), empty_sources, empty_blocks)
)
detached_right = (
detach_runtime_context(injection.get("content"), right_marker_dict)
if right_marker_dict is not None
else (injection.get("content"), empty_sources, empty_blocks)
)
if detached_left is not None and detached_right is not None:
left_content, left_sources, left_blocks = detached_left
right_content, right_sources, right_blocks = detached_right
merged_content = cls._merge_message_content(left_content, right_content)
context_blocks = [*left_blocks, *right_blocks]
if context_blocks:
merged_content, marker = reattach_runtime_context(
merged_content,
[*left_sources, *right_sources],
context_blocks,
)
internal_meta = dict(left_meta_dict) if left_meta_dict is not None else {}
if right_meta_dict is not None:
for key, value in right_meta_dict.items():
internal_meta.setdefault(key, value)
internal_meta[RUNTIME_CONTEXT_MESSAGE_META] = marker
merged["_meta"] = internal_meta
merged["content"] = merged_content
else:
merged["content"] = cls._merge_message_content(
merged.get("content"),
injection.get("content"),
)
followup_id = injection.get(PENDING_FOLLOWUP_ID_KEY)
if isinstance(followup_id, str) and followup_id:
existing = cast(object, merged.get(PENDING_FOLLOWUP_ID_KEY))
followup_ids = (
[existing]
if isinstance(existing, str)
else [
item
for item in cast(list[object], existing)
if isinstance(item, str)
]
if isinstance(existing, list)
else []
)
if followup_id not in followup_ids:
followup_ids.append(followup_id)
merged[PENDING_FOLLOWUP_ID_KEY] = followup_ids
messages[-1] = merged
continue
messages.append(injection)
async def _try_drain_injections( async def _try_drain_injections(
self, self,
@@ -425,7 +306,7 @@ class AgentRunner:
async def run(self, spec: AgentRunSpec) -> AgentRunResult: async def run(self, spec: AgentRunSpec) -> AgentRunResult:
hook = spec.hook or AgentHook() hook = spec.hook or AgentHook()
messages = self._initial_transcript(spec) messages, compaction = self._initial_transcript_and_compaction(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)
@@ -433,7 +314,7 @@ class AgentRunner:
try: try:
await hook.before_run(context) await hook.before_run(context)
result = await self._run_core(spec, hook, messages) result = await self._run_core(spec, hook, messages, compaction)
except asyncio.CancelledError as exc: except asyncio.CancelledError as exc:
context.messages = deepcopy(messages) context.messages = deepcopy(messages)
context.stop_reason = "cancelled" context.stop_reason = "cancelled"
@@ -478,23 +359,35 @@ class AgentRunner:
reset_llm_usage_source(llm_usage_source_token) reset_llm_usage_source(llm_usage_source_token)
@staticmethod @staticmethod
def _initial_transcript(spec: AgentRunSpec) -> list[dict[str, Any]]: def _initial_transcript_and_compaction(
"""Resolve exactly one supported source for the initial model transcript.""" spec: AgentRunSpec,
if spec.transcript_input is not None: ) -> tuple[list[dict[str, Any]], ContextCompactionState | None]:
"""Build the initial transcript and its optional compaction state."""
transcript_input = spec.transcript_input
if transcript_input is not None:
if spec.initial_messages is not None: if spec.initial_messages is not None:
raise ValueError("provide either transcript_input or initial_messages, not both") raise ValueError("provide either transcript_input or initial_messages, not both")
if spec.transcript_builder is None: transcript_builder = spec.transcript_builder
if transcript_builder is None:
raise ValueError("transcript_builder is required with transcript_input") raise ValueError("transcript_builder is required with transcript_input")
return list(spec.transcript_builder(spec.transcript_input)) return ContextCompactionState.from_transcript(
transcript_input,
transcript_builder,
spec.consolidate_history,
spec.consolidate_provider_compaction,
)
if spec.initial_messages is None: if spec.initial_messages is None:
raise ValueError("initial_messages is required without transcript_input") raise ValueError("initial_messages is required without transcript_input")
return list(spec.initial_messages) if spec.consolidate_history is not None:
raise ValueError("consolidate_history requires transcript_input")
return list(spec.initial_messages), None
async def _run_core( async def _run_core(
self, self,
spec: AgentRunSpec, spec: AgentRunSpec,
hook: AgentHook, hook: AgentHook,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
compaction: ContextCompactionState | None,
) -> AgentRunResult: ) -> AgentRunResult:
final_content: str | None = None final_content: str | None = None
tools_used: list[str] = [] tools_used: list[str] = []
@@ -530,9 +423,10 @@ class AgentRunner:
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,
) )
request_state = _ModelRequestState( request_state = ModelRequestState(
config=governance_config, config=governance_config,
conversation=conversation_state, conversation=conversation_state,
compaction=compaction,
) )
for iteration in range(spec.max_iterations): for iteration in range(spec.max_iterations):
@@ -542,9 +436,15 @@ class AgentRunner:
session_key=spec.session_key, session_key=spec.session_key,
) )
await hook.before_iteration(context) await hook.before_iteration(context)
request_message_count = len(messages)
request_messages = (
request_state.compaction.request_messages(messages)
if request_state.compaction is not None
else messages
)
response = await self._request_model( response = await self._request_model(
spec, spec,
messages, request_messages,
hook, hook,
context, context,
request_state=request_state, request_state=request_state,
@@ -553,6 +453,11 @@ class AgentRunner:
assert request_state.messages is not None assert request_state.messages is not None
messages_for_model = request_state.messages messages_for_model = request_state.messages
conversation_state.observe_response(response, messages) conversation_state.observe_response(response, messages)
if request_state.compaction is not None:
request_state.compaction.accept_request(
messages_for_model,
raw_boundary=request_message_count,
)
context.response = response context.response = response
context.tool_calls = list(response.tool_calls) context.tool_calls = list(response.tool_calls)
@@ -634,7 +539,7 @@ class AgentRunner:
messages.append(tool_message) messages.append(tool_message)
completed_tool_results.append(tool_message) completed_tool_results.append(tool_message)
checkpoint_model_messages = ( checkpoint_model_messages = (
self.context_governor.prepare_for_model( self.context_governor.prepare_messages_for_model(
governance_config, governance_config,
messages, messages,
) )
@@ -920,6 +825,12 @@ class AgentRunner:
had_injections=had_injections, had_injections=had_injections,
pending_stream_content=pending_stream_content, pending_stream_content=pending_stream_content,
provider_state=conversation_state.finish(messages), provider_state=conversation_state.finish(messages),
summary_checkpoint=(
request_state.compaction.summary_checkpoint
if request_state.compaction is not None
else None
),
provider_compaction_applied=request_state.provider_compaction_applied,
) )
def _build_request_kwargs( def _build_request_kwargs(
@@ -942,60 +853,6 @@ 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,
@@ -1003,13 +860,13 @@ class AgentRunner:
hook: AgentHook, hook: AgentHook,
context: AgentHookContext, context: AgentHookContext,
*, *,
request_state: _ModelRequestState, request_state: ModelRequestState,
malformed_retry: bool = False, malformed_retry: bool = False,
transcript: list[dict[str, Any]] | None, transcript: list[dict[str, Any]] | 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() tool_definitions = spec.tools.get_definitions()
messages, provider_context = self._prepare_model_request( messages, provider_context = await self.context_governor.prepare_request(
request_state, request_state,
messages, messages,
tool_definitions=tool_definitions, tool_definitions=tool_definitions,
@@ -1169,6 +1026,12 @@ class AgentRunner:
response.ttft_ms = max(0, round((first_output_at - request_started_at) * 1000)) response.ttft_ms = max(0, round((first_output_at - request_started_at) * 1000))
if generation_elapsed_s > 0: if generation_elapsed_s > 0:
response.generation_ms = max(1, round(generation_elapsed_s * 1000)) response.generation_ms = max(1, round(generation_elapsed_s * 1000))
await self.context_governor.summarize_provider_compaction(
request_state,
response,
current_request_boundary=(len(transcript) if transcript is not None else None),
)
request_state.provider_compaction_applied |= response.provider_compaction_applied
# chat_stream_with_retry may recover internally, so only fail unfinished # chat_stream_with_retry may recover internally, so only fail unfinished
# hosted calls after the provider returns its final error response. # hosted calls after the provider returns its final error response.
if response.finish_reason == "error": if response.finish_reason == "error":
@@ -1283,7 +1146,7 @@ class AgentRunner:
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
*, *,
request_state: _ModelRequestState, request_state: ModelRequestState,
transcript: list[dict[str, Any]], transcript: list[dict[str, Any]],
) -> LLMResponse: ) -> LLMResponse:
retry_messages = self._finalization_retry_messages(messages) retry_messages = self._finalization_retry_messages(messages)
@@ -1313,14 +1176,21 @@ class AgentRunner:
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
usage: LLMUsage | None, usage: LLMUsage | None,
*, *,
request_state: _ModelRequestState, request_state: ModelRequestState,
) -> tuple[str | None, LLMUsage | None]: ) -> tuple[str | None, LLMUsage | None]:
retry_messages = self._budget_exhausted_finalization_messages(messages) compaction = request_state.compaction
request_messages = (
compaction.request_messages(messages)
if compaction is not None
else messages
)
retry_messages = self._budget_exhausted_finalization_messages(request_messages)
try: try:
response = await self._request_no_tools( response = await self._request_no_tools(
spec, spec,
retry_messages, retry_messages,
request_state=request_state, request_state=request_state,
transcript=messages if compaction is not None else None,
) )
except Exception: except Exception:
logger.exception( logger.exception(
@@ -1358,10 +1228,10 @@ class AgentRunner:
spec: AgentRunSpec, spec: AgentRunSpec,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
*, *,
request_state: _ModelRequestState, request_state: ModelRequestState,
transcript: list[dict[str, Any]] | None = None, transcript: list[dict[str, Any]] | None = None,
) -> LLMResponse: ) -> LLMResponse:
messages, provider_context = self._prepare_model_request( messages, provider_context = await self.context_governor.prepare_request(
request_state, request_state,
messages, messages,
tool_definitions=None, tool_definitions=None,
@@ -1389,6 +1259,12 @@ class AgentRunner:
finish_reason="error", finish_reason="error",
error_kind="timeout", error_kind="timeout",
) )
await self.context_governor.summarize_provider_compaction(
request_state,
response,
current_request_boundary=(len(transcript) if transcript is not None else None),
)
request_state.provider_compaction_applied |= response.provider_compaction_applied
return response return response
@staticmethod @staticmethod
@@ -1453,7 +1329,7 @@ class AgentRunner:
def _record_request_usage( def _record_request_usage(
self, self,
spec: AgentRunSpec, spec: AgentRunSpec,
state: _ModelRequestState, state: ModelRequestState,
response: LLMResponse, response: LLMResponse,
) -> LLMUsage | None: ) -> LLMUsage | None:
assert state.messages is not None assert state.messages is not None
-5
View File
@@ -662,11 +662,6 @@ def _run_gateway(
if isinstance(message_tool, MessageTool) and suppress_token is not None: if isinstance(message_tool, MessageTool) and suppress_token is not None:
message_tool.reset_suppress_delivery(suppress_token) message_tool.reset_suppress_delivery(suppress_token)
# Keep a small tail of heartbeat history so the loop stays bounded.
session = agent.sessions.get_or_create("heartbeat")
session.retain_recent_legal_suffix(hb_cfg.keep_recent_messages)
agent.sessions.save(session)
if not resp or not resp.content: if not resp or not resp.content:
return return
-1
View File
@@ -329,7 +329,6 @@ class HeartbeatConfig(Base):
enabled: bool = True enabled: bool = True
interval_s: int = 30 * 60 # 30 minutes interval_s: int = 30 * 60 # 30 minutes
keep_recent_messages: int = 8
class ApiConfig(Base): class ApiConfig(Base):
@@ -34,6 +34,7 @@ from nanobot.providers.base import (
) )
from nanobot.providers.openai_responses import ( from nanobot.providers.openai_responses import (
ResponsesStreamCapture, ResponsesStreamCapture,
build_responses_compaction_state,
build_responses_state, build_responses_state,
consume_sdk_stream, consume_sdk_stream,
convert_tools, convert_tools,
@@ -410,6 +411,16 @@ class AzureOpenAIProvider(LLMProvider):
output_items=capture.output_items, output_items=capture.output_items,
usage=usage, usage=usage,
) )
result.provider_compaction_state = build_responses_compaction_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
output_items=capture.output_items,
)
result.provider_compaction_applied = (
result.provider_compaction_state is not None
)
if result.provider_compaction_applied:
result.provider_compaction_scope = "current_request"
return result return result
except Exception as e: except Exception as e:
return self._handle_error(e) return self._handle_error(e)
+17
View File
@@ -31,6 +31,7 @@ RETRY_AFTER_BUFFER = 1
RetryEventCallback = Callable[[str], Awaitable[None]] RetryEventCallback = Callable[[str], Awaitable[None]]
LLMCallObserver = Callable[["LLMCallRecord"], None] LLMCallObserver = Callable[["LLMCallRecord"], None]
ProviderCompactionScope = Literal["prior_context", "current_request"]
def resolve_stream_idle_timeout_s( def resolve_stream_idle_timeout_s(
@@ -563,6 +564,22 @@ class LLMResponse:
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc. reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking
provider_state: ProviderConversationState | None = field(default=None, repr=False) provider_state: ProviderConversationState | None = field(default=None, repr=False)
# True only when this response installed a new provider-native compaction
# boundary. Replaying an older compaction item does not set this flag.
provider_compaction_applied: bool = field(default=False, repr=False)
# State immediately after native compaction, before the normal response
# continues. An archive prompt can resume this state without replaying H.
provider_compaction_state: ProviderConversationState | None = field(
default=None,
repr=False,
)
# Which model input the native compaction state replaces. Providers that
# compact before attaching the current request delta report
# ``prior_context``; in-request compaction reports ``current_request``.
provider_compaction_scope: ProviderCompactionScope | None = field(
default=None,
repr=False,
)
# Routing wrappers may preserve or discard an incoming provider-owned # Routing wrappers may preserve or discard an incoming provider-owned
# continuation independently of the final fallback error's retry policy. # continuation independently of the final fallback error's retry policy.
preserve_provider_state_on_error: bool | None = field(default=None, repr=False) preserve_provider_state_on_error: bool | None = field(default=None, repr=False)
+6
View File
@@ -101,6 +101,12 @@ class ProviderConversationStateController:
) )
return context_tokens + max(0, delta_tokens) return context_tokens + max(0, delta_tokens)
def replace_transcript(self, messages: list[dict[str, Any]]) -> None:
"""Discard append-only provider state after a transcript rewrite."""
self._state = None
self._boundary = len(messages)
self._request_messages = []
def prepare_request( def prepare_request(
self, self,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
+22 -3
View File
@@ -31,6 +31,7 @@ from nanobot.providers.oauth_model_catalog import (
) )
from nanobot.providers.openai_responses import ( from nanobot.providers.openai_responses import (
ResponsesStreamCapture, ResponsesStreamCapture,
build_responses_compaction_state,
build_responses_state, build_responses_state,
consume_sse_with_reasoning, consume_sse_with_reasoning,
convert_tools, convert_tools,
@@ -137,6 +138,8 @@ class OpenAICodexProvider(LLMProvider):
body.update(self._extra_body) body.update(self._extra_body)
stage = "oauth_token" stage = "oauth_token"
native_compaction_applied = False
native_compaction_state: ProviderConversationState | None = None
try: try:
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy) token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
headers = _build_headers(cast(str, token.account_id), token.access) headers = _build_headers(cast(str, token.account_id), token.access)
@@ -187,9 +190,11 @@ class OpenAICodexProvider(LLMProvider):
and responses_state_context_tokens(sanitized_state) >= compact_threshold and responses_state_context_tokens(sanitized_state) >= compact_threshold
): ):
stage = "codex_compaction" stage = "codex_compaction"
history_items = responses_state_items(sanitized_state) or []
delta_items = input_items[len(history_items):]
compact_body = { compact_body = {
**body, **body,
"input": [*input_items, {"type": "compaction_trigger"}], "input": [*history_items, {"type": "compaction_trigger"}],
} }
try: try:
compact_result = await _send(compact_body, emit_deltas=False) compact_result = await _send(compact_body, emit_deltas=False)
@@ -205,9 +210,16 @@ class OpenAICodexProvider(LLMProvider):
}: }:
raise RuntimeError("Codex compaction returned no compaction item") raise RuntimeError("Codex compaction returned no compaction item")
body["input"] = [ body["input"] = [
*_retained_compaction_messages(input_items), *_retained_compaction_messages(history_items),
*compact_items, *compact_items,
*delta_items,
] ]
native_compaction_state = build_responses_compaction_state(
provider=self._responses_state_provider(),
model=_strip_model_prefix(model),
output_items=compact_items,
)
native_compaction_applied = True
except Exception as compact_error: except Exception as compact_error:
if is_compaction_compatibility_error(compact_error): if is_compaction_compatibility_error(compact_error):
self._native_compaction_available = False self._native_compaction_available = False
@@ -220,7 +232,14 @@ class OpenAICodexProvider(LLMProvider):
) )
stage = "codex_request" stage = "codex_request"
return await _send(body, emit_deltas=True) result = await _send(body, emit_deltas=True)
result.provider_compaction_applied = (
result.provider_compaction_applied or native_compaction_applied
)
if native_compaction_state is not None:
result.provider_compaction_state = native_compaction_state
result.provider_compaction_scope = "prior_context"
return result
except Exception as e: except Exception as e:
response = _codex_error_response(e) response = _codex_error_response(e)
exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__ exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__
@@ -36,6 +36,7 @@ from nanobot.providers.base import (
) )
from nanobot.providers.openai_responses import ( from nanobot.providers.openai_responses import (
ResponsesStreamCapture, ResponsesStreamCapture,
build_responses_compaction_state,
build_responses_state, build_responses_state,
consume_sdk_stream, consume_sdk_stream,
convert_tools, convert_tools,
@@ -2049,6 +2050,18 @@ class OpenAICompatProvider(LLMProvider):
output_items=capture.output_items, output_items=capture.output_items,
usage=usage, usage=usage,
) )
result.provider_compaction_state = (
build_responses_compaction_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
output_items=capture.output_items,
)
)
result.provider_compaction_applied = (
result.provider_compaction_state is not None
)
if result.provider_compaction_applied:
result.provider_compaction_scope = "current_request"
return result return result
except Exception as responses_error: except Exception as responses_error:
if self._spec and self._spec.name == "github_copilot": if self._spec and self._spec.name == "github_copilot":
@@ -18,6 +18,7 @@ from nanobot.providers.openai_responses.parsing import (
parse_response_output, parse_response_output,
) )
from nanobot.providers.openai_responses.state import ( from nanobot.providers.openai_responses.state import (
build_responses_compaction_state,
build_responses_state, build_responses_state,
is_compaction_compatibility_error, is_compaction_compatibility_error,
prepare_responses_input, prepare_responses_input,
@@ -40,6 +41,7 @@ __all__ = [
"is_replayable_finish_reason", "is_replayable_finish_reason",
"map_finish_reason", "map_finish_reason",
"parse_response_output", "parse_response_output",
"build_responses_compaction_state",
"build_responses_state", "build_responses_state",
"is_compaction_compatibility_error", "is_compaction_compatibility_error",
"prepare_responses_input", "prepare_responses_input",
+12 -1
View File
@@ -11,7 +11,10 @@ import httpx
from loguru import logger from loguru import logger
from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest, parse_tool_arguments from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest, parse_tool_arguments
from nanobot.providers.openai_responses.state import build_responses_state from nanobot.providers.openai_responses.state import (
build_responses_compaction_state,
build_responses_state,
)
FINISH_REASON_MAP = { FINISH_REASON_MAP = {
"completed": "stop", "completed": "stop",
@@ -655,6 +658,14 @@ def parse_response_output(
output_items=output, output_items=output,
usage=usage, usage=usage,
) )
result.provider_compaction_state = build_responses_compaction_state(
provider=state_provider,
model=state_model,
output_items=output,
)
result.provider_compaction_applied = result.provider_compaction_state is not None
if result.provider_compaction_applied:
result.provider_compaction_scope = "current_request"
return result return result
@@ -108,6 +108,28 @@ def build_responses_state(
) )
def build_responses_compaction_state(
*,
provider: str,
model: str,
output_items: list[dict[str, Any]],
) -> ProviderConversationState | None:
"""Return the state at the latest native compaction output boundary."""
latest = None
for index, item in enumerate(output_items):
if item.get("type") in _COMPACTION_ITEM_TYPES:
latest = index
if latest is None:
return None
return ProviderConversationState(
kind=RESPONSES_STATE_KIND,
provider=provider,
model=model,
version=RESPONSES_STATE_VERSION,
payload={_ITEMS_KEY: [deepcopy(output_items[latest])]},
)
def responses_state_items( def responses_state_items(
state: ProviderConversationState, state: ProviderConversationState,
) -> list[dict[str, Any]] | None: ) -> list[dict[str, Any]] | None:
+3 -3
View File
@@ -209,11 +209,11 @@ class RuntimeClient:
return self._loop.runtime_events.subscribe(handler, SessionTurnPersisted) return self._loop.runtime_events.subscribe(handler, SessionTurnPersisted)
async def compact_session(self, session_key: str) -> SessionSnapshot: async def compact_session(self, session_key: str) -> SessionSnapshot:
"""Run token consolidation for one session.""" """Archive one session through the shared idle-compaction path."""
session = self._loop.sessions.get_or_create(session_key) session = self._loop.sessions.get_or_create(session_key)
runtime = self._loop.runtime_for_session(session) runtime = self._loop.runtime_for_session(session)
await self._loop.consolidator.maybe_consolidate_by_tokens( await self._loop.consolidator.compact_idle_session(
session, session_key,
runtime=runtime, runtime=runtime,
) )
return snapshot_from_session(self._loop.sessions.get_or_create(session_key)) return snapshot_from_session(self._loop.sessions.get_or_create(session_key))
+14 -121
View File
@@ -27,7 +27,9 @@ from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META, RUNTIME_CONTEXT_HISTORY_META,
public_history_message, public_history_message,
) )
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.model_selection import SESSION_MODEL_PRESET_METADATA_KEY from nanobot.session.model_selection import SESSION_MODEL_PRESET_METADATA_KEY
from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT
from nanobot.utils.helpers import ( from nanobot.utils.helpers import (
content_with_media_breadcrumbs, content_with_media_breadcrumbs,
ensure_dir, ensure_dir,
@@ -262,12 +264,6 @@ def _metadata_title(metadata: object) -> str:
return strip_think(title) return strip_think(title)
@dataclass
class RetentionResult:
dropped: list[dict[str, Any]]
already_consolidated_count: int
@dataclass(frozen=True) @dataclass(frozen=True)
class SessionPolicy: class SessionPolicy:
"""Runtime rules that do not belong in durable session data.""" """Runtime rules that do not belong in durable session data."""
@@ -286,9 +282,7 @@ class Session:
created_at: datetime = field(default_factory=datetime.now) created_at: datetime = field(default_factory=datetime.now)
updated_at: datetime = field(default_factory=datetime.now) updated_at: datetime = field(default_factory=datetime.now)
metadata: dict[str, Any] = field(default_factory=dict) metadata: dict[str, Any] = field(default_factory=dict)
# Legacy storage name for the Memory ingestion watermark. New code should # Keep the legacy storage name while persisted sessions and SDK callers migrate.
# use ``last_archived`` so this progress is not confused with model-context
# compaction. Keep the field while persisted sessions and SDK callers migrate.
last_consolidated: int = 0 last_consolidated: int = 0
provider_state: ProviderConversationState | None = field(default=None, repr=False) provider_state: ProviderConversationState | None = field(default=None, repr=False)
policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False) policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False)
@@ -309,7 +303,7 @@ class Session:
@property @property
def last_archived(self) -> int: def last_archived(self) -> int:
"""Number of transcript messages already written to the Memory journal.""" """End of the latest committed Memory checkpoint."""
return self.last_consolidated return self.last_consolidated
@last_archived.setter @last_archived.setter
@@ -337,14 +331,17 @@ class Session:
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""Return recent replayable messages for LLM input. """Return recent replayable messages for LLM input.
A positive ``max_messages`` applies an explicit caller-owned count A committed in-turn checkpoint replaces its old prefix with the stored
limit. The normal model path relies on ``max_tokens`` instead. summary and resumes replay at a hidden continuation marker. A positive
``max_messages`` applies an additional caller-owned count limit.
""" """
replay_start = self.last_archived replay_start = self.last_archived
if replay_start: resumes_from_checkpoint = (
# ``last_archived`` is archive progress, not a replay boundary. replay_start < len(self.messages)
# Keep a small raw suffix for continuity, extending back to the user and is_hidden_history_message(self.messages[replay_start])
# that started an assistant/tool sequence when necessary. and self.messages[replay_start].get("content") == SUMMARY_CONTINUATION_TEXT
)
if replay_start and not resumes_from_checkpoint:
recent_start = recent_message_start_index( recent_start = recent_message_start_index(
self.messages, self.messages,
MIN_COMPACTED_REPLAY_MESSAGES, MIN_COMPACTED_REPLAY_MESSAGES,
@@ -485,110 +482,6 @@ class Session:
self.updated_at = datetime.now() self.updated_at = datetime.now()
self.metadata.pop("_last_summary", None) self.metadata.pop("_last_summary", None)
def retain_recent_legal_suffix(
self,
max_messages: int,
*,
extend_to_user: bool = False,
) -> RetentionResult:
"""Keep a legal recent suffix, optionally extending it back to a user turn.
Returns a RetentionResult with dropped messages and how many of those
were in the already-consolidated prefix. This method mutates
self.messages and self.last_archived in place.
"""
if max_messages <= 0:
dropped = list(self.messages)
lc = self.last_archived
self.clear()
return RetentionResult(
dropped=dropped,
already_consolidated_count=min(lc, len(dropped)),
)
if len(self.messages) <= max_messages:
return RetentionResult(
dropped=[],
already_consolidated_count=0,
)
original = list(self.messages)
before_lc = self.last_archived
start_idx = max(0, len(self.messages) - max_messages)
if extend_to_user:
recovered_user = next(
(i for i in range(start_idx, -1, -1) if self.messages[i].get("role") == "user"),
None,
)
if recovered_user is not None:
start_idx = recovered_user
if start_idx > 0 and self.messages[start_idx - 1].get("_channel_delivery"):
start_idx -= 1
retained = self.messages[start_idx:]
# Prefer starting at a user turn (or its preceding _channel_delivery) when one exists within the retained window.
first_user = next((i for i, m in enumerate(retained) if m.get("role") == "user"), None)
if first_user is not None:
if first_user > 0 and retained[first_user - 1].get("_channel_delivery"):
retained = retained[first_user - 1:]
else:
retained = retained[first_user:]
elif not extend_to_user:
# If the hard-capped tail is assistant/tool-only, anchor to the
# latest user in the full session and take a capped forward window.
latest_user = next(
(i for i in range(len(self.messages) - 1, -1, -1)
if self.messages[i].get("role") == "user"),
None,
)
if latest_user is not None:
retained = self.messages[latest_user: latest_user + max_messages]
# Mirror get_history(): avoid persisting orphan tool results at the front.
start = find_legal_message_start(retained)
if start:
retained = retained[start:]
# Hard-cap guarantee unless the caller requested user-turn extension.
if not extend_to_user and len(retained) > max_messages:
retained = retained[-max_messages:]
start = find_legal_message_start(retained)
if start:
retained = retained[start:]
# Compute actually-dropped messages using identity comparison so that
# even when retained is a non-contiguous slice of original (the else
# branch above), we never duplicate or lose messages.
retained_ids = set(id(m) for m in retained)
dropped = [m for m in original if id(m) not in retained_ids]
# Count how many dropped messages were in the already-consolidated
# prefix of the original list. This cannot be a simple min() because
# dropped may include messages from *after* the consolidated prefix
# (e.g. in the else branch).
already_consolidated = sum(
1 for i, m in enumerate(original)
if i < before_lc and id(m) not in retained_ids
)
# New last_archived = count of retained messages that were inside
# the old consolidated prefix.
new_lc = sum(
1 for i, m in enumerate(original)
if i < before_lc and id(m) in retained_ids
)
self.messages = retained
self.last_archived = new_lc
if dropped:
self.provider_state = None
self.updated_at = datetime.now()
return RetentionResult(
dropped=dropped,
already_consolidated_count=already_consolidated,
)
class SessionPayload(TypedDict): class SessionPayload(TypedDict):
key: str key: str
created_at: str | None created_at: str | None
@@ -2010,7 +1903,7 @@ class SessionManager:
user_index = 0 user_index = 0
found_target = False found_target = False
for message in source.messages: for message in source.messages:
if message.get("role") == "user": if message.get("role") == "user" and not is_hidden_history_message(message):
if user_index == before_user_index: if user_index == before_user_index:
found_target = True found_target = True
break break
+12
View File
@@ -3,15 +3,27 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Mapping from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from typing import TypedDict, cast from typing import TypedDict, cast
SUMMARY_CONTINUATION_TEXT = (
"Continue the active task from the working-memory checkpoint above."
)
class SessionSummary(TypedDict): class SessionSummary(TypedDict):
text: str text: str
last_active: str last_active: str
@dataclass(frozen=True, slots=True)
class SessionSummaryCheckpoint:
"""A replacement summary and the raw transcript boundary it covers."""
summary: str
transcript_boundary: int
def session_summary_from_metadata( def session_summary_from_metadata(
metadata: Mapping[str, object] | None, metadata: Mapping[str, object] | None,
*, *,
-1
View File
@@ -114,7 +114,6 @@ def system_settings_payload(
"heartbeat": { "heartbeat": {
"enabled": config.gateway.heartbeat.enabled, "enabled": config.gateway.heartbeat.enabled,
"interval_s": config.gateway.heartbeat.interval_s, "interval_s": config.gateway.heartbeat.interval_s,
"keep_recent_messages": config.gateway.heartbeat.keep_recent_messages,
}, },
"dream": { "dream": {
"schedule": defaults.dream.describe_schedule(), "schedule": defaults.dream.describe_schedule(),
-31
View File
@@ -8,7 +8,6 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.runner import AgentRunResult
from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
@@ -229,36 +228,6 @@ class TestAgentLoopTTLParam:
loop = _make_loop(tmp_path, session_ttl_minutes=0) loop = _make_loop(tmp_path, session_ttl_minutes=0)
assert loop.auto_compact._ttl == 0 assert loop.auto_compact._ttl == 0
@pytest.mark.asyncio
async def test_process_message_reads_history_with_token_budget(self, tmp_path):
"""_process_message should pass an auto-derived token budget to get_history."""
loop = _make_loop(tmp_path)
session = loop.sessions.get_or_create("cli:direct")
session.get_history = MagicMock(return_value=[])
loop.context.build_messages = MagicMock(return_value=[])
loop._run_agent_loop = AsyncMock(
return_value=AgentRunResult(
final_content="ok",
messages=[],
stop_reason="stop",
)
)
loop._save_turn = MagicMock()
msg = InboundMessage(
channel="cli",
sender_id="u1",
chat_id="direct",
content="hello",
)
await loop._process_message(msg)
session.get_history.assert_called_once()
kwargs = session.get_history.call_args.kwargs
assert isinstance(kwargs.get("max_tokens"), int)
assert kwargs["max_tokens"] > 0
assert set(kwargs) == {"max_tokens", "extend_to_user"}
class TestAutoCompact: class TestAutoCompact:
"""Test the _archive method.""" """Test the _archive method."""
+85 -263
View File
@@ -55,9 +55,7 @@ def runtime(mock_provider):
def consolidator(store): def consolidator(store):
sessions = MagicMock() sessions = MagicMock()
sessions.save = MagicMock() sessions.save = MagicMock()
# When maybe_consolidate_by_tokens refreshes the session reference via # Store sessions by key so refreshes observe the same test object.
# get_or_create(session.key), it should get back the same object the test
# passed in. Store sessions by key so the lookup is transparent.
_session_cache: dict[str, MagicMock] = {} _session_cache: dict[str, MagicMock] = {}
sessions.get_or_create = MagicMock(side_effect=lambda key: _session_cache.get(key, MagicMock())) sessions.get_or_create = MagicMock(side_effect=lambda key: _session_cache.get(key, MagicMock()))
sessions._session_cache = _session_cache sessions._session_cache = _session_cache
@@ -93,11 +91,17 @@ def _provider_state() -> ProviderConversationState:
def _build_test_messages(**kwargs): def _build_test_messages(**kwargs):
return [ system = "system prompt"
{"role": "system", "content": "system prompt"}, session_summary = kwargs.get("session_summary")
if session_summary:
system += f"\n\n[Archived Context Summary]\n{session_summary['text']}"
messages = [
{"role": "system", "content": system},
*kwargs["history"], *kwargs["history"],
{"role": "user", "content": kwargs["current_message"]},
] ]
if kwargs["current_message"] is not None:
messages.append({"role": "user", "content": kwargs["current_message"]})
return messages
async def _archive( async def _archive(
@@ -112,15 +116,85 @@ async def _archive(
messages, messages,
runtime=runtime, runtime=runtime,
session_key=session_key, session_key=session_key,
request_messages=_build_test_messages( history=[
history=messages, {"role": "system", "content": "system prompt"},
current_message="consolidate", *messages,
), ],
request_tools=[], request_tools=[],
previous_summary=previous_summary, previous_summary=previous_summary,
) )
class TestTurnTranscriptSummary:
async def test_uses_exact_accepted_prefix_and_existing_archiver(
self,
consolidator,
mock_provider,
runtime,
):
accepted = [
{"role": "system", "content": "stable system"},
{"role": "user", "content": "accepted history"},
]
tools = [{"type": "function", "function": {"name": "inspect"}}]
mock_provider.chat_with_retry.return_value = LLMResponse(
content="replacement checkpoint",
)
summary = await consolidator.summarize_transcript(
accepted,
"previous checkpoint",
runtime=runtime,
session_key="test:turn",
tools=tools,
)
assert summary == "replacement checkpoint"
call = mock_provider.chat_with_retry.await_args.kwargs
assert call["messages"][:-1] == accepted
assert call["messages"][-1]["role"] == "user"
assert "SNIP" in call["messages"][-1]["content"]
assert call["tools"] == tools
async def test_native_compaction_appends_only_archive_prompt(
self,
consolidator,
mock_provider,
runtime,
):
accepted = [
{"role": "system", "content": "stable system"},
{"role": "user", "content": "raw history must not be replayed"},
]
state = _provider_state()
mock_provider.can_resume_conversation_state.return_value = True
mock_provider.chat_with_retry.return_value = LLMResponse(
content="replacement checkpoint",
)
summary = await consolidator.summarize_provider_compaction(
state,
accepted,
"previous checkpoint",
runtime=runtime,
session_key="test:turn",
tools=[{"type": "function", "function": {"name": "inspect"}}],
)
assert summary == "replacement checkpoint"
call = mock_provider.chat_with_retry.await_args.kwargs
assert call["messages"][0] == accepted[0]
assert call["messages"][-1]["content"] == _ARCHIVE_PROMPT
assert accepted[1] not in call["messages"]
assert call["tools"] == []
provider_context = call["provider_context"]
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.payload == state.payload
assert provider_context.conversation_state.pending_messages == [
call["messages"][-1],
]
class TestConsolidatorSummarize: class TestConsolidatorSummarize:
def test_format_messages_keeps_media_only_user_turn(self): def test_format_messages_keeps_media_only_user_turn(self):
path = "/home/user/.nanobot/media/websocket/clip.mp4" path = "/home/user/.nanobot/media/websocket/clip.mp4"
@@ -379,32 +453,7 @@ class TestConsolidatorArchiveErrorHandling:
consolidator.store.raw_archive.assert_not_called() consolidator.store.raw_archive.assert_not_called()
class TestConsolidatorTokenBudget: class TestConsolidatorPromptEstimate:
async def test_prompt_below_threshold_does_not_consolidate(
self, consolidator, runtime
):
"""No consolidation when tokens are within budget."""
session = MagicMock()
session.last_archived = 0
session.messages = [{"role": "user", "content": "hi"}]
session.key = "test:key"
consolidator.sessions._session_cache[session.key] = session
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(100, "tiktoken"))
consolidator.archive_session = AsyncMock(return_value=True)
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
consolidator.archive_session.assert_not_called()
async def test_token_estimation_failure_propagates(self, consolidator, runtime):
session = Session(key="test:estimate-failure")
session.add_message("user", "hello")
consolidator.sessions._session_cache[session.key] = session
consolidator.estimate_session_prompt_tokens = MagicMock(
side_effect=RuntimeError("counter failed")
)
with pytest.raises(RuntimeError, match="counter failed"):
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
async def test_estimate_uses_full_unarchived_tail(self, consolidator, runtime): async def test_estimate_uses_full_unarchived_tail(self, consolidator, runtime):
"""Consolidation pressure must account for the full unarchived tail.""" """Consolidation pressure must account for the full unarchived tail."""
session = Session(key="test:full-tail") session = Session(key="test:full-tail")
@@ -443,129 +492,6 @@ class TestConsolidatorTokenBudget:
assert len(captured["history"]) == 8 assert len(captured["history"]) == 8
assert captured["history"][0]["content"] == "msg-2" assert captured["history"][0]["content"] == "msg-2"
async def test_token_overflow_appends_prompt_to_replay_prefix(
self,
consolidator,
mock_provider,
runtime,
):
consolidator._SAFETY_BUFFER = 0
session = Session(key="test:token-prefix")
session.provider_state = _provider_state()
session.messages = [
{
"role": "user" if i in {0, 50, 61} else "assistant",
"content": f"m{i}",
}
for i in range(70)
]
consolidator.sessions._session_cache[session.key] = session
consolidator.estimate_session_prompt_tokens = MagicMock(
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
)
consolidator.pick_consolidation_boundary = MagicMock(return_value=50)
consolidator.archiver._build_messages = MagicMock(side_effect=_build_test_messages)
mock_provider.estimate_prompt_tokens.return_value = (100, "test-counter")
mock_provider.chat_with_retry.return_value = LLMResponse(
content="Token overflow summary.",
finish_reason="stop",
)
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
request = mock_provider.chat_with_retry.await_args.kwargs
assert [message["content"] for message in request["messages"][1:-1]] == [
f"m{i}" for i in range(50)
]
assert request["messages"][-1]["content"] == _ARCHIVE_PROMPT
assert request["tools"] == []
assert "tool_choice" not in request
assert session.last_archived == 50
assert session.provider_state == _provider_state()
async def test_raw_archive_fallback_advances_archive_watermark(
self, consolidator, runtime
):
"""When archive() falls back to raw-archive (LLM failed), the cursor
must still advance. Otherwise the same chunk gets raw-archived again
on every subsequent maybe_consolidate_by_tokens() call, spamming
duplicate [RAW] entries into history.jsonl."""
consolidator._SAFETY_BUFFER = 0
session = Session(key="test:key")
session.provider_state = _provider_state()
session.messages = [
{"role": "user" if i in {0, 50} else "assistant", "content": f"m{i}"}
for i in range(70)
]
consolidator.sessions._session_cache[session.key] = session
consolidator.estimate_session_prompt_tokens = MagicMock(
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
)
consolidator.archive_session = AsyncMock(return_value="[RAW] checkpoint")
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
consolidator.archive_session.assert_awaited_once()
# The chunk is considered "materialized" (as a raw-archive breadcrumb),
# so the archive watermark must have moved past it without touching
# the provider-owned continuation state.
assert session.last_archived == 50
assert session.provider_state == _provider_state()
async def test_raw_archive_fallback_breaks_round_loop(
self, consolidator, runtime
):
"""A degraded LLM should not trigger more archive() calls within the
same maybe_consolidate_by_tokens invocation bail after one fallback."""
consolidator._SAFETY_BUFFER = 0
session = MagicMock()
session.last_archived = 0
session.key = "test:key"
session.messages = [
{"role": "user" if i in {0, 20, 40, 60} else "assistant", "content": f"m{i}"}
for i in range(70)
]
session.metadata = {}
consolidator.sessions._session_cache[session.key] = session
# Keep estimates high so the loop would otherwise run multiple rounds.
consolidator.estimate_session_prompt_tokens = MagicMock(
return_value=(1200, "tiktoken")
)
consolidator.archive_session = AsyncMock(return_value="[RAW] checkpoint")
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
# The fixed policy archives at most one prefix per call.
assert consolidator.archive_session.await_count == 1
async def test_boundary_respected_when_no_intermediate_user_turn(
self, consolidator, runtime
):
"""When boundary points past a long tool chain, the full chunk is archived."""
consolidator._SAFETY_BUFFER = 0
session = MagicMock()
session.last_archived = 0
session.key = "test:key"
session.messages = [
{
"role": "user" if i in {0, 61} else "assistant",
"content": f"m{i}",
}
for i in range(70)
]
consolidator.sessions._session_cache[session.key] = session
consolidator.estimate_session_prompt_tokens = MagicMock(
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
)
consolidator.archive_session = AsyncMock(return_value=True)
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
consolidator.archive_session.assert_awaited_once()
# The fixed recent tail expands backward to the user at idx=61.
assert session.last_archived == 61
class TestCompactIdleSession: class TestCompactIdleSession:
"""Idle compaction tests.""" """Idle compaction tests."""
@@ -1347,110 +1273,6 @@ class TestCompactIdleSession:
assert not lock.locked() assert not lock.locked()
class TestConsolidatorSessionRefresh:
"""Background consolidation must detect stale session references."""
@pytest.mark.asyncio
async def test_reloads_before_empty_session_guard(self, tmp_path):
"""A stale empty reference must not skip a non-empty cached session."""
from nanobot.agent.memory import Consolidator, MemoryStore
from nanobot.session.manager import Session, SessionManager
store = MemoryStore(tmp_path)
provider = MagicMock()
provider.chat_with_retry = AsyncMock(
return_value=MagicMock(content="summary", finish_reason="stop")
)
provider.generation = GenerationSettings(max_tokens=4096)
provider.estimate_prompt_tokens = MagicMock(return_value=(10, "test"))
runtime = LLMRuntime.capture(
provider,
"test-model",
context_window_tokens=128_000,
)
sessions = SessionManager(tmp_path)
consolidator = Consolidator(
store=store,
sessions=sessions,
build_messages=MagicMock(return_value=[]),
get_tool_definitions=MagicMock(return_value=[]),
)
fresh = sessions.get_or_create("cli:test")
fresh.add_message("user", "fresh message")
sessions.save(fresh)
stale_empty = Session(key="cli:test")
seen: dict[str, Session] = {}
def estimate(session: Session, *, runtime):
seen["session"] = session
return 10, "test"
consolidator.estimate_session_prompt_tokens = MagicMock(side_effect=estimate)
await consolidator.maybe_consolidate_by_tokens(
stale_empty,
runtime=runtime,
)
assert seen["session"] is fresh
@pytest.mark.asyncio
async def test_reloads_stale_session_after_compact(self, tmp_path):
"""After compact_idle_session replaces the session, a concurrent
maybe_consolidate_by_tokens with the old reference should use the
fresh session from cache instead of overwriting."""
from nanobot.agent.memory import Consolidator, MemoryStore
from nanobot.session.manager import SessionManager
store = MemoryStore(tmp_path)
provider = MagicMock()
provider.chat_with_retry = AsyncMock(
return_value=MagicMock(content="summary", finish_reason="stop")
)
provider.generation = GenerationSettings(max_tokens=4096)
provider.estimate_prompt_tokens = MagicMock(return_value=(10, "test"))
runtime = LLMRuntime.capture(
provider,
"test-model",
context_window_tokens=128_000,
)
sessions = SessionManager(tmp_path)
consolidator = Consolidator(
store=store,
sessions=sessions,
build_messages=MagicMock(return_value=[]),
get_tool_definitions=MagicMock(return_value=[]),
)
# Populate session with many messages
session = sessions.get_or_create("cli:test")
for i in range(20):
session.add_message("user", f"u{i}")
session.add_message("assistant", f"a{i}")
sessions.save(session)
# Simulate: background consolidation captures old reference
old_ref = session
await consolidator.compact_idle_session(
"cli:test",
runtime=runtime,
max_suffix=8,
)
await consolidator.maybe_consolidate_by_tokens(
old_ref,
runtime=runtime,
)
session_after = sessions.get_or_create("cli:test")
assert len(session_after.messages) == 40
assert session_after.last_archived == 40
assert len(session_after.get_history(max_messages=40)) == 8
class TestRawArchiveTruncation: class TestRawArchiveTruncation:
"""raw_archive() must cap entry size to avoid bloating history.jsonl.""" """raw_archive() must cap entry size to avoid bloating history.jsonl."""
+2 -18
View File
@@ -404,10 +404,9 @@ class TestEphemeralDirect:
with ( with (
patch("nanobot.agent.loop.SessionManager"), patch("nanobot.agent.loop.SessionManager"),
patch("nanobot.agent.loop.SubagentManager") as mock_sub, patch("nanobot.agent.loop.SubagentManager") as mock_sub,
patch("nanobot.agent.loop.Consolidator") as mock_consolidator_cls, patch("nanobot.agent.loop.Consolidator"),
): ):
mock_sub.return_value.cancel_by_session = AsyncMock(return_value=0) mock_sub.return_value.cancel_by_session = AsyncMock(return_value=0)
mock_consolidator_cls.return_value.maybe_consolidate_by_tokens = AsyncMock()
loop = AgentLoop( loop = AgentLoop(
bus=bus, bus=bus,
provider=provider, provider=provider,
@@ -493,20 +492,6 @@ class TestEphemeralDirect:
assert captured.get("ephemeral") is False assert captured.get("ephemeral") is False
async def test_ephemeral_skips_consolidator(self, tmp_path, _make_loop):
"""When ephemeral=True, consolidator.maybe_consolidate_by_tokens is not called."""
from unittest.mock import patch
loop, store = _make_loop
with patch.object(
loop.consolidator, "maybe_consolidate_by_tokens",
) as mock_consolidate:
await loop.process_direct(
"test", session_key="dream:consolidate-test", ephemeral=True,
)
mock_consolidate.assert_not_called()
async def test_ephemeral_response_reports_stop_reason(self, tmp_path, _make_loop): async def test_ephemeral_response_reports_stop_reason(self, tmp_path, _make_loop):
loop, store = _make_loop loop, store = _make_loop
loop.provider.chat_with_retry.return_value = LLMResponse( loop.provider.chat_with_retry.return_value = LLMResponse(
@@ -701,10 +686,9 @@ class TestEphemeralHooks:
with ( with (
patch("nanobot.agent.loop.SessionManager"), patch("nanobot.agent.loop.SessionManager"),
patch("nanobot.agent.loop.SubagentManager") as mock_sub, patch("nanobot.agent.loop.SubagentManager") as mock_sub,
patch("nanobot.agent.loop.Consolidator") as mock_consolidator_cls, patch("nanobot.agent.loop.Consolidator"),
): ):
mock_sub.return_value.cancel_by_session = AsyncMock(return_value=0) mock_sub.return_value.cancel_by_session = AsyncMock(return_value=0)
mock_consolidator_cls.return_value.maybe_consolidate_by_tokens = AsyncMock()
loop = AgentLoop( loop = AgentLoop(
bus=bus, bus=bus,
provider=provider, provider=provider,
+7 -9
View File
@@ -12,6 +12,7 @@ from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse from nanobot.providers.base import LLMResponse
from nanobot.session.manager import Session from nanobot.session.manager import Session
from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT
def _make_loop(tmp_path: Path, context_window_tokens: int = 200_000) -> AgentLoop: def _make_loop(tmp_path: Path, context_window_tokens: int = 200_000) -> AgentLoop:
@@ -66,13 +67,12 @@ def test_explicit_message_limit_still_starts_at_user_turn() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_message_replays_with_token_budget_only(tmp_path: Path) -> None: async def test_process_message_hands_complete_replay_to_runner(tmp_path: Path) -> None:
loop = _make_loop(tmp_path, context_window_tokens=32_768) loop = _make_loop(tmp_path, context_window_tokens=32_768)
loop.provider.chat_with_retry = AsyncMock( loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage=None) return_value=LLMResponse(content="ok", tool_calls=[], usage=None)
) )
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
with patch.object(session, "get_history", wraps=session.get_history) as get_history: with patch.object(session, "get_history", wraps=session.get_history) as get_history:
@@ -81,20 +81,16 @@ async def test_process_message_replays_with_token_budget_only(tmp_path: Path) ->
) )
assert result is not None assert result is not None
assert get_history.call_args.kwargs == { assert get_history.call_args.kwargs == {"extend_to_user": False}
"max_tokens": loop._replay_token_budget(loop.llm_runtime()),
"extend_to_user": False,
}
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_token_budget_keeps_current_user_as_replay_boundary(tmp_path: Path) -> None: async def test_runner_checkpoint_keeps_current_user_as_replay_boundary(tmp_path: Path) -> None:
loop = _make_loop(tmp_path, context_window_tokens=8_000) loop = _make_loop(tmp_path, context_window_tokens=8_000)
loop.provider.chat_with_retry = AsyncMock( loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage=None) return_value=LLMResponse(content="ok", tool_calls=[], usage=None)
) )
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
session.add_message("user", "old") session.add_message("user", "old")
@@ -117,4 +113,6 @@ async def test_token_budget_keeps_current_user_as_replay_boundary(tmp_path: Path
sent_messages = loop.provider.chat_with_retry.await_args.kwargs["messages"] sent_messages = loop.provider.chat_with_retry.await_args.kwargs["messages"]
sent_text = "\n".join(str(message.get("content")) for message in sent_messages) sent_text = "\n".join(str(message.get("content")) for message in sent_messages)
assert "new question" in sent_text assert "new question" in sent_text
assert "long older turn" not in sent_text assert [message["role"] for message in sent_messages] == ["system", "user", "user"]
assert sent_messages[1]["content"] == SUMMARY_CONTINUATION_TEXT
assert any(message.get("content") == "long older turn" for message in session.messages)
+90 -164
View File
@@ -4,7 +4,12 @@ import pytest
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse from nanobot.providers.base import (
GenerationSettings,
LLMResponse,
ProviderConversationState,
)
from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT
def _make_loop( def _make_loop(
@@ -14,7 +19,6 @@ def _make_loop(
context_window_tokens: int, context_window_tokens: int,
max_tokens: int = 0, max_tokens: int = 0,
) -> AgentLoop: ) -> AgentLoop:
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=max_tokens) provider.generation = GenerationSettings(max_tokens=max_tokens)
@@ -39,186 +43,108 @@ def _make_loop(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_prompt_below_threshold_does_not_consolidate(tmp_path) -> None: async def test_runner_pressure_commits_summary_and_current_delta(tmp_path) -> None:
loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=200) loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=2_000)
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign] loop.context_block_limit = 500
loop.provider.generation = GenerationSettings(max_tokens=100)
await loop.process_direct("hello", session_key="cli:test") loop.provider.can_resume_conversation_state.return_value = False
loop.consolidator.archive_session.assert_not_awaited()
@pytest.mark.asyncio
async def test_prompt_above_threshold_triggers_consolidation(tmp_path) -> None:
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
loop.consolidator.archive_session = AsyncMock(return_value=True) # 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")
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] loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
session.messages = [ session.messages = [
{"role": role, "content": f"{role[0]}{turn}"} {"role": role, "content": f"old-{role}-{turn}"}
for turn in range(10) for turn in range(6)
for role in ("user", "assistant") for role in ("user", "assistant")
] ]
loop.sessions.save(session) loop.sessions.save(session)
await loop.process_direct("hello", session_key="cli:test") def estimate(messages, _tools, _model):
contents = [str(message.get("content")) for message in messages]
if contents and "SNIP" in contents[-1]:
return 300, "test-counter"
if any(content.startswith("old-") for content in contents):
return 600, "test-counter"
return 100, "test-counter"
request_messages = loop.provider.chat_with_retry.await_args.kwargs["messages"] loop.provider.estimate_prompt_tokens.side_effect = estimate
system_prompt = request_messages[0]["content"] loop.provider.chat_with_retry = AsyncMock(side_effect=[
assert "FRESH_CHECKPOINT" in system_prompt LLMResponse(content="Current checkpoint.", tool_calls=[]),
assert all(message.get("content") != "u0" for message in request_messages) LLMResponse(content="done", tool_calls=[]),
assert loop.sessions.get_or_create("cli:test").last_archived == 12 ])
result = await loop.process_direct("continue the task", session_key="cli:test")
@pytest.mark.asyncio assert result.content == "done"
async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None: assert loop.provider.chat_with_retry.await_count == 2
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200) model_request = loop.provider.chat_with_retry.await_args_list[1].kwargs["messages"]
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign] assert "Current checkpoint." in model_request[0]["content"]
assert model_request[1]["content"] == SUMMARY_CONTINUATION_TEXT
session = loop.sessions.get_or_create("cli:test") assert model_request[2]["content"] == "continue the task"
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.consolidator.maybe_consolidate_by_tokens(
session,
runtime=loop.llm_runtime(),
)
archive_end = loop.consolidator.archive_session.await_args.kwargs["archive_end"]
archived_chunk = session.messages[:archive_end]
assert [message["content"] for message in archived_chunk] == [
"u0", "a0", "u1", "a1", "u2", "a2", "u3", "a3", "u4", "a4", "u5", "a5",
]
assert session.last_archived == 12
@pytest.mark.asyncio
async def test_consolidation_persists_summary_for_next_prepare_session(tmp_path) -> None:
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
loop.consolidator.archive_session = AsyncMock(return_value="User discussed project status.") # 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(5)
for role in ("user", "assistant")
]
loop.sessions.save(session)
def mock_estimate(_session, *, runtime):
return (500, "test")
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
await loop.consolidator.maybe_consolidate_by_tokens(
session,
runtime=loop.llm_runtime(),
)
reloaded = loop.sessions.get_or_create("cli:test") reloaded = loop.sessions.get_or_create("cli:test")
meta = reloaded.metadata.get("_last_summary") assert reloaded.messages[0]["content"] == "old-user-0"
assert meta is not None assert reloaded.metadata["_last_summary"]["text"] == "Current checkpoint."
assert meta["text"] == "User discussed project status." assert reloaded.messages[reloaded.last_archived]["content"] == (
SUMMARY_CONTINUATION_TEXT
reloaded, pending = loop.auto_compact.prepare_session(reloaded, "cli:test") )
assert pending is not None assert [message["content"] for message in reloaded.get_history()] == [
assert pending["text"] == "User discussed project status." SUMMARY_CONTINUATION_TEXT,
# _last_summary persists for restart survival. "continue the task",
assert "_last_summary" in reloaded.metadata "done",
]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> None: async def test_native_provider_compaction_commits_portable_terminal_checkpoint(
loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=200) tmp_path,
session = loop.sessions.get_or_create("cli:test") ) -> None:
loop.auto_compact.prepare_session = MagicMock( loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=2_000)
return_value=( session = loop.sessions.get_or_create("cli:native")
session,
{"text": "earlier context", "last_active": session.updated_at.isoformat()},
)
) # type: ignore[method-assign]
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None) # type: ignore[method-assign]
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
runtime = loop.llm_runtime()
await loop.process_direct("hello", session_key="cli:test", runtime=runtime)
loop.consolidator.maybe_consolidate_by_tokens.assert_any_await(
session,
runtime=runtime,
)
assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
assert all(
call.kwargs["runtime"] is runtime
for call in loop.consolidator.maybe_consolidate_by_tokens.call_args_list
)
@pytest.mark.asyncio
async def test_preflight_consolidation_before_llm_call(tmp_path) -> None:
"""Verify preflight consolidation runs before the LLM call in process_direct."""
order: list[str] = []
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
archived_session_keys: list[str | None] = []
async def track_consolidate(session, *, archive_end, runtime):
order.append("consolidate")
archived_session_keys.append(session.key)
return True
loop.consolidator.archive_session = track_consolidate # type: ignore[method-assign]
async def track_llm(*args, **kwargs):
order.append("llm")
return LLMResponse(content="ok", tool_calls=[])
loop.provider.chat_with_retry = track_llm
loop.provider.chat_stream_with_retry = track_llm
loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test")
session.messages = [ session.messages = [
{"role": role, "content": f"{role[0]}{turn}"} {"role": "user", "content": "accepted history"},
for turn in range(10) {"role": "assistant", "content": "accepted answer"},
for role in ("user", "assistant")
] ]
loop.sessions.save(session) loop.sessions.save(session)
call_count = [0] compacted_state = ProviderConversationState(
def mock_estimate(_session, *, runtime): kind="openai_responses",
call_count[0] += 1 provider="openai:test",
return (1000 if call_count[0] <= 1 else 80, "test") model="test-model",
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign] version=1,
payload={"items": [{"type": "compaction", "encrypted_content": "opaque"}]},
)
loop.provider.can_resume_conversation_state.return_value = True
loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="done",
provider_state=compacted_state,
provider_compaction_applied=True,
provider_compaction_state=compacted_state,
provider_compaction_scope="current_request",
))
loop.consolidator.summarize_provider_compaction = AsyncMock(
return_value="portable terminal checkpoint",
)
await loop.process_direct("hello", session_key="cli:test") result = await loop.process_direct("continue", session_key="cli:native")
assert "consolidate" in order assert result.content == "done"
assert "llm" in order summarize = loop.consolidator.summarize_provider_compaction
assert order.index("consolidate") < order.index("llm") summarize.assert_awaited_once()
assert archived_session_keys == ["cli:test"] assert summarize.await_args.args[0] == compacted_state
accepted = summarize.await_args.args[1]
accepted_contents = [message.get("content") for message in accepted]
assert "accepted history" in accepted_contents
assert "accepted answer" in accepted_contents
assert "continue" in accepted_contents
assert "done" not in accepted_contents
reloaded = loop.sessions.get_or_create("cli:native")
assert reloaded.provider_state is None
assert reloaded.metadata["_last_summary"]["text"] == (
"portable terminal checkpoint"
)
assert reloaded.messages[reloaded.last_archived]["content"] == (
SUMMARY_CONTINUATION_TEXT
)
assert [message["content"] for message in reloaded.get_history()] == [
SUMMARY_CONTINUATION_TEXT,
"done",
]
@@ -69,7 +69,6 @@ async def test_outbound_no_longer_carries_generated_media(
), ),
image_generation_provider_config=ProviderConfig(api_key="sk-or-test"), image_generation_provider_config=ProviderConfig(api_key="sk-or-test"),
) )
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
result = await loop._process_message( result = await loop._process_message(
InboundMessage( InboundMessage(
-15
View File
@@ -425,7 +425,6 @@ class TestToolEventProgress:
None, None,
), ),
) )
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -473,7 +472,6 @@ class TestToolEventProgress:
provider.chat_stream_with_retry = AsyncMock() provider.chat_stream_with_retry = AsyncMock()
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5")
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="whatsapp", channel="whatsapp",
@@ -512,7 +510,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -566,7 +563,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -611,7 +607,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -655,7 +650,6 @@ class TestToolEventProgress:
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.max_iterations = 1 loop.max_iterations = 1
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -747,7 +741,6 @@ class TestToolEventProgress:
) )
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -815,9 +808,6 @@ class TestToolEventProgress:
return "ok" return "ok"
loop.tools.execute = execute_tool loop.tools.execute = execute_tool
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign]
return_value=False
)
session_key = "websocket:chat-a" session_key = "websocket:chat-a"
session = loop.sessions.get_or_create(session_key) session = loop.sessions.get_or_create(session_key)
@@ -949,7 +939,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -1048,7 +1037,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -1132,7 +1120,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await asyncio.wait_for(loop._dispatch(InboundMessage( await asyncio.wait_for(loop._dispatch(InboundMessage(
channel="websocket", channel="websocket",
@@ -1181,7 +1168,6 @@ class TestToolEventProgress:
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
_attach_webui_runtime_events(loop, bus) _attach_webui_runtime_events(loop, bus)
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
captured: dict[str, object] = {} captured: dict[str, object] = {}
@@ -1268,7 +1254,6 @@ class TestToolEventProgress:
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[])) provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[]))
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
await loop._dispatch(InboundMessage( await loop._dispatch(InboundMessage(
channel="slack", channel="slack",
@@ -112,7 +112,6 @@ async def test_goal_command_can_implement_plan_from_prior_discussion(tmp_path):
LLMResponse(content="done", tool_calls=[], usage=None), LLMResponse(content="done", tool_calls=[], usage=None),
]) ])
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None)
session = loop.sessions.get_or_create("cli:direct") session = loop.sessions.get_or_create("cli:direct")
session.add_message("user", "Let's agree on the migration implementation.") session.add_message("user", "Let's agree on the migration implementation.")
session.add_message("assistant", "Use the staged migration plan and run integration tests.") session.add_message("assistant", "Use the staged migration plan and run integration tests.")
@@ -166,7 +165,6 @@ async def test_runtime_context_is_persisted_as_next_turn_prompt_prefix(tmp_path)
LLMResponse(content="second answer", usage=None), LLMResponse(content="second answer", usage=None),
]) ])
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None)
session = loop.sessions.get_or_create("cli:direct") session = loop.sessions.get_or_create("cli:direct")
provider_calls: list[str | None] = [] provider_calls: list[str | None] = []
@@ -220,7 +218,6 @@ async def test_webui_quote_reaches_model_without_leaking_into_public_history(tmp
provider.generation = GenerationSettings() provider.generation = GenerationSettings()
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="answer", usage=None)) provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="answer", usage=None))
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None)
session = loop.sessions.get_or_create("websocket:chat") session = loop.sessions.get_or_create("websocket:chat")
quote = webui_quote_runtime_context({ quote = webui_quote_runtime_context({
WEBUI_QUOTE_METADATA: "the selected answer excerpt", WEBUI_QUOTE_METADATA: "the selected answer excerpt",
@@ -265,7 +262,6 @@ async def test_runtime_context_provider_runs_once_across_tool_iterations(tmp_pat
LLMResponse(content="done", usage=None), LLMResponse(content="done", usage=None),
]) ])
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None)
provider_calls = 0 provider_calls = 0
async def provide_context(_request): async def provide_context(_request):
@@ -310,7 +306,6 @@ async def test_non_goal_direct_turn_cannot_reuse_prior_goal_command(tmp_path):
LLMResponse(content="handled as a one-time task", tool_calls=[], usage=None), LLMResponse(content="handled as a one-time task", tool_calls=[], usage=None),
]) ])
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=None)
session = loop.sessions.get_or_create("api:default") session = loop.sessions.get_or_create("api:default")
session.add_message("user", "/goal old completed request") session.add_message("user", "/goal old completed request")
session.add_message("assistant", "The old request is complete.") session.add_message("assistant", "The old request is complete.")
@@ -589,7 +584,6 @@ async def test_next_turn_after_llm_error_keeps_turn_boundary(tmp_path):
loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") loop = AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model")
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
first = await loop._process_message( first = await loop._process_message(
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="first question") InboundMessage(channel="cli", sender_id="user", chat_id="test", content="first question")
+58 -32
View File
@@ -45,6 +45,10 @@ from nanobot.session.recovery import (
RUNTIME_CHECKPOINT_KEY, RUNTIME_CHECKPOINT_KEY,
restore_runtime_checkpoint, restore_runtime_checkpoint,
) )
from nanobot.session.summary import (
SUMMARY_CONTINUATION_TEXT,
SessionSummaryCheckpoint,
)
from nanobot.session.turn_continuation import ( from nanobot.session.turn_continuation import (
INTERNAL_CONTINUATION_META, INTERNAL_CONTINUATION_META,
INTERNAL_CONTINUATION_RUN_STARTED_AT_META, INTERNAL_CONTINUATION_RUN_STARTED_AT_META,
@@ -506,6 +510,60 @@ def test_save_turn_keeps_multimodal_runtime_context_for_model_replay() -> None:
assert public_history_message(session.messages[0])["content"] == [] assert public_history_message(session.messages[0])["content"] == []
def test_save_turn_commits_summary_boundary_without_rewriting_raw_history() -> None:
loop = _mk_loop()
session = Session(key="test:summary-checkpoint")
session.add_message("user", "inspect the project")
messages = [
{"role": "system", "content": "system"},
{"role": "user", "content": "inspect the project"},
{
"role": "assistant",
"content": "",
"tool_calls": [{
"id": "call-1",
"type": "function",
"function": {"name": "inspect", "arguments": "{}"},
}],
},
{
"role": "tool",
"tool_call_id": "call-1",
"name": "inspect",
"content": "full current result",
},
{"role": "assistant", "content": "done"},
]
loop._save_turn(
session,
messages,
skip=2,
summary_checkpoint=SessionSummaryCheckpoint(
summary="Current working-memory checkpoint.",
transcript_boundary=2,
),
input_persisted_early=True,
)
assert [message["role"] for message in session.messages] == [
"user", "user", "assistant", "tool", "assistant",
]
assert session.messages[0]["content"] == "inspect the project"
assert session.messages[1]["content"] == SUMMARY_CONTINUATION_TEXT
assert session.messages[1]["_hidden_history"] is True
assert session.last_archived == 1
assert session.metadata["_last_summary"]["text"] == (
"Current working-memory checkpoint."
)
assert [message["content"] for message in session.get_history()] == [
SUMMARY_CONTINUATION_TEXT,
"",
"full current result",
"done",
]
def test_save_turn_acknowledges_every_merged_recovery_followup() -> None: def test_save_turn_acknowledges_every_merged_recovery_followup() -> None:
"""Persisting a merged injected row retires every durable follow-up ID.""" """Persisting a merged injected row retires every durable follow-up ID."""
loop = _mk_loop() loop = _mk_loop()
@@ -966,7 +1024,6 @@ async def test_runtime_checkpoint_keeps_provider_state_out_of_public_metadata(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None: async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None:
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._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign] loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
msg = InboundMessage(channel="feishu", sender_id="u1", chat_id="c1", content="persist me") msg = InboundMessage(channel="feishu", sender_id="u1", chat_id="c1", content="persist me")
@@ -986,7 +1043,6 @@ async def test_subagent_followup_stages_provider_state_before_turn_runs(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign] loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True loop.provider.can_resume_conversation_state.return_value = True
session = loop.sessions.get_or_create("cli:subagent-crash") session = loop.sessions.get_or_create("cli:subagent-crash")
@@ -1016,7 +1072,6 @@ async def test_subagent_followup_state_is_durable_before_prompt_assembly(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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.provider.can_resume_conversation_state.return_value = True loop.provider.can_resume_conversation_state.return_value = True
loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign] loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign]
side_effect=RuntimeError("prompt boom"), side_effect=RuntimeError("prompt boom"),
@@ -1049,7 +1104,6 @@ async def test_subagent_redelivery_does_not_duplicate_staged_provider_input(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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.provider.can_resume_conversation_state.return_value = True loop.provider.can_resume_conversation_state.return_value = True
build_system_prompt = loop.context.build_system_prompt build_system_prompt = loop.context.build_system_prompt
loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign] loop.context.build_system_prompt = MagicMock( # type: ignore[method-assign]
@@ -1101,7 +1155,6 @@ async def test_subagent_followup_clears_state_before_compatibility_failure(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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.provider.can_resume_conversation_state.side_effect = RuntimeError( loop.provider.can_resume_conversation_state.side_effect = RuntimeError(
"compatibility boom" "compatibility boom"
) )
@@ -1129,7 +1182,6 @@ async def test_subagent_followup_clears_state_before_compatibility_failure(
async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None: async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
loop._unified_session = True loop._unified_session = True
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign] loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
msg = InboundMessage( msg = InboundMessage(
@@ -1230,7 +1282,6 @@ async def test_process_message_persists_media_paths_on_user_turn(tmp_path: Path)
img_b.write_bytes(_PNG_1X1) img_b.write_bytes(_PNG_1X1)
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._run_agent_loop = AsyncMock(side_effect=RuntimeError("interrupt")) # type: ignore[method-assign] loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("interrupt")) # type: ignore[method-assign]
msg = InboundMessage( msg = InboundMessage(
@@ -1262,7 +1313,6 @@ async def test_process_message_persists_media_only_turn_without_text(tmp_path: P
img.write_bytes(_PNG_1X1) img.write_bytes(_PNG_1X1)
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._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign] loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
msg = InboundMessage( msg = InboundMessage(
@@ -1286,7 +1336,6 @@ async def test_process_message_persists_media_only_turn_without_text(tmp_path: P
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_message_does_not_duplicate_early_persisted_user_message(tmp_path: Path) -> None: async def test_process_message_does_not_duplicate_early_persisted_user_message(tmp_path: Path) -> None:
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._run_agent_loop = AsyncMock(return_value=_agent_run_result( loop._run_agent_loop = AsyncMock(return_value=_agent_run_result(
"done", "done",
[ [
@@ -1319,7 +1368,6 @@ async def test_internal_continuation_queues_turn_without_fake_user_history(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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]
session = loop.sessions.get_or_create("feishu:c-auto") session = loop.sessions.get_or_create("feishu:c-auto")
session.metadata[GOAL_STATE_KEY] = { session.metadata[GOAL_STATE_KEY] = {
"status": "active", "status": "active",
@@ -1388,7 +1436,6 @@ async def test_internal_continuation_preserves_streaming_route_metadata(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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]
session = loop.sessions.get_or_create("feishu:c-stream") session = loop.sessions.get_or_create("feishu:c-stream")
session.metadata[GOAL_STATE_KEY] = { session.metadata[GOAL_STATE_KEY] = {
"status": "active", "status": "active",
@@ -1462,7 +1509,6 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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]
session = loop.sessions.get_or_create("websocket:c-auto") session = loop.sessions.get_or_create("websocket:c-auto")
session.metadata[GOAL_STATE_KEY] = { session.metadata[GOAL_STATE_KEY] = {
"status": "active", "status": "active",
@@ -1526,7 +1572,6 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_message_keeps_delivery_chat_for_thread_session(tmp_path: Path) -> None: async def test_process_message_keeps_delivery_chat_for_thread_session(tmp_path: Path) -> None:
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.context.build_messages = MagicMock( # type: ignore[method-assign] loop.context.build_messages = MagicMock( # type: ignore[method-assign]
return_value=[ return_value=[
{"role": "system", "content": "system"}, {"role": "system", "content": "system"},
@@ -1565,7 +1610,6 @@ async def test_process_message_uses_explicit_session_for_goal_context(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
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]
chat_session = loop.sessions.get_or_create("websocket:chat-with-goal") chat_session = loop.sessions.get_or_create("websocket:chat-with-goal")
chat_session.metadata[GOAL_STATE_KEY] = { chat_session.metadata[GOAL_STATE_KEY] = {
"status": "active", "status": "active",
@@ -1713,7 +1757,6 @@ async def test_request_context_uses_effective_key_for_spawn_tool(tmp_path: Path)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(tmp_path: Path) -> None: async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(tmp_path: Path) -> None:
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.provider.chat_with_retry = AsyncMock(return_value=MagicMock()) # unused because _run_agent_loop is stubbed loop.provider.chat_with_retry = AsyncMock(return_value=MagicMock()) # unused because _run_agent_loop is stubbed
session = loop.sessions.get_or_create("feishu:c3") session = loop.sessions.get_or_create("feishu:c3")
@@ -1762,7 +1805,6 @@ async def test_stop_preserves_runtime_checkpoint_for_next_turn(tmp_path: Path) -
from nanobot.command.router import CommandContext from nanobot.command.router import CommandContext
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]
checkpoint_saved = asyncio.Event() checkpoint_saved = asyncio.Event()
@@ -1866,7 +1908,6 @@ async def test_stop_preserves_runtime_checkpoint_for_next_turn(tmp_path: Path) -
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_path: Path) -> None: async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_path: Path) -> None:
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]
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
session.add_message("user", "question") session.add_message("user", "question")
@@ -1913,11 +1954,6 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
assert request.metadata == {"subagent_task_id": "sub-1"} assert request.metadata == {"subagent_task_id": "sub-1"}
assert request.turn_id assert request.turn_id
record_runtime.assert_called_once_with("cli:test", runtime) record_runtime.assert_called_once_with("cli:test", runtime)
assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
assert all(
call.kwargs["runtime"] is runtime
for call in loop.consolidator.maybe_consolidate_by_tokens.call_args_list
)
initial_messages = seen["initial_messages"] initial_messages = seen["initial_messages"]
assert isinstance(initial_messages, list) assert isinstance(initial_messages, list)
non_system = [m for m in initial_messages if m.get("role") != "system"] non_system = [m for m in initial_messages if m.get("role") != "system"]
@@ -1952,7 +1988,6 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) -> None: async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) -> None:
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]
turn_usage = LLMUsage.reported(input_tokens=64, output_tokens=9) turn_usage = LLMUsage.reported(input_tokens=64, output_tokens=9)
async def fake_run_agent_loop(transcript_input, **_kwargs): async def fake_run_agent_loop(transcript_input, **_kwargs):
@@ -1978,9 +2013,6 @@ async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) ->
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) -> None: async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign]
return_value=False
)
async def fake_run_agent_loop(transcript_input, **_kwargs): async def fake_run_agent_loop(transcript_input, **_kwargs):
initial_messages = _assembled_messages(loop.context, transcript_input) initial_messages = _assembled_messages(loop.context, transcript_input)
@@ -2017,9 +2049,6 @@ async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) ->
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Path) -> None: async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign]
return_value=False
)
visited: list[str] = [] visited: list[str] = []
for name in ( for name in (
@@ -2081,7 +2110,6 @@ async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Pat
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp_path: Path) -> None: async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp_path: Path) -> None:
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]
async def fake_run_agent_loop(transcript_input, **_kwargs): async def fake_run_agent_loop(transcript_input, **_kwargs):
initial_messages = _assembled_messages(loop.context, transcript_input) initial_messages = _assembled_messages(loop.context, transcript_input)
@@ -2207,7 +2235,6 @@ async def test_request_context_passes_thread_session_key_to_spawn(tmp_path: Path
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(tmp_path: Path) -> None: async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(tmp_path: Path) -> None:
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]
thread_session = loop.sessions.get_or_create("slack:C123:1700.42") thread_session = loop.sessions.get_or_create("slack:C123:1700.42")
thread_session.add_message("user", "thread question") thread_session.add_message("user", "thread question")
@@ -2266,7 +2293,6 @@ async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(t
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_turn_after_unanswered_user_keeps_tool_call_pairing(tmp_path: Path) -> None: async def test_turn_after_unanswered_user_keeps_tool_call_pairing(tmp_path: Path) -> None:
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]
session = loop.sessions.get_or_create("feishu:c-merge") session = loop.sessions.get_or_create("feishu:c-merge")
session.add_message("user", "earlier question that never got an answer") session.add_message("user", "earlier question that never got an answer")
-2
View File
@@ -46,7 +46,6 @@ def _loop(tmp_path, responses: list[str], **kwargs) -> AgentLoop:
async def test_transient_session_keeps_history_without_persisting_or_durable_tools(tmp_path) -> None: async def test_transient_session_keeps_history_without_persisting_or_durable_tools(tmp_path) -> None:
loop = _loop(tmp_path, ["first answer", "second answer"]) loop = _loop(tmp_path, ["first answer", "second answer"])
loop.context.memory.write_memory("private durable memory") loop.context.memory.write_memory("private durable memory")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock()
key = "websocket:transient-test" key = "websocket:transient-test"
loop.sessions.get_or_create_transient( loop.sessions.get_or_create_transient(
key, key,
@@ -71,7 +70,6 @@ async def test_transient_session_keeps_history_without_persisting_or_durable_too
"assistant", "assistant",
] ]
assert loop.sessions.read_session_file(key) is None assert loop.sessions.read_session_file(key) is None
loop.consolidator.maybe_consolidate_by_tokens.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
+4 -1
View File
@@ -61,7 +61,10 @@ def test_initial_transcript_is_built_from_structured_turn_input() -> None:
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
) )
assert AgentRunner._initial_transcript(spec) == expected messages, compaction = AgentRunner._initial_transcript_and_compaction(spec)
assert messages == expected
assert compaction is None
transcript_builder.assert_called_once_with(transcript_input) transcript_builder.assert_called_once_with(transcript_input)
+488 -7
View File
@@ -7,13 +7,14 @@ 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 ( from nanobot.agent.context_governance import (
BACKFILL_CONTENT, BACKFILL_CONTENT,
ContextGovernanceConfig, ContextGovernanceConfig,
ContextGovernor, ContextGovernor,
ContextWindowExceededError, ContextWindowExceededError,
) )
from nanobot.agent.runner import AgentRunSpec from nanobot.agent.runner import AgentRunner, AgentRunSpec
from nanobot.config.schema import AgentDefaults from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import ( from nanobot.providers.base import (
LLMProvider, LLMProvider,
@@ -22,10 +23,23 @@ from nanobot.providers.base import (
ProviderConversationState, ProviderConversationState,
ToolCallRequest, ToolCallRequest,
) )
from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars _MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
def _build_transcript(transcript: TranscriptInput) -> list[dict]:
system = (
transcript.session_summary["text"]
if transcript.session_summary is not None
else "system"
)
messages = [{"role": "system", "content": system}, *transcript.history]
if transcript.current_message is not None:
messages.append({"role": transcript.current_role, "content": transcript.current_message})
return messages
def _governance_config( def _governance_config(
provider, provider,
tools, tools,
@@ -97,13 +111,16 @@ async def test_runner_locally_fits_oversized_initial_transcript(monkeypatch):
tools = MagicMock() tools = MagicMock()
tools.get_definitions.return_value = [] tools.get_definitions.return_value = []
old_content = "x" * 20_000 old_content = "x" * 20_000
monkeypatch.setattr( estimate = MagicMock(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain", side_effect=lambda _provider, _model, messages, _tools: (
lambda _provider, _model, messages, _tools: (
(600, "test-counter") (600, "test-counter")
if any(message.get("content") == old_content for message in messages) if any(message.get("content") == old_content for message in messages)
else (100, "test-counter") else (100, "test-counter")
), )
)
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
) )
result = await AgentRunner().run(make_run_spec( result = await AgentRunner().run(make_run_spec(
@@ -127,9 +144,441 @@ async def test_runner_locally_fits_oversized_initial_transcript(monkeypatch):
{"role": "system", "content": "system"}, {"role": "system", "content": "system"},
{"role": "user", "content": "continue"}, {"role": "user", "content": "continue"},
] ]
estimated_messages = [call.args[2] for call in estimate.call_args_list]
assert sum(
any(message.get("content") == old_content for message in messages)
for messages in estimated_messages
) == 1
assert len(estimated_messages) == 3
assert any(message.get("content") == old_content for message in result.messages) assert any(message.get("content") == old_content for message in result.messages)
async def test_runner_summarizes_history_and_preserves_current_input(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
prior_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"type": "message", "role": "assistant"}]},
)
candidate_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"type": "message", "role": "assistant", "fresh": True}]},
)
requests: list[tuple[list[dict], object]] = []
async def request(*, messages, provider_context, **_kwargs):
requests.append((messages, provider_context))
return LLMResponse(content="done", provider_state=candidate_state)
provider.chat_with_retry = request
tools = MagicMock()
tools.get_definitions.return_value = []
old_answer = "old answer " * 2_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_answer for message in messages)
else (100, "test-counter")
),
)
consolidate = AsyncMock(return_value="fresh checkpoint")
previous = {"text": "existing checkpoint", "last_active": "2026-08-30T00:00:00"}
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(
history=[
{"role": "user", "content": "old question"},
{"role": "assistant", "content": old_answer},
],
current_message="continue the current task",
session_summary=previous,
),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
provider_state=prior_state,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
consolidate.assert_awaited_once_with(
[
{"role": "system", "content": "existing checkpoint"},
{"role": "user", "content": "old question"},
{"role": "assistant", "content": old_answer},
],
"existing checkpoint",
)
assert requests[0][0] == [
{"role": "system", "content": "fresh checkpoint"},
{"role": "user", "content": SUMMARY_CONTINUATION_TEXT},
{"role": "user", "content": "continue the current task"},
]
assert requests[0][1].conversation_state is None
assert result.provider_state == candidate_state
assert result.summary_checkpoint is not None
assert result.summary_checkpoint.summary == "fresh checkpoint"
assert result.summary_checkpoint.transcript_boundary == 3
assert any(message.get("content") == old_answer for message in result.messages)
async def test_runner_rejects_oversized_delta_without_summarizable_history(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock()
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args, **_kwargs: (600, "test-counter"),
)
consolidate = AsyncMock(return_value=None)
with pytest.raises(ContextWindowExceededError):
await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(
history=[],
current_message="current input is the entire oversized delta",
),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
consolidate.assert_awaited_once_with(
[{"role": "system", "content": "system"}],
None,
)
provider.chat_with_retry.assert_not_awaited()
async def test_runner_governs_history_before_summarizing_it(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = False
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="done"))
tools = MagicMock()
tools.get_definitions.return_value = []
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda _provider, _model, messages, _tools: (
(100, "test-counter")
if messages[0].get("content") == "fresh checkpoint"
else (600, "test-counter")
),
)
consolidate = AsyncMock(return_value="fresh checkpoint")
await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(
history=[
{"role": "user", "content": "inspect"},
{
"role": "assistant",
"content": "",
"tool_calls": [{
"id": "call-missing",
"type": "function",
"function": {"name": "inspect", "arguments": "{}"},
}],
},
],
current_message="continue",
),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
summarized = consolidate.await_args.args[0]
assert [message["role"] for message in summarized] == [
"system", "user", "assistant", "tool",
]
assert summarized[-1]["tool_call_id"] == "call-missing"
assert summarized[-1]["content"] == BACKFILL_CONTENT
@pytest.mark.parametrize(
("scope", "expected_contents", "expected_boundary"),
[
("prior_context", ["system", "accepted question", "accepted answer"], 3),
(
"current_request",
["system", "accepted question", "accepted answer", "inspect the project"],
4,
),
],
)
async def test_native_compaction_uses_provider_request_boundary(
monkeypatch,
scope,
expected_contents,
expected_boundary,
):
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = False
compacted_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"type": "compaction", "encrypted_content": "opaque"}]},
)
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="call-1", name="inspect", arguments={})],
provider_compaction_applied=True,
provider_compaction_state=compacted_state,
provider_compaction_scope=scope,
),
LLMResponse(content="done"),
])
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="complete tool result")
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda *_args: (100, "test-counter"),
)
consolidate = AsyncMock(return_value="portable checkpoint")
consolidate_native = AsyncMock(return_value="portable checkpoint")
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(
history=[
{"role": "user", "content": "accepted question"},
{"role": "assistant", "content": "accepted answer"},
],
current_message="inspect the project",
),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
consolidate_provider_compaction=consolidate_native,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=2,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
consolidate.assert_not_awaited()
consolidate_native.assert_awaited_once()
assert consolidate_native.await_args.args[0] == compacted_state
assert [
message["content"] for message in consolidate_native.await_args.args[1]
] == expected_contents
assert consolidate_native.await_args.args[2] is None
assert result.summary_checkpoint is not None
assert result.summary_checkpoint.transcript_boundary == expected_boundary
assert result.provider_compaction_applied is True
assert any(message.get("content") == "inspect the project" for message in result.messages)
assert any(message.get("content") == "complete tool result" for message in result.messages)
async def test_runner_keeps_current_tool_exchange_outside_summary(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = False
responses = [
LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="call-1", name="inspect", arguments={})],
),
LLMResponse(content="done"),
]
requests: list[list[dict]] = []
async def request(*, messages, **_kwargs):
requests.append(messages)
return responses.pop(0)
provider.chat_with_retry = request
tools = MagicMock()
tools.get_definitions.return_value = []
full_result = "tool-result:" + ("x" * 4_000)
tools.execute = AsyncMock(return_value=full_result)
def estimate(_provider, _model, messages, _tools):
has_tool_result = any(message.get("role") == "tool" for message in messages)
has_old_system = any(
message.get("role") == "system" and message.get("content") == "system"
for message in messages
)
return (600 if has_tool_result and has_old_system else 100, "test-counter")
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
consolidate = AsyncMock(return_value="fresh checkpoint")
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(history=[], current_message="inspect the project"),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=2,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
consolidate.assert_awaited_once_with(
[
{"role": "system", "content": "system"},
{"role": "user", "content": "inspect the project"},
],
None,
)
assert [message["role"] for message in requests[1]] == [
"system", "user", "assistant", "tool",
]
assert requests[1][1]["content"] == SUMMARY_CONTINUATION_TEXT
assert requests[1][-1]["content"] == full_result
assert result.summary_checkpoint is not None
assert result.summary_checkpoint.transcript_boundary == 2
assert any(message.get("content") == full_result for message in result.messages)
async def test_repeated_pressure_advances_summary_boundary(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = False
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="call-1", name="inspect", arguments={})],
),
LLMResponse(
content=None,
tool_calls=[ToolCallRequest(id="call-2", name="inspect", arguments={})],
),
LLMResponse(content="done"),
])
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(side_effect=["result-1", "result-2"])
def estimate(_provider, _model, messages, _tools):
system = messages[0].get("content")
contents = {message.get("content") for message in messages}
if "result-2" in contents:
return (100 if system == "checkpoint-2" else 600, "test-counter")
if "result-1" in contents:
return (100 if system == "checkpoint-1" else 600, "test-counter")
return 100, "test-counter"
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
consolidate = AsyncMock(side_effect=["checkpoint-1", "checkpoint-2"])
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(history=[], current_message="inspect"),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert consolidate.await_count == 2
assert consolidate.await_args_list[0].args[1] is None
assert consolidate.await_args_list[1].args[1] == "checkpoint-1"
second_prefix = consolidate.await_args_list[1].args[0]
assert second_prefix[0]["content"] == "checkpoint-1"
assert any(message.get("content") == "result-1" for message in second_prefix)
assert result.final_content == "done"
assert result.summary_checkpoint is not None
assert result.summary_checkpoint.summary == "checkpoint-2"
assert result.summary_checkpoint.transcript_boundary == 4
async def test_runner_refuses_checkpoint_that_cannot_fit_with_delta(monkeypatch):
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="unexpected"))
tools = MagicMock()
tools.get_definitions.return_value = []
old_answer = "old answer"
current_input = "current input must remain intact"
def estimate(_provider, _model, messages, _tools):
contents = {message.get("content") for message in messages}
if old_answer in contents or current_input in contents:
return 600, "test-counter"
return 100, "test-counter"
monkeypatch.setattr(
"nanobot.agent.context_governance.estimate_prompt_tokens_chain",
estimate,
)
consolidate = AsyncMock(return_value="small checkpoint")
with pytest.raises(ContextWindowExceededError):
await AgentRunner().run(make_run_spec(
provider,
initial_messages=None,
transcript_input=TranscriptInput(
history=[{"role": "assistant", "content": old_answer}],
current_message=current_input,
),
transcript_builder=_build_transcript,
consolidate_history=consolidate,
tools=tools,
model="test-model",
context_window_tokens=2_000,
context_block_limit=500,
max_tokens=100,
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
summarized = consolidate.await_args.args[0]
assert all(message.get("content") != current_input for message in summarized)
provider.chat_with_retry.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatch): async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatch):
from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.agent.hook import AgentHook, AgentHookContext
@@ -145,7 +594,7 @@ async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatc
"nanobot.agent.context_governance.estimate_prompt_tokens_chain", "nanobot.agent.context_governance.estimate_prompt_tokens_chain",
lambda _provider, _model, messages, _tools: ( lambda _provider, _model, messages, _tools: (
(2_000, "test-counter") (2_000, "test-counter")
if any(message.get("content") == oversized for message in messages) if any(oversized in str(message.get("content")) for message in messages)
else (100, "test-counter") else (100, "test-counter")
), ),
) )
@@ -491,6 +940,39 @@ def test_changed_messages_use_local_estimate_after_reported_usage(monkeypatch):
estimate.assert_called_once() estimate.assert_called_once()
def test_resumed_provider_context_avoids_full_transcript_estimate(monkeypatch):
provider = MagicMock(spec=LLMProvider)
tools = MagicMock()
tools.get_definitions.return_value = []
spec = make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "pending delta"}],
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("resumed provider context must be authoritative")
),
)
pressure = ContextGovernor().request_pressure(
_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(),
request_context_tokens=100,
)
assert pressure is None
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runner_counts_resumed_provider_state_before_dispatch(monkeypatch): async def test_runner_counts_resumed_provider_state_before_dispatch(monkeypatch):
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
@@ -832,7 +1314,6 @@ async def test_backfill_repairs_model_context_without_shifting_save_turn_boundar
model="test-model", model="test-model",
) )
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
session.messages = [ session.messages = [
+53 -26
View File
@@ -750,25 +750,27 @@ async def test_pending_injection_resolves_its_own_runtime_context(tmp_path):
), ),
] ]
injected = [message for message in result.messages if message.get("role") == "user"][-1] injected = [message for message in result.messages if message.get("role") == "user"][-2:]
assert "follow-up from the second speaker" in str(injected["content"]) assert str(injected[0]["content"]).startswith("follow-up from the second speaker\n\n")
assert str(injected[1]["content"]).startswith("another follow-up\n\n")
model_messages = provider.chat_with_retry.await_args_list[-1].kwargs["messages"] model_messages = provider.chat_with_retry.await_args_list[-1].kwargs["messages"]
assert "telegram | group-1 | user-b | message-2" in str(model_messages) assert "telegram | group-1 | user-b | message-2" in str(model_messages)
assert "Bob | topic-7" in str(model_messages) assert "Bob | topic-7" in str(model_messages)
assert "telegram | group-1 | user-c | message-3" in str(model_messages) assert "telegram | group-1 | user-c | message-3" in str(model_messages)
assert "Carol | topic-7" in str(model_messages) assert "Carol | topic-7" in str(model_messages)
assert injected["_meta"][RUNTIME_CONTEXT_MESSAGE_META]["sources"] == [ assert all(
"identity", message["_meta"][RUNTIME_CONTEXT_MESSAGE_META]["sources"] == ["identity"]
"identity", for message in injected
] )
loop._save_turn(session, result.messages, skip=1) loop._save_turn(session, result.messages, skip=1)
persisted = [message for message in session.messages if message.get("role") == "user"][-1] persisted = [message for message in session.messages if message.get("role") == "user"][-2:]
assert "telegram | group-1 | user-b | message-2" in str(persisted["content"]) assert "telegram | group-1 | user-b | message-2" in str(persisted[0]["content"])
assert "telegram | group-1 | user-c | message-3" in str(persisted["content"]) assert "telegram | group-1 | user-c | message-3" in str(persisted[1]["content"])
assert public_history_message(persisted)["content"] == ( assert [public_history_message(message)["content"] for message in persisted] == [
"follow-up from the second speaker\n\nanother follow-up" "follow-up from the second speaker",
) "another follow-up",
]
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -835,8 +837,8 @@ async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_p
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runner_merges_multiple_injected_user_messages_without_losing_media(): async def test_model_request_merges_injected_user_messages_without_losing_media():
"""Multiple injected follow-ups should not create lossy consecutive user messages.""" """The model copy may merge follow-ups while the raw transcript keeps each event."""
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
provider = MagicMock() provider = MagicMock()
@@ -895,10 +897,17 @@ async def test_runner_merges_multiple_injected_user_messages_without_losing_medi
for block in injected["content"] for block in injected["content"]
if isinstance(block, dict) if isinstance(block, dict)
) )
assert [message["content"] for message in result.messages[-3:-1]] == [
[
{"type": "image_url", "image_url": {"url": "data:image/png;base64,abc"}},
{"type": "text", "text": "look at this"},
],
"and answer briefly",
]
def test_runner_merge_keeps_all_recovery_followup_ids() -> None: def test_runner_append_keeps_recovery_followups_separate() -> None:
"""Merged follow-ups stay acknowledged together after a later save.""" """Each raw follow-up keeps its own recovery identity."""
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY
@@ -908,10 +917,12 @@ def test_runner_merge_keeps_all_recovery_followup_ids() -> None:
[{"role": "user", "content": "second", PENDING_FOLLOWUP_ID_KEY: "two"}], [{"role": "user", "content": "second", PENDING_FOLLOWUP_ID_KEY: "two"}],
) )
assert messages[-1][PENDING_FOLLOWUP_ID_KEY] == ["one", "two"] assert [message["content"] for message in messages] == ["first", "second"]
assert [message[PENDING_FOLLOWUP_ID_KEY] for message in messages] == ["one", "two"]
def test_runner_merge_preserves_runtime_markers_with_media() -> None: def test_model_request_merge_preserves_runtime_markers_with_media() -> None:
from nanobot.agent.context_governance import ContextGovernor
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
from nanobot.runtime_context import ( from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META, RUNTIME_CONTEXT_HISTORY_META,
@@ -948,8 +959,9 @@ def test_runner_merge_preserves_runtime_markers_with_media() -> None:
}, },
]) ])
assert len(messages) == 1 assert len(messages) == 2
merged = messages[0] merged = ContextGovernor._merge_adjacent_user_messages_for_model(messages)[0]
assert len(messages) == 2
assert "private first" in str(merged["content"]) assert "private first" in str(merged["content"])
assert "private second" in str(merged["content"]) assert "private second" in str(merged["content"])
persisted = { persisted = {
@@ -1681,15 +1693,17 @@ async def test_drain_injections_after_recoverable_tool_error():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_drain_injections_on_llm_error(): async def test_drain_injections_on_llm_error():
"""Pending injections should be drained when the LLM returns an error finish_reason.""" """A follow-up after an error stays raw and reaches the next model request."""
from nanobot.agent.runner import AgentRunner from nanobot.agent.runner import AgentRunner
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
provider = MagicMock() provider = MagicMock()
call_count = {"n": 0} call_count = {"n": 0}
requests: list[list[dict]] = []
async def chat_with_retry(*, messages, **kwargs): async def chat_with_retry(*, messages, **kwargs):
call_count["n"] += 1 call_count["n"] += 1
requests.append(messages)
if call_count["n"] == 1: if call_count["n"] == 1:
return LLMResponse( return LLMResponse(
content=None, content=None,
@@ -1713,11 +1727,20 @@ async def test_drain_injections_on_llm_error():
runner = AgentRunner() runner = AgentRunner()
result = await runner.run(make_run_spec(provider, result = await runner.run(make_run_spec(provider,
initial_messages=[ initial_messages=None,
transcript_input=TranscriptInput(
history=[
{"role": "user", "content": "hello"}, {"role": "user", "content": "hello"},
{"role": "assistant", "content": "previous response"}, {"role": "assistant", "content": "previous response"},
{"role": "user", "content": "trigger error"}, {"role": "user", "content": "trigger error"},
], ],
current_message=None,
),
transcript_builder=lambda transcript: [
{"role": "system", "content": "system"},
*transcript.history,
],
consolidate_history=AsyncMock(return_value=None),
tools=tools, tools=tools,
model="test-model", model="test-model",
max_iterations=5, max_iterations=5,
@@ -1727,11 +1750,15 @@ async def test_drain_injections_on_llm_error():
assert result.had_injections is True assert result.had_injections is True
assert result.final_content == "recovered answer" assert result.final_content == "recovered answer"
injected = [ assert "follow-up after LLM error" in str(requests[1])
m for m in result.messages assert [
if m.get("role") == "user" and "follow-up after LLM error" in str(m.get("content", "")) message["content"]
for message in result.messages
if message.get("role") == "user"
][-2:] == [
"trigger error",
"follow-up after LLM error",
] ]
assert len(injected) == 1
@pytest.mark.asyncio @pytest.mark.asyncio
+36 -170
View File
@@ -1,10 +1,11 @@
from nanobot.providers.base import ProviderConversationState
from nanobot.runtime_context import ( from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META, RUNTIME_CONTEXT_HISTORY_META,
RuntimeContextBlock, RuntimeContextBlock,
append_runtime_context, append_runtime_context,
) )
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
from nanobot.session.manager import Session, SessionManager from nanobot.session.manager import Session, SessionManager
from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT
def _assert_no_orphans(history: list[dict]) -> None: def _assert_no_orphans(history: list[dict]) -> None:
@@ -136,58 +137,6 @@ def test_legitimate_tool_pairs_preserved_after_trim():
assert history[0]["role"] == "user" assert history[0]["role"] == "user"
def test_retain_recent_legal_suffix_keeps_recent_messages():
session = Session(key="test:trim")
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
session.retain_recent_legal_suffix(4)
assert len(session.messages) == 4
assert session.messages[0]["content"] == "msg6"
assert session.messages[-1]["content"] == "msg9"
def test_retain_recent_legal_suffix_adjusts_last_archived():
session = Session(key="test:trim-cons")
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
session.last_archived = 7
session.retain_recent_legal_suffix(4)
assert len(session.messages) == 4
assert session.last_archived == 1
def test_retain_recent_legal_suffix_zero_clears_session():
session = Session(key="test:trim-zero")
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
session.last_archived = 5
session.retain_recent_legal_suffix(0)
assert session.messages == []
assert session.last_archived == 0
def test_retain_recent_legal_suffix_keeps_legal_tool_boundary():
session = Session(key="test:trim-tools")
session.messages.append({"role": "user", "content": "old"})
session.messages.extend(_tool_turn("old", 0))
session.messages.append({"role": "user", "content": "keep"})
session.messages.extend(_tool_turn("keep", 0))
session.messages.append({"role": "assistant", "content": "done"})
session.retain_recent_legal_suffix(4)
history = session.get_history(max_messages=500)
_assert_no_orphans(history)
assert history[0]["role"] == "user"
assert history[0]["content"] == "keep"
# --- last_archived > 0 --- # --- last_archived > 0 ---
def test_orphan_trim_with_last_archived(): def test_orphan_trim_with_last_archived():
@@ -635,6 +584,40 @@ def test_fork_session_allows_index_equal_to_user_count(tmp_path):
assert [m["content"] for m in forked.messages] == ["round1", "answer1"] assert [m["content"] for m in forked.messages] == ["round1", "answer1"]
def test_fork_session_user_index_ignores_hidden_checkpoint_anchor(tmp_path):
manager = SessionManager(tmp_path)
source = manager.get_or_create("websocket:source")
source.add_message("user", "round1")
source.add_message("assistant", "answer1")
source.add_message("user", "round2")
source.add_message(
"user",
SUMMARY_CONTINUATION_TEXT,
**{HIDDEN_HISTORY_META: True},
)
source.add_message("assistant", "answer2")
source.last_archived = 3
source.metadata["_last_summary"] = {"text": "round1 and round2"}
manager.save(source)
forked = manager.fork_session_before_user_index(
"websocket:source",
"websocket:fork",
2,
)
assert forked is not None
assert [message["content"] for message in forked.messages] == [
"round1",
"answer1",
"round2",
SUMMARY_CONTINUATION_TEXT,
"answer2",
]
assert forked.last_archived == 3
assert forked.metadata["_last_summary"]["text"] == "round1 and round2"
def test_fork_session_drops_summary_when_fork_point_is_inside_archived_prefix(tmp_path): def test_fork_session_drops_summary_when_fork_point_is_inside_archived_prefix(tmp_path):
manager = SessionManager(tmp_path) manager = SessionManager(tmp_path)
source = manager.get_or_create("websocket:source") source = manager.get_or_create("websocket:source")
@@ -756,44 +739,6 @@ def test_get_history_recovers_user_when_token_slice_would_be_assistant_only(monk
assert [m["content"] for m in history] == ["u2", "a2"] assert [m["content"] for m in history] == ["u2", "a2"]
def test_retain_recent_legal_suffix_hard_cap_with_long_non_user_chain():
session = Session(key="test:hard-cap-chain")
session.messages.append({"role": "user", "content": "u0"})
session.messages.append(
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "c1", "type": "function", "function": {"name": "x", "arguments": "{}"}}
],
}
)
for i in range(12):
session.messages.append({"role": "assistant", "content": f"a{i}"})
session.retain_recent_legal_suffix(6)
assert len(session.messages) <= 6
def test_retain_recent_legal_suffix_can_extend_to_user_for_long_recent_turn():
session = Session(key="test:extend-to-user")
session.messages.append({"role": "user", "content": "old"})
session.messages.append({"role": "assistant", "content": "old answer"})
session.messages.append({"role": "user", "content": "record this"})
for i in range(4):
session.messages.extend(_tool_turn("recent", i))
session.messages.append({"role": "assistant", "content": "done"})
session.retain_recent_legal_suffix(8, extend_to_user=True)
assert len(session.messages) > 8
assert session.messages[0]["content"] == "record this"
assert session.messages[-1]["content"] == "done"
history = session.get_history(max_messages=500)
_assert_no_orphans(history)
def test_get_history_can_extend_to_user_for_long_recent_turn(): def test_get_history_can_extend_to_user_for_long_recent_turn():
session = Session(key="test:history-extend-to-user") session = Session(key="test:history-extend-to-user")
session.messages.append({"role": "user", "content": "old"}) session.messages.append({"role": "user", "content": "old"})
@@ -828,82 +773,3 @@ def test_get_history_extend_to_user_keeps_newer_user_inside_window():
assert [m["content"] for m in history] == ["new question", "new answer"] assert [m["content"] for m in history] == ["new question", "new answer"]
_assert_no_orphans(history) _assert_no_orphans(history)
def test_retain_recent_legal_suffix_returns_dropped_messages():
"""retain_recent_legal_suffix returns the actually-dropped messages."""
session = Session(
key="test:return-dropped",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
),
)
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
result = session.retain_recent_legal_suffix(4)
assert len(result.dropped) == 6
assert [m["content"] for m in result.dropped] == [f"msg{i}" for i in range(6)]
assert len(session.messages) == 4
assert result.already_consolidated_count == 0
assert session.provider_state is None
def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
"""No messages dropped → empty list returned."""
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
session = Session(key="test:no-drop", provider_state=state)
for i in range(3):
session.messages.append({"role": "user", "content": f"msg{i}"})
result = session.retain_recent_legal_suffix(4)
assert result.dropped == []
assert result.already_consolidated_count == 0
assert len(session.messages) == 3
assert session.provider_state is state
def test_retain_recent_legal_suffix_returns_all_on_zero():
"""max_messages=0 clears session and returns all messages."""
session = Session(key="test:zero-return")
for i in range(5):
session.messages.append({"role": "user", "content": f"msg{i}"})
session.last_archived = 3
result = session.retain_recent_legal_suffix(0)
assert len(result.dropped) == 5
assert result.already_consolidated_count == 3
assert session.messages == []
def test_retain_recent_legal_suffix_last_archived_correct_in_else_branch():
"""last_archived should count retained messages from the old archived prefix."""
session = Session(key="test:else-lc-correct")
# 20 messages: u0..u9, a0..a9
for i in range(10):
session.messages.append({"role": "user", "content": f"u{i}"})
for i in range(10):
session.messages.append({"role": "assistant", "content": f"a{i}"})
session.last_archived = 12 # u0..u9, a0, a1 archived
result = session.retain_recent_legal_suffix(4)
# Retained messages start from latest user (u9) + max_messages forward
# so retained = [u9, a0..a9][:4] → but these are from original indices 9..12
# Of those, indices 9,10,11 are < 12 (before_lc), so new_lc = 3
assert session.last_archived == 3
# already_cons should count dropped messages with original index < 12
assert result.already_consolidated_count == 9
-225
View File
@@ -1,225 +0,0 @@
from nanobot.session.manager import Session
def _assert_no_orphans(history: list[dict]) -> None:
declared = {
tc["id"]
for m in history
if m.get("role") == "assistant"
for tc in (m.get("tool_calls") or [])
}
orphans = [
m.get("tool_call_id")
for m in history
if m.get("role") == "tool" and m.get("tool_call_id") not in declared
]
assert orphans == [], f"orphan tool_call_ids: {orphans}"
def _delivery(content: str) -> dict:
return {"role": "assistant", "content": content, "_channel_delivery": True}
def _tool_turn(prefix: str, idx: int) -> list[dict]:
return [
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": f"{prefix}_{idx}_a",
"type": "function",
"function": {"name": "x", "arguments": "{}"},
},
{
"id": f"{prefix}_{idx}_b",
"type": "function",
"function": {"name": "y", "arguments": "{}"},
},
],
},
{"role": "tool", "tool_call_id": f"{prefix}_{idx}_a", "name": "x", "content": "ok"},
{"role": "tool", "tool_call_id": f"{prefix}_{idx}_b", "name": "y", "content": "ok"},
]
def _contents(messages: list[dict]) -> list[str]:
return [m.get("content") for m in messages]
def _has_delivery(messages: list[dict]) -> bool:
return any(m.get("_channel_delivery") for m in messages)
# --- Hard-cap trimming must preserve a proactive delivery the user replied to ---
def test_retain_hard_cap_keeps_delivery_before_user():
session = Session(key="test:cap-delivery")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append(_delivery("Remember to drink water"))
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "great"})
session.retain_recent_legal_suffix(3)
assert _has_delivery(session.messages), "delivery dropped by hard-cap trim"
assert _contents(session.messages) == [
"Remember to drink water",
"ok",
"great",
]
def test_retain_hard_cap_matches_get_history_boundary():
"""The trimmed suffix must start on the same message as get_history()."""
session = Session(key="test:cap-boundary")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append(_delivery("You have 3 pending tasks"))
session.messages.append({"role": "user", "content": "show them"})
session.messages.append({"role": "assistant", "content": "done"})
expected = session.get_history(max_messages=3)
session.retain_recent_legal_suffix(3)
assert _contents(session.messages) == _contents(expected)
def test_retain_extend_to_user_keeps_delivery_before_recovered_user():
session = Session(key="test:extend-delivery")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append({"role": "assistant", "content": "work"})
session.messages.append(_delivery("Reminder: deploy at 17:00"))
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "a1"})
session.messages.append({"role": "assistant", "content": "a2"})
session.messages.append({"role": "assistant", "content": "a3"})
session.retain_recent_legal_suffix(3, extend_to_user=True)
assert _has_delivery(session.messages), "delivery dropped by extend_to_user trim"
assert session.messages[0]["content"] == "Reminder: deploy at 17:00"
assert session.messages[-1]["content"] == "a3"
def test_retain_extend_to_user_matches_get_history_boundary():
session = Session(key="test:extend-boundary")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append({"role": "assistant", "content": "work"})
session.messages.append(_delivery("Reminder: review the draft"))
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "a1"})
session.messages.append({"role": "assistant", "content": "a2"})
session.messages.append({"role": "assistant", "content": "a3"})
expected = session.get_history(max_messages=3, extend_to_user=True)
session.retain_recent_legal_suffix(3, extend_to_user=True)
assert _contents(session.messages) == _contents(expected)
def test_retain_extend_to_user_does_not_extend_delivery_only_tail():
session = Session(key="test:extend-no-user")
for i in range(4):
session.messages.append(_delivery(f"notification {i}"))
session.retain_recent_legal_suffix(3, extend_to_user=True)
assert _contents(session.messages) == [
"notification 1",
"notification 2",
"notification 3",
]
# --- Only the immediately-preceding delivery is part of the anchor ---
def test_retain_keeps_only_immediate_delivery():
session = Session(key="test:multi-delivery")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append(_delivery("old scheduled note"))
session.messages.append(_delivery("new scheduled note"))
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "great"})
session.retain_recent_legal_suffix(3)
kept = _contents(session.messages)
assert kept == ["new scheduled note", "ok", "great"], kept
def test_retain_drops_delivery_not_adjacent_to_anchor_user():
"""A delivery that does not immediately precede the retained user turn is
not part of the anchor and should not be force-retained."""
session = Session(key="test:nonadjacent")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append(_delivery("unrelated scheduled note"))
session.messages.append({"role": "assistant", "content": "reply"})
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "great"})
session.retain_recent_legal_suffix(2)
assert not _has_delivery(session.messages)
assert _contents(session.messages) == ["ok", "great"]
def test_compact_probe_keeps_delivery_in_visible_suffix():
"""compact_idle_session() trims a probe copy with extend_to_user=True; the
visible suffix it keeps must still contain the delivery message."""
tail = [
{"role": "user", "content": "setup"},
{"role": "assistant", "content": "work"},
_delivery("Reminder: deploy at 17:00"),
{"role": "user", "content": "ok"},
{"role": "assistant", "content": "a1"},
{"role": "assistant", "content": "a2"},
{"role": "assistant", "content": "a3"},
]
probe = Session(key="test:probe", messages=tail)
probe.retain_recent_legal_suffix(3, extend_to_user=True)
assert _has_delivery(probe.messages)
assert probe.messages[0]["content"] == "Reminder: deploy at 17:00"
# --- Trimming must stay coherent with the rest of replay ---
def test_retain_then_replay_keeps_delivery_and_no_orphans():
session = Session(key="test:replay-after-trim")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append(_delivery("You have 3 pending tasks"))
session.messages.append({"role": "user", "content": "show them"})
session.messages.extend(_tool_turn("cur", 0))
session.messages.append({"role": "assistant", "content": "done"})
session.retain_recent_legal_suffix(6)
assert _has_delivery(session.messages)
history = session.get_history(max_messages=500)
_assert_no_orphans(history)
assert any(m.get("content") == "You have 3 pending tasks" for m in history)
def test_retain_keeps_delivery_when_user_inside_window():
"""When the capped window already contains a user, its immediately
preceding delivery must stay attached to it."""
session = Session(key="test:window-user")
session.messages.append({"role": "user", "content": "setup"})
session.messages.append({"role": "assistant", "content": "a0"})
session.messages.append(_delivery("Reminder"))
session.messages.append({"role": "user", "content": "ok"})
session.messages.append({"role": "assistant", "content": "a1"})
session.messages.append({"role": "assistant", "content": "a2"})
expected = session.get_history(max_messages=4)
session.retain_recent_legal_suffix(4)
assert _has_delivery(session.messages)
assert _contents(session.messages) == _contents(expected)
+1 -111
View File
@@ -28,7 +28,7 @@ from nanobot.command.router import CommandContext, CommandRouter
from nanobot.config.schema import AgentDefaults, Config from nanobot.config.schema import AgentDefaults, Config
from nanobot.providers.base import GenerationSettings from nanobot.providers.base import GenerationSettings
from nanobot.session.keys import UNIFIED_SESSION_KEY from nanobot.session.keys import UNIFIED_SESSION_KEY
from nanobot.session.manager import Session, SessionManager from nanobot.session.manager import SessionManager
from nanobot.utils.llm_runtime import LLMRuntime from nanobot.utils.llm_runtime import LLMRuntime
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -334,116 +334,6 @@ class TestCmdNewUnifiedSession:
assert len(sessions.get_or_create("discord:999").messages) == 1 assert len(sessions.get_or_create("discord:999").messages) == 1
# ---------------------------------------------------------------------------
# TestConsolidationUnaffectedByUnifiedSession — consolidation is key-agnostic
# ---------------------------------------------------------------------------
class TestConsolidationUnaffectedByUnifiedSession:
"""maybe_consolidate_by_tokens() behaviour is identical regardless of session key."""
@pytest.mark.asyncio
async def test_consolidation_skips_empty_session_for_unified_key(self):
"""Empty unified:default session → consolidation exits immediately, archive not called."""
from nanobot.agent.memory import Consolidator, MemoryStore
store = MagicMock(spec=MemoryStore)
mock_provider = MagicMock()
mock_provider.chat_with_retry = AsyncMock(return_value=MagicMock(content="summary"))
runtime = _runtime(mock_provider)
# Use spec= so MagicMock doesn't auto-generate AsyncMock for non-async methods,
# which would leave unawaited coroutines and trigger RuntimeWarning.
sessions = MagicMock(spec=SessionManager)
consolidator = Consolidator(
store=store,
sessions=sessions,
build_messages=MagicMock(return_value=[]),
get_tool_definitions=MagicMock(return_value=[]),
)
consolidator.archive_session = AsyncMock()
session = Session(key="unified:default")
session.messages = []
sessions.get_or_create.return_value = session
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
consolidator.archive_session.assert_not_called()
@pytest.mark.asyncio
async def test_consolidation_behaviour_identical_for_any_key(self):
"""Archive call count is the same for 'telegram:123' and 'unified:default'
under identical token conditions."""
from nanobot.agent.memory import Consolidator, MemoryStore
archive_calls: dict[str, int] = {}
for key in ("telegram:123", "unified:default"):
store = MagicMock(spec=MemoryStore)
mock_provider = MagicMock()
mock_provider.chat_with_retry = AsyncMock(return_value=MagicMock(content="summary"))
runtime = _runtime(mock_provider)
sessions = MagicMock(spec=SessionManager)
consolidator = Consolidator(
store=store,
sessions=sessions,
build_messages=MagicMock(return_value=[]),
get_tool_definitions=MagicMock(return_value=[]),
)
session = Session(key=key)
session.messages = [] # empty → exits immediately for both keys
sessions.get_or_create.return_value = session
consolidator.archive_session = AsyncMock()
await consolidator.maybe_consolidate_by_tokens(
session,
runtime=runtime,
)
archive_calls[key] = consolidator.archive_session.call_count
assert archive_calls["telegram:123"] == archive_calls["unified:default"] == 0
@pytest.mark.asyncio
async def test_consolidation_triggers_when_over_budget_unified_key(self):
"""When tokens exceed budget, consolidation attempts to find a boundary —
behaviour is identical to any other session key."""
from nanobot.agent.memory import Consolidator, MemoryStore
store = MagicMock(spec=MemoryStore)
mock_provider = MagicMock()
runtime = _runtime(mock_provider)
sessions = MagicMock(spec=SessionManager)
consolidator = Consolidator(
store=store,
sessions=sessions,
build_messages=MagicMock(return_value=[]),
get_tool_definitions=MagicMock(return_value=[]),
)
session = Session(key="unified:default")
session.messages = [{"role": "user", "content": "msg"}]
sessions.get_or_create.return_value = session
# Simulate over-budget: estimated > budget
consolidator.estimate_session_prompt_tokens = MagicMock(return_value=(950, "tiktoken"))
# No valid boundary found → returns gracefully without archiving
consolidator.pick_consolidation_boundary = MagicMock(return_value=None)
consolidator.archive_session = AsyncMock()
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
# estimate was called (consolidation was attempted)
consolidator.estimate_session_prompt_tokens.assert_called_once_with(
session,
runtime=runtime,
)
# but archive was not called (no valid boundary)
consolidator.archive_session.assert_not_called()
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# TestStopCommandWithUnifiedSession — /stop command integration # TestStopCommandWithUnifiedSession — /stop command integration
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+2 -23
View File
@@ -1837,12 +1837,6 @@ def test_agent_workspace_override_wins_over_config_workspace(mock_agent_runtime,
assert passed_config.workspace_path == workspace_path assert passed_config.workspace_path == workspace_path
def test_heartbeat_retains_recent_messages_by_default():
config = Config()
assert config.gateway.heartbeat.keep_recent_messages == 8
@pytest.mark.parametrize( @pytest.mark.parametrize(
"content, expected", "content, expected",
[ [
@@ -2101,7 +2095,7 @@ def _patch_cli_command_runtime(
monkeypatch.setattr("nanobot.config.paths.get_cron_dir", get_cron_dir) monkeypatch.setattr("nanobot.config.paths.get_cron_dir", get_cron_dir)
def test_heartbeat_empty_response_still_retains_recent_messages( def test_heartbeat_empty_response_is_not_evaluated(
monkeypatch, tmp_path: Path, monkeypatch, tmp_path: Path,
) -> None: ) -> None:
config_file = _write_instance_config(tmp_path) config_file = _write_instance_config(tmp_path)
@@ -2119,21 +2113,9 @@ def test_heartbeat_empty_response_still_retains_recent_messages(
bus.publish_outbound = AsyncMock() bus.publish_outbound = AsyncMock()
seen: dict[str, object] = {} seen: dict[str, object] = {}
class _FakeSession:
def retain_recent_legal_suffix(self, limit: int) -> None:
seen["retained_limit"] = limit
class _FakeSessionManager: class _FakeSessionManager:
def __init__(self, _workspace: Path) -> None: def __init__(self, _workspace: Path) -> None:
self.session = _FakeSession() pass
seen["heartbeat_session"] = self.session
def get_or_create(self, key: str) -> _FakeSession:
seen["session_key"] = key
return self.session
def save(self, session: _FakeSession) -> None:
seen["saved_session"] = session
def list_sessions(self) -> list[dict[str, str]]: def list_sessions(self) -> list[dict[str, str]]:
return [{"key": "telegram:u1"}] return [{"key": "telegram:u1"}]
@@ -2199,9 +2181,6 @@ def test_heartbeat_empty_response_still_retains_recent_messages(
response = asyncio.run(cron.on_job(CronJob(id="heartbeat", name="heartbeat"))) response = asyncio.run(cron.on_job(CronJob(id="heartbeat", name="heartbeat")))
assert response is None assert response is None
assert seen["session_key"] == "heartbeat"
assert seen["retained_limit"] == config.gateway.heartbeat.keep_recent_messages
assert seen["saved_session"] is seen["heartbeat_session"]
def test_webui_yes_creates_config_and_enables_local_websocket( def test_webui_yes_creates_config_and_enables_local_websocket(
+9
View File
@@ -13,3 +13,12 @@ def test_gateway_restart_mode_accepts_camel_alias():
def test_gateway_restart_mode_rejects_unknown_value(): def test_gateway_restart_mode_rejects_unknown_value():
with pytest.raises(ValueError): with pytest.raises(ValueError):
GatewayConfig(restart_mode="service") GatewayConfig(restart_mode="service")
def test_heartbeat_ignores_removed_retention_limit():
config = Config.model_validate(
{"gateway": {"heartbeat": {"keepRecentMessages": 8}}}
)
heartbeat = config.model_dump(by_alias=True)["gateway"]["heartbeat"]
assert "keepRecentMessages" not in heartbeat
+29 -5
View File
@@ -22,7 +22,10 @@ from nanobot.providers.openai_codex_provider import (
_request_codex, _request_codex,
_should_retry_status, _should_retry_status,
) )
from nanobot.providers.openai_responses import build_responses_state from nanobot.providers.openai_responses import (
build_responses_state,
responses_state_items,
)
from nanobot.providers.registry import find_by_name from nanobot.providers.registry import find_by_name
@@ -811,12 +814,33 @@ async def test_codex_compacts_state_at_ninety_percent_before_next_request(
) )
assert response.content == "done" assert response.content == "done"
assert len(bodies) == 2 assert response.provider_compaction_applied is True
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"} assert response.provider_compaction_state is not None
assert bodies[1]["input"][-1] == { assert response.provider_compaction_scope == "prior_context"
assert responses_state_items(response.provider_compaction_state) == [{
"type": "compaction", "type": "compaction",
"encrypted_content": "compacted opaque state", "encrypted_content": "compacted opaque state",
} }]
assert len(bodies) == 2
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"}
assert not any(
item.get("role") == "user"
and "new question" in str(item.get("content"))
for item in bodies[0]["input"]
)
assert {
"type": "compaction",
"encrypted_content": "compacted opaque state",
} in bodies[1]["input"]
assert bodies[1]["input"].index({
"type": "compaction",
"encrypted_content": "compacted opaque state",
}) < next(
index
for index, item in enumerate(bodies[1]["input"])
if item.get("role") == "user"
and "new question" in str(item.get("content"))
)
assert not any( assert not any(
item.get("type") == "reasoning" item.get("type") == "reasoning"
for item in bodies[1]["input"] for item in bodies[1]["input"]
+39
View File
@@ -712,6 +712,45 @@ class TestParseResponseOutput:
assert result.provider_state is not None assert result.provider_state is not None
assert responses_state_items(result.provider_state) == [*input_items, *output] assert responses_state_items(result.provider_state) == [*input_items, *output]
def test_marks_only_a_new_response_compaction(self):
compacted = parse_response_output(
{
"output": [
{"type": "compaction", "encrypted_content": "opaque"},
{"type": "message", "role": "assistant", "content": "done"},
],
"status": "completed",
"usage": {},
},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "old"}],
)
replayed = parse_response_output(
{
"output": [
{"type": "message", "role": "assistant", "content": "continued"},
],
"status": "completed",
"usage": {},
},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[
{"type": "compaction", "encrypted_content": "opaque"},
],
)
assert compacted.provider_compaction_applied is True
assert compacted.provider_compaction_state is not None
assert compacted.provider_compaction_scope == "current_request"
assert responses_state_items(compacted.provider_compaction_state) == [
{"type": "compaction", "encrypted_content": "opaque"},
]
assert replayed.provider_compaction_applied is False
assert replayed.provider_compaction_state is None
assert replayed.provider_compaction_scope is None
class TestResponsesConversationState: class TestResponsesConversationState:
def test_server_compaction_prunes_superseded_prefix(self): def test_server_compaction_prunes_superseded_prefix(self):
+5 -6
View File
@@ -1412,7 +1412,6 @@ async def test_sessions_ingest_imports_transcript_without_running_model(tmp_path
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
bot._loop.process_direct = AsyncMock() bot._loop.process_direct = AsyncMock()
bot._loop.consolidator.maybe_consolidate_by_tokens = AsyncMock()
snapshot = await bot.sessions.ingest( snapshot = await bot.sessions.ingest(
"sdk:history", "sdk:history",
@@ -1442,7 +1441,6 @@ async def test_sessions_ingest_imports_transcript_without_running_model(tmp_path
assert snapshot.messages[0]["source"] == "longmemeval" assert snapshot.messages[0]["source"] == "longmemeval"
assert snapshot.messages[1]["source"] == "longmemeval" assert snapshot.messages[1]["source"] == "longmemeval"
bot._loop.process_direct.assert_not_called() bot._loop.process_direct.assert_not_called()
bot._loop.consolidator.maybe_consolidate_by_tokens.assert_not_called()
reloaded = bot.sessions.get("sdk:history") reloaded = bot.sessions.get("sdk:history")
assert reloaded is not None assert reloaded is not None
@@ -1635,12 +1633,13 @@ async def test_runtime_helpers_expose_model_workspace_and_compact(tmp_path):
runtime = bot._loop.llm_runtime() runtime = bot._loop.llm_runtime()
bot._loop.runtime_for_session = MagicMock(return_value=runtime) # type: ignore[method-assign] bot._loop.runtime_for_session = MagicMock(return_value=runtime) # type: ignore[method-assign]
bot._loop.consolidator.maybe_consolidate_by_tokens = AsyncMock() compact_session = AsyncMock()
bot._loop.consolidator.compact_idle_session = compact_session
snapshot = await bot.runtime.compact_session("sdk:history") snapshot = await bot.runtime.compact_session("sdk:history")
assert snapshot.key == "sdk:history" assert snapshot.key == "sdk:history"
assert ( compact_session.assert_awaited_once_with(
bot._loop.consolidator.maybe_consolidate_by_tokens.await_args.kwargs["runtime"] "sdk:history",
is runtime runtime=runtime,
) )
assert bot.runtime.model == bot._loop.model assert bot.runtime.model == bot._loop.model
assert bot.runtime.workspace == tmp_path assert bot.runtime.workspace == tmp_path
-1
View File
@@ -714,7 +714,6 @@ export interface SettingsPayload {
heartbeat: { heartbeat: {
enabled: boolean; enabled: boolean;
interval_s: number; interval_s: number;
keep_recent_messages: number;
}; };
dream: { dream: {
schedule: string; schedule: string;
-3
View File
@@ -131,7 +131,6 @@ function baseSettingsPayload() {
heartbeat: { heartbeat: {
enabled: true, enabled: true,
interval_s: 1800, interval_s: 1800,
keep_recent_messages: 8,
}, },
dream: { dream: {
schedule: "every 2h", schedule: "every 2h",
@@ -2470,7 +2469,6 @@ describe("App layout", () => {
heartbeat: { heartbeat: {
enabled: true, enabled: true,
interval_s: 1800, interval_s: 1800,
keep_recent_messages: 8,
}, },
dream: { dream: {
schedule: "every 2h", schedule: "every 2h",
@@ -2960,7 +2958,6 @@ describe("App layout", () => {
heartbeat: { heartbeat: {
enabled: true, enabled: true,
interval_s: 1800, interval_s: 1800,
keep_recent_messages: 8,
}, },
dream: { dream: {
schedule: "every 2h", schedule: "every 2h",
-1
View File
@@ -91,7 +91,6 @@ export function settingsPayload(): SettingsPayload {
heartbeat: { heartbeat: {
enabled: true, enabled: true,
interval_s: 1800, interval_s: 1800,
keep_recent_messages: 8,
}, },
dream: { dream: {
schedule: "every 2h", schedule: "every 2h",
-1
View File
@@ -368,7 +368,6 @@ function modelSettings(model: string, provider: string): SettingsPayload {
heartbeat: { heartbeat: {
enabled: true, enabled: true,
interval_s: 1800, interval_s: 1800,
keep_recent_messages: 8,
}, },
dream: { dream: {
schedule: "every 2h", schedule: "every 2h",