From 802e1bc0674b2281459c47729485cc01f101ecd7 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Fri, 11 Sep 2026 19:29:08 -0700 Subject: [PATCH] fix(context): recover required persisted reminders safely Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 20 + amplifier_module_context_simple/__init__.py | 450 ++++++++++++++++++-- tests/test_request_overlays.py | 336 +++++++++++++++ 3 files changed, 763 insertions(+), 43 deletions(-) create mode 100644 tests/test_request_overlays.py diff --git a/README.md b/README.md index 7fca880..525fab7 100644 --- a/README.md +++ b/README.md @@ -104,6 +104,26 @@ Compaction triggers when token usage reaches the configured threshold (default: - **Recent messages**: Last N% of messages (configurable via `protected_recent`) - **Tool pairs**: Tool_use and tool_result messages are treated as atomic units +Persisted ephemeral reminders are recognized only from their structured metadata +(`ephemeral`, `persisted`, and `reminder_placement`), never from their text. +The first and last human user messages are used as compaction anchors; if a +legacy transcript has no human user message, the first and last user messages +remain the fallback. The first human anchor remains stubbable only at level 8. + +### Request overlays + +`get_messages_for_request_with_overlays(overlays, *, provider=None)` is an +additive duck-typed convenience method for an orchestrator that must ensure a +trusted persisted reminder appears in one request. Each overlay supplies the +canonical reminder message and a `pre_user` or `tail` placement. An optional +`overlay_content` string lets the orchestrator supply framing appropriate to +that placement, without changing the canonical message. If every required +reminder is visible, the normal request view is returned unchanged. Otherwise, +all currently required reminders are assembled once at their requested +placements, within the finite request budget. Recovery neither changes +canonical history nor commits overlay-driven sticky decisions. An unsafe +placement or required content that cannot fit raises `RequiredRequestOverlayError`. + ### Compaction Phases 1. **Phase 1 - Tool Result Truncation**: Older tool results are truncated to reduce token usage diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 4435d13..4476b0a 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -40,10 +40,10 @@ __amplifier_module_type__ = "context" import logging -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Sequence from datetime import UTC, datetime from sys import maxsize -from typing import Any +from typing import Any, Literal, NotRequired, TypedDict from amplifier_core import ModuleCoordinator, TextBlock @@ -99,6 +99,37 @@ ) +class RequestOverlay(TypedDict): + message: dict[str, Any] + placement: Literal["pre_user", "tail"] + overlay_content: NotRequired[str] + + +class RequiredRequestOverlayError(RuntimeError): + """A required request-only overlay cannot safely fit in this request.""" + + +def _is_persisted_ephemeral_reminder(msg: dict[str, Any]) -> bool: + """Return whether *msg* is the trusted persisted reminder envelope.""" + metadata = msg.get("metadata") + return ( + msg.get("role") == "user" + and isinstance(metadata, dict) + and metadata.get("ephemeral") is True + and metadata.get("persisted") is True + and metadata.get("reminder_placement") in ("pre_user", "tail") + ) + + +def _is_human_user_message(msg: dict[str, Any]) -> bool: + """A user message that is neither a tool result nor a trusted reminder.""" + return ( + msg.get("role") == "user" + and not msg.get("tool_call_id") + and not _is_persisted_ephemeral_reminder(msg) + ) + + def _carries_loaded_tool_state(msg: dict[str, Any]) -> bool: """True if removing this message would silently unload tools (TS:854).""" meta = msg.get("metadata") or {} @@ -348,6 +379,7 @@ def __init__( # Reported in compaction stats / notice so the LLM sees the total # accumulated effect, not just the most recent escalation step. self._sticky_level: int = 0 + self._transient_compaction = False @staticmethod def _encode_tool_result_text(text: str) -> bytes: @@ -418,7 +450,9 @@ def _limit_tool_result_text_at_ingress( remaining_text_bytes = self.max_tool_result_bytes - marker_bytes if isinstance(content, str): - limited_content = self._utf8_safe_prefix(content, remaining_text_bytes) + marker + limited_content = ( + self._utf8_safe_prefix(content, remaining_text_bytes) + marker + ) else: limited_content: list[Any] = [] marker_added = False @@ -433,7 +467,9 @@ def _limit_tool_result_text_at_ingress( prefix = self._utf8_safe_prefix(text, remaining_text_bytes) prefix_bytes = len(self._encode_tool_result_text(prefix)) if prefix_bytes == len(self._encode_tool_result_text(text)): - limited_content.append(self._copy_text_block_with_text(block, prefix)) + limited_content.append( + self._copy_text_block_with_text(block, prefix) + ) remaining_text_bytes -= prefix_bytes continue @@ -461,11 +497,15 @@ def _limit_tool_result_text_at_ingress( } return {**message, "content": limited_content}, event_data - async def _emit_tool_result_ingress_truncation(self, event_data: dict[str, Any]) -> None: + async def _emit_tool_result_ingress_truncation( + self, event_data: dict[str, Any] + ) -> None: """Emit ingress observability without ever blocking tool-result admission.""" if self._hooks is not None: try: - await self._hooks.emit("context:tool_result_ingress_truncated", event_data) + await self._hooks.emit( + "context:tool_result_ingress_truncated", event_data + ) return except Exception: emission_status = "event emission failed" @@ -580,32 +620,283 @@ async def get_messages_for_request( Returns: Messages ready for LLM request, compacted if necessary. """ - budget = self._calculate_budget(token_budget, provider) + effective_budget = self._effective_budget(token_budget, provider) + working_messages = await self._working_messages_for_request() + return await self._request_view(working_messages, effective_budget) + + async def get_messages_for_request_with_overlays( + self, + overlays: Sequence[RequestOverlay], + *, + provider: Any | None = None, + ) -> list[dict[str, Any]]: + """Build one request view with required, request-only reminder overlays. - # Reserve token budget for potential compaction notice (if enabled) - effective_budget = budget - if self.compaction_notice_enabled: - effective_budget = budget - self.compaction_notice_token_reserve - if effective_budget <= 0: - # Misconfiguration guard: if the reserve consumes the entire budget - # (or more), _should_compact's `budget > 0` check would silently - # force usage to 0, disabling compaction entirely rather than - # loudly failing. Fall back to the full budget instead of a - # non-positive effective budget - a reserve that swallows the - # whole context is not a valid state to compact against. - logger.warning( - f"compaction_notice_token_reserve ({self.compaction_notice_token_reserve:,}) " - f">= budget ({budget:,}); ignoring reserve for this request to avoid " - f"silently disabling compaction (effective budget would be {effective_budget:,})" + An overlay refers to an exact trusted persisted reminder already in the + canonical history. If normal compaction omitted or stubbed it, the + reminder is reintroduced for this request only; canonical history and + sticky decisions remain unchanged. + """ + required = self._validate_request_overlays(overlays) + request_budget = self._calculate_budget(None, provider) + effective_budget = self._effective_budget(request_budget, None) + working_messages = await self._working_messages_for_request() + normal_view = await self._request_view(working_messages, effective_budget) + + if all( + self._reminder_content_is_visible(normal_view, content) + for _, _, content, _ in required + ): + return normal_view + + # In recovery only, reserve the required messages once, outside the + # compactable source. Otherwise a visible required body is charged both + # as history and as an overlay, or lost in the transient pass. + normal_notice = self._trailing_compaction_notice(normal_view) + material_budget = request_budget - sum( + self._estimate_tokens( + [self._request_overlay_message(message, placement, overlay_content)] + ) + for message, placement, _, overlay_content in required + ) + if normal_notice is not None: + material_budget -= self._estimate_tokens([normal_notice]) + if material_budget <= 0: + raise RequiredRequestOverlayError( + "required request overlays leave no material context budget" + ) + + material = [ + message + for message in working_messages + if not any( + message == required_message for required_message, _, _, _ in required + ) + ] + transient_view = await self._compact_ephemeral_transient( + material_budget, material + ) + view = list(transient_view) + self._insert_request_overlays(view, required, normal_notice) + + if not all( + self._reminder_content_is_visible(view, overlay_content) + for _, _, _, overlay_content in required + ): + raise RequiredRequestOverlayError( + "required request overlays could not be restored" + ) + if self._estimate_tokens(view) > request_budget: + raise RequiredRequestOverlayError( + "required request overlays cannot fit within the request budget" + ) + return self._strip_internal_metadata(view) + + @staticmethod + def _reminder_literal_content(message: dict[str, Any]) -> str | None: + """Return a supported literal reminder body without decoding envelopes.""" + content = message.get("content") + if isinstance(content, str): + return content + if isinstance(content, TextBlock): + return content.text + if isinstance(content, list) and len(content) == 1: + block = content[0] + if isinstance(block, TextBlock): + return block.text + if ( + isinstance(block, dict) + and block.get("type") == "text" + and isinstance(block.get("text"), str) + ): + return block["text"] + return None + + def _validate_request_overlays( + self, overlays: Sequence[RequestOverlay] + ) -> list[tuple[dict[str, Any], str, str, str]]: + """Validate and de-duplicate overlays before building a request view.""" + if isinstance(overlays, str | bytes) or not isinstance(overlays, Sequence): + raise ValueError("overlays must be a sequence of overlay mappings") + + required: list[tuple[dict[str, Any], str, str, str]] = [] + seen_content: set[str] = set() + for overlay in overlays: + if not isinstance(overlay, dict): + raise ValueError("each request overlay must be a mapping") + message = overlay.get("message") + placement = overlay.get("placement") + if ( + not isinstance(message, dict) + or not _is_persisted_ephemeral_reminder(message) + or placement not in ("pre_user", "tail") + ): + raise ValueError( + "request overlay must contain a trusted reminder and placement" ) - effective_budget = budget - else: - logger.debug( - f"Reserved {self.compaction_notice_token_reserve} tokens for potential notice " - f"(effective budget: {effective_budget:,})" + content = self._reminder_literal_content(message) + if content is None: + raise ValueError( + "request overlay reminder content must be literal text" + ) + overlay_content = overlay.get("overlay_content", content) + if not isinstance(overlay_content, str) or not overlay_content: + raise ValueError( + "request overlay content must be non-empty literal text" ) + if not any( + candidate == message and _is_persisted_ephemeral_reminder(candidate) + for candidate in self.messages + ): + raise ValueError( + "request overlay reminder must exactly match canonical history" + ) + if content not in seen_content: + required.append((message, placement, content, overlay_content)) + seen_content.add(content) + return required + + def _reminder_content_is_visible( + self, messages: list[dict[str, Any]], content: str + ) -> bool: + """Whether a full trusted or fresh overlay body is present in a view.""" + return any( + self._reminder_literal_content(message) == content + and ( + _is_persisted_ephemeral_reminder(message) + or ( + message.get("role") == "user" + and isinstance(message.get("metadata"), dict) + and message["metadata"].get("ephemeral") is True + and message["metadata"].get("persisted") is not True + ) + ) + for message in messages + ) - # Determine working messages based on whether factory is set + @staticmethod + def _trailing_compaction_notice( + messages: list[dict[str, Any]], + ) -> dict[str, Any] | None: + """Return this request's normal compaction notice, if one was emitted.""" + if not messages: + return None + metadata = messages[-1].get("metadata") + if ( + isinstance(metadata, dict) + and metadata.get("source") == "context-compaction" + and metadata.get("ephemeral") is True + ): + return messages[-1] + return None + + @staticmethod + def _request_overlay_message( + message: dict[str, Any], placement: str, overlay_content: str | None = None + ) -> dict[str, Any]: + """Copy a canonical reminder as a fresh request-only ephemeral message.""" + metadata = dict(message["metadata"]) + metadata.pop("persisted", None) + metadata.pop("_seq", None) + metadata["ephemeral"] = True + metadata["reminder_placement"] = placement + return { + **message, + "content": message["content"] + if overlay_content is None + else overlay_content, + "metadata": metadata, + } + + @staticmethod + def _tail_request_placement_is_safe(messages: list[dict[str, Any]]) -> bool: + """A tail user message is safe only when all tool calls have results.""" + pending: set[str] = set() + for message in messages: + if message.get("role") == "assistant" and message.get("tool_calls"): + for call in message["tool_calls"]: + if not isinstance(call, dict): + return False + call_id = call.get("id") or call.get("tool_call_id") + if not call_id: + return False + pending.add(call_id) + elif message.get("role") == "tool": + call_id = message.get("tool_call_id") + if call_id: + pending.discard(call_id) + return not pending + + def _insert_request_overlays( + self, + view: list[dict[str, Any]], + missing: list[tuple[dict[str, Any], str, str, str]], + normal_notice: dict[str, Any] | None, + ) -> None: + """Insert fresh overlays without splitting an assistant/tool-result group.""" + pre_user = [ + self._request_overlay_message(message, placement, overlay_content) + for message, placement, _, overlay_content in missing + if placement == "pre_user" + ] + tail = [ + self._request_overlay_message(message, placement, overlay_content) + for message, placement, _, overlay_content in missing + if placement == "tail" + ] + anchor = next( + ( + i + for i in range(len(view) - 1, -1, -1) + if _is_human_user_message(view[i]) + ), + None, + ) + needs_safe_end = ( + tail or (pre_user and anchor is None) or normal_notice is not None + ) + if needs_safe_end and not self._tail_request_placement_is_safe(view): + raise RequiredRequestOverlayError( + "required request overlay has no safe placement around pending tool calls" + ) + if pre_user: + if anchor is None: + view.extend(pre_user) + else: + view[anchor:anchor] = pre_user + if normal_notice is not None: + view.append(normal_notice) + view.extend(tail) + + def _effective_budget(self, token_budget: int | None, provider: Any | None) -> int: + """Calculate the normal request budget, including the notice reserve.""" + budget = self._calculate_budget(token_budget, provider) + if not self.compaction_notice_enabled: + return budget + + effective_budget = budget - self.compaction_notice_token_reserve + if effective_budget <= 0: + # Misconfiguration guard: if the reserve consumes the entire budget + # (or more), _should_compact's `budget > 0` check would silently + # force usage to 0, disabling compaction entirely rather than + # loudly failing. Fall back to the full budget instead of a + # non-positive effective budget - a reserve that swallows the + # whole context is not a valid state to compact against. + logger.warning( + f"compaction_notice_token_reserve ({self.compaction_notice_token_reserve:,}) " + f">= budget ({budget:,}); ignoring reserve for this request to avoid " + f"silently disabling compaction (effective budget would be {effective_budget:,})" + ) + return budget + + logger.debug( + f"Reserved {self.compaction_notice_token_reserve} tokens for potential notice " + f"(effective budget: {effective_budget:,})" + ) + return effective_budget + + async def _working_messages_for_request(self) -> list[dict[str, Any]]: + """Build the one fresh source view used for one request.""" if self._system_prompt_factory: # Factory mode: get fresh system content, exclude stored system messages # BUT preserve hook-injected system messages (they have metadata.source = "hook") @@ -620,15 +911,19 @@ async def get_messages_for_request( if msg.get("role") != "system" or (msg.get("metadata") or {}).get("source") == "hook" ] - working_messages = [system_message] + conversation_messages logger.debug( f"System prompt factory produced {len(system_content):,} chars, " f"{len(conversation_messages)} conversation messages" ) - else: - # Static mode: use messages as-is (may include stored system messages) - working_messages = list(self.messages) + return [system_message] + conversation_messages + # Static mode: use messages as-is (may include stored system messages). + return list(self.messages) + + async def _request_view( + self, working_messages: list[dict[str, Any]], effective_budget: int + ) -> list[dict[str, Any]]: + """Apply the normal compaction policy to an already-built source view.""" token_count, meter_source, estimated_tokens = self._measure_working_tokens( working_messages ) @@ -880,7 +1175,9 @@ def _exceeds_threshold(self, estimated_tokens: int, budget: int) -> bool: self.token_meter == TOKEN_METER_ACTUAL and self._last_measured_prompt_tokens is not None ): - return (self._last_measured_prompt_tokens / budget) >= self.compact_threshold + return ( + self._last_measured_prompt_tokens / budget + ) >= self.compact_threshold return (estimated_tokens / budget) >= self.compact_threshold def _measure_working_tokens( @@ -980,6 +1277,44 @@ async def _on_llm_response(self, event: str, data: dict[str, Any]) -> Any: # entire compaction decision on every single get_messages_for_request() # call. + def _snapshot_compaction_state( + self, + ) -> tuple[set[int], set[int], set[int], int, dict[str, Any] | None]: + """Take the explicit mutable state a request-only pass must not retain.""" + return ( + set(self._removed_seqs), + set(self._truncated_seqs), + set(self._stubbed_seqs), + self._sticky_level, + self._last_compaction_stats, + ) + + def _restore_compaction_state( + self, + state: tuple[set[int], set[int], set[int], int, dict[str, Any] | None], + ) -> None: + """Restore state after transient overlay recovery.""" + ( + self._removed_seqs, + self._truncated_seqs, + self._stubbed_seqs, + self._sticky_level, + self._last_compaction_stats, + ) = state + + async def _compact_ephemeral_transient( + self, budget: int, source_messages: list[dict[str, Any]] + ) -> list[dict[str, Any]]: + """Compact a request-only view without sticky changes, stats, or events.""" + state = self._snapshot_compaction_state() + previous_transient = self._transient_compaction + self._transient_compaction = True + try: + return await self._compact_ephemeral(budget, source_messages) + finally: + self._restore_compaction_state(state) + self._transient_compaction = previous_transient + @staticmethod def _extract_seq(msg: dict[str, Any]) -> int | None: """Return a message's stable sequence id, or None if it has none. @@ -1423,7 +1758,10 @@ async def _compact_ephemeral( # === LEVEL 8: Stub first user message + remove old stubs (extreme pressure) === max_level_reached = 8 - # Find first user message and stub it if not already stubbed + # Find human anchors; persisted reminders must not become the + # "original task" merely because they were admitted first. + first_human_idx = None + last_human_idx = None first_user_idx = None last_user_idx = None for i, msg in enumerate(working_messages): @@ -1431,15 +1769,25 @@ async def _compact_ephemeral( if first_user_idx is None: first_user_idx = i last_user_idx = i + if _is_human_user_message(msg): + if first_human_idx is None: + first_human_idx = i + last_human_idx = i + first_anchor_idx = ( + first_human_idx if first_human_idx is not None else first_user_idx + ) + last_anchor_idx = ( + last_human_idx if last_human_idx is not None else last_user_idx + ) # Stub first user message (previously protected) - but NEVER if it's also the last # The last user message is the current intent and must always be preserved - if first_user_idx is not None and first_user_idx != last_user_idx: - first_msg = working_messages[first_user_idx] + if first_anchor_idx is not None and first_anchor_idx != last_anchor_idx: + first_msg = working_messages[first_anchor_idx] if not first_msg.get("_stubbed"): content = first_msg.get("content", "") if isinstance(content, str) and len(content) > 80: - working_messages[first_user_idx] = self._stub_user_message( + working_messages[first_anchor_idx] = self._stub_user_message( first_msg ) # Sticky: record before `first_msg` var is superseded. @@ -1461,7 +1809,7 @@ async def _compact_ephemeral( for i, msg in enumerate(working_messages) if msg.get("_stubbed") and i < protected_boundary # Outside protected recent zone - and i != last_user_idx # Never remove last user message + and i != last_anchor_idx # Never remove the current anchor ] stubs_removed = 0 @@ -1625,7 +1973,10 @@ def _remove_messages_with_protection( i for i, msg in enumerate(messages) if msg.get("role") == "user" } - # Find first and last user message indices (always fully protected from stubbing too) + # Find human anchors. If history contains only persisted reminders, + # retain the original role-user behavior as a safe fallback. + first_human_idx = None + last_human_idx = None first_user_idx = None last_user_idx = None for i, msg in enumerate(messages): @@ -1633,6 +1984,16 @@ def _remove_messages_with_protection( if first_user_idx is None: first_user_idx = i last_user_idx = i + if _is_human_user_message(msg): + if first_human_idx is None: + first_human_idx = i + last_human_idx = i + first_anchor_idx = ( + first_human_idx if first_human_idx is not None else first_user_idx + ) + last_anchor_idx = ( + last_human_idx if last_human_idx is not None else last_user_idx + ) # Always protect system messages for i, msg in enumerate(messages): @@ -1659,8 +2020,8 @@ def _remove_messages_with_protection( # We don't add it to protected_indices so it can be stubbed at Level 8 # Always protect the LAST user message (current context) - if last_user_idx is not None: - protected_indices.add(last_user_idx) + if last_anchor_idx is not None: + protected_indices.add(last_anchor_idx) # Protect last N% of messages (using the passed protection level) protected_boundary = int(len(messages) * (1 - protected_recent)) @@ -1752,8 +2113,8 @@ def _remove_messages_with_protection( i for i in user_message_indices if i not in protected_indices - and i != first_user_idx # Protected from stubbing at levels 1-7 - and i != last_user_idx # Always protected (never stubbed) + and i != first_anchor_idx # Protected from stubbing at levels 1-7 + and i != last_anchor_idx # Always protected (never stubbed) and not messages[i].get("_stubbed") # Don't re-stub ] ) @@ -1953,6 +2314,9 @@ async def _finalize_compaction_with_stats( f"{cause}." ) + if self._transient_compaction: + return final_messages + # Cumulative high-water mark across ALL escalations ever, not just # this one -- monotonic, never goes backward. This is what feeds the # compaction notice, so its content only changes when a genuinely diff --git a/tests/test_request_overlays.py b/tests/test_request_overlays.py new file mode 100644 index 0000000..f6724af --- /dev/null +++ b/tests/test_request_overlays.py @@ -0,0 +1,336 @@ +"""Focused regression coverage for persisted reminders and request overlays.""" + +from __future__ import annotations + +import copy + +import pytest + +from amplifier_module_context_simple import ( + RequiredRequestOverlayError, + SimpleContextManager, + _is_human_user_message, + _is_persisted_ephemeral_reminder, +) + + +def _reminder(content: str, placement: str = "pre_user") -> dict: + return { + "role": "user", + "content": content, + "metadata": { + "ephemeral": True, + "persisted": True, + "reminder_placement": placement, + }, + } + + +async def _hidden_reminder( + content: str = "restore me " * 20, placement: str = "pre_user" +) -> tuple[SimpleContextManager, dict]: + context = SimpleContextManager( + max_tokens=10_000, compact_threshold=0, compaction_notice_enabled=False + ) + await context.add_message({"role": "user", "content": "current human"}) + await context.add_message(_reminder(content, placement)) + canonical = context.messages[-1] + context._stubbed_seqs.add(canonical["metadata"].get("_seq")) + return context, canonical + + +@pytest.mark.parametrize("protected_recent", [0.30, 0.18, 0.09]) +def test_removal_levels_use_human_anchors(protected_recent: float): + """L3/L5/L7 preserve human anchors; role-user synthetic stays retained.""" + context = SimpleContextManager(compaction_notice_enabled=False) + messages = [ + _reminder("synthetic " * 30), + {"role": "user", "content": "first human " * 30}, + {"role": "assistant", "content": "old assistant " * 30}, + {"role": "user", "content": "middle human " * 30}, + {"role": "user", "content": "current human " * 30}, + ] + result, _, _, _ = context._remove_messages_with_protection( + messages, target_tokens=0, protected_recent=protected_recent, system_tokens=0 + ) + contents = [message["content"] for message in result] + assert "first human " * 30 in contents + assert "current human " * 30 in contents + assert any( + message.get("_stubbed") and message.get("metadata", {}).get("persisted") is True + for message in result + ) + + +@pytest.mark.asyncio +async def test_level_eight_keeps_current_human_protected(): + context = SimpleContextManager(compaction_notice_enabled=False) + await context.add_message(_reminder("synthetic " * 60)) + await context.add_message({"role": "user", "content": "first human " * 60}) + await context.add_message({"role": "user", "content": "current human " * 60}) + view = await context._compact_ephemeral(20) + # Level 8 relaxes the first-human rule; the current human intent remains + # protected even under an otherwise irreducible all-user transcript. + assert "current human " * 60 in [message["content"] for message in view] + + +@pytest.mark.parametrize( + "message, expected", + [ + (_reminder("trusted"), True), + ({"role": "user", "content": "", "metadata": {}}, False), + ({"role": "user", "content": "x", "metadata": None}, False), + ({"role": "user", "content": "x", "metadata": "malformed"}, False), + ({"role": "assistant", "content": "x", "metadata": {}}, False), + ], +) +def test_synthetic_detector_is_metadata_only_and_safe(message: dict, expected: bool): + assert _is_persisted_ephemeral_reminder(message) is expected + + +@pytest.mark.asyncio +async def test_set_messages_preserves_trusted_reminder_classification(): + context = SimpleContextManager(compaction_notice_enabled=False) + await context.set_messages( + [_reminder("restored"), {"role": "user", "content": "human"}] + ) + assert _is_persisted_ephemeral_reminder(context.messages[0]) + assert _is_human_user_message(context.messages[1]) + + +@pytest.mark.asyncio +async def test_visible_overlay_matches_normal_view_and_calls_factory_once(): + async def populated() -> tuple[SimpleContextManager, dict, list[int]]: + context = SimpleContextManager( + max_tokens=10_000, compaction_notice_enabled=False + ) + calls = [0] + + async def factory() -> str: + calls[0] += 1 + return "fresh system" + + await context.set_system_prompt_factory(factory) + await context.add_message({"role": "user", "content": "human"}) + await context.add_message(_reminder("visible")) + return context, copy.deepcopy(context.messages[-1]), calls + + normal, _, normal_calls = await populated() + overlaid, overlay, overlay_calls = await populated() + assert ( + await overlaid.get_messages_for_request_with_overlays( + [{"message": overlay, "placement": "tail"}] + ) + == await normal.get_messages_for_request() + ) + assert overlay_calls == normal_calls == [1] + + +@pytest.mark.asyncio +async def test_missing_overlay_is_request_only_and_transient_state_matches_normal(): + normal, _ = await _hidden_reminder() + recovering, canonical = await _hidden_reminder() + await normal.get_messages_for_request() + baseline_state = normal._snapshot_compaction_state() + view = await recovering.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "pre_user"}] + ) + restored = next( + message for message in view if message["content"] == "restore me " * 20 + ) + assert restored["metadata"].get("persisted") is not True + assert recovering._snapshot_compaction_state() == baseline_state + assert canonical["metadata"]["persisted"] is True + assert "restore me " * 20 not in [ + message["content"] for message in await recovering.get_messages_for_request() + ] + + +@pytest.mark.asyncio +async def test_overlays_reinsert_all_required_and_preserve_current_placement(): + context, first = await _hidden_reminder("first " * 20) + await context.add_message(_reminder("second " * 20, "tail")) + second = context.messages[-1] + context._stubbed_seqs.add(second["metadata"].get("_seq")) + view = await context.get_messages_for_request_with_overlays( + [ + {"message": copy.deepcopy(first), "placement": "pre_user"}, + {"message": copy.deepcopy(first), "placement": "tail"}, + {"message": copy.deepcopy(second), "placement": "tail"}, + ] + ) + contents = [message["content"] for message in view] + assert contents.count("first " * 20) == contents.count("second " * 20) == 1 + assert contents.index("first " * 20) < contents.index("current human") + assert contents[-1] == "second " * 20 + + +@pytest.mark.asyncio +async def test_overlay_rejects_untrusted_or_noncanonical_messages(): + context, canonical = await _hidden_reminder() + untrusted = copy.deepcopy(canonical) + untrusted["metadata"].pop("persisted") + changed = copy.deepcopy(canonical) + changed["content"] = "different" + for invalid in (untrusted, changed): + with pytest.raises(ValueError): + await context.get_messages_for_request_with_overlays( + [{"message": invalid, "placement": "pre_user"}] + ) + + +@pytest.mark.asyncio +async def test_tail_overlay_refuses_partial_tool_group_and_rolls_back(): + context, canonical = await _hidden_reminder("must be tail " * 20, "tail") + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [{"id": "answered"}, {"id": "pending"}], + } + ) + await context.add_message( + {"role": "tool", "tool_call_id": "answered", "content": "ok"} + ) + await context.get_messages_for_request() + before = context._snapshot_compaction_state() + with pytest.raises(RequiredRequestOverlayError): + await context.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "tail"}] + ) + assert context._snapshot_compaction_state() == before + + +@pytest.mark.asyncio +async def test_recovery_charges_visible_required_body_and_notice_only_once(): + context, missing = await _hidden_reminder("missing policy " * 60) + await context.add_message(_reminder("visible policy " * 60, "tail")) + visible = context.messages[-1] + # Produce an actual normal compaction notice before choosing a tight cap. + context.compaction_notice_enabled = True + context.compaction_notice_token_reserve = 500 + normal = await context.get_messages_for_request() + notice = context._trailing_compaction_notice(normal) + assert notice is not None + expected = [ + next(m for m in context.messages if m["content"] == "current human"), + context._request_overlay_message(missing, "pre_user"), + notice, + context._request_overlay_message(visible, "tail"), + ] + # Raw budget admits each body exactly once, but not another notice reserve. + context.max_tokens = context._estimate_tokens(expected) + 100 + canonical_before = copy.deepcopy(context.messages) + view = await context.get_messages_for_request_with_overlays( + [ + {"message": copy.deepcopy(missing), "placement": "pre_user"}, + {"message": copy.deepcopy(visible), "placement": "tail"}, + ] + ) + assert context._estimate_tokens(view) <= context.max_tokens + assert sum(m["content"] == missing["content"] for m in view) == 1 + assert sum(m["content"] == visible["content"] for m in view) == 1 + assert context.messages == canonical_before + + +@pytest.mark.asyncio +async def test_recovery_refuses_required_content_larger_than_entire_budget(): + context, canonical = await _hidden_reminder("large policy " * 300) + context.max_tokens = 100 + await context.get_messages_for_request() + before = context._snapshot_compaction_state() + with pytest.raises(RequiredRequestOverlayError, match="budget"): + await context.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "tail"}] + ) + assert context._snapshot_compaction_state() == before + + +@pytest.mark.asyncio +async def test_recovery_factory_runs_once_and_counts_system_floor(): + context, canonical = await _hidden_reminder() + calls = 0 + + async def factory(): + nonlocal calls + calls += 1 + return "system policy " * 100 + + await context.set_system_prompt_factory(factory) + context.max_tokens = 100 + with pytest.raises(RequiredRequestOverlayError, match="budget"): + await context.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "pre_user"}] + ) + assert calls == 1 + + +@pytest.mark.asyncio +async def test_recovery_accepts_a_completed_multi_tool_group(): + context, canonical = await _hidden_reminder("current policy " * 20, "tail") + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [{"id": "first"}, {"id": "second"}], + } + ) + for call_id in ("first", "second"): + await context.add_message( + {"role": "tool", "tool_call_id": call_id, "content": "ok"} + ) + view = await context.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "tail"}] + ) + assert view[-1]["content"] == canonical["content"] + assert [m["tool_call_id"] for m in view if m["role"] == "tool"] == [ + "first", + "second", + ] + + +@pytest.mark.asyncio +async def test_visible_plaintext_echo_does_not_suppress_trusted_recovery(): + context, canonical = await _hidden_reminder() + await context.add_message({"role": "assistant", "content": canonical["content"]}) + view = await context.get_messages_for_request_with_overlays( + [{"message": copy.deepcopy(canonical), "placement": "tail"}] + ) + assert view[-1]["role"] == "user" + assert view[-1]["content"] == canonical["content"] + assert view[-1]["metadata"].get("persisted") is not True + + +@pytest.mark.asyncio +async def test_recovery_can_use_current_placement_framing_without_editing_history(): + context, canonical = await _hidden_reminder("original pre-user frame " * 10) + before = copy.deepcopy(canonical) + replacement = "Current tail frame: " + "required policy " * 10 + view = await context.get_messages_for_request_with_overlays( + [ + { + "message": copy.deepcopy(canonical), + "placement": "tail", + "overlay_content": replacement, + } + ] + ) + assert view[-1]["content"] == replacement + assert view[-1]["metadata"]["reminder_placement"] == "tail" + assert canonical == before + + +@pytest.mark.asyncio +async def test_recovery_budgets_the_supplied_framing_not_just_canonical_text(): + context, canonical = await _hidden_reminder() + context.max_tokens = 1000 + with pytest.raises(RequiredRequestOverlayError, match="budget"): + await context.get_messages_for_request_with_overlays( + [ + { + "message": copy.deepcopy(canonical), + "placement": "tail", + "overlay_content": "long current frame " * 1000, + } + ] + )