diff --git a/README.md b/README.md index 5a3e497..28f7a19 100644 --- a/README.md +++ b/README.md @@ -64,6 +64,26 @@ The existing `orchestrator:provider_budget` event exposes each preflight's attempt, result, estimate, allowance, and requested context budget to mounted observability consumers. +### Provider-reported overflow recovery + +A provider may optionally expose synchronous `recover_context_overflow` to +turn its own `ContextLengthError` into one smaller context target. The loop +uses it only before a normal or finalization request has yielded an SDK chunk +or returned a response, only when retention explicitly supports `hard_fit`, +and only once for that outbound generation. The recovered view replays the +already-resolved retention and request overlays without rerunning hooks or +tools; a finalization retry keeps `tool_choice="none"`. + +Recovery feedback must be a strict budget dictionary: non-boolean integer +fields, an observed input above its allowance, and a positive target smaller +than the failed context estimate. A retry is preflighted with the same complete +options and the same or lower wire output cap. A fitting preflight retries; +`None` is an explicitly unproven retry authorized by the server rejection; an +over-budget, malformed, cancelled, or second-overflow path propagates without +another send. The generic loop neither parses provider error messages nor +knows provider-private feedback formats. `orchestrator:provider_overflow_recovery` +records only the scalar recovery result. + When a provider reports its effective output cap, an oversized compacted view is also preflighted with progressively smaller response reserves: 50%, 40%, 30%, 20%, and 10% of the original cap, then 1,000 tokens. These are local diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 74f9852..15800a5 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -13,10 +13,17 @@ import logging import re import time -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Mapping +from functools import partial from typing import Any, ClassVar -from amplifier_core import ContextLengthError, HookRegistry, HookResult, ModuleCoordinator, ToolResult +from amplifier_core import ( + ContextLengthError, + HookRegistry, + HookResult, + ModuleCoordinator, + ToolResult, +) from amplifier_core.events import ( CANCEL_COMPLETED, CANCEL_REQUESTED, @@ -748,6 +755,7 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = "orchestrator:goal_progress", # /goal auto-continue loop progress (see docs/designs/goal-command.md) "orchestrator:budget_warning", # Layer 1 call budget at budget_warn_ratio (see _execute_stream) "orchestrator:provider_budget", # Provider request-budget preflight result (see _execute_stream) + "orchestrator:provider_overflow_recovery", # Bounded provider-authorized input-overflow retry ], ) @@ -3345,6 +3353,35 @@ def retention_accepts_hard_fit() -> bool: for parameter in parameters ) + def accepts_named_request_options(method: Any) -> bool: + """Whether an optional provider method explicitly accepts options. + + A provider with only ``**kwargs`` keeps its pre-existing call + shape. As with ``hard_fit`` above, inspect separately from the + invocation: an implementation ``TypeError`` is never a cue to + retry the call without the keyword. + """ + try: + parameter = inspect.signature(method).parameters.get("request_options") + except (TypeError, ValueError): + return False + return parameter is not None and parameter.kind in ( + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.KEYWORD_ONLY, + ) + + 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) + if ( + not callable(recovery) + or retaining_getter is None + or not retention_accepts_hard_fit() + or (coordinator and coordinator.cancellation.is_cancelled) + ): + return None + return recovery + async def request_messages( retain_contents: list[str], *, @@ -3486,6 +3523,8 @@ async def check_request_budget( base_messages: list[dict[str, Any]], *, attempt: int, + request_options: Mapping[str, Any] | None = None, + allow_unproven: bool = False, ) -> tuple[int | None, int | None]: """Return a smaller context budget and effective output cap when available. @@ -3499,16 +3538,19 @@ async def check_request_budget( """ request_budget = getattr(provider, "request_budget", None) if not budget_capable or not callable(request_budget): - if attempt: + if attempt and not allow_unproven: raise ContextLengthError( "Provider request_budget capability was unavailable after " "reporting a concrete budget" ) return None, None context_estimate = sum(len(str(message)) // 4 for message in base_messages) - decision = request_budget(request, context_estimate=context_estimate) + budget_kwargs: dict[str, Any] = {"context_estimate": context_estimate} + if accepts_named_request_options(request_budget): + budget_kwargs["request_options"] = request_options + decision = request_budget(request, **budget_kwargs) if decision is None: - if attempt: + if attempt and not allow_unproven: raise ContextLengthError( "Provider request_budget capability was unavailable after " "reporting a concrete budget" @@ -3596,6 +3638,7 @@ async def try_reduced_output( *, original_output_cap: int | None, tool_choice: str | None = None, + request_options: Mapping[str, Any] | None = None, ) -> ChatRequest | None: """Preflight unchanged input at bounded lower output reserves.""" for cap in output_cap_candidates(original_output_cap): @@ -3605,7 +3648,10 @@ async def try_reduced_output( max_output_tokens=cap, ) smaller_budget, _ = await check_request_budget( - candidate_request, base_messages, attempt=1 + candidate_request, + base_messages, + attempt=1, + request_options=request_options, ) if smaller_budget is None: if cap < 10_000: @@ -3619,6 +3665,150 @@ async def try_reduced_output( return candidate_request return None + async def recover_context_overflow( + failed_request: ChatRequest, + error: ContextLengthError, + base_messages: list[dict[str, Any]], + *, + retain_contents: list[str], + turn_start_view_block: str | None, + request_injection: tuple[str, bool] | None, + pending_injections: list[dict[str, Any]], + tool_choice: str | None, + finalization_overlay: dict[str, Any] | None, + request_options: Mapping[str, Any] | None, + effective_output_cap: int | None, + ) -> ChatRequest | None: + """Build one provider-authorized hard-fit retry, or decline safely. + + The provider owns overflow classification. This generic layer + only validates its public decision and replays the already-resolved + request view; it never invokes hooks, tools, or provider-specific + error parsing. + """ + recovery = get_context_overflow_recovery() + if recovery is None: + return None + + failed_call_end = time.monotonic() + context_estimate = sum(len(str(message)) // 4 for message in base_messages) + recovery_kwargs: dict[str, Any] = {"context_estimate": context_estimate} + if accepts_named_request_options(recovery): + recovery_kwargs["request_options"] = request_options + decision = recovery(failed_request, error, **recovery_kwargs) + if decision is None: + logger.warning("Provider overflow recovery unavailable for this error") + await hooks.emit( + "orchestrator:provider_overflow_recovery", + {"result": "unavailable"}, + ) + return None + required = ( + "estimated_input_tokens", + "input_limit_tokens", + "context_token_budget", + ) + valid = isinstance(decision, dict) and all( + isinstance(decision.get(key), int) + and not isinstance(decision.get(key), bool) + and decision[key] >= 0 + for key in required + ) + if not valid: + logger.warning("Provider overflow recovery declined invalid feedback") + await hooks.emit( + "orchestrator:provider_overflow_recovery", + {"result": "invalid"}, + ) + return None + + observed = decision["estimated_input_tokens"] + allowance = decision["input_limit_tokens"] + target = decision["context_token_budget"] + decision_cap = decision.get("max_output_tokens") + valid_cap = decision_cap is None or ( + isinstance(decision_cap, int) + and not isinstance(decision_cap, bool) + and decision_cap > 0 + and ( + effective_output_cap is None + or decision_cap <= effective_output_cap + ) + ) + if ( + observed <= allowance + or allowance <= 0 + or target <= 0 + or target >= context_estimate + or not valid_cap + ): + logger.warning("Provider overflow recovery declined unsafe feedback") + await hooks.emit( + "orchestrator:provider_overflow_recovery", + {"result": "invalid"}, + ) + return None + + if coordinator and coordinator.cancellation.is_cancelled: + return None + rebuilt_base_messages = list( + await request_messages( + retain_contents, token_budget=target, hard_fit=True + ) + ) + rebuilt_messages = ( + _replay_request_overlays( + rebuilt_base_messages, + turn_start_view_block=turn_start_view_block, + request_injection=request_injection, + pending_injections=pending_injections, + ) + if self._ephemeral_injection_mode == "tail" + else list(rebuilt_base_messages) + ) + if finalization_overlay is not None: + rebuilt_messages.append(finalization_overlay) + if effective_output_cap is None: + retry_output_cap = decision_cap + elif decision_cap is None: + retry_output_cap = effective_output_cap + else: + retry_output_cap = min(effective_output_cap, decision_cap) + rebuilt_request = build_chat_request( + rebuilt_messages, + tool_choice=tool_choice, + max_output_tokens=retry_output_cap, + ) + smaller_budget, _ = await check_request_budget( + rebuilt_request, + rebuilt_base_messages, + attempt=1, + request_options=request_options, + # Valid server rejection authorizes one unknown preflight retry. + allow_unproven=True, + ) + if smaller_budget is not None: + logger.warning("Provider overflow recovery rebuilt request remains oversized") + await hooks.emit( + "orchestrator:provider_overflow_recovery", + {"result": "oversized"}, + ) + return None + if coordinator and coordinator.cancellation.is_cancelled: + return None + logger.warning("Retrying provider generation after input-overflow recovery") + await hooks.emit( + "orchestrator:provider_overflow_recovery", + {"result": "retry"}, + ) + # A rejected request is still a provider call. Preserve the normal + # minimum spacing before issuing the one authorized retry. + self._last_provider_call_end = failed_call_end + await self._apply_rate_limit_delay(hooks, iteration) + if coordinator and coordinator.cancellation.is_cancelled: + return None + return rebuilt_request + # --- Turn-start reminder assembly (reminder-redesign-spec.md, # W1.2, Option D). Hoists iteration 1's provider:request emit to # BEFORE the user prompt is appended below, so the merged @@ -4122,6 +4312,14 @@ async def exit_for_cancellation() -> None: # Clear pending injections after applying (both modes) self._pending_ephemeral_injections = [] + stream_provider = callable(getattr(provider, "stream", None)) + # `provider.stream` has no option-kwargs contract. Preflight only + # with options the selected dispatch path will actually receive. + request_options: dict[str, Any] = ( + {"extended_thinking": True} + if self.extended_thinking and not stream_provider + else {} + ) chat_request = build_chat_request(message_dicts) logger.info( f"[ORCHESTRATOR] ChatRequest created with {len(tools) if tools else 0} tools" @@ -4132,8 +4330,12 @@ async def exit_for_cancellation() -> None: ) smaller_context_budget, original_output_cap = await check_request_budget( - chat_request, base_message_dicts, attempt=0 + 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: raise ContextLengthError( @@ -4146,6 +4348,7 @@ async def exit_for_cancellation() -> None: hard_fit=True, ) ) + dispatch_base_messages = rebuilt_base_messages rebuilt_messages = ( _replay_request_overlays( rebuilt_base_messages, @@ -4158,13 +4361,17 @@ async def exit_for_cancellation() -> None: ) rebuilt_request = build_chat_request(rebuilt_messages) next_context_budget, _ = await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=1 + rebuilt_request, + rebuilt_base_messages, + attempt=1, + request_options=request_options, ) if next_context_budget is not None: reduced_output_request = await try_reduced_output( rebuilt_messages, rebuilt_base_messages, original_output_cap=original_output_cap, + request_options=request_options, ) if reduced_output_request is not None: chat_request = reduced_output_request @@ -4186,6 +4393,7 @@ async def exit_for_cancellation() -> None: hard_fit=True, ) ) + dispatch_base_messages = rebuilt_base_messages rebuilt_messages = ( _replay_request_overlays( rebuilt_base_messages, @@ -4199,7 +4407,10 @@ async def exit_for_cancellation() -> None: rebuilt_request = build_chat_request(rebuilt_messages) if ( await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=2 + rebuilt_request, + rebuilt_base_messages, + attempt=2, + request_options=request_options, ) )[0] is not None: raise ContextLengthError( @@ -4213,18 +4424,45 @@ async def exit_for_cancellation() -> None: await self._apply_rate_limit_delay(hooks, iteration) # Check if provider supports streaming - if hasattr(provider, "stream"): + if stream_provider: # Use streaming if available self._llm_calls_this_turn += 1 - async for chunk in self._stream_from_provider( - provider, - chat_request, - context, - tools, - hooks, - coordinator, - provider_name=provider_name, - ): + if get_context_overflow_recovery() is not None: + response_stream = self._stream_with_overflow_recovery( + provider, + chat_request, + context, + tools, + hooks, + coordinator, + provider_name=provider_name, + recover_overflow=partial( + recover_context_overflow, + chat_request, + base_messages=dispatch_base_messages, + retain_contents=retained_contents, + turn_start_view_block=replay_turn_start_block, + request_injection=replay_request_injection, + pending_injections=replay_pending_injections, + tool_choice=None, + finalization_overlay=None, + request_options=request_options, + effective_output_cap=( + chat_request.max_output_tokens or original_output_cap + ), + ), + ) + else: + response_stream = self._stream_from_provider( + provider, + chat_request, + context, + tools, + hooks, + coordinator, + provider_name=provider_name, + ) + async for chunk in response_stream: # Check for immediate cancellation between chunks if coordinator and coordinator.cancellation.is_immediate: # Clear pending steers: immediate cancellation ends the turn, @@ -4255,13 +4493,43 @@ async def exit_for_cancellation() -> None: break else: # Fallback to non-streaming - # Build kwargs for provider - kwargs = {} - if self.extended_thinking: - kwargs["extended_thinking"] = True try: self._llm_calls_this_turn += 1 - response = await provider.complete(chat_request, **kwargs) + response = await provider.complete(chat_request, **request_options) + except ContextLengthError as error: + recovered_request = await recover_context_overflow( + chat_request, + error, + dispatch_base_messages, + retain_contents=retained_contents, + turn_start_view_block=replay_turn_start_block, + request_injection=replay_request_injection, + pending_injections=replay_pending_injections, + tool_choice=None, + finalization_overlay=None, + request_options=request_options, + effective_output_cap=( + chat_request.max_output_tokens or original_output_cap + ), + ) + if recovered_request is None: + await hooks.emit( + PROVIDER_ERROR, + { + "provider": provider_name, + "error": { + "type": type(error).__name__, + "msg": str(error), + }, + "retryable": error.retryable, + "status_code": error.status_code, + }, + ) + raise + self._llm_calls_this_turn += 1 + response = await provider.complete( + recovered_request, **request_options + ) except LLMError as e: await hooks.emit( PROVIDER_ERROR, @@ -4859,12 +5127,19 @@ async def exit_for_cancellation() -> None: # result in the existing transcript stay valid. The portable # choice prevents new calls; this finalization path never # parses or dispatches a tool response. + 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 + 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: raise ContextLengthError( @@ -4878,6 +5153,7 @@ async def exit_for_cancellation() -> None: hard_fit=True, ) ) + dispatch_base_messages = rebuilt_base_messages rebuilt_messages = ( _replay_request_overlays( rebuilt_base_messages, @@ -4893,7 +5169,10 @@ async def exit_for_cancellation() -> None: rebuilt_messages, tool_choice="none" ) next_context_budget, _ = await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=1 + rebuilt_request, + rebuilt_base_messages, + attempt=1, + request_options=request_options, ) if next_context_budget is not None: reduced_output_request = await try_reduced_output( @@ -4901,6 +5180,7 @@ async def exit_for_cancellation() -> None: rebuilt_base_messages, original_output_cap=original_output_cap, tool_choice="none", + request_options=request_options, ) if reduced_output_request is not None: max_iter_chat_request = reduced_output_request @@ -4922,6 +5202,7 @@ async def exit_for_cancellation() -> None: hard_fit=True, ) ) + dispatch_base_messages = rebuilt_base_messages rebuilt_messages = ( _replay_request_overlays( rebuilt_base_messages, @@ -4938,7 +5219,10 @@ async def exit_for_cancellation() -> None: ) if ( await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=2 + rebuilt_request, + rebuilt_base_messages, + attempt=2, + request_options=request_options, ) )[0] is not None: raise ContextLengthError( @@ -4949,12 +5233,34 @@ async def exit_for_cancellation() -> None: else: max_iter_chat_request = rebuilt_request - kwargs = {} - if self.extended_thinking: - kwargs["extended_thinking"] = True - self._llm_calls_this_turn += 1 - response = await provider.complete(max_iter_chat_request, **kwargs) + try: + response = await provider.complete( + max_iter_chat_request, **request_options + ) + except ContextLengthError as error: + recovered_request = await recover_context_overflow( + max_iter_chat_request, + error, + dispatch_base_messages, + retain_contents=final_retained_contents, + turn_start_view_block=None, + request_injection=final_replay_request_injection, + pending_injections=final_replay_pending_injections, + tool_choice="none", + finalization_overlay=finalization_overlay, + request_options=request_options, + effective_output_cap=( + max_iter_chat_request.max_output_tokens + or original_output_cap + ), + ) + if recovered_request is None: + raise + self._llm_calls_this_turn += 1 + response = await provider.complete( + recovered_request, **request_options + ) response_text = getattr(response, "text", None) if not isinstance(response_text, str) or not response_text: response_text = self._extract_text_from_content( @@ -5070,6 +5376,63 @@ async def exit_for_cancellation() -> None: }, ) + async def _close_async_iterator(self, iterator: Any, *, description: str) -> None: + """Close a provider iterator when it offers asynchronous cleanup.""" + closer = getattr(iterator, "aclose", None) + if callable(closer): + try: + await closer() + except Exception: + logger.debug("Unable to close %s", description, exc_info=True) + + async def _stream_with_overflow_recovery( + self, + provider, + chat_request, + context, + tools, + hooks, + coordinator=None, + provider_name=None, + recover_overflow=None, + ) -> AsyncIterator[str]: + """Forward a stream, retrying exactly once before its first SDK chunk.""" + while True: + saw_provider_chunk = False + + def mark_provider_chunk() -> None: + nonlocal saw_provider_chunk + saw_provider_chunk = True + + stream = self._stream_from_provider( + provider, + chat_request, + context, + tools, + hooks, + coordinator, + provider_name=provider_name, + on_provider_chunk=mark_provider_chunk, + ) + try: + async for chunk in stream: + yield chunk + return + except ContextLengthError as error: + if saw_provider_chunk or recover_overflow is None: + raise + recovered_request = await recover_overflow(error) + if recovered_request is None: + raise + self._llm_calls_this_turn += 1 + chat_request = recovered_request + # A prospective outbound generation has one recovery allowance. + recover_overflow = None + finally: + await self._close_async_iterator( + stream, description="provider stream wrapper" + ) + async def _stream_from_provider( self, provider, @@ -5079,6 +5442,7 @@ async def _stream_from_provider( hooks, coordinator=None, provider_name=None, + on_provider_chunk=None, ) -> AsyncIterator[str]: """Stream tokens from provider that supports streaming. @@ -5121,31 +5485,38 @@ async def _stream_from_provider( ) raise - async for chunk in stream_iter: - # Check for immediate cancellation between chunks - if coordinator and coordinator.cancellation.is_immediate: - # Add partial response to context before exiting - if full_response: - await context.add_message( - {"role": "assistant", "content": full_response} - ) - return + try: + async for chunk in stream_iter: + if on_provider_chunk is not None: + on_provider_chunk() + # Check for immediate cancellation between chunks + if coordinator and coordinator.cancellation.is_immediate: + # Add partial response to context before exiting + if full_response: + await context.add_message( + {"role": "assistant", "content": full_response} + ) + return - # Skip non-text block deltas (e.g. thinking block streaming chunks). - # Providers that stream extended-thinking models include a block_type - # field so callers can distinguish thinking deltas from text deltas. - # Without this guard, thinking content leaks into full_response and - # ultimately into parse_json extraction downstream. - chunk_block_type = chunk.get("block_type") - if chunk_block_type and chunk_block_type != "text": - continue + # Skip non-text block deltas (e.g. thinking block streaming chunks). + # Providers that stream extended-thinking models include a block_type + # field so callers can distinguish thinking deltas from text deltas. + # Without this guard, thinking content leaks into full_response and + # ultimately into parse_json extraction downstream. + chunk_block_type = chunk.get("block_type") + if chunk_block_type and chunk_block_type != "text": + continue - token = chunk.get("content", "") - if token: - yield token - full_response += token - if self.stream_delay: - await asyncio.sleep(self.stream_delay) + token = chunk.get("content", "") + if token: + yield token + full_response += token + if self.stream_delay: + await asyncio.sleep(self.stream_delay) + finally: + await self._close_async_iterator( + stream_iter, description="provider stream iterator" + ) # Add complete message to context if full_response: diff --git a/tests/test_provider_budget_guard.py b/tests/test_provider_budget_guard.py index 210b531..1dff01c 100644 --- a/tests/test_provider_budget_guard.py +++ b/tests/test_provider_budget_guard.py @@ -1080,6 +1080,7 @@ async def test_mount_discovers_provider_budget_observability_event() -> None: ) events = events_contributor() assert events.count("orchestrator:provider_budget") == 1 + assert events.count("orchestrator:provider_overflow_recovery") == 1 assert { "execution:start", "execution:end", diff --git a/tests/test_provider_overflow_recovery.py b/tests/test_provider_overflow_recovery.py new file mode 100644 index 0000000..b1d1ee0 --- /dev/null +++ b/tests/test_provider_overflow_recovery.py @@ -0,0 +1,458 @@ +"""Focused generic Loop coverage for bounded provider overflow recovery.""" + +from __future__ import annotations + +import pytest +from amplifier_core import ContextLengthError + +from amplifier_module_loop_streaming import StreamingOrchestrator +from tests.test_ephemeral_cache_persist_mode import ( + MockResponse, + OneShotTool, + RequestCapturingProvider, + ScriptedHooks, + ToolCallStub, +) +from tests.test_provider_budget_guard import ( + HardFitBudgetContext, + _retaining_coordinator, +) + + +def _fit(cap: int = 128) -> dict[str, int]: + return { + "estimated_input_tokens": 1, + "input_limit_tokens": 10, + "context_token_budget": 0, + "max_output_tokens": cap, + } + + +def _overflow(target: int = 1, cap: int = 128) -> dict[str, int]: + return { + "estimated_input_tokens": 100, + "input_limit_tokens": 10, + "context_token_budget": target, + "max_output_tokens": cap, + } + + +class RecoveringProvider(RequestCapturingProvider): + def __init__(self, *, recovery: object = None, fail_call: int = 1) -> None: + super().__init__() + self.recovery = _overflow() if recovery is None else recovery + self.fail_call = fail_call + self.complete_calls = 0 + self.budget_options: list[object] = [] + self.recovery_options: list[object] = [] + + def request_budget(self, request, *, context_estimate: int, request_options=None): + self.budget_options.append(request_options) + # The post-rebuild None is authorized only by a valid server overflow. + return _fit() if len(self.budget_options) == 1 else None + + def recover_context_overflow( + self, request, error, *, context_estimate: int, request_options=None + ): + self.recovery_options.append(request_options) + return self.recovery + + async def complete(self, request, **kwargs): + self.complete_calls += 1 + self.requests.append(request) + if self.complete_calls == self.fail_call: + raise ContextLengthError("provider rejected input") + return MockResponse(text="recovered") + + +@pytest.mark.asyncio +async def test_normal_overflow_recovers_once_with_named_options_and_preserved_overlay() -> None: + context = HardFitBudgetContext() + context._messages.append({"role": "assistant", "content": "old history" * 100}) + provider = RecoveringProvider() + hooks = ScriptedHooks({}) + + result = await StreamingOrchestrator( + { + "extended_thinking": True, + "ephemeral_injection_mode": "tail", + "reminder_placement": "tail", + } + ).execute( + "work", + context, + {"main": provider}, + {}, + hooks, + _retaining_coordinator(context), + ) + + assert result == "recovered" + assert provider.complete_calls == 2 + assert provider.budget_options == [ + {"extended_thinking": True}, + {"extended_thinking": True}, + ] + assert provider.recovery_options == [{"extended_thinking": True}] + assert context.hard_fit_calls == [False, True] + assert provider.requests[1].max_output_tokens == 128 + assert [name for name, _ in hooks.emitted].count("provider:request") == 1 + + +@pytest.mark.asyncio +async def test_kwargs_only_budget_and_recovery_keep_legacy_call_shape() -> None: + class KwargsOnlyProvider(RecoveringProvider): + def request_budget(self, request, **kwargs): + self.budget_options.append(dict(kwargs)) + return _fit() if len(self.budget_options) == 1 else None + + def recover_context_overflow(self, request, error, **kwargs): + self.recovery_options.append(dict(kwargs)) + return self.recovery + + context = HardFitBudgetContext() + provider = KwargsOnlyProvider() + await StreamingOrchestrator({"extended_thinking": True}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert provider.budget_options == [ + {"context_estimate": provider.budget_options[0]["context_estimate"]}, + {"context_estimate": provider.budget_options[1]["context_estimate"]}, + ] + assert provider.recovery_options == [ + {"context_estimate": provider.recovery_options[0]["context_estimate"]} + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "feedback", + [ + None, + {"estimated_input_tokens": True, "input_limit_tokens": 10, "context_token_budget": 1}, + {"estimated_input_tokens": 10, "input_limit_tokens": 10, "context_token_budget": 1}, + {"estimated_input_tokens": 100, "input_limit_tokens": 0, "context_token_budget": 1}, + _overflow(target=0), + _overflow(target=999), + _overflow(cap=129), + ], +) +async def test_unavailable_invalid_or_unsafe_recovery_feedback_never_retries(feedback) -> None: + context = HardFitBudgetContext() + provider = RecoveringProvider() + provider.recovery = feedback + hooks = ScriptedHooks({}) + + with pytest.raises(ContextLengthError, match="provider rejected input"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + hooks, + _retaining_coordinator(context), + ) + + assert provider.complete_calls == 1 + assert context.hard_fit_calls == [False] + assert [ + data["result"] + for name, data in hooks.emitted + if name == "orchestrator:provider_overflow_recovery" + ] == ["unavailable" if feedback is None else "invalid"] + + +@pytest.mark.asyncio +async def test_recovery_honors_a_lower_provider_output_cap() -> None: + context = HardFitBudgetContext() + provider = RecoveringProvider(recovery=_overflow(cap=64)) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert provider.requests[-1].max_output_tokens == 64 + + +@pytest.mark.asyncio +async def test_cancellation_during_recovery_never_sends_a_retry() -> None: + context = HardFitBudgetContext() + coordinator = _retaining_coordinator(context) + + class CancellingProvider(RecoveringProvider): + def recover_context_overflow(self, request, error, *, context_estimate: int, request_options=None): + coordinator.cancellation.is_cancelled = True + return _overflow() + + provider = CancellingProvider() + with pytest.raises(ContextLengthError, match="provider rejected input"): + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert provider.complete_calls == 1 + + +@pytest.mark.asyncio +async def test_recovery_applies_rate_limit_delay_before_retry(monkeypatch) -> None: + context = HardFitBudgetContext() + provider = RecoveringProvider() + rate_limit_calls: list[tuple[int, bool]] = [] + + async def record_rate_limit_delay(self, _hooks, iteration: int) -> None: + rate_limit_calls.append((iteration, self._last_provider_call_end is not None)) + + monkeypatch.setattr( + StreamingOrchestrator, "_apply_rate_limit_delay", record_rate_limit_delay + ) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert rate_limit_calls == [(1, False), (1, True)] + + +@pytest.mark.asyncio +async def test_cancellation_during_recovery_rate_limit_never_sends_a_retry( + monkeypatch, +) -> None: + context = HardFitBudgetContext() + coordinator = _retaining_coordinator(context) + provider = RecoveringProvider() + + async def cancel_recovery_retry(self, _hooks, _iteration: int) -> None: + if self._last_provider_call_end is not None: + coordinator.cancellation.is_cancelled = True + + monkeypatch.setattr( + StreamingOrchestrator, "_apply_rate_limit_delay", cancel_recovery_retry + ) + + with pytest.raises(ContextLengthError, match="provider rejected input"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + + assert provider.complete_calls == 1 + + +@pytest.mark.asyncio +async def test_unrecoverable_normal_overflow_keeps_provider_error_event() -> None: + class NoRecoveryProvider(RequestCapturingProvider): + async def complete(self, request, **kwargs): + self.requests.append(request) + raise ContextLengthError("unrecoverable") + + context = HardFitBudgetContext() + hooks = ScriptedHooks({}) + provider = NoRecoveryProvider() + with pytest.raises(ContextLengthError, match="unrecoverable"): + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, hooks, _retaining_coordinator(context) + ) + + assert [name for name, _ in hooks.emitted].count("provider:error") == 1 + + +class RecoveringStreamProvider(RecoveringProvider): + def __init__(self, *, after_chunk: bool = False) -> None: + super().__init__() + self.after_chunk = after_chunk + self.stream_calls = 0 + + async def stream(self, request, *, tools): + self.stream_calls += 1 + self.requests.append(request) + if self.stream_calls == 1: + if self.after_chunk: + yield {"block_type": "thinking", "content": "hidden"} + raise ContextLengthError("stream rejected input") + yield {"content": "stream recovered"} + + +@pytest.mark.asyncio +async def test_stream_recovers_only_before_its_first_provider_chunk() -> None: + context = HardFitBudgetContext() + provider = RecoveringStreamProvider() + + result = await StreamingOrchestrator({"extended_thinking": True}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert result == "stream recovered" + assert provider.stream_calls == 2 + assert provider.budget_options == [{}, {}] + assert provider.recovery_options == [{}] + + +@pytest.mark.asyncio +async def test_stream_without_recovery_uses_the_direct_stream_path(monkeypatch) -> None: + class DirectStreamProvider(RequestCapturingProvider): + async def stream(self, request, *, tools): + self.requests.append(request) + yield {"content": "direct"} + + orchestrator = StreamingOrchestrator({}) + monkeypatch.setattr( + orchestrator, + "_stream_with_overflow_recovery", + lambda *_args, **_kwargs: pytest.fail("unavailable recovery wrapped a stream"), + ) + context = HardFitBudgetContext() + + result = await orchestrator.execute( + "work", + context, + {"main": DirectStreamProvider()}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert result == "direct" + + +@pytest.mark.asyncio +async def test_closing_a_recovery_stream_closes_its_provider_iterator() -> None: + from amplifier_core.message_models import ChatRequest, Message + + class ClosableStream: + def __init__(self) -> None: + self.closed = False + self.sent = False + + def __aiter__(self): + return self + + async def __anext__(self): + if self.sent: + raise StopAsyncIteration + self.sent = True + return {"content": "first"} + + async def aclose(self) -> None: + self.closed = True + + class ClosableStreamProvider: + def __init__(self) -> None: + self.iterator = ClosableStream() + + def stream(self, _request, *, tools): + return self.iterator + + orchestrator = StreamingOrchestrator({}) + provider = ClosableStreamProvider() + response_stream = orchestrator._stream_with_overflow_recovery( + provider, + ChatRequest(messages=[Message(role="user", content="work")]), + HardFitBudgetContext(), + {}, + ScriptedHooks({}), + recover_overflow=None, + ) + + assert await anext(response_stream) == "first" + assert not provider.iterator.closed + await response_stream.aclose() + assert provider.iterator.closed + + +@pytest.mark.asyncio +async def test_stream_overflow_after_nontext_chunk_does_not_recover() -> None: + context = HardFitBudgetContext() + provider = RecoveringStreamProvider(after_chunk=True) + + with pytest.raises(ContextLengthError, match="stream rejected input"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert provider.stream_calls == 1 + assert provider.recovery_options == [] + + +@pytest.mark.asyncio +async def test_finalization_overflow_recovers_once_and_keeps_tool_choice_none() -> None: + class FinalizationProvider(RecoveringProvider): + def parse_tool_calls(self, response): + return [ToolCallStub()] if self.complete_calls == 1 else [] + + def request_budget(self, request, *, context_estimate: int, request_options=None): + self.budget_options.append(request_options) + return _fit() if len(self.budget_options) <= 2 else None + + context = HardFitBudgetContext() + provider = FinalizationProvider(recovery=_overflow(cap=64), fail_call=2) + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert provider.complete_calls == 3 + assert provider.requests[-2].tool_choice == "none" + assert provider.requests[-1].tool_choice == "none" + assert provider.requests[-1].max_output_tokens == 64 + + +@pytest.mark.asyncio +async def test_finalization_invalid_recovery_feedback_closes_tool_turn() -> None: + class FinalizationProvider(RecoveringProvider): + def parse_tool_calls(self, response): + return [ToolCallStub()] if self.complete_calls == 1 else [] + + context = HardFitBudgetContext() + provider = FinalizationProvider(recovery=None, fail_call=2) + # Explicitly make recovery unavailable rather than treating None as the + # constructor's valid default. + provider.recovery = None + with pytest.raises(ContextLengthError, match="provider rejected input"): + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert provider.complete_calls == 2 + assert context._messages[-1] == { + "role": "assistant", + "content": "The final response could not be generated because the context is too long.", + } \ No newline at end of file