From aa2c264d4526e5ebeb80befcbc8e2689024a7c18 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Sun, 20 Sep 2026 02:30:35 -0700 Subject: [PATCH] fix: fit measured output reserves Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 20 ++ amplifier_module_context_simple/__init__.py | 85 ++++++- tests/test_measured_request_view.py | 243 +++++++++++++++++++- 3 files changed, 346 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index cc9d4f8..5242944 100644 --- a/README.md +++ b/README.md @@ -374,6 +374,26 @@ returned transaction is committed by the orchestrator immediately before it dispatches the already-counted request; otherwise it rolls back staged sticky decisions. +After all eight legal Context reduction rungs, a compatible orchestrator may +optionally supply `fit_output(view, attempt)`. Context calls it only when the +final provider-derived hard input estimate still exceeds the limit, passing +the exact notice-inclusive public base view and the exact final counted +attempt. The +orchestrator owns any bounded output-cap ladder, request cloning, reminders, +and options; a successful result returns a newly counted +`{dispatch, budget_decision, count_calls}` envelope. It must carry a usable +provider measurement, a hard-safe input estimate, and a positive count-call +total. `None` means no legal output fit, so Context preserves the existing +fail-closed `ContextLengthError` and rolls back staged decisions. + +This adds no Context protection relaxation and never retries generation. The +Context ladder performs at most nine base counts (the original candidate plus +eight legal rungs); the compatible Loop bounds output fitting to at most six +additional counts. Output relief is reported as `reduced_output`, never as +input-compaction target success. It helps only when the provider's input +allowance grows with a smaller output reserve. Independent input ceilings +remain binding; bounded recounts cannot make oversized required input fit. + ### Legacy actual-meter path Only the **escalation gate** (whether to compact at all, and whether a diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index ba0cf1c..767a1a5 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -916,6 +916,39 @@ async def _count_measured_view( raise TypeError("count_view must return {dispatch, budget_decision}") return public_view, envelope, self._measured_budget_decision(envelope) + async def _fit_measured_output( + self, + fit_output: Callable[ + [list[dict[str, Any]], dict[str, Any]], + Awaitable[dict[str, Any] | None], + ], + view: list[dict[str, Any]], + attempt: dict[str, Any], + ) -> tuple[dict[str, Any], tuple[int, int, int, str], int] | None: + """Ask the request owner for a bounded output-reserve fit.""" + fitted = await fit_output(view, attempt) + if fitted is None: + return None + if not isinstance(fitted, dict) or "dispatch" not in fitted: + raise TypeError("fit_output must return {dispatch, budget_decision, count_calls}") + extra_count_calls = fitted.get("count_calls") + if ( + not isinstance(extra_count_calls, int) + or isinstance(extra_count_calls, bool) + or extra_count_calls <= 0 + ): + raise TypeError("fit_output count_calls must be a positive integer") + measurement = self._measured_budget_decision(fitted) + if measurement is None: + raise TypeError("fit_output must return a usable provider_count measurement") + _, estimated, limit, _ = measurement + if estimated > limit: + raise ContextLengthError( + "Context cannot fit protected content within the provider input " + "limit; required content was not discarded." + ) + return fitted, measurement, extra_count_calls + def _measured_apply_rung( self, rung: int, @@ -1072,13 +1105,20 @@ async def get_measured_request_view( provider: Any, retain_contents: list[str], count_view: Callable[[list[dict[str, Any]]], Awaitable[dict[str, Any]]], + fit_output: Callable[ + [list[dict[str, Any]], dict[str, Any]], + Awaitable[dict[str, Any] | None], + ] + | None = None, ) -> dict[str, Any]: """Build and reduce one complete request using provider-native counts. The loop owns request construction. Context owns only canonical history, its one materialized policy snapshot, and the established eight legal reduction rungs. The returned envelope is the exact - callback object that counted the final provider-facing request. + callback object that counted the final provider-facing request. After + all eight rungs, an optional request-owner callback may reduce only + its output reserve and return a newly counted envelope. """ if self.token_meter != TOKEN_METER_ACTUAL: raise RuntimeError("measured request views require token_meter='actual'") @@ -1320,6 +1360,49 @@ async def get_measured_request_view( measured_after, estimated_after, limit_after, measured_source = final_measurement if estimated_after > limit_after: + if fit_output is not None: + fitted = await self._fit_measured_output( + fit_output, final_view, final_attempt + ) + if fitted is not None: + fitted_attempt, fitted_measurement, extra_count_calls = fitted + ( + fitted_after, + _fitted_estimated, + _fitted_limit, + fitted_source, + ) = fitted_measurement + fitted_count_calls = count_calls + extra_count_calls + if changed_any: + self._stage_measured_stats( + initial_view=initial, + final_view=candidate, + strategy_level=last_changed_rung, + budget=effective_budget, + measured_before=before, + measured_after=fitted_after, + measurement_source=fitted_source, + trigger=trigger, + target=target, + outcome="reduced_output", + count_calls=fitted_count_calls, + ) + else: + # Output relief alone creates no sticky context state. + transaction.rollback() + transaction = None + return { + "base_view": final_view, + "final_attempt": fitted_attempt, + "outcome": "reduced_output", + "measured_before": before, + "measured_after": fitted_after, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": fitted_count_calls, + "transaction": transaction, + } transaction.rollback() raise ContextLengthError( "Context cannot fit protected content within the provider input " diff --git a/tests/test_measured_request_view.py b/tests/test_measured_request_view.py index 44ff662..f688568 100644 --- a/tests/test_measured_request_view.py +++ b/tests/test_measured_request_view.py @@ -745,4 +745,245 @@ async def count_view(_view): ) assert calls == 3 assert not context._truncated_seqs - assert not context._removed_seqs \ No newline at end of file + assert not context._removed_seqs + + +async def _hard_oversize_context(*, reducible=False): + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.8, + target_usage=0.5, + protected_tool_results=0, + compaction_notice_enabled=False, + token_meter="actual", + ) + await context.add_message({"role": "user", "content": "first"}) + if reducible: + for index in range(4): + await context.add_message( + {"role": "tool", "content": "x" * 800, "name": str(index)} + ) + await context.add_message({"role": "user", "content": "last"}) + return context + + +@pytest.mark.asyncio +async def test_output_fit_is_opt_in_and_preserves_existing_terminal_failure(): + context = await _hard_oversize_context() + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert calls == 1 + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_output_fit_returns_exact_fitted_attempt_without_staging_noop_context(): + context = await _hard_oversize_context() + canonical = await context.get_messages() + old_stats = {"old": "stats"} + context._last_compaction_stats = old_stats + counted_attempt = { + "dispatch": {"output_cap": 64_000}, + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + fitted_attempt = { + "dispatch": {"output_cap": 32_000}, + "budget_decision": _decision(900, estimated=900, limit=1_000), + "count_calls": 2, + } + received = {} + + async def count_view(_view): + return counted_attempt + + async def fit_output(view, attempt): + received["view"] = view + received["attempt"] = attempt + return fitted_attempt + + result = await context.get_measured_request_view( + provider=None, + retain_contents=[], + count_view=count_view, + fit_output=fit_output, + ) + + assert result["outcome"] == "reduced_output" + assert result["base_view"] is received["view"] + assert received["attempt"] is counted_attempt + assert result["final_attempt"] is fitted_attempt + assert result["measured_before"] == result["measured_after"] == 900 + assert result["count_calls"] == 3 + assert result["transaction"] is None + assert await context.get_messages() == canonical + assert context._last_compaction_stats is old_stats + + +@pytest.mark.asyncio +async def test_output_fit_none_rolls_back_and_keeps_terminal_failure(): + context = await _hard_oversize_context() + + async def count_view(_view): + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + + async def fit_output(_view, _attempt): + return None + + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await context.get_measured_request_view( + provider=None, + retain_contents=[], + count_view=count_view, + fit_output=fit_output, + ) + + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("fitted", "error"), + [ + ( + { + "budget_decision": _decision(900, estimated=900, limit=1_000), + "count_calls": 1, + }, + TypeError, + ), + ( + { + "dispatch": object(), + "budget_decision": _decision(900, estimated=900, limit=1_000), + "count_calls": True, + }, + TypeError, + ), + ({"dispatch": object(), "budget_decision": None, "count_calls": 1}, TypeError), + ( + { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + "count_calls": 1, + }, + ContextLengthError, + ), + ], +) +async def test_output_fit_rejects_invalid_or_still_oversized_results(fitted, error): + context = await _hard_oversize_context() + + async def count_view(_view): + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + + async def fit_output(_view, _attempt): + return fitted + + with pytest.raises(error): + await context.get_measured_request_view( + provider=None, + retain_contents=[], + count_view=count_view, + fit_output=fit_output, + ) + + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [RuntimeError("fit failed"), asyncio.CancelledError()]) +async def test_output_fit_errors_and_cancellation_rollback_staged_compaction(error): + context = await _hard_oversize_context(reducible=True) + + async def count_view(_view): + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + + async def fit_output(_view, _attempt): + raise error + + with pytest.raises(type(error)): + await context.get_measured_request_view( + provider=None, + retain_contents=[], + count_view=count_view, + fit_output=fit_output, + ) + + assert not context._removed_seqs + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_output_fit_stages_changed_compaction_until_rollback_or_commit(): + async def build_candidate(): + context = await _hard_oversize_context(reducible=True) + counted_views = [] + raw_attempts = [] + fit_attempt = { + "dispatch": {"output_cap": 32_000}, + "budget_decision": _decision(900, estimated=900, limit=1_000), + "count_calls": 2, + } + + async def count_view(view): + counted_views.append(view) + attempt = { + "dispatch": {"output_cap": 64_000}, + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + raw_attempts.append(attempt) + return attempt + + async def fit_output(view, attempt): + assert view is counted_views[-1] + assert attempt is raw_attempts[-1] + return fit_attempt + + result = await context.get_measured_request_view( + provider=None, + retain_contents=[], + count_view=count_view, + fit_output=fit_output, + ) + assert result["outcome"] == "reduced_output" + assert result["final_attempt"] is fit_attempt + assert result["count_calls"] == len(counted_views) + 2 + assert result["transaction"] is not None + assert context._last_compaction_stats["outcome"] == "reduced_output" + assert ( + context._last_compaction_stats["count_calls"] == result["count_calls"] + ) + return context, result + + rolled_back_context, rolled_back = await build_candidate() + rolled_back["transaction"].rollback() + assert not rolled_back_context._removed_seqs + assert not rolled_back_context._truncated_seqs + assert rolled_back_context._last_compaction_stats is None + + committed_context, committed = await build_candidate() + assert await committed["transaction"].commit() + assert committed_context._last_compaction_stats["outcome"] == "reduced_output" \ No newline at end of file