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
20 changes: 20 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
85 changes: 84 additions & 1 deletion amplifier_module_context_simple/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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'")
Expand Down Expand Up @@ -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 "
Expand Down
243 changes: 242 additions & 1 deletion tests/test_measured_request_view.py
Original file line number Diff line number Diff line change
Expand Up @@ -745,4 +745,245 @@ async def count_view(_view):
)
assert calls == 3
assert not context._truncated_seqs
assert not context._removed_seqs
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"
Loading