diff --git a/README.md b/README.md index 2ab7559..61bd438 100644 --- a/README.md +++ b/README.md @@ -251,23 +251,31 @@ a child session -- this module has no concept of "delegation" itself; it only counts main-loop LLM calls against whatever `max_iterations` it was given, root session or child). -**Exhaustion is a normal turn ending, not an error.** When `iteration` -reaches `max_iterations`, the loop makes **one additional** tool-less -`provider.complete()` call with an injected -`` asking the agent to -wrap up and summarize -- so a budget of `N` permits at most `N + 1` -main-loop provider calls, not `N`. This wrap-up call ends the turn with a -normal return (`execution:end` fires, no exception raised), so the caller's -usual persistence path runs and the resulting transcript is complete and -resumable. +**Exhaustion is a normal turn ending, not an error.** A response that ends +naturally (no tool call and no pending steer) returns immediately, including +on iteration `max_iterations`; it makes no duplicate provider call and is +not marked budget-exhausted. Only when the hard limit prevents a required +continuation does the loop make one additional `provider.complete()` call +with an injected `` asking +the agent to wrap up and summarize. Thus a budget of `N` permits at most +`N + 1` main-loop provider calls, while a natural completion uses exactly the +calls it needed. + +That final request retains the ordinary generic tool declarations (including +provider-native declarations needed to validate preceding tool history), but +sets the portable `tool_choice="none"`. A compliant final response retains +safe text/thinking blocks and provider metadata in the transcript while +rendering normalized text; any unexpected tool call is neither dispatched nor +persisted structurally, so finalization cannot extend the iteration budget or +leave unpaired tool state. `ORCHESTRATOR_COMPLETE`'s payload always carries a `metadata` bag: ```python { - "llm_calls": 300, # main-loop iterations actually used (not goal-loop internal calls -- see below) + "llm_calls": 301, # actual main-loop provider calls, including any one finalization call (not goal-loop internal calls -- see below) "llm_call_budget": 300, # the max_iterations this turn ran under, or None if unlimited - "budget_exhausted": True, # whether this turn hit the budget (vs. finishing early) + "budget_exhausted": True, # whether the budget prevented a needed continuation "resumable": True, # whether this exit path guarantees the transcript was persisted } ``` diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 5e16195..2d29ba1 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -1137,6 +1137,10 @@ def __init__(self, config: dict[str, Any]): # AttributeError. self._tool_calls_this_turn: int = 0 self._goal_turn_evidence: dict[str, Any] | None = None + # Actual main-loop provider calls for the active turn. This differs + # from the bounded iteration count when a forced finalization call is + # needed after the loop has spent its budget. + self._llm_calls_this_turn: int = 0 # Layer 1 call-budget bookkeeping (spec: 298-replacement). Both are # per-execute()-call state, reset in _execute_one_turn alongside # _tool_calls_this_turn. Initialized here so a fresh instance never @@ -1709,6 +1713,7 @@ async def _execute_one_turn( # "did this turn run any tools" count once this method returns. self._tool_calls_this_turn = 0 self._goal_turn_evidence = {"tools": []} if goal_turn is not None else None + self._llm_calls_this_turn = 0 # Layer 1 call-budget bookkeeping, reset per turn alongside # _tool_calls_this_turn above (spec: 298-replacement). self._budget_exhausted = False @@ -1786,7 +1791,7 @@ async def _execute_one_turn( # module has no visibility into whether an upstream cancellation # (e.g. a delegate's hard wall-clock deadline) skips the save. "metadata": { - "llm_calls": iteration_count, + "llm_calls": self._llm_calls_this_turn, "llm_call_budget": ( self.max_iterations if self.max_iterations != -1 else None ), @@ -3353,10 +3358,22 @@ async def _execute_stream( await context.add_message({"role": "user", "content": prompt}) iteration = 0 - - while self.max_iterations == -1 or iteration < self.max_iterations: - # Check for cancellation at iteration start - if coordinator and coordinator.cancellation.is_cancelled: + # A bounded loop needs a finalization call only when the final + # budgeted response explicitly required another model turn. A normal + # no-tool break is already a complete answer, even when it happens on + # the final permitted iteration. + natural_completion = False + continuation_needed = False + + async def close_finalization_tool_turn(message: str) -> None: + """Close a capped tool-result turn when finalization cannot answer.""" + messages = await context.get_messages() + if messages and messages[-1].get("role") == "tool": + await context.add_message({"role": "assistant", "content": message}) + + async def exit_for_cancellation() -> None: + """Emit the normal cancellation lifecycle before ending this turn.""" + if coordinator: # Emit cancel:requested on first detection and trigger cleanup callbacks if not self._cancel_requested_emitted: self._cancel_requested_emitted = True @@ -3370,7 +3387,7 @@ async def _execute_stream( ) try: await coordinator.cancellation.trigger_callbacks() - except Exception as e: + except Exception as e: # noqa: BLE001 logger.warning(f"Error in cancellation callbacks: {e}") # Emit cancel:completed โ€” orchestrator is exiting due to cancellation await hooks.emit( @@ -3381,11 +3398,14 @@ async def _execute_stream( "turn_count": iteration, }, ) - # Don't yield more content, just exit. - # Clear any pending steers so they cannot leak into the next turn - # (cancellation means "stop now" โ€” stale steers have no next injection - # point and must not silently ride a future, unrelated turn). (spec ยง5.2) - self._steering_queue.clear() + # Cancellation means "stop now": queued steers cannot silently + # carry into a future, unrelated turn. + self._steering_queue.clear() + + while self.max_iterations == -1 or iteration < self.max_iterations: + # Check for cancellation at iteration start + if coordinator and coordinator.cancellation.is_cancelled: + await exit_for_cancellation() return iteration += 1 @@ -3809,6 +3829,7 @@ async def _execute_stream( # Check if provider supports streaming if hasattr(provider, "stream"): # Use streaming if available + self._llm_calls_this_turn += 1 async for chunk in self._stream_from_provider( provider, chat_request, @@ -3835,13 +3856,16 @@ async def _execute_stream( if await self._has_pending_tools(context): # Process tools await self._process_tools(context, tools, hooks) + continuation_needed = True continue else: # Last-drain edge: if a steer arrived during the final generation, # loop once more so the model acts on it this turn. The top-of- # iteration drain performs the actual injection. if not self._steering_queue.is_empty: + continuation_needed = True continue + natural_completion = True break else: # Fallback to non-streaming @@ -3850,6 +3874,7 @@ async def _execute_stream( if self.extended_thinking: kwargs["extended_thinking"] = True try: + self._llm_calls_this_turn += 1 response = await provider.complete(chat_request, **kwargs) except LLMError as e: await hooks.emit( @@ -4010,9 +4035,16 @@ async def _execute_stream( # loop once more so the model acts on it this turn. The top-of- # iteration drain performs the actual injection. if not self._steering_queue.is_empty: + continuation_needed = True continue + natural_completion = True break + # A tool response needs another model turn to consume its + # corresponding results. This remains true even if the bounded + # loop cannot enter that next iteration. + continuation_needed = True + # Add assistant message with tool calls # Store structured content blocks (preserves reasoning state, thinking blocks, etc.) # Extract text for display/logging only @@ -4265,8 +4297,21 @@ async def _execute_stream( } ) - # Check if we exceeded max iterations (only if not unlimited) - if self.max_iterations != -1 and iteration >= self.max_iterations: + # Add exactly one finalization call only when the bounded budget + # prevented a continuation. A normal no-tool break at the cap is a + # natural completion and must return without a duplicate provider call. + if ( + self.max_iterations != -1 + and iteration >= self.max_iterations + and continuation_needed + and not natural_completion + ): + if coordinator and coordinator.cancellation.is_cancelled: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return # Layer 1 call-budget bookkeeping (spec: 298-replacement). Read by # _execute_one_turn's status precedence and the payload's # metadata bag -- this is the ONLY place it is set to True; it is @@ -4275,7 +4320,7 @@ async def _execute_stream( logger.warning(f"Max iterations ({self.max_iterations}) reached") # Inject system reminder to agent before returning - await hooks.emit( + finalization_result = await hooks.emit( PROVIDER_REQUEST, { "provider": provider_name, @@ -4283,10 +4328,45 @@ async def _execute_stream( "max_reached": True, }, ) + if coordinator: + finalization_result = await coordinator.process_hook_result( + finalization_result, "provider:request", "orchestrator" + ) + if coordinator.cancellation.is_cancelled: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return + if finalization_result.action == "deny": + denial_text = f"Operation denied: {finalization_result.reason}" + await close_finalization_tool_turn(denial_text) + yield (denial_text, iteration) + return + + if coordinator and coordinator.cancellation.is_cancelled: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return + + # A steer that required this finalization must be part of its one + # allowed provider call, not discarded at the next turn boundary. + await self._drain_steering(context, hooks, iteration) # Get one final response with the reminder (via _execute_stream helper) message_dicts = await context.get_messages_for_request(provider=provider) message_dicts = list(message_dicts) + # The finalization hook and context assembly both await. Check + # again before contacting the provider so a concurrent + # cancellation cannot buy an unrequested final provider call. + if coordinator and coordinator.cancellation.is_cancelled: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + await exit_for_cancellation() + return # Enveloped (tail variant) and metadata-stamped. This message is # VIEW-ONLY -- appended to message_dicts, never persisted via # context.add_message -- so it carries no "persisted" key: under @@ -4315,18 +4395,18 @@ async def _execute_stream( # Convert dicts to ChatRequest messages_objects = [Message(**msg) for msg in message_dicts] - # Deliberately tool-less (spec: 298-replacement D3 -- "one - # final tool-less wrap-up LLM call"). Offering tools here - # would let the model emit a tool_calls-only response with - # no text; this code path never calls parse_tool_calls() or - # processes a tool_call, so that response's content would be - # empty and the wrap-up would silently produce no summary -- - # defeating Layer 1's whole point (success criterion: "a - # non-empty agent-authored response"). Forcing tools=None - # guarantees the model must answer in text. + # Preserve the normal declarations, including provider-native + # specifications, so any assistant tool call and paired tool + # result in the existing transcript stay valid. The portable + # choice prevents new calls; this finalization path never + # parses or dispatches a tool response. + tools_list = ( + [_build_tool_spec(tool) for tool in tools.values()] if tools else None + ) max_iter_chat_request = ChatRequest( messages=messages_objects, - tools=None, + tools=tools_list, + tool_choice="none", reasoning_effort=self.config.get("reasoning_effort"), ) @@ -4334,19 +4414,82 @@ async def _execute_stream( if self.extended_thinking: kwargs["extended_thinking"] = True + self._llm_calls_this_turn += 1 response = await provider.complete(max_iter_chat_request, **kwargs) - content = ( - response.content if hasattr(response, "content") else str(response) - ) + response_text = getattr(response, "text", None) + if not isinstance(response_text, str) or not response_text: + response_text = self._extract_text_from_content( + getattr(response, "content", None) + ) - if content: - # Yield the final response - async for token in self._tokenize_stream(content): + if response_text: + async for token in self._tokenize_stream(response_text): yield (token, iteration) - # Add to context - await context.add_message({"role": "assistant", "content": content}) + # Do not persist opaque provider state for a response that + # requested tools the capped loop will not execute: it + # could require unpaired tool results on replay. Otherwise + # retain the normal no-tool response shape, but only the + # safe text/thinking blocks. + response_compliant = not provider.parse_tool_calls(response) + response_content = getattr(response, "content", None) + if isinstance(response_content, list): + content_dicts = [ + block.model_dump() + if hasattr(block, "model_dump") + else block + for block in response_content + ] + safe_content = [ + block_dict + for block_dict in content_dicts + if isinstance(block_dict, dict) + and getattr( + block_dict.get("type"), "value", block_dict.get("type") + ) + in ("text", "thinking") + ] + response_compliant = response_compliant and len( + safe_content + ) == len(content_dicts) + if not response_compliant: + safe_content = [] + else: + safe_content = [] + if safe_content: + assistant_msg = { + "role": "assistant", + "content": safe_content, + } + for block_dict in safe_content: + if ( + getattr( + block_dict["type"], "value", block_dict["type"] + ) + == "thinking" + ): + assistant_msg["thinking_block"] = block_dict + break + else: + assistant_msg = { + "role": "assistant", + "content": response_text, + } + + if response_compliant and getattr(response, "metadata", None): + assistant_msg["metadata"] = response.metadata + await context.add_message(assistant_msg) + else: + await close_finalization_tool_turn( + "The final response could not be generated." + ) + + except asyncio.CancelledError: + await close_finalization_tool_turn( + "The previous operation was cancelled. Results from completed tools have been preserved." + ) + raise except LLMError as e: await hooks.emit( PROVIDER_ERROR, @@ -4358,6 +4501,9 @@ async def _execute_stream( }, ) logger.error(f"Error getting final response after max iterations: {e}") + await close_finalization_tool_turn( + "The final response could not be generated." + ) except Exception as e: await hooks.emit( PROVIDER_ERROR, @@ -4367,6 +4513,9 @@ async def _execute_stream( }, ) logger.error(f"Error getting final response after max iterations: {e}") + await close_finalization_tool_turn( + "The final response could not be generated." + ) # Emit execution end await hooks.emit( @@ -4479,17 +4628,21 @@ def _extract_text_from_content(self, content) -> str: # message_models.ThinkingBlock which uses .thinking). The hasattr-based # filter was letting thinking text leak into the response string, which # pollutes parse_json extraction in downstream recipe steps. + blocks = [content] if isinstance(content, dict) else content text_parts = [] - for block in content: + for block in blocks: # Explicit type check โ€” works for both enum (ContentBlockType.TEXT) # and plain-str "type" fields (e.g., message_models.TextBlock). - block_type = getattr(block, "type", None) + block_type = ( + block.get("type") if isinstance(block, dict) else getattr(block, "type", None) + ) # Handle both enum (block_type.value == "text") and raw str ("text") type_value = ( getattr(block_type, "value", block_type) if block_type else None ) - if type_value == "text" and hasattr(block, "text"): - text_parts.append(block.text) + text = block.get("text") if isinstance(block, dict) else getattr(block, "text", None) + if type_value == "text" and isinstance(text, str): + text_parts.append(text) # Thinking blocks, tool_use blocks, etc. are all correctly excluded. return "\n\n".join(text_parts) diff --git a/tests/test_call_budget.py b/tests/test_call_budget.py index ca9651b..e6051fe 100644 --- a/tests/test_call_budget.py +++ b/tests/test_call_budget.py @@ -21,6 +21,7 @@ from __future__ import annotations +import asyncio from typing import Any, ClassVar import pytest @@ -128,15 +129,70 @@ async def execute(self, arguments): return ToolResult(success=True, output="ok") +class NativeishTool(MockTool): + """A generic stand-in for a provider-native tool declaration.""" + + name = "nativeish_tool" + native_tool_spec: ClassVar[dict] = { + "type": "nativeish", + "display_width": 1024, + } + + +class CountingNativeishTool(NativeishTool): + def __init__(self) -> None: + self.executions = 0 + + async def execute(self, arguments): + self.executions += 1 + return ToolResult(success=True, output="ok") + + +class TypedTextBlock: + """Minimal typed text block, matching the ContentBlock shape.""" + + type = "text" + + def __init__(self, text: str) -> None: + self.text = text + + def model_dump(self) -> dict[str, str]: + return {"type": self.type, "text": self.text} + + +class TypedThinkingBlock: + """Minimal typed thinking block with the replay signature.""" + + type = "thinking" + + def __init__(self, thinking: str, signature: str) -> None: + self.thinking = thinking + self.signature = signature + + def model_dump(self) -> dict[str, str]: + return { + "type": self.type, + "thinking": self.thinking, + "signature": self.signature, + } + + class MockTurnResponse: """A plain conversational-turn response (non-streaming path).""" - def __init__(self, text: str = "", tool_calls: list | None = None) -> None: + def __init__( + self, + text: str = "", + tool_calls: list | None = None, + *, + content: Any | None = None, + metadata: dict[str, Any] | None = None, + ) -> None: self.text = text - self.content = text + self.content = text if content is None else content self.content_blocks = None self.usage = None - self.metadata = None + self.metadata = metadata self._intended_tool_calls = tool_calls or [] @@ -225,6 +281,74 @@ async def test_normal_success_when_budget_not_hit(self) -> None: assert provider.call_count == 3 +# --------------------------------------------------------------------------- +# Bounded natural completion -- no unnecessary wrap-up +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestBoundedNaturalCompletion: + async def test_cap_one_plain_text_is_one_successful_call(self) -> None: + """A natural final answer at the cap must not trigger a duplicate call.""" + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = [MockTurnResponse(text="complete answer")] + + result = await orch.execute( + prompt="answer once", + context=ctx, + providers={"main": provider}, + tools={}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "complete answer" + assert provider.call_count == 1 + event = hooks.orchestrator_complete_events()[-1] + assert event["status"] == "success" + assert event["metadata"]["budget_exhausted"] is False + assert event["metadata"]["llm_calls"] == 1 + + async def test_cap_two_native_tool_then_final_text_is_two_calls(self) -> None: + """A tool result followed by a natural answer at the cap needs no wrap-up.""" + orch = _make_orchestrator({"max_iterations": 2}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = [ + MockTurnResponse( + tool_calls=[MockToolCall(call_id="native-call", name="nativeish_tool")] + ), + MockTurnResponse(text="natural final answer"), + ] + + result = await orch.execute( + prompt="use the native tool", + context=ctx, + providers={"main": provider}, + tools={"nativeish_tool": NativeishTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "natural final answer" + assert provider.call_count == 2 + assert provider.requests[0].tools is not None + assert provider.requests[0].tools[0].model_dump()["type"] == "nativeish" + assert any( + message.get("role") == "tool" and message.get("name") == "nativeish_tool" + for message in ctx._messages + ) + event = hooks.orchestrator_complete_events()[-1] + assert event["status"] == "success" + assert event["metadata"]["budget_exhausted"] is False + + # --------------------------------------------------------------------------- # T1.2 / T1.3 / T1.4 / T1.5 -- Budget reached exactly # --------------------------------------------------------------------------- @@ -257,9 +381,9 @@ async def test_exactly_n_plus_one_provider_calls(self) -> None: async def test_wrapup_reminder_and_no_further_loop(self) -> None: """T1.3: the final request's last message carries the orchestrator-loop-limit system-reminder, and the wrap-up call never - triggers a further loop iteration (tool-less by construction -- see - the exhaustion branch, which forces tools=None on the wrap-up - ChatRequest).""" + triggers a further loop iteration. Tools remain declared so native + assistant/tool-result history stays valid, but tool_choice disables + further dispatch.""" orch = _make_orchestrator({"max_iterations": 2}) ctx = MockContext() hooks = MockHooks() @@ -286,13 +410,44 @@ async def test_wrapup_reminder_and_no_further_loop(self) -> None: last_message = wrapup_request.messages[-1] assert last_message.role == "user" assert 'source="orchestrator-loop-limit"' in last_message.content - # Tool-less: the wrap-up ChatRequest carries no tools regardless of - # what tools were available during the turn. - assert wrapup_request.tools is None + # Preserve the normal tool declarations, including native shapes, but + # force the portable no-tool choice so the final response cannot + # request work the capped loop will not execute. + assert wrapup_request.tools is not None + assert wrapup_request.tool_choice == "none" # Exactly one wrap-up call -- no further loop, even though the fake # wrap-up response included a tool call. assert provider.call_count == 3 # 2 budgeted + 1 wrap-up, no more + async def test_cap_one_tool_call_wrapup_keeps_native_tools_but_disables_them(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = [ + MockTurnResponse( + tool_calls=[MockToolCall(call_id="native-call", name="nativeish_tool")] + ) + ] + provider.wrapup_response = MockTurnResponse(text="wrapped up") + + result = await orch.execute( + prompt="use one tool", + context=ctx, + providers={"main": provider}, + tools={"nativeish_tool": NativeishTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "wrapped up" + assert provider.call_count == 2 + wrapup_request = provider.requests[-1] + assert wrapup_request.tools is not None + assert wrapup_request.tools[0].model_dump()["type"] == "nativeish" + assert wrapup_request.tool_choice == "none" + async def test_status_is_budget_exhausted(self) -> None: """T1.4.""" orch = _make_orchestrator({"max_iterations": 2}) @@ -340,7 +495,9 @@ async def test_metadata_shape_is_exact(self) -> None: "budget_exhausted", "resumable", } - assert metadata["llm_calls"] == 2 + # The forced finalization is an actual third provider call, even + # though only two calls were within the bounded loop iteration budget. + assert metadata["llm_calls"] == 3 assert metadata["llm_call_budget"] == 2 assert metadata["budget_exhausted"] is True assert metadata["resumable"] is True @@ -381,6 +538,212 @@ async def test_wrapup_message_appended_to_context(self) -> None: assert len(assistant_messages) >= 1 +# --------------------------------------------------------------------------- +# Wrap-up response handling +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestWrapupResponseHandling: + @pytest.mark.parametrize( + ("content", "expected"), + [ + ("string summary", "string summary"), + ([TypedTextBlock("typed summary")], "typed summary"), + ([{"type": "text", "text": "dict summary"}], "dict summary"), + ], + ids=["string", "typed-block", "dict-block"], + ) + async def test_final_response_renders_normalized_text_and_persists_safe_blocks( + self, content: Any, expected: str + ) -> None: + """Finalization renders text while retaining safe structured content.""" + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + provider.wrapup_response = MockTurnResponse(content=content) + + result = await orch.execute( + prompt="force a summary", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == expected + final_message = [m for m in ctx._messages if m.get("role") == "assistant"][-1] + assert final_message == { + "role": "assistant", + "content": ( + [ + block.model_dump() if hasattr(block, "model_dump") else block + for block in content + ] + if isinstance(content, list) + else expected + ), + } + + async def test_final_response_preserves_safe_structured_content_and_metadata( + self, + ) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + thinking = TypedThinkingBlock("private reasoning", "replay-signature") + text = TypedTextBlock("structured final text") + provider.wrapup_response = MockTurnResponse( + content=[thinking, text], + metadata={"provider_state": "opaque"}, + ) + + result = await orch.execute( + prompt="force a structured summary", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "structured final text" + final_message = [m for m in ctx._messages if m.get("role") == "assistant"][-1] + assert final_message == { + "role": "assistant", + "content": [thinking.model_dump(), text.model_dump()], + "thinking_block": thinking.model_dump(), + "metadata": {"provider_state": "opaque"}, + } + + async def test_final_tool_calls_are_never_dispatched_or_structurally_persisted( + self, + ) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + tool = CountingNativeishTool() + provider.turn_queue = [ + MockTurnResponse( + tool_calls=[MockToolCall(call_id="budget-tool", name="nativeish_tool")] + ) + ] + provider.wrapup_response = MockTurnResponse( + text="final text only", + tool_calls=[MockToolCall(call_id="unsupported-final", name="nativeish_tool")], + content=[ + TypedThinkingBlock("must not persist", "invalid-signature"), + TypedTextBlock("final text only"), + {"type": "tool_call", "name": "nativeish_tool"}, + ], + metadata={"must_not_persist": True}, + ) + + result = await orch.execute( + prompt="run then summarize", + context=ctx, + providers={"main": provider}, + tools={"nativeish_tool": tool}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "final text only" + assert provider.call_count == 2 + assert tool.executions == 1 + final_message = [m for m in ctx._messages if m.get("role") == "assistant"][-1] + assert final_message == {"role": "assistant", "content": "final text only"} + + +# --------------------------------------------------------------------------- +# Bounded streaming and steering finalization +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestBoundedStreamingAndSteering: + async def test_streaming_natural_completion_at_cap_is_successful(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + + class StreamingProvider: + calls = 0 + + async def stream(self, chat_request, tools=None): # noqa: ANN001,ANN201 + self.calls += 1 + yield {"content": "streamed answer"} + + async def complete(self, chat_request, **kwargs): # noqa: ANN001,ANN201 + pytest.fail("natural streaming completion must not wrap up") + + provider = StreamingProvider() + result = await orch.execute( + prompt="stream once", + context=ctx, + providers={"main": provider}, + tools={}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "streamed answer" + assert provider.calls == 1 + event = hooks.orchestrator_complete_events()[-1] + assert event["status"] == "success" + assert event["metadata"]["budget_exhausted"] is False + + async def test_pending_steer_at_cap_requires_finalization(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + + class SteeringProvider(FakeProvider): + async def complete(self, chat_request, **kwargs): # noqa: ANN001,ANN201 + response = await super().complete(chat_request, **kwargs) + if self.call_count == 1: + orch.steer("respond to this before ending") + return response + + provider = SteeringProvider() + provider.turn_queue = [MockTurnResponse(text="first answer")] + provider.wrapup_response = MockTurnResponse(text="forced final answer") + + result = await orch.execute( + prompt="answer", + context=ctx, + providers={"main": provider}, + tools={}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "first answerforced final answer" + assert provider.call_count == 2 + user_contents = [ + message.content + for message in provider.requests[-1].messages + if message.role == "user" + ] + assert "respond to this before ending" in user_contents + assert any("maximum number of iterations" in content for content in user_contents) + assert orch._steering_queue.is_empty + event = hooks.orchestrator_complete_events()[-1] + assert event["status"] == "budget_exhausted" + assert event["metadata"]["budget_exhausted"] is True + + # --------------------------------------------------------------------------- # T1.7 / T1.8 -- 80% warning # --------------------------------------------------------------------------- @@ -574,6 +937,206 @@ async def fake_execute_stream(*args, **kwargs): assert events[-1]["status"] == "error" +# --------------------------------------------------------------------------- +# Cancellation while finalization is pending +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestFinalizationCancellation: + async def test_cancellation_before_first_call_never_calls_provider(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + coordinator.cancellation.is_cancelled = True + provider = FakeProvider() + + await orch.execute( + prompt="cancel now", + context=ctx, + providers={"main": provider}, + tools={}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert provider.call_count == 0 + assert hooks.orchestrator_complete_events()[-1]["status"] == "cancelled" + + async def test_cancellation_during_finalization_skips_wrapup_provider_call(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + coordinator = MockCoordinator() + + class HooksThatCancelAtFinalization(MockHooks): + async def emit( + self, event_name: str, payload: dict | None = None + ) -> MockHookResult: + result = await super().emit(event_name, payload) + if event_name == "provider:request" and (payload or {}).get("max_reached"): + coordinator.cancellation.is_cancelled = True + return result + + hooks = HooksThatCancelAtFinalization() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + + await orch.execute( + prompt="cancel during finalization", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert provider.call_count == 1 + assert hooks.orchestrator_complete_events()[-1]["status"] == "cancelled" + assert ctx._messages[-1] == { + "role": "assistant", + "content": "The previous operation was cancelled. Results from completed tools have been preserved.", + } + + async def test_finalization_denial_closes_a_tool_result_with_yielded_text(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + + class HooksThatDenyAtFinalization(MockHooks): + def __init__(self) -> None: + super().__init__() + self.finalizing = False + + async def emit( + self, event_name: str, payload: dict | None = None + ) -> MockHookResult: + self.finalizing = event_name == "provider:request" and bool( + (payload or {}).get("max_reached") + ) + return await super().emit(event_name, payload) + + class DenyingCoordinator(MockCoordinator): + async def process_hook_result(self, result, *args, **kwargs): + if hooks.finalizing: + return type("DenyResult", (), {"action": "deny", "reason": "stop"})() + return result + + hooks = HooksThatDenyAtFinalization() + coordinator = DenyingCoordinator() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + + result = await orch.execute( + prompt="deny finalization", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "Operation denied: stop" + assert provider.call_count == 1 + assert ctx._messages[-1] == { + "role": "assistant", + "content": "Operation denied: stop", + } + + async def test_cancelled_finalization_provider_closes_tool_turn_then_propagates( + self, + ) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + provider.wrapup_should_raise = asyncio.CancelledError("cancelled finalization") + + with pytest.raises(asyncio.CancelledError, match="cancelled finalization"): + await orch.execute( + prompt="cancel finalization provider", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert ctx._messages[-1] == { + "role": "assistant", + "content": "The previous operation was cancelled. Results from completed tools have been preserved.", + } + + async def test_cancellation_wins_over_finalization_hook_denial(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + + class CountingCancellation(MockCancellation): + def __init__(self) -> None: + self.callback_count = 0 + + async def trigger_callbacks(self) -> None: + self.callback_count += 1 + + class HooksThatDenyAtFinalization(MockHooks): + def __init__(self) -> None: + super().__init__() + self.finalizing = False + + async def emit( + self, event_name: str, payload: dict | None = None + ) -> MockHookResult: + self.finalizing = event_name == "provider:request" and bool( + (payload or {}).get("max_reached") + ) + return await super().emit(event_name, payload) + + class CancellingDenyCoordinator(MockCoordinator): + def __init__(self, hooks: HooksThatDenyAtFinalization) -> None: + super().__init__() + self.cancellation = CountingCancellation() + self.hooks = hooks + + async def process_hook_result(self, result, *args, **kwargs): + if not self.hooks.finalizing: + return result + self.cancellation.is_cancelled = True + return type("DenyResult", (), {"action": "deny", "reason": "stop"})() + + class SteeringProvider(FakeProvider): + async def complete(self, chat_request, **kwargs): + response = await super().complete(chat_request, **kwargs) + if self.call_count == 1: + orch.steer("discard this steer after cancellation") + return response + + hooks = HooksThatDenyAtFinalization() + coordinator = CancellingDenyCoordinator(hooks) + provider = SteeringProvider() + provider.turn_queue = _looping_turns(1) + + await orch.execute( + prompt="cancel during denied finalization", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert provider.call_count == 1 + assert coordinator.cancellation.callback_count == 1 + assert hooks.events("cancel:requested") + assert hooks.events("cancel:completed") + assert orch._steering_queue.is_empty + assert hooks.orchestrator_complete_events()[-1]["status"] == "cancelled" + assert ctx._messages[-1] == { + "role": "assistant", + "content": "The previous operation was cancelled. Results from completed tools have been preserved.", + } + + # --------------------------------------------------------------------------- # T1.12 -- Wrap-up LLM failure # --------------------------------------------------------------------------- @@ -582,12 +1145,12 @@ async def fake_execute_stream(*args, **kwargs): @pytest.mark.asyncio class TestWrapupFailure: async def test_wrapup_provider_error_does_not_crash(self) -> None: - orch = _make_orchestrator({"max_iterations": 2}) + orch = _make_orchestrator({"max_iterations": 1}) ctx = MockContext() hooks = MockHooks() coordinator = MockCoordinator() provider = FakeProvider() - provider.turn_queue = _looping_turns(2) + provider.turn_queue = _looping_turns(1) provider.wrapup_should_raise = RuntimeError("provider exploded") result = await orch.execute( @@ -610,3 +1173,73 @@ async def test_wrapup_provider_error_does_not_crash(self) -> None: provider_errors = hooks.events("provider:error") assert len(provider_errors) == 1 assert provider_errors[0]["error"]["type"] == "RuntimeError" + assert ctx._messages[-1] == { + "role": "assistant", + "content": "The final response could not be generated.", + } + + async def test_empty_final_response_closes_a_tool_result_turn(self) -> None: + orch = _make_orchestrator({"max_iterations": 1}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = _looping_turns(1) + provider.wrapup_response = MockTurnResponse() + + result = await orch.execute( + prompt="return an empty finalization", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + assert result == "" + assert ctx._messages[-1] == { + "role": "assistant", + "content": "The final response could not be generated.", + } + + +@pytest.mark.asyncio +class TestCallCountsOutsideForcedFinalization: + @pytest.mark.parametrize( + ("max_iterations", "responses"), + [ + (1, [MockTurnResponse(text="natural cap completion")]), + ( + -1, + [ + MockTurnResponse( + tool_calls=[MockToolCall(call_id="unlimited-tool")] + ), + MockTurnResponse(text="unlimited completion"), + ], + ), + ], + ids=["natural-at-cap", "unlimited"], + ) + async def test_normal_paths_count_only_actual_provider_calls( + self, max_iterations: int, responses: list[MockTurnResponse] + ) -> None: + orch = _make_orchestrator({"max_iterations": max_iterations}) + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator() + provider = FakeProvider() + provider.turn_queue = responses + + await orch.execute( + prompt="complete normally", + context=ctx, + providers={"main": provider}, + tools={"mock_tool": MockTool()}, + hooks=hooks, # type: ignore[arg-type] + coordinator=coordinator, # type: ignore[arg-type] + ) + + metadata = hooks.orchestrator_complete_events()[-1]["metadata"] + assert metadata["llm_calls"] == provider.call_count + assert metadata["budget_exhausted"] is False