From d95e757dab0718c1ac0607f71d2bef73149d2d95 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:15:08 -0700 Subject: [PATCH 1/3] feat(loop): dispatch provider-counted context views Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 16 + amplifier_module_loop_streaming/__init__.py | 388 ++++++++++++++++++-- tests/test_measured_compaction_runtime.py | 127 +++++++ 3 files changed, 506 insertions(+), 25 deletions(-) create mode 100644 tests/test_measured_compaction_runtime.py diff --git a/README.md b/README.md index dfc5297..69d2ca8 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,22 @@ The existing `orchestrator:provider_budget` event exposes each preflight's attempt, result, estimate, allowance, and requested context budget to mounted observability consumers. +### Provider-count measured compaction + +When both sides advertise the optional measured contracts — +`context.measured_request_view` and +`request_budget:provider_count` — the loop lets Context compact against the +provider's native input count. Context returns the exact `ChatRequest` it +counted, and the loop sends that same object after committing the selected +retention transaction. This path is additive: contexts and providers without +both capabilities keep the legacy request-budget preflight and estimate-based +retention behavior. + +For non-streaming foreground responses, the optional +`context.foreground_usage` capability records only normalized successful +response usage. A counted streaming request records its final provider count; +streams without a provider count do not invent usage. + ### Provider-reported overflow recovery A provider may optionally expose synchronous `recover_context_overflow` to diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 7d3dac3..f6daee4 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -3316,6 +3316,15 @@ async def _execute_stream( candidate = get_capability("context.request_retention") if callable(candidate): retaining_getter = candidate + measured_view_getter = None + foreground_usage_claim = None + if callable(get_capability): + candidate = get_capability("context.measured_request_view") + if callable(candidate): + measured_view_getter = candidate + candidate = get_capability("context.foreground_usage") + if callable(candidate): + foreground_usage_claim = candidate if ( self._ephemeral_injection_mode == "persist" and retaining_getter is None @@ -3370,6 +3379,17 @@ def accepts_named_request_options(method: Any) -> bool: inspect.Parameter.KEYWORD_ONLY, ) + def provider_advertises(capability: str) -> bool: + """Check an optional provider capability without vendor knowledge.""" + get_info = getattr(provider, "get_info", None) + if not callable(get_info): + return False + try: + capabilities = getattr(get_info(), "capabilities", ()) + except Exception: # Optional capability discovery must be additive. + return False + return capability in capabilities + def get_context_overflow_recovery() -> Any | None: """Return the provider recovery method when this request can use it.""" recovery = getattr(provider, "recover_context_overflow", None) @@ -3469,6 +3489,11 @@ async def request_messages( provider_name = name break budget_capable = callable(getattr(provider, "request_budget", None)) + measured_compaction_capable = ( + measured_view_getter is not None + and budget_capable + and provider_advertises("request_budget:provider_count") + ) # Pure observability. `basis` names WHY this provider won: # "pinned" when the conversation-scope pin decided it (capability @@ -3601,6 +3626,105 @@ async def check_request_budget( return None, output_cap return target, output_cap + async def count_measured_request( + request: ChatRequest, + base_view: list[dict[str, Any]], + *, + request_options: Mapping[str, Any] | None, + ) -> dict[str, Any] | None: + """Count one frozen provider-facing request for the Context contract. + + Unlike the legacy guard this deliberately makes no fit decision: + context-simple owns the raw-count trigger, target, and eight-rung + progression. Retaining the old envelope validation here keeps a + malformed advertised provider budget fail-closed. + """ + request_budget = getattr(provider, "request_budget", None) + if not callable(request_budget): + return None + context_estimate = sum(len(str(message)) // 4 for message in base_view) + kwargs: dict[str, Any] = {"context_estimate": context_estimate} + if accepts_named_request_options(request_budget): + kwargs["request_options"] = request_options + decision = request_budget(request, **kwargs) + if inspect.isawaitable(decision): + decision = await decision + if decision is None: + return None + required = ( + "estimated_input_tokens", + "input_limit_tokens", + "context_token_budget", + ) + if not isinstance(decision, dict) or any( + isinstance(decision.get(key), bool) + or not isinstance(decision.get(key), int) + or decision[key] < 0 + for key in required + ): + raise ContextLengthError( + "Provider request_budget returned an invalid budget decision" + ) + output_cap = decision.get("max_output_tokens") + if output_cap is not None and ( + isinstance(output_cap, bool) + or not isinstance(output_cap, int) + or output_cap <= 0 + ): + raise ContextLengthError( + "Provider request_budget returned an invalid max_output_tokens" + ) + return decision + + def record_foreground_usage(response: Any) -> None: + """Persist only a normalized successful foreground response usage.""" + nonlocal foreground_usage_claim + usage = getattr(response, "usage", None) + if usage is None or foreground_usage_claim is None: + return + if hasattr(usage, "model_dump"): + usage = usage.model_dump() + elif not isinstance(usage, dict): + usage = vars(usage) + if not isinstance(usage, dict): + return + input_tokens = usage.get("input_tokens") + cache_write_tokens = usage.get("cache_write_tokens", 0) + if ( + isinstance(input_tokens, bool) + or not isinstance(input_tokens, int) + or input_tokens < 0 + or isinstance(cache_write_tokens, bool) + or not isinstance(cache_write_tokens, int) + or cache_write_tokens < 0 + ): + return + recorder = foreground_usage_claim() + if callable(recorder): + recorder( + input_tokens=input_tokens, + cache_write_tokens=cache_write_tokens, + ) + + def record_measured_stream_usage(measured_result: dict[str, Any] | None) -> None: + """A counted stream has no normalized final response usage object.""" + if foreground_usage_claim is None or not isinstance(measured_result, dict): + return + attempt = measured_result.get("final_attempt") + decision = attempt.get("budget_decision") if isinstance(attempt, dict) else None + measurement = decision.get("measurement") if isinstance(decision, dict) else None + count = measurement.get("input_tokens") if isinstance(measurement, dict) else None + if ( + isinstance(count, bool) + or not isinstance(count, int) + or count < 0 + or measurement.get("kind") != "provider_count" + ): + return + recorder = foreground_usage_claim() + if callable(recorder): + recorder(input_tokens=count) + def output_cap_candidates(original: int | None) -> list[int]: """Return the bounded lossless output-reserve ladder.""" if original is None or original <= 1_000: @@ -4062,11 +4186,6 @@ async def exit_for_cancellation() -> None: verify_admitted=True, ) retained_contents.append(content) - # Get messages for LLM request (context handles compaction internally) - # Pass provider for dynamic budget calculation based on model's context window - message_dicts = await request_messages(retained_contents) - message_dicts = list(message_dicts) # Convert to list for modification - base_message_dicts = list(message_dicts) replay_turn_start_block = ( self._turn_start_view_block if self._ephemeral_injection_mode == "tail" and iteration == 1 @@ -4085,6 +4204,19 @@ async def exit_for_cancellation() -> None: if self._ephemeral_injection_mode == "tail" else [] ) + # Measured compaction must start from Context's sticky snapshot, not + # from an ordinary estimate-driven view. The loop still admits all + # current reminders above, and captures their request-only replay + # plan once below; the count callback will apply that frozen plan to + # every candidate Context asks it to measure. + if measured_compaction_capable: + message_dicts: list[dict[str, Any]] = [] + base_message_dicts: list[dict[str, Any]] = [] + else: + # Get messages for LLM request (context handles compaction internally) + # Pass provider for dynamic budget calculation based on model's context window + message_dicts = list(await request_messages(retained_contents)) + base_message_dicts = list(message_dicts) # Splice the turn-start reminder block into the request VIEW # (reminder-redesign-spec.md, W1.2). Only reachable when @@ -4152,7 +4284,7 @@ async def exit_for_cancellation() -> None: _, changed = await self._persist_reminder( context, result.context_injection, tail=True ) - if changed: + if changed and not measured_compaction_capable: message_dicts = list(await request_messages([])) base_message_dicts = list(message_dicts) # Check if we should append to last tool result @@ -4255,7 +4387,7 @@ async def exit_for_cancellation() -> None: verify_admitted=retaining_getter is not None, ) retained_contents.append(content) - if changed or retaining_getter is not None: + if (changed or retaining_getter is not None) and not measured_compaction_capable: message_dicts = list( await request_messages(retained_contents) ) @@ -4322,7 +4454,95 @@ async def exit_for_cancellation() -> None: if self.extended_thinking and not stream_provider else {} ) - chat_request = build_chat_request(message_dicts) + measured_transaction = None + measured_count = None + if measured_compaction_capable: + async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: + candidate_messages = ( + _replay_request_overlays( + base_view, + turn_start_view_block=replay_turn_start_block, + request_injection=replay_request_injection, + pending_injections=replay_pending_injections, + ) + if self._ephemeral_injection_mode == "tail" + else list(base_view) + ) + candidate_request = build_chat_request(candidate_messages) + return { + "dispatch": candidate_request, + "budget_decision": await count_measured_request( + candidate_request, + base_view, + request_options=request_options, + ), + } + + measured_result = await measured_view_getter( + provider=provider, + retain_contents=retained_contents, + count_view=count_view, + ) + if not isinstance(measured_result, dict): + raise TypeError("context.measured_request_view returned a non-dictionary result") + chat_request = measured_result.get("final_attempt", {}).get("dispatch") + if not isinstance(chat_request, ChatRequest): + raise TypeError( + "context.measured_request_view did not return a ChatRequest dispatch" + ) + measured_transaction = measured_result.get("transaction") + measured_count = measured_result + smaller_context_budget = None + original_output_cap = None + budget_decision = measured_result.get("final_attempt", {}).get( + "budget_decision" + ) + if isinstance(budget_decision, dict) and isinstance( + budget_decision.get("measurement"), dict + ): + measurement = budget_decision["measurement"] + if ( + measurement.get("kind") == "provider_count" + and isinstance(measurement.get("input_tokens"), int) + and not isinstance(measurement.get("input_tokens"), bool) + ): + estimated = budget_decision["estimated_input_tokens"] + allowance = budget_decision["input_limit_tokens"] + await hooks.emit( + "orchestrator:provider_budget", + { + "attempt": measured_result.get("count_calls", 1) - 1, + "context_estimate": sum( + len(str(message)) // 4 + for message in measured_result.get("base_view", []) + ), + "estimated_input_tokens": estimated, + "input_limit_tokens": allowance, + "context_token_budget": budget_decision[ + "context_token_budget" + ], + "result": ( + "fits" if estimated <= allowance else "oversized" + ), + "measurement_kind": measurement["kind"], + "measurement_source": measurement.get("source"), + "measured_before": measured_result.get("measured_before"), + "measured_after": measured_result.get("measured_after"), + "policy_budget": measured_result.get("policy_budget"), + "trigger": measured_result.get("trigger"), + "target": measured_result.get("target"), + "outcome": measured_result.get("outcome"), + "count_calls": measured_result.get("count_calls"), + }, + ) + else: + chat_request = build_chat_request(message_dicts) + smaller_context_budget, original_output_cap = await check_request_budget( + chat_request, + base_message_dicts, + attempt=0, + request_options=request_options, + ) logger.info( f"[ORCHESTRATOR] ChatRequest created with {len(tools) if tools else 0} tools" ) @@ -4331,12 +4551,6 @@ async def exit_for_cancellation() -> None: f"[ORCHESTRATOR] Tool names: {[t.name for t in tools.values()]}" ) - smaller_context_budget, original_output_cap = await check_request_budget( - chat_request, - base_message_dicts, - attempt=0, - request_options=request_options, - ) dispatch_base_messages = base_message_dicts if smaller_context_budget is not None: if smaller_context_budget <= 0: @@ -4424,10 +4638,21 @@ async def exit_for_cancellation() -> None: # Apply rate limit delay before provider call await self._apply_rate_limit_delay(hooks, iteration) + # A measured Context transaction becomes visible only at the + # selected dispatch boundary. Cancellation during the delay must + # leave its staged sticky decisions and telemetry untouched. + if coordinator and coordinator.cancellation.is_cancelled: + if measured_transaction is not None: + measured_transaction.rollback() + await exit_for_cancellation() + return + if measured_transaction is not None: + await measured_transaction.commit() # Check if provider supports streaming if stream_provider: # Use streaming if available + record_measured_stream_usage(measured_count) self._llm_calls_this_turn += 1 if get_context_overflow_recovery() is not None: response_stream = self._stream_with_overflow_recovery( @@ -4555,6 +4780,7 @@ async def exit_for_cancellation() -> None: # Update rate limit timestamp after non-streaming response self._last_provider_call_end = time.monotonic() + record_foreground_usage(response) # Emit content block events if present content_blocks = getattr(response, "content_blocks", None) @@ -5081,8 +5307,14 @@ async def exit_for_cancellation() -> None: if pending_body: await self._persist_reminder(context, pending_body, tail=True) self._pending_ephemeral_injections.clear() - message_dicts = list(await request_messages(final_retained_contents)) - base_message_dicts = list(message_dicts) + if measured_compaction_capable: + # Context's measured capability owns the snapshot; do not make + # an estimate-driven finalization view before asking it. + message_dicts: list[dict[str, Any]] = [] + base_message_dicts: list[dict[str, Any]] = [] + else: + message_dicts = list(await request_messages(final_retained_contents)) + base_message_dicts = list(message_dicts) if self._ephemeral_injection_mode == "tail": message_dicts = _replay_request_overlays( message_dicts, @@ -5132,15 +5364,110 @@ async def exit_for_cancellation() -> None: request_options: dict[str, Any] = {} if self.extended_thinking: request_options["extended_thinking"] = True - max_iter_chat_request = build_chat_request( - message_dicts, tool_choice="none" - ) - smaller_context_budget, original_output_cap = await check_request_budget( - max_iter_chat_request, - base_message_dicts, - attempt=0, - request_options=request_options, - ) + final_measured_transaction = None + if measured_compaction_capable: + async def count_final_view( + base_view: list[dict[str, Any]], + ) -> dict[str, Any]: + candidate_messages = ( + _replay_request_overlays( + base_view, + turn_start_view_block=None, + request_injection=final_replay_request_injection, + pending_injections=final_replay_pending_injections, + ) + if self._ephemeral_injection_mode == "tail" + else list(base_view) + ) + candidate_messages.append(finalization_overlay) + candidate_request = build_chat_request( + candidate_messages, tool_choice="none" + ) + return { + "dispatch": candidate_request, + "budget_decision": await count_measured_request( + candidate_request, + base_view, + request_options=request_options, + ), + } + + measured_result = await measured_view_getter( + provider=provider, + retain_contents=final_retained_contents, + count_view=count_final_view, + ) + if not isinstance(measured_result, dict): + raise TypeError( + "context.measured_request_view returned a non-dictionary result" + ) + max_iter_chat_request = measured_result.get( + "final_attempt", {} + ).get("dispatch") + if not isinstance(max_iter_chat_request, ChatRequest): + raise TypeError( + "context.measured_request_view did not return a ChatRequest dispatch" + ) + final_measured_transaction = measured_result.get("transaction") + smaller_context_budget = None + original_output_cap = None + budget_decision = measured_result.get("final_attempt", {}).get( + "budget_decision" + ) + if isinstance(budget_decision, dict) and isinstance( + budget_decision.get("measurement"), dict + ): + measurement = budget_decision["measurement"] + if ( + measurement.get("kind") == "provider_count" + and isinstance(measurement.get("input_tokens"), int) + and not isinstance(measurement.get("input_tokens"), bool) + ): + estimated = budget_decision["estimated_input_tokens"] + allowance = budget_decision["input_limit_tokens"] + await hooks.emit( + "orchestrator:provider_budget", + { + "attempt": measured_result.get("count_calls", 1) - 1, + "context_estimate": sum( + len(str(message)) // 4 + for message in measured_result.get("base_view", []) + ), + "estimated_input_tokens": estimated, + "input_limit_tokens": allowance, + "context_token_budget": budget_decision[ + "context_token_budget" + ], + "result": ( + "fits" + if estimated <= allowance + else "oversized" + ), + "measurement_kind": measurement["kind"], + "measurement_source": measurement.get("source"), + "measured_before": measured_result.get( + "measured_before" + ), + "measured_after": measured_result.get("measured_after"), + "policy_budget": measured_result.get("policy_budget"), + "trigger": measured_result.get("trigger"), + "target": measured_result.get("target"), + "outcome": measured_result.get("outcome"), + "count_calls": measured_result.get("count_calls"), + }, + ) + else: + max_iter_chat_request = build_chat_request( + message_dicts, tool_choice="none" + ) + smaller_context_budget, original_output_cap = ( + await check_request_budget( + max_iter_chat_request, + base_message_dicts, + attempt=0, + request_options=request_options, + ) + ) dispatch_base_messages = base_message_dicts if smaller_context_budget is not None: if smaller_context_budget <= 0: @@ -5235,6 +5562,16 @@ async def exit_for_cancellation() -> None: else: max_iter_chat_request = rebuilt_request + if coordinator and coordinator.cancellation.is_cancelled: + if final_measured_transaction is not None: + final_measured_transaction.rollback() + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return + if final_measured_transaction is not None: + await final_measured_transaction.commit() self._llm_calls_this_turn += 1 try: response = await provider.complete( @@ -5263,6 +5600,7 @@ async def exit_for_cancellation() -> None: response = await provider.complete( recovered_request, **request_options ) + record_foreground_usage(response) response_text = getattr(response, "text", None) if not isinstance(response_text, str) or not response_text: response_text = self._extract_text_from_content( diff --git a/tests/test_measured_compaction_runtime.py b/tests/test_measured_compaction_runtime.py new file mode 100644 index 0000000..2d405c9 --- /dev/null +++ b/tests/test_measured_compaction_runtime.py @@ -0,0 +1,127 @@ +"""Contract coverage for Context's optional measured request-view capability.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from amplifier_module_loop_streaming import StreamingOrchestrator +from tests.test_ephemeral_cache_persist_mode import ( + MockCoordinator, + MockResponse, + RequestCapturingProvider, + ScriptedHookResult, + ScriptedHooks, +) +from tests.test_provider_budget_guard import BudgetContext + + +def _provider_count(count: int = 5) -> dict[str, object]: + return { + "estimated_input_tokens": count + 4, + "input_limit_tokens": 100, + "context_token_budget": 0, + "measurement": { + "kind": "provider_count", + "source": "test.provider.count", + "input_tokens": count, + }, + } + + +class _Transaction: + def __init__(self) -> None: + self.committed = 0 + self.rolled_back = 0 + + async def commit(self) -> None: + self.committed += 1 + + def rollback(self) -> None: + self.rolled_back += 1 + + +class _MeasuredContext(BudgetContext): + def __init__(self) -> None: + super().__init__() + self.measured_calls: list[tuple[object, list[str]]] = [] + self.transaction = _Transaction() + + async def get_measured_request_view(self, *, provider, retain_contents, count_view): + base_view = list(self._messages) + attempt = await count_view(base_view) + self.measured_calls.append((provider, list(retain_contents))) + decision = attempt["budget_decision"] + count = decision["measurement"]["input_tokens"] + return { + "base_view": base_view, + "final_attempt": attempt, + "outcome": "not_needed", + "measured_before": count, + "measured_after": count, + "policy_budget": 100, + "trigger": 80.0, + "target": 50, + "count_calls": 1, + "transaction": self.transaction, + } + + +class _MeasuredProvider(RequestCapturingProvider): + def __init__(self) -> None: + super().__init__() + self.budget_calls: list[object] = [] + + def get_info(self): + return SimpleNamespace(capabilities=["request_budget:provider_count"]) + + def request_budget(self, request, *, context_estimate, request_options=None): + self.budget_calls.append(request) + return _provider_count() + + async def complete(self, request, **kwargs): + self.requests.append(request) + response = MockResponse(text="ok") + response.usage = SimpleNamespace(input_tokens=17, cache_write_tokens=3) + return response + + +@pytest.mark.asyncio +async def test_measured_view_counts_and_dispatches_the_identical_request_once() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = MockCoordinator() + recorded: list[tuple[int, int]] = [] + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + coordinator.register_capability( + "context.foreground_usage", + lambda: lambda *, input_tokens, cache_write_tokens=0: recorded.append( + (input_tokens, cache_write_tokens) + ), + ) + hooks = ScriptedHooks( + { + "provider:request": ScriptedHookResult( + action="inject_context", + ephemeral=True, + context_injection="COUNTED-OVERLAY", + ) + } + ) + + await StreamingOrchestrator( + {"ephemeral_injection_mode": "tail", "reminder_placement": "tail"} + ).execute("current request", context, {"main": provider}, {}, hooks, coordinator) + + assert len(context.measured_calls) == 1 + assert context.request_calls == [] + assert provider.budget_calls == [provider.requests[0]] + assert "COUNTED-OVERLAY" in "\n".join( + message.content for message in provider.requests[0].messages + ) + assert context.transaction.committed == 1 + assert context.transaction.rolled_back == 0 + assert recorded == [(17, 3)] \ No newline at end of file From 977b8d7eb6cf3be18df18f2748ffd7b98594722e Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:51:00 -0700 Subject: [PATCH 2/3] fix(loop): harden measured compaction dispatch Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_loop_streaming/__init__.py | 345 ++++++++++--- tests/test_measured_compaction_runtime.py | 530 +++++++++++++++++++- 2 files changed, 791 insertions(+), 84 deletions(-) diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index f6daee4..df523b3 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -13,6 +13,7 @@ import logging import re import time +import weakref from collections.abc import AsyncIterator, Mapping from functools import partial from typing import Any, ClassVar @@ -1215,6 +1216,13 @@ def __init__(self, config: dict[str, Any]): # from the bounded iteration count when a forced finalization call is # needed after the loop has spent its budget. self._llm_calls_this_turn: int = 0 + # Context owns the meter's persistent claim; Loop retains the exact + # returned recorder per Context so every foreground path (including a + # later no-count stream) addresses the original scoped meter. Weak + # keys avoid retaining completed Context instances. + self._foreground_usage_recorders: weakref.WeakKeyDictionary[Any, Any] = ( + weakref.WeakKeyDictionary() + ) # Layer 1 call-budget bookkeeping (spec: 298-replacement). Both are # per-execute()-call state, reset in _execute_one_turn alongside # _tool_calls_this_turn. Initialized here so a fresh instance never @@ -3386,7 +3394,8 @@ def provider_advertises(capability: str) -> bool: return False try: capabilities = getattr(get_info(), "capabilities", ()) - except Exception: # Optional capability discovery must be additive. + except (AttributeError, TypeError, ValueError): + # Optional capability discovery must be additive. return False return capability in capabilities @@ -3665,6 +3674,33 @@ async def count_measured_request( raise ContextLengthError( "Provider request_budget returned an invalid budget decision" ) + validate_measured_budget_decision(decision) + return decision + + def validate_measured_budget_decision( + decision: Any, *, require_hard_fit: bool = False + ) -> int | None: + """Validate the old guard and return usable optional count truth.""" + required = ( + "estimated_input_tokens", + "input_limit_tokens", + "context_token_budget", + ) + if not isinstance(decision, dict) or any( + isinstance(decision.get(key), bool) + or not isinstance(decision.get(key), int) + or decision[key] < 0 + for key in required + ): + raise ContextLengthError( + "Provider request_budget returned an invalid budget decision" + ) + estimated = decision["estimated_input_tokens"] + allowance = decision["input_limit_tokens"] + if require_hard_fit and estimated > allowance: + raise ContextLengthError( + "Provider request exceeds its input budget at measured dispatch" + ) output_cap = decision.get("max_output_tokens") if output_cap is not None and ( isinstance(output_cap, bool) @@ -3674,22 +3710,69 @@ async def count_measured_request( raise ContextLengthError( "Provider request_budget returned an invalid max_output_tokens" ) - return decision + measurement = decision.get("measurement") + if not isinstance(measurement, dict): + return None + count = measurement.get("input_tokens") + source = measurement.get("source") + if ( + measurement.get("kind") != "provider_count" + or isinstance(count, bool) + or not isinstance(count, int) + or count < 0 + or not isinstance(source, str) + or not source + or estimated < count + ): + return None + return count - def record_foreground_usage(response: Any) -> None: - """Persist only a normalized successful foreground response usage.""" - nonlocal foreground_usage_claim + def measured_final_count(measured_result: Any) -> int | None: + """Validate a Context-produced final attempt before dispatching it.""" + if not isinstance(measured_result, dict): + raise TypeError("context.measured_request_view returned a non-dictionary result") + attempt = measured_result.get("final_attempt") + if not isinstance(attempt, dict): + raise TypeError("context.measured_request_view returned an invalid final attempt") + decision = attempt.get("budget_decision") + if decision is None: + return None + return validate_measured_budget_decision( + decision, require_hard_fit=True + ) + + def foreground_usage_recorder() -> Any | None: + """Claim one Context's foreground meter immediately before dispatch.""" + if foreground_usage_claim is None: + return None + recorder = self._foreground_usage_recorders.get(context) + if callable(recorder): + return recorder + recorder = foreground_usage_claim() + if callable(recorder): + self._foreground_usage_recorders[context] = recorder + return recorder + return None + + def record_foreground_usage(recorder: Any | None, response: Any) -> None: + """Persist a response usage or explicitly mark an owned value stale.""" usage = getattr(response, "usage", None) - if usage is None or foreground_usage_claim is None: + if not callable(recorder): + return + if usage is None: + recorder(input_tokens=None) return if hasattr(usage, "model_dump"): usage = usage.model_dump() elif not isinstance(usage, dict): usage = vars(usage) if not isinstance(usage, dict): + recorder(input_tokens=None) return input_tokens = usage.get("input_tokens") cache_write_tokens = usage.get("cache_write_tokens", 0) + if cache_write_tokens is None: + cache_write_tokens = 0 if ( isinstance(input_tokens, bool) or not isinstance(input_tokens, int) @@ -3698,32 +3781,43 @@ def record_foreground_usage(response: Any) -> None: or not isinstance(cache_write_tokens, int) or cache_write_tokens < 0 ): + recorder(input_tokens=None) return - recorder = foreground_usage_claim() - if callable(recorder): - recorder( - input_tokens=input_tokens, - cache_write_tokens=cache_write_tokens, - ) + recorder( + input_tokens=input_tokens, + cache_write_tokens=cache_write_tokens, + ) - def record_measured_stream_usage(measured_result: dict[str, Any] | None) -> None: - """A counted stream has no normalized final response usage object.""" - if foreground_usage_claim is None or not isinstance(measured_result, dict): + def record_unavailable_stream_usage() -> None: + """Keep a prior scoped meter stale rather than reopening generic hooks.""" + if foreground_usage_claim is None: return - attempt = measured_result.get("final_attempt") - decision = attempt.get("budget_decision") if isinstance(attempt, dict) else None - measurement = decision.get("measurement") if isinstance(decision, dict) else None - count = measurement.get("input_tokens") if isinstance(measurement, dict) else None - if ( - isinstance(count, bool) - or not isinstance(count, int) - or count < 0 - or measurement.get("kind") != "provider_count" - ): + if context not in self._foreground_usage_recorders: + logger.debug( + "No provider count for an unclaimed stream; foreground usage " + "ownership remains unavailable." + ) return - recorder = foreground_usage_claim() + recorder = foreground_usage_recorder() if callable(recorder): - recorder(input_tokens=count) + recorder(input_tokens=None) + logger.debug( + "No provider count for a previously claimed stream; retained " + "foreground usage is stale." + ) + + async def await_with_measured_rollback( + awaitable: Any, transaction: Any | None + ) -> Any: + """Propagate an interrupted post-selection await after rollback.""" + completed = False + try: + result = await awaitable + completed = True + return result + finally: + if not completed and transaction is not None: + transaction.rollback() def output_cap_candidates(original: int | None) -> list[int]: """Return the bounded lossless output-reserve ladder.""" @@ -4455,15 +4549,22 @@ async def exit_for_cancellation() -> None: else {} ) measured_transaction = None - measured_count = None + measured_stream_count = None if measured_compaction_capable: - async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: + async def count_view( + base_view: list[dict[str, Any]], + *, + turn_start_view_block: str | None = replay_turn_start_block, + request_injection: tuple[str, bool] | None = replay_request_injection, + pending_injections: list[dict[str, Any]] = replay_pending_injections, + options: Mapping[str, Any] = request_options, + ) -> dict[str, Any]: candidate_messages = ( _replay_request_overlays( base_view, - turn_start_view_block=replay_turn_start_block, - request_injection=replay_request_injection, - pending_injections=replay_pending_injections, + turn_start_view_block=turn_start_view_block, + request_injection=request_injection, + pending_injections=pending_injections, ) if self._ephemeral_injection_mode == "tail" else list(base_view) @@ -4474,7 +4575,7 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: "budget_decision": await count_measured_request( candidate_request, base_view, - request_options=request_options, + request_options=options, ), } @@ -4485,30 +4586,36 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: ) if not isinstance(measured_result, dict): raise TypeError("context.measured_request_view returned a non-dictionary result") - chat_request = measured_result.get("final_attempt", {}).get("dispatch") - if not isinstance(chat_request, ChatRequest): - raise TypeError( - "context.measured_request_view did not return a ChatRequest dispatch" - ) measured_transaction = measured_result.get("transaction") - measured_count = measured_result + try: + chat_request = measured_result.get("final_attempt", {}).get("dispatch") + if not isinstance(chat_request, ChatRequest): + raise TypeError( + "context.measured_request_view did not return a ChatRequest dispatch" + ) + measured_stream_count = measured_final_count(measured_result) + measured_base_view = measured_result.get("base_view") + if not isinstance(measured_base_view, list) or not all( + isinstance(message, dict) for message in measured_base_view + ): + raise TypeError( + "context.measured_request_view returned an invalid base_view" + ) + except BaseException: + if measured_transaction is not None: + measured_transaction.rollback() + raise smaller_context_budget = None original_output_cap = None budget_decision = measured_result.get("final_attempt", {}).get( "budget_decision" ) - if isinstance(budget_decision, dict) and isinstance( - budget_decision.get("measurement"), dict - ): + if measured_stream_count is not None and isinstance(budget_decision, dict): measurement = budget_decision["measurement"] - if ( - measurement.get("kind") == "provider_count" - and isinstance(measurement.get("input_tokens"), int) - and not isinstance(measurement.get("input_tokens"), bool) - ): - estimated = budget_decision["estimated_input_tokens"] - allowance = budget_decision["input_limit_tokens"] - await hooks.emit( + estimated = budget_decision["estimated_input_tokens"] + allowance = budget_decision["input_limit_tokens"] + await await_with_measured_rollback( + hooks.emit( "orchestrator:provider_budget", { "attempt": measured_result.get("count_calls", 1) - 1, @@ -4534,7 +4641,14 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: "outcome": measured_result.get("outcome"), "count_calls": measured_result.get("count_calls"), }, - ) + ), + measured_transaction, + ) + if coordinator and coordinator.cancellation.is_cancelled: + if measured_transaction is not None: + measured_transaction.rollback() + await exit_for_cancellation() + return else: chat_request = build_chat_request(message_dicts) smaller_context_budget, original_output_cap = await check_request_budget( @@ -4551,7 +4665,9 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: f"[ORCHESTRATOR] Tool names: {[t.name for t in tools.values()]}" ) - dispatch_base_messages = base_message_dicts + dispatch_base_messages = ( + measured_base_view if measured_compaction_capable else base_message_dicts + ) if smaller_context_budget is not None: if smaller_context_budget <= 0: raise ContextLengthError( @@ -4637,7 +4753,9 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: chat_request = rebuilt_request # Apply rate limit delay before provider call - await self._apply_rate_limit_delay(hooks, iteration) + await await_with_measured_rollback( + self._apply_rate_limit_delay(hooks, iteration), measured_transaction + ) # A measured Context transaction becomes visible only at the # selected dispatch boundary. Cancellation during the delay must # leave its staged sticky decisions and telemetry untouched. @@ -4647,12 +4765,32 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: await exit_for_cancellation() return if measured_transaction is not None: - await measured_transaction.commit() + committed = await await_with_measured_rollback( + measured_transaction.commit( + is_cancelled=( + lambda: bool( + coordinator and coordinator.cancellation.is_cancelled + ) + ) + ), + measured_transaction, + ) + if not committed or ( + coordinator and coordinator.cancellation.is_cancelled + ): + measured_transaction.rollback() + await exit_for_cancellation() + return # Check if provider supports streaming if stream_provider: # Use streaming if available - record_measured_stream_usage(measured_count) + if measured_stream_count is None: + record_unavailable_stream_usage() + else: + recorder = foreground_usage_recorder() + if callable(recorder): + recorder(input_tokens=measured_stream_count) self._llm_calls_this_turn += 1 if get_context_overflow_recovery() is not None: response_stream = self._stream_with_overflow_recovery( @@ -4678,6 +4816,7 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: chat_request.max_output_tokens or original_output_cap ), ), + on_unconfirmed_stream=record_unavailable_stream_usage, ) else: response_stream = self._stream_from_provider( @@ -4688,6 +4827,7 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: hooks, coordinator, provider_name=provider_name, + on_stream_error=record_unavailable_stream_usage, ) async for chunk in response_stream: # Check for immediate cancellation between chunks @@ -4720,6 +4860,7 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: break else: # Fallback to non-streaming + foreground_recorder = foreground_usage_recorder() try: self._llm_calls_this_turn += 1 response = await provider.complete(chat_request, **request_options) @@ -4780,7 +4921,7 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: # Update rate limit timestamp after non-streaming response self._last_provider_call_end = time.monotonic() - record_foreground_usage(response) + record_foreground_usage(foreground_recorder, response) # Emit content block events if present content_blocks = getattr(response, "content_blocks", None) @@ -5368,18 +5509,23 @@ async def count_view(base_view: list[dict[str, Any]]) -> dict[str, Any]: if measured_compaction_capable: async def count_final_view( base_view: list[dict[str, Any]], + *, + request_injection: tuple[str, bool] | None = final_replay_request_injection, + pending_injections: list[dict[str, Any]] = final_replay_pending_injections, + overlay: dict[str, Any] = finalization_overlay, + options: Mapping[str, Any] = request_options, ) -> dict[str, Any]: candidate_messages = ( _replay_request_overlays( base_view, turn_start_view_block=None, - request_injection=final_replay_request_injection, - pending_injections=final_replay_pending_injections, + request_injection=request_injection, + pending_injections=pending_injections, ) if self._ephemeral_injection_mode == "tail" else list(base_view) ) - candidate_messages.append(finalization_overlay) + candidate_messages.append(overlay) candidate_request = build_chat_request( candidate_messages, tool_choice="none" ) @@ -5388,7 +5534,7 @@ async def count_final_view( "budget_decision": await count_measured_request( candidate_request, base_view, - request_options=request_options, + request_options=options, ), } @@ -5401,6 +5547,7 @@ async def count_final_view( raise TypeError( "context.measured_request_view returned a non-dictionary result" ) + final_measured_transaction = measured_result.get("transaction") max_iter_chat_request = measured_result.get( "final_attempt", {} ).get("dispatch") @@ -5408,24 +5555,28 @@ async def count_final_view( raise TypeError( "context.measured_request_view did not return a ChatRequest dispatch" ) - final_measured_transaction = measured_result.get("transaction") + measured_final_count(measured_result) + final_measured_base_view = measured_result.get("base_view") + if not isinstance(final_measured_base_view, list) or not all( + isinstance(message, dict) for message in final_measured_base_view + ): + raise TypeError( + "context.measured_request_view returned an invalid base_view" + ) smaller_context_budget = None original_output_cap = None budget_decision = measured_result.get("final_attempt", {}).get( "budget_decision" ) - if isinstance(budget_decision, dict) and isinstance( - budget_decision.get("measurement"), dict + if ( + measured_final_count(measured_result) is not None + and isinstance(budget_decision, dict) ): measurement = budget_decision["measurement"] - if ( - measurement.get("kind") == "provider_count" - and isinstance(measurement.get("input_tokens"), int) - and not isinstance(measurement.get("input_tokens"), bool) - ): - estimated = budget_decision["estimated_input_tokens"] - allowance = budget_decision["input_limit_tokens"] - await hooks.emit( + estimated = budget_decision["estimated_input_tokens"] + allowance = budget_decision["input_limit_tokens"] + await await_with_measured_rollback( + hooks.emit( "orchestrator:provider_budget", { "attempt": measured_result.get("count_calls", 1) - 1, @@ -5455,7 +5606,9 @@ async def count_final_view( "outcome": measured_result.get("outcome"), "count_calls": measured_result.get("count_calls"), }, - ) + ), + final_measured_transaction, + ) else: max_iter_chat_request = build_chat_request( message_dicts, tool_choice="none" @@ -5468,7 +5621,11 @@ async def count_final_view( request_options=request_options, ) ) - dispatch_base_messages = base_message_dicts + dispatch_base_messages = ( + final_measured_base_view + if measured_compaction_capable + else base_message_dicts + ) if smaller_context_budget is not None: if smaller_context_budget <= 0: raise ContextLengthError( @@ -5562,6 +5719,10 @@ async def count_final_view( else: max_iter_chat_request = rebuilt_request + await await_with_measured_rollback( + self._apply_rate_limit_delay(hooks, iteration), + final_measured_transaction, + ) if coordinator and coordinator.cancellation.is_cancelled: if final_measured_transaction is not None: final_measured_transaction.rollback() @@ -5571,7 +5732,27 @@ async def count_final_view( await exit_for_cancellation() return if final_measured_transaction is not None: - await final_measured_transaction.commit() + committed = await await_with_measured_rollback( + final_measured_transaction.commit( + is_cancelled=( + lambda: bool( + coordinator + and coordinator.cancellation.is_cancelled + ) + ) + ), + final_measured_transaction, + ) + if not committed or ( + coordinator and coordinator.cancellation.is_cancelled + ): + final_measured_transaction.rollback() + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return + foreground_recorder = foreground_usage_recorder() self._llm_calls_this_turn += 1 try: response = await provider.complete( @@ -5600,7 +5781,7 @@ async def count_final_view( response = await provider.complete( recovered_request, **request_options ) - record_foreground_usage(response) + record_foreground_usage(foreground_recorder, response) response_text = getattr(response, "text", None) if not isinstance(response_text, str) or not response_text: response_text = self._extract_text_from_content( @@ -5706,6 +5887,9 @@ async def count_final_view( await close_finalization_tool_turn( "The final response could not be generated." ) + finally: + if final_measured_transaction is not None: + final_measured_transaction.rollback() # Emit execution end await hooks.emit( @@ -5735,6 +5919,7 @@ async def _stream_with_overflow_recovery( coordinator=None, provider_name=None, recover_overflow=None, + on_unconfirmed_stream=None, ) -> AsyncIterator[str]: """Forward a stream, retrying exactly once before its first SDK chunk.""" while True: @@ -5753,6 +5938,7 @@ def mark_provider_chunk() -> None: coordinator, provider_name=provider_name, on_provider_chunk=mark_provider_chunk, + on_stream_error=on_unconfirmed_stream, ) try: async for chunk in stream: @@ -5783,6 +5969,7 @@ async def _stream_from_provider( coordinator=None, provider_name=None, on_provider_chunk=None, + on_stream_error=None, ) -> AsyncIterator[str]: """Stream tokens from provider that supports streaming. @@ -5804,6 +5991,10 @@ async def _stream_from_provider( tools_list = list(tools.values()) if tools else [] try: stream_iter = provider.stream(chat_request, tools=tools_list) + except ContextLengthError: + if on_stream_error is not None: + on_stream_error() + raise except LLMError as e: await hooks.emit( PROVIDER_ERROR, @@ -5853,6 +6044,10 @@ async def _stream_from_provider( full_response += token if self.stream_delay: await asyncio.sleep(self.stream_delay) + except ContextLengthError: + if on_stream_error is not None: + on_stream_error() + raise finally: await self._close_async_iterator( stream_iter, description="provider stream iterator" diff --git a/tests/test_measured_compaction_runtime.py b/tests/test_measured_compaction_runtime.py index 2d405c9..61c48cd 100644 --- a/tests/test_measured_compaction_runtime.py +++ b/tests/test_measured_compaction_runtime.py @@ -2,19 +2,23 @@ from __future__ import annotations +import asyncio from types import SimpleNamespace import pytest +from amplifier_core import ContextLengthError from amplifier_module_loop_streaming import StreamingOrchestrator from tests.test_ephemeral_cache_persist_mode import ( MockCoordinator, MockResponse, + NRoundToolProvider, + OneShotTool, RequestCapturingProvider, ScriptedHookResult, ScriptedHooks, ) -from tests.test_provider_budget_guard import BudgetContext +from tests.test_provider_budget_guard import HardFitBudgetContext, _retaining_coordinator def _provider_count(count: int = 5) -> dict[str, object]: @@ -34,23 +38,36 @@ class _Transaction: def __init__(self) -> None: self.committed = 0 self.rolled_back = 0 + self.terminal = False - async def commit(self) -> None: + async def commit(self, *, is_cancelled=None) -> bool: + if self.terminal: + return self.committed == 1 + if is_cancelled is not None and is_cancelled(): + self.rollback() + return False self.committed += 1 + self.terminal = True + return True def rollback(self) -> None: + if self.terminal: + return self.rolled_back += 1 + self.terminal = True -class _MeasuredContext(BudgetContext): +class _MeasuredContext(HardFitBudgetContext): def __init__(self) -> None: super().__init__() self.measured_calls: list[tuple[object, list[str]]] = [] - self.transaction = _Transaction() + self.transactions: list[_Transaction] = [] async def get_measured_request_view(self, *, provider, retain_contents, count_view): base_view = list(self._messages) attempt = await count_view(base_view) + transaction = _Transaction() + self.transactions.append(transaction) self.measured_calls.append((provider, list(retain_contents))) decision = attempt["budget_decision"] count = decision["measurement"]["input_tokens"] @@ -64,14 +81,16 @@ async def get_measured_request_view(self, *, provider, retain_contents, count_vi "trigger": 80.0, "target": 50, "count_calls": 1, - "transaction": self.transaction, + "transaction": transaction, } class _MeasuredProvider(RequestCapturingProvider): - def __init__(self) -> None: + def __init__(self, *, fail_first: bool = False) -> None: super().__init__() self.budget_calls: list[object] = [] + self.complete_calls = 0 + self.fail_first = fail_first def get_info(self): return SimpleNamespace(capabilities=["request_budget:provider_count"]) @@ -80,8 +99,20 @@ def request_budget(self, request, *, context_estimate, request_options=None): self.budget_calls.append(request) return _provider_count() + def recover_context_overflow( + self, request, error, *, context_estimate, request_options=None + ): + return { + "estimated_input_tokens": 100, + "input_limit_tokens": 10, + "context_token_budget": 1, + } + async def complete(self, request, **kwargs): + self.complete_calls += 1 self.requests.append(request) + if self.fail_first and self.complete_calls == 1: + raise ContextLengthError("provider rejected input") response = MockResponse(text="ok") response.usage = SimpleNamespace(input_tokens=17, cache_write_tokens=3) return response @@ -122,6 +153,487 @@ async def test_measured_view_counts_and_dispatches_the_identical_request_once() assert "COUNTED-OVERLAY" in "\n".join( message.content for message in provider.requests[0].messages ) - assert context.transaction.committed == 1 - assert context.transaction.rolled_back == 0 - assert recorded == [(17, 3)] \ No newline at end of file + assert [transaction.committed for transaction in context.transactions] == [1] + assert [transaction.rolled_back for transaction in context.transactions] == [0] + assert recorded == [(17, 3)] + + +@pytest.mark.asyncio +async def test_measured_overflow_recovery_uses_context_base_view_and_retries_once() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider(fail_first=True) + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + await StreamingOrchestrator({}).execute( + "long enough request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.complete_calls == 2 + assert len(provider.budget_calls) == 2 + assert len(provider.requests) == 2 + assert context.request_calls[-1][1] == 1 + + +@pytest.mark.asyncio +async def test_measured_transaction_rolls_back_when_budget_event_sets_cancellation() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + class CancellingHooks(ScriptedHooks): + async def emit(self, event, payload=None): + result = await super().emit(event, payload) + if event == "orchestrator:provider_budget": + coordinator.cancellation.is_cancelled = True + return result + + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + CancellingHooks({}), + coordinator, + ) + + assert provider.requests == [] + assert context.transactions[0].committed == 0 + assert context.transactions[0].rolled_back == 1 + + +@pytest.mark.asyncio +async def test_task_cancellation_during_measured_budget_event_rolls_back() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + started = asyncio.Event() + + class BlockingHooks(ScriptedHooks): + async def emit(self, event, payload=None): + if event == "orchestrator:provider_budget": + started.set() + await asyncio.Event().wait() + return await super().emit(event, payload) + + task = asyncio.create_task( + StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + BlockingHooks({}), + coordinator, + ) + ) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert provider.requests == [] + assert context.transactions[0].rolled_back == 1 + + + +@pytest.mark.asyncio +async def test_measured_rate_delay_cancellation_rolls_back_without_sdk_call( + monkeypatch, +) -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + async def cancel_during_delay(_self, _hooks, _iteration) -> None: + coordinator.cancellation.is_cancelled = True + + monkeypatch.setattr( + StreamingOrchestrator, "_apply_rate_limit_delay", cancel_during_delay + ) + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.requests == [] + assert context.transactions[0].rolled_back == 1 + + +class _MeasuredFinalizingProvider(NRoundToolProvider): + def __init__(self) -> None: + super().__init__(n_tool_rounds=1) + self.budget_calls: list[object] = [] + + def get_info(self): + return SimpleNamespace(capabilities=["request_budget:provider_count"]) + + def request_budget(self, request, *, context_estimate, request_options=None): + self.budget_calls.append(request) + return _provider_count() + + def recover_context_overflow( + self, request, error, *, context_estimate, request_options=None + ): + return { + "estimated_input_tokens": 100, + "input_limit_tokens": 10, + "context_token_budget": 1, + } + + async def complete(self, request, **kwargs): + response = await super().complete(request, **kwargs) + if self.call_count == 2: + raise ContextLengthError("finalization rejected input") + return response + + +@pytest.mark.asyncio +async def test_measured_finalization_commit_veto_closes_tool_turn_without_send() -> None: + coordinator = _retaining_coordinator(_MeasuredContext()) + + class CancellingFinalContext(_MeasuredContext): + async def get_measured_request_view(self, **kwargs): + result = await super().get_measured_request_view(**kwargs) + if len(self.transactions) == 2: + class CancellingTransaction(_Transaction): + async def commit(self, *, is_cancelled=None) -> bool: + coordinator.cancellation.is_cancelled = True + return await super().commit(is_cancelled=is_cancelled) + + transaction = CancellingTransaction() + self.transactions[-1] = transaction + result["transaction"] = transaction + return result + + context = CancellingFinalContext() + coordinator.register_capability( + "context.request_retention", context.retaining_view + ) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + provider = _MeasuredFinalizingProvider() + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "current request", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.call_count == 1 + assert [transaction.committed for transaction in context.transactions] == [1, 0] + assert [transaction.rolled_back for transaction in context.transactions] == [0, 1] + + +@pytest.mark.asyncio +async def test_measured_finalization_recovery_uses_context_base_view_once() -> None: + context = _MeasuredContext() + provider = _MeasuredFinalizingProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "current request", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.call_count == 3 + assert len(provider.budget_calls) == 3 + assert context.request_calls[-1][1] == 1 + assert [transaction.committed for transaction in context.transactions] == [1, 1] + + +@pytest.mark.asyncio +async def test_none_cache_write_usage_is_normalized_to_zero() -> None: + class CoreUsageShape: + def model_dump(self): + return {"input_tokens": 17, "cache_write_tokens": None} + + class NoneCacheWriteProvider(_MeasuredProvider): + async def complete(self, request, **kwargs): + response = await super().complete(request, **kwargs) + response.usage = CoreUsageShape() + return response + + context = _MeasuredContext() + provider = NoneCacheWriteProvider() + coordinator = MockCoordinator() + recorded: list[tuple[int, int]] = [] + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + coordinator.register_capability( + "context.foreground_usage", + lambda: lambda *, input_tokens, cache_write_tokens=0: recorded.append( + (input_tokens, cache_write_tokens) + ), + ) + + await StreamingOrchestrator({}).execute( + "current request", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert recorded == [(17, 0)] + + +@pytest.mark.asyncio +async def test_previously_claimed_no_count_stream_marks_usage_stale() -> None: + class UncountedStreamProvider: + def __init__(self) -> None: + self.requests: list[object] = [] + + async def stream(self, request, *, tools): + self.requests.append(request) + yield {"content": "streamed"} + + context = _MeasuredContext() + coordinator = MockCoordinator() + recorded: list[tuple[object, object]] = [] + claims = 0 + + def claim_usage(): + nonlocal claims + claims += 1 + + def recorder(*, input_tokens, cache_write_tokens=0): + recorded.append((input_tokens, cache_write_tokens)) + + return recorder + + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + coordinator.register_capability( + "context.foreground_usage", + claim_usage, + ) + loop = StreamingOrchestrator({}) + + await loop.execute( + "counted complete", + context, + {"main": _MeasuredProvider()}, + {}, + ScriptedHooks({}), + coordinator, + ) + stream_provider = UncountedStreamProvider() + await loop.execute( + "uncounted stream", + context, + {"stream": stream_provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert len(stream_provider.requests) == 1 + assert claims == 1 + assert recorded == [(17, 3), (None, 0)] + + +@pytest.mark.asyncio +async def test_never_claimed_no_count_stream_does_not_advertise_ownership() -> None: + class UncountedStreamProvider: + async def stream(self, request, *, tools): + yield {"content": "streamed"} + + context = _MeasuredContext() + coordinator = MockCoordinator() + recorded: list[tuple[object, object]] = [] + coordinator.register_capability( + "context.foreground_usage", + lambda: lambda *, input_tokens, cache_write_tokens=0: recorded.append( + (input_tokens, cache_write_tokens) + ), + ) + + await StreamingOrchestrator({}).execute( + "uncounted stream", + context, + {"stream": UncountedStreamProvider()}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert recorded == [] + + +@pytest.mark.asyncio +async def test_rejected_measured_stream_marks_its_count_stale_before_retry() -> None: + class RecoveringMeasuredStreamProvider(_MeasuredProvider): + def __init__(self) -> None: + super().__init__() + self.stream_calls = 0 + + async def stream(self, request, *, tools): + self.stream_calls += 1 + self.requests.append(request) + if self.stream_calls == 1: + raise ContextLengthError("stream rejected input") + yield {"content": "recovered"} + + context = _MeasuredContext() + provider = RecoveringMeasuredStreamProvider() + coordinator = _retaining_coordinator(context) + recorded: list[tuple[object, object]] = [] + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + coordinator.register_capability( + "context.foreground_usage", + lambda: lambda *, input_tokens, cache_write_tokens=0: recorded.append( + (input_tokens, cache_write_tokens) + ), + ) + + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.stream_calls == 2 + assert recorded == [(5, 0), (None, 0)] + + +@pytest.mark.asyncio +async def test_synchronously_rejected_measured_stream_marks_count_stale_before_retry() -> None: + class SynchronousRejectingStreamProvider(_MeasuredProvider): + def __init__(self) -> None: + super().__init__() + self.stream_calls = 0 + + def stream(self, request, *, tools): + self.stream_calls += 1 + self.requests.append(request) + if self.stream_calls == 1: + raise ContextLengthError("stream rejected input") + + async def recovered(): + yield {"content": "recovered"} + + return recovered() + + context = _MeasuredContext() + provider = SynchronousRejectingStreamProvider() + coordinator = _retaining_coordinator(context) + recorded: list[tuple[object, object]] = [] + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + coordinator.register_capability( + "context.foreground_usage", + lambda: lambda *, input_tokens, cache_write_tokens=0: recorded.append( + (input_tokens, cache_write_tokens) + ), + ) + + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.stream_calls == 2 + assert recorded == [(5, 0), (None, 0)] + + +@pytest.mark.asyncio +async def test_measured_final_dispatch_rejects_hard_oversize_before_sdk_call() -> None: + class HardOversizeContext(_MeasuredContext): + async def get_measured_request_view(self, **kwargs): + result = await super().get_measured_request_view(**kwargs) + result["final_attempt"] = dict(result["final_attempt"]) + result["final_attempt"]["budget_decision"] = { + **_provider_count(5), + "estimated_input_tokens": 101, + "input_limit_tokens": 100, + } + return result + + context = HardOversizeContext() + provider = _MeasuredProvider() + coordinator = MockCoordinator() + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + with pytest.raises(ContextLengthError, match="input budget at measured dispatch"): + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.requests == [] + assert context.transactions[0].rolled_back == 1 + + +@pytest.mark.asyncio +async def test_invalid_measured_dispatch_rolls_back_its_staged_transaction() -> None: + class InvalidDispatchContext(_MeasuredContext): + async def get_measured_request_view(self, **kwargs): + result = await super().get_measured_request_view(**kwargs) + result["final_attempt"] = dict(result["final_attempt"]) + result["final_attempt"]["dispatch"] = object() + return result + + context = InvalidDispatchContext() + provider = _MeasuredProvider() + coordinator = MockCoordinator() + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + + with pytest.raises(TypeError, match="did not return a ChatRequest"): + await StreamingOrchestrator({}).execute( + "current request", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.requests == [] + assert context.transactions[0].rolled_back == 1 From a9544aa70f16e3dca38596c56963fac2c072c7ec Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 22:31:20 -0700 Subject: [PATCH 3/3] test(loop): clarify measured stream verification Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 6 ++++-- tests/test_measured_compaction_runtime.py | 5 ++++- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 69d2ca8..4fe6cb7 100644 --- a/README.md +++ b/README.md @@ -78,8 +78,10 @@ retention behavior. For non-streaming foreground responses, the optional `context.foreground_usage` capability records only normalized successful -response usage. A counted streaming request records its final provider count; -streams without a provider count do not invent usage. +response usage. A counted streaming request records its selected provider count. +If a pre-chunk overflow requires a stream retry, that reading is marked stale +rather than attributed to the replacement request. Streams without a provider +count do not invent usage; an already-owned reading is marked stale. ### Provider-reported overflow recovery diff --git a/tests/test_measured_compaction_runtime.py b/tests/test_measured_compaction_runtime.py index 61c48cd..33e50c7 100644 --- a/tests/test_measured_compaction_runtime.py +++ b/tests/test_measured_compaction_runtime.py @@ -18,7 +18,10 @@ ScriptedHookResult, ScriptedHooks, ) -from tests.test_provider_budget_guard import HardFitBudgetContext, _retaining_coordinator +from tests.test_provider_budget_guard import ( + HardFitBudgetContext, + _retaining_coordinator, +) def _provider_count(count: int = 5) -> dict[str, object]: