diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 38f7e2e..7da87dc 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -670,7 +670,7 @@ class StreamingOrchestrator: "(e.g. a soak period, waiting for an external event) that cannot " "pass within this session no matter what the assistant does.\n" "STRUCTURE-LOCKED -- the condition applies a universal " - 'requirement over a set that contains a member which is ' + "requirement over a set that contains a member which is " 'structurally exempt or unreachable (e.g. "all N sites" when one ' "site cannot produce the required measurement).\n" "HISTORY-LOCKED -- the condition constrains the transcript's own " @@ -948,9 +948,7 @@ def __init__(self, config: dict[str, Any]): # the full history); the summary call only needs the CURRENT # state, which the tail of a long run already establishes. `<= 0` # disables the cap (send the whole list, pre-existing behavior). - self.goal_summary_max_reasons = int( - config.get("goal_summary_max_reasons", 20) - ) + self.goal_summary_max_reasons = int(config.get("goal_summary_max_reasons", 20)) # Per-execute()-call cache for `_resolve_goal_model`'s result (see # that method's CRITICAL PERF note) -- reset to None at the top of # execute() so each run re-resolves once, not on every turn. @@ -994,6 +992,10 @@ def __init__(self, config: dict[str, Any]): # callers of the tool-execution methods never hit an # AttributeError. self._tool_calls_this_turn: int = 0 + # Identical-failure counter for the circuit breaker. Keyed on + # (tool, arguments, error) so only a call that fails the SAME way + # counts -- see `_apply_failure_breaker`. + self._repeated_failures: dict[str, int] = {} # Store ephemeral injections from tool:post hooks for next iteration self._pending_ephemeral_injections: list[dict[str, Any]] = [] # Track whether cancel:requested has been emitted for the current execution @@ -1289,14 +1291,16 @@ async def execute( # re-running a test after a fix) can look repetitive # too. See TestDualConditionStallTrip. try: - is_stalled, stall_detail, stall_verdict = ( - await self._judge_stall( - goal, - providers, - hooks, - coordinator, - trigger=stall_trigger, - ) + ( + is_stalled, + stall_detail, + stall_verdict, + ) = await self._judge_stall( + goal, + providers, + hooks, + coordinator, + trigger=stall_trigger, ) except Exception as e: # Fail open: a flaky judge call must never itself @@ -1633,8 +1637,7 @@ def _goal_stall_escalation_prompt( explanation = self._GOAL_STALL_VERDICT_EXPLANATIONS.get(verdict or "") if explanation: verdict_clause = ( - f" A reviewing judge classified this as {verdict}: " - f"{explanation}." + f" A reviewing judge classified this as {verdict}: {explanation}." ) return ( @@ -2090,9 +2093,7 @@ async def _judge_stall( ) else: recent_reasons = ( - goal["reasons"][-goal["no_tool_turns"] :] - if goal.get("reasons") - else [] + goal["reasons"][-goal["no_tool_turns"] :] if goal.get("reasons") else [] ) system_prompt = self._GOAL_STALL_SYSTEM_PROMPT_IDLE activity_clause = "the assistant took no tool actions at all" @@ -2282,9 +2283,7 @@ def _goal_summary_fallback(goal: dict[str, Any], final_state: str) -> str | None return "evaluator failed" return None - def _cap_reasons_for_summary( - self, reasons: list[str] - ) -> tuple[list[str], int]: + def _cap_reasons_for_summary(self, reasons: list[str]) -> tuple[list[str], int]: """Bound how many of ``goal["reasons"]`` are shipped to the summary model (see ``_summarize_goal_run``). @@ -2386,9 +2385,7 @@ async def _summarize_goal_run( else: reasons = goal.get("reasons", []) kept_reasons, omitted_count = self._cap_reasons_for_summary(reasons) - reasons_lines = [ - f"{i + 1}. {r}" for i, r in enumerate(kept_reasons) - ] + reasons_lines = [f"{i + 1}. {r}" for i, r in enumerate(kept_reasons)] if omitted_count: reasons_lines.insert( 0, f"(earliest {omitted_count} reasons omitted)" @@ -2708,137 +2705,158 @@ async def _execute_stream( # Emit execution start await hooks.emit("execution:start", {"prompt": prompt}) + try: + # Reset rate limit tracking for new session + self._last_provider_call_end = None - # Reset rate limit tracking for new session - self._last_provider_call_end = None + # Add user message + await context.add_message({"role": "user", "content": prompt}) - # Add user message - await context.add_message({"role": "user", "content": prompt}) + # Select provider + provider = self._select_provider(providers) + if not provider: + yield ("Error: No providers available", 0) + return - # Select provider - provider = self._select_provider(providers) - if not provider: - yield ("Error: No providers available", 0) - return + # Find provider name for event emission + provider_name = None + for name, prov in providers.items(): + if prov is provider: + provider_name = name + break - # Find provider name for event emission - provider_name = None - for name, prov in providers.items(): - if prov is provider: - provider_name = name - break - - # Pure observability. `basis` names WHY this provider won: - # "pinned" when the conversation-scope pin decided it (capability - # `conversation.provider_pin`), else "priority" -- the unpinned - # path, unchanged. Reading the pin here is sound because - # `_select_provider` above honors it whenever set and RAISES if a - # pin no longer resolves, so reaching this line with a pin set - # means the pin is what selected `provider`. - # - # The main conversation loop never sets an explicit model override - # (see the ChatRequest built below), so the model that will - # ACTUALLY be used is the provider's own default -- read locally - # through the kernel's Provider contract - # (`get_info().defaults["model"]`, no I/O), never via a network - # call and never via a vendor-specific attribute. See - # `_provider_default_model`; None there means "the provider could - # not tell us", not a guess. - await hooks.emit( - PROVIDER_RESOLVE, - { - "provider": provider_name, - "model": self._provider_default_model(provider), - "basis": ( - "pinned" if self._pinned_provider_name is not None else "priority" - ), - "scope": "conversation", - }, - ) + # Pure observability. `basis` names WHY this provider won: + # "pinned" when the conversation-scope pin decided it (capability + # `conversation.provider_pin`), else "priority" -- the unpinned + # path, unchanged. Reading the pin here is sound because + # `_select_provider` above honors it whenever set and RAISES if a + # pin no longer resolves, so reaching this line with a pin set + # means the pin is what selected `provider`. + # + # The main conversation loop never sets an explicit model override + # (see the ChatRequest built below), so the model that will + # ACTUALLY be used is the provider's own default -- read locally + # through the kernel's Provider contract + # (`get_info().defaults["model"]`, no I/O), never via a network + # call and never via a vendor-specific attribute. See + # `_provider_default_model`; None there means "the provider could + # not tell us", not a guess. + await hooks.emit( + PROVIDER_RESOLVE, + { + "provider": provider_name, + "model": self._provider_default_model(provider), + "basis": ( + "pinned" + if self._pinned_provider_name is not None + else "priority" + ), + "scope": "conversation", + }, + ) - iteration = 0 + 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: - # Emit cancel:requested on first detection and trigger cleanup callbacks - if not self._cancel_requested_emitted: - self._cancel_requested_emitted = True + while self.max_iterations == -1 or iteration < self.max_iterations: + # Check for cancellation at iteration start + if coordinator and coordinator.cancellation.is_cancelled: + # Emit cancel:requested on first detection and trigger cleanup callbacks + if not self._cancel_requested_emitted: + self._cancel_requested_emitted = True + await hooks.emit( + CANCEL_REQUESTED, + { + "orchestrator": "loop-streaming", + "state": str(coordinator.cancellation.state), + "turn_count": iteration, + }, + ) + try: + await coordinator.cancellation.trigger_callbacks() + except Exception as e: + logger.warning(f"Error in cancellation callbacks: {e}") + # Emit cancel:completed — orchestrator is exiting due to cancellation await hooks.emit( - CANCEL_REQUESTED, + CANCEL_COMPLETED, { "orchestrator": "loop-streaming", - "state": str(coordinator.cancellation.state), + "was_immediate": coordinator.cancellation.is_immediate, "turn_count": iteration, }, ) - try: - await coordinator.cancellation.trigger_callbacks() - except Exception as e: - logger.warning(f"Error in cancellation callbacks: {e}") - # Emit cancel:completed — orchestrator is exiting due to cancellation - await hooks.emit( - CANCEL_COMPLETED, - { - "orchestrator": "loop-streaming", - "was_immediate": coordinator.cancellation.is_immediate, - "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() - return + # 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() + return - iteration += 1 + iteration += 1 - # Mid-turn steering: drain queued user messages BEFORE building the request, - # so they are part of this iteration's provider call. At iteration 1 this is - # "before the first LLM call"; at iteration N>1 this is "after the prior tool - # round, before the next provider call" — the single natural boundary. - await self._drain_steering(context, hooks, iteration) + # Mid-turn steering: drain queued user messages BEFORE building the request, + # so they are part of this iteration's provider call. At iteration 1 this is + # "before the first LLM call"; at iteration N>1 this is "after the prior tool + # round, before the next provider call" — the single natural boundary. + await self._drain_steering(context, hooks, iteration) - # Emit provider request BEFORE getting messages (allows hook injections) - result = await hooks.emit( - PROVIDER_REQUEST, {"provider": provider_name, "iteration": iteration} - ) - if coordinator: - result = await coordinator.process_hook_result( - result, "provider:request", "orchestrator" + # Emit provider request BEFORE getting messages (allows hook injections) + result = await hooks.emit( + PROVIDER_REQUEST, + {"provider": provider_name, "iteration": iteration}, ) - if result.action == "deny": - yield (f"Operation denied: {result.reason}", iteration) - return + if coordinator: + result = await coordinator.process_hook_result( + result, "provider:request", "orchestrator" + ) + if result.action == "deny": + yield (f"Operation denied: {result.reason}", iteration) + return - # Get messages for LLM request (context handles compaction internally) - # Pass provider for dynamic budget calculation based on model's context window - message_dicts = await context.get_messages_for_request(provider=provider) - message_dicts = list(message_dicts) # Convert to list for modification + # Get messages for LLM request (context handles compaction internally) + # Pass provider for dynamic budget calculation based on model's context window + message_dicts = await context.get_messages_for_request( + provider=provider + ) + message_dicts = list(message_dicts) # Convert to list for modification - # Append ephemeral injection if present (temporary, not stored) - if ( - result.action == "inject_context" - and result.ephemeral - and result.context_injection - ): - # Check if we should append to last tool result - if result.append_to_last_tool_result and len(message_dicts) > 0: - last_msg = message_dicts[-1] - # Append to last message if it's a tool result - if last_msg.get("role") == "tool": - # Append to existing content - original_content = last_msg.get("content", "") - message_dicts[-1] = { - **last_msg, - "content": f"{original_content}\n\n{result.context_injection}", - } - logger.debug( - "Appended ephemeral injection to last tool result message" - ) + # Append ephemeral injection if present (temporary, not stored) + if ( + result.action == "inject_context" + and result.ephemeral + and result.context_injection + ): + # Check if we should append to last tool result + if result.append_to_last_tool_result and len(message_dicts) > 0: + last_msg = message_dicts[-1] + # Append to last message if it's a tool result + if last_msg.get("role") == "tool": + # Append to existing content + original_content = last_msg.get("content", "") + message_dicts[-1] = { + **last_msg, + "content": f"{original_content}\n\n{result.context_injection}", + } + logger.debug( + "Appended ephemeral injection to last tool result message" + ) + else: + # Fall back to new message if last message isn't a tool result + # metadata.ephemeral marks this as regenerated-per-turn content + # so the provider never places a prompt-cache breakpoint on it + # (see amplifier_module_provider_anthropic._count_trailing_ephemeral_messages). + message_dicts.append( + { + "role": result.context_injection_role, + "content": result.context_injection, + "metadata": {"ephemeral": True}, + } + ) + logger.debug( + f"Last message role is '{last_msg.get('role')}', not 'tool' - " + "created new message for injection" + ) else: - # Fall back to new message if last message isn't a tool result + # Default behavior: append as new message # metadata.ephemeral marks this as regenerated-per-turn content # so the provider never places a prompt-cache breakpoint on it # (see amplifier_module_provider_anthropic._count_trailing_ephemeral_messages). @@ -2849,40 +2867,39 @@ async def _execute_stream( "metadata": {"ephemeral": True}, } ) - logger.debug( - f"Last message role is '{last_msg.get('role')}', not 'tool' - " - "created new message for injection" - ) - else: - # Default behavior: append as new message - # metadata.ephemeral marks this as regenerated-per-turn content - # so the provider never places a prompt-cache breakpoint on it - # (see amplifier_module_provider_anthropic._count_trailing_ephemeral_messages). - message_dicts.append( - { - "role": result.context_injection_role, - "content": result.context_injection, - "metadata": {"ephemeral": True}, - } - ) - # Apply pending ephemeral injections from tool:post hooks - if self._pending_ephemeral_injections: - for injection in self._pending_ephemeral_injections: - if ( - injection.get("append_to_last_tool_result") - and len(message_dicts) > 0 - ): - last_msg = message_dicts[-1] - if last_msg.get("role") == "tool": - original_content = last_msg.get("content", "") - message_dicts[-1] = { - **last_msg, - "content": f"{original_content}\n\n{injection['content']}", - } - logger.debug( - "Applied pending ephemeral injection to last tool result" - ) + # Apply pending ephemeral injections from tool:post hooks + if self._pending_ephemeral_injections: + for injection in self._pending_ephemeral_injections: + if ( + injection.get("append_to_last_tool_result") + and len(message_dicts) > 0 + ): + last_msg = message_dicts[-1] + if last_msg.get("role") == "tool": + original_content = last_msg.get("content", "") + message_dicts[-1] = { + **last_msg, + "content": f"{original_content}\n\n{injection['content']}", + } + logger.debug( + "Applied pending ephemeral injection to last tool result" + ) + else: + # metadata.ephemeral marks this as regenerated-per-turn + # content so the provider never places a prompt-cache + # breakpoint on it (see + # amplifier_module_provider_anthropic._count_trailing_ephemeral_messages). + message_dicts.append( + { + "role": injection["role"], + "content": injection["content"], + "metadata": {"ephemeral": True}, + } + ) + logger.debug( + "Last message not a tool result, created new message for injection" + ) else: # metadata.ephemeral marks this as regenerated-per-turn # content so the provider never places a prompt-cache @@ -2896,223 +2913,286 @@ async def _execute_stream( } ) logger.debug( - "Last message not a tool result, created new message for injection" + "Applied pending ephemeral injection as new message" ) - else: - # metadata.ephemeral marks this as regenerated-per-turn - # content so the provider never places a prompt-cache - # breakpoint on it (see - # amplifier_module_provider_anthropic._count_trailing_ephemeral_messages). - message_dicts.append( - { - "role": injection["role"], - "content": injection["content"], - "metadata": {"ephemeral": True}, - } - ) - logger.debug( - "Applied pending ephemeral injection as new message" - ) - # Clear pending injections after applying - self._pending_ephemeral_injections = [] + # Clear pending injections after applying + self._pending_ephemeral_injections = [] - # Convert dicts to ChatRequest for provider - messages_objects = [Message(**msg) for msg in message_dicts] + # Convert dicts to ChatRequest for provider + messages_objects = [Message(**msg) for msg in message_dicts] - # Convert tools to ToolSpec format for ChatRequest - tools_list = None - if tools: - tools_list = [_build_tool_spec(t) for t in tools.values()] + # Convert tools to ToolSpec format for ChatRequest + tools_list = None + if tools: + tools_list = [_build_tool_spec(t) for t in tools.values()] - chat_request = ChatRequest( - messages=messages_objects, - tools=tools_list, - reasoning_effort=self.config.get("reasoning_effort"), - ) - logger.info( - f"[ORCHESTRATOR] ChatRequest created with {len(tools_list) if tools_list else 0} tools" - ) - if tools_list: - logger.debug( - f"[ORCHESTRATOR] Tool names: {[t.name for t in tools_list]}" + chat_request = ChatRequest( + messages=messages_objects, + tools=tools_list, + reasoning_effort=self.config.get("reasoning_effort"), ) - - # Apply rate limit delay before provider call - await self._apply_rate_limit_delay(hooks, iteration) - - # Check if provider supports streaming - if hasattr(provider, "stream"): - # Use streaming if available - async for chunk in self._stream_from_provider( - provider, - chat_request, - context, - tools, - hooks, - coordinator, - provider_name=provider_name, - ): - # Check for immediate cancellation between chunks - if coordinator and coordinator.cancellation.is_immediate: - # Clear pending steers: immediate cancellation ends the turn, - # and any steer queued during streaming must not leak into a - # future turn — matching the other cancellation exits. (spec §5.2) - self._steering_queue.clear() - return - yield (chunk, iteration) - - # Update rate limit timestamp after streaming completes - self._last_provider_call_end = time.monotonic() - - # Check for tool calls after streaming - # This is simplified - real implementation would parse during stream - if await self._has_pending_tools(context): - # Process tools - await self._process_tools(context, tools, hooks) - 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: - continue - break - else: - # Fallback to non-streaming - # Build kwargs for provider - kwargs = {} - if self.extended_thinking: - kwargs["extended_thinking"] = True - try: - response = await provider.complete(chat_request, **kwargs) - except LLMError as e: - await hooks.emit( - PROVIDER_ERROR, - { - "provider": provider_name, - "error": {"type": type(e).__name__, "msg": str(e)}, - "retryable": e.retryable, - "status_code": e.status_code, - }, - ) - raise - except Exception as e: - await hooks.emit( - PROVIDER_ERROR, - { - "provider": provider_name, - "error": {"type": type(e).__name__, "msg": str(e)}, - }, + logger.info( + f"[ORCHESTRATOR] ChatRequest created with {len(tools_list) if tools_list else 0} tools" + ) + if tools_list: + logger.debug( + f"[ORCHESTRATOR] Tool names: {[t.name for t in tools_list]}" ) - raise - # Update rate limit timestamp after non-streaming response - self._last_provider_call_end = time.monotonic() - - # Emit content block events if present - content_blocks = getattr(response, "content_blocks", None) - if content_blocks: - total_blocks = len(content_blocks) - for idx, block in enumerate(content_blocks): - # Emit block start + # Apply rate limit delay before provider call + await self._apply_rate_limit_delay(hooks, iteration) + + # Check if provider supports streaming + if hasattr(provider, "stream"): + # Use streaming if available + async for chunk in self._stream_from_provider( + provider, + chat_request, + context, + tools, + hooks, + coordinator, + provider_name=provider_name, + ): + # Check for immediate cancellation between chunks + if coordinator and coordinator.cancellation.is_immediate: + # Clear pending steers: immediate cancellation ends the turn, + # and any steer queued during streaming must not leak into a + # future turn — matching the other cancellation exits. (spec §5.2) + self._steering_queue.clear() + return + yield (chunk, iteration) + + # Update rate limit timestamp after streaming completes + self._last_provider_call_end = time.monotonic() + + # Check for tool calls after streaming + # This is simplified - real implementation would parse during stream + if await self._has_pending_tools(context): + # Process tools + await self._process_tools(context, tools, hooks) + 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: + continue + break + else: + # Fallback to non-streaming + # Build kwargs for provider + kwargs = {} + if self.extended_thinking: + kwargs["extended_thinking"] = True + try: + response = await provider.complete(chat_request, **kwargs) + except LLMError as e: await hooks.emit( - CONTENT_BLOCK_START, + PROVIDER_ERROR, { - "block_type": block.type.value, - "block_index": idx, - "total_blocks": total_blocks, - "metadata": getattr(block, "raw", None), + "provider": provider_name, + "error": {"type": type(e).__name__, "msg": str(e)}, + "retryable": e.retryable, + "status_code": e.status_code, }, ) - - # Emit block end with complete block, usage, and total count - event_data = { - "block_index": idx, - "total_blocks": total_blocks, - "block": block.to_dict(), - } - if response.usage: - event_data["usage"] = response.usage.model_dump() - await hooks.emit(CONTENT_BLOCK_END, event_data) - elif response.content and isinstance(response.content, list): - # Fallback for providers that populate response.content - # (Pydantic ContentBlock models) but not content_blocks - # (raw SDK objects). Synthesize content_block events so - # downstream hooks (e.g. streaming-ui token usage) fire. - total_blocks = len(response.content) - for idx, block in enumerate(response.content): - block_dict = ( - block.model_dump() - if hasattr(block, "model_dump") - else block - ) - block_type = ( - block_dict.get("type", "text") - if isinstance(block_dict, dict) - else "text" - ) + raise + except Exception as e: await hooks.emit( - CONTENT_BLOCK_START, + PROVIDER_ERROR, { - "block_type": block_type, - "block_index": idx, - "total_blocks": total_blocks, + "provider": provider_name, + "error": {"type": type(e).__name__, "msg": str(e)}, }, ) - event_data = { - "block_index": idx, - "total_blocks": total_blocks, - "block": block_dict, - } - if response.usage: - event_data["usage"] = response.usage.model_dump() - await hooks.emit(CONTENT_BLOCK_END, event_data) + raise + + # Update rate limit timestamp after non-streaming response + self._last_provider_call_end = time.monotonic() + + # Emit content block events if present + content_blocks = getattr(response, "content_blocks", None) + if content_blocks: + total_blocks = len(content_blocks) + for idx, block in enumerate(content_blocks): + # Emit block start + await hooks.emit( + CONTENT_BLOCK_START, + { + "block_type": block.type.value, + "block_index": idx, + "total_blocks": total_blocks, + "metadata": getattr(block, "raw", None), + }, + ) - # Parse tool calls - tool_calls = provider.parse_tool_calls(response) + # Emit block end with complete block, usage, and total count + event_data = { + "block_index": idx, + "total_blocks": total_blocks, + "block": block.to_dict(), + } + if response.usage: + event_data["usage"] = response.usage.model_dump() + await hooks.emit(CONTENT_BLOCK_END, event_data) + elif response.content and isinstance(response.content, list): + # Fallback for providers that populate response.content + # (Pydantic ContentBlock models) but not content_blocks + # (raw SDK objects). Synthesize content_block events so + # downstream hooks (e.g. streaming-ui token usage) fire. + total_blocks = len(response.content) + for idx, block in enumerate(response.content): + block_dict = ( + block.model_dump() + if hasattr(block, "model_dump") + else block + ) + block_type = ( + block_dict.get("type", "text") + if isinstance(block_dict, dict) + else "text" + ) + await hooks.emit( + CONTENT_BLOCK_START, + { + "block_type": block_type, + "block_index": idx, + "total_blocks": total_blocks, + }, + ) + event_data = { + "block_index": idx, + "total_blocks": total_blocks, + "block": block_dict, + } + if response.usage: + event_data["usage"] = response.usage.model_dump() + await hooks.emit(CONTENT_BLOCK_END, event_data) + + # Parse tool calls + tool_calls = provider.parse_tool_calls(response) + + if not tool_calls: + # Extract text content from response for streaming + # Use .text field if available (e.g., OpenAI provider), otherwise extract from content blocks + if hasattr(response, "text") and response.text: + response_text = response.text + else: + response_text = self._extract_text_from_content( + response.content + ) + + # Stream the final response token by token + async for token in self._tokenize_stream(response_text): + yield (token, iteration) + + # Store structured content from response.content (our Pydantic models) + # This preserves reasoning state, thinking blocks, etc. + # response.content = list of our ContentBlock models (TextBlock, ThinkingBlock, etc.) + # response.content_blocks = raw SDK objects (for streaming events only) + response_content = getattr(response, "content", None) + if response_content and isinstance(response_content, list): + # Convert ContentBlock objects to dicts for serialization + content_dicts = [ + block.model_dump() + if hasattr(block, "model_dump") + else block + for block in response_content + ] + logger.info( + f"[ORCHESTRATOR] Storing {len(content_dicts)} content blocks" + ) + for i, block_dict in enumerate(content_dicts): + logger.info( + f"[ORCHESTRATOR] Block {i}: type={block_dict.get('type')}, has_content={'content' in block_dict}" + ) + assistant_msg = { + "role": "assistant", + "content": content_dicts, + } + else: + assistant_msg = { + "role": "assistant", + "content": response_text, + } - if not tool_calls: - # Extract text content from response for streaming - # Use .text field if available (e.g., OpenAI provider), otherwise extract from content blocks + # Preserve thinking blocks for Anthropic extended thinking (backward compat) + # Use response_content (our Pydantic models) not content_blocks (raw SDK objects) + if response_content and isinstance(response_content, list): + for block in response_content: + block_type = getattr(block, "type", None) + type_value = ( + getattr(block_type, "value", block_type) + if block_type + else None + ) + if type_value == "thinking": + # Store the thinking block as dict to preserve signature + assistant_msg["thinking_block"] = ( + block.model_dump() + if hasattr(block, "model_dump") + else None + ) + break + + # Preserve provider metadata (provider-agnostic passthrough) + # This enables providers to maintain state across steps (e.g., OpenAI reasoning items) + if hasattr(response, "metadata") and response.metadata: + assistant_msg["metadata"] = response.metadata + + await context.add_message(assistant_msg) + # 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: + continue + break + + # Add assistant message with tool calls + # Store structured content blocks (preserves reasoning state, thinking blocks, etc.) + # Extract text for display/logging only if hasattr(response, "text") and response.text: response_text = response.text else: - response_text = self._extract_text_from_content( - response.content + response_text = ( + self._extract_text_from_content(response.content) + if response.content + else "" ) - # Stream the final response token by token - async for token in self._tokenize_stream(response_text): - yield (token, iteration) - # Store structured content from response.content (our Pydantic models) - # This preserves reasoning state, thinking blocks, etc. - # response.content = list of our ContentBlock models (TextBlock, ThinkingBlock, etc.) - # response.content_blocks = raw SDK objects (for streaming events only) response_content = getattr(response, "content", None) if response_content and isinstance(response_content, list): - # Convert ContentBlock objects to dicts for serialization - content_dicts = [ - block.model_dump() - if hasattr(block, "model_dump") - else block - for block in response_content - ] - logger.info( - f"[ORCHESTRATOR] Storing {len(content_dicts)} content blocks" - ) - for i, block_dict in enumerate(content_dicts): - logger.info( - f"[ORCHESTRATOR] Block {i}: type={block_dict.get('type')}, has_content={'content' in block_dict}" - ) assistant_msg = { "role": "assistant", - "content": content_dicts, + "content": [ + block.model_dump() + if hasattr(block, "model_dump") + else block + for block in response_content + ], + "tool_calls": [ + { + "id": tc.id, + "tool": tc.name, + "arguments": tc.arguments, + } + for tc in tool_calls + ], } else: assistant_msg = { "role": "assistant", "content": response_text, + "tool_calls": [ + { + "id": tc.id, + "tool": tc.name, + "arguments": tc.arguments, + } + for tc in tool_calls + ], } # Preserve thinking blocks for Anthropic extended thinking (backward compat) @@ -3140,133 +3220,108 @@ async def _execute_stream( assistant_msg["metadata"] = response.metadata await context.add_message(assistant_msg) - # 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: - continue - break - # Add assistant message with tool calls - # Store structured content blocks (preserves reasoning state, thinking blocks, etc.) - # Extract text for display/logging only - if hasattr(response, "text") and response.text: - response_text = response.text - else: - response_text = ( - self._extract_text_from_content(response.content) - if response.content - else "" - ) + # Process tool calls in parallel (user guidance: assume parallel intent) + # Execute tools concurrently, but add results to context sequentially for determinism + import uuid - # Store structured content from response.content (our Pydantic models) - response_content = getattr(response, "content", None) - if response_content and isinstance(response_content, list): - assistant_msg = { - "role": "assistant", - "content": [ - block.model_dump() - if hasattr(block, "model_dump") - else block - for block in response_content - ], - "tool_calls": [ - { - "id": tc.id, - "tool": tc.name, - "arguments": tc.arguments, - } - for tc in tool_calls - ], - } - else: - assistant_msg = { - "role": "assistant", - "content": response_text, - "tool_calls": [ - { - "id": tc.id, - "tool": tc.name, - "arguments": tc.arguments, - } - for tc in tool_calls - ], - } + parallel_group_id = str(uuid.uuid4()) - # Preserve thinking blocks for Anthropic extended thinking (backward compat) - # Use response_content (our Pydantic models) not content_blocks (raw SDK objects) - if response_content and isinstance(response_content, list): - for block in response_content: - block_type = getattr(block, "type", None) - type_value = ( - getattr(block_type, "value", block_type) - if block_type - else None + # Execute all tools in parallel (no context updates inside) + # Wrap in try/except for CancelledError to handle immediate cancellation + tool_tasks = [ + self._execute_tool_only( + tc, tools, hooks, parallel_group_id, coordinator ) - if type_value == "thinking": - # Store the thinking block as dict to preserve signature - assistant_msg["thinking_block"] = ( - block.model_dump() - if hasattr(block, "model_dump") - else None - ) - break - - # Preserve provider metadata (provider-agnostic passthrough) - # This enables providers to maintain state across steps (e.g., OpenAI reasoning items) - if hasattr(response, "metadata") and response.metadata: - assistant_msg["metadata"] = response.metadata + for tc in tool_calls + ] - await context.add_message(assistant_msg) - - # Process tool calls in parallel (user guidance: assume parallel intent) - # Execute tools concurrently, but add results to context sequentially for determinism - import uuid - - parallel_group_id = str(uuid.uuid4()) - - # Execute all tools in parallel (no context updates inside) - # Wrap in try/except for CancelledError to handle immediate cancellation - tool_tasks = [ - self._execute_tool_only( - tc, tools, hooks, parallel_group_id, coordinator - ) - for tc in tool_calls - ] - - try: - tool_results = await asyncio.gather(*tool_tasks) - except asyncio.CancelledError: - # Immediate cancellation (second Ctrl+C) - synthesize cancelled results - # for ALL tool_calls to maintain tool_use/tool_result pairing - logger.info( - "Tool execution cancelled - synthesizing cancelled results" - ) - for tc in tool_calls: + try: + tool_results = await asyncio.gather(*tool_tasks) + except asyncio.CancelledError: + # Immediate cancellation (second Ctrl+C) - synthesize cancelled results + # for ALL tool_calls to maintain tool_use/tool_result pairing + logger.info( + "Tool execution cancelled - synthesizing cancelled results" + ) + for tc in tool_calls: + await context.add_message( + { + "role": "tool", + "name": tc.name, + "tool_call_id": tc.id, + "content": f'{{"error": "Tool execution was cancelled by user", "cancelled": true, "tool": "{tc.name}"}}', + } + ) + # Emit cancel events before re-raising so hooks receive them + if coordinator and not self._cancel_requested_emitted: + self._cancel_requested_emitted = True + await hooks.emit( + CANCEL_REQUESTED, + { + "orchestrator": "loop-streaming", + "state": str(coordinator.cancellation.state), + "turn_count": iteration, + }, + ) + try: + await coordinator.cancellation.trigger_callbacks() + except Exception as e: + logger.warning(f"Error in cancellation callbacks: {e}") + if coordinator: + await hooks.emit( + CANCEL_COMPLETED, + { + "orchestrator": "loop-streaming", + "was_immediate": coordinator.cancellation.is_immediate, + "turn_count": iteration, + }, + ) + # Write synthetic assistant message to close the turn. + # Without this, transcript has tool_results without a closing assistant + # message, triggering FM3 (incomplete_assistant_turn) on resume. await context.add_message( { - "role": "tool", - "name": tc.name, - "tool_call_id": tc.id, - "content": f'{{"error": "Tool execution was cancelled by user", "cancelled": true, "tool": "{tc.name}"}}', + "role": "assistant", + "content": "The previous operation was cancelled. Results from completed tools have been preserved.", } ) - # Emit cancel events before re-raising so hooks receive them - if coordinator and not self._cancel_requested_emitted: - self._cancel_requested_emitted = True - await hooks.emit( - CANCEL_REQUESTED, - { - "orchestrator": "loop-streaming", - "state": str(coordinator.cancellation.state), - "turn_count": iteration, - }, - ) - try: - await coordinator.cancellation.trigger_callbacks() - except Exception as e: - logger.warning(f"Error in cancellation callbacks: {e}") - if coordinator: + # Re-raise to let the cancellation propagate. + # Clear pending steers first — a steer queued during tool + # execution must not leak into any future turn. (spec §5.2) + self._steering_queue.clear() + raise + + # Check for cancellation after tools complete (graceful cancellation) + if coordinator and coordinator.cancellation.is_cancelled: + # MUST add tool results to context before returning + # Otherwise we leave orphaned tool_calls without matching tool_results + # which violates provider API contracts (Anthropic, OpenAI) + for tool_call_id, tool_name, content in tool_results: + await context.add_message( + { + "role": "tool", + "name": tool_name, + "tool_call_id": tool_call_id, + "content": content, + } + ) + # Emit cancel:requested on first detection and trigger cleanup callbacks + if not self._cancel_requested_emitted: + self._cancel_requested_emitted = True + await hooks.emit( + CANCEL_REQUESTED, + { + "orchestrator": "loop-streaming", + "state": str(coordinator.cancellation.state), + "turn_count": iteration, + }, + ) + try: + await coordinator.cancellation.trigger_callbacks() + except Exception as e: + logger.warning(f"Error in cancellation callbacks: {e}") + # Emit cancel:completed — orchestrator is exiting due to cancellation await hooks.emit( CANCEL_COMPLETED, { @@ -3275,26 +3330,23 @@ async def _execute_stream( "turn_count": iteration, }, ) - # Write synthetic assistant message to close the turn. - # Without this, transcript has tool_results without a closing assistant - # message, triggering FM3 (incomplete_assistant_turn) on resume. - await context.add_message( - { - "role": "assistant", - "content": "The previous operation was cancelled. Results from completed tools have been preserved.", - } - ) - # Re-raise to let the cancellation propagate. - # Clear pending steers first — a steer queued during tool - # execution must not leak into any future turn. (spec §5.2) - self._steering_queue.clear() - raise + # Write synthetic assistant message to close the turn. + # Without this, transcript has tool_results without a closing assistant + # message, triggering FM3 (incomplete_assistant_turn) on resume. + await context.add_message( + { + "role": "assistant", + "content": "The previous operation was cancelled. Results from completed tools have been preserved.", + } + ) + # Exit the loop - orchestrator complete event will be emitted in execute(). + # Clear pending steers: cancellation closes the turn; any steer that + # arrived after the last injection point must not ride a future turn. (spec §5.2) + self._steering_queue.clear() + return - # Check for cancellation after tools complete (graceful cancellation) - if coordinator and coordinator.cancellation.is_cancelled: - # MUST add tool results to context before returning - # Otherwise we leave orphaned tool_calls without matching tool_results - # which violates provider API contracts (Anthropic, OpenAI) + # Add all results to context in original order (sequential, deterministic) + # Note: Context manager handles compaction internally when get_messages_for_request() is called for tool_call_id, tool_name, content in tool_results: await context.add_message( { @@ -3304,140 +3356,110 @@ async def _execute_stream( "content": content, } ) - # Emit cancel:requested on first detection and trigger cleanup callbacks - if not self._cancel_requested_emitted: - self._cancel_requested_emitted = True - await hooks.emit( - CANCEL_REQUESTED, - { - "orchestrator": "loop-streaming", - "state": str(coordinator.cancellation.state), - "turn_count": iteration, - }, - ) - try: - await coordinator.cancellation.trigger_callbacks() - except Exception as e: - logger.warning(f"Error in cancellation callbacks: {e}") - # Emit cancel:completed — orchestrator is exiting due to cancellation - await hooks.emit( - CANCEL_COMPLETED, - { - "orchestrator": "loop-streaming", - "was_immediate": coordinator.cancellation.is_immediate, - "turn_count": iteration, - }, - ) - # Write synthetic assistant message to close the turn. - # Without this, transcript has tool_results without a closing assistant - # message, triggering FM3 (incomplete_assistant_turn) on resume. - await context.add_message( - { - "role": "assistant", - "content": "The previous operation was cancelled. Results from completed tools have been preserved.", - } - ) - # Exit the loop - orchestrator complete event will be emitted in execute(). - # Clear pending steers: cancellation closes the turn; any steer that - # arrived after the last injection point must not ride a future turn. (spec §5.2) - self._steering_queue.clear() - return - - # Add all results to context in original order (sequential, deterministic) - # Note: Context manager handles compaction internally when get_messages_for_request() is called - for tool_call_id, tool_name, content in tool_results: - await context.add_message( - { - "role": "tool", - "name": tool_name, - "tool_call_id": tool_call_id, - "content": content, - } - ) - # Check if we exceeded max iterations (only if not unlimited) - if self.max_iterations != -1 and iteration >= self.max_iterations: - logger.warning(f"Max iterations ({self.max_iterations}) reached") - - # Inject system reminder to agent before returning - await hooks.emit( - PROVIDER_REQUEST, - { - "provider": provider_name, - "iteration": iteration, - "max_reached": True, - }, - ) - - # 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) - message_dicts.append( - { - "role": "user", - "content": """ -You have reached the maximum number of iterations for this turn. Please provide a response to the user now, summarizing your progress and noting what remains to be done. You can continue in the next turn if needed. - -DO NOT mention this iteration limit or reminder to the user explicitly. Simply wrap up naturally. -""", - } - ) + # Check if we exceeded max iterations (only if not unlimited) + if self.max_iterations != -1 and iteration >= self.max_iterations: + logger.warning(f"Max iterations ({self.max_iterations}) reached") - try: - # Convert dicts to ChatRequest - messages_objects = [Message(**msg) for msg in message_dicts] + # Inject system reminder to agent before returning + await hooks.emit( + PROVIDER_REQUEST, + { + "provider": provider_name, + "iteration": iteration, + "max_reached": True, + }, + ) - # Convert tools to ToolSpec format for ChatRequest - tools_list = None - if tools: - tools_list = [_build_tool_spec(t) for t in tools.values()] + # 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) + message_dicts.append( + { + "role": "user", + "content": """ + You have reached the maximum number of iterations for this turn. Please provide a response to the user now, summarizing your progress and noting what remains to be done. You can continue in the next turn if needed. - max_iter_chat_request = ChatRequest( - messages=messages_objects, - tools=tools_list, - reasoning_effort=self.config.get("reasoning_effort"), + DO NOT mention this iteration limit or reminder to the user explicitly. Simply wrap up naturally. + """, + } ) - kwargs = {} - if self.extended_thinking: - kwargs["extended_thinking"] = True + try: + # Convert dicts to ChatRequest + messages_objects = [Message(**msg) for msg in message_dicts] + + # Convert tools to ToolSpec format for ChatRequest + tools_list = None + if tools: + tools_list = [_build_tool_spec(t) for t in tools.values()] + + max_iter_chat_request = ChatRequest( + messages=messages_objects, + tools=tools_list, + reasoning_effort=self.config.get("reasoning_effort"), + ) - response = await provider.complete(max_iter_chat_request, **kwargs) - content = ( - response.content if hasattr(response, "content") else str(response) - ) + kwargs = {} + if self.extended_thinking: + kwargs["extended_thinking"] = True - if content: - # Yield the final response - async for token in self._tokenize_stream(content): - yield (token, iteration) + response = await provider.complete(max_iter_chat_request, **kwargs) + content = ( + response.content + if hasattr(response, "content") + else str(response) + ) - # Add to context - await context.add_message({"role": "assistant", "content": content}) + if content: + # Yield the final response + async for token in self._tokenize_stream(content): + yield (token, iteration) - except LLMError as e: - await hooks.emit( - PROVIDER_ERROR, - { - "provider": provider_name, - "error": {"type": type(e).__name__, "msg": str(e)}, - "retryable": e.retryable, - "status_code": e.status_code, - }, - ) - logger.error(f"Error getting final response after max iterations: {e}") - except Exception as e: - await hooks.emit( - PROVIDER_ERROR, - { - "provider": provider_name, - "error": {"type": type(e).__name__, "msg": str(e)}, - }, - ) - logger.error(f"Error getting final response after max iterations: {e}") + # Add to context + await context.add_message( + {"role": "assistant", "content": content} + ) - # Emit execution end - await hooks.emit("execution:end", {}) + except LLMError as e: + await hooks.emit( + PROVIDER_ERROR, + { + "provider": provider_name, + "error": {"type": type(e).__name__, "msg": str(e)}, + "retryable": e.retryable, + "status_code": e.status_code, + }, + ) + logger.error( + f"Error getting final response after max iterations: {e}" + ) + except Exception as e: + await hooks.emit( + PROVIDER_ERROR, + { + "provider": provider_name, + "error": {"type": type(e).__name__, "msg": str(e)}, + }, + ) + logger.error( + f"Error getting final response after max iterations: {e}" + ) + + finally: + # A `finally` inside the generator, not a patch at each early + # return. Four returns between here and the old emit site + # (:2793, :2813, :2961, :3344 pre-change) each skipped it, and a + # fifth would have too. This also covers the consumer breaking + # out of its `async for`: Python raises GeneratorExit at the + # suspended yield, and `finally` still runs. + # + # Symptom it fixes: 12 of 27 executions never emitted an end, + # exactly matching the 12 cancellations, leaving the turn state + # machine stuck in 'executing'. + await hooks.emit("execution:end", {}) async def _stream_from_provider( self, @@ -3611,6 +3633,57 @@ async def _execute_tool( tool_call, tools, context, hooks, coordinator ) + _FAILURE_BREAKER_THRESHOLD = 3 + + def _apply_failure_breaker(self, tool_call, result): + """Tell the model when a call keeps failing the exact same way. + + Observed in session ``eec9ae98``: the SAME failing ``read_file`` call + was issued 13 times with nothing intervening. Whitespace-malformed + arguments fail deterministically -- same input, same error, forever -- + so an agent that cannot notice the repetition burns tokens and + wall-clock producing nothing. + + Keyed on (tool, arguments, error) and NOT on arguments alone, because + legitimate repeats exist: polling a file being written, ``git status`` + in a loop, retrying after fixing something externally. Only a call that + fails the SAME way counts toward the trip. + + The trip is SURFACED to the model, never silently dropped -- a silent + breaker is the same class of bug as a silent argument rewrite. The tool + still ran and its real error is preserved; a note is appended. + """ + if getattr(result, "success", True): + return result + + error = getattr(result, "error", None) or {} + message = error.get("message", "") if isinstance(error, dict) else str(error) + try: + arguments = json.dumps(tool_call.arguments, sort_keys=True, default=str) + except (TypeError, ValueError): + arguments = str(tool_call.arguments) + key = f"{tool_call.name}\x00{arguments}\x00{message}" + + count = self._repeated_failures.get(key, 0) + 1 + self._repeated_failures[key] = count + if count < self._FAILURE_BREAKER_THRESHOLD: + return result + + logger.warning( + f"Identical failure repeated {count}x for tool '{tool_call.name}'; " + f"surfacing a breaker note to the model" + ) + return ToolResult( + success=False, + error={ + "message": ( + f"{message}\n\n[This exact call to `{tool_call.name}` has now failed " + f"{count} times with an identical error. Repeating it will not " + f"succeed. Change the arguments or take a different approach.]" + ) + }, + ) + async def _execute_tool_only( self, tool_call, @@ -3646,6 +3719,28 @@ async def _execute_tool_only( f"Denied by hook: {pre_result.reason}", ) + # Adopt hook-modified arguments. The kernel NORMALIZES a + # `modify` result away: `emit()` returns `action="continue"` + # with the modified payload in `data`, so a check for + # `action == "modify"` can never fire and the documented + # pattern in ORCHESTRATOR_CONTRACT.md is unreachable code. + # Reading `data` unconditionally is the only consumption that + # works -- `emit()` always populates it. + # + # Without this, every `tool:pre` handler that rewrites input + # (argument normalization, path jailing, secret scrubbing) is a + # silent no-op: the hook runs, returns its correction, and the + # original arguments execute anyway. `tool:post` already honors + # modifications; this makes `tool:pre` symmetric. + # `isinstance(dict)` and not just truthiness: a partial or + # mocked result can carry a non-dict `data`, and adopting from + # it would replace real arguments with nonsense. + hook_data = pre_result.data + if isinstance(hook_data, dict): + hook_input = hook_data.get("tool_input") + if hook_input is not None and hook_input != tool_call.arguments: + tool_call.arguments = hook_input + # Get tool tool = tools.get(tool_call.name) if not tool: @@ -3711,6 +3806,8 @@ async def _execute_tool_only( asyncio.current_task(), None ) + result = self._apply_failure_breaker(tool_call, result) + # Serialize result for logging result_data = ( result.model_dump() if hasattr(result, "model_dump") else str(result) @@ -3826,6 +3923,28 @@ async def _execute_tool_with_result( response_added = True return {"success": False, "error": f"Denied: {pre_result.reason}"} + # Adopt hook-modified arguments. The kernel NORMALIZES a + # `modify` result away: `emit()` returns `action="continue"` + # with the modified payload in `data`, so a check for + # `action == "modify"` can never fire and the documented + # pattern in ORCHESTRATOR_CONTRACT.md is unreachable code. + # Reading `data` unconditionally is the only consumption that + # works -- `emit()` always populates it. + # + # Without this, every `tool:pre` handler that rewrites input + # (argument normalization, path jailing, secret scrubbing) is a + # silent no-op: the hook runs, returns its correction, and the + # original arguments execute anyway. `tool:post` already honors + # modifications; this makes `tool:pre` symmetric. + # `isinstance(dict)` and not just truthiness: a partial or + # mocked result can carry a non-dict `data`, and adopting from + # it would replace real arguments with nonsense. + hook_data = pre_result.data + if isinstance(hook_data, dict): + hook_input = hook_data.get("tool_input") + if hook_input is not None and hook_input != tool_call.arguments: + tool_call.arguments = hook_input + # Get tool tool = tools.get(tool_call.name) if not tool: @@ -3869,6 +3988,8 @@ async def _execute_tool_with_result( asyncio.current_task(), None ) + result = self._apply_failure_breaker(tool_call, result) + # Serialize result for logging result_data = ( result.model_dump() if hasattr(result, "model_dump") else str(result) diff --git a/pyproject.toml b/pyproject.toml index 6826251..d8a6c92 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,12 +34,18 @@ allow-direct-references = true [dependency-groups] dev = [ "amplifier-core", + # tests/test_goal_loop.py imports ProviderPreference from amplifier-foundation. + # Without it the whole file fails to COLLECT, which takes the suite down at + # collection time rather than failing one test -- so `uv run pytest` on a + # clean checkout errored out entirely. + "amplifier-foundation", "pytest>=9.0.3", "pytest-asyncio>=1.0.0", ] [tool.uv.sources] amplifier-core = { git = "https://github.com/microsoft/amplifier-core", branch = "main" } +amplifier-foundation = { git = "https://github.com/microsoft/amplifier-foundation", branch = "main" } [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/tests/test_execution_end_invariant.py b/tests/test_execution_end_invariant.py new file mode 100644 index 0000000..5cf1a64 --- /dev/null +++ b/tests/test_execution_end_invariant.py @@ -0,0 +1,125 @@ +"""``execution:end`` must fire even when the turn exits early. + +Regression cover for session ``eec9ae98``: **27 ``execution:start`` events, 15 +``execution:end``**. The 12 missing ends matched the 12 cancellations exactly, +on both the kernel event log and the UI event stream, leaving the turn state +machine stuck in "executing" and never unwinding. + +Cause: ``_execute_stream`` emitted the end event on its last line, and several +paths returned before reaching it -- the graceful-cancellation exit, immediate +cancellation between chunks, a denied ``provider:request``, and "no providers +available". A consumer breaking out of its ``async for`` skipped it too. + +The fix is a ``finally`` inside the generator rather than an emit bolted onto +each early return: it covers the paths that exist, the paths nobody has written +yet, and ``GeneratorExit``. +""" + +from __future__ import annotations + +from typing import Any + +import pytest +from amplifier_core import HookRegistry + + +def _orchestrator() -> Any: + from amplifier_module_loop_streaming import StreamingOrchestrator + + return StreamingOrchestrator(config={}) + + +class _Context: + """Enough context surface to reach the early return under test.""" + + def __init__(self) -> None: + self.messages: list[dict[str, Any]] = [] + + async def add_message(self, message: dict[str, Any]) -> None: + self.messages.append(message) + + async def get_messages(self) -> list[dict[str, Any]]: + return list(self.messages) + + async def get_messages_for_request(self) -> list[dict[str, Any]]: + return list(self.messages) + + +def _recording_hooks() -> tuple[HookRegistry, list[str]]: + hooks = HookRegistry() + seen: list[str] = [] + + async def record(event: str, data: Any) -> None: + del data + seen.append(event) + + hooks.register("execution:start", record) + hooks.register("execution:end", record) + return hooks, seen + + +@pytest.mark.asyncio +async def test_execution_end_fires_on_the_no_provider_early_return() -> None: + """The simplest early exit that lands after ``execution:start``. + + Before the fix this path emitted a start with no matching end, which is the + shape that left the turn state machine stuck. + """ + hooks, seen = _recording_hooks() + orchestrator = _orchestrator() + + tokens = [ + token + async for token, _iteration in orchestrator._execute_stream( + "do the thing", _Context(), {}, {}, hooks + ) + ] + + assert any("No providers available" in token for token in tokens), ( + f"fixture did not reach the intended early return; got {tokens!r}" + ) + assert seen.count("execution:start") == 1 + assert seen.count("execution:end") == 1, ( + f"execution:end did not fire on an early return: {seen}" + ) + + +@pytest.mark.asyncio +async def test_execution_end_fires_when_the_consumer_stops_reading() -> None: + """A consumer that breaks out of ``async for`` must still close the turn. + + This is the cancellation shape from the incident: the turn stops because + something upstream stopped listening, not because the loop ran to + completion. Python raises ``GeneratorExit`` at the suspended yield, so the + ``finally`` runs -- an emit bolted onto each ``return`` would not have. + """ + hooks, seen = _recording_hooks() + orchestrator = _orchestrator() + + stream = orchestrator._execute_stream("do the thing", _Context(), {}, {}, hooks) + async for _token, _iteration in stream: + break # stop reading after the first token + await stream.aclose() + + assert seen.count("execution:start") == 1 + assert seen.count("execution:end") == 1, ( + f"execution:end did not fire when the consumer stopped reading: {seen}" + ) + + +@pytest.mark.asyncio +async def test_every_start_has_exactly_one_end_across_repeated_turns() -> None: + """The invariant the incident violated, stated directly: 27 starts, 15 ends.""" + hooks, seen = _recording_hooks() + orchestrator = _orchestrator() + + for _ in range(5): + async for _token, _iteration in orchestrator._execute_stream( + "do the thing", _Context(), {}, {}, hooks + ): + pass + + assert seen.count("execution:start") == 5 + assert seen.count("execution:end") == 5, ( + f"starts and ends are unbalanced across turns: {seen}" + ) diff --git a/tests/test_failure_circuit_breaker.py b/tests/test_failure_circuit_breaker.py new file mode 100644 index 0000000..e24909a --- /dev/null +++ b/tests/test_failure_circuit_breaker.py @@ -0,0 +1,168 @@ +"""Stop an agent re-issuing a call that keeps failing the exact same way. + +In session ``eec9ae98`` the SAME failing ``read_file`` call was issued **13 +times** with nothing intervening. Whitespace-malformed arguments fail +deterministically -- same input, same error, forever -- so an agent that cannot +notice the repetition burns tokens and wall-clock producing nothing. 37 distinct +tool inputs were issued more than once; 73 calls were redundant. + +The breaker keys on **(tool, arguments, error)** and deliberately NOT on +arguments alone. Legitimate repeats exist: polling a file being written, +``git status`` in a loop, retrying after fixing something externally. Only a +call that fails the SAME way counts toward the trip. + +The trip is SURFACED to the model, never silently dropped -- a silent breaker is +the same class of bug as a silent argument rewrite. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from amplifier_core import HookRegistry, ToolResult +from amplifier_module_loop_streaming import StreamingOrchestrator + +THRESHOLD = StreamingOrchestrator._FAILURE_BREAKER_THRESHOLD + + +def _orchestrator() -> StreamingOrchestrator: + return StreamingOrchestrator(config={}) + + +def _call(arguments: dict[str, Any] | None = None, name: str = "read_file") -> Any: + tool_call = MagicMock() + tool_call.id = "call-1" + tool_call.name = name + tool_call.arguments = arguments if arguments is not None else {"path": " /tmp/x "} + return tool_call + + +def _failure(message: str = "Path not found: /tmp/x ") -> ToolResult: + return ToolResult(success=False, error={"message": message}) + + +def _message(result: ToolResult) -> str: + return (result.error or {}).get("message", "") + + +def test_the_first_failures_pass_through_untouched() -> None: + """Below the threshold the model sees exactly what the tool said.""" + orchestrator = _orchestrator() + call = _call() + + for _ in range(THRESHOLD - 1): + out = orchestrator._apply_failure_breaker(call, _failure()) + assert _message(out) == "Path not found: /tmp/x " + assert "has now failed" not in _message(out) + + +def test_the_same_failure_repeated_trips_the_breaker() -> None: + """The defect: 13 identical failures with nothing intervening.""" + orchestrator = _orchestrator() + call = _call() + + for _ in range(THRESHOLD - 1): + orchestrator._apply_failure_breaker(call, _failure()) + tripped = orchestrator._apply_failure_breaker(call, _failure()) + + note = _message(tripped) + assert "Path not found: /tmp/x " in note, "the tool's real error must be preserved" + assert f"has now failed {THRESHOLD} times" in note + assert "read_file" in note + assert "different approach" in note, ( + "the note must tell the model what to do instead" + ) + + +def test_a_different_error_for_the_same_input_does_not_count() -> None: + """Same call, different failure, is not the loop this guards against.""" + orchestrator = _orchestrator() + call = _call() + + for i in range(THRESHOLD * 2): + out = orchestrator._apply_failure_breaker( + call, _failure(f"transient error {i}") + ) + assert "has now failed" not in _message(out) + + +def test_different_arguments_do_not_count_toward_each_other() -> None: + """Two paths that each fail once are not one call failing twice.""" + orchestrator = _orchestrator() + + for i in range(THRESHOLD * 2): + call = _call({"path": f" /tmp/{i} "}) + out = orchestrator._apply_failure_breaker(call, _failure("Path not found")) + assert "has now failed" not in _message(out) + + +def test_the_same_failure_from_a_different_tool_does_not_count() -> None: + orchestrator = _orchestrator() + + for name in ("read_file", "write_file", "glob", "grep"): + call = _call(name=name) + out = orchestrator._apply_failure_breaker(call, _failure("Path not found")) + assert "has now failed" not in _message(out) + + +def test_success_never_trips_and_is_returned_unchanged() -> None: + """Polling a file being written must not be mistaken for a stuck loop.""" + orchestrator = _orchestrator() + call = _call() + success = ToolResult(success=True, data={"content": "ok"}) + + for _ in range(THRESHOLD * 3): + assert orchestrator._apply_failure_breaker(call, success) is success + + +def test_unhashable_arguments_do_not_break_dispatch() -> None: + """Arguments are not guaranteed to be JSON-serialisable.""" + orchestrator = _orchestrator() + call = _call({"path": object()}) + + for _ in range(THRESHOLD): + out = orchestrator._apply_failure_breaker(call, _failure()) + assert "has now failed" in _message(out) + + +class _AlwaysFailingTool: + @property + def name(self) -> str: + return "read_file" + + async def execute(self, arguments: Any) -> ToolResult: + del arguments + return _failure() + + +def _coordinator() -> Any: + coordinator = MagicMock() + coordinator._tool_dispatch_contexts = {} + coordinator.cancellation.register_tool_start = MagicMock() + coordinator.cancellation.register_tool_complete = MagicMock() + result = MagicMock() + result.action = "continue" + result.data = None + coordinator.process_hook_result = AsyncMock(return_value=result) + return coordinator + + +@pytest.mark.asyncio +async def test_the_breaker_reaches_the_model_through_real_dispatch() -> None: + """End to end on the parallel path: the note lands in the tool content.""" + orchestrator = _orchestrator() + tools = {"read_file": _AlwaysFailingTool()} + + contents: list[str] = [] + for _ in range(THRESHOLD): + _id, _name, content = await orchestrator._execute_tool_only( + _call(), tools, HookRegistry(), "group-1", _coordinator() + ) + contents.append(content) + + assert "has now failed" not in contents[0] + assert "has now failed" in contents[-1], ( + f"the breaker note never reached the model: {contents[-1]!r}" + ) diff --git a/tests/test_tool_pre_input_adoption.py b/tests/test_tool_pre_input_adoption.py new file mode 100644 index 0000000..2fc07fb --- /dev/null +++ b/tests/test_tool_pre_input_adoption.py @@ -0,0 +1,159 @@ +"""A ``tool:pre`` hook that rewrites arguments must actually change what runs. + +Both dispatch paths emitted ``tool:pre``, passed the result through +``coordinator.process_hook_result``, branched on ``action == "deny"`` -- and +then executed the ORIGINAL arguments regardless. Every hook that rewrites input +(argument normalization, path jailing, secret scrubbing) was a silent no-op: +it ran, returned its correction, and was ignored. + +The subtlety that made this easy to get wrong: the kernel NORMALIZES ``modify`` +away. ``emit()`` returns ``action="continue"`` with the modified payload in +``data`` (see ``hooks.rs``), so the pattern documented in +``ORCHESTRATOR_CONTRACT.md`` -- ``if result.action == "modify"`` -- is +unreachable code and can never fire. Reading ``data`` unconditionally is the +only consumption that works. + +``tool:post`` already honored modifications. This makes ``tool:pre`` symmetric. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from amplifier_core import HookRegistry, ToolResult +from amplifier_module_loop_streaming import StreamingOrchestrator + +PADDED = {"action": " create ", "path": " /tmp/x "} +CLEANED = {"action": "create", "path": "/tmp/x"} + + +class _RecordingTool: + """Captures exactly what arguments reached ``execute``.""" + + def __init__(self) -> None: + self.seen: list[Any] = [] + + @property + def name(self) -> str: + return "todo" + + async def execute(self, arguments: Any) -> ToolResult: + self.seen.append(arguments) + return ToolResult(success=True, data={"ok": True}) + + +def _tool_call(arguments: dict[str, Any] | None = None) -> Any: + call = MagicMock() + call.id = "call-1" + call.name = "todo" + call.arguments = dict(PADDED if arguments is None else arguments) + return call + + +def _coordinator(hook_data: Any) -> Any: + """A coordinator whose ``tool:pre`` result carries *hook_data* as ``.data``.""" + coordinator = MagicMock() + coordinator._tool_dispatch_contexts = {} + coordinator.cancellation.register_tool_start = MagicMock() + coordinator.cancellation.register_tool_complete = MagicMock() + result = MagicMock() + result.action = "continue" + result.data = hook_data + coordinator.process_hook_result = AsyncMock(return_value=result) + return coordinator + + +class _Context: + def __init__(self) -> None: + self.messages: list[dict[str, Any]] = [] + + async def add_message(self, message: dict[str, Any]) -> None: + self.messages.append(message) + + +@pytest.mark.asyncio +async def test_parallel_path_adopts_rewritten_arguments() -> None: + tool = _RecordingTool() + call = _tool_call() + + await StreamingOrchestrator(config={})._execute_tool_only( + call, + {"todo": tool}, + HookRegistry(), + "group-1", + _coordinator({"tool_input": CLEANED}), + ) + + assert tool.seen == [CLEANED], ( + f"the tool ran with {tool.seen!r}; a tool:pre hook's correction was discarded" + ) + + +@pytest.mark.asyncio +async def test_sequential_path_adopts_rewritten_arguments() -> None: + """The second dispatch site -- the one that runs when tools are not batched.""" + tool = _RecordingTool() + call = _tool_call() + + await StreamingOrchestrator(config={})._execute_tool_with_result( + call, + {"todo": tool}, + _Context(), + HookRegistry(), + _coordinator({"tool_input": CLEANED}), + ) + + assert tool.seen == [CLEANED] + + +@pytest.mark.asyncio +async def test_an_unmodified_payload_is_not_treated_as_a_rewrite() -> None: + """The kernel round-trips the payload, so ``data`` is always a NEW object. + + Equal content must not count as a modification, or every call would look + like it had been rewritten. + """ + tool = _RecordingTool() + call = _tool_call() + original = call.arguments + + await StreamingOrchestrator(config={})._execute_tool_only( + call, + {"todo": tool}, + HookRegistry(), + "group-1", + _coordinator({"tool_input": dict(PADDED)}), # equal content, different object + ) + + assert tool.seen == [PADDED] + assert call.arguments is original, "arguments were replaced by an equal copy" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "hook_data", + [ + pytest.param(None, id="no-data"), + pytest.param({}, id="empty-data"), + pytest.param({"tool_name": "todo"}, id="data-without-tool-input"), + pytest.param(MagicMock(), id="non-dict-data"), + pytest.param("not a mapping", id="string-data"), + ], +) +async def test_a_partial_result_never_corrupts_the_arguments(hook_data: Any) -> None: + """Adopt only from a real mapping. + + A partial or mocked hook result can carry a non-dict ``data``; reading it + loosely would replace real arguments with nonsense, which is a worse + failure than the no-op this fix removes. + """ + tool = _RecordingTool() + call = _tool_call() + + await StreamingOrchestrator(config={})._execute_tool_only( + call, {"todo": tool}, HookRegistry(), "group-1", _coordinator(hook_data) + ) + + assert tool.seen == [PADDED], f"arguments were corrupted to {tool.seen!r}" diff --git a/uv.lock b/uv.lock index e673424..277248a 100644 --- a/uv.lock +++ b/uv.lock @@ -14,6 +14,15 @@ dependencies = [ { name = "typing-extensions" }, ] +[[package]] +name = "amplifier-foundation" +version = "1.0.0" +source = { git = "https://github.com/microsoft/amplifier-foundation?branch=main#3f9a6e28c8f92b36f83c44942cca18927cb3eeae" } +dependencies = [ + { name = "amplifier-core" }, + { name = "pyyaml" }, +] + [[package]] name = "amplifier-module-loop-streaming" version = "1.0.0" @@ -22,6 +31,7 @@ source = { editable = "." } [package.dev-dependencies] dev = [ { name = "amplifier-core" }, + { name = "amplifier-foundation" }, { name = "pytest" }, { name = "pytest-asyncio" }, ] @@ -31,6 +41,7 @@ dev = [ [package.metadata.requires-dev] dev = [ { name = "amplifier-core", git = "https://github.com/microsoft/amplifier-core?branch=main" }, + { name = "amplifier-foundation", git = "https://github.com/microsoft/amplifier-foundation?branch=main" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-asyncio", specifier = ">=1.0.0" }, ]