diff --git a/README.md b/README.md index 917c7cb..5a3e497 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,13 @@ 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. +If a provider exposes `request_budget` but returns literal `None` on the +initial preflight, the loop treats that request as unavailable and uses the +same normal dispatch path. Once a concrete budget has required a rebuild or +output-reserve probe, `None` (or a removed capability) is a local +capability-loss error: it never counts as a fit or permits an unchecked SDK +request. + When `context.request_retention` advertises its optional `hard_fit` keyword, each provider-forced rebuild forwards `hard_fit=True`, allowing the context to target the provider's requested budget directly. The second rebuild runs only diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 0385a87..74f9852 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -3487,12 +3487,33 @@ async def check_request_budget( *, attempt: int, ) -> tuple[int | None, int | None]: - """Return a smaller context budget and effective output cap when available.""" + """Return a smaller context budget and effective output cap when available. + + ``None`` is a compatibility result only for the initial preflight: + it means this provider cannot make a trustworthy decision for this + request, so dispatch follows the historical no-capability path. + Every later probe is reachable only after a concrete decision made + a rebuild or output-reserve retry necessary. Losing the capability + then must fail locally rather than turn an unknown result into a + fit and send an unchecked request. + """ request_budget = getattr(provider, "request_budget", None) if not budget_capable or not callable(request_budget): + if attempt: + 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) + if decision is None: + if attempt: + raise ContextLengthError( + "Provider request_budget capability was unavailable after " + "reporting a concrete budget" + ) + return None, None required = ( "estimated_input_tokens", "input_limit_tokens", diff --git a/tests/test_provider_budget_context_simple_runtime.py b/tests/test_provider_budget_context_simple_runtime.py index e5c13e0..61f3943 100644 --- a/tests/test_provider_budget_context_simple_runtime.py +++ b/tests/test_provider_budget_context_simple_runtime.py @@ -213,7 +213,7 @@ async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: # The character-based ordinary estimate stays below context-simple's # real provider-derived budget; max_tokens is only a fallback here. # Each supplementary Han character serializes as four UTF-8 bytes, so - # the provider's payload preflight forces the first hard-fit rebuild. + # the calibrated provider preflight forces the first hard-fit rebuild. max_tokens=500_000, compact_threshold=0.99, target_usage=0.50, @@ -239,6 +239,18 @@ async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: ) loop = StreamingOrchestrator({}) + # Cold byte estimates no longer force compaction. Establish calibration + # through a real completed transport response before testing warm hard fit. + await provider.complete( + ChatRequest( + messages=[Message(role="user", content="CALIBRATION-ONLY-CONTROL")], + max_output_tokens=1024, + ) + ) + assert len(client.responses.calls) == 1 + assert client.responses.hard_fit_counts_at_dispatch == [0] + assert "gpt-5-mini" in provider._budget_calibration + bulk = "REMOVED-BULK-MARKER:" + ("\U00020000" * 40_000) await context.add_message({"role": "assistant", "content": bulk}) @@ -275,18 +287,18 @@ async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: prompt, context, {"openai": provider}, {}, hooks, coordinator ) - # Every call in this list reached the fake SDK. The first accepted call - # followed the provider-directed rebuild; the four subsequent calls prove - # ordinary fetches do not resurrect the canonical bulk history. - assert len(client.responses.calls) >= 5 - payloads = [_payload_text(params) for params in client.responses.calls] + # Exclude only the explicit calibration control above. The first workload + # call followed the provider-directed rebuild; the four subsequent calls + # prove ordinary fetches do not resurrect the canonical bulk history. + assert len(client.responses.calls) >= 6 + payloads = [_payload_text(params) for params in client.responses.calls[1:]] assert all("REMOVED-BULK-MARKER" not in payload for payload in payloads) first_payload = payloads[0] assert "ORIGINAL-HUMAN" in first_payload assert "REQUIRED-REMINDER" in first_payload assert context.hard_fit_calls[:2] == [False, True] assert context.hard_fit_calls.count(True) == 1 - assert client.responses.hard_fit_counts_at_dispatch == [1] * len(payloads) + assert client.responses.hard_fit_counts_at_dispatch[1:] == [1] * len(payloads) post_force_payloads = payloads[1:] assert post_force_payloads assert all("ORIGINAL-HUMAN" in payload for payload in post_force_payloads) diff --git a/tests/test_provider_budget_guard.py b/tests/test_provider_budget_guard.py index 2503b8c..210b531 100644 --- a/tests/test_provider_budget_guard.py +++ b/tests/test_provider_budget_guard.py @@ -40,12 +40,12 @@ def _output_decision( class BudgetProvider(RequestCapturingProvider): - def __init__(self, decisions: list[dict[str, int]]) -> None: + def __init__(self, decisions: list[object]) -> None: super().__init__() self.decisions = list(decisions) self.budget_calls: list[tuple[object, int]] = [] - def request_budget(self, request, *, context_estimate: int) -> dict[str, int]: + def request_budget(self, request, *, context_estimate: int) -> object: self.budget_calls.append((request, context_estimate)) return self.decisions.pop(0) @@ -257,6 +257,27 @@ async def test_fitting_budget_dispatches_the_original_request_once() -> None: assert context.request_calls == [([], None)] +@pytest.mark.asyncio +async def test_initial_unavailable_budget_keeps_normal_dispatch_without_budget_event() -> None: + context = BudgetContext() + provider = BudgetProvider([None]) + hooks = ScriptedHooks({}) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + hooks, + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 1 + assert len(provider.budget_calls) == 1 + assert context.request_calls == [([], None)] + assert [name for name, _ in hooks.emitted].count("orchestrator:provider_budget") == 0 + + @pytest.mark.asyncio async def test_one_smaller_retained_view_is_rechecked_and_dispatches_once() -> None: context = BudgetContext() @@ -596,6 +617,52 @@ async def test_non_decreasing_second_budget_makes_no_extra_rebuild_or_sdk_call() assert provider.requests == [] +@pytest.mark.asyncio +async def test_unavailable_budget_after_rebuild_fails_without_sdk_dispatch() -> None: + context = BudgetContext() + provider = BudgetProvider([_decision(100, 10, 7), None]) + + with pytest.raises(ContextLengthError, match="unavailable after reporting a concrete budget"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.budget_calls) == 2 + assert [budget for _, budget in context.request_calls] == [None, 7] + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_unavailable_budget_during_output_probe_fails_without_sdk_dispatch() -> None: + context = BudgetContext() + provider = BudgetProvider( + [ + _output_decision(100, 10, 7, 128_000), + _output_decision(50, 10, 1, 128_000), + None, + ] + ) + + with pytest.raises(ContextLengthError, match="unavailable after reporting a concrete budget"): + 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] + assert provider.requests == [] + + @pytest.mark.asyncio @pytest.mark.parametrize( "decision", @@ -605,7 +672,8 @@ async def test_non_decreasing_second_budget_makes_no_extra_rebuild_or_sdk_call() {"estimated_input_tokens": -1, "input_limit_tokens": 10, "context_token_budget": 7}, {"estimated_input_tokens": float("nan"), "input_limit_tokens": 10, "context_token_budget": 7}, {"estimated_input_tokens": "1", "input_limit_tokens": 10, "context_token_budget": 7}, - None, + [], + 0, ], ) async def test_malformed_budget_result_fails_before_dispatch(decision) -> None: @@ -625,6 +693,27 @@ async def test_malformed_budget_result_fails_before_dispatch(decision) -> None: assert provider.requests == [] +@pytest.mark.asyncio +async def test_budget_probe_exception_propagates_without_sdk_dispatch() -> None: + class ExplodingBudgetProvider(RequestCapturingProvider): + def request_budget(self, request, *, context_estimate: int) -> object: + raise RuntimeError("budget probe exploded") + + provider = ExplodingBudgetProvider() + + with pytest.raises(RuntimeError, match="budget probe exploded"): + await StreamingOrchestrator({}).execute( + "work", + BudgetContext(), + {"main": provider}, + {}, + ScriptedHooks({}), + MockCoordinator(), + ) + + assert provider.requests == [] + + @pytest.mark.asyncio async def test_budget_replay_keeps_tail_overlay_once_without_rerunning_hooks() -> None: context = BudgetContext() @@ -743,12 +832,12 @@ async def test_pending_tool_overlay_is_replayed_once_after_budget_rebuild() -> N class FinalizingBudgetProvider(NRoundToolProvider): - def __init__(self, decisions: list[dict[str, int]] | None = None) -> None: + def __init__(self, decisions: list[object] | None = None) -> None: super().__init__(n_tool_rounds=1) self.budget_calls: list[object] = [] self.decisions = decisions or [_decision(1, 10, 0), _decision(1, 10, 0)] - def request_budget(self, request, *, context_estimate: int) -> dict[str, int]: + def request_budget(self, request, *, context_estimate: int) -> object: self.budget_calls.append(request) return self.decisions.pop(0) @@ -772,6 +861,86 @@ async def test_finalization_request_is_budget_checked_before_dispatch() -> None: assert provider.requests[-1].tool_choice == "none" +@pytest.mark.asyncio +async def test_finalization_initial_unavailable_budget_keeps_normal_dispatch() -> None: + context = BudgetContext() + provider = FinalizingBudgetProvider([None, None]) + + 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) == 2 + assert provider.requests[-1].tool_choice == "none" + + +@pytest.mark.asyncio +async def test_finalization_unavailable_budget_after_rebuild_closes_tool_turn() -> None: + context = HardFitBudgetContext() + provider = FinalizingBudgetProvider( + [_decision(1, 10, 0), _decision(100, 10, 7), None] + ) + + with pytest.raises(ContextLengthError, match="unavailable after reporting a concrete budget"): + 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) == 3 + assert context.hard_fit_calls == [False, False, True] + 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 +async def test_finalization_unavailable_budget_during_output_probe_closes_tool_turn() -> None: + context = HardFitBudgetContext() + provider = FinalizingBudgetProvider( + [ + _decision(1, 10, 0), + _output_decision(100, 10, 7, 128_000), + _output_decision(50, 10, 1, 128_000), + None, + ] + ) + + with pytest.raises(ContextLengthError, match="unavailable after reporting a concrete budget"): + 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] + 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 async def test_forced_finalization_rebuild_forwards_hard_fit_only_at_rebuild() -> None: context = HardFitBudgetContext()