mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-02 17:22:06 +03:00
refactor(agent): let runner own context compaction (#5568)
* refactor(agent): consolidate accepted history under pressure * fix(agent): align provider and session compaction * refactor(agent): simplify runner context compaction * refactor(agent): remove background token consolidation * fix(agent): keep injected transcript messages distinct * refactor(agent): unify native compaction summaries * fix(agent): preserve native compaction boundary * fix(agent): unify context compaction paths * fix(agent): preserve exact compaction request boundaries
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
+119
-38
@@ -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:
|
||||
|
||||
+164
-124
@@ -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,
|
||||
|
||||
+81
-205
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]],
|
||||
|
||||
@@ -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__
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
+14
-121
@@ -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
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -714,7 +714,6 @@ export interface SettingsPayload {
|
||||
heartbeat: {
|
||||
enabled: boolean;
|
||||
interval_s: number;
|
||||
keep_recent_messages: number;
|
||||
};
|
||||
dream: {
|
||||
schedule: string;
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -91,7 +91,6 @@ export function settingsPayload(): SettingsPayload {
|
||||
heartbeat: {
|
||||
enabled: true,
|
||||
interval_s: 1800,
|
||||
keep_recent_messages: 8,
|
||||
},
|
||||
dream: {
|
||||
schedule: "every 2h",
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user