Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
23 changes: 22 additions & 1 deletion amplifier_module_loop_streaming/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
26 changes: 19 additions & 7 deletions tests/test_provider_budget_context_simple_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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})

Expand Down Expand Up @@ -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)
Expand Down
179 changes: 174 additions & 5 deletions tests/test_provider_budget_guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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",
Expand All @@ -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:
Expand All @@ -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()
Expand Down Expand Up @@ -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)

Expand All @@ -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()
Expand Down
Loading