diff --git a/nanobot/agent/context.py b/nanobot/agent/context.py index 5d49f25d1..6d280a431 100644 --- a/nanobot/agent/context.py +++ b/nanobot/agent/context.py @@ -235,7 +235,7 @@ class ContextBuilder: def build_messages( self, history: list[dict[str, Any]], - current_message: str, + current_message: str | None, *, media: list[str] | None = None, channel: str | None = None, @@ -259,6 +259,8 @@ class ContextBuilder: workspace=workspace, include_memory=include_memory, ) + if current_message is None: + return messages current = messages[-1] if len(messages) < 2 or messages[-2].get("role") != current.get("role"): return messages diff --git a/nanobot/agent/context_governance.py b/nanobot/agent/context_governance.py index 086f4e723..afcc69f82 100644 --- a/nanobot/agent/context_governance.py +++ b/nanobot/agent/context_governance.py @@ -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. -It may return copied messages or persisted-result placeholders, but it must not -mutate an existing session history list in place. +This module owns model-facing message shaping, request pressure, H/delta +compaction state, and tool-result content normalization. It may return copied +messages or persisted-result placeholders, but it must not mutate an existing +session history list in place. """ 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 typing import TYPE_CHECKING, Any, cast 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 ( estimate_message_tokens, estimate_prompt_tokens_chain, @@ -27,6 +51,16 @@ if TYPE_CHECKING: from nanobot.agent.tools.registry import ToolRegistry 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 # read_file is the recovery path for persisted results; exempting it prevents persist->read->persist loops. TOOL_RESULT_OFFLOAD_EXEMPT_TOOLS = frozenset({"read_file"}) @@ -85,8 +119,204 @@ class ContextGovernanceConfig: 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: - """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( self, @@ -115,17 +345,31 @@ class ContextGovernor: ) updated = self.drop_orphan_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: - return updated + return messages budget = self.input_budget(config) estimated, source = estimate_prompt_tokens_chain( config.provider, config.model, - updated, + messages, tool_definitions, ) if budget > 0 and estimated <= budget: - return updated + return messages raise ContextWindowExceededError( session_key=config.session_key, estimated_tokens=estimated, @@ -133,6 +377,41 @@ class ContextGovernor: 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( self, config: ContextGovernanceConfig, @@ -144,27 +423,15 @@ class ContextGovernor: request_context_tokens: int | None = None, ) -> tuple[list[dict[str, Any]], bool]: """Fit the request when its measured or estimated input is pressured.""" - if not config.context_window_tokens: - return messages, False - budget = self.input_budget(config) - if ( - request_context_tokens is None - and usage_matches_messages - and usage is not None - and usage.context_tokens is not None - ): - pressured = budget <= 0 or usage.context_tokens >= budget - else: - estimated, _ = estimate_prompt_tokens_chain( - config.provider, - config.model, - messages, - tool_definitions, - ) - if request_context_tokens is not None: - estimated = max(estimated, request_context_tokens) - pressured = budget <= 0 or estimated >= budget - if not pressured: + pressure = self.request_pressure( + config, + messages, + usage, + usage_matches_messages=usage_matches_messages, + tool_definitions=tool_definitions, + request_context_tokens=request_context_tokens, + ) + if pressure is None: return messages, False return self.fit_to_budget( config, @@ -172,6 +439,201 @@ class ContextGovernor: tool_definitions=tool_definitions, ), 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 def input_budget(config: ContextGovernanceConfig) -> int: if not config.context_window_tokens: @@ -424,14 +886,15 @@ class ContextGovernor: if budget <= 0: return messages - estimate, _ = estimate_prompt_tokens_chain( - config.provider, - config.model, - messages, - tool_definitions, - ) - if not force and estimate <= budget: - return messages + if not force: + estimate, _ = estimate_prompt_tokens_chain( + config.provider, + config.model, + messages, + tool_definitions, + ) + if estimate <= budget: + return messages system_messages = [dict(msg) for msg in messages if msg.get("role") == "system"] non_system = [dict(msg) for msg in messages if msg.get("role") != "system"] diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index f10434f1d..0131adbd6 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -13,6 +13,7 @@ import weakref from collections.abc import Coroutine, Iterable, Mapping from contextlib import AbstractContextManager, ExitStack, nullcontext, suppress from dataclasses import dataclass, field +from datetime import datetime from enum import Enum, auto from functools import partial from pathlib import Path @@ -93,7 +94,11 @@ from nanobot.session.recovery import ( restore_pending_interruption, 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.utils.cancellation import task_is_cancelling from nanobot.utils.document import reference_non_image_attachments @@ -161,6 +166,8 @@ class TurnContext: pending_queue: asyncio.Queue[InboundMessage] | None = None pending_summary: SessionSummary | None = None + summary_checkpoint: SessionSummaryCheckpoint | None = None + provider_compaction_applied: bool = False ephemeral: bool = False run_extra_hooks_for_ephemeral: bool = False @@ -923,19 +930,6 @@ class AgentLoop: return 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( self, transcript_input: TranscriptInput, @@ -1186,6 +1180,26 @@ class AgentLoop: provider_retry_mode=self.provider_retry_mode, retry_wait_callback=on_retry_wait, 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, terminal_injection_callback=_wait_for_pending, # 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: await ctx.on_runtime_admitted(runtime) 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( session, ctx.session_key, @@ -1899,11 +1907,7 @@ class AgentLoop: session = ctx.require_session() is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent" - _hist_kwargs: dict[str, Any] = { - "max_tokens": self._replay_token_budget(runtime), - "extend_to_user": is_subagent, - } - ctx.history = session.get_history(**_hist_kwargs) + ctx.history = session.get_history(extend_to_user=is_subagent) stored_state = session.provider_state subagent_followup_persisted = False if is_subagent: @@ -2021,6 +2025,8 @@ class AgentLoop: ) ctx.final_content = result.final_content 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 if ( ctx.kind is TurnKind.USER @@ -2034,7 +2040,6 @@ class AgentLoop: await turn_continuation.maybe_continue_turn(ctx) async def _persist_turn(self, ctx: TurnContext) -> None: - runtime = ctx.require_runtime() session = ctx.require_session() turn_continuation.prepare_save_boundary(ctx) @@ -2060,15 +2065,18 @@ class AgentLoop: self._save_turn( session, ctx.all_messages, ctx.save_skip, 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) - if not ctx.ephemeral: - self.schedule_background( - self.consolidator.maybe_consolidate_by_tokens( - session, - runtime=runtime, - ) - ) self._clear_pending_user_turn(session) self._clear_runtime_checkpoint(session) self.sessions.save(session) @@ -2142,6 +2150,55 @@ class AgentLoop: 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( self, session: Session, @@ -2149,10 +2206,10 @@ class AgentLoop: skip: int, *, turn_latency_ms: int | None = None, + summary_checkpoint: SessionSummaryCheckpoint | None = None, + input_persisted_early: bool = False, ) -> None: - """Save new-turn messages into session, truncating large tool results.""" - from datetime import datetime - + """Commit new-turn messages and an optional summary boundary.""" declared_tool_call_ids = { str(tc["id"]) for m in session.messages @@ -2169,8 +2226,30 @@ class AgentLoop: } last_assistant_idx: int | None = None saved_followup_ids: set[str] = set() - for m in messages[skip:]: - entry = dict(m) + checkpoint_boundary = self._validated_checkpoint_boundary( + 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_ids = ( [followup_id_value] @@ -2249,6 +2328,8 @@ class AgentLoop: for tc in (cast(dict[str, Any], tc_value),) 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: session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms) if saved_followup_ids: diff --git a/nanobot/agent/memory.py b/nanobot/agent/memory.py index e23a4c286..51a726dad 100644 --- a/nanobot/agent/memory.py +++ b/nanobot/agent/memory.py @@ -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 # 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 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.session.manager import ( 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 @@ -780,6 +781,20 @@ class MemoryArchiver: ) -> str: """Persist the failed chunk and return a bounded replacement checkpoint.""" 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) if not previous_summary: return truncate_text_to_tokens(raw, token_limit) @@ -806,35 +821,94 @@ class MemoryArchiver: async def archive( self, - messages: list[dict[str, Any]], + source_messages: list[dict[str, Any]], *, runtime: LLMRuntime, session_key: str, - request_messages: list[dict[str, Any]], + history: list[dict[str, Any]], request_tools: list[dict[str, Any]], previous_summary: str | None = None, + input_token_budget: int | None = None, + fallback_max_tokens: int | None = None, + provider_state: ProviderConversationState | None = None, ) -> str | None: - """Execute a prepared archive request and persist its result.""" - if not messages: + """Append the archive prompt to H and persist its summary.""" + if not source_messages: return None def raw_fallback() -> str: return self._raw_checkpoint( - messages, + source_messages, session_key=session_key, 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: with llm_usage_source("dream"): response = await runtime.provider.chat_with_retry( model=runtime.model, messages=request_messages, - tools=request_tools, + tools=call_tools, temperature=runtime.generation.temperature, max_tokens=runtime.generation.max_tokens, reasoning_effort=runtime.generation.reasoning_effort, + provider_context=provider_context, ) except Exception: 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 - 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( messages, session_key=session.key, previous_summary=previous_summary, 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( key=session.key, messages=list(session.messages[:archive_end]), @@ -908,47 +979,37 @@ class MemoryArchiver: "Memory archive cannot replay the full chunk for {}; raw-dumping", session.key, ) - return raw_fallback() - prompt = render_template("agent/consolidator_archive.md", strip=True) + return self._raw_checkpoint( + 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 workspace: Path | None = None if self._resolve_prompt_context is not None: channel, workspace = self._resolve_prompt_context(session) - request_messages = self._build_messages( + history_messages = self._build_messages( history=history, - current_message=prompt, + current_message=None, channel=channel, session_summary=session_summary, workspace=workspace, ) 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( messages, runtime=runtime, session_key=session.key, - request_messages=request_messages, + history=history_messages, request_tools=tools, previous_summary=previous_summary, + input_token_budget=input_token_budget, ) 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 @@ -978,22 +1039,73 @@ class Consolidator: """Return the shared consolidation lock for one session.""" return self._locks.setdefault(session_key, asyncio.Lock()) - def pick_consolidation_boundary( + async def summarize_transcript( self, - session: Session, - ) -> int | None: - """Return the fixed user-led boundary before the recent replay tail.""" - if not session.messages: + accepted_messages: list[dict[str, Any]], + previous_summary: str | None, + *, + 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 - boundary = max(0, len(session.messages) - MIN_COMPACTED_REPLAY_MESSAGES) - while boundary > 0 and session.messages[boundary].get("role") != "user": - boundary -= 1 - if ( - boundary <= session.last_archived - or session.messages[boundary].get("role") != "user" - ): + + max_output_tokens = max(0, runtime.generation.max_tokens) + input_token_budget = runtime.context_window_tokens - max_output_tokens + checkpoint_tokens = min( + max_output_tokens, + 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 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 def _full_replay_history( @@ -1058,7 +1170,7 @@ class Consolidator: archive_end: int, runtime: LLMRuntime, ) -> str | None: - """Compatibility wrapper for the extracted MemoryArchiver.""" + """Archive one captured session range through the shared Memory path.""" return await self.archiver.archive_session( session, archive_end=archive_end, @@ -1066,78 +1178,6 @@ class Consolidator: 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( self, session_key: str, diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index a19a9d9ae..c658d448b 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -16,8 +16,13 @@ from loguru import logger from nanobot.agent.context import TranscriptInput from nanobot.agent.context_governance import ( + ContextCompactionState, ContextGovernanceConfig, ContextGovernor, + HistoryConsolidator, + ModelRequestState, + ProviderCompactionConsolidator, + TranscriptBuilder, ) from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext from nanobot.agent.tools.execution import execute_tool_calls @@ -32,20 +37,10 @@ from nanobot.providers.base import ( LLMProvider, 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.recovery import PENDING_FOLLOWUP_ID_KEY +from nanobot.providers.conversation_state import ProviderConversationStateController +from nanobot.session.summary import SessionSummaryCheckpoint from nanobot.utils.helpers import ( build_assistant_message, estimate_message_tokens, @@ -67,7 +62,6 @@ ContinuationCallback = Callable[[], str | None] RetryWaitCallback = Callable[[str], Awaitable[None]] CheckpointCallback = Callable[[dict[str, Any]], Awaitable[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." _ARREARAGE_ERROR_MESSAGE = ( @@ -113,6 +107,8 @@ class AgentRunSpec: provider_retry_mode: str = "standard" retry_wait_callback: RetryWaitCallback | None = None checkpoint_callback: CheckpointCallback | None = None + consolidate_history: HistoryConsolidator | None = None + consolidate_provider_compaction: ProviderCompactionConsolidator | None = None injection_callback: InjectionCallback | None = None terminal_injection_callback: InjectionCallback | 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. pending_stream_content: str | None = None provider_state: ProviderConversationState | None = field(default=None, repr=False) - - -@dataclass(slots=True) -class _ModelRequestState: - """Per-run state used to govern the next provider request.""" - - config: ContextGovernanceConfig - conversation: ProviderConversationStateController - usage: LLMUsage | None = None - messages: list[dict[str, Any]] | None = None - tool_definitions: list[dict[str, Any]] | None = None + summary_checkpoint: SessionSummaryCheckpoint | None = field(default=None, repr=False) + provider_compaction_applied: bool = field(default=False, repr=False) class AgentRunner: @@ -157,118 +144,12 @@ class AgentRunner: self.context_governor = ContextGovernor() @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( - cls, messages: list[dict[str, Any]], injections: list[dict[str, Any]], ) -> None: - """Append injected user messages while preserving role alternation.""" - for injection in 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) + """Append injected messages without rewriting the raw transcript.""" + messages.extend(injections) async def _try_drain_injections( self, @@ -425,7 +306,7 @@ class AgentRunner: async def run(self, spec: AgentRunSpec) -> AgentRunResult: hook = spec.hook or AgentHook() - messages = self._initial_transcript(spec) + messages, compaction = self._initial_transcript_and_compaction(spec) context = AgentRunHookContext(messages=deepcopy(messages)) llm_usage_source_token = bind_llm_usage_source( spec.llm_usage_source or source_from_session_key(spec.session_key) @@ -433,7 +314,7 @@ class AgentRunner: try: 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: context.messages = deepcopy(messages) context.stop_reason = "cancelled" @@ -478,23 +359,35 @@ class AgentRunner: reset_llm_usage_source(llm_usage_source_token) @staticmethod - def _initial_transcript(spec: AgentRunSpec) -> list[dict[str, Any]]: - """Resolve exactly one supported source for the initial model transcript.""" - if spec.transcript_input is not None: + def _initial_transcript_and_compaction( + spec: AgentRunSpec, + ) -> 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: 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") - 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: 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( self, spec: AgentRunSpec, hook: AgentHook, messages: list[dict[str, Any]], + compaction: ContextCompactionState | None, ) -> AgentRunResult: final_content: str | None = None tools_used: list[str] = [] @@ -530,9 +423,10 @@ class AgentRunner: context_block_limit=spec.context_block_limit, max_tokens=spec.runtime.generation.max_tokens, ) - request_state = _ModelRequestState( + request_state = ModelRequestState( config=governance_config, conversation=conversation_state, + compaction=compaction, ) for iteration in range(spec.max_iterations): @@ -542,9 +436,15 @@ class AgentRunner: session_key=spec.session_key, ) 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( spec, - messages, + request_messages, hook, context, request_state=request_state, @@ -553,6 +453,11 @@ class AgentRunner: assert request_state.messages is not None messages_for_model = request_state.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.tool_calls = list(response.tool_calls) @@ -634,7 +539,7 @@ class AgentRunner: messages.append(tool_message) completed_tool_results.append(tool_message) checkpoint_model_messages = ( - self.context_governor.prepare_for_model( + self.context_governor.prepare_messages_for_model( governance_config, messages, ) @@ -920,6 +825,12 @@ class AgentRunner: had_injections=had_injections, pending_stream_content=pending_stream_content, 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( @@ -942,60 +853,6 @@ class AgentRunner: kwargs["reasoning_effort"] = generation.reasoning_effort 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( self, spec: AgentRunSpec, @@ -1003,13 +860,13 @@ class AgentRunner: hook: AgentHook, context: AgentHookContext, *, - request_state: _ModelRequestState, + request_state: ModelRequestState, malformed_retry: bool = False, transcript: list[dict[str, Any]] | None, ) -> LLMResponse: timeout_s = self._resolve_llm_timeout_s(spec) tool_definitions = spec.tools.get_definitions() - messages, provider_context = self._prepare_model_request( + messages, provider_context = await self.context_governor.prepare_request( request_state, messages, tool_definitions=tool_definitions, @@ -1169,6 +1026,12 @@ class AgentRunner: response.ttft_ms = max(0, round((first_output_at - request_started_at) * 1000)) if generation_elapsed_s > 0: 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 # hosted calls after the provider returns its final error response. if response.finish_reason == "error": @@ -1283,7 +1146,7 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], *, - request_state: _ModelRequestState, + request_state: ModelRequestState, transcript: list[dict[str, Any]], ) -> LLMResponse: retry_messages = self._finalization_retry_messages(messages) @@ -1313,14 +1176,21 @@ class AgentRunner: messages: list[dict[str, Any]], usage: LLMUsage | None, *, - request_state: _ModelRequestState, + request_state: ModelRequestState, ) -> 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: response = await self._request_no_tools( spec, retry_messages, request_state=request_state, + transcript=messages if compaction is not None else None, ) except Exception: logger.exception( @@ -1358,10 +1228,10 @@ class AgentRunner: spec: AgentRunSpec, messages: list[dict[str, Any]], *, - request_state: _ModelRequestState, + request_state: ModelRequestState, transcript: list[dict[str, Any]] | None = None, ) -> LLMResponse: - messages, provider_context = self._prepare_model_request( + messages, provider_context = await self.context_governor.prepare_request( request_state, messages, tool_definitions=None, @@ -1389,6 +1259,12 @@ class AgentRunner: finish_reason="error", 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 @staticmethod @@ -1453,7 +1329,7 @@ class AgentRunner: def _record_request_usage( self, spec: AgentRunSpec, - state: _ModelRequestState, + state: ModelRequestState, response: LLMResponse, ) -> LLMUsage | None: assert state.messages is not None diff --git a/nanobot/cli/gateway_runtime.py b/nanobot/cli/gateway_runtime.py index 38360d7df..295c99ce7 100644 --- a/nanobot/cli/gateway_runtime.py +++ b/nanobot/cli/gateway_runtime.py @@ -662,11 +662,6 @@ def _run_gateway( if isinstance(message_tool, MessageTool) and suppress_token is not None: 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: return diff --git a/nanobot/config/schema.py b/nanobot/config/schema.py index c32d0fca3..8438352f4 100644 --- a/nanobot/config/schema.py +++ b/nanobot/config/schema.py @@ -329,7 +329,6 @@ class HeartbeatConfig(Base): enabled: bool = True interval_s: int = 30 * 60 # 30 minutes - keep_recent_messages: int = 8 class ApiConfig(Base): diff --git a/nanobot/providers/azure_openai_provider.py b/nanobot/providers/azure_openai_provider.py index cf851fb05..deb2aca91 100644 --- a/nanobot/providers/azure_openai_provider.py +++ b/nanobot/providers/azure_openai_provider.py @@ -34,6 +34,7 @@ from nanobot.providers.base import ( ) from nanobot.providers.openai_responses import ( ResponsesStreamCapture, + build_responses_compaction_state, build_responses_state, consume_sdk_stream, convert_tools, @@ -410,6 +411,16 @@ class AzureOpenAIProvider(LLMProvider): output_items=capture.output_items, 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 except Exception as e: return self._handle_error(e) diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index 0b9b9ef5e..f4a89ffe4 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -31,6 +31,7 @@ RETRY_AFTER_BUFFER = 1 RetryEventCallback = Callable[[str], Awaitable[None]] LLMCallObserver = Callable[["LLMCallRecord"], None] +ProviderCompactionScope = Literal["prior_context", "current_request"] def resolve_stream_idle_timeout_s( @@ -563,6 +564,22 @@ class LLMResponse: reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc. thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking 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 # continuation independently of the final fallback error's retry policy. preserve_provider_state_on_error: bool | None = field(default=None, repr=False) diff --git a/nanobot/providers/conversation_state.py b/nanobot/providers/conversation_state.py index 7c70a1848..c9ea247f1 100644 --- a/nanobot/providers/conversation_state.py +++ b/nanobot/providers/conversation_state.py @@ -101,6 +101,12 @@ class ProviderConversationStateController: ) 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( self, messages: list[dict[str, Any]], diff --git a/nanobot/providers/openai_codex_provider.py b/nanobot/providers/openai_codex_provider.py index 42b8e360d..f2cfb0689 100644 --- a/nanobot/providers/openai_codex_provider.py +++ b/nanobot/providers/openai_codex_provider.py @@ -31,6 +31,7 @@ from nanobot.providers.oauth_model_catalog import ( ) from nanobot.providers.openai_responses import ( ResponsesStreamCapture, + build_responses_compaction_state, build_responses_state, consume_sse_with_reasoning, convert_tools, @@ -137,6 +138,8 @@ class OpenAICodexProvider(LLMProvider): body.update(self._extra_body) stage = "oauth_token" + native_compaction_applied = False + native_compaction_state: ProviderConversationState | None = None try: token = await asyncio.to_thread(get_codex_token, proxy=self.proxy) 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 ): stage = "codex_compaction" + history_items = responses_state_items(sanitized_state) or [] + delta_items = input_items[len(history_items):] compact_body = { **body, - "input": [*input_items, {"type": "compaction_trigger"}], + "input": [*history_items, {"type": "compaction_trigger"}], } try: compact_result = await _send(compact_body, emit_deltas=False) @@ -205,9 +210,16 @@ class OpenAICodexProvider(LLMProvider): }: raise RuntimeError("Codex compaction returned no compaction item") body["input"] = [ - *_retained_compaction_messages(input_items), + *_retained_compaction_messages(history_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: if is_compaction_compatibility_error(compact_error): self._native_compaction_available = False @@ -220,7 +232,14 @@ class OpenAICodexProvider(LLMProvider): ) 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: response = _codex_error_response(e) exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__ diff --git a/nanobot/providers/openai_compat_provider.py b/nanobot/providers/openai_compat_provider.py index 54a9bb0b9..1a7cca084 100644 --- a/nanobot/providers/openai_compat_provider.py +++ b/nanobot/providers/openai_compat_provider.py @@ -36,6 +36,7 @@ from nanobot.providers.base import ( ) from nanobot.providers.openai_responses import ( ResponsesStreamCapture, + build_responses_compaction_state, build_responses_state, consume_sdk_stream, convert_tools, @@ -2049,6 +2050,18 @@ class OpenAICompatProvider(LLMProvider): output_items=capture.output_items, 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 except Exception as responses_error: if self._spec and self._spec.name == "github_copilot": diff --git a/nanobot/providers/openai_responses/__init__.py b/nanobot/providers/openai_responses/__init__.py index dfc371972..9b95c418f 100644 --- a/nanobot/providers/openai_responses/__init__.py +++ b/nanobot/providers/openai_responses/__init__.py @@ -18,6 +18,7 @@ from nanobot.providers.openai_responses.parsing import ( parse_response_output, ) from nanobot.providers.openai_responses.state import ( + build_responses_compaction_state, build_responses_state, is_compaction_compatibility_error, prepare_responses_input, @@ -40,6 +41,7 @@ __all__ = [ "is_replayable_finish_reason", "map_finish_reason", "parse_response_output", + "build_responses_compaction_state", "build_responses_state", "is_compaction_compatibility_error", "prepare_responses_input", diff --git a/nanobot/providers/openai_responses/parsing.py b/nanobot/providers/openai_responses/parsing.py index b5676a70a..aced8b7d6 100644 --- a/nanobot/providers/openai_responses/parsing.py +++ b/nanobot/providers/openai_responses/parsing.py @@ -11,7 +11,10 @@ import httpx from loguru import logger 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 = { "completed": "stop", @@ -655,6 +658,14 @@ def parse_response_output( output_items=output, 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 diff --git a/nanobot/providers/openai_responses/state.py b/nanobot/providers/openai_responses/state.py index a424c66d8..2941a0b90 100644 --- a/nanobot/providers/openai_responses/state.py +++ b/nanobot/providers/openai_responses/state.py @@ -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( state: ProviderConversationState, ) -> list[dict[str, Any]] | None: diff --git a/nanobot/sdk/clients.py b/nanobot/sdk/clients.py index 08ff1b9db..59a534b48 100644 --- a/nanobot/sdk/clients.py +++ b/nanobot/sdk/clients.py @@ -209,11 +209,11 @@ class RuntimeClient: return self._loop.runtime_events.subscribe(handler, SessionTurnPersisted) 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) runtime = self._loop.runtime_for_session(session) - await self._loop.consolidator.maybe_consolidate_by_tokens( - session, + await self._loop.consolidator.compact_idle_session( + session_key, runtime=runtime, ) return snapshot_from_session(self._loop.sessions.get_or_create(session_key)) diff --git a/nanobot/session/manager.py b/nanobot/session/manager.py index 411efc702..5c39f680d 100644 --- a/nanobot/session/manager.py +++ b/nanobot/session/manager.py @@ -27,7 +27,9 @@ from nanobot.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, 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.summary import SUMMARY_CONTINUATION_TEXT from nanobot.utils.helpers import ( content_with_media_breadcrumbs, ensure_dir, @@ -262,12 +264,6 @@ def _metadata_title(metadata: object) -> str: return strip_think(title) -@dataclass -class RetentionResult: - dropped: list[dict[str, Any]] - already_consolidated_count: int - - @dataclass(frozen=True) class SessionPolicy: """Runtime rules that do not belong in durable session data.""" @@ -286,9 +282,7 @@ class Session: created_at: datetime = field(default_factory=datetime.now) updated_at: datetime = field(default_factory=datetime.now) metadata: dict[str, Any] = field(default_factory=dict) - # Legacy storage name for the Memory ingestion watermark. New code should - # use ``last_archived`` so this progress is not confused with model-context - # compaction. Keep the field while persisted sessions and SDK callers migrate. + # Keep the legacy storage name while persisted sessions and SDK callers migrate. last_consolidated: int = 0 provider_state: ProviderConversationState | None = field(default=None, repr=False) policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False) @@ -309,7 +303,7 @@ class Session: @property 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 @last_archived.setter @@ -337,14 +331,17 @@ class Session: ) -> list[dict[str, Any]]: """Return recent replayable messages for LLM input. - A positive ``max_messages`` applies an explicit caller-owned count - limit. The normal model path relies on ``max_tokens`` instead. + A committed in-turn checkpoint replaces its old prefix with the stored + 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 - if replay_start: - # ``last_archived`` is archive progress, not a replay boundary. - # Keep a small raw suffix for continuity, extending back to the user - # that started an assistant/tool sequence when necessary. + resumes_from_checkpoint = ( + replay_start < len(self.messages) + and is_hidden_history_message(self.messages[replay_start]) + and self.messages[replay_start].get("content") == SUMMARY_CONTINUATION_TEXT + ) + if replay_start and not resumes_from_checkpoint: recent_start = recent_message_start_index( self.messages, MIN_COMPACTED_REPLAY_MESSAGES, @@ -485,110 +482,6 @@ class Session: self.updated_at = datetime.now() 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): key: str created_at: str | None @@ -2010,7 +1903,7 @@ class SessionManager: user_index = 0 found_target = False 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: found_target = True break diff --git a/nanobot/session/summary.py b/nanobot/session/summary.py index 42bada0ed..2c990c13e 100644 --- a/nanobot/session/summary.py +++ b/nanobot/session/summary.py @@ -3,15 +3,27 @@ from __future__ import annotations from collections.abc import Mapping +from dataclasses import dataclass from datetime import datetime from typing import TypedDict, cast +SUMMARY_CONTINUATION_TEXT = ( + "Continue the active task from the working-memory checkpoint above." +) class SessionSummary(TypedDict): text: 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( metadata: Mapping[str, object] | None, *, diff --git a/nanobot/webui/settings_system.py b/nanobot/webui/settings_system.py index fa4a5f7b2..8900cbc9e 100644 --- a/nanobot/webui/settings_system.py +++ b/nanobot/webui/settings_system.py @@ -114,7 +114,6 @@ def system_settings_payload( "heartbeat": { "enabled": config.gateway.heartbeat.enabled, "interval_s": config.gateway.heartbeat.interval_s, - "keep_recent_messages": config.gateway.heartbeat.keep_recent_messages, }, "dream": { "schedule": defaults.dream.describe_schedule(), diff --git a/tests/agent/test_auto_compact.py b/tests/agent/test_auto_compact.py index 96e2e97ff..c1cfc89c6 100644 --- a/tests/agent/test_auto_compact.py +++ b/tests/agent/test_auto_compact.py @@ -8,7 +8,6 @@ from unittest.mock import AsyncMock, MagicMock import pytest from nanobot.agent.loop import AgentLoop -from nanobot.agent.runner import AgentRunResult from nanobot.agent.tools.registry import ToolRegistry from nanobot.bus.events import InboundMessage from nanobot.bus.queue import MessageBus @@ -229,36 +228,6 @@ class TestAgentLoopTTLParam: loop = _make_loop(tmp_path, session_ttl_minutes=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: """Test the _archive method.""" diff --git a/tests/agent/test_consolidator.py b/tests/agent/test_consolidator.py index 712dc2568..6f5d11829 100644 --- a/tests/agent/test_consolidator.py +++ b/tests/agent/test_consolidator.py @@ -55,9 +55,7 @@ def runtime(mock_provider): def consolidator(store): sessions = MagicMock() sessions.save = MagicMock() - # When maybe_consolidate_by_tokens refreshes the session reference via - # 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. + # Store sessions by key so refreshes observe the same test object. _session_cache: dict[str, MagicMock] = {} sessions.get_or_create = MagicMock(side_effect=lambda key: _session_cache.get(key, MagicMock())) sessions._session_cache = _session_cache @@ -93,11 +91,17 @@ def _provider_state() -> ProviderConversationState: def _build_test_messages(**kwargs): - return [ - {"role": "system", "content": "system prompt"}, + system = "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"], - {"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( @@ -112,15 +116,85 @@ async def _archive( messages, runtime=runtime, session_key=session_key, - request_messages=_build_test_messages( - history=messages, - current_message="consolidate", - ), + history=[ + {"role": "system", "content": "system prompt"}, + *messages, + ], request_tools=[], 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: def test_format_messages_keeps_media_only_user_turn(self): path = "/home/user/.nanobot/media/websocket/clip.mp4" @@ -379,32 +453,7 @@ class TestConsolidatorArchiveErrorHandling: consolidator.store.raw_archive.assert_not_called() -class TestConsolidatorTokenBudget: - 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) - +class TestConsolidatorPromptEstimate: async def test_estimate_uses_full_unarchived_tail(self, consolidator, runtime): """Consolidation pressure must account for the full unarchived tail.""" session = Session(key="test:full-tail") @@ -443,129 +492,6 @@ class TestConsolidatorTokenBudget: assert len(captured["history"]) == 8 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: """Idle compaction tests.""" @@ -1347,110 +1273,6 @@ class TestCompactIdleSession: 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: """raw_archive() must cap entry size to avoid bloating history.jsonl.""" diff --git a/tests/agent/test_dream.py b/tests/agent/test_dream.py index 15f1080ee..ea11c01e4 100644 --- a/tests/agent/test_dream.py +++ b/tests/agent/test_dream.py @@ -404,10 +404,9 @@ class TestEphemeralDirect: with ( patch("nanobot.agent.loop.SessionManager"), 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_consolidator_cls.return_value.maybe_consolidate_by_tokens = AsyncMock() loop = AgentLoop( bus=bus, provider=provider, @@ -493,20 +492,6 @@ class TestEphemeralDirect: 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): loop, store = _make_loop loop.provider.chat_with_retry.return_value = LLMResponse( @@ -701,10 +686,9 @@ class TestEphemeralHooks: with ( patch("nanobot.agent.loop.SessionManager"), 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_consolidator_cls.return_value.maybe_consolidate_by_tokens = AsyncMock() loop = AgentLoop( bus=bus, provider=provider, diff --git a/tests/agent/test_history_replay.py b/tests/agent/test_history_replay.py index e9e7ce634..b4158b88a 100644 --- a/tests/agent/test_history_replay.py +++ b/tests/agent/test_history_replay.py @@ -12,6 +12,7 @@ from nanobot.bus.events import InboundMessage from nanobot.bus.queue import MessageBus from nanobot.providers.base import LLMResponse 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: @@ -66,13 +67,12 @@ def test_explicit_message_limit_still_starts_at_user_turn() -> None: @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.provider.chat_with_retry = AsyncMock( return_value=LLMResponse(content="ok", tool_calls=[], usage=None) ) 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") 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 get_history.call_args.kwargs == { - "max_tokens": loop._replay_token_budget(loop.llm_runtime()), - "extend_to_user": False, - } + assert get_history.call_args.kwargs == {"extend_to_user": False} @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.provider.chat_with_retry = AsyncMock( return_value=LLMResponse(content="ok", tool_calls=[], usage=None) ) 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.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_text = "\n".join(str(message.get("content")) for message in sent_messages) 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) diff --git a/tests/agent/test_loop_consolidation_tokens.py b/tests/agent/test_loop_consolidation_tokens.py index 0b304d971..f443e106b 100644 --- a/tests/agent/test_loop_consolidation_tokens.py +++ b/tests/agent/test_loop_consolidation_tokens.py @@ -4,7 +4,12 @@ import pytest from nanobot.agent.loop import AgentLoop 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( @@ -14,7 +19,6 @@ def _make_loop( context_window_tokens: int, max_tokens: int = 0, ) -> AgentLoop: - from nanobot.providers.base import GenerationSettings provider = MagicMock() provider.get_default_model.return_value = "test-model" provider.generation = GenerationSettings(max_tokens=max_tokens) @@ -39,186 +43,108 @@ def _make_loop( @pytest.mark.asyncio -async def test_prompt_below_threshold_does_not_consolidate(tmp_path) -> None: - loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=200) - loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign] - - await loop.process_direct("hello", session_key="cli:test") - - 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") - ) +async def test_runner_pressure_commits_summary_and_current_delta(tmp_path) -> None: + loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=2_000) + loop.context_block_limit = 500 + loop.provider.generation = GenerationSettings(max_tokens=100) + loop.provider.can_resume_conversation_state.return_value = False loop.schedule_background = lambda coro: coro.close() # type: ignore[method-assign] session = loop.sessions.get_or_create("cli:test") session.messages = [ - {"role": role, "content": f"{role[0]}{turn}"} - for turn in range(10) + {"role": role, "content": f"old-{role}-{turn}"} + for turn in range(6) for role in ("user", "assistant") ] 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"] - system_prompt = request_messages[0]["content"] - assert "FRESH_CHECKPOINT" in system_prompt - assert all(message.get("content") != "u0" for message in request_messages) - assert loop.sessions.get_or_create("cli:test").last_archived == 12 + loop.provider.estimate_prompt_tokens.side_effect = estimate + loop.provider.chat_with_retry = AsyncMock(side_effect=[ + LLMResponse(content="Current checkpoint.", tool_calls=[]), + LLMResponse(content="done", tool_calls=[]), + ]) + result = await loop.process_direct("continue the task", session_key="cli:test") -@pytest.mark.asyncio -async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None: - loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200) - loop.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.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(), - ) + assert result.content == "done" + assert loop.provider.chat_with_retry.await_count == 2 + model_request = loop.provider.chat_with_retry.await_args_list[1].kwargs["messages"] + assert "Current checkpoint." in model_request[0]["content"] + assert model_request[1]["content"] == SUMMARY_CONTINUATION_TEXT + assert model_request[2]["content"] == "continue the task" reloaded = loop.sessions.get_or_create("cli:test") - meta = reloaded.metadata.get("_last_summary") - assert meta is not None - assert meta["text"] == "User discussed project status." - - reloaded, pending = loop.auto_compact.prepare_session(reloaded, "cli:test") - assert pending is not None - assert pending["text"] == "User discussed project status." - # _last_summary persists for restart survival. - assert "_last_summary" in reloaded.metadata + assert reloaded.messages[0]["content"] == "old-user-0" + assert reloaded.metadata["_last_summary"]["text"] == "Current checkpoint." + assert reloaded.messages[reloaded.last_archived]["content"] == ( + SUMMARY_CONTINUATION_TEXT + ) + assert [message["content"] for message in reloaded.get_history()] == [ + SUMMARY_CONTINUATION_TEXT, + "continue the task", + "done", + ] @pytest.mark.asyncio -async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> None: - loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=200) - session = loop.sessions.get_or_create("cli:test") - loop.auto_compact.prepare_session = MagicMock( - return_value=( - 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") +async def test_native_provider_compaction_commits_portable_terminal_checkpoint( + tmp_path, +) -> None: + loop = _make_loop(tmp_path, estimated_tokens=100, context_window_tokens=2_000) + session = loop.sessions.get_or_create("cli:native") session.messages = [ - {"role": role, "content": f"{role[0]}{turn}"} - for turn in range(10) - for role in ("user", "assistant") + {"role": "user", "content": "accepted history"}, + {"role": "assistant", "content": "accepted answer"}, ] loop.sessions.save(session) - call_count = [0] - def mock_estimate(_session, *, runtime): - call_count[0] += 1 - return (1000 if call_count[0] <= 1 else 80, "test") - loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign] + compacted_state = ProviderConversationState( + kind="openai_responses", + provider="openai:test", + model="test-model", + 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 "llm" in order - assert order.index("consolidate") < order.index("llm") - assert archived_session_keys == ["cli:test"] + assert result.content == "done" + summarize = loop.consolidator.summarize_provider_compaction + summarize.assert_awaited_once() + 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", + ] diff --git a/tests/agent/test_loop_image_generation_media.py b/tests/agent/test_loop_image_generation_media.py index cfcc3b2cd..f57fd9196 100644 --- a/tests/agent/test_loop_image_generation_media.py +++ b/tests/agent/test_loop_image_generation_media.py @@ -69,7 +69,6 @@ async def test_outbound_no_longer_carries_generated_media( ), 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( InboundMessage( diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index b286fb569..cc18f1540 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -425,7 +425,6 @@ class TestToolEventProgress: None, ), ) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] await loop._dispatch(InboundMessage( channel="websocket", @@ -473,7 +472,6 @@ class TestToolEventProgress: provider.chat_stream_with_retry = AsyncMock() loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5") 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( channel="whatsapp", @@ -512,7 +510,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -566,7 +563,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -611,7 +607,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -655,7 +650,6 @@ class TestToolEventProgress: _attach_webui_runtime_events(loop, bus) loop.max_iterations = 1 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( channel="websocket", @@ -747,7 +741,6 @@ class TestToolEventProgress: ) _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -815,9 +808,6 @@ class TestToolEventProgress: return "ok" 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 = 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") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -1048,7 +1037,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -1132,7 +1120,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") _attach_webui_runtime_events(loop, bus) 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( channel="websocket", @@ -1181,7 +1168,6 @@ class TestToolEventProgress: loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") _attach_webui_runtime_events(loop, bus) 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] = {} @@ -1268,7 +1254,6 @@ class TestToolEventProgress: 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.tools.get_definitions = MagicMock(return_value=[]) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] await loop._dispatch(InboundMessage( channel="slack", diff --git a/tests/agent/test_loop_runner_integration.py b/tests/agent/test_loop_runner_integration.py index af2776da3..e99b8b6cd 100644 --- a/tests/agent/test_loop_runner_integration.py +++ b/tests/agent/test_loop_runner_integration.py @@ -112,7 +112,6 @@ async def test_goal_command_can_implement_plan_from_prior_discussion(tmp_path): LLMResponse(content="done", tool_calls=[], usage=None), ]) 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.add_message("user", "Let's agree on the migration implementation.") 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), ]) 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") 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.chat_with_retry = AsyncMock(return_value=LLMResponse(content="answer", usage=None)) 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") quote = webui_quote_runtime_context({ 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), ]) 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 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), ]) 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.add_message("user", "/goal old completed request") 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.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( InboundMessage(channel="cli", sender_id="user", chat_id="test", content="first question") diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index 7059da77b..25a391c44 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -45,6 +45,10 @@ from nanobot.session.recovery import ( RUNTIME_CHECKPOINT_KEY, restore_runtime_checkpoint, ) +from nanobot.session.summary import ( + SUMMARY_CONTINUATION_TEXT, + SessionSummaryCheckpoint, +) from nanobot.session.turn_continuation import ( INTERNAL_CONTINUATION_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"] == [] +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: """Persisting a merged injected row retires every durable follow-up ID.""" loop = _mk_loop() @@ -966,7 +1024,6 @@ async def test_runtime_checkpoint_keeps_provider_state_out_of_public_metadata( @pytest.mark.asyncio async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None: 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] 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, ) -> None: 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.provider.can_resume_conversation_state.return_value = True 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, ) -> None: 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.context.build_system_prompt = MagicMock( # type: ignore[method-assign] side_effect=RuntimeError("prompt boom"), @@ -1049,7 +1104,6 @@ async def test_subagent_redelivery_does_not_duplicate_staged_provider_input( tmp_path: Path, ) -> None: 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 build_system_prompt = loop.context.build_system_prompt 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, ) -> None: 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( "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: loop = _make_full_loop(tmp_path) 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] 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) 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] 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) 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] msg = InboundMessage( @@ -1286,7 +1336,6 @@ async def test_process_message_persists_media_only_turn_without_text(tmp_path: P @pytest.mark.asyncio async def test_process_message_does_not_duplicate_early_persisted_user_message(tmp_path: Path) -> None: 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( "done", [ @@ -1319,7 +1368,6 @@ async def test_internal_continuation_queues_turn_without_fake_user_history( tmp_path: Path, ) -> None: loop = _make_full_loop(tmp_path) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] session = loop.sessions.get_or_create("feishu:c-auto") session.metadata[GOAL_STATE_KEY] = { "status": "active", @@ -1388,7 +1436,6 @@ async def test_internal_continuation_preserves_streaming_route_metadata( tmp_path: Path, ) -> None: 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.metadata[GOAL_STATE_KEY] = { "status": "active", @@ -1462,7 +1509,6 @@ async def test_websocket_internal_continuation_keeps_single_visible_run( tmp_path: Path, ) -> None: 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.metadata[GOAL_STATE_KEY] = { "status": "active", @@ -1526,7 +1572,6 @@ async def test_websocket_internal_continuation_keeps_single_visible_run( @pytest.mark.asyncio async def test_process_message_keeps_delivery_chat_for_thread_session(tmp_path: Path) -> None: 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] return_value=[ {"role": "system", "content": "system"}, @@ -1565,7 +1610,6 @@ async def test_process_message_uses_explicit_session_for_goal_context( tmp_path: Path, ) -> None: 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.metadata[GOAL_STATE_KEY] = { "status": "active", @@ -1713,7 +1757,6 @@ async def test_request_context_uses_effective_key_for_spawn_tool(tmp_path: Path) @pytest.mark.asyncio 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.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 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 loop = _make_full_loop(tmp_path) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] checkpoint_saved = asyncio.Event() @@ -1866,7 +1908,6 @@ async def test_stop_preserves_runtime_checkpoint_for_next_turn(tmp_path: Path) - @pytest.mark.asyncio async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_path: Path) -> None: 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.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.turn_id 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"] assert isinstance(initial_messages, list) 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 async def test_turn_usage_is_persisted_with_the_saved_session(tmp_path: Path) -> None: 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) 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 async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) -> None: 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): 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 async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Path) -> None: loop = _make_full_loop(tmp_path) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign] - return_value=False - ) visited: list[str] = [] for name in ( @@ -2081,7 +2110,6 @@ async def test_system_subagent_followup_uses_common_turn_lifecycle(tmp_path: Pat @pytest.mark.asyncio async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp_path: Path) -> None: loop = _make_full_loop(tmp_path) - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] async def fake_run_agent_loop(transcript_input, **_kwargs): 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 async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(tmp_path: Path) -> None: 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.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 async def test_turn_after_unanswered_user_keeps_tool_call_pairing(tmp_path: Path) -> None: 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.add_message("user", "earlier question that never got an answer") diff --git a/tests/agent/test_loop_session_policy.py b/tests/agent/test_loop_session_policy.py index 59c372e42..860d3e0c7 100644 --- a/tests/agent/test_loop_session_policy.py +++ b/tests/agent/test_loop_session_policy.py @@ -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: loop = _loop(tmp_path, ["first answer", "second answer"]) loop.context.memory.write_memory("private durable memory") - loop.consolidator.maybe_consolidate_by_tokens = AsyncMock() key = "websocket:transient-test" loop.sessions.get_or_create_transient( key, @@ -71,7 +70,6 @@ async def test_transient_session_keeps_history_without_persisting_or_durable_too "assistant", ] assert loop.sessions.read_session_file(key) is None - loop.consolidator.maybe_consolidate_by_tokens.assert_not_awaited() @pytest.mark.asyncio diff --git a/tests/agent/test_runner_core.py b/tests/agent/test_runner_core.py index af3d243ce..67ef542c4 100644 --- a/tests/agent/test_runner_core.py +++ b/tests/agent/test_runner_core.py @@ -61,7 +61,10 @@ def test_initial_transcript_is_built_from_structured_turn_input() -> None: 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) diff --git a/tests/agent/test_runner_governance.py b/tests/agent/test_runner_governance.py index fe09bc272..a3ddb4c93 100644 --- a/tests/agent/test_runner_governance.py +++ b/tests/agent/test_runner_governance.py @@ -7,13 +7,14 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from agent.runner_helpers import make_run_spec +from nanobot.agent.context import TranscriptInput from nanobot.agent.context_governance import ( BACKFILL_CONTENT, ContextGovernanceConfig, ContextGovernor, ContextWindowExceededError, ) -from nanobot.agent.runner import AgentRunSpec +from nanobot.agent.runner import AgentRunner, AgentRunSpec from nanobot.config.schema import AgentDefaults from nanobot.providers.base import ( LLMProvider, @@ -22,10 +23,23 @@ from nanobot.providers.base import ( ProviderConversationState, ToolCallRequest, ) +from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT _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( provider, tools, @@ -97,13 +111,16 @@ async def test_runner_locally_fits_oversized_initial_transcript(monkeypatch): tools = MagicMock() tools.get_definitions.return_value = [] old_content = "x" * 20_000 - monkeypatch.setattr( - "nanobot.agent.context_governance.estimate_prompt_tokens_chain", - lambda _provider, _model, messages, _tools: ( + estimate = MagicMock( + side_effect=lambda _provider, _model, messages, _tools: ( (600, "test-counter") if any(message.get("content") == old_content for message in messages) else (100, "test-counter") - ), + ) + ) + monkeypatch.setattr( + "nanobot.agent.context_governance.estimate_prompt_tokens_chain", + estimate, ) 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": "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) +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 async def test_runner_governs_messages_added_by_before_iteration_hook(monkeypatch): 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", lambda _provider, _model, messages, _tools: ( (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") ), ) @@ -491,6 +940,39 @@ def test_changed_messages_use_local_estimate_after_reported_usage(monkeypatch): 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 async def test_runner_counts_resumed_provider_state_before_dispatch(monkeypatch): 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", ) 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.messages = [ diff --git a/tests/agent/test_runner_injections.py b/tests/agent/test_runner_injections.py index c7763f9de..b0f6c3110 100644 --- a/tests/agent/test_runner_injections.py +++ b/tests/agent/test_runner_injections.py @@ -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] - assert "follow-up from the second speaker" in str(injected["content"]) + injected = [message for message in result.messages if message.get("role") == "user"][-2:] + 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"] assert "telegram | group-1 | user-b | message-2" 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 "Carol | topic-7" in str(model_messages) - assert injected["_meta"][RUNTIME_CONTEXT_MESSAGE_META]["sources"] == [ - "identity", - "identity", - ] + assert all( + message["_meta"][RUNTIME_CONTEXT_MESSAGE_META]["sources"] == ["identity"] + for message in injected + ) loop._save_turn(session, result.messages, skip=1) - persisted = [message for message in session.messages if message.get("role") == "user"][-1] - assert "telegram | group-1 | user-b | message-2" in str(persisted["content"]) - assert "telegram | group-1 | user-c | message-3" in str(persisted["content"]) - assert public_history_message(persisted)["content"] == ( - "follow-up from the second speaker\n\nanother follow-up" - ) + persisted = [message for message in session.messages if message.get("role") == "user"][-2:] + assert "telegram | group-1 | user-b | message-2" in str(persisted[0]["content"]) + assert "telegram | group-1 | user-c | message-3" in str(persisted[1]["content"]) + assert [public_history_message(message)["content"] for message in persisted] == [ + "follow-up from the second speaker", + "another follow-up", + ] @pytest.mark.asyncio @@ -835,8 +837,8 @@ async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_p @pytest.mark.asyncio -async def test_runner_merges_multiple_injected_user_messages_without_losing_media(): - """Multiple injected follow-ups should not create lossy consecutive user messages.""" +async def test_model_request_merges_injected_user_messages_without_losing_media(): + """The model copy may merge follow-ups while the raw transcript keeps each event.""" from nanobot.agent.runner import AgentRunner provider = MagicMock() @@ -895,10 +897,17 @@ async def test_runner_merges_multiple_injected_user_messages_without_losing_medi for block in injected["content"] 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: - """Merged follow-ups stay acknowledged together after a later save.""" +def test_runner_append_keeps_recovery_followups_separate() -> None: + """Each raw follow-up keeps its own recovery identity.""" from nanobot.agent.runner import AgentRunner 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"}], ) - 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.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, @@ -948,8 +959,9 @@ def test_runner_merge_preserves_runtime_markers_with_media() -> None: }, ]) - assert len(messages) == 1 - merged = messages[0] + assert len(messages) == 2 + merged = ContextGovernor._merge_adjacent_user_messages_for_model(messages)[0] + assert len(messages) == 2 assert "private first" in str(merged["content"]) assert "private second" in str(merged["content"]) persisted = { @@ -1681,15 +1693,17 @@ async def test_drain_injections_after_recoverable_tool_error(): @pytest.mark.asyncio 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.bus.events import InboundMessage provider = MagicMock() call_count = {"n": 0} + requests: list[list[dict]] = [] async def chat_with_retry(*, messages, **kwargs): call_count["n"] += 1 + requests.append(messages) if call_count["n"] == 1: return LLMResponse( content=None, @@ -1713,11 +1727,20 @@ async def test_drain_injections_on_llm_error(): runner = AgentRunner() result = await runner.run(make_run_spec(provider, - initial_messages=[ - {"role": "user", "content": "hello"}, - {"role": "assistant", "content": "previous response"}, - {"role": "user", "content": "trigger error"}, + initial_messages=None, + transcript_input=TranscriptInput( + history=[ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "previous response"}, + {"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, model="test-model", max_iterations=5, @@ -1727,11 +1750,15 @@ async def test_drain_injections_on_llm_error(): assert result.had_injections is True assert result.final_content == "recovered answer" - injected = [ - m for m in result.messages - if m.get("role") == "user" and "follow-up after LLM error" in str(m.get("content", "")) + assert "follow-up after LLM error" in str(requests[1]) + assert [ + 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 diff --git a/tests/agent/test_session_manager_history.py b/tests/agent/test_session_manager_history.py index 95de72502..ebf9b6f2e 100644 --- a/tests/agent/test_session_manager_history.py +++ b/tests/agent/test_session_manager_history.py @@ -1,10 +1,11 @@ -from nanobot.providers.base import ProviderConversationState from nanobot.runtime_context import ( RUNTIME_CONTEXT_HISTORY_META, RuntimeContextBlock, append_runtime_context, ) +from nanobot.session.history_visibility import HIDDEN_HISTORY_META from nanobot.session.manager import Session, SessionManager +from nanobot.session.summary import SUMMARY_CONTINUATION_TEXT 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" -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 --- 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"] +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): manager = SessionManager(tmp_path) 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"] -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(): session = Session(key="test:history-extend-to-user") 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_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 diff --git a/tests/agent/test_session_retention.py b/tests/agent/test_session_retention.py deleted file mode 100644 index 06e1ac216..000000000 --- a/tests/agent/test_session_retention.py +++ /dev/null @@ -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) diff --git a/tests/agent/test_unified_session.py b/tests/agent/test_unified_session.py index a54f27271..2edc55200 100644 --- a/tests/agent/test_unified_session.py +++ b/tests/agent/test_unified_session.py @@ -28,7 +28,7 @@ from nanobot.command.router import CommandContext, CommandRouter from nanobot.config.schema import AgentDefaults, Config from nanobot.providers.base import GenerationSettings 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 # --------------------------------------------------------------------------- @@ -334,116 +334,6 @@ class TestCmdNewUnifiedSession: 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 # --------------------------------------------------------------------------- diff --git a/tests/cli/test_commands.py b/tests/cli/test_commands.py index 08b1a3d4d..d628b9cca 100644 --- a/tests/cli/test_commands.py +++ b/tests/cli/test_commands.py @@ -1837,12 +1837,6 @@ def test_agent_workspace_override_wins_over_config_workspace(mock_agent_runtime, 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( "content, expected", [ @@ -2101,7 +2095,7 @@ def _patch_cli_command_runtime( 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, ) -> None: config_file = _write_instance_config(tmp_path) @@ -2119,21 +2113,9 @@ def test_heartbeat_empty_response_still_retains_recent_messages( bus.publish_outbound = AsyncMock() seen: dict[str, object] = {} - class _FakeSession: - def retain_recent_legal_suffix(self, limit: int) -> None: - seen["retained_limit"] = limit - class _FakeSessionManager: def __init__(self, _workspace: Path) -> None: - self.session = _FakeSession() - 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 + pass def list_sessions(self) -> list[dict[str, str]]: 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"))) 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( diff --git a/tests/config/test_gateway_config.py b/tests/config/test_gateway_config.py index c48602079..e8323a32d 100644 --- a/tests/config/test_gateway_config.py +++ b/tests/config/test_gateway_config.py @@ -13,3 +13,12 @@ def test_gateway_restart_mode_accepts_camel_alias(): def test_gateway_restart_mode_rejects_unknown_value(): with pytest.raises(ValueError): 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 diff --git a/tests/providers/test_openai_codex_provider.py b/tests/providers/test_openai_codex_provider.py index e7c2b05a9..687ecc884 100644 --- a/tests/providers/test_openai_codex_provider.py +++ b/tests/providers/test_openai_codex_provider.py @@ -22,7 +22,10 @@ from nanobot.providers.openai_codex_provider import ( _request_codex, _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 @@ -811,12 +814,33 @@ async def test_codex_compacts_state_at_ninety_percent_before_next_request( ) assert response.content == "done" - assert len(bodies) == 2 - assert bodies[0]["input"][-1] == {"type": "compaction_trigger"} - assert bodies[1]["input"][-1] == { + assert response.provider_compaction_applied is True + assert response.provider_compaction_state is not None + assert response.provider_compaction_scope == "prior_context" + assert responses_state_items(response.provider_compaction_state) == [{ "type": "compaction", "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( item.get("type") == "reasoning" for item in bodies[1]["input"] diff --git a/tests/providers/test_openai_responses.py b/tests/providers/test_openai_responses.py index abb4005f8..cc6158940 100644 --- a/tests/providers/test_openai_responses.py +++ b/tests/providers/test_openai_responses.py @@ -712,6 +712,45 @@ class TestParseResponseOutput: assert result.provider_state is not None 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: def test_server_compaction_prunes_superseded_prefix(self): diff --git a/tests/test_nanobot_facade.py b/tests/test_nanobot_facade.py index be6dc4a46..efb3c8dba 100644 --- a/tests/test_nanobot_facade.py +++ b/tests/test_nanobot_facade.py @@ -1412,7 +1412,6 @@ async def test_sessions_ingest_imports_transcript_without_running_model(tmp_path config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) bot._loop.process_direct = AsyncMock() - bot._loop.consolidator.maybe_consolidate_by_tokens = AsyncMock() snapshot = await bot.sessions.ingest( "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[1]["source"] == "longmemeval" bot._loop.process_direct.assert_not_called() - bot._loop.consolidator.maybe_consolidate_by_tokens.assert_not_called() reloaded = bot.sessions.get("sdk:history") 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() 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") assert snapshot.key == "sdk:history" - assert ( - bot._loop.consolidator.maybe_consolidate_by_tokens.await_args.kwargs["runtime"] - is runtime + compact_session.assert_awaited_once_with( + "sdk:history", + runtime=runtime, ) assert bot.runtime.model == bot._loop.model assert bot.runtime.workspace == tmp_path diff --git a/webui/src/lib/types.ts b/webui/src/lib/types.ts index ac20e4712..22c05d7d2 100644 --- a/webui/src/lib/types.ts +++ b/webui/src/lib/types.ts @@ -714,7 +714,6 @@ export interface SettingsPayload { heartbeat: { enabled: boolean; interval_s: number; - keep_recent_messages: number; }; dream: { schedule: string; diff --git a/webui/src/tests/app-layout.test.tsx b/webui/src/tests/app-layout.test.tsx index 706485a4a..d73d62b83 100644 --- a/webui/src/tests/app-layout.test.tsx +++ b/webui/src/tests/app-layout.test.tsx @@ -131,7 +131,6 @@ function baseSettingsPayload() { heartbeat: { enabled: true, interval_s: 1800, - keep_recent_messages: 8, }, dream: { schedule: "every 2h", @@ -2470,7 +2469,6 @@ describe("App layout", () => { heartbeat: { enabled: true, interval_s: 1800, - keep_recent_messages: 8, }, dream: { schedule: "every 2h", @@ -2960,7 +2958,6 @@ describe("App layout", () => { heartbeat: { enabled: true, interval_s: 1800, - keep_recent_messages: 8, }, dream: { schedule: "every 2h", diff --git a/webui/src/tests/settings-test-utils.tsx b/webui/src/tests/settings-test-utils.tsx index 6dec851f9..fb4c8741d 100644 --- a/webui/src/tests/settings-test-utils.tsx +++ b/webui/src/tests/settings-test-utils.tsx @@ -91,7 +91,6 @@ export function settingsPayload(): SettingsPayload { heartbeat: { enabled: true, interval_s: 1800, - keep_recent_messages: 8, }, dream: { schedule: "every 2h", diff --git a/webui/src/tests/thread-shell.test.tsx b/webui/src/tests/thread-shell.test.tsx index 8868482e3..7a16aef9b 100644 --- a/webui/src/tests/thread-shell.test.tsx +++ b/webui/src/tests/thread-shell.test.tsx @@ -368,7 +368,6 @@ function modelSettings(model: string, provider: string): SettingsPayload { heartbeat: { enabled: true, interval_s: 1800, - keep_recent_messages: 8, }, dream: { schedule: "every 2h",