diff --git a/README.md b/README.md index d6f58d1..c60827b 100644 --- a/README.md +++ b/README.md @@ -31,7 +31,10 @@ Provides streaming orchestration that delivers LLM responses token-by-token for - Token-level streaming from provider - Real-time response delivery -- **Parallel tool execution**: Multiple tool calls execute concurrently +- **Parallel tool execution by default**: Multiple ordinary tool calls execute + concurrently. A provider-marked native-toolset batch executes in response + order and stops after an error, so stateful native computer actions cannot + race. - Deterministic context updates: Results added in original order - Progressive rendering - Interruptible generation diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index b1f488f..de91950 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -2740,6 +2740,7 @@ async def _judge_stall( model=model_override, metadata={"stream": False}, max_output_tokens=self._GOAL_INTERNAL_CALL_MAX_TOKENS, + reasoning_effort="high", ) request_result = await hooks.emit( @@ -2787,14 +2788,16 @@ async def _judge_stall( # hooks-session-naming's identical opt-out, same file/line cited # above). # - # `role_config` (from the resolved model role's `ProviderPreference`, - # e.g. `{"reasoning_effort": "high"}`) is forwarded as complete() - # kwargs first -- `extended_thinking=False` is applied AFTER so it - # always wins regardless of what the routing matrix's per-role - # config carries. + # Fresh utility calls use the compatible `high` effort overlay. This + # prevents a main-conversation xhigh/max role setting from becoming + # incompatible with their explicit thinking opt-out. Preserve all + # other role config (including credentials and source metadata). complete_kwargs: dict[str, Any] = dict(role_config) + complete_kwargs.pop("thinking_budget_tokens", None) + complete_kwargs.pop("thinking_display", None) if model_override: complete_kwargs["model"] = model_override + complete_kwargs["effort"] = "high" complete_kwargs["extended_thinking"] = False try: response = await provider.complete(chat_request, **complete_kwargs) @@ -3012,6 +3015,7 @@ async def _summarize_goal_run( model=model_override, metadata={"stream": False}, max_output_tokens=self._GOAL_INTERNAL_CALL_MAX_TOKENS, + reasoning_effort="high", ) request_result = await hooks.emit( @@ -3042,12 +3046,14 @@ async def _summarize_goal_run( # simply not setting it is not sufficient (a session-level # provider config can force thinking on regardless). # - # `role_config` forwarded first, extended_thinking=False applied - # after so it always wins (see the matching comment in - # _judge_stall / _resolve_goal_model). + # See _judge_stall: utility calls use a compatible `high` effort + # overlay while retaining unrelated role config. complete_kwargs: dict[str, Any] = dict(role_config) + complete_kwargs.pop("thinking_budget_tokens", None) + complete_kwargs.pop("thinking_display", None) if model_override: complete_kwargs["model"] = model_override + complete_kwargs["effort"] = "high" complete_kwargs["extended_thinking"] = False response = await provider.complete(chat_request, **complete_kwargs) summary_text = "" @@ -3188,6 +3194,7 @@ async def _evaluate_goal( model=model_override, metadata={"stream": False}, max_output_tokens=self._GOAL_INTERNAL_CALL_MAX_TOKENS, + reasoning_effort="high", ) # Mirror _execute_stream's provider:request instrumentation so hooks @@ -3223,12 +3230,14 @@ async def _evaluate_goal( # confirmed (via real-session telemetry) to have run with thinking # enabled and a 32000-token budget in session e97e192b. # - # `role_config` forwarded first, extended_thinking=False applied - # after so it always wins (see the matching comment in - # _judge_stall / _resolve_goal_model). + # See _judge_stall: utility calls use a compatible `high` effort + # overlay while retaining unrelated role config. complete_kwargs: dict[str, Any] = dict(role_config) + complete_kwargs.pop("thinking_budget_tokens", None) + complete_kwargs.pop("thinking_display", None) if model_override: complete_kwargs["model"] = model_override + complete_kwargs["effort"] = "high" complete_kwargs["extended_thinking"] = False try: response = await provider.complete(chat_request, **complete_kwargs) @@ -3356,6 +3365,8 @@ async def _execute_stream( retaining_getter = candidate measured_view_getter = None foreground_usage_claim = None + record_final_request = None + bind_signed_response = None if callable(get_capability): candidate = get_capability("context.measured_request_view") if callable(candidate): @@ -3363,6 +3374,31 @@ async def _execute_stream( candidate = get_capability("context.foreground_usage") if callable(candidate): foreground_usage_claim = candidate + candidate = get_capability("context.final_request_record") + if callable(candidate): + record_final_request = candidate + candidate = get_capability("context.signed_replay") + if callable(candidate): + bind_signed_response = candidate + + # These capabilities form one handoff, not two independent features: + # Context must both snapshot the exact request about to reach the + # provider and associate its admitted assistant response with that + # snapshot. A mixed-version Context gets the legacy path unchanged. + if not (callable(record_final_request) and callable(bind_signed_response)): + record_final_request = None + bind_signed_response = None + + def record_dispatch(request: ChatRequest) -> Any | None: + """Synchronously snapshot one final foreground provider dispatch.""" + if record_final_request is None: + return None + return record_final_request(request) + + def bind_admitted_response(ticket: Any | None) -> None: + """Synchronously associate an admitted assistant turn with its request.""" + if ticket is not None and bind_signed_response is not None: + bind_signed_response(ticket) if ( self._ephemeral_injection_mode == "persist" and retaining_getter is None @@ -5022,6 +5058,8 @@ async def count_view( ), ), on_unconfirmed_stream=record_unavailable_stream_usage, + record_dispatch=record_dispatch, + bind_admitted_response=bind_admitted_response, ) else: response_stream = self._stream_from_provider( @@ -5033,6 +5071,8 @@ async def count_view( coordinator, provider_name=provider_name, on_stream_error=record_unavailable_stream_usage, + record_dispatch=record_dispatch, + bind_admitted_response=bind_admitted_response, ) async for chunk in response_stream: # Check for immediate cancellation between chunks @@ -5066,8 +5106,10 @@ async def count_view( else: # Fallback to non-streaming foreground_recorder = foreground_usage_recorder() + response_ticket = None try: self._llm_calls_this_turn += 1 + response_ticket = record_dispatch(chat_request) response = await provider.complete(chat_request, **request_options) except ContextLengthError as error: recovered_request = await recover_context_overflow( @@ -5100,6 +5142,7 @@ async def count_view( ) raise self._llm_calls_this_turn += 1 + response_ticket = record_dispatch(recovered_request) response = await provider.complete( recovered_request, **request_options ) @@ -5259,6 +5302,7 @@ async def count_view( assistant_msg["metadata"] = response.metadata await context.add_message(assistant_msg) + bind_admitted_response(response_ticket) # 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. @@ -5344,29 +5388,82 @@ async def count_view( assistant_msg["metadata"] = response.metadata await context.add_message(assistant_msg) + bind_admitted_response(response_ticket) - # Process tool calls in parallel (user guidance: assume parallel intent) - # Execute tools concurrently, but add results to context sequentially for determinism + # Process tool calls in parallel by default. A provider can mark + # native-toolset calls as one ordered batch; one such marker makes + # the complete response batch sequential so response order is + # preserved even when it also contains an ordinary tool call. + # Results are always added to context in response order. import uuid parallel_group_id = str(uuid.uuid4()) - - # Execute all tools in parallel (no context updates inside). - # Materialized as Tasks rather than bare coroutines so the - # cancellation handler below can ask each one INDIVIDUALLY - # whether it finished. A completed sibling's result must never - # be overwritten by another task's cancellation. - tool_tasks = [ - asyncio.ensure_future( - self._execute_tool_only( - tc, tools, hooks, parallel_group_id, coordinator - ) - ) + sequential_batch = any( + getattr(tc, "_amplifier_execution_mode", None) == "sequential" for tc in tool_calls - ] + ) + tool_tasks: list[asyncio.Task] = [] + tool_result_errors: list[bool] | None = None try: - tool_results = await asyncio.gather(*tool_tasks) + if sequential_batch: + # Native computer members mutate one shared UI. They must + # run in exactly the provider response order; dispatching + # the later action before the prior action settles makes + # coordinate/state-dependent actions unsafe. + tool_results = [] + tool_result_errors = [] + failed_tool_call_id: str | None = None + for tc in tool_calls: + if failed_tool_call_id is not None: + tool_results.append( + ( + tc.id, + tc.name, + json.dumps( + { + "error": ( + "Skipped because a prior sequential " + "tool call failed" + ), + "skipped": True, + "failed_tool_call_id": failed_tool_call_id, + "tool": tc.name, + } + ), + ) + ) + tool_result_errors.append(True) + continue + + tool_call_id, tool_name, content, is_error = ( + await self._execute_tool_only( + tc, + tools, + hooks, + parallel_group_id, + coordinator, + include_error_status=True, + ) + ) + tool_results.append((tool_call_id, tool_name, content)) + tool_result_errors.append(is_error) + if is_error: + failed_tool_call_id = tool_call_id + else: + # Materialized as Tasks rather than bare coroutines so the + # cancellation handler below can ask each one INDIVIDUALLY + # whether it finished. A completed sibling's result must + # never be overwritten by another task's cancellation. + tool_tasks = [ + asyncio.ensure_future( + self._execute_tool_only( + tc, tools, hooks, parallel_group_id, coordinator + ) + ) + for tc in tool_calls + ] + tool_results = await asyncio.gather(*tool_tasks) except asyncio.CancelledError: # Cancellation reached this batch. Two ways in: # (a) the enclosing task was cancelled (second Ctrl+C) -- @@ -5388,32 +5485,55 @@ async def count_view( # tool_result pairing is preserved (the property the # blanket overwrite was protecting). preserved = 0 - for tc, task in zip(tool_calls, tool_tasks): - content: str | None = None - if task.done() and not task.cancelled(): - task_exc = task.exception() - if task_exc is None: - # Completed -- keep ITS OWN output verbatim. - _done_id, _done_name, content = task.result() + if sequential_batch: + # No sequential task is ever left running: the active + # await propagated cancellation and later calls have not + # started. Preserve each already-settled result, then pair + # every remaining call with a cancelled result. + for index, tc in enumerate(tool_calls): + if index < len(tool_results): + _done_id, _done_name, content = tool_results[index] preserved += 1 + is_error = bool(tool_result_errors[index]) else: - # Finished by raising something that is not a - # cancellation; report that, not "cancelled". - content = f"Internal error executing tool: {task_exc!s}" - elif not task.done(): - # Still in flight. Stop it rather than leaving it - # running unobserved past the cancelled turn. - task.cancel() - if content is None: - content = f'{{"error": "Tool execution was cancelled by user", "cancelled": true, "tool": "{tc.name}"}}' - await context.add_message( - { + content = f'{{"error": "Tool execution was cancelled by user", "cancelled": true, "tool": "{tc.name}"}}' + is_error = False + tool_message = { "role": "tool", "name": tc.name, "tool_call_id": tc.id, "content": content, } - ) + if is_error: + tool_message["is_error"] = True + await context.add_message(tool_message) + else: + for tc, task in zip(tool_calls, tool_tasks): + content: str | None = None + if task.done() and not task.cancelled(): + task_exc = task.exception() + if task_exc is None: + # Completed -- keep ITS OWN output verbatim. + _done_id, _done_name, content = task.result() + preserved += 1 + else: + # Finished by raising something that is not a + # cancellation; report that, not "cancelled". + content = f"Internal error executing tool: {task_exc!s}" + elif not task.done(): + # Still in flight. Stop it rather than leaving it + # running unobserved past the cancelled turn. + task.cancel() + if content is None: + content = f'{{"error": "Tool execution was cancelled by user", "cancelled": true, "tool": "{tc.name}"}}' + await context.add_message( + { + "role": "tool", + "name": tc.name, + "tool_call_id": tc.id, + "content": content, + } + ) logger.info( "Tool execution cancelled - preserved %d of %d completed tool result(s)", preserved, @@ -5465,15 +5585,18 @@ async def count_view( # 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, - } - ) + for index, (tool_call_id, tool_name, content) in enumerate( + tool_results + ): + tool_message = { + "role": "tool", + "name": tool_name, + "tool_call_id": tool_call_id, + "content": content, + } + if tool_result_errors is not None and tool_result_errors[index]: + tool_message["is_error"] = True + await context.add_message(tool_message) # Emit cancel:requested on first detection and trigger cleanup callbacks if not self._cancel_requested_emitted: self._cancel_requested_emitted = True @@ -5515,15 +5638,16 @@ async def count_view( # 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, - } - ) + for index, (tool_call_id, tool_name, content) in enumerate(tool_results): + tool_message = { + "role": "tool", + "name": tool_name, + "tool_call_id": tool_call_id, + "content": content, + } + if tool_result_errors is not None and tool_result_errors[index]: + tool_message["is_error"] = True + await context.add_message(tool_message) # tool:post precedes this ordered append. Only now is the # settled batch ready for a host-owned durable checkpoint. @@ -5990,7 +6114,9 @@ async def count_final_view( return foreground_recorder = foreground_usage_recorder() self._llm_calls_this_turn += 1 + response_ticket = None try: + response_ticket = record_dispatch(max_iter_chat_request) response = await provider.complete( max_iter_chat_request, **request_options ) @@ -6014,6 +6140,7 @@ async def count_final_view( if recovered_request is None: raise self._llm_calls_this_turn += 1 + response_ticket = record_dispatch(recovered_request) response = await provider.complete( recovered_request, **request_options ) @@ -6082,6 +6209,7 @@ async def count_final_view( if response_compliant and getattr(response, "metadata", None): assistant_msg["metadata"] = response.metadata await context.add_message(assistant_msg) + bind_admitted_response(response_ticket) else: await close_finalization_tool_turn( "The final response could not be generated." @@ -6162,6 +6290,8 @@ async def _stream_with_overflow_recovery( provider_name=None, recover_overflow=None, on_unconfirmed_stream=None, + record_dispatch=None, + bind_admitted_response=None, ) -> AsyncIterator[str]: """Forward a stream, retrying exactly once before its first SDK chunk.""" while True: @@ -6181,6 +6311,8 @@ def mark_provider_chunk() -> None: provider_name=provider_name, on_provider_chunk=mark_provider_chunk, on_stream_error=on_unconfirmed_stream, + record_dispatch=record_dispatch, + bind_admitted_response=bind_admitted_response, ) try: async for chunk in stream: @@ -6212,6 +6344,8 @@ async def _stream_from_provider( provider_name=None, on_provider_chunk=None, on_stream_error=None, + record_dispatch=None, + bind_admitted_response=None, ) -> AsyncIterator[str]: """Stream tokens from provider that supports streaming. @@ -6231,7 +6365,10 @@ async def _stream_from_provider( # Convert tools dict to list for provider tools_list = list(tools.values()) if tools else [] + response_ticket = None try: + if record_dispatch is not None: + response_ticket = record_dispatch(chat_request) stream_iter = provider.stream(chat_request, tools=tools_list) except ContextLengthError: if on_stream_error is not None: @@ -6269,6 +6406,8 @@ async def _stream_from_provider( await context.add_message( {"role": "assistant", "content": full_response} ) + if bind_admitted_response is not None: + bind_admitted_response(response_ticket) return # Skip non-text block deltas (e.g. thinking block streaming chunks). @@ -6298,6 +6437,8 @@ async def _stream_from_provider( # Add complete message to context if full_response: await context.add_message({"role": "assistant", "content": full_response}) + if bind_admitted_response is not None: + bind_admitted_response(response_ticket) def _extract_text_from_content(self, content) -> str: """Extract text from content blocks. @@ -6401,7 +6542,9 @@ async def _execute_tool_only( hooks: HookRegistry, parallel_group_id: str, coordinator: ModuleCoordinator | None = None, - ) -> tuple[str, str, str]: + *, + include_error_status: bool = False, + ) -> tuple[str, str, str] | tuple[str, str, str, bool]: """Execute a single tool in parallel without adding to context. Returns (tool_call_id, name, content) tuple. @@ -6423,6 +6566,12 @@ async def _execute_tool_only( was false for exactly that case, and it is what made the gather site look safe while it was overwriting completed work. """ + def outcome( + content: str, is_error: bool + ) -> tuple[str, str, str] | tuple[str, str, str, bool]: + result = (tool_call.id, tool_call.name, content) + return (*result, is_error) if include_error_status else result + try: # Pre-tool hook pre_result = await hooks.emit( @@ -6439,11 +6588,7 @@ async def _execute_tool_only( pre_result, "tool:pre", tool_call.name ) if pre_result.action == "deny": - return ( - tool_call.id, - tool_call.name, - f"Denied by hook: {pre_result.reason}", - ) + return outcome(f"Denied by hook: {pre_result.reason}", True) # Get tool tool = tools.get(tool_call.name) @@ -6458,7 +6603,7 @@ async def _execute_tool_only( "parallel_group_id": parallel_group_id, }, ) - return (tool_call.id, tool_call.name, error_msg) + return outcome(error_msg, True) # Register tool with cancellation token for visibility if coordinator: @@ -6566,7 +6711,7 @@ async def _execute_tool_only( content = str(modified_result) else: content = result.get_serialized_output() - return (tool_call.id, tool_call.name, content) + return outcome(content, not result.success) except Exception as e: # Safety net: errors become error messages @@ -6581,7 +6726,7 @@ async def _execute_tool_only( "parallel_group_id": parallel_group_id, }, ) - return (tool_call.id, tool_call.name, error_msg) + return outcome(error_msg, True) async def _execute_tool_with_result( self, diff --git a/tests/test_context_request_dispatch.py b/tests/test_context_request_dispatch.py new file mode 100644 index 0000000..198a3c2 --- /dev/null +++ b/tests/test_context_request_dispatch.py @@ -0,0 +1,342 @@ +"""Contract coverage for the paired Context final-dispatch handoff.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import pytest +from amplifier_core import ContextLengthError + +from amplifier_module_loop_streaming import StreamingOrchestrator +from tests.test_ephemeral_cache_persist_mode import ( + MockContext, + MockCoordinator, + MockResponse, + NRoundToolProvider, + OneShotTool, + RequestCapturingProvider, + ScriptedHookResult, + ScriptedHooks, +) +from tests.test_measured_compaction_runtime import _MeasuredContext, _MeasuredProvider +from tests.test_provider_budget_guard import ( + HardFitBudgetContext, + _retaining_coordinator, +) +from tests.test_provider_overflow_recovery import RecoveringProvider + + +@dataclass(frozen=True) +class OpaqueTicket: + sequence: int + + +class DispatchHandoff: + """Small synchronous Context-capability fake; tickets are deliberately opaque.""" + + def __init__(self) -> None: + self.requests: list[object] = [] + self.bound: list[OpaqueTicket] = [] + + def record_final_request(self, request) -> OpaqueTicket: + self.requests.append(request) + return OpaqueTicket(len(self.requests)) + + def bind_signed_response(self, ticket: OpaqueTicket) -> None: + self.bound.append(ticket) + + +def _register_handoff(coordinator: MockCoordinator, handoff: DispatchHandoff) -> None: + coordinator.register_capability( + "context.final_request_record", handoff.record_final_request + ) + coordinator.register_capability( + "context.signed_replay", handoff.bind_signed_response + ) + + +@pytest.mark.asyncio +async def test_records_final_tail_overlay_and_binds_after_assistant_admission() -> None: + context = MockContext() + provider = RequestCapturingProvider() + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + hooks = ScriptedHooks( + { + "provider:request": ScriptedHookResult( + action="inject_context", + ephemeral=True, + context_injection="FINAL-REQUEST-OVERLAY", + ) + } + ) + + await StreamingOrchestrator( + {"ephemeral_injection_mode": "tail", "reminder_placement": "tail"} + ).execute("work", context, {"main": provider}, {}, hooks, coordinator) + + assert handoff.requests == provider.requests + assert "FINAL-REQUEST-OVERLAY" in "\n".join( + message.content for message in handoff.requests[0].messages + ) + assert handoff.bound == [OpaqueTicket(1)] + + +@pytest.mark.asyncio +async def test_binds_tool_use_assistant_response_after_admission() -> None: + context = MockContext() + provider = NRoundToolProvider(n_tool_rounds=1) + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + coordinator, + ) + + assert handoff.requests == provider.requests + assert handoff.bound == [OpaqueTicket(1), OpaqueTicket(2)] + + +@pytest.mark.asyncio +async def test_binds_thinking_for_completed_and_tool_use_assistant_turns() -> None: + class ThinkingBlock: + type = "thinking" + + def model_dump(self): + return {"type": "thinking", "thinking": "private", "signature": "sig"} + + class ThinkingToolProvider(NRoundToolProvider): + async def complete(self, chat_request, **kwargs): + self.call_count += 1 + self.requests.append(chat_request) + response = MockResponse(text=f"round {self.call_count}") + response.content = [ThinkingBlock()] + return response + + context = MockContext() + provider = ThinkingToolProvider(n_tool_rounds=1) + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + coordinator, + ) + + assistant_messages = [ + message + for message in context.add_message_calls + if message.get("role") == "assistant" + ] + assert all("thinking_block" in message for message in assistant_messages) + assert handoff.bound == [OpaqueTicket(1), OpaqueTicket(2)] + + +@pytest.mark.asyncio +async def test_finalization_dispatch_records_and_binds_its_accepted_response() -> None: + context = MockContext() + provider = NRoundToolProvider(n_tool_rounds=1) + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + ScriptedHooks({}), + coordinator, + ) + + assert handoff.requests == provider.requests + assert handoff.requests[-1].tool_choice == "none" + assert handoff.bound == [OpaqueTicket(1), OpaqueTicket(2)] + + +@pytest.mark.asyncio +async def test_measured_dispatch_records_only_after_transaction_commit() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + committed_at_record: list[int] = [] + + def record_after_commit(request) -> OpaqueTicket: + committed_at_record.append(context.transactions[-1].committed) + return handoff.record_final_request(request) + + coordinator.register_capability("context.final_request_record", record_after_commit) + + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert committed_at_record == [1] + assert handoff.requests == provider.requests + assert handoff.bound == [OpaqueTicket(1)] + + +@pytest.mark.asyncio +async def test_measured_cancellation_before_dispatch_records_nothing() -> None: + context = _MeasuredContext() + provider = _MeasuredProvider() + coordinator = _retaining_coordinator(context) + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + class CancellingHooks(ScriptedHooks): + async def emit(self, event, payload=None): + result = await super().emit(event, payload) + if event == "orchestrator:provider_budget": + coordinator.cancellation.is_cancelled = True + return result + + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, CancellingHooks({}), coordinator + ) + + assert provider.requests == [] + assert context.transactions[0].rolled_back == 1 + assert handoff.requests == [] + assert handoff.bound == [] + + +@pytest.mark.asyncio +async def test_hard_limit_preflight_records_nothing() -> None: + class OverBudgetProvider(RequestCapturingProvider): + def request_budget(self, request, *, context_estimate, request_options=None): + return { + "estimated_input_tokens": 100, + "input_limit_tokens": 10, + "context_token_budget": 0, + } + + context = MockContext() + provider = OverBudgetProvider() + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + with pytest.raises(ContextLengthError, match="cannot retain a smaller context"): + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert provider.requests == [] + assert handoff.requests == [] + assert handoff.bound == [] + + +@pytest.mark.asyncio +async def test_overflow_retry_records_and_binds_only_corrected_request() -> None: + context = HardFitBudgetContext() + provider = RecoveringProvider() + coordinator = _retaining_coordinator(context) + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert handoff.requests == provider.requests + assert len(handoff.requests) == 2 + assert handoff.bound == [OpaqueTicket(2)] + + +@pytest.mark.asyncio +async def test_missing_either_handoff_capability_makes_no_handoff_calls() -> None: + for registered in ("record", "bind"): + context = MockContext() + provider = RequestCapturingProvider() + coordinator = MockCoordinator() + handoff = DispatchHandoff() + if registered == "record": + coordinator.register_capability( + "context.final_request_record", handoff.record_final_request + ) + else: + coordinator.register_capability( + "context.signed_replay", handoff.bind_signed_response + ) + + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert provider.requests + assert handoff.requests == [] + assert handoff.bound == [] + + +@pytest.mark.asyncio +async def test_failed_assistant_admission_does_not_bind_ticket() -> None: + class FailingAssistantContext(MockContext): + async def add_message(self, message): + if message.get("role") == "assistant": + raise RuntimeError("assistant admission failed") + await super().add_message(message) + + context = FailingAssistantContext() + provider = RequestCapturingProvider() + coordinator = MockCoordinator() + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + with pytest.raises(RuntimeError, match="assistant admission failed"): + await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert handoff.requests == provider.requests + assert handoff.bound == [] + + +@pytest.mark.asyncio +async def test_stream_overflow_retry_binds_only_retry_ticket() -> None: + class StreamRecoveringProvider(RecoveringProvider): + def stream(self, request, *, tools): + self.complete_calls += 1 + self.requests.append(request) + + async def chunks(): + if self.complete_calls == 1: + raise ContextLengthError("provider rejected input") + yield {"content": "recovered"} + + return chunks() + + context = HardFitBudgetContext() + provider = StreamRecoveringProvider() + coordinator = _retaining_coordinator(context) + handoff = DispatchHandoff() + _register_handoff(coordinator, handoff) + + result = await StreamingOrchestrator({}).execute( + "work", context, {"main": provider}, {}, ScriptedHooks({}), coordinator + ) + + assert result == "recovered" + assert handoff.requests == provider.requests + assert len(handoff.requests) == 2 + assert handoff.bound == [OpaqueTicket(2)] \ No newline at end of file diff --git a/tests/test_goal_loop.py b/tests/test_goal_loop.py index cdd72f2..96de80f 100644 --- a/tests/test_goal_loop.py +++ b/tests/test_goal_loop.py @@ -2089,10 +2089,10 @@ async def test_pref_config_forwarded_as_kwargs_extended_thinking_wins( self, ) -> None: """A resolved role's `ProviderPreference.config` (e.g. - `{"reasoning_effort": "high"}`) must be forwarded as `complete()` + `{"reasoning_effort": "high"}`) must preserve unrelated `complete()` kwargs -- this is how the delegate path applies per-role config. - `extended_thinking=False` must still win even when the role's own - config tries to turn it on (DEFECT 4 invariant). + The utility-specific effort and thinking opt-out still take + precedence. """ orch = _make_orchestrator() ctx = MockContext() @@ -2102,7 +2102,10 @@ async def test_pref_config_forwarded_as_kwargs_extended_thinking_wins( ProviderPreference( provider="main", model="routed-fast-model", - config={"reasoning_effort": "high", "extended_thinking": True}, + config={ + "reasoning_effort": "high", + "extended_thinking": True, + }, ) ] ) @@ -2122,6 +2125,7 @@ async def test_pref_config_forwarded_as_kwargs_extended_thinking_wins( kwargs = provider.eval_call_kwargs[0] assert kwargs.get("reasoning_effort") == "high" + assert kwargs.get("effort") == "high" # extended_thinking=False always wins, even though the role's own # config tried to set it True. assert kwargs.get("extended_thinking") is False @@ -3190,6 +3194,101 @@ async def test_zero_reasons_is_zero(self) -> None: @pytest.mark.asyncio class TestInternalCallsDoNotStream: + @pytest.mark.parametrize( + "role_config", + [ + { + "reasoning_effort": "xhigh", + "thinking_budget_tokens": 32_000, + "thinking_display": "enabled", + "credential": "utility-role-credential", + "source": "goal-utility-test", + }, + { + "effort": "max", + "thinking_budget_tokens": 32_000, + "thinking_display": "enabled", + "credential": "utility-role-credential", + "source": "goal-utility-test", + }, + ], + ) + async def test_fresh_utility_calls_overlay_high_effort_and_drop_thinking_config( + self, role_config: dict[str, Any] + ) -> None: + """Fresh /goal utilities must not inherit incompatible thinking config.""" + orch = _make_orchestrator() + ctx = MockContext() + hooks = MockHooks() + coordinator = MockCoordinator( + capabilities={ + "model_role_resolver": FakeModelRoleResolver( + preferences=[ + ProviderPreference( + provider="main", + model="routed-fast-model", + config=role_config, + ) + ] + ) + } + ) + provider = FakeProvider() + provider.eval_queue.append((True, "looks satisfied")) + provider.judge_queue.append((True, "same blocker again")) + + await ctx.add_message({"role": "user", "content": "do the thing"}) + await orch._evaluate_goal( + "the thing is done", + ctx, + {"main": provider}, + hooks, # type: ignore[arg-type] + coordinator, # type: ignore[arg-type] + ) + await orch._judge_stall( + { + "condition": "solved", + "turns_used": 2, + "last_reason": None, + "cap": None, + "reasons": ["blocked: x", "blocked: x"], + "no_tool_turns": 2, + }, + {"main": provider}, + hooks, # type: ignore[arg-type] + coordinator, # type: ignore[arg-type] + ) + await orch._summarize_goal_run( + { + "condition": "solved", + "turns_used": 3, + "last_reason": "still blocked", + "cap": None, + "reasons": ["blocked: x"], + "no_tool_turns": 3, + }, + {"main": provider}, + hooks, # type: ignore[arg-type] + coordinator, # type: ignore[arg-type] + final_state="stalled", + ) + + utility_calls = [ + (provider.eval_call_requests[0], provider.eval_call_kwargs[0]), + (provider.judge_call_requests[0], provider.judge_call_kwargs[0]), + (provider.summary_call_requests[0], provider.summary_call_kwargs[0]), + ] + for request, kwargs in utility_calls: + assert request.reasoning_effort == "high" + assert kwargs["effort"] == "high" + assert kwargs["extended_thinking"] is False + assert "thinking_budget_tokens" not in kwargs + assert "thinking_display" not in kwargs + assert kwargs["credential"] == "utility-role-credential" + assert kwargs["source"] == "goal-utility-test" + assert role_config["thinking_budget_tokens"] == 32_000 + assert role_config["thinking_display"] == "enabled" + async def test_evaluate_goal_sets_stream_false_and_disables_thinking( self, ) -> None: diff --git a/tests/test_sequential_tool_batches.py b/tests/test_sequential_tool_batches.py new file mode 100644 index 0000000..b2f98c0 --- /dev/null +++ b/tests/test_sequential_tool_batches.py @@ -0,0 +1,277 @@ +"""Regression tests for provider-marked sequential native-toolset batches.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from amplifier_core import ToolResult +from amplifier_core.message_models import ChatResponse, TextBlock, ToolCall +from amplifier_core.testing import EventRecorder, MockContextManager + +from amplifier_module_loop_streaming import StreamingOrchestrator + + +class _Provider: + """Return one tool response followed by a normal completion.""" + + def __init__(self, tool_calls: list[ToolCall]) -> None: + self._responses = [ + ChatResponse(content=[TextBlock(text="running tools")], tool_calls=tool_calls), + ChatResponse(content=[TextBlock(text="done")]), + ] + + async def complete(self, request, **kwargs): # noqa: ANN001, ANN201 + return self._responses.pop(0) + + def parse_tool_calls(self, response): # noqa: ANN001, ANN201 + return response.tool_calls or [] + + +class _Tool: + description = "test tool" + input_schema = {"type": "object", "properties": {}} + + def __init__(self, name: str, execute) -> None: # noqa: ANN001 + self.name = name + self._execute = execute + + async def execute(self, arguments): # noqa: ANN001, ANN201 + return await self._execute(arguments) + + +async def _run_batch( + tool_calls: list[ToolCall], tools: dict[str, _Tool] +) -> MockContextManager: + context = MockContextManager() + await StreamingOrchestrator({"stream_delay": 0}).execute( + prompt="run tools", + context=context, + providers={"default": _Provider(tool_calls)}, + tools=tools, + hooks=EventRecorder(), + ) + return context + + +def _tool_messages(context: MockContextManager) -> list[dict]: + return [message for message in context.messages if message["role"] == "tool"] + + +@pytest.mark.asyncio +async def test_marked_toolcalls_use_pydantic_extra_and_execute_in_response_order() -> None: + """A core ToolCall extra makes the whole response batch ordered.""" + + events: list[str] = [] + + async def first(arguments): # noqa: ANN001 + events.append("first:start") + await asyncio.sleep(0) + events.append("first:finish") + return ToolResult(success=True, output="first complete") + + async def second(arguments): # noqa: ANN001 + # This assertion fails if the loop starts this action before the first + # native action has settled. + assert events == ["first:start", "first:finish"] + events.append("second:start") + return ToolResult(success=True, output="second complete") + + calls = [ + ToolCall( + id="call-first", + name="first", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ToolCall( + id="call-second", + name="second", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ] + + assert getattr(calls[0], "_amplifier_execution_mode") == "sequential" + context = await _run_batch( + calls, + { + "first": _Tool("first", first), + "second": _Tool("second", second), + }, + ) + + assert events == ["first:start", "first:finish", "second:start"] + assert [message["tool_call_id"] for message in _tool_messages(context)] == [ + "call-first", + "call-second", + ] + + +@pytest.mark.asyncio +async def test_sequential_failure_marks_failure_and_skips_later_actions() -> None: + """A failed native action stops the batch and preserves all result pairs.""" + + executed: list[str] = [] + + async def succeeds(arguments): # noqa: ANN001 + executed.append("succeeds") + return ToolResult(success=True, output="ok") + + async def fails(arguments): # noqa: ANN001 + executed.append("fails") + return ToolResult(success=False, error={"message": "native action failed"}) + + async def must_not_run(arguments): # noqa: ANN001 + executed.append("must_not_run") + return ToolResult(success=True, output="unexpected") + + calls = [ + ToolCall( + id="call-ok", + name="succeeds", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ToolCall( + id="call-fail", + name="fails", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ToolCall( + id="call-skipped", + name="must_not_run", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ] + context = await _run_batch( + calls, + { + "succeeds": _Tool("succeeds", succeeds), + "fails": _Tool("fails", fails), + "must_not_run": _Tool("must_not_run", must_not_run), + }, + ) + + messages = _tool_messages(context) + assert executed == ["succeeds", "fails"] + assert [message["tool_call_id"] for message in messages] == [ + "call-ok", + "call-fail", + "call-skipped", + ] + assert "is_error" not in messages[0] + assert messages[1]["is_error"] is True + assert messages[2]["is_error"] is True + assert "native action failed" in messages[1]["content"] + assert "Skipped because a prior sequential tool call failed" in messages[2]["content"] + assert '"failed_tool_call_id": "call-fail"' in messages[2]["content"] + + +@pytest.mark.asyncio +async def test_sequential_cancellation_preserves_settled_actions_and_pairs_rest() -> None: + """Cancellation keeps completed work and pairs active/unstarted calls.""" + + second_started = asyncio.Event() + third_executed = False + + async def first(arguments): # noqa: ANN001 + return ToolResult(success=True, output="first complete") + + async def waits_for_cancellation(arguments): # noqa: ANN001 + second_started.set() + await asyncio.Event().wait() + + async def must_not_run(arguments): # noqa: ANN001 + nonlocal third_executed + third_executed = True + return ToolResult(success=True, output="unexpected") + + calls = [ + ToolCall( + id="call-first", + name="first", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ToolCall( + id="call-active", + name="waits", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ToolCall( + id="call-never-started", + name="third", + arguments={}, + _amplifier_execution_mode="sequential", + ), + ] + context = MockContextManager() + task = asyncio.create_task( + StreamingOrchestrator({"stream_delay": 0}).execute( + prompt="run tools", + context=context, + providers={"default": _Provider(calls)}, + tools={ + "first": _Tool("first", first), + "waits": _Tool("waits", waits_for_cancellation), + "third": _Tool("third", must_not_run), + }, + hooks=EventRecorder(), + ) + ) + await asyncio.wait_for(second_started.wait(), timeout=1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + messages = _tool_messages(context) + assert third_executed is False + assert [message["tool_call_id"] for message in messages] == [ + "call-first", + "call-active", + "call-never-started", + ] + assert messages[0]["content"] == "first complete" + assert '"cancelled": true' in messages[1]["content"] + assert '"cancelled": true' in messages[2]["content"] + + +@pytest.mark.asyncio +async def test_unmarked_batch_retains_concurrent_execution() -> None: + """The marker is opt-in; ordinary batches still start all tools together.""" + + started: set[str] = set() + both_started = asyncio.Event() + + async def concurrent_tool(name: str, arguments): # noqa: ANN001 + started.add(name) + if len(started) == 2: + both_started.set() + await both_started.wait() + return ToolResult(success=True, output=name) + + calls = [ + ToolCall(id="call-a", name="a", arguments={}), + ToolCall(id="call-b", name="b", arguments={}), + ] + context = await asyncio.wait_for( + _run_batch( + calls, + { + "a": _Tool("a", lambda arguments: concurrent_tool("a", arguments)), + "b": _Tool("b", lambda arguments: concurrent_tool("b", arguments)), + }, + ), + timeout=1, + ) + + assert started == {"a", "b"} + assert [message["tool_call_id"] for message in _tool_messages(context)] == [ + "call-a", + "call-b", + ] \ No newline at end of file