diff --git a/README.md b/README.md index dfc5297..4fe6cb7 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,24 @@ 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 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 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..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 @@ -3316,6 +3324,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 +3387,18 @@ 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 (AttributeError, TypeError, ValueError): + # 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 +3498,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 +3635,190 @@ 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" + ) + 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) + or not isinstance(output_cap, int) + or output_cap <= 0 + ): + raise ContextLengthError( + "Provider request_budget returned an invalid max_output_tokens" + ) + 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 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 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) + or input_tokens < 0 + or isinstance(cache_write_tokens, bool) + or not isinstance(cache_write_tokens, int) + or cache_write_tokens < 0 + ): + recorder(input_tokens=None) + return + recorder( + input_tokens=input_tokens, + cache_write_tokens=cache_write_tokens, + ) + + def record_unavailable_stream_usage() -> None: + """Keep a prior scoped meter stale rather than reopening generic hooks.""" + if foreground_usage_claim is None: + return + 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_recorder() + if callable(recorder): + 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.""" if original is None or original <= 1_000: @@ -4062,11 +4280,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 +4298,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 +4378,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 +4481,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 +4548,115 @@ 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_stream_count = None + if measured_compaction_capable: + 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=turn_start_view_block, + request_injection=request_injection, + pending_injections=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=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") + measured_transaction = measured_result.get("transaction") + 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 measured_stream_count is not None and isinstance(budget_decision, dict): + measurement = budget_decision["measurement"] + 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, + "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"), + }, + ), + 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( + 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,13 +4665,9 @@ 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 = ( + measured_base_view if measured_compaction_capable else base_message_dicts ) - dispatch_base_messages = base_message_dicts if smaller_context_budget is not None: if smaller_context_budget <= 0: raise ContextLengthError( @@ -4423,11 +4753,44 @@ async def exit_for_cancellation() -> None: 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. + 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: + 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 + 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( @@ -4453,6 +4816,7 @@ async def exit_for_cancellation() -> None: chat_request.max_output_tokens or original_output_cap ), ), + on_unconfirmed_stream=record_unavailable_stream_usage, ) else: response_stream = self._stream_from_provider( @@ -4463,6 +4827,7 @@ async def exit_for_cancellation() -> None: 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 @@ -4495,6 +4860,7 @@ async def exit_for_cancellation() -> None: 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) @@ -4555,6 +4921,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(foreground_recorder, response) # Emit content block events if present content_blocks = getattr(response, "content_blocks", None) @@ -5081,8 +5448,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,16 +5505,127 @@ 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]], + *, + 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=request_injection, + pending_injections=pending_injections, + ) + if self._ephemeral_injection_mode == "tail" + else list(base_view) + ) + candidate_messages.append(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=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" + ) + final_measured_transaction = measured_result.get("transaction") + 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" + ) + 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 ( + measured_final_count(measured_result) is not None + and isinstance(budget_decision, dict) + ): + measurement = budget_decision["measurement"] + 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, + "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"), + }, + ), + final_measured_transaction, + ) + 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 = ( + final_measured_base_view + if measured_compaction_capable + else base_message_dicts ) - dispatch_base_messages = base_message_dicts if smaller_context_budget is not None: if smaller_context_budget <= 0: raise ContextLengthError( @@ -5235,6 +5719,40 @@ async def exit_for_cancellation() -> None: 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() + 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: + 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( @@ -5263,6 +5781,7 @@ async def exit_for_cancellation() -> None: response = await provider.complete( recovered_request, **request_options ) + 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( @@ -5368,6 +5887,9 @@ async def exit_for_cancellation() -> None: 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( @@ -5397,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: @@ -5415,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: @@ -5445,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. @@ -5466,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, @@ -5515,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 new file mode 100644 index 0000000..33e50c7 --- /dev/null +++ b/tests/test_measured_compaction_runtime.py @@ -0,0 +1,642 @@ +"""Contract coverage for Context's optional measured request-view capability.""" + +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 ( + HardFitBudgetContext, + _retaining_coordinator, +) + + +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 + self.terminal = False + + 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(HardFitBudgetContext): + def __init__(self) -> None: + super().__init__() + self.measured_calls: list[tuple[object, list[str]]] = [] + 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"] + 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": transaction, + } + + +class _MeasuredProvider(RequestCapturingProvider): + 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"]) + + 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 + + +@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 [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