From d6ab7cc44c2abc4165247a5552dd990e35ac7d64 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Mon, 14 Sep 2026 22:51:04 -0700 Subject: [PATCH 1/2] fix: stabilize retained request views and hard-fit budgets Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 8 + amplifier_module_context_simple/__init__.py | 250 +++++++++++------- tests/test_request_retention.py | 229 +++++++++++++++- .../test_sticky_compaction_and_tail_notice.py | 39 +++ 4 files changed, 432 insertions(+), 94 deletions(-) diff --git a/README.md b/README.md index 4282ddd..69689ac 100644 --- a/README.md +++ b/README.md @@ -122,6 +122,14 @@ compaction state. Protection takes precedence over the compaction target. A protected tool cohort can leave a view above that target; it is not a strict native-token ceiling. +After a compaction decision, every later provider-facing request reuses the +same reduced view before it is measured; canonical history remains complete. +The capability also accepts `hard_fit=True` for a provider-directed forced +rebuild: it targets the supplied effective request budget rather than applying +`target_usage` again. This is an additive capability keyword for orchestrators, +not a user configuration setting; ordinary `token_budget` calls keep their +existing `target_usage` semantics. + ### 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 cfc8830..e2e894f 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -347,6 +347,10 @@ def __init__( self._system_prompt_factory: Callable[[], Awaitable[str]] | None = None self._request_retained_contents: frozenset[str] = frozenset() self._request_protected_seqs: set[int] = set() + # Request-retention's optional provider-directed fit is deliberately + # transient. It is set only by get_messages_for_request_retaining() + # and restored before that capability returns or raises. + self._request_hard_fit = False # --- Sticky compaction decision state --- # Compaction decisions (remove / truncate / stub) are keyed by a @@ -578,6 +582,7 @@ async def get_messages_for_request_retaining( retain_contents: list[str], provider: Any | None = None, token_budget: int | None = None, + hard_fit: bool = False, ) -> list[dict[str, Any]]: """Optional capability: retain current persisted injections for one view. @@ -585,9 +590,12 @@ async def get_messages_for_request_retaining( ephemeral=True and persisted=True. Only the newest matching copy is protected. This neither admits messages nor changes their lifetime; the caller supplies the current delivery requirements on every call. + `hard_fit=True` is a capability-only request to target the supplied + effective budget rather than the normal target_usage fraction. """ previous_contents = self._request_retained_contents previous_seqs = self._request_protected_seqs + previous_hard_fit = self._request_hard_fit decisions = ( self._removed_seqs.copy(), self._truncated_seqs.copy(), @@ -598,6 +606,7 @@ async def get_messages_for_request_retaining( ) try: self._request_retained_contents = frozenset(retain_contents) + self._request_hard_fit = hard_fit return await self.get_messages_for_request(token_budget, provider) except BaseException: # A failed/cancelled assembly must not commit a reduction that was @@ -614,6 +623,7 @@ async def get_messages_for_request_retaining( finally: self._request_retained_contents = previous_contents self._request_protected_seqs = previous_seqs + self._request_hard_fit = previous_hard_fit def _protected_sequences(self, messages: list[dict[str, Any]]) -> set[int]: humans = [msg for msg in messages if _is_human_message(msg)] @@ -641,13 +651,40 @@ def _is_request_protected(self, msg: dict[str, Any]) -> bool: def _check_retained_budget( self, messages: list[dict[str, Any]], budget: int ) -> None: - if self._request_retained_contents and self._estimate_tokens(messages) > budget: + if ( + (self._request_retained_contents or self._request_hard_fit) + and self._estimate_tokens(messages) > budget + ): raise ContextLengthError( "Context cannot fit the current injections and protected conversation " "within the estimated input budget; shorten the active instructions " "or use a larger context window. Required content was not discarded." ) + @staticmethod + def _tail_has_unanswered_tool_calls(messages: list[dict[str, Any]]) -> bool: + """Return whether the trailing tool-call group is incomplete. + + A tail can end on the assistant tool-call message itself or on one of + several sibling tool results. A notice is safe only after every call + declared by that final assistant message has a result. + """ + trailing_result_ids: set[str] = set() + for message in reversed(messages): + if message.get("role") == "tool": + tool_call_id = message.get("tool_call_id") + if tool_call_id: + trailing_result_ids.add(tool_call_id) + continue + if message.get("role") != "assistant" or not message.get("tool_calls"): + return False + declared_ids = { + tool_call.get("id") or tool_call.get("tool_call_id") + for tool_call in message["tool_calls"] + } + return bool(declared_ids - trailing_result_ids) + return False + async def get_messages_for_request( self, token_budget: int | None = None, @@ -663,9 +700,8 @@ async def get_messages_for_request( Applies EPHEMERAL compaction if needed - returns a NEW list without modifying self.messages. The original history is always preserved. - If compaction occurs and notice is enabled, a system-reminder is inserted - at position 1 (after main system message) to inform the LLM about what - was compacted. + If compaction occurs and notice is enabled, a tail system-reminder + informs the LLM about what was compacted without changing its prefix. Args: token_budget: Optional explicit token limit (deprecated, prefer provider). @@ -734,8 +770,22 @@ async def get_messages_for_request( effective_budget, ) + # Materialize every recorded sticky decision before measuring or deciding + # whether more compaction is needed. Canonical history stays untouched: + # only this request view receives the old removal/truncation/stubbing + # transforms. New messages have no decision and remain visible. + system_messages = [ + dict(msg) for msg in working_messages if msg.get("role") == "system" + ] + sticky_view = system_messages + self._apply_sticky_decisions( + [msg for msg in working_messages if msg.get("role") != "system"] + ) + has_sticky_decisions = bool( + self._removed_seqs or self._truncated_seqs or self._stubbed_seqs + ) + token_count, meter_source, estimated_tokens = self._measure_working_tokens( - working_messages + sticky_view ) self._last_token_meter_stats = { "mode": self.token_meter, @@ -747,12 +797,24 @@ async def get_messages_for_request( "ratio": (token_count / effective_budget) if effective_budget > 0 else None, } - # Check if compaction needed (using effective budget with notice reserve deducted) + # Check whether this request needs a new compaction escalation. A + # provider-directed hard fit applies at the full forced budget, not at + # the ordinary compact_threshold or target_usage fraction of it. retained_view_over_budget = ( bool(self._request_retained_contents) and estimated_tokens > effective_budget ) - if self._should_compact(token_count, effective_budget) or retained_view_over_budget: + hard_fit_view_over_budget = ( + self._request_hard_fit and estimated_tokens > effective_budget + ) + should_compact = ( + self._should_compact(token_count, effective_budget) + or retained_view_over_budget + or hard_fit_view_over_budget + ) + previous_compaction_stats = self._last_compaction_stats + deferred_hard_fit_stats: dict[str, Any] | None = None + if should_compact: # Compact EPHEMERALLY - returns new list, working_messages unchanged compacted = await self._compact_ephemeral( effective_budget, working_messages @@ -760,91 +822,76 @@ async def get_messages_for_request( logger.info( f"Ephemeral compaction: {len(working_messages)} -> {len(compacted)} messages for this request" ) + if ( + self._request_hard_fit + and self._last_compaction_stats is not previous_compaction_stats + ): + deferred_hard_fit_stats = self._last_compaction_stats - # Append compaction notice at the TAIL if enabled and level threshold met. - # - # CRITICAL (prompt cache stability): this notice must never be inserted - # into the prefix. Two things make the tail the only safe placement: - # - # 1. role: previously this was "system", which -- for the Anthropic - # provider -- gets extracted OUT of the conversation entirely and - # merged into the single top-level system content block (see - # provider-anthropic's `_complete_chat_request`: `system_msgs = [m - # for m in request.messages if m.role == "system"]`, combined by - # `_format_system_with_cache`). That means a "system"-role notice - # inserted anywhere -- even at the tail -- would change the system - # block's text on every compaction, busting the SYSTEM cache - # breakpoint too, not just the conversation-region one. Using - # "user" keeps this message in the conversation region, where the - # provider's ephemeral-exclusion logic can see and skip it. - # 2. metadata.ephemeral=True + tail position: the Anthropic provider's - # `_count_trailing_ephemeral_messages` walks backward from the end - # of the conversation and excludes trailing messages carrying - # `metadata.ephemeral=True` from cache-breakpoint placement. A - # "system"-role message is excluded from that walk entirely (it - # never reaches the conversation list), and content anywhere - # other than the tail is not "trailing" and would still corrupt - # the cached prefix. Tail + ephemeral=True + role != "system" is - # the only combination the existing provider fix recognizes. + elif has_sticky_decisions: + # Sticky decisions are sender-view state, not only an input to a + # future escalation. Return that stable reduced view even below the + # ordinary threshold (including inactive/direct fetches). + compacted = sticky_view + else: + self._check_retained_budget(working_messages, budget) + return self._strip_internal_metadata(working_messages) + + # Append compaction notice at the TAIL if enabled and level threshold met. + # CRITICAL (prompt cache stability): this notice must never be inserted + # into the prefix. Two things make the tail the only safe placement: # - # The notice content itself only changes when a NEW compaction - # escalation actually occurs (see _compact_ephemeral's sticky decision - # state) -- so on calls between escalations, this tail addition is - # byte-identical, and everything before it (the real prefix) is - # completely undisturbed either way. - if self.compaction_notice_enabled and self._last_compaction_stats: - level = self._last_compaction_stats.get("strategy_level", 0) - # GUARD: never append into an unanswered tool_calls turn. - # - # Appending at the tail is what makes the notice cache-safe - # (above), but the tail is not always a safe place to stand: if - # the view ends with an assistant message carrying tool_calls, - # its tool results have not been added yet, and a user-role - # notice would land BETWEEN the tool call and its results. - # Providers reject or mishandle that interleaving -- the same - # tool_use/tool_result atomicity the compaction levels work hard - # to preserve. (This is new exposure from the move to the tail; - # the old index-1 insert could never land here.) - # - # Skip rather than reposition: placing it before the assistant - # message would put it back INSIDE the prefix, re-introducing - # exactly the cache-busting this fix exists to prevent. Skipping - # costs nothing -- the notice is derived from sticky stats that - # persist, so it reappears on the very next request once the - # tool results have arrived and the tail is a safe place again. - if compacted and compacted[-1].get("tool_calls"): + # 1. role: previously this was "system", which -- for the Anthropic + # provider -- gets extracted OUT of the conversation entirely and + # merged into the single top-level system content block. Using "user" + # keeps it in the conversation region, where ephemeral exclusion can + # recognize it. + # 2. metadata.ephemeral=True + tail position: the provider's trailing + # ephemeral walk-back can exclude it without corrupting the cached + # prefix. + # + # The notice comes from cumulative sticky stats, so a stable view gets + # exactly one byte-identical notice on every safe request. + if self.compaction_notice_enabled and self._last_compaction_stats: + level = self._last_compaction_stats.get("strategy_level", 0) + # Do not interleave a user-role notice between an unanswered + # assistant tool call and its result. Sticky stats persist, so it + # safely reappears once the result reaches the history. + if self._tail_has_unanswered_tool_calls(compacted): + logger.debug( + "Skipping compaction notice this request: the trailing " + "tool-call group has unanswered calls; the notice would " + "interleave between tool_use and tool_result. It will be " + "appended on the next request instead." + ) + elif level >= self.compaction_notice_min_level: + notice = self._format_compaction_notice() + if notice: + compacted.append( + { + "role": "user", + "content": notice, + "metadata": { + "source": "context-compaction", + "ephemeral": True, + }, + } + ) logger.debug( - "Skipping compaction notice this request: view ends with " - "an assistant message with unanswered tool_calls; the " - "notice would interleave between tool_use and tool_result. " - "It will be appended on the next request instead." + f"Appended compaction notice at tail (level {level}, " + f"verbosity: {self.compaction_notice_verbosity})" ) - elif level >= self.compaction_notice_min_level: - notice = self._format_compaction_notice() - if notice: - compacted.append( - { - "role": "user", - "content": notice, - "metadata": { - "source": "context-compaction", - "ephemeral": True, - }, - } - ) - logger.debug( - f"Appended compaction notice at tail (level {level}, " - f"verbosity: {self.compaction_notice_verbosity})" - ) - # Strip internal bookkeeping at the module boundary -- everything - # above this point (sticky decisions, token accounting) still runs - # on messages carrying `_seq`; only what leaves has it removed. - self._check_retained_budget(compacted, budget) - return self._strip_internal_metadata(compacted) - - self._check_retained_budget(working_messages, budget) - return self._strip_internal_metadata(working_messages) + # Strip internal bookkeeping at the module boundary -- everything above + # this point (sticky decisions, token accounting) still runs on messages + # carrying `_seq`; only what leaves has it removed. + self._check_retained_budget(compacted, budget) + if deferred_hard_fit_stats is not None and self._hooks is not None: + try: + await self._hooks.emit("context:compaction", deferred_hard_fit_stats) + except Exception as e: + logger.warning(f"Could not emit compaction event: {e}") + return self._strip_internal_metadata(compacted) # Metadata keys that are internal bookkeeping only and must never cross # the module boundary into a provider-facing view. `_seq` is sticky @@ -1196,7 +1243,7 @@ async def _compact_ephemeral( source_messages if source_messages is not None else self.messages ) self._request_protected_seqs = self._protected_sequences(messages_to_compact) - target_tokens = int(budget * self.target_usage) + target_tokens = budget if self._request_hard_fit else int(budget * self.target_usage) old_count = len(messages_to_compact) old_tokens = self._estimate_tokens(messages_to_compact) @@ -1293,7 +1340,9 @@ async def _compact_ephemeral( # enough) -- this module fires the escalation honestly, but the # sizing of that escalation is only as good as the estimator was # before this meter existed. - needs_escalation = self._exceeds_threshold(current_tokens, budget) + needs_escalation = self._exceeds_threshold(current_tokens, budget) or ( + self._request_hard_fit and current_tokens > budget + ) if not needs_escalation: # Sticky state alone already keeps us under the threshold that # triggered compaction in the first place -- nothing NEW needs @@ -1997,7 +2046,13 @@ def _check_tool_pair_removable( for tc in assistant_msg.get("tool_calls", []): tc_id = tc.get("id") or tc.get("tool_call_id") if tc_id: - for k in tool_call_id_to_indices.get(tc_id, []): + result_indices = tool_call_id_to_indices.get(tc_id, []) + # An unanswered call must remain until its result arrives. + # Removing it now would make a subsequently admitted result + # orphaned in the provider-facing view. + if not result_indices: + all_removable = False + for k in result_indices: if k in protected_indices: all_removable = False else: @@ -2083,6 +2138,16 @@ async def _finalize_compaction_with_stats( f"{cause}." ) + # A provider-directed hard fit must fail before committing observable + # compaction state. The request-retention wrapper then restores its + # local decisions, and hook consumers never see a phantom compaction. + if self._request_hard_fit and final_tokens > budget: + raise ContextLengthError( + "Context cannot fit the current injections and protected conversation " + "within the estimated input budget; shorten the active instructions " + "or use a larger context window. Required content was not discarded." + ) + # 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 @@ -2113,8 +2178,9 @@ async def _finalize_compaction_with_stats( } self._last_compaction_stats = stats - # Emit event if hooks available - if self._hooks is not None: + # A hard fit delays emission until its tail notice is added and the + # complete provider-facing view passes its final budget check. + if self._hooks is not None and not self._request_hard_fit: try: await self._hooks.emit("context:compaction", stats) except Exception as e: diff --git a/tests/test_request_retention.py b/tests/test_request_retention.py index 8a41413..c16d694 100644 --- a/tests/test_request_retention.py +++ b/tests/test_request_retention.py @@ -144,7 +144,15 @@ async def test_impossible_retention_fails_before_returning_an_overfull_view(): @pytest.mark.asyncio async def test_post_compaction_failure_rolls_back_request_state_and_decisions(): - context = SimpleContextManager(max_tokens=1000, compaction_notice_enabled=False) + emitted: list[tuple[str, dict]] = [] + + class Hooks: + async def emit(self, event: str, data: dict) -> None: + emitted.append((event, data)) + + context = SimpleContextManager( + max_tokens=1000, compaction_notice_enabled=False, hooks=Hooks() + ) required = "Active retention requirement." await context.add_message(human("First human request.")) await context.add_message(reminder(required)) @@ -158,7 +166,9 @@ async def test_post_compaction_failure_rolls_back_request_state_and_decisions(): await context.add_message(human("Latest human correction.")) with pytest.raises(ContextLengthError, match="Required content was not discarded"): - await context.get_messages_for_request_retaining(retain_contents=[required]) + await context.get_messages_for_request_retaining( + retain_contents=[required], hard_fit=True + ) assert context._last_compaction_stats is None assert context._last_token_meter_stats is None @@ -166,6 +176,221 @@ async def test_post_compaction_failure_rolls_back_request_state_and_decisions(): assert not context._truncated_seqs assert not context._stubbed_seqs assert context._sticky_level == 0 + assert context._request_hard_fit is False + assert not emitted + + +@pytest.mark.asyncio +async def test_hard_fit_notice_failure_rolls_back_before_emitting_compaction(): + """A rejected final view must not publish a compaction that was rolled back.""" + emitted: list[tuple[str, dict]] = [] + + class Hooks: + async def emit(self, event: str, data: dict) -> None: + emitted.append((event, data)) + + context = SimpleContextManager( + max_tokens=1000, + compact_threshold=0.5, + target_usage=0.5, + protected_recent=0.2, + compaction_notice_enabled=True, + compaction_notice_token_reserve=1, + hooks=Hooks(), + ) + await context.add_message(human("Original human request.")) + for i in range(20): + await context.add_message( + {"role": "assistant", "content": f"Historical bulk {i}: " + "x" * 300} + ) + await context.add_message(human("Latest human correction.")) + context._format_compaction_notice = lambda: "notice " * 1000 + + with pytest.raises(ContextLengthError, match="Required content was not discarded"): + await context.get_messages_for_request_retaining( + retain_contents=[], token_budget=1000, hard_fit=True + ) + + assert not emitted + assert context._last_compaction_stats is None + assert not context._removed_seqs + assert context._request_hard_fit is False + + +@pytest.mark.asyncio +async def test_hard_fit_targets_forced_budget_and_stable_view_survives_inactive_fetches(): + """A provider-forced fit is full-budget, while ordinary explicit budgets are not.""" + hard, body = await pressured_context() + canonical = copy.deepcopy(await hard.get_messages()) + + forced = await hard.get_messages_for_request_retaining( + retain_contents=[body], + token_budget=1500, + hard_fit=True, + ) + assert hard._last_compaction_stats is not None + assert hard._last_compaction_stats["target_tokens"] == 1500 + assert hard._request_hard_fit is False + assert await hard.get_messages() == canonical + removed_contents = { + message["content"] + for message in canonical + if message["metadata"]["_seq"] in hard._removed_seqs + } + assert removed_contents, "Setup must record an old sticky removal" + + # The forced view is now the stable provider-facing view even when a later + # caller has no active retention requirement and uses the ordinary getter. + ordinary = await hard.get_messages_for_request() + inactive = await hard.get_messages_for_request_retaining(retain_contents=[]) + assert ordinary == forced == inactive + assert not removed_contents & {message.get("content") for message in ordinary} + + normal, normal_body = await pressured_context() + ordinary_budget = await normal.get_messages_for_request_retaining( + retain_contents=[normal_body], + token_budget=1500, + ) + assert normal._last_compaction_stats is not None + assert normal._last_compaction_stats["target_tokens"] == 750 + assert normal._request_hard_fit is False + assert normal._estimate_tokens(ordinary_budget) <= normal.max_tokens + + empty, _ = await pressured_context() + await empty.get_messages_for_request_retaining( + retain_contents=[], token_budget=1500, hard_fit=True + ) + assert empty._last_compaction_stats is not None + assert empty._last_compaction_stats["target_tokens"] == 1500 + assert empty._request_hard_fit is False + + +@pytest.mark.asyncio +async def test_hard_fit_keeps_new_human_reminder_and_current_tool_result(): + """New growth is visible, but old sticky removals never reinflate.""" + context = SimpleContextManager( + max_tokens=1800, + protected_tool_results=1, + compaction_notice_enabled=False, + ) + first = "Original human task: retain this complete prompt." + required = "" + "active policy " * 150 + "" + await context.add_message(human(first)) + await context.add_message(reminder(required)) + for i in range(30): + await context.add_message( + {"role": "assistant", "content": f"Historical bulk {i}: " + "x" * 900} + ) + + await context.get_messages_for_request_retaining( + retain_contents=[required], token_budget=1500, hard_fit=True + ) + removed_contents = { + message["content"] + for message in await context.get_messages() + if message["metadata"]["_seq"] in context._removed_seqs + } + assert removed_contents, "Setup must record an old sticky removal" + + latest = "Latest human correction: preserve current tool output." + tool_output = "CURRENT_TOOL_RESULT " + "z" * 500 + await context.add_message(human(latest)) + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [ + {"id": "current_call", "type": "function", "function": {"name": "read"}} + ], + } + ) + await context.add_message( + {"role": "tool", "tool_call_id": "current_call", "content": tool_output} + ) + + view = await context.get_messages_for_request_retaining( + retain_contents=[required], token_budget=1500, hard_fit=True + ) + contents = {message.get("content") for message in view} + assert {first, latest, required, tool_output}.issubset(contents) + assert not removed_contents & contents + assert context._request_hard_fit is False + + +@pytest.mark.asyncio +async def test_hard_fit_with_empty_retention_targets_the_full_forced_budget(): + """An empty requirement set must not silently fall back to the half target.""" + context, _ = await pressured_context() + forced_budget = 1500 + + await context.get_messages_for_request_retaining( + retain_contents=[], + token_budget=forced_budget, + hard_fit=True, + ) + + assert context._last_compaction_stats is not None + assert context._last_compaction_stats["target_tokens"] == forced_budget + assert context._request_hard_fit is False + + +@pytest.mark.asyncio +async def test_hard_fit_protects_appended_human_reminder_and_tool_result(): + """A provider-directed fit keeps current protected growth and old reductions.""" + context = SimpleContextManager( + max_tokens=1800, + protected_tool_results=1, + compaction_notice_enabled=False, + ) + first = "Original human task: retain this complete prompt." + required = "" + "active policy " * 150 + "" + forced_budget = 1500 + await context.add_message(human(first)) + await context.add_message(reminder(required)) + for i in range(30): + await context.add_message( + {"role": "assistant", "content": f"Historical bulk {i}: " + "x" * 900} + ) + + await context.get_messages_for_request_retaining( + retain_contents=[required], + token_budget=forced_budget, + hard_fit=True, + ) + removed_contents = { + message["content"] + for message in await context.get_messages() + if message["metadata"]["_seq"] in context._removed_seqs + } + assert removed_contents, "Setup must record an old sticky removal" + + latest = "Latest human correction: preserve current tool output." + tool_output = "CURRENT_TOOL_RESULT " + "z" * 500 + await context.add_message(human(latest)) + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [ + {"id": "current_call", "type": "function", "function": {"name": "read"}} + ], + } + ) + await context.add_message( + {"role": "tool", "tool_call_id": "current_call", "content": tool_output} + ) + + view = await context.get_messages_for_request_retaining( + retain_contents=[required], + token_budget=forced_budget, + hard_fit=True, + ) + contents = {message.get("content") for message in view} + assert {first, latest, required, tool_output}.issubset(contents) + assert not removed_contents & contents + assert context._last_compaction_stats is not None + assert context._last_compaction_stats["target_tokens"] == forced_budget + assert context._request_hard_fit is False @pytest.mark.asyncio diff --git a/tests/test_sticky_compaction_and_tail_notice.py b/tests/test_sticky_compaction_and_tail_notice.py index 8cdabf7..50c7153 100644 --- a/tests/test_sticky_compaction_and_tail_notice.py +++ b/tests/test_sticky_compaction_and_tail_notice.py @@ -826,6 +826,45 @@ async def test_notice_returns_once_tool_results_arrive(): "Notice must not sit directly after an unanswered tool_calls message" ) + # A repeated ordinary fetch must rebuild exactly one stable tail notice, + # rather than retaining the prior ephemeral notice and duplicating it. + repeated_view = await context.get_messages_for_request() + assert repeated_view == resumed_view + assert len(_notices(repeated_view)) == 1 + + +@pytest.mark.asyncio +async def test_notice_waits_for_all_results_of_a_multi_call_group(): + """A partial multi-result group is still an unsafe tail for a notice.""" + context = _make_context() + await _fill_until_compacted(context) + await context.get_messages_for_request() + + await context.add_message({"role": "user", "content": "run both tools"}) + await context.add_message( + { + "role": "assistant", + "content": "", + "tool_calls": [ + {"id": "call_one", "type": "function", "function": {"name": "one"}}, + {"id": "call_two", "type": "function", "function": {"name": "two"}}, + ], + } + ) + await context.add_message( + {"role": "tool", "tool_call_id": "call_one", "content": "first result"} + ) + + partial_view = await context.get_messages_for_request() + assert not _notices(partial_view) + + await context.add_message( + {"role": "tool", "tool_call_id": "call_two", "content": "second result"} + ) + complete_view = await context.get_messages_for_request() + assert len(_notices(complete_view)) == 1 + assert complete_view[-1] is _notices(complete_view)[0] + # --------------------------------------------------------------------------- # (e) `_seq` is internal bookkeeping and must not cross the module boundary From f619f75909530f6b3fdcd4d6d8c430fc44d5e36c Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Mon, 14 Sep 2026 23:09:37 -0700 Subject: [PATCH 2/2] fix: commit hard-fit compaction at event delivery boundary Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 11 +- amplifier_module_context_simple/__init__.py | 49 +++++++-- tests/test_request_retention.py | 115 ++++++++++++++++++++ 3 files changed, 163 insertions(+), 12 deletions(-) diff --git a/README.md b/README.md index 69689ac..5f55ee3 100644 --- a/README.md +++ b/README.md @@ -116,8 +116,8 @@ Quoting reminder XML in an ordinary human prompt does not change its identity. This protects delivery in the request view; it does not change message roles, pin every historical reminder, or rewrite canonical history. A missing required body or an irreducible required set that cannot fit raises `ContextLengthError` -instead of silently dropping instructions. Failed assembly restores the prior -compaction state. +instead of silently dropping instructions. Failed assembly before compaction +event delivery restores the prior compaction state. Protection takes precedence over the compaction target. A protected tool cohort can leave a view above that target; it is not a strict native-token ceiling. @@ -130,6 +130,13 @@ rebuild: it targets the supplied effective request budget rather than applying not a user configuration setting; ordinary `token_budget` calls keep their existing `target_usage` semantics. +For a hard-fit request, sticky compaction decisions and their accounting roll +back if assembly fails or is cancelled before the final, notice-inclusive view +starts `context:compaction` event delivery. Once delivery starts, that validated +compaction is retained even if the caller is cancelled. The event marks this +compaction commit boundary only; it does **not** imply that a provider request +was dispatched. + ### 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 e2e894f..5f5136d 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -351,6 +351,11 @@ def __init__( # transient. It is set only by get_messages_for_request_retaining() # and restored before that capability returns or raises. self._request_hard_fit = False + # A retention request rolls back its sticky decisions until its + # validated hard-fit compaction event is handed to hooks. This flag is + # request-local and is always restored by the retention wrapper. + self._request_compaction_delivery_started = False + self._request_retention_depth = 0 # --- Sticky compaction decision state --- # Compaction decisions (remove / truncate / stub) are keyed by a @@ -592,10 +597,16 @@ async def get_messages_for_request_retaining( the caller supplies the current delivery requirements on every call. `hard_fit=True` is a capability-only request to target the supplied effective budget rather than the normal target_usage fraction. + + A cancelled or failed assembly rolls back sticky decisions and + accounting until validated hard-fit compaction delivery begins. Starting + that delivery commits the view state; it does not imply an LLM/provider + request was dispatched. """ previous_contents = self._request_retained_contents previous_seqs = self._request_protected_seqs previous_hard_fit = self._request_hard_fit + previous_delivery_started = self._request_compaction_delivery_started decisions = ( self._removed_seqs.copy(), self._truncated_seqs.copy(), @@ -604,26 +615,40 @@ async def get_messages_for_request_retaining( self._last_compaction_stats, self._last_token_meter_stats, ) + self._request_retention_depth += 1 try: self._request_retained_contents = frozenset(retain_contents) self._request_hard_fit = hard_fit + self._request_compaction_delivery_started = False return await self.get_messages_for_request(token_budget, provider) except BaseException: - # A failed/cancelled assembly must not commit a reduction that was - # never delivered. Canonical history is unchanged throughout. - ( - self._removed_seqs, - self._truncated_seqs, - self._stubbed_seqs, - self._sticky_level, - self._last_compaction_stats, - self._last_token_meter_stats, - ) = decisions + # A failed/cancelled assembly rolls back state until the validated + # hard-fit compaction event starts delivery. Canonical history is + # unchanged throughout; beginning delivery is not provider dispatch. + if not self._request_compaction_delivery_started: + ( + self._removed_seqs, + self._truncated_seqs, + self._stubbed_seqs, + self._sticky_level, + self._last_compaction_stats, + self._last_token_meter_stats, + ) = decisions raise finally: self._request_retained_contents = previous_contents self._request_protected_seqs = previous_seqs self._request_hard_fit = previous_hard_fit + self._request_retention_depth -= 1 + if self._request_retention_depth: + # A nested call propagates delivery to its outer handler, + # which owns the eventual restoration of transient state. + self._request_compaction_delivery_started = ( + previous_delivery_started + or self._request_compaction_delivery_started + ) + else: + self._request_compaction_delivery_started = previous_delivery_started def _protected_sequences(self, messages: list[dict[str, Any]]) -> set[int]: humans = [msg for msg in messages if _is_human_message(msg)] @@ -888,6 +913,10 @@ async def get_messages_for_request( self._check_retained_budget(compacted, budget) if deferred_hard_fit_stats is not None and self._hooks is not None: try: + # This is the request-retention commit boundary. It follows all + # final, notice-inclusive budget checks but does not mean the + # provider request has been dispatched. + self._request_compaction_delivery_started = True await self._hooks.emit("context:compaction", deferred_hard_fit_stats) except Exception as e: logger.warning(f"Could not emit compaction event: {e}") diff --git a/tests/test_request_retention.py b/tests/test_request_retention.py index c16d694..09714e3 100644 --- a/tests/test_request_retention.py +++ b/tests/test_request_retention.py @@ -1,5 +1,6 @@ """Required request content must survive actual compaction, not just storage.""" +import asyncio import copy import pytest @@ -217,6 +218,120 @@ async def emit(self, event: str, data: dict) -> None: assert context._request_hard_fit is False +@pytest.mark.asyncio +async def test_pre_delivery_cancellation_rolls_back_compaction_request_state(): + """Cancellation before delivery leaves neither an event nor sticky state.""" + emitted: list[tuple[str, dict]] = [] + compaction_suspended = asyncio.Event() + release_compaction = asyncio.Event() + + class Hooks: + async def emit(self, event: str, data: dict) -> None: + emitted.append((event, data)) + + context, body = await pressured_context() + context._hooks = Hooks() + canonical = copy.deepcopy(await context.get_messages()) + previous_state = ( + context._removed_seqs.copy(), + context._truncated_seqs.copy(), + context._stubbed_seqs.copy(), + context._sticky_level, + context._last_compaction_stats, + context._last_token_meter_stats, + ) + original_compact = context._compact_ephemeral + + async def suspend_after_compaction(*args): + compacted = await original_compact(*args) + compaction_suspended.set() + await release_compaction.wait() + return compacted + + context._compact_ephemeral = suspend_after_compaction + task = asyncio.create_task( + context.get_messages_for_request_retaining( + retain_contents=[body], token_budget=1500, hard_fit=True + ) + ) + try: + await asyncio.wait_for(compaction_suspended.wait(), timeout=1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + release_compaction.set() + if not task.done(): + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert not emitted + assert ( + context._removed_seqs, + context._truncated_seqs, + context._stubbed_seqs, + context._sticky_level, + context._last_compaction_stats, + context._last_token_meter_stats, + ) == previous_state + assert context._request_retained_contents == frozenset() + assert context._request_protected_seqs == set() + assert context._request_hard_fit is False + assert context._request_compaction_delivery_started is False + assert await context.get_messages() == canonical + + +@pytest.mark.asyncio +async def test_delivery_boundary_cancellation_commits_compaction_state(): + """Cancellation after event delivery begins keeps the validated hard fit.""" + emitted: list[tuple[str, dict]] = [] + delivery_started = asyncio.Event() + release_delivery = asyncio.Event() + + class Hooks: + async def emit(self, event: str, data: dict) -> None: + emitted.append((event, data)) + delivery_started.set() + await release_delivery.wait() + + context, body = await pressured_context() + context._hooks = Hooks() + canonical = copy.deepcopy(await context.get_messages()) + task = asyncio.create_task( + context.get_messages_for_request_retaining( + retain_contents=[body], token_budget=1500, hard_fit=True + ) + ) + try: + await asyncio.wait_for(delivery_started.wait(), timeout=1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + release_delivery.set() + if not task.done(): + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert [event for event, _ in emitted] == ["context:compaction"] + assert context._last_compaction_stats is not None + assert context._last_token_meter_stats is not None + assert context._removed_seqs or context._truncated_seqs or context._stubbed_seqs + assert context._sticky_level > 0 + assert context._request_retained_contents == frozenset() + assert context._request_protected_seqs == set() + assert context._request_hard_fit is False + assert context._request_compaction_delivery_started is False + assert await context.get_messages() == canonical + + ordinary = await context.get_messages_for_request() + assert ordinary == await context.get_messages_for_request() + assert len(ordinary) < len(canonical) + assert [event for event, _ in emitted] == ["context:compaction"] + + @pytest.mark.asyncio async def test_hard_fit_targets_forced_budget_and_stable_view_survives_inactive_fetches(): """A provider-forced fit is full-budget, while ordinary explicit budgets are not."""