diff --git a/README.md b/README.md index 28f7a19..dfc5297 100644 --- a/README.md +++ b/README.md @@ -38,13 +38,14 @@ Provides streaming orchestration that delivers LLM responses token-by-token for ### Provider budget preflight -When a provider exposes the optional synchronous `request_budget` capability, -the loop checks the fully assembled request before dispatch. An oversized -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 a provider exposes the optional `request_budget` capability, the loop +awaits an awaitable result (or accepts a synchronous result) before checking +the fully assembled request for dispatch. An oversized 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. If a provider exposes `request_budget` but returns literal `None` on the initial preflight, the loop treats that request as unavailable and uses the diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 15800a5..7d3dac3 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -3549,6 +3549,8 @@ async def check_request_budget( if accepts_named_request_options(request_budget): budget_kwargs["request_options"] = request_options decision = request_budget(request, **budget_kwargs) + if inspect.isawaitable(decision): + decision = await decision if decision is None: if attempt and not allow_unproven: raise ContextLengthError( diff --git a/tests/test_provider_budget_guard.py b/tests/test_provider_budget_guard.py index 1dff01c..a09d1d0 100644 --- a/tests/test_provider_budget_guard.py +++ b/tests/test_provider_budget_guard.py @@ -50,6 +50,12 @@ def request_budget(self, request, *, context_estimate: int) -> object: return self.decisions.pop(0) +class AsyncBudgetProvider(BudgetProvider): + async def request_budget(self, request, *, context_estimate: int) -> object: + self.budget_calls.append((request, context_estimate)) + return self.decisions.pop(0) + + class BudgetContext(MockContext): """Context double that records ordinary and retention-budget requests.""" @@ -257,6 +263,25 @@ async def test_fitting_budget_dispatches_the_original_request_once() -> None: assert context.request_calls == [([], None)] +@pytest.mark.asyncio +async def test_awaitable_fitting_budget_dispatches_the_original_request_once() -> None: + context = BudgetContext() + provider = AsyncBudgetProvider([_decision(5, 5, 0)]) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 1 + assert len(provider.budget_calls) == 1 + assert context.request_calls == [([], None)] + + @pytest.mark.asyncio async def test_initial_unavailable_budget_keeps_normal_dispatch_without_budget_event() -> None: context = BudgetContext() @@ -322,6 +347,26 @@ async def test_forced_normal_rebuild_forwards_hard_fit_only_to_modern_retention( assert context.hard_fit_calls == [False, True] +@pytest.mark.asyncio +async def test_awaitable_oversize_budget_rebuilds_with_hard_fit_then_dispatches() -> None: + context = HardFitBudgetContext() + context._messages.append({"role": "assistant", "content": "history" * 200}) + provider = AsyncBudgetProvider([_decision(100, 10, 7), _decision(9, 10, 0)]) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({}), + _retaining_coordinator(context), + ) + + assert len(provider.requests) == 1 + assert len(provider.budget_calls) == 2 + assert context.hard_fit_calls == [False, True] + + @pytest.mark.asyncio async def test_forced_rebuild_forwards_hard_fit_to_kwargs_retention() -> None: context = KwargsBudgetContext() @@ -637,6 +682,32 @@ async def test_unavailable_budget_after_rebuild_fails_without_sdk_dispatch() -> assert provider.requests == [] +@pytest.mark.asyncio +async def test_awaitable_unavailable_budget_after_rebuild_fails_without_sdk_dispatch() -> None: + context = HardFitBudgetContext() + provider = AsyncBudgetProvider([_decision(100, 10, 7), None]) + hooks = ScriptedHooks({}) + + with pytest.raises(ContextLengthError, match="unavailable after reporting a concrete budget"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + hooks, + _retaining_coordinator(context), + ) + + assert len(provider.budget_calls) == 2 + assert context.hard_fit_calls == [False, True] + assert provider.requests == [] + assert [ + payload["result"] + for event, payload in hooks.emitted + if event == "orchestrator:provider_budget" + ] == ["oversized"] + + @pytest.mark.asyncio async def test_unavailable_budget_during_output_probe_fails_without_sdk_dispatch() -> None: context = BudgetContext()