diff --git a/README.md b/README.md index 5c51bfb..f14159f 100644 --- a/README.md +++ b/README.md @@ -105,6 +105,24 @@ retention transaction. This path is additive: contexts and providers without both capabilities keep the legacy request-budget preflight and estimate-based retention behavior. +If the measured Context getter explicitly accepts `fit_output`, the loop also +offers its lossless output-reserve ladder after Context exhausts its eight legal +reduction rungs and still measures a hard oversize. Each of at most six extra +counts uses a deep clone of the frozen request, preserving tools, overlays, +options, and tool choice. Only the output cap changes, plus a view-only warning +below 10,000 output tokens that is included **before** recounting. The exact +accepted counted request is dispatched; hooks, tools, and the system-prompt +factory are not rerun. + +A missing/unusable count during this fitting fails closed. A provider that +does not honor the requested output cap also fails locally. An independent +input ceiling can leave every rung oversized; protected content is never +discarded to force a fit. Older measured getters that lack the keyword keep +their existing behavior. Output fitting can help only when the provider's +reported input allowance grows as its output reserve shrinks. For independent +input ceilings, the bounded probes may add count calls without finding a fit; +this change does not solve an input-only overflow. + For non-streaming foreground responses, the optional `context.foreground_usage` capability records only normalized successful response usage. A counted streaming request records its selected provider count. @@ -112,6 +130,17 @@ If a pre-chunk overflow requires a stream retry, that reading is marked stale rather than attributed to the replacement request. Streams without a provider count do not invent usage; an already-owned reading is marked stale. +### Failed goal turns + +A failed conversational turn ends the active `/goal`, flushes any pending error +completion as final, and emits terminal goal progress without calling the +evaluator, stall judge, or summary model. Task cancellation propagates after +terminal goal progress; unlike a cooperative stop, it creates no completion. +This applies to initial, continuation, and escalation turns. Cleanup diagnostics +are best effort (including cancellation during those diagnostics); the original +turn exception still propagates. No successful response, on-disk persistence, +or absence of earlier provider calls is claimed by this cleanup. + ### Provider-reported overflow recovery A provider may optionally expose synchronous `recover_context_overflow` to diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 9befe0f..0fa9b69 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -47,6 +47,10 @@ logger = logging.getLogger(__name__) +class _MeasuredOutputFitCancelled(BaseException): + """Unwind Context's staged fit into Loop's cooperative cancellation path.""" + + def _build_tool_spec(tool: Any) -> ToolSpec: """Build a `ToolSpec` for one mounted tool, preserving model-native form. @@ -1422,9 +1426,37 @@ async def execute( self._ensure_goal_defaults(initial_goal) goal_turn = (initial_goal["turns_used"] + 1) if initial_goal else None - full_response = await self._execute_one_turn( - prompt, context, providers, tools, hooks, coordinator, goal_turn=goal_turn - ) + async def run_turn(turn_prompt: str, *, goal_turn: int | None) -> str: + try: + return await self._execute_one_turn( + turn_prompt, context, providers, tools, hooks, coordinator, + goal_turn=goal_turn, + ) + except (Exception, asyncio.CancelledError) as error: + if goal_turn is not None and coordinator is not None: + failed_goal = coordinator.session_state.get("goal") + coordinator.session_state["goal"] = None + # Each diagnostic is best effort. Neither may replace the + # original turn failure or prevent the other from flushing. + try: + await self._flush_pending_complete(goal_final=True) + except (Exception, asyncio.CancelledError): + logger.warning("Failed to emit final errored goal completion") + if failed_goal: + state = "cancelled" if isinstance(error, asyncio.CancelledError) else "error" + try: + await hooks.emit( + "orchestrator:goal_progress", + self._goal_progress_payload( + failed_goal, state=state, + reason="Goal turn did not complete; automatic continuation stopped.", + ), + ) + except (Exception, asyncio.CancelledError): + logger.warning("Failed to emit terminal goal progress") + raise + + full_response = await run_turn(prompt, goal_turn=goal_turn) if coordinator is None: return full_response @@ -1700,13 +1732,8 @@ async def execute( trigger=stall_trigger or "idle", verdict=stall_verdict, ) - full_response = await self._execute_one_turn( + full_response = await run_turn( stall_prompt, - context, - providers, - tools, - hooks, - coordinator, goal_turn=goal["turns_used"] + 1, ) is_continuation_turn = True @@ -1747,13 +1774,8 @@ async def execute( ) goal["continuations"] += 1 - full_response = await self._execute_one_turn( + full_response = await run_turn( reason, - context, - providers, - tools, - hooks, - coordinator, goal_turn=goal["turns_used"] + 1, ) is_continuation_turn = True @@ -3956,6 +3978,89 @@ async def try_reduced_output( return candidate_request return None + async def fit_measured_output( + base_view: list[dict[str, Any]], + attempt: dict[str, Any], + *, + request_options: Mapping[str, Any] | None, + ) -> dict[str, Any] | None: + """Recount a frozen protected view at the existing legal output rungs.""" + original = attempt.get("dispatch") + if not isinstance(original, ChatRequest): + raise TypeError("Measured output fitting requires a ChatRequest") + decision = attempt.get("budget_decision") + if validate_measured_budget_decision(decision) is None: + raise ContextLengthError("Measured output fitting requires a provider count") + original_cap = decision.get("max_output_tokens") + if original.max_output_tokens is not None and original_cap is not None: + original_cap = min(original_cap, original.max_output_tokens) + for count_calls, cap in enumerate(output_cap_candidates(original_cap), 1): + if coordinator and coordinator.cancellation.is_cancelled: + raise _MeasuredOutputFitCancelled() + # Clone the already-assembled request, not the underlying history: + # preserve overlays, tools, choice, metadata, and frozen options. + candidate = original.model_copy(deep=True) + candidate.max_output_tokens = cap + candidate.messages.extend( + Message(**message) for message in degraded_output_warning(cap) + ) + decision = await count_measured_request( + candidate, base_view, request_options=request_options + ) + if coordinator and coordinator.cancellation.is_cancelled: + raise _MeasuredOutputFitCancelled() + if decision is None or validate_measured_budget_decision(decision) is None: + raise ContextLengthError( + "Provider count became unavailable during measured output fitting" + ) + if decision.get("max_output_tokens") != cap: + raise ContextLengthError( + "Provider did not honor the measured output cap" + ) + if decision["estimated_input_tokens"] <= decision["input_limit_tokens"]: + return { + "dispatch": candidate, + "budget_decision": decision, + "count_calls": count_calls, + } + return None + + async def request_measured_view( + retain_contents: list[str], count_view: Any, + request_options: Mapping[str, Any] | None, + ) -> dict[str, Any]: + kwargs = { + "provider": provider, + "retain_contents": retain_contents, + "count_view": count_view, + } + # Explicit keyword negotiation keeps both mixed-version pairings + # on their existing path. Never catch an implementation TypeError. + try: + parameter = inspect.signature(measured_view_getter).parameters.get("fit_output") + except (TypeError, ValueError): + parameter = None + if parameter is not None and parameter.kind in ( + inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD, + ): + kwargs["fit_output"] = partial( + fit_measured_output, request_options=request_options + ) + return await measured_view_getter(**kwargs) + + async def emit_measured_output_degradation(result: dict[str, Any], transaction) -> None: + if result.get("outcome") != "reduced_output": + return + request = result["final_attempt"]["dispatch"] + if request.max_output_tokens is not None and request.max_output_tokens < 10_000: + await await_with_measured_rollback( + hooks.emit( + "orchestrator:context_degradation", + {"mode": "reduced_output", "max_output_tokens": request.max_output_tokens}, + ), + transaction, + ) + async def recover_context_overflow( failed_request: ChatRequest, error: ContextLengthError, @@ -4650,11 +4755,13 @@ async def count_view( ), } - measured_result = await measured_view_getter( - provider=provider, - retain_contents=retained_contents, - count_view=count_view, - ) + try: + measured_result = await request_measured_view( + retained_contents, count_view, request_options, + ) + except _MeasuredOutputFitCancelled: + await exit_for_cancellation() + return if not isinstance(measured_result, dict): raise TypeError("context.measured_request_view returned a non-dictionary result") measured_transaction = measured_result.get("transaction") @@ -4715,6 +4822,7 @@ async def count_view( ), measured_transaction, ) + await emit_measured_output_degradation(measured_result, measured_transaction) if coordinator and coordinator.cancellation.is_cancelled: if measured_transaction is not None: measured_transaction.rollback() @@ -5632,10 +5740,8 @@ async def count_final_view( ), } - measured_result = await measured_view_getter( - provider=provider, - retain_contents=final_retained_contents, - count_view=count_final_view, + measured_result = await request_measured_view( + final_retained_contents, count_final_view, request_options, ) if not isinstance(measured_result, dict): raise TypeError( @@ -5703,6 +5809,9 @@ async def count_final_view( ), final_measured_transaction, ) + await emit_measured_output_degradation( + measured_result, final_measured_transaction + ) else: max_iter_chat_request = build_chat_request( message_dicts, tool_choice="none" @@ -5945,6 +6054,12 @@ async def count_final_view( "The final response could not be generated." ) + except _MeasuredOutputFitCancelled: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return except asyncio.CancelledError: await close_finalization_tool_turn( "The previous operation was cancelled. Results from completed tools have been preserved." diff --git a/tests/test_goal_loop.py b/tests/test_goal_loop.py index 43b5d6a..cdd72f2 100644 --- a/tests/test_goal_loop.py +++ b/tests/test_goal_loop.py @@ -18,6 +18,7 @@ from __future__ import annotations +import asyncio import json from collections.abc import Callable from types import SimpleNamespace @@ -3903,3 +3904,182 @@ async def _run(config: dict) -> str: assert "EXIT_STATUS_END code=0" in after_prompt # Tool-call arguments are now visible too. assert "/var/log/huge.log" in after_prompt + + +# --------------------------------------------------------------------------- +# 15. Goal-turn execution failures must finalize the failed deferred turn once, +# clear state, and still propagate the original failure. +# --------------------------------------------------------------------------- + + +class _ConversationalFailureProvider(FakeProvider): + """Uses the real loop; only the selected conversational transport fails.""" + + def __init__(self, failures: dict[int, BaseException]) -> None: + super().__init__() + self.failures = failures + self.conversation_calls = 0 + + async def complete(self, chat_request, **kwargs): + system = next((m for m in chat_request.messages if m.role == "system"), None) + system_text = system.content if system and isinstance(system.content, str) else "" + if not any( + marker in system_text + for marker in ("tool-less evaluator", "tool-less judge", "single, short line for a developer") + ): + self.conversation_calls += 1 + failure = self.failures.get(self.conversation_calls) + if failure is not None: + raise failure + return await super().complete(chat_request, **kwargs) + + +def _active_goal() -> dict[str, Any]: + return {"condition": "solve it", "turns_used": 0, "last_reason": None, "cap": None} + + +def _goal_completions(hooks: MockHooks) -> list[dict]: + return hooks.orchestrator_complete_events() + + +@pytest.mark.asyncio +class TestGoalTurnFailureCleanup: + async def test_first_goal_turn_failure_propagates_identity_and_cannot_leak_a_later_completion( + self, + ) -> None: + error = RuntimeError("first conversational transport failed") + orch = _make_orchestrator() + ctx, hooks, coordinator = MockContext(), MockHooks(), MockCoordinator() + provider = _ConversationalFailureProvider({1: error}) + coordinator.session_state["goal"] = _active_goal() + + with pytest.raises(RuntimeError) as raised: + await orch.execute("solve it", ctx, {"main": provider}, {}, hooks, coordinator) + + assert raised.value is error + assert coordinator.session_state["goal"] is None + assert orch._pending_orchestrator_complete is None + completions = _goal_completions(hooks) + assert len(completions) == 1 + assert completions[0]["status"] == "error" + assert completions[0]["goal_final"] is True + progress = hooks.goal_progress_events() + assert len(progress) == 1 + assert progress[0]["state"] == "error" + assert provider.eval_call_requests == [] + assert provider.summary_call_count == 0 + + # A later ordinary invocation must create only its own completion; it + # must never flush the prior failed goal turn a second time. + provider.turn_queue.append(MockTurnResponse(text="ordinary later success")) + assert await orch.execute("later", ctx, {"main": provider}, {}, hooks, coordinator) == ( + "ordinary later success" + ) + assert len(_goal_completions(hooks)) == 2 + assert _goal_completions(hooks)[1]["goal_turn"] is None + + async def test_continuation_failure_keeps_earlier_completion_nonfinal_and_stops_evaluation( + self, + ) -> None: + error = RuntimeError("continuation transport failed") + orch = _make_orchestrator() + ctx, hooks, coordinator = MockContext(), MockHooks(), MockCoordinator() + provider = _ConversationalFailureProvider({2: error}) + provider.turn_queue.append(MockTurnResponse(text="first answer")) + provider.eval_queue.append((False, "continue with a different approach")) + coordinator.session_state["goal"] = _active_goal() + + with pytest.raises(RuntimeError) as raised: + await orch.execute("solve it", ctx, {"main": provider}, {}, hooks, coordinator) + + assert raised.value is error + assert coordinator.session_state["goal"] is None + assert orch._pending_orchestrator_complete is None + completions = _goal_completions(hooks) + assert [(item["status"], item["goal_final"]) for item in completions] == [ + ("success", False), + ("error", True), + ] + assert len(provider.eval_call_requests) == 1 + assert provider.summary_call_count == 0 + assert [item["state"] for item in hooks.goal_progress_events()] == [ + "continuing", + "error", + ] + + async def test_stall_escalation_turn_failure_stops_without_summary_or_extra_judge( + self, + ) -> None: + error = RuntimeError("escalation transport failed") + orch = _make_orchestrator({"goal_stall_threshold": 1, "goal_busy_stall_window": 100}) + ctx, hooks, coordinator = MockContext(), MockHooks(), MockCoordinator() + provider = _ConversationalFailureProvider({3: error}) + provider.turn_queue.extend([MockTurnResponse(text="initial"), MockTurnResponse(text="idle")]) + provider.eval_queue.extend([(False, "blocked"), (False, "still blocked")]) + provider.judge_queue.append((True, "static blocker", "NOT_DEMONSTRATED")) + coordinator.session_state["goal"] = _active_goal() + + with pytest.raises(RuntimeError) as raised: + await orch.execute("solve it", ctx, {"main": provider}, {}, hooks, coordinator) + + assert raised.value is error + assert coordinator.session_state["goal"] is None + assert orch._pending_orchestrator_complete is None + assert len(provider.eval_call_requests) == 2 + assert len(provider.judge_call_requests) == 1 + assert provider.summary_call_count == 0 + completions = _goal_completions(hooks) + assert [(item["status"], item["goal_final"]) for item in completions] == [ + ("success", False), + ("success", False), + ("error", True), + ] + assert [item["state"] for item in hooks.goal_progress_events()] == [ + "continuing", + "continuing", + "error", + ] + + @pytest.mark.parametrize("event", ["orchestrator:complete", "orchestrator:goal_progress"]) + @pytest.mark.parametrize("hook_error", [RuntimeError, asyncio.CancelledError]) + async def test_terminal_diagnostic_hook_failures_do_not_mask_original_turn_error( + self, event, hook_error, + ) -> None: + error = RuntimeError("original turn error") + + class FailingTerminalHooks(MockHooks): + async def emit(self, event_name: str, payload: dict | None = None): + result = await super().emit(event_name, payload) + if event_name == event: + raise hook_error("diagnostic hook failed") + return result + + orch = _make_orchestrator() + ctx, hooks, coordinator = MockContext(), FailingTerminalHooks(), MockCoordinator() + provider = _ConversationalFailureProvider({1: error}) + coordinator.session_state["goal"] = _active_goal() + + with pytest.raises(RuntimeError) as raised: + await orch.execute("solve it", ctx, {"main": provider}, {}, hooks, coordinator) + + assert raised.value is error + assert coordinator.session_state["goal"] is None + assert orch._pending_orchestrator_complete is None + assert len(_goal_completions(hooks)) == 1 + assert len(hooks.goal_progress_events()) == 1 + + async def test_cancelled_active_goal_clears_state_and_propagates_cancellation(self) -> None: + error = asyncio.CancelledError() + orch = _make_orchestrator() + ctx, hooks, coordinator = MockContext(), MockHooks(), MockCoordinator() + provider = _ConversationalFailureProvider({1: error}) + coordinator.session_state["goal"] = _active_goal() + + with pytest.raises(asyncio.CancelledError) as raised: + await orch.execute("solve it", ctx, {"main": provider}, {}, hooks, coordinator) + + assert raised.value is error + assert coordinator.session_state["goal"] is None + assert orch._pending_orchestrator_complete is None + assert _goal_completions(hooks) == [] + assert [event["state"] for event in hooks.goal_progress_events()] == ["cancelled"] diff --git a/tests/test_measured_output_fitting.py b/tests/test_measured_output_fitting.py new file mode 100644 index 0000000..8de506a --- /dev/null +++ b/tests/test_measured_output_fitting.py @@ -0,0 +1,680 @@ +"""Joint regressions for measured output-reserve fitting. + +These run the real StreamingOrchestrator with the real context-simple measured +view. Only the provider transport/count endpoint and hooks are stubbed: the +assertions therefore cover the complete Context -> Loop contract rather than a +fake measured-view result. +""" + +from __future__ import annotations + +import asyncio +from copy import deepcopy +from types import SimpleNamespace +from typing import Any + +import pytest +from amplifier_core import ContextLengthError, ToolResult + +pytest.importorskip("amplifier_module_context_simple") + +from amplifier_module_context_simple import SimpleContextManager +from amplifier_module_loop_streaming import StreamingOrchestrator + + +class _HookResult: + action = "continue" + reason = None + ephemeral = False + context_injection = None + context_injection_role = "system" + append_to_last_tool_result = False + data = None + + +class _Hooks: + def __init__(self) -> None: + self.events: list[tuple[str, dict[str, Any]]] = [] + + async def emit(self, name: str, payload: dict[str, Any] | None = None): + self.events.append((name, payload or {})) + if name == "provider:request": + result = _HookResult() + result.action = "inject_context" + result.ephemeral = True + result.context_injection = "ACTIVE-OVERLAY" + return result + return _HookResult() + + def payloads(self, name: str) -> list[dict[str, Any]]: + return [payload for event, payload in self.events if event == name] + + +class _Cancellation: + is_cancelled = False + is_immediate = False + state = "running" + + def register_tool_start(self, *_args) -> None: + pass + + def register_tool_complete(self, *_args) -> None: + pass + + async def trigger_callbacks(self) -> None: + pass + + +class _Coordinator: + def __init__(self) -> None: + self.cancellation = _Cancellation() + self.session_state: dict[str, Any] = {} + self.capabilities: dict[str, Any] = {} + + def register_capability(self, name: str, capability: Any) -> None: + self.capabilities[name] = capability + + def get_capability(self, name: str) -> Any: + return self.capabilities.get(name) + + async def process_hook_result(self, result, *_args): + return result + + +class _Tool: + name = "preserved_tool" + description = "records a real tool declaration in the counted request" + input_schema = {"type": "object", "properties": {}} + + async def execute(self, _arguments) -> ToolResult: + return ToolResult(success=True, output="tool result") + + +class _MeasuredTransport: + """Native-count provider whose answer is a deterministic function of cap.""" + + def __init__(self, *, fit_at: int | None) -> None: + self.fit_at = fit_at + self.counted: list[Any] = [] + self.counted_options: list[dict[str, Any] | None] = [] + self.requests: list[Any] = [] + self.complete_options: list[dict[str, Any]] = [] + + def get_info(self): + return SimpleNamespace( + capabilities=["request_budget:provider_count"], + defaults={"context_window": 100_000, "max_output_tokens": 64_000}, + ) + + @staticmethod + def _warning_present(request) -> bool: + return any( + "orchestrator-context-degraded" in str(message.content) + for message in request.messages + ) + + def request_budget(self, request, *, context_estimate: int, request_options=None): + self.counted.append(request) + self.counted_options.append(request_options) + cap = request.max_output_tokens or 64_000 + warning = self._warning_present(request) + # The synthetic native count describes input, not reserved output. + # Only the warning increases it; the model's allowance changes by cap. + estimated = 89_503 if self.fit_at == 1_000 else 45_000 + if warning: + estimated += 2 # 6,400 output allows 89,504, so the warning matters. + allowance = 100_000 - cap - 4_096 + if self.fit_at is None: + allowance = min(allowance, 30_000 - 4_096) + return { + "estimated_input_tokens": estimated, + "input_limit_tokens": allowance, + "context_token_budget": 30_000, + "max_output_tokens": cap, + "measurement": { + "kind": "provider_count", + "source": "test.measured-output", + "input_tokens": estimated, + }, + } + + async def complete(self, request, **kwargs): + self.requests.append(request) + self.complete_options.append(kwargs) + return SimpleNamespace(text="accepted", content=None, usage=None) + + def parse_tool_calls(self, _response): + return [] + + +async def _context_with_protected_markers(*, hooks=None) -> tuple[SimpleContextManager, list[int]]: + """Create a real actual-meter Context with no legal input reductions.""" + context = SimpleContextManager( + max_tokens=200_000, + compact_threshold=0.99, + target_usage=0.50, + protected_recent=1.0, + protected_tool_results=1, + compaction_notice_enabled=False, + token_meter="actual", + hooks=hooks, + ) + factory_calls = [0] + + async def factory() -> str: + factory_calls[0] += 1 + return "SYSTEM-FIRST-MARKER" + + await context.set_system_prompt_factory(factory) + await context.add_message({"role": "user", "content": "FIRST-USER-MARKER"}) + await context.add_message( + {"role": "developer", "content": "DEVELOPER-REQUIRED-MARKER"} + ) + await context.add_message( + { + "role": "assistant", + "content": "TOOL-PAIR-ASSISTANT-MARKER", + "tool_calls": [ + {"id": "kept-tool-call", "name": "preserved_tool", "arguments": {}} + ], + } + ) + await context.add_message( + { + "role": "tool", + "name": "preserved_tool", + "tool_call_id": "kept-tool-call", + "content": "TOOL-PAIR-RESULT-MARKER", + } + ) + await context.add_message( + { + "role": "user", + "content": "REQUIRED-REMINDER-MARKER", + "metadata": {"ephemeral": True, "source": "hook"}, + } + ) + return context, factory_calls + + +def _wire_without_cap(request) -> dict[str, Any]: + """A deep snapshot of every request field except the intended cap change.""" + payload = request.model_dump(mode="python") + payload.pop("max_output_tokens", None) + return deepcopy(payload) + + +async def _run(provider: _MeasuredTransport, *, stream: bool = False): + context, factory_calls = await _context_with_protected_markers() + hooks = _Hooks() + coordinator = _Coordinator() + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + loop = StreamingOrchestrator( + { + "extended_thinking": True, + "ephemeral_injection_mode": "tail", + "reminder_placement": "tail", + } + ) + if stream: + async def stream_method(request, *, tools): + provider.requests.append(request) + yield {"content": "streamed"} + + provider.stream = stream_method # type: ignore[attr-defined] + result = await loop.execute( + "LATEST-USER-MARKER", + context, + {"main": provider}, + {"preserved_tool": _Tool()}, + hooks, + coordinator, + ) + return result, context, hooks, coordinator, factory_calls + + +@pytest.mark.asyncio +async def test_measured_fit_dispatches_the_exact_32k_counted_request_without_input_loss() -> None: + """64k hard-oversize can retain all input by lowering only to the 32k rung.""" + provider = _MeasuredTransport(fit_at=32_000) + result, context, hooks, _coordinator, factory_calls = await _run(provider) + + assert result == "accepted" + assert factory_calls == [1] + assert len(provider.counted) == 2 # 64k initial count, then legal 50% rung. + assert [request.max_output_tokens for request in provider.counted] == [None, 32_000] + assert len(provider.requests) == 1 + counted_final = provider.counted[-1] + assert provider.requests[0] is counted_final + assert _wire_without_cap(provider.counted[0]) == _wire_without_cap(counted_final) + assert provider.complete_options == [{"extended_thinking": True}] + assert provider.counted_options[-1] == {"extended_thinking": True} + body = "\n".join(str(message.content) for message in provider.requests[0].messages) + for marker in ( + "SYSTEM-FIRST-MARKER", + "DEVELOPER-REQUIRED-MARKER", + "TOOL-PAIR-ASSISTANT-MARKER", + "TOOL-PAIR-RESULT-MARKER", + "REQUIRED-REMINDER-MARKER", + "FIRST-USER-MARKER", + "ACTIVE-OVERLAY", + "LATEST-USER-MARKER", + ): + assert marker in body + assert provider.requests[0].tools is not None + assert context._last_compaction_stats is None # Output relief alone is non-sticky. + assert hooks.payloads("orchestrator:context_degradation") == [] + assert len(hooks.payloads("provider:request")) == 1 + + +@pytest.mark.asyncio +async def test_measured_fit_streams_once_with_the_fitted_counted_request() -> None: + provider = _MeasuredTransport(fit_at=32_000) + result, _context, _hooks, _coordinator, _factory_calls = await _run(provider, stream=True) + + assert result == "streamed" + assert len(provider.requests) == 1 + assert provider.requests[0] is provider.counted[-1] + assert provider.requests[0].max_output_tokens == 32_000 + assert provider.complete_options == [] + + +@pytest.mark.asyncio +async def test_measured_fit_fails_after_exactly_six_legal_rungs_without_sdk_dispatch() -> None: + """An immutable protected input gets one initial count plus all six rungs.""" + provider = _MeasuredTransport(fit_at=None) + context, _factory_calls = await _context_with_protected_markers() + hooks = _Hooks() + coordinator = _Coordinator() + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + loop = StreamingOrchestrator( + {"ephemeral_injection_mode": "tail", "reminder_placement": "tail"} + ) + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await loop.execute( + "LATEST-USER-MARKER", context, {"main": provider}, {}, hooks, coordinator + ) + + assert len(provider.counted) == 7 + assert [request.max_output_tokens for request in provider.counted] == [ + None, + 32_000, + 25_600, + 19_200, + 12_800, + 6_400, + 1_000, + ] + assert provider.requests == [] + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_severe_measured_fit_counts_one_view_warning_and_emits_one_degradation_event() -> None: + provider = _MeasuredTransport(fit_at=1_000) + result, context, hooks, _coordinator, _factory_calls = await _run(provider) + + assert result == "accepted" + assert len(provider.counted) == 7 + severe = [ + request for request in provider.counted + if request.max_output_tokens is not None and request.max_output_tokens < 10_000 + ] + assert [request.max_output_tokens for request in severe] == [6_400, 1_000] + assert all(_MeasuredTransport._warning_present(request) for request in severe) + assert provider.requests[0].max_output_tokens == 1_000 + assert sum( + "orchestrator-context-degraded" in str(message.content) + for message in provider.requests[0].messages + ) == 1 + assert hooks.payloads("orchestrator:context_degradation") == [ + {"mode": "reduced_output", "max_output_tokens": 1_000} + ] + canonical = await context.get_messages() + assert not any( + "orchestrator-context-degraded" in str(message.get("content")) + for message in canonical + ) + + +@pytest.mark.asyncio +async def test_measured_finalization_recounts_to_low_cap_keeps_tools_and_disables_new_calls() -> None: + class ToolCall: + id = "only-tool" + name = "preserved_tool" + arguments: dict[str, Any] = {} + + class FinalizingTransport(_MeasuredTransport): + def request_budget(self, request, *, context_estimate: int, request_options=None): + # The conversational request fits at 32k, but the complete tool + # transcript plus finalization overlay needs the 1k floor. + old_fit = self.fit_at + self.fit_at = 1_000 if request.tool_choice == "none" else 32_000 + try: + return super().request_budget( + request, + context_estimate=context_estimate, + request_options=request_options, + ) + finally: + self.fit_at = old_fit + + async def complete(self, request, **kwargs): + self.requests.append(request) + self.complete_options.append(kwargs) + if len(self.requests) == 1: + return SimpleNamespace( + text="calling tool", + content=None, + usage=None, + _calls=[ToolCall()], + ) + return SimpleNamespace(text="final answer", content=None, usage=None, _calls=[]) + + def parse_tool_calls(self, response): + return response._calls + + provider = FinalizingTransport(fit_at=32_000) + context, _factory_calls = await _context_with_protected_markers() + hooks, coordinator = _Hooks(), _Coordinator() + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + loop = StreamingOrchestrator( + { + "max_iterations": 1, + "extended_thinking": True, + "ephemeral_injection_mode": "tail", + "reminder_placement": "tail", + } + ) + + assert await loop.execute( + "LATEST-USER-MARKER", + context, + {"main": provider}, + {"preserved_tool": _Tool()}, + hooks, + coordinator, + ) == "final answer" + + assert len(provider.requests) == 2 + final = provider.requests[-1] + assert final.max_output_tokens == 1_000 + assert final.tool_choice == "none" + assert final.tools is not None + assert provider.complete_options == [{"extended_thinking": True}] * 2 + assert [request.max_output_tokens for request in provider.counted[:2]] == [None, 32_000] + # Finalization may first apply legal history reductions at the original cap. + final_counts = [request for request in provider.counted if request.tool_choice == "none"] + assert sum(request.max_output_tokens is None for request in final_counts) == 2 + assert [request.max_output_tokens for request in final_counts[-6:]] == [ + 32_000, + 25_600, + 19_200, + 12_800, + 6_400, + 1_000, + ] + assert final is final_counts[-1] + assert all(request.tools == final.tools for request in final_counts) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["none", "absent", "malformed"]) +async def test_measured_recount_loss_fails_closed_and_leaves_no_sticky_context(failure: str) -> None: + class RecountFailureTransport(_MeasuredTransport): + def request_budget(self, request, *, context_estimate: int, request_options=None): + if request.max_output_tokens is not None: + self.counted.append(request) + self.counted_options.append(request_options) + if failure == "none": + return None + if failure == "absent": + return { + "estimated_input_tokens": 100_000, + "input_limit_tokens": 90_000, + "context_token_budget": 30_000, + "max_output_tokens": request.max_output_tokens, + } + return { + "estimated_input_tokens": 100_000, + "input_limit_tokens": 90_000, + "context_token_budget": 30_000, + "max_output_tokens": 0, + "measurement": {"kind": "provider_count", "source": "bad", "input_tokens": 45_000}, + } + return super().request_budget( + request, context_estimate=context_estimate, request_options=request_options + ) + + provider = RecountFailureTransport(fit_at=None) + context, _factory_calls = await _context_with_protected_markers() + hooks = _Hooks() + # Build explicitly here so the failed call still lets us inspect the real Context. + coordinator = _Coordinator() + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + loop = StreamingOrchestrator({"ephemeral_injection_mode": "tail", "reminder_placement": "tail"}) + with pytest.raises(ContextLengthError): + await loop.execute("CURRENT", context, {"main": provider}, {}, hooks, coordinator) + + assert provider.requests == [] + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("phase", ["counter", "budget_event", "delay"]) +async def test_measured_output_fit_cancellation_rolls_back_before_dispatch(phase) -> None: + context, _ = await _context_with_protected_markers() + context.protected_recent = 0 + await context.add_message({"role": "assistant", "content": "discardable old detail"}) + coordinator = _Coordinator() + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + + class CancellingTransport(_MeasuredTransport): + async def request_budget(self, request, **kwargs): + result = super().request_budget(request, **kwargs) + if phase == "counter" and request.max_output_tokens is not None: + raise asyncio.CancelledError() + return result + + class CancellingHooks(_Hooks): + async def emit(self, event, payload=None): + result = await super().emit(event, payload) + if phase == "budget_event" and event == "orchestrator:provider_budget": + coordinator.cancellation.is_cancelled = True + return result + + provider = CancellingTransport(fit_at=32_000) + loop = StreamingOrchestrator({}) + + async def cancelling_delay(*_args): + raise asyncio.CancelledError() + + if phase == "delay": + loop._apply_rate_limit_delay = cancelling_delay + call = loop.execute("CURRENT", context, {"main": provider}, {}, CancellingHooks(), coordinator) + if phase == "budget_event": + await call + else: + with pytest.raises(asyncio.CancelledError): + await call + assert provider.requests == [] + assert len(provider.counted) >= 3 # A reduction rung was staged before the fit. + assert not context._removed_seqs + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_measured_output_fit_commits_staged_compaction_before_dispatch() -> None: + hooks, coordinator = _Hooks(), _Coordinator() + context, _ = await _context_with_protected_markers(hooks=hooks) + context.protected_recent = 0 + await context.add_message({"role": "assistant", "content": "discardable old detail"}) + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + + class CommitObservingProvider(_MeasuredTransport): + async def complete(self, request, **kwargs): + assert context._removed_seqs + assert context._last_compaction_stats["outcome"] == "reduced_output" + assert len(hooks.payloads("context:compaction")) == 1 + return await super().complete(request, **kwargs) + + provider = CommitObservingProvider(fit_at=32_000) + await StreamingOrchestrator({}).execute("CURRENT", context, {"main": provider}, {}, hooks, coordinator) + assert len(provider.counted) >= 3 + assert provider.requests == [provider.counted[-1]] + assert context._last_compaction_stats["count_calls"] == len(provider.counted) + assert any(m.get("content") == "discardable old detail" for m in await context.get_messages()) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("finalization", [False, True]) +@pytest.mark.parametrize("when", ["before_fit", "during_fit"]) +async def test_cooperative_cancel_during_output_fit_uses_normal_cancel_lifecycle( + finalization, when, +) -> None: + from types import SimpleNamespace + + context, _ = await _context_with_protected_markers() + context.protected_recent = 0 + await context.add_message({"role": "assistant", "content": "discardable old detail"}) + coordinator = _Coordinator() + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + hooks = _Hooks() + + class CancellingProvider(_MeasuredTransport): + def request_budget(self, request, **kwargs): + # For finalization, complete a normal tool turn before the failure. + if finalization and request.tool_choice != "none": + self.counted.append(request) + self.counted_options.append(kwargs.get("request_options")) + return { + "estimated_input_tokens": 1, "input_limit_tokens": 100_000, + "context_token_budget": 1, "max_output_tokens": 64_000, + "measurement": {"kind": "provider_count", "source": "test", "input_tokens": 1}, + } + result = super().request_budget(request, **kwargs) + if when == "during_fit" and request.max_output_tokens is not None: + coordinator.cancellation.is_cancelled = True + elif when == "before_fit" and context._removed_seqs: + coordinator.cancellation.is_cancelled = True + return result + + async def complete(self, request, **kwargs): + self.requests.append(request) + return SimpleNamespace(text="", content=None, usage=None) + + def parse_tool_calls(self, response): + return [SimpleNamespace(id="new-tool", name="preserved_tool", arguments={})] + + provider = CancellingProvider(fit_at=None) + loop = StreamingOrchestrator({"max_iterations": 1}) + await loop.execute( + "CURRENT", context, {"main": provider}, {"preserved_tool": _Tool()}, hooks, coordinator + ) + assert len(provider.requests) == (1 if finalization else 0) + assert len(hooks.payloads("cancel:requested")) == 1 + assert len(hooks.payloads("cancel:completed")) == 1 + assert hooks.payloads("orchestrator:complete")[-1]["status"] == "cancelled" + assert not context._removed_seqs + assert context._last_compaction_stats is None + caps = [r.max_output_tokens for r in provider.counted if r.max_output_tokens is not None] + assert caps == ([32_000] if when == "during_fit" else []) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("original_cap", [None, 1_000]) +async def test_measured_fit_requires_a_reducible_reported_cap(original_cap) -> None: + class NoLadderTransport(_MeasuredTransport): + def request_budget(self, request, **kwargs): + result = super().request_budget(request, **kwargs) + if original_cap is None: + result.pop("max_output_tokens") + else: + result["max_output_tokens"] = original_cap + return result + + provider = NoLadderTransport(fit_at=None) + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await _run(provider) + assert all(request.max_output_tokens is None for request in provider.counted) + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_provider_ignoring_reduced_output_cap_is_not_dispatched() -> None: + class IgnoringCapTransport(_MeasuredTransport): + def request_budget(self, request, **kwargs): + result = super().request_budget(request, **kwargs) + result["max_output_tokens"] = 64_000 + return result + + provider = IgnoringCapTransport(fit_at=32_000) + with pytest.raises(ContextLengthError, match="did not honor"): + await _run(provider) + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_old_measured_getter_is_called_without_the_new_keyword() -> None: + context, _ = await _context_with_protected_markers() + coordinator = _Coordinator() + calls = [] + + async def old_getter(*, provider, retain_contents, count_view): + calls.append(True) + return await context.get_measured_request_view( + provider=provider, retain_contents=retain_contents, count_view=count_view + ) + + coordinator.register_capability("context.measured_request_view", old_getter) + provider = _MeasuredTransport(fit_at=32_000) + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await StreamingOrchestrator({}).execute("CURRENT", context, {"main": provider}, {}, _Hooks(), coordinator) + assert calls == [True] + assert all(request.max_output_tokens is None for request in provider.counted) + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_negotiated_callback_typeerror_propagates_without_legacy_retry() -> None: + context, _ = await _context_with_protected_markers() + coordinator = _Coordinator() + calls = [] + failure = TypeError("callback implementation failed") + + async def broken_getter(*, provider, retain_contents, count_view, fit_output=None): + calls.append(fit_output) + raise failure + + coordinator.register_capability("context.measured_request_view", broken_getter) + provider = _MeasuredTransport(fit_at=32_000) + with pytest.raises(TypeError) as raised: + await StreamingOrchestrator({}).execute("CURRENT", context, {"main": provider}, {}, _Hooks(), coordinator) + assert raised.value is failure + assert len(calls) == 1 and callable(calls[0]) + assert provider.requests == [] + + +@pytest.mark.asyncio +async def test_measured_protected_floor_error_finalizes_goal_without_generation() -> None: + context, _ = await _context_with_protected_markers() + coordinator = _Coordinator() + coordinator.register_capability("context.measured_request_view", context.get_measured_request_view) + coordinator.session_state["goal"] = {"condition": "finish", "turns_used": 0, "cap": None} + hooks = _Hooks() + provider = _MeasuredTransport(fit_at=None) + loop = StreamingOrchestrator({}) + with pytest.raises(ContextLengthError): + await loop.execute("CURRENT", context, {"main": provider}, {}, hooks, coordinator) + assert provider.requests == [] # Includes evaluators and summaries, not just main calls. + assert coordinator.session_state["goal"] is None + assert loop._pending_orchestrator_complete is None + complete = hooks.payloads("orchestrator:complete") + assert len(complete) == 1 + assert complete[0]["status"] == "error" and complete[0]["goal_final"] is True + assert [p["state"] for p in hooks.payloads("orchestrator:goal_progress")] == ["error"] diff --git a/tests/test_provider_budget_context_simple_runtime.py b/tests/test_provider_budget_context_simple_runtime.py index 61f3943..330477c 100644 --- a/tests/test_provider_budget_context_simple_runtime.py +++ b/tests/test_provider_budget_context_simple_runtime.py @@ -211,7 +211,7 @@ async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: coordinator = _Coordinator(hooks) context = _RecordingContext( # The character-based ordinary estimate stays below context-simple's - # real provider-derived budget; max_tokens is only a fallback here. + # real provider-derived budget; max_tokens does not lower it here. # Each supplementary Han character serializes as four UTF-8 bytes, so # the calibrated provider preflight forces the first hard-fit rebuild. max_tokens=500_000, @@ -251,7 +251,9 @@ async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None: assert client.responses.hard_fit_counts_at_dispatch == [0] assert "gpt-5-mini" in provider._budget_calibration - bulk = "REMOVED-BULK-MARKER:" + ("\U00020000" * 40_000) + # Stay oversized after mini's documented input ceiling was corrected to + # 272k: 70k four-byte characters, while chars/4 stays below Context's budget. + bulk = "REMOVED-BULK-MARKER:" + ("\U00020000" * 70_000) await context.add_message({"role": "assistant", "content": bulk}) await loop.execute(