From 37b2c87e3f796ead6ef19a30e9c69679f914c61e Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 20:00:46 -0700 Subject: [PATCH 1/4] fix(context): omit misleading token totals from compaction notices Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_context_simple/__init__.py | 6 ++-- .../test_sticky_compaction_and_tail_notice.py | 36 +++++++++++++++++++ 2 files changed, 38 insertions(+), 4 deletions(-) diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index d4cc625..088378b 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -2404,9 +2404,8 @@ def _format_compaction_notice(self) -> str: removed = stats.get("messages_removed", 0) stubbed = stats.get("user_messages_stubbed", 0) truncated = stats.get("messages_truncated", 0) - old_tokens = stats.get("before_tokens", 0) - new_tokens = stats.get("after_tokens", 0) - target_tokens = stats.get("target_tokens", 0) + # Local chars/4 statistics are not provider input counts; keep them + # in diagnostic stats rather than presenting them as tokens to the model. protected_recent = stats.get("protected_recent", 0.0) protected_tool_results = stats.get("protected_tool_results", 0) @@ -2417,7 +2416,6 @@ def _format_compaction_notice(self) -> str: - Strategy level: {level}/8 - Messages: {old_count} → {new_count} ({removed} removed, {stubbed} stubbed) - Tool results: {truncated} truncated -- Tokens: {old_tokens:,} → {new_tokens:,} (target: {target_tokens:,}) What was preserved: - All system messages (your instructions and identity) diff --git a/tests/test_sticky_compaction_and_tail_notice.py b/tests/test_sticky_compaction_and_tail_notice.py index 50c7153..7d90e4e 100644 --- a/tests/test_sticky_compaction_and_tail_notice.py +++ b/tests/test_sticky_compaction_and_tail_notice.py @@ -232,6 +232,42 @@ async def test_notice_is_still_visible_and_informative(): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("verbosity", ["minimal", "normal", "verbose"]) +async def test_notice_omits_token_totals_but_keeps_diagnostic_stats(verbosity): + """Local compaction estimates must not be presented as provider token counts.""" + context = _make_context(compaction_notice_verbosity=verbosity) + await _fill_until_compacted(context) + + messages = await context.get_messages_for_request() + stats = context._last_compaction_stats + assert stats is not None, "The fixture must actually compact" + for key in ("before_tokens", "after_tokens", "target_tokens"): + assert isinstance(stats[key], int) and stats[key] > 0 + + notices = [ + message + for message in messages + if (message.get("metadata") or {}).get("source") == "context-compaction" + ] + assert len(notices) == 1 + notice = notices[0] + assert messages[-1] is notice + assert notice["role"] == "user" + assert notice["metadata"]["ephemeral"] is True + content = notice["content"] + assert '' in content + assert "Context has been compacted" in content + assert "- Tokens:" not in content + assert "(target:" not in content + if verbosity != "minimal": + assert "- Strategy level:" in content + assert "- Messages:" in content + assert "- Tool results:" in content + assert "What was preserved:" in content + assert "What may be affected:" in content + + @pytest.mark.asyncio async def test_sticky_removal_decisions_never_reversed(): """Once a message is decided as removed, it must stay removed on every From 36688f52c803958037b53715fddf1fca2e7c9ba1 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:08:16 -0700 Subject: [PATCH 2/4] feat(context): add measured request view Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 22 +- amplifier_module_context_simple/__init__.py | 665 +++++++++++++++++++- tests/test_measured_request_view.py | 324 ++++++++++ 3 files changed, 1002 insertions(+), 9 deletions(-) create mode 100644 tests/test_measured_request_view.py diff --git a/README.md b/README.md index c83a3b8..cc9d4f8 100644 --- a/README.md +++ b/README.md @@ -33,7 +33,7 @@ Provides straightforward in-memory conversation context management. This is the - No persistence across sessions - Automatic compaction when approaching token limit (keeps system messages + last 10 messages) - **Preserves tool pairs as atomic units** during compaction (data integrity guarantee) -- **Optional real-usage token meter** (`token_meter: "actual"`, default off) drives the compaction trigger from real provider usage instead of the built-in estimator -- see [Real-usage token meter](#real-usage-token-meter-token_meter) below +- **Optional real-usage token meter** (`token_meter: "actual"`, default off) can negotiate a complete-request provider count with a compatible orchestrator; older orchestrators retain the legacy actual-meter path -- see [Real-usage token meter](#real-usage-token-meter-token_meter) below ## Configuration @@ -356,7 +356,25 @@ own reported usage instead of guessing. This module ports that same full pre-existing test suite unchanged. An unrecognized `token_meter` value logs a warning and falls back to `"estimate"` rather than raising. -### Known, accepted limitation +### Negotiated provider-count path + +When `token_meter: "actual"` is paired with an orchestrator and provider that +both advertise the optional measured-request capability, Context supplies a +notice-inclusive candidate view and the orchestrator builds and counts the +complete request, including its own tools and overlays. Context uses the +provider's raw input count for both the configured trigger and target, applies +at most the existing eight legal rungs, and returns the exact final counted +envelope for dispatch. A protected floor is reported as such rather than being +called target success; a known request over the provider hard limit fails +closed. + +This is negotiated rather than a Context-side provider call: Context never +imports provider or Loop types, and cannot count overlays it does not own. The +returned transaction is committed by the orchestrator immediately before it +dispatches the already-counted request; otherwise it rolls back staged sticky +decisions. + +### Legacy actual-meter path Only the **escalation gate** (whether to compact at all, and whether a sticky escalation needs to advance) uses the real measurement in `"actual"` diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 088378b..37492fc 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -259,6 +259,14 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = register_capability( "context.request_retention", context.get_messages_for_request_retaining ) + register_capability("context.foreground_usage", context.claim_foreground_usage) + # This is deliberately absent in estimate mode. The loop must opt in + # only when both its provider and Context can make the same complete + # request count; the legacy getter remains the estimate-mode path. + if token_meter == TOKEN_METER_ACTUAL: + register_capability( + "context.measured_request_view", context.get_measured_request_view + ) logger.info(f"Mounted SimpleContextManager (token_meter={token_meter!r})") async def cleanup() -> None: @@ -412,6 +420,12 @@ def __init__( # "estimate" mode -- see README "Real-usage token meter". self._last_measured_prompt_tokens: int | None = None self._last_token_meter_stats: dict[str, Any] | None = None + # A Loop that understands foreground ownership claims this once, then + # writes only its successful foreground response usage through the + # recorder. Generic llm:response remains a compatibility fallback + # until that persistent claim, but can never overwrite owned data. + self._foreground_usage_claimed = False + self._foreground_usage_stale = False self._system_prompt_factory: Callable[[], Awaitable[str]] | None = None self._request_retained_contents: frozenset[str] = frozenset() self._request_protected_seqs: set[int] = set() @@ -424,6 +438,9 @@ def __init__( # request-local and is always restored by the retention wrapper. self._request_compaction_delivery_started = False self._request_retention_depth = 0 + # The measured capability stages the legacy fallback too, so a + # callback cancellation cannot publish a compaction no request sent. + self._defer_compaction_delivery = False # --- Sticky compaction decision state --- # Compaction decisions (remove / truncate / stub) are keyed by a @@ -759,6 +776,565 @@ async def get_messages_for_request_retaining( else: self._request_compaction_delivery_started = previous_delivery_started + def claim_foreground_usage(self) -> Callable[..., bool]: + """Claim the persistent foreground usage meter and return its recorder. + + This capability is intentionally not a temporary dispatch lease: + helper/naming/evaluator responses can arrive after a prompt completes, + so a ContextVar or an around-send flag would reopen the generic hook + race. The first claim discards an unowned hook reading; later claims + are idempotent and return the same stateful recorder. + """ + if not self._foreground_usage_claimed: + self._foreground_usage_claimed = True + self._last_measured_prompt_tokens = None + self._foreground_usage_stale = False + return self._record_foreground_usage + + def _record_foreground_usage( + self, *, input_tokens: Any, cache_write_tokens: Any = 0 + ) -> bool: + """Record one scoped foreground response, refusing malformed usage.""" + valid = ( + isinstance(input_tokens, int) + and not isinstance(input_tokens, bool) + and input_tokens >= 0 + and isinstance(cache_write_tokens, int) + and not isinstance(cache_write_tokens, bool) + and cache_write_tokens >= 0 + ) + if not valid: + # Keep an earlier owned reading, but state plainly that it is not a + # fresh scoped value. Turning malformed input into zero would be a + # false measurement and could suppress necessary compaction. + self._foreground_usage_stale = True + logger.debug( + "context-simple: foreground usage was malformed; retaining the " + "previous owned reading as stale" + ) + return False + self._last_measured_prompt_tokens = input_tokens + cache_write_tokens + self._foreground_usage_stale = False + return True + + @staticmethod + def _is_nonnegative_int(value: Any) -> bool: + return isinstance(value, int) and not isinstance(value, bool) and value >= 0 + + def _measured_transaction(self) -> "_MeasuredViewTransaction": + """Snapshot all sticky decisions which a measured candidate may stage.""" + return _MeasuredViewTransaction(self) + + def _append_compaction_notice( + self, messages: list[dict[str, Any]] + ) -> list[dict[str, Any]]: + """Return the one safe, qualitative tail notice for a measured view.""" + result = list(messages) + if ( + self.compaction_notice_enabled + and self._last_compaction_stats + and self._last_compaction_stats.get("strategy_level", 0) + >= self.compaction_notice_min_level + and not self._tail_has_unanswered_tool_calls(result) + ): + notice = self._format_compaction_notice() + if notice: + result.append( + { + "role": "user", + "content": notice, + "metadata": { + "source": "context-compaction", + "ephemeral": True, + }, + } + ) + return result + + def _measured_budget_decision( + self, envelope: dict[str, Any] + ) -> tuple[int, int, int, str] | None: + """Validate the additive native-count fields or report unavailability. + + A non-dictionary callback envelope is a broken capability contract, not + an unavailable measurement. Missing/malformed *advertised* + measurement, however, is explicitly an unavailable count so an older + provider can retain its compatibility path. + """ + decision = envelope.get("budget_decision") + if decision is None: + return None + if not isinstance(decision, dict): + raise TypeError("count_view returned a non-dictionary budget_decision") + measurement = decision.get("measurement") + if not isinstance(measurement, dict): + return None + count = measurement.get("input_tokens") + estimated = decision.get("estimated_input_tokens") + limit = decision.get("input_limit_tokens") + if ( + measurement.get("kind") != "provider_count" + or not isinstance(measurement.get("source"), str) + or not measurement["source"] + or not self._is_nonnegative_int(count) + or not self._is_nonnegative_int(estimated) + or not self._is_nonnegative_int(limit) + or estimated < count + ): + return None + return count, estimated, limit, measurement["source"] + + async def _count_measured_view( + self, count_view: Callable[[list[dict[str, Any]]], Awaitable[dict[str, Any]]], + view: list[dict[str, Any]], + ) -> tuple[ + list[dict[str, Any]], + dict[str, Any], + tuple[int, int, int, str] | None, + ]: + """Count exactly the provider-facing candidate, without reassembly.""" + public_view = self._strip_internal_metadata(self._append_compaction_notice(view)) + envelope = await count_view(public_view) + if not isinstance(envelope, dict) or "dispatch" not in envelope: + raise TypeError("count_view must return {dispatch, budget_decision}") + return public_view, envelope, self._measured_budget_decision(envelope) + + def _measured_apply_rung( + self, rung: int, view: list[dict[str, Any]] + ) -> tuple[list[dict[str, Any]], bool]: + """Apply one existing legal compaction rung without local success claims.""" + systems = [dict(msg) for msg in view if msg.get("role") == "system"] + working = [dict(msg) for msg in view if msg.get("role") != "system"] + before = ( + self._removed_seqs.copy(), + self._truncated_seqs.copy(), + self._stubbed_seqs.copy(), + ) + tool_indices = [ + index for index, message in enumerate(working) if message.get("role") == "tool" + ] + protected_tools = self._protected_tool_indices(tool_indices) + + def truncate(indices: list[int]) -> None: + # Unlike the legacy local-estimate helper, this exhausts the + # rung. Only an actual transformation becomes a sticky decision, + # so a short/protected wave remains a genuine no-op. + for index in indices: + if index in protected_tools: + continue + message = working[index] + if message.get("_truncated"): + continue + reduced = self._truncate_tool_result(message) + if reduced is message: + continue + working[index] = reduced + self._record_truncated(message) + + if rung == 1: + truncate(tool_indices[: int(len(tool_indices) * 0.25)]) + elif rung == 2: + truncate( + tool_indices[ + int(len(tool_indices) * 0.25) : int(len(tool_indices) * 0.50) + ] + ) + elif rung in (3, 5, 7): + protection = { + 3: self.protected_recent, + 5: self.protected_recent * 0.6, + 7: self.protected_recent * 0.3, + }[rung] + working, _, _, _ = self._remove_messages_with_protection( + working, + -1, # exhaust legal candidates; count_view decides success + protected_recent=protection, + system_tokens=self._estimate_tokens(systems), + ) + elif rung == 4: + truncate( + tool_indices[ + int(len(tool_indices) * 0.50) : int(len(tool_indices) * 0.75) + ] + ) + elif rung == 6: + truncate(tool_indices) + elif rung == 8: + user_indices = [ + index + for index, message in enumerate(working) + if message.get("role") == "user" + ] + first_user = user_indices[0] if user_indices else None + last_user = user_indices[-1] if user_indices else None + # This is the legacy level-eight machine-prefix rule, deliberately + # retaining actual human boundaries and request-retained bodies. + if first_user is not None and first_user != last_user: + first = working[first_user] + reduced = self._stub_user_message(first) + if ( + not self._is_request_protected(first) + and not _is_human_message(first) + and not first.get("_stubbed") + and reduced is not first + ): + working[first_user] = reduced + self._record_stubbed(first) + protected_boundary = int(len(working) * (1 - self.protected_recent * 0.3)) + for index, message in list(enumerate(working)): + if ( + index >= protected_boundary + or index == last_user + or not message.get("_stubbed") + or self._is_request_protected(message) + ): + continue + self._record_removed(message) + working = [ + message + for message in working + if self._extract_seq(message) not in self._removed_seqs + ] + else: # Defensive boundary for this private fixed eight-rung policy. + raise ValueError(f"unknown measured compaction rung {rung}") + + after = (self._removed_seqs, self._truncated_seqs, self._stubbed_seqs) + changed = after != before + return systems + working, changed + + def _stage_measured_stats( + self, + *, + initial_view: list[dict[str, Any]], + final_view: list[dict[str, Any]], + strategy_level: int, + budget: int, + measured_before: int, + measured_after: int, + measurement_source: str, + trigger: float, + target: int, + outcome: str, + count_calls: int, + ) -> None: + """Stage telemetry for a candidate; transaction commit makes it sticky.""" + self._sticky_level = max(self._sticky_level, strategy_level) + self._last_compaction_stats = { + "before_tokens": self._estimate_tokens(initial_view), + "after_tokens": self._estimate_tokens(final_view), + "before_messages": len(initial_view), + "after_messages": len(final_view), + "messages_removed": len(self._removed_seqs), + "messages_truncated": len(self._truncated_seqs), + "user_messages_stubbed": len(self._stubbed_seqs), + "system_messages_preserved": sum( + message.get("role") == "system" for message in final_view + ), + "strategy_level": self._sticky_level, + "budget": budget, + "target_tokens": target, + "protected_recent": self.protected_recent, + "protected_tool_results": self.protected_tool_results, + "measurement_kind": "provider_count", + "measurement_source": measurement_source, + "measured_before": measured_before, + "measured_after": measured_after, + "trigger": trigger, + "policy_budget": budget, + "outcome": outcome, + "count_calls": count_calls, + } + + async def get_measured_request_view( + self, + *, + provider: Any, + retain_contents: list[str], + count_view: Callable[[list[dict[str, Any]]], Awaitable[dict[str, Any]]], + ) -> dict[str, Any]: + """Build and reduce one complete request using provider-native counts. + + The loop owns request construction. Context owns only canonical + history, its one materialized policy snapshot, and the established + eight legal reduction rungs. The returned envelope is the exact + callback object that counted the final provider-facing request. + """ + if self.token_meter != TOKEN_METER_ACTUAL: + raise RuntimeError("measured request views require token_meter='actual'") + + previous_contents = self._request_retained_contents + previous_seqs = self._request_protected_seqs + previous_hard_fit = self._request_hard_fit + previous_delivery = self._request_compaction_delivery_started + cycle_snapshot = ( + self._removed_seqs.copy(), + self._truncated_seqs.copy(), + self._stubbed_seqs.copy(), + self._sticky_level, + self._last_compaction_stats, + self._last_token_meter_stats, + ) + transaction: _MeasuredViewTransaction | None = None + try: + self._request_retained_contents = frozenset(retain_contents) + self._request_hard_fit = False + self._request_compaction_delivery_started = False + + # Materialize the dynamic factory exactly once. The callback adds + # all loop-owned overlays to every candidate from this same base. + if self._system_prompt_factory: + system_content = await self._system_prompt_factory() + working = [{"role": "system", "content": system_content}] + [ + message + for message in self.messages + if message.get("role") != "system" + or (message.get("metadata") or {}).get("source") == "hook" + ] + else: + working = list(self.messages) + + self._request_protected_seqs = self._protected_sequences(working) + budget = self._calculate_budget(None, provider) + if self.compaction_notice_enabled: + effective_budget = budget - self.compaction_notice_token_reserve + if effective_budget <= 0: + logger.warning( + "compaction_notice_token_reserve consumes measured request " + "budget; ignoring it for this request" + ) + effective_budget = budget + else: + effective_budget = budget + self._last_effective_budget = effective_budget + trigger = effective_budget * self.compact_threshold + target = int(effective_budget * self.target_usage) + initial = [ + *[dict(message) for message in working if message.get("role") == "system"], + *self._apply_sticky_decisions( + [message for message in working if message.get("role") != "system"] + ), + ] + initial_view, initial_attempt, initial_measurement = await self._count_measured_view( + count_view, initial + ) + if initial_measurement is None: + # Native count availability is dynamic. Preserve the old + # actual-meter/estimate behavior from this already-created + # snapshot; never call the factory or replay loop overlays. + legacy_count, _, legacy_estimate = self._measure_working_tokens(initial) + legacy_needed = self._should_compact( + legacy_count, effective_budget + ) or ( + bool(self._request_retained_contents) + and legacy_estimate > effective_budget + ) + if legacy_needed: + transaction = _MeasuredViewTransaction(self, snapshot=cycle_snapshot) + self._defer_compaction_delivery = True + try: + fallback = await self._compact_ephemeral(effective_budget, working) + finally: + self._defer_compaction_delivery = False + changed = ( + self._removed_seqs, + self._truncated_seqs, + self._stubbed_seqs, + ) != cycle_snapshot[:3] + if changed: + fallback_view, fallback_attempt, fallback_measurement = ( + await self._count_measured_view(count_view, fallback) + ) + if ( + fallback_measurement is not None + and fallback_measurement[1] > fallback_measurement[2] + ): + transaction.rollback() + raise ContextLengthError( + "Context cannot fit protected content within the " + "provider input limit." + ) + return { + "base_view": fallback_view, + "final_attempt": fallback_attempt, + "outcome": "measurement_unavailable", + "measured_before": None, + "measured_after": None, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": 2, + "transaction": transaction, + } + transaction.rollback() + transaction = None + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "measurement_unavailable", + "measured_before": None, + "measured_after": None, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": 1, + "transaction": None, + } + + before, estimated, limit, measurement_source = initial_measurement + needs_compaction = before >= trigger or estimated > limit + if not needs_compaction: + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "not_needed", + "measured_before": before, + "measured_after": before, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": 1, + "transaction": None, + } + + transaction = _MeasuredViewTransaction(self, snapshot=cycle_snapshot) + candidate = initial + final_view = initial_view + final_attempt = initial_attempt + final_measurement = initial_measurement + count_calls = 1 + changed_any = False + last_changed_rung = 0 + for rung in range(1, 9): + candidate, changed = self._measured_apply_rung(rung, candidate) + if not changed: + continue + changed_any = True + last_changed_rung = rung + # Stage qualitative notice/statistics before counting this + # complete candidate; never append a post-count outcome notice. + self._stage_measured_stats( + initial_view=initial, + final_view=candidate, + strategy_level=rung, + budget=effective_budget, + measured_before=before, + measured_after=final_measurement[0], + measurement_source=final_measurement[3], + trigger=trigger, + target=target, + outcome="compacting", + count_calls=count_calls + 1, + ) + final_view, final_attempt, final_measurement = await self._count_measured_view( + count_view, candidate + ) + count_calls += 1 + if final_measurement is None: + transaction.rollback() + if estimated > limit: + raise ContextLengthError( + "Context cannot fit protected content within the provider " + "input limit after its recount became unavailable." + ) + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "measurement_unavailable", + "measured_before": before, + "measured_after": before, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": count_calls, + "transaction": None, + } + measured_after, estimated_after, limit_after, measured_source = final_measurement + if measured_after <= target and estimated_after <= limit_after: + self._stage_measured_stats( + initial_view=initial, + final_view=candidate, + strategy_level=rung, + budget=effective_budget, + measured_before=before, + measured_after=measured_after, + measurement_source=measured_source, + trigger=trigger, + target=target, + outcome="target_reached", + count_calls=count_calls, + ) + return { + "base_view": final_view, + "final_attempt": final_attempt, + "outcome": "target_reached", + "measured_before": before, + "measured_after": measured_after, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": count_calls, + "transaction": transaction, + } + + measured_after, estimated_after, limit_after, measured_source = final_measurement + if estimated_after > limit_after: + transaction.rollback() + raise ContextLengthError( + "Context cannot fit protected content within the provider input " + "limit; required content was not discarded." + ) + if not changed_any: + # A protected floor is a real answer to the one exact probe, + # not a new compaction decision. Do not manufacture sticky + # stats/notices/events (or a transaction) from a no-op walk. + transaction.rollback() + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "protected_floor", + "measured_before": before, + "measured_after": before, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": count_calls, + "transaction": None, + } + self._stage_measured_stats( + initial_view=initial, + final_view=candidate, + strategy_level=last_changed_rung, + budget=effective_budget, + measured_before=before, + measured_after=measured_after, + measurement_source=measured_source, + trigger=trigger, + target=target, + outcome="protected_floor", + count_calls=count_calls, + ) + return { + "base_view": final_view, + "final_attempt": final_attempt, + "outcome": "protected_floor", + "measured_before": before, + "measured_after": measured_after, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": count_calls, + "transaction": transaction, + } + except BaseException: + if transaction is not None: + transaction.rollback() + raise + finally: + self._request_retained_contents = previous_contents + self._request_protected_seqs = previous_seqs + self._request_hard_fit = previous_hard_fit + self._request_compaction_delivery_started = previous_delivery + def _protected_sequences(self, messages: list[dict[str, Any]]) -> set[int]: humans = [msg for msg in messages if _is_human_message(msg)] protected = [humans[0], humans[-1]] if humans else [] @@ -929,6 +1505,7 @@ async def get_messages_for_request( "used_tokens": token_count, "estimated_tokens": estimated_tokens, "measured_tokens": self._last_measured_prompt_tokens, + "foreground_usage_stale": self._foreground_usage_stale, "budget": effective_budget, "ratio": (token_count / effective_budget) if effective_budget > 0 else None, } @@ -1131,6 +1708,9 @@ async def clear(self) -> None: self._last_compaction_stats = None self._last_measured_prompt_tokens = None self._last_token_meter_stats = None + # Ownership is a capability claim, not a turn-local meter value. A + # cleared Context must not permit utility hook traffic to take it back. + self._foreground_usage_stale = self._foreground_usage_claimed logger.info("Context cleared") async def should_compact(self) -> bool: @@ -1214,7 +1794,12 @@ def _measure_working_tokens( self.token_meter == TOKEN_METER_ACTUAL and self._last_measured_prompt_tokens is not None ): - return self._last_measured_prompt_tokens, "measured", estimated_tokens + source = ( + "owned_stale" + if self._foreground_usage_claimed and self._foreground_usage_stale + else "measured" + ) + return self._last_measured_prompt_tokens, source, estimated_tokens return estimated_tokens, "estimate", estimated_tokens async def _on_llm_response(self, event: str, data: dict[str, Any]) -> Any: @@ -1251,16 +1836,27 @@ async def _on_llm_response(self, event: str, data: dict[str, Any]) -> Any: """ from amplifier_core.models import HookResult + if self._foreground_usage_claimed: + # The Loop has claimed responsibility for successful foreground + # responses. This broad hook also sees utility replies and may + # run after the foreground turn, so it must become permanently + # observational once ownership exists. + return HookResult(action="continue") + usage = (data or {}).get("usage") or {} input_tokens = usage.get("input_tokens") - cache_write_tokens = usage.get("cache_write_tokens") or 0 - if isinstance(input_tokens, int | float): - total = int(input_tokens) + int(cache_write_tokens) + cache_write_tokens = ( + usage["cache_write_tokens"] if "cache_write_tokens" in usage else 0 + ) + if self._is_nonnegative_int(input_tokens) and self._is_nonnegative_int( + cache_write_tokens + ): + total = input_tokens + cache_write_tokens self._last_measured_prompt_tokens = total logger.debug( f"context-simple: token_meter recorded real usage from " - f"llm:response -- input_tokens={int(input_tokens):,} + " - f"cache_write_tokens={int(cache_write_tokens):,} = {total:,} total" + f"llm:response -- input_tokens={input_tokens:,} + " + f"cache_write_tokens={cache_write_tokens:,} = {total:,} total" ) else: logger.debug( @@ -2322,7 +2918,11 @@ async def _finalize_compaction_with_stats( # 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: + if ( + self._hooks is not None + and not self._request_hard_fit + and not self._defer_compaction_delivery + ): try: await self._hooks.emit("context:compaction", stats) except Exception as e: @@ -2587,3 +3187,54 @@ def _derive_budget(self, token_budget: int | None, provider: Any | None) -> int: def _estimate_tokens(self, messages: list[dict[str, Any]]) -> int: """Rough token estimation (chars / 4).""" return sum(len(str(msg)) // 4 for msg in messages) + + +class _MeasuredViewTransaction: + """Idempotent commit/rollback boundary for staged measured compaction.""" + + def __init__( + self, + context: SimpleContextManager, + snapshot: tuple[set[int], set[int], set[int], int, dict[str, Any] | None, dict[str, Any] | None] + | None = None, + ): + self._context = context + self._snapshot = snapshot or ( + context._removed_seqs.copy(), + context._truncated_seqs.copy(), + context._stubbed_seqs.copy(), + context._sticky_level, + context._last_compaction_stats, + context._last_token_meter_stats, + ) + self._terminal = False + + async def commit(self) -> None: + """Accept the selected candidate immediately before Loop dispatch.""" + if self._terminal: + return + if self._context._hooks is not None and self._context._last_compaction_stats: + try: + await self._context._hooks.emit( + "context:compaction", self._context._last_compaction_stats + ) + except Exception as error: + logger.warning(f"Could not emit measured compaction event: {error}") + # Keep rollback possible while event delivery is awaitable. In + # particular, cancellation here means Loop never dispatches and its + # finally block must restore the staged decisions. + self._terminal = True + + def rollback(self) -> None: + """Restore pre-candidate state unless the transaction was committed.""" + if self._terminal: + return + ( + self._context._removed_seqs, + self._context._truncated_seqs, + self._context._stubbed_seqs, + self._context._sticky_level, + self._context._last_compaction_stats, + self._context._last_token_meter_stats, + ) = self._snapshot + self._terminal = True diff --git a/tests/test_measured_request_view.py b/tests/test_measured_request_view.py new file mode 100644 index 0000000..89ac042 --- /dev/null +++ b/tests/test_measured_request_view.py @@ -0,0 +1,324 @@ +"""Focused contract tests for Context's provider-count request-view capability.""" + +import asyncio + +import pytest +from amplifier_core.llm_errors import ContextLengthError + +from amplifier_module_context_simple import SimpleContextManager, mount + + +class _Coordinator: + def __init__(self): + self.capabilities = {} + + async def mount(self, kind, instance): + self.context = instance + + def register_capability(self, name, capability): + self.capabilities[name] = capability + + +def _decision(count, *, estimated=None, limit=1_000): + return { + "estimated_input_tokens": count if estimated is None else estimated, + "input_limit_tokens": limit, + "measurement": { + "kind": "provider_count", + "source": "test.provider.count", + "input_tokens": count, + }, + } + + +@pytest.mark.asyncio +async def test_measured_capability_is_actual_mode_only_and_foreground_is_always_additive(): + estimate = _Coordinator() + await mount(estimate, {"token_meter": "estimate"}) + assert "context.measured_request_view" not in estimate.capabilities + assert callable(estimate.capabilities["context.foreground_usage"]) + + actual = _Coordinator() + await mount(actual, {"token_meter": "actual"}) + assert callable(actual.capabilities["context.measured_request_view"]) + + +@pytest.mark.asyncio +async def test_provider_count_controls_trigger_target_and_returns_exact_final_envelope(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first human boundary"}) + for index in range(4): + await context.add_message( + {"role": "tool", "tool_call_id": str(index), "content": "x" * 800} + ) + await context.add_message({"role": "user", "content": "latest human boundary"}) + + calls = [] + + async def count_view(view): + calls.append(view) + # chars/4 remains very different from this native count. The first + # native count triggers, and one legal truncation rung reaches target. + count = 900 if not any(message.get("_truncated") for message in view) else 400 + return {"dispatch": object(), "budget_decision": _decision(count)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "target_reached" + assert result["measured_before"] == 900 + assert result["measured_after"] == 400 + assert result["trigger"] == 800 + assert result["target"] == 500 + assert result["count_calls"] == 2 + assert result["final_attempt"]["dispatch"] is not None + assert result["base_view"] is calls[-1] + assert result["transaction"] is not None + assert context._last_compaction_stats["measurement_source"] == "test.provider.count" + await result["transaction"].commit() + + +@pytest.mark.asyncio +async def test_protected_floor_with_no_legal_change_has_one_count_and_no_sticky_event_state(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "only protected human prompt"}) + calls = 0 + + async def count_view(view): + nonlocal calls + calls += 1 + return {"dispatch": object(), "budget_decision": _decision(900)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "protected_floor" + assert result["count_calls"] == calls == 1 + assert result["transaction"] is None + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_unavailable_recount_rolls_back_to_original_counted_view(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for index in range(4): + await context.add_message({"role": "tool", "content": "x" * 800, "name": str(index)}) + await context.add_message({"role": "user", "content": "last"}) + views = [] + + async def count_view(view): + views.append(view) + decision = _decision(900) if len(views) == 1 else None + return {"dispatch": object(), "budget_decision": decision} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "measurement_unavailable" + assert result["base_view"] is views[0] + assert result["count_calls"] == 2 + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_initial_unavailable_uses_one_snapshot_legacy_fallback_and_recounts_it(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + factory_calls = 0 + + async def factory(): + nonlocal factory_calls + factory_calls += 1 + return "dynamic system" + + await context.set_system_prompt_factory(factory) + calls = [] + + async def count_view(view): + calls.append(view) + return {"dispatch": object(), "budget_decision": None} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "measurement_unavailable" + assert result["count_calls"] == len(calls) == 2 + assert factory_calls == 1 + assert result["transaction"] is not None + result["transaction"].rollback() + assert not context._truncated_seqs + + +@pytest.mark.asyncio +async def test_unavailable_recount_of_known_hard_oversize_fails_closed(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000) + if calls == 1 + else None, + } + + with pytest.raises(ContextLengthError, match="recount became unavailable"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert not context._truncated_seqs + + +@pytest.mark.asyncio +async def test_cancellation_from_count_callback_propagates_without_sticky_state(): + context = SimpleContextManager(max_tokens=1_000, token_meter="actual") + await context.add_message({"role": "user", "content": "prompt"}) + + async def cancelled_count(_view): + raise asyncio.CancelledError + + with pytest.raises(asyncio.CancelledError): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=cancelled_count + ) + assert not context._removed_seqs + assert not context._truncated_seqs + + +@pytest.mark.asyncio +async def test_cancellation_during_commit_leaves_staged_decisions_rollbackable(): + started = asyncio.Event() + release = asyncio.Event() + + class Hooks: + async def emit(self, _event, _data): + started.set() + await release.wait() + + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + hooks=Hooks(), + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + + async def count_view(view): + count = 400 if any(message.get("_truncated") for message in view) else 900 + return {"dispatch": object(), "budget_decision": _decision(count)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + task = asyncio.create_task(result["transaction"].commit()) + await asyncio.wait_for(started.wait(), timeout=1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + release.set() + result["transaction"].rollback() + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_foreground_claim_blocks_generic_hook_after_owned_recording(): + context = SimpleContextManager(token_meter="actual") + await context._on_llm_response("llm:response", {"usage": {"input_tokens": 17}}) + recorder = context.claim_foreground_usage() + assert context._last_measured_prompt_tokens is None + assert recorder(input_tokens=20, cache_write_tokens=3) + await context._on_llm_response("llm:response", {"usage": {"input_tokens": 999}}) + assert context._last_measured_prompt_tokens == 23 + + +@pytest.mark.asyncio +async def test_foreground_recorder_rejects_malformed_usage_without_replacing_owned_reading(): + context = SimpleContextManager() + recorder = context.claim_foreground_usage() + assert recorder(input_tokens=4) + assert not recorder(input_tokens=True) + assert context._last_measured_prompt_tokens == 4 + assert context._foreground_usage_stale is True + + +@pytest.mark.asyncio +async def test_owned_stale_reading_is_labeled_not_fresh_measured_usage(): + context = SimpleContextManager(token_meter="actual", max_tokens=10_000) + recorder = context.claim_foreground_usage() + assert recorder(input_tokens=4) + assert not recorder(input_tokens=-1) + await context.add_message({"role": "user", "content": "prompt"}) + await context.get_messages_for_request() + assert context._last_token_meter_stats["source"] == "owned_stale" + assert context._last_token_meter_stats["foreground_usage_stale"] is True + + +@pytest.mark.asyncio +async def test_unowned_generic_hook_keeps_legacy_meter_behavior(): + context = SimpleContextManager() + await context._on_llm_response("llm:response", {"usage": {"input_tokens": 12}}) + assert context._last_measured_prompt_tokens == 12 + + +@pytest.mark.asyncio +async def test_unowned_hook_rejects_malformed_cache_write_without_crashing(): + context = SimpleContextManager() + await context._on_llm_response("llm:response", {"usage": {"input_tokens": 12}}) + await context._on_llm_response( + "llm:response", {"usage": {"input_tokens": 99, "cache_write_tokens": "bad"}} + ) + assert context._last_measured_prompt_tokens == 12 \ No newline at end of file From 5bb2858f17dbcb9c1c0fdc91e45d75e4e5ec9c15 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:32:40 -0700 Subject: [PATCH 3/4] fix(context): fail closed measured safety Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_context_simple/__init__.py | 119 ++++++-- tests/test_measured_request_view.py | 314 +++++++++++++++++++- 2 files changed, 407 insertions(+), 26 deletions(-) diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 37492fc..a132da9 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -854,24 +854,33 @@ def _append_compaction_notice( def _measured_budget_decision( self, envelope: dict[str, Any] ) -> tuple[int, int, int, str] | None: - """Validate the additive native-count fields or report unavailability. + """Validate old hard-safety fields and additive native-count fields. A non-dictionary callback envelope is a broken capability contract, not an unavailable measurement. Missing/malformed *advertised* measurement, however, is explicitly an unavailable count so an older - provider can retain its compatibility path. + provider can retain its compatibility path. The pre-existing + estimated-input/limit pair remains mandatory whenever a decision is + supplied, regardless of measurement availability. """ decision = envelope.get("budget_decision") if decision is None: return None if not isinstance(decision, dict): raise TypeError("count_view returned a non-dictionary budget_decision") + estimated = decision.get("estimated_input_tokens") + limit = decision.get("input_limit_tokens") + if not self._is_nonnegative_int(estimated) or not self._is_nonnegative_int( + limit + ): + raise TypeError( + "count_view budget_decision requires nonnegative integer " + "estimated_input_tokens and input_limit_tokens" + ) measurement = decision.get("measurement") if not isinstance(measurement, dict): return None count = measurement.get("input_tokens") - estimated = decision.get("estimated_input_tokens") - limit = decision.get("input_limit_tokens") if ( measurement.get("kind") != "provider_count" or not isinstance(measurement.get("source"), str) @@ -884,6 +893,14 @@ def _measured_budget_decision( return None return count, estimated, limit, measurement["source"] + def _measured_hard_oversize(self, envelope: dict[str, Any]) -> bool: + """Whether a validated callback envelope exceeds its old hard limit.""" + decision = envelope.get("budget_decision") + return ( + isinstance(decision, dict) + and decision["estimated_input_tokens"] > decision["input_limit_tokens"] + ) + async def _count_measured_view( self, count_view: Callable[[list[dict[str, Any]]], Awaitable[dict[str, Any]]], view: list[dict[str, Any]], @@ -900,7 +917,10 @@ async def _count_measured_view( return public_view, envelope, self._measured_budget_decision(envelope) def _measured_apply_rung( - self, rung: int, view: list[dict[str, Any]] + self, + rung: int, + view: list[dict[str, Any]], + developer_seqs: set[int], ) -> tuple[list[dict[str, Any]], bool]: """Apply one existing legal compaction rung without local success claims.""" systems = [dict(msg) for msg in view if msg.get("role") == "system"] @@ -950,6 +970,7 @@ def truncate(indices: list[int]) -> None: -1, # exhaust legal candidates; count_view decides success protected_recent=protection, system_tokens=self._estimate_tokens(systems), + additional_protected_seqs=developer_seqs, ) elif rung == 4: truncate( @@ -973,7 +994,7 @@ def truncate(indices: list[int]) -> None: first = working[first_user] reduced = self._stub_user_message(first) if ( - not self._is_request_protected(first) + not self._is_request_protected(first, developer_seqs) and not _is_human_message(first) and not first.get("_stubbed") and reduced is not first @@ -986,7 +1007,7 @@ def truncate(indices: list[int]) -> None: index >= protected_boundary or index == last_user or not message.get("_stubbed") - or self._is_request_protected(message) + or self._is_request_protected(message, developer_seqs) ): continue self._record_removed(message) @@ -1093,6 +1114,12 @@ async def get_measured_request_view( else: working = list(self.messages) + developer_seqs = { + seq + for message in working + if message.get("role") == "developer" + and (seq := self._extract_seq(message)) is not None + } self._request_protected_seqs = self._protected_sequences(working) budget = self._calculate_budget(None, provider) if self.compaction_notice_enabled: @@ -1111,13 +1138,19 @@ async def get_measured_request_view( initial = [ *[dict(message) for message in working if message.get("role") == "system"], *self._apply_sticky_decisions( - [message for message in working if message.get("role") != "system"] + [message for message in working if message.get("role") != "system"], + additional_protected_seqs=developer_seqs, ), ] initial_view, initial_attempt, initial_measurement = await self._count_measured_view( count_view, initial ) if initial_measurement is None: + if self._measured_hard_oversize(initial_attempt): + raise ContextLengthError( + "Context cannot fit protected content within the provider " + "input limit." + ) # Native count availability is dynamic. Preserve the old # actual-meter/estimate behavior from this already-created # snapshot; never call the factory or replay loop overlays. @@ -1132,7 +1165,11 @@ async def get_measured_request_view( transaction = _MeasuredViewTransaction(self, snapshot=cycle_snapshot) self._defer_compaction_delivery = True try: - fallback = await self._compact_ephemeral(effective_budget, working) + fallback = await self._compact_ephemeral( + effective_budget, + working, + additional_protected_seqs=developer_seqs, + ) finally: self._defer_compaction_delivery = False changed = ( @@ -1141,13 +1178,10 @@ async def get_measured_request_view( self._stubbed_seqs, ) != cycle_snapshot[:3] if changed: - fallback_view, fallback_attempt, fallback_measurement = ( + fallback_view, fallback_attempt, _fallback_measurement = ( await self._count_measured_view(count_view, fallback) ) - if ( - fallback_measurement is not None - and fallback_measurement[1] > fallback_measurement[2] - ): + if self._measured_hard_oversize(fallback_attempt): transaction.rollback() raise ContextLengthError( "Context cannot fit protected content within the " @@ -1181,6 +1215,7 @@ async def get_measured_request_view( } before, estimated, limit, measurement_source = initial_measurement + known_hard_oversize = estimated > limit needs_compaction = before >= trigger or estimated > limit if not needs_compaction: return { @@ -1205,7 +1240,9 @@ async def get_measured_request_view( changed_any = False last_changed_rung = 0 for rung in range(1, 9): - candidate, changed = self._measured_apply_rung(rung, candidate) + candidate, changed = self._measured_apply_rung( + rung, candidate, developer_seqs + ) if not changed: continue changed_any = True @@ -1230,8 +1267,12 @@ async def get_measured_request_view( ) count_calls += 1 if final_measurement is None: + known_hard_oversize = ( + known_hard_oversize + or self._measured_hard_oversize(final_attempt) + ) transaction.rollback() - if estimated > limit: + if known_hard_oversize: raise ContextLengthError( "Context cannot fit protected content within the provider " "input limit after its recount became unavailable." @@ -1249,6 +1290,7 @@ async def get_measured_request_view( "transaction": None, } measured_after, estimated_after, limit_after, measured_source = final_measurement + known_hard_oversize = known_hard_oversize or estimated_after > limit_after if measured_after <= target and estimated_after <= limit_after: self._stage_measured_stats( initial_view=initial, @@ -1355,8 +1397,19 @@ def _protected_sequences(self, messages: list[dict[str, Any]]) -> set[int]: raise ValueError("Requested retained injection is not in admitted history") return {seq for msg in protected if (seq := self._extract_seq(msg)) is not None} - def _is_request_protected(self, msg: dict[str, Any]) -> bool: - return self._extract_seq(msg) in self._request_protected_seqs + def _is_request_protected( + self, + msg: dict[str, Any], + additional_protected_seqs: set[int] | None = None, + ) -> bool: + seq = self._extract_seq(msg) + return ( + seq in self._request_protected_seqs + or ( + additional_protected_seqs is not None + and seq in additional_protected_seqs + ) + ) def _check_retained_budget( self, messages: list[dict[str, Any]], budget: int @@ -1913,7 +1966,9 @@ def _record_stubbed(self, msg: dict[str, Any]) -> None: self._stubbed_seqs.add(seq) def _apply_sticky_decisions( - self, messages: list[dict[str, Any]] + self, + messages: list[dict[str, Any]], + additional_protected_seqs: set[int] | None = None, ) -> list[dict[str, Any]]: """Cheaply (O(n), no search) re-apply every previously-recorded compaction decision to a fresh copy of `messages`. @@ -1935,7 +1990,7 @@ def _apply_sticky_decisions( result: list[dict[str, Any]] = [] for msg in messages: seq = self._extract_seq(msg) - if self._is_request_protected(msg): + if self._is_request_protected(msg, additional_protected_seqs): result.append(dict(msg)) continue if seq is not None and seq in self._removed_seqs: @@ -1949,7 +2004,10 @@ def _apply_sticky_decisions( return result async def _compact_ephemeral( - self, budget: int, source_messages: list[dict[str, Any]] | None = None + self, + budget: int, + source_messages: list[dict[str, Any]] | None = None, + additional_protected_seqs: set[int] | None = None, ) -> list[dict[str, Any]]: """ Compact the context EPHEMERALLY using progressive interleaved strategy. @@ -2017,7 +2075,10 @@ async def _compact_ephemeral( # escalation has happened: deterministic given the same input messages, # so it reproduces exactly what was returned last time for anything not # newly appended -- which is what keeps the shared prefix byte-stable. - working_messages = self._apply_sticky_decisions(non_system_messages) + working_messages = self._apply_sticky_decisions( + non_system_messages, + additional_protected_seqs=additional_protected_seqs, + ) # === UNITS CONVENTION: TOTAL vs TOTAL, everywhere in this path === # @@ -2182,6 +2243,7 @@ async def _compact_ephemeral( target_tokens, protected_recent=level3_protection, system_tokens=system_tokens, + additional_protected_seqs=additional_protected_seqs, ) ) total_removed += removed @@ -2248,6 +2310,7 @@ async def _compact_ephemeral( target_tokens, protected_recent=level5_protection, system_tokens=system_tokens, + additional_protected_seqs=additional_protected_seqs, ) ) total_removed += removed @@ -2311,6 +2374,7 @@ async def _compact_ephemeral( target_tokens, protected_recent=level7_protection, system_tokens=system_tokens, + additional_protected_seqs=additional_protected_seqs, ) ) total_removed += removed @@ -2341,7 +2405,9 @@ async def _compact_ephemeral( if ( first_user_idx is not None and first_user_idx != last_user_idx - and not self._is_request_protected(working_messages[first_user_idx]) + and not self._is_request_protected( + working_messages[first_user_idx], additional_protected_seqs + ) and not _is_human_message(working_messages[first_user_idx]) ): first_msg = working_messages[first_user_idx] @@ -2371,7 +2437,7 @@ async def _compact_ephemeral( if msg.get("_stubbed") and i < protected_boundary # Outside protected recent zone and i != last_user_idx # Never remove last user message - and not self._is_request_protected(msg) + and not self._is_request_protected(msg, additional_protected_seqs) ] stubs_removed = 0 @@ -2510,6 +2576,7 @@ def _remove_messages_with_protection( target_tokens: int, protected_recent: float, system_tokens: int, + additional_protected_seqs: set[int] | None = None, ) -> tuple[list[dict[str, Any]], int, int, int]: """ Remove oldest messages with specified protection level. @@ -2565,7 +2632,9 @@ def _remove_messages_with_protection( ) protected_indices.update( - i for i, msg in enumerate(messages) if self._is_request_protected(msg) + i + for i, msg in enumerate(messages) + if self._is_request_protected(msg, additional_protected_seqs) ) if first_user_idx is not None: protected_indices.add(first_user_idx) diff --git a/tests/test_measured_request_view.py b/tests/test_measured_request_view.py index 89ac042..3d6f8ff 100644 --- a/tests/test_measured_request_view.py +++ b/tests/test_measured_request_view.py @@ -229,6 +229,7 @@ async def cancelled_count(_view): provider=None, retain_contents=[], count_view=cancelled_count ) assert not context._removed_seqs + assert not context._truncated_seqs @@ -321,4 +322,315 @@ async def test_unowned_hook_rejects_malformed_cache_write_without_crashing(): await context._on_llm_response( "llm:response", {"usage": {"input_tokens": 99, "cache_write_tokens": "bad"}} ) - assert context._last_measured_prompt_tokens == 12 \ No newline at end of file + assert context._last_measured_prompt_tokens == 12 + + +@pytest.mark.asyncio +async def test_measured_view_keeps_developer_instructions_at_protected_floor(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_recent=0, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + developer = "developer instruction " * 200 + await context.add_message({"role": "developer", "content": developer}) + await context.add_message({"role": "user", "content": "first human"}) + await context.add_message({"role": "assistant", "content": "old removable reply"}) + await context.add_message({"role": "user", "content": "latest human"}) + views = [] + + async def count_view(view): + views.append(view) + # The count derives from the instruction's presence: removing it would + # falsely make the candidate appear to reach target. + has_full_developer = any( + message.get("role") == "developer" + and message.get("content") == developer + for message in view + ) + return { + "dispatch": object(), + "budget_decision": _decision(900 if has_full_developer else 400), + } + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "protected_floor" + assert result["measured_before"] == result["measured_after"] == 900 + assert result["count_calls"] == len(views) == 2 + assert all( + any( + message.get("role") == "developer" + and message.get("content") == developer + for message in view + ) + for view in views + ) + assert (await context.get_messages())[0]["content"] == developer + result["transaction"].rollback() + + +@pytest.mark.asyncio +async def test_initial_unavailable_legacy_fallback_keeps_developer_instructions(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + protected_recent=0, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + developer = "developer instruction " * 200 + await context.add_message({"role": "developer", "content": developer}) + await context.add_message({"role": "user", "content": "first human"}) + await context.add_message({"role": "assistant", "content": "old removable reply"}) + await context.add_message({"role": "user", "content": "latest human"}) + views = [] + + async def unavailable_count(view): + views.append(view) + return {"dispatch": object(), "budget_decision": None} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=unavailable_count + ) + + assert result["outcome"] == "measurement_unavailable" + assert result["count_calls"] == len(views) == 2 + assert any( + message.get("role") == "developer" and message.get("content") == developer + for message in result["base_view"] + ) + assert (await context.get_messages())[0]["content"] == developer + result["transaction"].rollback() + assert not context._removed_seqs + + +@pytest.mark.asyncio +async def test_developer_protection_is_scoped_to_the_measured_capability(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + protected_recent=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + developer = "developer instruction " * 200 + await context.add_message({"role": "developer", "content": developer}) + await context.add_message({"role": "user", "content": "first human"}) + await context.add_message({"role": "user", "content": "latest human"}) + + async def count_view(_view): + return {"dispatch": object(), "budget_decision": _decision(0)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert result["outcome"] == "not_needed" + + legacy_view = await context.get_messages_for_request(1_000) + assert all(message.get("role") != "developer" for message in legacy_view) + + +@pytest.mark.asyncio +async def test_no_op_early_rungs_continue_to_a_later_legal_reduction(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for index in range(3): + await context.add_message({"role": "tool", "content": "x" * 800, "name": str(index)}) + await context.add_message({"role": "user", "content": "last"}) + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + return { + "dispatch": object(), + "budget_decision": _decision(900 if calls == 1 else 400), + } + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + # Three tool results leave rung one empty; rung two must still be reached. + assert result["outcome"] == "target_reached" + assert result["count_calls"] == calls == 2 + assert context._truncated_seqs + result["transaction"].rollback() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("field", "value"), + [ + ("estimated_input_tokens", True), + ("estimated_input_tokens", -1), + ("estimated_input_tokens", None), + ("input_limit_tokens", True), + ("input_limit_tokens", -1), + ("input_limit_tokens", None), + ], +) +async def test_invalid_old_hard_safety_fields_are_contract_errors(field, value): + context = SimpleContextManager(max_tokens=1_000, token_meter="actual") + + async def count_view(_view): + decision = _decision(10) + if value is None: + decision.pop(field) + else: + decision[field] = value + return {"dispatch": object(), "budget_decision": decision} + + with pytest.raises(TypeError, match="estimated_input_tokens and input_limit_tokens"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + +@pytest.mark.asyncio +async def test_soft_safe_malformed_measurement_remains_unavailable(): + context = SimpleContextManager(max_tokens=1_000, token_meter="actual") + + async def count_view(_view): + return { + "dispatch": object(), + "budget_decision": { + "estimated_input_tokens": 1_000, + "input_limit_tokens": 1_000, + "measurement": {"kind": "wrong", "source": "test.provider.count"}, + }, + } + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "measurement_unavailable" + assert result["count_calls"] == 1 + assert result["transaction"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "measurement", + [ + None, + {"kind": "wrong", "source": "test.provider.count", "input_tokens": 1}, + ], +) +async def test_initial_unavailable_measurement_cannot_bypass_known_hard_oversize( + measurement, +): + context = SimpleContextManager(max_tokens=1_000, token_meter="actual") + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + return { + "dispatch": object(), + "budget_decision": { + "estimated_input_tokens": 1_001, + "input_limit_tokens": 1_000, + "measurement": measurement, + }, + } + + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert calls == 1 + + +@pytest.mark.asyncio +async def test_fallback_unavailable_measurement_cannot_bypass_known_hard_oversize(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + if calls == 1: + return {"dispatch": object(), "budget_decision": None} + return { + "dispatch": object(), + "budget_decision": { + "estimated_input_tokens": 1_001, + "input_limit_tokens": 1_000, + "measurement": None, + }, + } + + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert calls == 2 + assert not context._truncated_seqs + assert not context._removed_seqs + + +@pytest.mark.asyncio +async def test_late_known_hard_oversize_fails_closed_when_next_recount_is_unavailable(): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + if calls == 1: + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=999), + } + if calls == 2: + return { + "dispatch": object(), + "budget_decision": _decision(800, estimated=1_001), + } + return {"dispatch": object(), "budget_decision": None} + + with pytest.raises(ContextLengthError, match="recount became unavailable"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert calls == 3 + assert not context._truncated_seqs + assert not context._removed_seqs \ No newline at end of file From 039f8689a96f5e9d26d4d65e35a7746a82f18cb0 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:39:38 -0700 Subject: [PATCH 4/4] fix(context): make measured commit cancellation-aware Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_context_simple/__init__.py | 38 +++++-- tests/test_measured_request_view.py | 112 ++++++++++++++++++++ 2 files changed, 139 insertions(+), 11 deletions(-) diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index a132da9..ba0cf1c 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -3261,6 +3261,10 @@ def _estimate_tokens(self, messages: list[dict[str, Any]]) -> int: class _MeasuredViewTransaction: """Idempotent commit/rollback boundary for staged measured compaction.""" + _OPEN = "open" + _COMMITTED = "committed" + _ROLLED_BACK = "rolled_back" + def __init__( self, context: SimpleContextManager, @@ -3276,12 +3280,19 @@ def __init__( context._last_compaction_stats, context._last_token_meter_stats, ) - self._terminal = False + self._state = self._OPEN - async def commit(self) -> None: - """Accept the selected candidate immediately before Loop dispatch.""" - if self._terminal: - return + async def commit(self, *, is_cancelled: Callable[[], bool] | None = None) -> bool: + """Accept the selected candidate immediately before Loop dispatch. + + Event delivery remains part of the open transaction. A caller can + therefore veto a candidate after the await, while its snapshot is + still available to rollback. + """ + if self._state == self._COMMITTED: + return True + if self._state == self._ROLLED_BACK: + return False if self._context._hooks is not None and self._context._last_compaction_stats: try: await self._context._hooks.emit( @@ -3289,14 +3300,19 @@ async def commit(self) -> None: ) except Exception as error: logger.warning(f"Could not emit measured compaction event: {error}") - # Keep rollback possible while event delivery is awaitable. In - # particular, cancellation here means Loop never dispatches and its - # finally block must restore the staged decisions. - self._terminal = True + if self._state != self._OPEN: + return False + if is_cancelled is not None and is_cancelled(): + self.rollback() + return False + if self._state != self._OPEN: + return False + self._state = self._COMMITTED + return True def rollback(self) -> None: """Restore pre-candidate state unless the transaction was committed.""" - if self._terminal: + if self._state != self._OPEN: return ( self._context._removed_seqs, @@ -3306,4 +3322,4 @@ def rollback(self) -> None: self._context._last_compaction_stats, self._context._last_token_meter_stats, ) = self._snapshot - self._terminal = True + self._state = self._ROLLED_BACK diff --git a/tests/test_measured_request_view.py b/tests/test_measured_request_view.py index 3d6f8ff..44ff662 100644 --- a/tests/test_measured_request_view.py +++ b/tests/test_measured_request_view.py @@ -31,6 +31,34 @@ def _decision(count, *, estimated=None, limit=1_000): } +async def _staged_transaction(*, hooks=None): + """Create a measured candidate whose staged state needs a transaction.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + hooks=hooks, + ) + await context.add_message({"role": "user", "content": "first"}) + for _ in range(4): + await context.add_message({"role": "tool", "content": "x" * 800}) + await context.add_message({"role": "user", "content": "last"}) + + async def count_view(view): + count = 400 if any(message.get("_truncated") for message in view) else 900 + return {"dispatch": object(), "budget_decision": _decision(count)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + assert result["transaction"] is not None + assert context._truncated_seqs + return context, result["transaction"] + + @pytest.mark.asyncio async def test_measured_capability_is_actual_mode_only_and_foreground_is_always_additive(): estimate = _Coordinator() @@ -275,6 +303,90 @@ async def count_view(view): assert context._last_compaction_stats is None +@pytest.mark.asyncio +async def test_commit_veto_after_awaited_event_restores_the_snapshot(): + started = asyncio.Event() + release = asyncio.Event() + cancelled = False + + class Hooks: + async def emit(self, _event, _data): + started.set() + await release.wait() + + context, transaction = await _staged_transaction(hooks=Hooks()) + task = asyncio.create_task( + transaction.commit(is_cancelled=lambda: cancelled) + ) + await asyncio.wait_for(started.wait(), timeout=1) + cancelled = True + release.set() + + assert await task is False + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_concurrent_rollback_during_event_commit_wins(): + started = asyncio.Event() + release = asyncio.Event() + + class Hooks: + async def emit(self, _event, _data): + started.set() + await release.wait() + + context, transaction = await _staged_transaction(hooks=Hooks()) + task = asyncio.create_task(transaction.commit()) + await asyncio.wait_for(started.wait(), timeout=1) + transaction.rollback() + release.set() + + assert await task is False + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_transaction_terminal_commit_and_rollback_results_are_idempotent(): + committed_context, committed = await _staged_transaction() + assert await committed.commit() is True + committed.rollback() + assert await committed.commit() is True + assert committed_context._truncated_seqs + + rolled_back_context, rolled_back = await _staged_transaction() + rolled_back.rollback() + rolled_back.rollback() + assert await rolled_back.commit() is False + assert not rolled_back_context._truncated_seqs + assert rolled_back_context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_commit_with_default_cancellation_predicate_accepts_candidate(): + context, transaction = await _staged_transaction() + + assert await transaction.commit(is_cancelled=None) is True + assert context._truncated_seqs + + +@pytest.mark.asyncio +async def test_commit_predicate_error_propagates_while_snapshot_is_rollbackable(): + context, transaction = await _staged_transaction() + + def predicate_error(): + raise TypeError("predicate failed") + + with pytest.raises(TypeError, match="predicate failed"): + await transaction.commit(is_cancelled=predicate_error) + + transaction.rollback() + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + @pytest.mark.asyncio async def test_foreground_claim_blocks_generic_hook_after_owned_recording(): context = SimpleContextManager(token_meter="actual")