From d5b7727ece08fd89276cf08d578e01aac1d06d77 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 15 Sep 2026 06:28:26 -0700 Subject: [PATCH] fix: retry a provider-directed context rebuild once Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 16 +- amplifier_module_loop_streaming/__init__.py | 228 ++++++++++++++--- ..._provider_budget_context_simple_runtime.py | 67 +++++ tests/test_provider_budget_guard.py | 231 +++++++++++++++++- 4 files changed, 503 insertions(+), 39 deletions(-) diff --git a/README.md b/README.md index d36eac7..917c7cb 100644 --- a/README.md +++ b/README.md @@ -40,22 +40,30 @@ Provides streaming orchestration that delivers LLM responses token-by-token for When a provider exposes the optional synchronous `request_budget` capability, the loop checks the fully assembled request before dispatch. An oversized -request gets exactly one smaller, retention-aware context view; required +request gets up to two smaller, retention-aware context views; required reminders and request-only injections are replayed without running hooks or draining pending state again. If that view still cannot fit, the loop raises locally and makes no SDK request. Providers without the capability retain the existing request and dispatch behavior. When `context.request_retention` advertises its optional `hard_fit` keyword, -that one provider-forced rebuild forwards `hard_fit=True`, allowing the context -to target the provider's requested budget directly. Older retention +each provider-forced rebuild forwards `hard_fit=True`, allowing the context to +target the provider's requested budget directly. The second rebuild runs only +when the provider requests a strictly smaller budget. Older retention capabilities, uninspectable dynamic callables, and the generic context fallback keep their existing `provider`/`retain_contents`/`token_budget` assembly; the -second preflight and provider's final payload guard remain the safety boundary. +final preflight and provider's final payload guard remain the safety boundary. The existing `orchestrator:provider_budget` event exposes each preflight's attempt, result, estimate, allowance, and requested context budget to mounted observability consumers. +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 +preflights: they do not resend a provider request or omit input. At caps below +10,000 tokens, a view-only system reminder asks the model to tell the user that +the session is degraded and to recommend a new session for substantial work. + ## Configuration ```toml diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index c4d0cc7..0385a87 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -3462,7 +3462,10 @@ async def request_messages( ) def build_chat_request( - message_dicts: list[dict[str, Any]], *, tool_choice: str | None = None + message_dicts: list[dict[str, Any]], + *, + tool_choice: str | None = None, + max_output_tokens: int | None = None, ) -> ChatRequest: tools_list = [_build_tool_spec(tool) for tool in tools.values()] if tools else None kwargs: dict[str, Any] = { @@ -3472,6 +3475,8 @@ def build_chat_request( } if tool_choice is not None: kwargs["tool_choice"] = tool_choice + if max_output_tokens is not None: + kwargs["max_output_tokens"] = max_output_tokens return ChatRequest( **kwargs ) @@ -3481,11 +3486,11 @@ async def check_request_budget( base_messages: list[dict[str, Any]], *, attempt: int, - ) -> int | None: - """Return one requested smaller context budget, or ``None`` when it fits.""" + ) -> tuple[int | None, int | None]: + """Return a smaller context budget and effective output cap when available.""" request_budget = getattr(provider, "request_budget", None) if not budget_capable or not callable(request_budget): - return None + return None, None context_estimate = sum(len(str(message)) // 4 for message in base_messages) decision = request_budget(request, context_estimate=context_estimate) required = ( @@ -3506,6 +3511,15 @@ async def check_request_budget( estimated = decision["estimated_input_tokens"] allowance = decision["input_limit_tokens"] target = decision["context_token_budget"] + 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" + ) fits = estimated <= allowance await hooks.emit( "orchestrator:provider_budget", @@ -3519,12 +3533,70 @@ async def check_request_budget( }, ) if fits: - return None - if target <= 0: - raise ContextLengthError( - "Provider request exceeds its input budget and cannot retain a smaller context" + return None, output_cap + return target, output_cap + + def output_cap_candidates(original: int | None) -> list[int]: + """Return the bounded lossless output-reserve ladder.""" + if original is None or original <= 1_000: + return [] + candidates: list[int] = [] + for fraction in (0.50, 0.40, 0.30, 0.20, 0.10): + cap = max(1_000, int(original * fraction)) + if cap < original and (not candidates or cap < candidates[-1]): + candidates.append(cap) + if candidates[-1:] != [1_000]: + candidates.append(1_000) + return candidates + + def degraded_output_warning(cap: int) -> list[dict[str, Any]]: + """Add a view-only warning only at the severe output-cap tier.""" + if cap >= 10_000: + return [] + return [ + { + "role": "system", + "content": ( + "\n" + "This request is running with a severely reduced response budget " + "because the conversation is near the provider context limit. " + "Give the user a concise answer, explain that the session is in " + "a degraded state, and recommend starting a new session for " + "substantial further work.\n" + "" + ), + "metadata": {"ephemeral": True}, + } + ] + + async def try_reduced_output( + message_dicts: list[dict[str, Any]], + base_messages: list[dict[str, Any]], + *, + original_output_cap: int | None, + tool_choice: str | None = None, + ) -> ChatRequest | None: + """Preflight unchanged input at bounded lower output reserves.""" + for cap in output_cap_candidates(original_output_cap): + candidate_request = build_chat_request( + message_dicts + degraded_output_warning(cap), + tool_choice=tool_choice, + max_output_tokens=cap, ) - return target + smaller_budget, _ = await check_request_budget( + candidate_request, base_messages, attempt=1 + ) + if smaller_budget is None: + if cap < 10_000: + await hooks.emit( + "orchestrator:context_degradation", + { + "mode": "reduced_output", + "max_output_tokens": cap, + }, + ) + return candidate_request + return None # --- Turn-start reminder assembly (reminder-redesign-spec.md, # W1.2, Option D). Hoists iteration 1's provider:request emit to @@ -4038,10 +4110,14 @@ async def exit_for_cancellation() -> None: f"[ORCHESTRATOR] Tool names: {[t.name for t in tools.values()]}" ) - smaller_context_budget = await check_request_budget( + smaller_context_budget, original_output_cap = await check_request_budget( chat_request, base_message_dicts, attempt=0 ) if smaller_context_budget is not None: + if smaller_context_budget <= 0: + raise ContextLengthError( + "Provider request exceeds its input budget and cannot retain a smaller context" + ) rebuilt_base_messages = list( await request_messages( retained_contents, @@ -4060,15 +4136,57 @@ async def exit_for_cancellation() -> None: else rebuilt_base_messages ) rebuilt_request = build_chat_request(rebuilt_messages) - if ( - await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=1 - ) - ) is not None: - raise ContextLengthError( - "Provider request remains over budget after one context rebuild" + next_context_budget, _ = await check_request_budget( + rebuilt_request, rebuilt_base_messages, attempt=1 + ) + 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, ) - chat_request = rebuilt_request + if reduced_output_request is not None: + chat_request = reduced_output_request + else: + if next_context_budget <= 0: + raise ContextLengthError( + "Provider request exceeds its input budget and cannot retain " + "a smaller context" + ) + if next_context_budget >= smaller_context_budget: + raise ContextLengthError( + "Provider request remains over budget and a second " + "context budget would not reduce it" + ) + rebuilt_base_messages = list( + await request_messages( + retained_contents, + token_budget=next_context_budget, + hard_fit=True, + ) + ) + rebuilt_messages = ( + _replay_request_overlays( + rebuilt_base_messages, + 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 rebuilt_base_messages + ) + rebuilt_request = build_chat_request(rebuilt_messages) + if ( + await check_request_budget( + rebuilt_request, rebuilt_base_messages, attempt=2 + ) + )[0] is not None: + raise ContextLengthError( + "Provider request remains over budget after two context rebuilds" + ) + chat_request = rebuilt_request + else: + chat_request = rebuilt_request # Apply rate limit delay before provider call await self._apply_rate_limit_delay(hooks, iteration) @@ -4723,10 +4841,15 @@ async def exit_for_cancellation() -> None: max_iter_chat_request = build_chat_request( message_dicts, tool_choice="none" ) - smaller_context_budget = await check_request_budget( + smaller_context_budget, original_output_cap = await check_request_budget( max_iter_chat_request, base_message_dicts, attempt=0 ) if smaller_context_budget is not None: + if smaller_context_budget <= 0: + raise ContextLengthError( + "Provider request exceeds its input budget and cannot retain " + "a smaller context" + ) rebuilt_base_messages = list( await request_messages( final_retained_contents, @@ -4748,16 +4871,62 @@ async def exit_for_cancellation() -> None: rebuilt_request = build_chat_request( rebuilt_messages, tool_choice="none" ) - if ( - await check_request_budget( - rebuilt_request, rebuilt_base_messages, attempt=1 - ) - ) is not None: - raise ContextLengthError( - "Provider finalization request remains over budget " - "after one context rebuild" + next_context_budget, _ = await check_request_budget( + rebuilt_request, rebuilt_base_messages, attempt=1 + ) + 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, + tool_choice="none", ) - max_iter_chat_request = rebuilt_request + if reduced_output_request is not None: + max_iter_chat_request = reduced_output_request + else: + if next_context_budget <= 0: + raise ContextLengthError( + "Provider request exceeds its input budget and cannot retain " + "a smaller context" + ) + if next_context_budget >= smaller_context_budget: + raise ContextLengthError( + "Provider finalization request remains over budget and a " + "second context budget would not reduce it" + ) + rebuilt_base_messages = list( + await request_messages( + final_retained_contents, + token_budget=next_context_budget, + hard_fit=True, + ) + ) + rebuilt_messages = ( + _replay_request_overlays( + rebuilt_base_messages, + 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(rebuilt_base_messages) + ) + rebuilt_messages.append(finalization_overlay) + rebuilt_request = build_chat_request( + rebuilt_messages, tool_choice="none" + ) + if ( + await check_request_budget( + rebuilt_request, rebuilt_base_messages, attempt=2 + ) + )[0] is not None: + raise ContextLengthError( + "Provider finalization request remains over budget " + "after two context rebuilds" + ) + max_iter_chat_request = rebuilt_request + else: + max_iter_chat_request = rebuilt_request kwargs = {} if self.extended_thinking: @@ -4840,6 +5009,9 @@ async def exit_for_cancellation() -> None: ) raise except ContextLengthError: + await close_finalization_tool_turn( + "The final response could not be generated because the context is too long." + ) raise except LLMError as e: await hooks.emit( diff --git a/tests/test_provider_budget_context_simple_runtime.py b/tests/test_provider_budget_context_simple_runtime.py index d4b6e94..e5c13e0 100644 --- a/tests/test_provider_budget_context_simple_runtime.py +++ b/tests/test_provider_budget_context_simple_runtime.py @@ -20,6 +20,7 @@ from amplifier_module_context_simple import SimpleContextManager from amplifier_module_loop_streaming import StreamingOrchestrator from amplifier_module_provider_openai import OpenAIProvider +from tests.test_ephemeral_cache_persist_mode import RequestCapturingProvider class _Cancellation: @@ -134,10 +135,76 @@ def __init__(self, hard_fit_calls: list[bool]) -> None: self.responses = _InMemoryResponses(hard_fit_calls) +class _TwoRebuildProvider(RequestCapturingProvider): + """Budget double that forces two local rebuilds before accepting a request.""" + + def __init__(self) -> None: + super().__init__() + self.budget_calls: list[tuple[ChatRequest, int]] = [] + self._decisions = [ + { + "estimated_input_tokens": 100, + "input_limit_tokens": 10, + "context_token_budget": 5_000, + }, + { + "estimated_input_tokens": 50, + "input_limit_tokens": 10, + "context_token_budget": 1_000, + }, + { + "estimated_input_tokens": 9, + "input_limit_tokens": 10, + "context_token_budget": 0, + }, + ] + + def request_budget( + self, request: ChatRequest, *, context_estimate: int + ) -> dict[str, int]: + self.budget_calls.append((request, context_estimate)) + return self._decisions.pop(0) + + def _payload_text(params: dict) -> str: return json.dumps(params, ensure_ascii=False, sort_keys=True) +@pytest.mark.asyncio +async def test_two_hard_fits_keep_real_context_requirements() -> None: + hooks = _StableReminderHooks() + coordinator = _Coordinator(hooks) + context = _RecordingContext( + max_tokens=100_000, + compact_threshold=0.99, + target_usage=0.50, + protected_recent=0.10, + protected_tool_results=1, + truncate_chars=64, + compaction_notice_enabled=True, + ) + coordinator.register_capability( + "context.request_retention", context.get_messages_for_request_retaining + ) + provider = _TwoRebuildProvider() + bulk = "TWO-HARD-FIT-BULK:" + ("history" * 10_000) + await context.add_message({"role": "assistant", "content": bulk}) + + await StreamingOrchestrator({}).execute( + "CURRENT-HUMAN", context, {"budget": provider}, {}, hooks, coordinator + ) + + assert len(provider.requests) == 1 + assert len(provider.budget_calls) == 3 + assert context.hard_fit_calls == [False, True, True] + request_bodies = "\n".join(message.content for message in provider.requests[0].messages) + assert "TWO-HARD-FIT-BULK" not in request_bodies + assert "CURRENT-HUMAN" in request_bodies + assert "REQUIRED-REMINDER" in request_bodies + canonical = await context.get_messages() + assert any(message.get("content") == bulk for message in canonical) + + @pytest.mark.asyncio async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: hooks = _StableReminderHooks() diff --git a/tests/test_provider_budget_guard.py b/tests/test_provider_budget_guard.py index 4e756e8..2503b8c 100644 --- a/tests/test_provider_budget_guard.py +++ b/tests/test_provider_budget_guard.py @@ -30,6 +30,15 @@ def _decision(estimated: int, limit: int, target: int) -> dict[str, int]: } +def _output_decision( + estimated: int, limit: int, target: int, output_cap: int +) -> dict[str, int]: + return { + **_decision(estimated, limit, target), + "max_output_tokens": output_cap, + } + + class BudgetProvider(RequestCapturingProvider): def __init__(self, decisions: list[dict[str, int]]) -> None: super().__init__() @@ -472,11 +481,107 @@ async def test_irreducible_budget_makes_no_sdk_call() -> None: @pytest.mark.asyncio -async def test_second_oversize_after_one_rebuild_makes_no_sdk_call() -> None: +async def test_second_oversize_rebuilds_again_and_dispatches_once() -> None: + context = BudgetContext() + provider = BudgetProvider( + [_decision(100, 10, 7), _decision(50, 10, 1), _decision(9, 10, 0)] + ) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.budget_calls) == 3 + assert [budget for _, budget in context.request_calls] == [None, 7, 1] + assert context.request_calls[0][0] == context.request_calls[1][0] == context.request_calls[2][0] + assert len(provider.requests) == 1 + + +@pytest.mark.asyncio +async def test_output_cap_ladder_preserves_input_and_warns_model_at_one_thousand() -> None: + context = BudgetContext() + context._messages.append({"role": "assistant", "content": "history" * 200}) + provider = BudgetProvider( + [ + _output_decision(100, 10, 7, 128_000), + _output_decision(50, 10, 0, 128_000), + _output_decision(50, 10, 1, 64_000), + _output_decision(50, 10, 1, 51_200), + _output_decision(50, 10, 1, 38_400), + _output_decision(50, 10, 1, 25_600), + _output_decision(50, 10, 1, 12_800), + _output_decision(9, 10, 0, 1_000), + ] + ) + body = "REQUIRED" + hooks = ScriptedHooks({"provider:request": _injection(body)}) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + hooks, + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 1 + request = provider.requests[0] + assert request.max_output_tokens == 1_000 + request_bodies = "\n".join(message.content for message in request.messages) + assert body in request_bodies + assert "history" not in request_bodies + assert "orchestrator-context-degraded" in request_bodies + assert "recommend starting a new session" in request_bodies + assert [request.max_output_tokens for request, _ in provider.budget_calls] == [ + None, + None, + 64_000, + 51_200, + 38_400, + 25_600, + 12_800, + 1_000, + ] + assert [name for name, _ in hooks.emitted].count("provider:request") == 1 + assert [name for name, _ in hooks.emitted].count( + "orchestrator:context_degradation" + ) == 1 + + +@pytest.mark.asyncio +async def test_third_oversize_after_two_rebuilds_makes_no_sdk_call() -> None: + context = BudgetContext() + provider = BudgetProvider( + [_decision(100, 10, 7), _decision(50, 10, 1), _decision(20, 10, 1)] + ) + + with pytest.raises(ContextLengthError, match="after two context rebuilds"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.budget_calls) == 3 + assert [budget for _, budget in context.request_calls] == [None, 7, 1] + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_non_decreasing_second_budget_makes_no_extra_rebuild_or_sdk_call() -> None: context = BudgetContext() - provider = BudgetProvider([_decision(100, 10, 7), _decision(50, 10, 1)]) + provider = BudgetProvider([_decision(100, 10, 7), _decision(50, 10, 7)]) - with pytest.raises(ContextLengthError, match="remains over budget"): + with pytest.raises(ContextLengthError, match="would not reduce"): await StreamingOrchestrator({}).execute( "work", context, @@ -488,7 +593,6 @@ async def test_second_oversize_after_one_rebuild_makes_no_sdk_call() -> None: assert len(provider.budget_calls) == 2 assert [budget for _, budget in context.request_calls] == [None, 7] - assert context.request_calls[0][0] == context.request_calls[1][0] assert provider.requests == [] @@ -525,7 +629,9 @@ async def test_malformed_budget_result_fails_before_dispatch(decision) -> None: async def test_budget_replay_keeps_tail_overlay_once_without_rerunning_hooks() -> None: context = BudgetContext() context._messages.append({"role": "assistant", "content": "history" * 200}) - provider = BudgetProvider([_decision(100, 10, 7), _decision(9, 10, 0)]) + provider = BudgetProvider( + [_decision(100, 10, 7), _decision(50, 10, 1), _decision(9, 10, 0)] + ) body = "ONCE" hooks = ScriptedHooks({"provider:request": _injection(body)}) @@ -545,7 +651,7 @@ async def test_budget_replay_keeps_tail_overlay_once_without_rerunning_hooks() - provider_requests = [name for name, _ in hooks.emitted if name == "provider:request"] assert provider_requests == ["provider:request"] assert context.legacy_calls == [] - assert [budget for _, budget in context.request_calls] == [None, 7] + assert [budget for _, budget in context.request_calls] == [None, 7, 1] @pytest.mark.asyncio @@ -685,6 +791,70 @@ async def test_forced_finalization_rebuild_forwards_hard_fit_only_at_rebuild() - assert context.hard_fit_calls == [False, False, True] +@pytest.mark.asyncio +async def test_finalization_rebuilds_again_before_its_single_sdk_dispatch() -> None: + context = HardFitBudgetContext() + provider = FinalizingBudgetProvider( + [ + _decision(1, 10, 0), + _decision(100, 10, 7), + _decision(50, 10, 1), + _decision(1, 10, 0), + ] + ) + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 2 + assert len(provider.budget_calls) == 4 + assert context.hard_fit_calls == [False, False, True, True] + assert provider.requests[-1].tool_choice == "none" + + +@pytest.mark.asyncio +async def test_finalization_uses_output_cap_ladder_before_another_context_rebuild() -> None: + context = HardFitBudgetContext() + provider = FinalizingBudgetProvider( + [ + _output_decision(1, 10, 0, 128_000), + _output_decision(100, 10, 7, 128_000), + _output_decision(50, 10, 0, 128_000), + _output_decision(50, 10, 0, 64_000), + _output_decision(50, 10, 0, 51_200), + _output_decision(50, 10, 0, 38_400), + _output_decision(50, 10, 0, 25_600), + _output_decision(50, 10, 0, 12_800), + _output_decision(9, 10, 0, 1_000), + ] + ) + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 2 + assert len(provider.budget_calls) == 9 + assert context.hard_fit_calls == [False, False, True] + final_request = provider.requests[-1] + assert final_request.tool_choice == "none" + assert final_request.max_output_tokens == 1_000 + assert "orchestrator-context-degraded" in "\n".join( + message.content for message in final_request.messages + ) + + class AnthropicStyleAssemblyProvider(RequestCapturingProvider): """Non-budget control: preserve ordinary assembled requests for other providers.""" @@ -769,13 +939,60 @@ async def test_finalization_irreducible_budget_skips_its_sdk_dispatch() -> None: assert len(provider.requests) == 1 assert len(provider.budget_calls) == 2 + messages = await context.get_messages() + assert messages[-2]["role"] == "tool" + assert messages[-1] == { + "role": "assistant", + "content": "The final response could not be generated because the context is too long.", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("final_target", [0, 1]) +async def test_finalization_rejection_after_two_rebuilds_closes_tool_turn( + final_target: int, +) -> None: + context = HardFitBudgetContext() + provider = FinalizingBudgetProvider( + [ + _decision(1, 10, 0), + _decision(100, 10, 7), + _decision(50, 10, 1), + _decision(20, 10, final_target), + ] + ) + + with pytest.raises(ContextLengthError): + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 1 + assert len(provider.budget_calls) == 4 + assert context.hard_fit_calls == [False, False, True, True] + messages = await context.get_messages() + assert messages[-2]["role"] == "tool" + assert messages[-1]["role"] == "assistant" + assert messages[-1]["content"] == ( + "The final response could not be generated because the context is too long." + ) @pytest.mark.asyncio async def test_finalization_replays_current_and_pending_tail_overlays_once() -> None: context = BudgetContext() provider = FinalizingBudgetProvider( - [_decision(1, 10, 0), _decision(100, 10, 7), _decision(1, 10, 0)] + [ + _decision(1, 10, 0), + _decision(100, 10, 7), + _decision(50, 10, 1), + _decision(1, 10, 0), + ] ) direct = ScriptedHookResult( action="inject_context",