From 24ce7501c6f23967e1228958c19c3a0779de75fb Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 29 Sep 2026 16:25:32 -0700 Subject: [PATCH 1/3] feat: execute marked native tool batches sequentially Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- README.md | 5 +- amplifier_module_loop_streaming/__init__.py | 207 ++++++++++----- tests/test_sequential_tool_batches.py | 277 ++++++++++++++++++++ 3 files changed, 426 insertions(+), 63 deletions(-) create mode 100644 tests/test_sequential_tool_batches.py 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..20a34b8 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -5345,28 +5345,80 @@ async def count_view( 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 + # 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 +5440,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 +5540,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 +5593,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. @@ -6401,7 +6480,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 +6504,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 +6526,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 +6541,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 +6649,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 +6664,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_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 From 8ce68dd8e2480f2804365ac9e31ef9c44e308c88 Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Tue, 29 Sep 2026 22:18:13 -0700 Subject: [PATCH 2/3] fix: isolate goal utility reasoning settings Fresh internal goal judge, evaluator, and summary calls now request high reasoning effort and explicitly disable extended thinking without inheriting incompatible thinking budget or display settings. Main conversational requests remain unchanged. Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_loop_streaming/__init__.py | 31 ++++-- tests/test_goal_loop.py | 107 +++++++++++++++++++- 2 files changed, 123 insertions(+), 15 deletions(-) diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 20a34b8..cea2b7f 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) 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: From 4d8e4106c9f30e2f7dea53414ad8adb72cca345b Mon Sep 17 00:00:00 2001 From: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> Date: Wed, 30 Sep 2026 09:39:10 -0700 Subject: [PATCH 3/3] fix: record final context request dispatch Generated with Amplifier Co-Authored-By: Amplifier <240397093+microsoft-amplifier@users.noreply.github.com> --- amplifier_module_loop_streaming/__init__.py | 53 +++ tests/test_context_request_dispatch.py | 342 ++++++++++++++++++++ 2 files changed, 395 insertions(+) create mode 100644 tests/test_context_request_dispatch.py diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index cea2b7f..de91950 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -3365,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): @@ -3372,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 @@ -5031,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( @@ -5042,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 @@ -5075,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( @@ -5109,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 ) @@ -5268,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. @@ -5353,6 +5388,7 @@ 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 by default. A provider can mark # native-toolset calls as one ordered batch; one such marker makes @@ -6078,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 ) @@ -6102,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 ) @@ -6170,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." @@ -6250,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: @@ -6269,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: @@ -6300,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. @@ -6319,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: @@ -6357,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). @@ -6386,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. 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