refactor(memory): decouple archival from provider state (#5565)

* refactor(memory): decouple archival from provider state

* test(memory): remove obsolete consolidation offset coverage
This commit is contained in:
chengyongru
2026-08-27 21:21:15 +08:00
committed by GitHub
parent 4d204ba077
commit 3c61fef7e8
17 changed files with 519 additions and 883 deletions
+1 -1
View File
@@ -48,7 +48,7 @@ class AutoCompact:
def _has_unarchived_messages(self, key: str) -> bool:
session = self.sessions.get_or_create(key)
return session.last_consolidated < len(session.messages)
return session.last_archived < len(session.messages)
@classmethod
def _is_internal_session(cls, key: str) -> bool:
+185 -127
View File
@@ -1,4 +1,4 @@
"""Memory system: pure file I/O store and lightweight Consolidator."""
"""Memory storage, transcript archiving, and legacy consolidation coordination."""
# Tool schemas are installed by the ``@tool_parameters`` class decorator at
# runtime; static analyzers cannot observe that it clears ``parameters`` from
@@ -785,7 +785,7 @@ class MemoryStore:
# ---------------------------------------------------------------------------
# Consolidator — lightweight token-budget triggered consolidation
# Memory ingestion and legacy context-pressure coordination
# ---------------------------------------------------------------------------
# Individual history.jsonl writers cap their own payloads tightly; the
@@ -796,8 +796,165 @@ _ARCHIVE_SUMMARY_MAX_CHARS = 8_000 # LLM-produced consolidation summary
_HISTORY_ENTRY_HARD_CAP = 64_000 # emergency cap in append_history
class MemoryArchiver:
"""Write durable transcript batches to the Memory ingestion journal.
The archiver deliberately has no SessionManager dependency: it may read a
captured transcript batch and append to history.jsonl, but it cannot mutate
provider continuation state or advance a session watermark.
"""
def __init__(
self,
store: MemoryStore,
build_messages: Callable[..., list[dict[str, Any]]],
get_tool_definitions: Callable[[], list[dict[str, Any]]],
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
unified_session: bool = False,
) -> None:
self.store = store
self._build_messages = build_messages
self._get_tool_definitions = get_tool_definitions
self._resolve_prompt_context = resolve_prompt_context
self.unified_session = unified_session
async def archive(
self,
messages: list[dict[str, Any]],
*,
runtime: LLMRuntime,
session_key: str,
request_messages: list[dict[str, Any]],
request_tools: list[dict[str, Any]],
) -> str | None:
"""Execute a prepared archive request and persist its result."""
if not messages:
return None
try:
with llm_usage_source("dream"):
response = await runtime.provider.chat_with_retry(
model=runtime.model,
messages=request_messages,
tools=request_tools,
tool_choice="none",
temperature=runtime.generation.temperature,
max_tokens=runtime.generation.max_tokens,
reasoning_effort=runtime.generation.reasoning_effort,
)
except Exception:
logger.warning("Memory archive provider call failed, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
if response.finish_reason in {"error", "length"}:
logger.warning(
"Memory archive provider did not complete ({}), raw-dumping to history",
response.finish_reason,
)
self.store.raw_archive(messages, session_key=session_key)
return None
if response.has_tool_calls is True:
logger.warning("Memory archive provider returned tool calls, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
summary = response.content
if not summary or not summary.strip():
logger.warning("Memory archive provider returned no summary, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
if summary.strip() == "(nothing)":
return "(nothing)"
self.store.append_history(
summary,
max_chars=_ARCHIVE_SUMMARY_MAX_CHARS,
session_key=session_key,
)
return summary
async def archive_session(
self,
session: Session,
*,
archive_end: int,
runtime: LLMRuntime,
input_token_budget: int,
) -> str | None:
"""Archive a captured session prefix without mutating the session."""
messages = list(session.messages[session.last_archived:archive_end])
if not messages:
return None
if input_token_budget <= 0:
logger.debug(
"Memory archive has no safe input budget for {}; raw-dumping",
session.key,
)
self.store.raw_archive(messages, session_key=session.key)
return None
prefix = Session(
key=session.key,
messages=list(session.messages[:archive_end]),
last_consolidated=session.last_archived,
)
history = prefix.get_history(max_tokens=input_token_budget)
archive_history = Session(
key=session.key,
messages=messages,
).get_history()
if not archive_history or history[-len(archive_history):] != archive_history:
logger.debug(
"Memory archive cannot replay the full chunk for {}; raw-dumping",
session.key,
)
self.store.raw_archive(messages, session_key=session.key)
return None
prompt = render_template(
"agent/consolidator_archive.md",
strip=True,
archive_count=len(archive_history),
)
channel = session.key.split(":", 1)[0] if ":" in session.key else None
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=history,
current_message=prompt,
channel=channel,
session_summary=session_summary_from_metadata(
session.metadata,
fallback_last_active=session.updated_at,
),
workspace=workspace,
session_key=session.key,
unified_session=self.unified_session,
)
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,
)
self.store.raw_archive(messages, session_key=session.key)
return None
return await self.archive(
messages,
runtime=runtime,
session_key=session.key,
request_messages=request_messages,
request_tools=tools,
)
class Consolidator:
"""Summarize compacted messages into history.jsonl."""
"""Legacy context-pressure coordinator backed by a MemoryArchiver."""
_MAX_CONSOLIDATION_ROUNDS = 5
@@ -820,6 +977,13 @@ class Consolidator:
self._build_messages = build_messages
self._get_tool_definitions = get_tool_definitions
self._resolve_prompt_context = resolve_prompt_context
self.archiver = MemoryArchiver(
store=store,
build_messages=build_messages,
get_tool_definitions=get_tool_definitions,
resolve_prompt_context=resolve_prompt_context,
unified_session=unified_session,
)
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
weakref.WeakValueDictionary()
)
@@ -834,7 +998,7 @@ class Consolidator:
tokens_to_remove: int,
) -> tuple[int, int] | None:
"""Pick a user-turn boundary that removes enough old prompt tokens."""
start = session.last_consolidated
start = session.last_archived
if start >= len(session.messages) or tokens_to_remove <= 0:
return None
@@ -912,48 +1076,14 @@ class Consolidator:
request_messages: list[dict[str, Any]],
request_tools: list[dict[str, Any]],
) -> str | None:
"""Execute a prepared consolidation request and persist its result."""
if not messages:
return None
try:
with llm_usage_source("dream"):
response = await runtime.provider.chat_with_retry(
model=runtime.model,
messages=request_messages,
tools=request_tools,
tool_choice="none",
temperature=runtime.generation.temperature,
max_tokens=runtime.generation.max_tokens,
reasoning_effort=runtime.generation.reasoning_effort,
)
except Exception:
logger.warning("Consolidation provider call failed, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
if response.finish_reason in {"error", "length"}:
logger.warning(
"Consolidation provider did not complete ({}), raw-dumping to history",
response.finish_reason,
)
self.store.raw_archive(messages, session_key=session_key)
return None
if response.has_tool_calls is True:
logger.warning("Consolidation provider returned tool calls, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
summary = response.content
if not summary or not summary.strip():
logger.warning("Consolidation provider returned no summary, raw-dumping to history")
self.store.raw_archive(messages, session_key=session_key)
return None
if summary.strip() == "(nothing)":
return "(nothing)"
self.store.append_history(
summary,
max_chars=_ARCHIVE_SUMMARY_MAX_CHARS,
"""Compatibility wrapper for the extracted MemoryArchiver."""
return await self.archiver.archive(
messages,
runtime=runtime,
session_key=session_key,
request_messages=request_messages,
request_tools=request_tools,
)
return summary
async def archive_session(
self,
@@ -962,82 +1092,12 @@ class Consolidator:
archive_end: int,
runtime: LLMRuntime,
) -> str | None:
"""Archive a session prefix by appending a consolidation instruction."""
messages = list(session.messages[session.last_consolidated:archive_end])
if not messages:
return None
budget = self._input_token_budget(runtime)
if budget <= 0:
logger.debug(
"Consolidation has no safe input budget for {}; raw-dumping",
session.key,
)
self.store.raw_archive(messages, session_key=session.key)
return None
prefix = Session(
key=session.key,
messages=list(session.messages[:archive_end]),
last_consolidated=session.last_consolidated,
)
history = prefix.get_history(max_tokens=budget)
archive_history = Session(
key=session.key,
messages=messages,
).get_history()
if (
not archive_history
or history[-len(archive_history):] != archive_history
):
logger.debug(
"Consolidation cannot replay the full chunk for {}; raw-dumping",
session.key,
)
self.store.raw_archive(messages, session_key=session.key)
return None
prompt = render_template(
"agent/consolidator_archive.md",
strip=True,
archive_count=len(archive_history),
)
channel = session.key.split(":", 1)[0] if ":" in session.key else None
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=history,
current_message=prompt,
channel=channel,
session_summary=session_summary_from_metadata(
session.metadata,
fallback_last_active=session.updated_at,
),
workspace=workspace,
session_key=session.key,
unified_session=self.unified_session,
)
tools = self._get_tool_definitions()
estimated, source = estimate_prompt_tokens_chain(
runtime.provider,
runtime.model,
request_messages,
tools,
)
if estimated > budget:
logger.debug(
"Consolidation prefix exceeds budget for {}; raw-dumping: {}/{} via {}",
session.key,
estimated,
budget,
source,
)
self.store.raw_archive(messages, session_key=session.key)
return None
return await self.archive(
messages,
"""Compatibility wrapper for the extracted MemoryArchiver."""
return await self.archiver.archive_session(
session,
archive_end=archive_end,
runtime=runtime,
session_key=session.key,
request_messages=request_messages,
request_tools=tools,
input_token_budget=self._input_token_budget(runtime),
)
async def maybe_consolidate_by_tokens(
@@ -1074,14 +1134,14 @@ class Consolidator:
self._persist_last_summary(session, last_summary)
return
if estimated < budget:
unconsolidated_count = len(session.messages) - session.last_consolidated
unarchived_count = len(session.messages) - session.last_archived
logger.debug(
"Token consolidation idle {}: {}/{} via {}, msgs={}",
session.key,
estimated,
runtime.context_window_tokens,
source,
unconsolidated_count,
unarchived_count,
)
self._persist_last_summary(session, last_summary)
return
@@ -1101,7 +1161,7 @@ class Consolidator:
end_idx = boundary[0]
chunk = session.messages[session.last_consolidated:end_idx]
chunk = session.messages[session.last_archived:end_idx]
if not chunk:
break
@@ -1125,8 +1185,7 @@ class Consolidator:
# would just emit duplicate [RAW] entries.
if summary:
last_summary = summary
session.last_consolidated = end_idx
session.provider_state = None
session.last_archived = end_idx
self.sessions.save(session)
if not summary:
# LLM is degraded — stop hammering it this call;
@@ -1170,7 +1229,7 @@ class Consolidator:
self.sessions.invalidate(session_key)
session = self.sessions.get_or_create(session_key)
archive_start = session.last_consolidated
archive_start = session.last_archived
messages_to_archive = list(session.messages[archive_start:])
if not messages_to_archive:
return ""
@@ -1191,8 +1250,7 @@ class Consolidator:
# A turn can append while the provider call is in flight. Advance only
# through the captured batch so new messages remain eligible next time.
session.last_consolidated = archive_end
session.provider_state = None
session.last_archived = archive_end
self.sessions.save(session)
visible = session.get_history(
+1 -1
View File
@@ -311,7 +311,7 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
snapshot = list(session.messages)
archive_snapshot = None
runtime = None
if session.last_consolidated < len(snapshot):
if session.last_archived < len(snapshot):
runtime = ctx.runtime or loop.runtime_for_session(session)
archive_snapshot = replace(
session,
+39 -25
View File
@@ -82,6 +82,15 @@ def _json_object(value: object) -> dict[str, Any]:
return cast(dict[str, Any], value)
def _archive_offset(data: dict[str, Any]) -> int:
"""Read the Memory archive watermark across the field-name migration."""
for key in ("last_archived", "last_consolidated"):
offset = cast(object, data.get(key))
if isinstance(offset, int) and not isinstance(offset, bool):
return offset
return 0
# TODO(0.3.2): Remove the write_stdin replay migration after 0.3.1.
def _migrate_legacy_exec_arguments(container: dict[str, Any]) -> bool:
raw_arguments = cast(object, container.get("arguments"))
@@ -277,7 +286,10 @@ 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)
last_consolidated: int = 0 # Number of messages already consolidated to files
# 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.
last_consolidated: int = 0
provider_state: ProviderConversationState | None = field(default=None, repr=False)
policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False)
@@ -295,6 +307,15 @@ class Session:
):
self.last_consolidated = 0
@property
def last_archived(self) -> int:
"""Number of transcript messages already written to the Memory journal."""
return self.last_consolidated
@last_archived.setter
def last_archived(self, value: int) -> None:
self.last_consolidated = value
def add_message(self, role: str, content: str, **kwargs: Any) -> None:
"""Add a message to the session."""
msg = {
@@ -319,9 +340,9 @@ class Session:
A positive ``max_messages`` applies an explicit caller-owned count
limit. The normal model path relies on ``max_tokens`` instead.
"""
replay_start = self.last_consolidated
replay_start = self.last_archived
if replay_start:
# ``last_consolidated`` is archive progress, not a replay boundary.
# ``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.
recent_start = recent_message_start_index(
@@ -335,8 +356,8 @@ class Session:
if max_messages <= 0:
start_idx = 0
else:
unarchived_count = len(self.messages) - self.last_consolidated
if replay_start < self.last_consolidated and unarchived_count < max_messages:
unarchived_count = len(self.messages) - self.last_archived
if replay_start < self.last_archived and unarchived_count < max_messages:
# The archived replay suffix can exceed the nominal count when one
# tool-heavy turn spans the boundary. Preserve that complete turn.
start_idx = 0
@@ -459,7 +480,7 @@ class Session:
def clear(self) -> None:
"""Clear all messages and reset session to initial state."""
self.messages = []
self.last_consolidated = 0
self.last_archived = 0
self.provider_state = None
self.updated_at = datetime.now()
self.metadata.pop("_last_summary", None)
@@ -474,11 +495,11 @@ class Session:
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_consolidated in place.
self.messages and self.last_archived in place.
"""
if max_messages <= 0:
dropped = list(self.messages)
lc = self.last_consolidated
lc = self.last_archived
self.clear()
return RetentionResult(
dropped=dropped,
@@ -491,7 +512,7 @@ class Session:
)
original = list(self.messages)
before_lc = self.last_consolidated
before_lc = self.last_archived
start_idx = max(0, len(self.messages) - max_messages)
if extend_to_user:
@@ -551,7 +572,7 @@ class Session:
if i < before_lc and id(m) not in retained_ids
)
# New last_consolidated = count of retained messages that were inside
# New last_archived = count of retained messages that were inside
# the old consolidated prefix.
new_lc = sum(
1 for i, m in enumerate(original)
@@ -559,7 +580,7 @@ class Session:
)
self.messages = retained
self.last_consolidated = new_lc
self.last_archived = new_lc
if dropped:
self.provider_state = None
self.updated_at = datetime.now()
@@ -1167,12 +1188,7 @@ class JsonlSessionStore:
if isinstance(updated_at_value, str) and updated_at_value
else None
)
offset = cast(object, data.get("last_consolidated", 0))
last_consolidated = (
offset
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
last_consolidated = _archive_offset(data)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
provider_state = ProviderConversationState.from_private_record(
data.get("state")
@@ -1254,12 +1270,7 @@ class JsonlSessionStore:
if isinstance(updated_at_value, str) and updated_at_value:
with suppress(ValueError):
updated_at = datetime.fromisoformat(updated_at_value)
offset = cast(object, data.get("last_consolidated", 0))
last_consolidated = (
offset
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
last_consolidated = _archive_offset(data)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
candidate = ProviderConversationState.from_private_record(
data.get("state")
@@ -1419,6 +1430,9 @@ class JsonlSessionStore:
"created_at": session.created_at.isoformat(),
"updated_at": session.updated_at.isoformat(),
"metadata": session.metadata,
"last_archived": session.last_archived,
# Keep old nanobot releases able to read sessions written
# during the field-name migration.
"last_consolidated": session.last_consolidated,
}
f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n")
@@ -2011,8 +2025,8 @@ class SessionManager:
for key in _FORK_VOLATILE_METADATA_KEYS:
metadata.pop(key, None)
last_consolidated = min(source.last_consolidated, len(copied))
if source.last_consolidated > len(copied):
last_consolidated = min(source.last_archived, len(copied))
if source.last_archived > len(copied):
metadata.pop("_last_summary", None)
last_consolidated = 0
+1 -1
View File
@@ -44,7 +44,7 @@ def session_context_payload(session: Session) -> dict[str, Any]:
"schema_version": 1,
"session_key": session.key,
"total_messages": len(session.messages),
"archived_messages": min(session.last_consolidated, len(session.messages)),
"archived_messages": min(session.last_archived, len(session.messages)),
"replay_messages": len(replay),
"estimated_replay_tokens": replay_tokens,
"estimated_summary_tokens": summary_tokens,