From 7aa97237eaaca18844ea66a204f736fdde8f5ba3 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Tue, 29 Sep 2026 14:58:34 +0200 Subject: [PATCH 1/3] Python: Add gates and buffering to ResponseStream --- .../core/agent_framework/_agent_hooks.py | 87 ++- .../packages/core/agent_framework/_agents.py | 2 +- .../core/agent_framework/_middleware.py | 75 +- .../packages/core/agent_framework/_types.py | 711 ++++++++++++------ .../core/tests/core/test_agent_hooks.py | 36 +- python/packages/core/tests/core/test_types.py | 366 +++++++++ python/samples/02-agents/middleware/README.md | 4 +- .../override_result_with_middleware.py | 144 ++-- .../middleware/usage_tracking_middleware.py | 4 +- python/samples/02-agents/response_stream.py | 249 ++++-- 10 files changed, 1244 insertions(+), 434 deletions(-) diff --git a/python/packages/core/agent_framework/_agent_hooks.py b/python/packages/core/agent_framework/_agent_hooks.py index bf2f2aaaa14..a547d147b76 100644 --- a/python/packages/core/agent_framework/_agent_hooks.py +++ b/python/packages/core/agent_framework/_agent_hooks.py @@ -69,9 +69,8 @@ outside the tool seam falls back to deferring both — fail-closed, matching the pre-ownership behavior. -Streaming is supported **fail-closed by buffering** via -:meth:`ResponseStream.buffered_and_gated`: the model/agent stream is fully consumed -internally, middleware stream hooks are applied to the buffered content, the +Streaming is supported **fail-closed by buffering**: the model/agent stream is fully +consumed internally, middleware stream hooks are applied to the buffered content, the ``post_model_call`` / ``output`` verdict is applied to the finalized result, and only then are the (possibly transformed) updates released to the consumer. No partial content ever egresses ahead of a verdict (spec §12.1/§12.1a ``buffered_output: true`` @@ -103,7 +102,7 @@ import json import logging import uuid -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence from contextvars import ContextVar from dataclasses import dataclass from typing import TYPE_CHECKING, Any, NoReturn, cast @@ -1231,7 +1230,10 @@ async def _process_streaming( f"got {type(inner).__name__}." ) context.result = self._gated_agent_stream( - state, cast("ResponseStream[AgentResponseUpdate, AgentResponse[Any]]", inner), gate_handle + context, + state, + cast("ResponseStream[AgentResponseUpdate, AgentResponse[Any]]", inner), + gate_handle, ) # The gated stream now owns the shutdown emission. shutdown_reason = None @@ -1255,6 +1257,7 @@ async def _process_streaming( def _gated_agent_stream( self, + context: AgentContext, state: _RunState, inner: ResponseStream[AgentResponseUpdate, AgentResponse[Any]], gate_handle: _RunPersistenceGate, @@ -1262,18 +1265,19 @@ def _gated_agent_stream( """Guard a streaming run with a fail-closed buffered gate. The run is fully consumed (with the run state and the persistence gate active), - middleware stream hooks are applied by the combinator before the verdict, the + middleware stream hooks are applied by the buffered enforcement helper before the verdict, the ``output`` verdict is applied to the finalized response, and deferred - persistence is released only after the verdict permits. The combinator owns the + persistence is released only after the verdict permits. The helper owns the no-divergence rule: a transformed output (or any applied stream hooks) re-derives the released updates from the verdicted response. """ - async def _consume() -> tuple[Sequence[AgentResponseUpdate], AgentResponse[Any]]: + @contextlib.asynccontextmanager + async def _consumption_scope() -> AsyncGenerator[None]: run_token = _RUN_STATE.set(state) try: with gate_handle: - final = await inner.get_final_response() + yield except asyncio.CancelledError: await self._emit_shutdown(state, "cancelled") raise @@ -1290,11 +1294,8 @@ async def _consume() -> tuple[Sequence[AgentResponseUpdate], AgentResponse[Any]] raise finally: _RUN_STATE.reset(run_token) - return list(inner.updates), final - async def _gate( - _updates: list[AgentResponseUpdate], final: AgentResponse[Any] - ) -> tuple[AgentResponse[Any], bool]: + async def _transform(final: AgentResponse[Any]) -> AgentResponse[Any] | None: from agent_hooks import InterceptionBlocked try: @@ -1304,9 +1305,7 @@ async def _gate( "the output interception point was not emitted." ) transformed = await self._emit_output(state, final) - await gate_handle.flush() - await self._emit_shutdown(state, "completed") - return final, transformed + return final if transformed else None except InterceptionBlocked: gate_handle.drop() await self._emit_shutdown(state, "error") @@ -1318,12 +1317,23 @@ async def _gate( await self._emit_shutdown(state, "error") raise - return cast( - "ResponseStream[AgentResponseUpdate, AgentResponse[Any]]", - cast(Any, ResponseStream).buffered_and_gated( - consume=_consume, gate=_gate, rederive=_agent_updates_from_response - ), - ) + async def _release(_: AgentResponse[Any]) -> None: + try: + await gate_handle.flush() + await self._emit_shutdown(state, "completed") + except asyncio.CancelledError: + await self._emit_shutdown(state, "cancelled") + raise + except BaseException: + await self._emit_shutdown(state, "error") + raise + + context.stream_result_transforms.append(_transform) + context.stream_result_gates_after.append(_release) + context.stream_buffer_updates = True + context.stream_result_to_updates = _agent_updates_from_response + context.stream_consumption_context_manager_factories.append(_consumption_scope) + return inner class _AgentHooksChatMiddleware(_AgentHooksMiddlewareBase, ChatMiddleware): @@ -1367,7 +1377,11 @@ async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[ result = context.result if isinstance(result, ResponseStream): context.result = self._gated_chat_stream( - state, model_id, cast("ResponseStream[ChatResponseUpdate, ChatResponse[Any]]", result), gate + context, + state, + model_id, + cast("ResponseStream[ChatResponseUpdate, ChatResponse[Any]]", result), + gate, ) elif isinstance(result, ChatResponse): try: @@ -1406,6 +1420,7 @@ async def _emit_post_model_call(self, state: _RunState, model_id: str, response: def _gated_chat_stream( self, + context: ChatContext, state: _RunState, model_id: str, inner: ResponseStream[ChatResponseUpdate, ChatResponse[Any]], @@ -1417,16 +1432,16 @@ def _gated_chat_stream( emitted, and nothing (updates or tool calls) is released beforehand. A deny raises before any update egresses, and per-service-call history persistence deferred by the run persistence gate is released only after the verdict - permits. The combinator owns the no-divergence rule: a transformed response + permits. The buffered enforcement helper owns the no-divergence rule: a transformed response (or any applied stream hooks) re-derives the released updates from it. """ - async def _consume() -> tuple[Sequence[ChatResponseUpdate], ChatResponse[Any]]: + @contextlib.asynccontextmanager + async def _consumption_scope() -> AsyncGenerator[None]: with gate_handle: - response = await inner.get_final_response() - return list(inner.updates), response + yield - async def _gate(_updates: list[ChatResponseUpdate], final: ChatResponse[Any]) -> tuple[ChatResponse[Any], bool]: + async def _transform(final: ChatResponse[Any]) -> ChatResponse[Any] | None: from agent_hooks import InterceptionBlocked if not isinstance(final, ChatResponse): @@ -1436,20 +1451,22 @@ async def _gate(_updates: list[ChatResponseUpdate], final: ChatResponse[Any]) -> ) try: changed = await self._emit_post_model_call(state, model_id, final) + return final if changed else None except InterceptionBlocked: # §6.1: the deferred per-service-call persistence for the denied # response is dropped, never executed. gate_handle.drop() raise + + async def _release(_: ChatResponse[Any]) -> None: await gate_handle.flush() - return final, changed - return cast( - "ResponseStream[ChatResponseUpdate, ChatResponse[Any]]", - cast(Any, ResponseStream).buffered_and_gated( - consume=_consume, gate=_gate, rederive=_chat_updates_from_response - ), - ) + context.stream_result_transforms.append(_transform) + context.stream_result_gates_after.append(_release) + context.stream_buffer_updates = True + context.stream_result_to_updates = _chat_updates_from_response + context.stream_consumption_context_manager_factories.append(_consumption_scope) + return inner class _AgentHooksFunctionMiddleware(_AgentHooksMiddlewareBase, FunctionMiddleware): diff --git a/python/packages/core/agent_framework/_agents.py b/python/packages/core/agent_framework/_agents.py index 11e317dcb80..db264535aa1 100644 --- a/python/packages/core/agent_framework/_agents.py +++ b/python/packages/core/agent_framework/_agents.py @@ -787,7 +787,7 @@ async def _agent_wrapper(ctx: FunctionInvocationContext, **kwargs: Any) -> str: # The callback is a host-facing observer: feed it the *released* # updates by consuming the stream, never by registering a transform # hook on it. Hooks can end up applied to buffered content ahead of an - # egress gate's verdict (see ResponseStream.buffered_and_gated), so a + # egress gate's verdict, so a # hook-registered observer could see denied or unredacted content. async for update in stream: callback_result = stream_callback(update) diff --git a/python/packages/core/agent_framework/_middleware.py b/python/packages/core/agent_framework/_middleware.py index ee307e7c632..1b473b9a515 100644 --- a/python/packages/core/agent_framework/_middleware.py +++ b/python/packages/core/agent_framework/_middleware.py @@ -285,7 +285,12 @@ def __init__( Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]] ] | None = None, - stream_result_hooks: Sequence[Callable[[AgentResponse], AgentResponse | Awaitable[AgentResponse]]] + stream_result_hooks: Sequence[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ] | None = None, stream_cleanup_hooks: Sequence[Callable[[], Awaitable[None] | None]] | None = None, ) -> None: @@ -324,9 +329,19 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self.stream_transform_hooks = list(stream_transform_hooks or []) - self.stream_result_hooks = list(stream_result_hooks or []) + self.stream_update_transforms = list(stream_transform_hooks or []) + self.stream_result_transforms = list(stream_result_hooks or []) + # Compatibility aliases for the original middleware context field names. + self.stream_transform_hooks = self.stream_update_transforms + self.stream_result_hooks = self.stream_result_transforms self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) + self.stream_update_gates_before: list[Callable[[AgentResponseUpdate], object]] = [] + self.stream_update_gates_after: list[Callable[[AgentResponseUpdate], object]] = [] + self.stream_result_gates_before: list[Callable[[AgentResponse[Any]], object]] = [] + self.stream_result_gates_after: list[Callable[[AgentResponse[Any]], object]] = [] + self.stream_buffer_updates = False + self.stream_result_to_updates: Callable[[AgentResponse[Any]], Sequence[AgentResponseUpdate]] | None = None + self.stream_consumption_context_manager_factories: list[Callable[[], Any]] = [] # Set by egress-enforcement middleware (agent-hooks): the run-persistence gate # covering this pipeline's run. The final handler offers it for adoption by # the run it starts (see _sessions._offer_run_persistence_gate_claim), so the @@ -618,7 +633,13 @@ def __init__( Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]] ] | None = None, - stream_result_hooks: Sequence[Callable[[ChatResponse], ChatResponse | Awaitable[ChatResponse]]] | None = None, + stream_result_hooks: Sequence[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ] + | None = None, stream_cleanup_hooks: Sequence[Callable[[], Awaitable[None] | None]] | None = None, ) -> None: """Initialize the ChatContext. @@ -648,9 +669,19 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self.stream_transform_hooks = list(stream_transform_hooks or []) - self.stream_result_hooks = list(stream_result_hooks or []) + self.stream_update_transforms = list(stream_transform_hooks or []) + self.stream_result_transforms = list(stream_result_hooks or []) + # Compatibility aliases for the original middleware context field names. + self.stream_transform_hooks = self.stream_update_transforms + self.stream_result_hooks = self.stream_result_transforms self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) + self.stream_update_gates_before: list[Callable[[ChatResponseUpdate], object]] = [] + self.stream_update_gates_after: list[Callable[[ChatResponseUpdate], object]] = [] + self.stream_result_gates_before: list[Callable[[ChatResponse[Any]], object]] = [] + self.stream_result_gates_after: list[Callable[[ChatResponse[Any]], object]] = [] + self.stream_buffer_updates = False + self.stream_result_to_updates: Callable[[ChatResponse[Any]], Sequence[ChatResponseUpdate]] | None = None + self.stream_consumption_context_manager_factories: list[Callable[[], Any]] = [] self._message_replacements: list[tuple[Message, tuple[Message, ...]]] = [] self._fallback_reconciliation_messages: list[Message] | None = None @@ -1299,10 +1330,22 @@ async def current_handler() -> None: await first_handler() if context.result and isinstance(context.result, ResponseStream): - for hook in context.stream_transform_hooks: + if context.stream_buffer_updates: + context.result.buffer_updates(result_to_updates=context.stream_result_to_updates) + for factory in context.stream_consumption_context_manager_factories: + context.result.with_consumption_context_manager(factory) + for hook in context.stream_update_transforms: context.result.with_transform_hook(hook) - for result_hook in context.stream_result_hooks: + for result_hook in context.stream_result_transforms: context.result.with_result_hook(result_hook) + for gate in context.stream_update_gates_before: + context.result.with_update_gate(gate, phase="before_transform") + for gate in context.stream_update_gates_after: + context.result.with_update_gate(gate, phase="after_transform") + for gate in context.stream_result_gates_before: + context.result.with_result_gate(gate, phase="before_transform") + for gate in context.stream_result_gates_after: + context.result.with_result_gate(gate, phase="after_transform") for cleanup_hook in context.stream_cleanup_hooks: context.result.with_cleanup_hook(cleanup_hook) return context.result @@ -1483,10 +1526,22 @@ async def current_handler() -> None: await first_handler() if context.result and isinstance(context.result, ResponseStream): - for hook in context.stream_transform_hooks: + if context.stream_buffer_updates: + context.result.buffer_updates(result_to_updates=context.stream_result_to_updates) + for factory in context.stream_consumption_context_manager_factories: + context.result.with_consumption_context_manager(factory) + for hook in context.stream_update_transforms: context.result.with_transform_hook(hook) - for result_hook in context.stream_result_hooks: + for result_hook in context.stream_result_transforms: context.result.with_result_hook(result_hook) + for gate in context.stream_update_gates_before: + context.result.with_update_gate(gate, phase="before_transform") + for gate in context.stream_update_gates_after: + context.result.with_update_gate(gate, phase="after_transform") + for gate in context.stream_result_gates_before: + context.result.with_result_gate(gate, phase="before_transform") + for gate in context.stream_result_gates_after: + context.result.with_result_gate(gate, phase="after_transform") for cleanup_hook in context.stream_cleanup_hooks: context.result.with_cleanup_hook(cleanup_hook) return context.result diff --git a/python/packages/core/agent_framework/_types.py b/python/packages/core/agent_framework/_types.py index cc2c948e707..b4609360b9e 100644 --- a/python/packages/core/agent_framework/_types.py +++ b/python/packages/core/agent_framework/_types.py @@ -3365,9 +3365,23 @@ def __init__( stream: AsyncIterable[UpdateT] | Awaitable[AsyncIterable[UpdateT]], *, finalizer: Callable[[Sequence[UpdateT]], FinalT | Awaitable[FinalT]] | None = None, - transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] | None = None, - cleanup_hooks: list[Callable[[], Awaitable[None] | None]] | None = None, - result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] | None = None, + transform_hooks: Sequence[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] | None = None, + cleanup_hooks: Sequence[Callable[[], Awaitable[None] | None]] | None = None, + result_hooks: Sequence[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] | None = None, + update_gates: Mapping[ + Literal["before_transform", "after_transform"], + Sequence[Callable[[UpdateT], object]], + ] + | None = None, + result_gates: Mapping[ + Literal["before_transform", "after_transform"], + Sequence[Callable[[FinalT], object]], + ] + | None = None, + update_transforms: Sequence[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] | None = None, + result_transforms: Sequence[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] | None = None, + stream_updates: bool = True, + result_to_updates: Callable[[FinalT], Sequence[UpdateT]] | None = None, ) -> None: """A Async Iterable stream of updates. @@ -3379,27 +3393,79 @@ def __init__( transform_hooks: Optional list of callables that transform each update as it is yielded. cleanup_hooks: Optional list of callables that run after the stream is fully consumed (before finalizer). result_hooks: Optional list of callables that transform the final result (after finalizer). + update_gates: Optional update gates grouped by whether they run before or after update transforms. + Gates return ``None`` to permit the update and raise to block it. + result_gates: Optional final-result gates grouped by whether they run before or after result transforms. + Gates return ``None`` to permit the result and raise to block it. + update_transforms: Preferred alias for ``transform_hooks``. Both cannot be supplied. + result_transforms: Preferred alias for ``result_hooks``. Both cannot be supplied. + stream_updates: Whether processed updates are yielded as they arrive. When ``False``, updates are held + until the stream has finalized and every gate has passed. + result_to_updates: Rebuilds buffered updates from a transformed final result. Only valid when + ``stream_updates`` is ``False`` and required when a result transform returns a replacement. """ + if transform_hooks is not None and update_transforms is not None: + raise ValueError("Cannot specify both transform_hooks and update_transforms.") + if result_hooks is not None and result_transforms is not None: + raise ValueError("Cannot specify both result_hooks and result_transforms.") + if stream_updates and result_to_updates is not None: + raise ValueError("result_to_updates is only valid when stream_updates is False.") + + update_gate_keys = set(update_gates or ()) + unknown_update_gate_keys = update_gate_keys - {"before_transform", "after_transform"} + if unknown_update_gate_keys: + formatted = ", ".join(sorted(unknown_update_gate_keys)) + raise ValueError(f"Unknown update_gates phase(s): {formatted}.") + result_gate_keys = set(result_gates or ()) + unknown_result_gate_keys = result_gate_keys - {"before_transform", "after_transform"} + if unknown_result_gate_keys: + formatted = ", ".join(sorted(unknown_result_gate_keys)) + raise ValueError(f"Unknown result_gates phase(s): {formatted}.") + self._stream_source = stream self._finalizer = finalizer self._stream: AsyncIterable[UpdateT] | None = None self._iterator: AsyncIterator[UpdateT] | None = None self._updates: list[UpdateT] = [] self._consumed: bool = False + self._result_prepared: bool = False self._finalized: bool = False self._final_result: FinalT | None = None - self._transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] = ( - transform_hooks if transform_hooks is not None else [] + self._transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] = list( + update_transforms if update_transforms is not None else transform_hooks or () + ) + self._result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] = list( + result_transforms if result_transforms is not None else result_hooks or () + ) + self._update_gates_before: list[Callable[[UpdateT], object]] = list( + (update_gates or {}).get("before_transform", ()) + ) + self._update_gates_after: list[Callable[[UpdateT], object]] = list( + (update_gates or {}).get("after_transform", ()) ) - self._result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] = ( - result_hooks if result_hooks is not None else [] + self._result_gates_before: list[Callable[[FinalT], object]] = list( + (result_gates or {}).get("before_transform", ()) ) - self._cleanup_hooks: list[Callable[[], Awaitable[None] | None]] = ( - cleanup_hooks if cleanup_hooks is not None else [] + self._result_gates_after: list[Callable[[FinalT], object]] = list( + (result_gates or {}).get("after_transform", ()) ) + self._cleanup_hooks: list[Callable[[], Awaitable[None] | None]] = list(cleanup_hooks or ()) self._cleanup_run: bool = False self._stream_error: Exception | None = None + self._stream_updates = stream_updates + self._result_to_updates = result_to_updates + self._result_was_transformed = False + self._released_updates: list[UpdateT] = [] + self._buffered_candidate_updates: list[UpdateT] = [] + self._buffered_output_updates: list[UpdateT] = [] + self._buffered_output_index = 0 + self._buffered_materialized = False + self._buffered_materializing = False + self._buffered_materialization_error: Exception | None = None + self._return_final_after_buffer_release_error = False + self._content_pipeline_started = False + self._content_hooks_sealed = False self._inner_stream: ResponseStream[Any, Any] | None = None self._inner_stream_source: ResponseStream[Any, Any] | Awaitable[ResponseStream[Any, Any]] | None = None self._wrap_inner: bool = False @@ -3407,6 +3473,12 @@ def __init__( self._flat_map_update: Callable[[Any], Iterable[UpdateT] | Awaitable[Iterable[UpdateT]]] | None = None self._pending_mapped_updates: list[UpdateT] = [] self._pull_context_manager_factories: list[Callable[[], contextlib.AbstractContextManager[Any]]] = [] + self._consumption_context_manager_factories: list[ + Callable[ + [], + contextlib.AbstractContextManager[Any] | contextlib.AbstractAsyncContextManager[Any], + ] + ] = [] def map( self, @@ -3532,11 +3604,15 @@ def buffered_and_gated( gate: Callable[[list[UpdateT], FinalT], Awaitable[tuple[FinalT, bool]]], rederive: Callable[[FinalT], Sequence[UpdateT]], ) -> ResponseStream[UpdateT, FinalT]: - """Create a fully buffered stream whose content is finalized by a gate. + """Create a fully buffered stream whose content is finalized by a legacy gate. + + .. deprecated:: + Configure ``update_gates``, ``result_gates``, transforms, and + ``stream_updates=False`` on :class:`ResponseStream` instead. - This combinator exists for egress-gating middleware (e.g. policy enforcement) - that must apply a verdict to a run's *complete* content before anything is - released, and it states the hook-ordering contract in one place: + This compatibility adapter preserves the previous combined gate/transform + contract. New gates should only return ``None`` or raise; content replacement + belongs in transforms. 1. On the first pull, ``consume`` runs and produces the buffered updates and the finalized result. Nothing has egressed yet. @@ -3573,10 +3649,54 @@ def buffered_and_gated( Returns: A sealed, fully buffered ResponseStream. """ - return cast( - "ResponseStream[UpdateT, FinalT]", - cast(Any, _GatedResponseStream).create_buffered_and_gated(consume, gate, rederive), + warnings.warn( + "ResponseStream.buffered_and_gated() is deprecated; configure update_gates, result_gates, transforms, " + "and stream_updates=False on ResponseStream instead.", + DeprecationWarning, + stacklevel=2, ) + return cls._buffered_and_gated(consume, gate, rederive) + + @classmethod + def _buffered_and_gated( + cls, + consume: Callable[[], Awaitable[tuple[Sequence[UpdateT], FinalT]]], + gate: Callable[[list[UpdateT], FinalT], Awaitable[tuple[FinalT, bool]]], + rederive: Callable[[FinalT], Sequence[UpdateT]], + ) -> ResponseStream[UpdateT, FinalT]: + """Translate the legacy combined gate contract onto a normal ResponseStream.""" + holder: dict[str, Any] = {} + + async def _source() -> AsyncGenerator[UpdateT]: + stream = cast("ResponseStream[UpdateT, FinalT]", holder["stream"]) + hooks_applied = bool(stream._transform_hooks or stream._result_hooks) + + async def _legacy_gate_transform(final: FinalT) -> FinalT | None: + gated_final, gate_transformed = await gate(list(stream._buffered_candidate_updates), final) + if hooks_applied or gate_transformed: + stream._result_to_updates = rederive + return gated_final + return None + + # This runs after every transform registered before the first pull, + # matching the legacy hook-before-gate ordering. + stream._result_hooks.append(_legacy_gate_transform) + updates, final = await consume() + holder["final"] = final + for update in updates: + yield update + + def _finalizer(_: Sequence[UpdateT]) -> FinalT: + return cast(FinalT, holder["final"]) + + stream: ResponseStream[UpdateT, FinalT] = cls( + _source(), + finalizer=_finalizer, + stream_updates=False, + ) + stream._return_final_after_buffer_release_error = True + holder["stream"] = stream + return stream async def _get_stream(self) -> AsyncIterable[UpdateT]: if self._stream is None: @@ -3595,57 +3715,243 @@ async def _get_stream(self) -> AsyncIterable[UpdateT]: def __aiter__(self) -> ResponseStream[UpdateT, FinalT]: return self - async def _record_update(self, update: UpdateT) -> UpdateT: + def _start_content_pipeline(self) -> None: + if self._content_pipeline_started: + return + self._content_pipeline_started = True + if ( + not self._stream_updates + or self._update_gates_before + or self._update_gates_after + or self._result_gates_before + or self._result_gates_after + ): + self._content_hooks_sealed = True + + async def _run_gates( + self, + gates: Sequence[Callable[[Any], object]], + value: Any, + *, + target: str, + ) -> None: + for gate in gates: + gate_result = gate(value) + if isawaitable(gate_result): + gate_result = await gate_result + if gate_result is not None: + raise TypeError( + f"ResponseStream {target} gates must return None or raise; use a {target} transform " + "to replace content." + ) + + async def _apply_transforms( + self, + transforms: Sequence[Callable[[Any], Any | Awaitable[Any | None] | None]], + value: Any, + ) -> tuple[Any, bool]: + transformed = False + for transform in transforms: + transformed_value = transform(value) + if isawaitable(transformed_value): + transformed_value = await transformed_value + if transformed_value is not None: + value = transformed_value + transformed = True + return value, transformed + + async def _record_update(self, update: UpdateT, *, run_after_gates: bool) -> UpdateT: + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._update_gates_before), + update, + target="update", + ) self._updates.append(update) - for hook in self._transform_hooks: - hooked = hook(update) - if isawaitable(hooked): - hooked = await hooked - if hooked is not None: - update = cast(UpdateT, hooked) + update, _ = await self._apply_transforms( + cast(Sequence[Callable[[Any], Any | Awaitable[Any | None] | None]], self._transform_hooks), + update, + ) + if run_after_gates: + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._update_gates_after), + update, + target="update", + ) return update - async def __anext__(self) -> UpdateT: + async def _pull_next_update(self, *, run_after_gates: bool) -> UpdateT: while True: - try: - if self._pending_mapped_updates: - return await self._record_update(self._pending_mapped_updates.pop(0)) - - with contextlib.ExitStack() as stack: - for factory in self._pull_context_manager_factories: - stack.enter_context(factory()) - # Resolve the underlying stream inside the pull contexts so that any - # spans/contexts created during stream resolution (e.g. inner chat - # completion spans created on the first pull of a wrapped agent stream) - # inherit the active context (e.g. an outer agent invoke span). - if self._iterator is None: - stream = await self._get_stream() - self._iterator = stream.__aiter__() - update: UpdateT = await self._iterator.__anext__() - - if self._flat_map_update is not None: - mapped_updates = self._flat_map_update(update) - if isawaitable(mapped_updates): - mapped_updates = await mapped_updates - self._pending_mapped_updates.extend(mapped_updates) - continue - if self._map_update is not None: - update = self._map_update(update) # type: ignore[assignment] - if isawaitable(update): - update = await update - return await self._record_update(update) - except StopAsyncIteration: - self._consumed = True - await self._run_cleanup_hooks() - await self.get_final_response() - raise - except Exception as exc: - self._stream_error = exc - try: + if self._pending_mapped_updates: + return await self._record_update( + self._pending_mapped_updates.pop(0), + run_after_gates=run_after_gates, + ) + + with contextlib.ExitStack() as stack: + for factory in self._pull_context_manager_factories: + stack.enter_context(factory()) + # Resolve the underlying stream inside the pull contexts so that any + # spans/contexts created during stream resolution (e.g. inner chat + # completion spans created on the first pull of a wrapped agent stream) + # inherit the active context (e.g. an outer agent invoke span). + if self._iterator is None: + stream = await self._get_stream() + self._iterator = stream.__aiter__() + update: UpdateT = await self._iterator.__anext__() + if self._flat_map_update is not None: + mapped_updates = self._flat_map_update(update) + if isawaitable(mapped_updates): + mapped_updates = await mapped_updates + self._pending_mapped_updates.extend(mapped_updates) + continue + if self._map_update is not None: + update = self._map_update(update) # type: ignore[assignment] + if isawaitable(update): + update = await update + return await self._record_update(update, run_after_gates=run_after_gates) + + async def _handle_stream_error(self, exc: Exception) -> None: + self._stream_error = exc + try: + await self._run_cleanup_hooks() + finally: + self._stream_error = None + + async def _prepare_final_result(self) -> None: + if self._result_prepared: + return + + inner_result: Any = None + if self._wrap_inner: + if self._inner_stream is None: + await self._resolve_stream_with_pull_contexts() + if self._inner_stream is None: + raise RuntimeError("Inner stream not available") + inner_result = await self._inner_stream.get_final_response() + + result: Any + if self._finalizer is not None: + result = self._finalizer(self._updates) + if isawaitable(result): + result = await result + elif self._wrap_inner: + result = inner_result + else: + result = list(self._updates) + + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._result_gates_before), + result, + target="result", + ) + result, self._result_was_transformed = await self._apply_transforms( + cast(Sequence[Callable[[Any], Any | Awaitable[Any | None] | None]], self._result_hooks), + result, + ) + self._final_result = result + self._result_prepared = True + + async def _complete_final_result(self) -> None: + if self._finalized: + return + await self._prepare_final_result() + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._result_gates_after), + self._final_result, + target="result", + ) + self._finalized = True + + async def _finish_consumption(self) -> None: + if not self._consumed: + self._consumed = True + await self._run_cleanup_hooks() + await self._complete_final_result() + + async def _ensure_buffered_materialized(self) -> None: + self._start_content_pipeline() + if self._buffered_materialized: + return + if self._buffered_materialization_error is not None: + raise self._buffered_materialization_error + if self._buffered_materializing: + raise RuntimeError("ResponseStream does not support concurrent buffered consumption.") + + self._buffered_materializing = True + try: + async with contextlib.AsyncExitStack() as stack: + for factory in self._consumption_context_manager_factories: + manager = factory() + if isinstance(manager, contextlib.AbstractAsyncContextManager): + await stack.enter_async_context(manager) + else: + stack.enter_context(manager) + + buffered_updates: list[UpdateT] = [] + while True: + try: + buffered_updates.append(await self._pull_next_update(run_after_gates=True)) + except StopAsyncIteration: + break + + self._buffered_candidate_updates = buffered_updates + if not self._consumed: + self._consumed = True await self._run_cleanup_hooks() - finally: - self._stream_error = None + await self._prepare_final_result() + + await self._complete_final_result() + released_updates: list[UpdateT] + if self._result_to_updates is not None: + released_updates = list(self._result_to_updates(cast(FinalT, self._final_result))) + for update in released_updates: + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._update_gates_after), + update, + target="update", + ) + elif self._result_was_transformed: + raise RuntimeError( + "A buffered ResponseStream result transform returned a replacement, but result_to_updates " + "was not configured." + ) + else: + released_updates = buffered_updates + + self._buffered_output_updates = released_updates + self._buffered_materialized = True + except Exception as exc: + try: + await self._handle_stream_error(exc) + except Exception as cleanup_exc: + self._buffered_materialization_error = cleanup_exc raise + self._buffered_materialization_error = exc + raise + finally: + self._buffered_materializing = False + + async def __anext__(self) -> UpdateT: + self._start_content_pipeline() + if not self._stream_updates: + await self._ensure_buffered_materialized() + if self._buffered_output_index >= len(self._buffered_output_updates): + raise StopAsyncIteration + update = self._buffered_output_updates[self._buffered_output_index] + self._buffered_output_index += 1 + return update + + try: + update = await self._pull_next_update(run_after_gates=True) + if self._update_gates_before or self._update_gates_after: + self._released_updates.append(update) + return update + except StopAsyncIteration: + await self._finish_consumption() + raise + except Exception as exc: + await self._handle_stream_error(exc) + raise async def close(self) -> None: """Close the active iterator and run cleanup hooks. @@ -3705,115 +4011,113 @@ async def get_final_response(self) -> FinalT: This ensures that post-processing hooks registered on the inner stream (e.g., context provider notifications) are still executed even when the stream is wrapped/mapped. """ - if self._wrap_inner: - if self._inner_stream is None: - # Use _resolve_stream_with_pull_contexts() so that any spans/contexts - # created while resolving the awaitable (e.g. inner telemetry spans) - # inherit the same active context as iterator pulls. This also handles - # the case where _stream_source and _inner_stream_source are the same - # coroutine (e.g., from from_awaitable), avoiding double-await errors. - await self._resolve_stream_with_pull_contexts() - if self._inner_stream is None: - raise RuntimeError("Inner stream not available") - if not self._finalized and not self._consumed: - # Consume outer stream (which delegates to inner) if not already consumed - async for _ in self: - pass - - # Re-check: __anext__ auto-finalization may have already finalized this stream - if not self._finalized: - # This ensures inner post-processing (e.g., context provider notifications) runs - # Skip if inner stream was already finalized (e.g., via auto-finalization on iteration) - if not self._inner_stream._finalized: - inner_stream = self._inner_stream - inner_result: Any - if inner_stream._finalizer is not None: - inner_finalizer = inner_stream._finalizer - inner_result = inner_finalizer(inner_stream._updates) - if isawaitable(inner_result): - inner_result = await inner_result - else: - inner_result = list(inner_stream._updates) - - # Run inner stream's result hooks - inner_hooks = cast(list[Callable[[Any], Any | Awaitable[Any] | None]], inner_stream._result_hooks) - for hook in inner_hooks: - hooked_result = hook(inner_result) - if isawaitable(hooked_result): - hooked_result = await hooked_result - if hooked_result is not None: - inner_result = hooked_result - inner_stream._final_result = inner_result - inner_stream._finalized = True - else: - inner_result = self._inner_stream._final_result - - # Now finalize the outer stream with its own finalizer - # If outer has no finalizer, use inner's result (preserves from_awaitable behavior) - outer_result: Any - if self._finalizer is not None: - outer_result = self._finalizer(self._updates) - if isawaitable(outer_result): - outer_result = await outer_result - else: - # No outer finalizer - use inner's finalized result - outer_result = inner_result - - # Apply outer's result_hooks - outer_hooks = cast(list[Callable[[Any], Any | Awaitable[Any] | None]], self._result_hooks) - for hook in outer_hooks: - outer_hook_result = hook(outer_result) - if isawaitable(outer_hook_result): - outer_hook_result = await outer_hook_result - if outer_hook_result is not None: - outer_result = outer_hook_result - self._final_result = outer_result - self._finalized = True + if not self._stream_updates: + if self._buffered_materializing and self._consumed: + await self._prepare_final_result() + return self._final_result # type: ignore[return-value] + try: + await self._ensure_buffered_materialized() + except Exception: + if not self._return_final_after_buffer_release_error or not self._finalized: + raise return self._final_result # type: ignore[return-value] if not self._finalized and not self._consumed: async for _ in self: pass - # Re-check: __anext__ auto-finalization may have already finalized this stream if not self._finalized: - result: Any - if self._finalizer is not None: - result = self._finalizer(self._updates) - if isawaitable(result): - result = await result - else: - result = list(self._updates) - - final_hooks = cast(list[Callable[[Any], Any | Awaitable[Any] | None]], self._result_hooks) - for hook in final_hooks: - final_hook_result = hook(result) - if isawaitable(final_hook_result): - final_hook_result = await final_hook_result - if final_hook_result is not None: - result = final_hook_result - self._final_result = result - self._finalized = True + await self._complete_final_result() return self._final_result # type: ignore[return-value] - def with_transform_hook( + def _ensure_content_configuration_mutable(self) -> None: + if self._content_pipeline_started: + raise RuntimeError("Cannot change ResponseStream gates or buffering after consumption has started.") + + @staticmethod + def _validate_gate_phase(phase: Literal["before_transform", "after_transform"]) -> None: + if phase not in {"before_transform", "after_transform"}: + raise ValueError(f"Unknown gate phase: {phase}.") + + def with_update_gate( + self, + gate: Callable[[UpdateT], object], + *, + phase: Literal["before_transform", "after_transform"] = "after_transform", + ) -> ResponseStream[UpdateT, FinalT]: + """Register a blocking gate for streamed updates.""" + self._ensure_content_configuration_mutable() + self._validate_gate_phase(phase) + gates = self._update_gates_before if phase == "before_transform" else self._update_gates_after + gates.append(gate) + return self + + def with_result_gate( + self, + gate: Callable[[FinalT], object], + *, + phase: Literal["before_transform", "after_transform"] = "after_transform", + ) -> ResponseStream[UpdateT, FinalT]: + """Register a blocking gate for the finalized result.""" + self._ensure_content_configuration_mutable() + self._validate_gate_phase(phase) + gates = self._result_gates_before if phase == "before_transform" else self._result_gates_after + gates.append(gate) + return self + + def with_update_transform( self, hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None], ) -> ResponseStream[UpdateT, FinalT]: - """Register a transform hook executed for each update during iteration.""" + """Register a transform executed for each update during iteration.""" + if self._content_hooks_sealed: + raise RuntimeError( + "Cannot register an update transform: content is sealed after gated or buffered consumption starts." + ) self._transform_hooks.append(hook) return self - def with_result_hook( + def with_result_transform( self, hook: Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None], ) -> ResponseStream[UpdateT, FinalT]: - """Register a result hook executed after finalization.""" + """Register a transform executed after finalization.""" + if self._content_hooks_sealed: + raise RuntimeError( + "Cannot register a result transform: content is sealed after gated or buffered consumption starts." + ) self._result_hooks.append(hook) + self._result_prepared = False self._finalized = False self._final_result = None return self + def with_transform_hook( + self, + hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None], + ) -> ResponseStream[UpdateT, FinalT]: + """Register an update transform using the legacy hook name.""" + return self.with_update_transform(hook) + + def with_result_hook( + self, + hook: Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a result transform using the legacy hook name.""" + return self.with_result_transform(hook) + + def buffer_updates( + self, + *, + result_to_updates: Callable[[FinalT], Sequence[UpdateT]] | None = None, + ) -> ResponseStream[UpdateT, FinalT]: + """Hold updates until finalization and all configured gates succeed.""" + self._ensure_content_configuration_mutable() + self._stream_updates = False + if result_to_updates is not None: + self._result_to_updates = result_to_updates + return self + def with_cleanup_hook( self, hook: Callable[[], Awaitable[None] | None], @@ -3841,6 +4145,18 @@ def with_pull_context_manager( self._pull_context_manager_factories.append(cm_factory) return self + def with_consumption_context_manager( + self, + cm_factory: Callable[ + [], + contextlib.AbstractContextManager[Any] | contextlib.AbstractAsyncContextManager[Any], + ], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a context manager around complete buffered consumption and finalization.""" + self._ensure_content_configuration_mutable() + self._consumption_context_manager_factories.append(cm_factory) + return self + async def _run_cleanup_hooks(self) -> None: if self._cleanup_run: return @@ -3852,106 +4168,13 @@ async def _run_cleanup_hooks(self) -> None: @property def updates(self) -> Sequence[UpdateT]: + if not self._stream_updates: + return self._buffered_output_updates + if self._update_gates_before or self._update_gates_after: + return self._released_updates return self._updates -class _GatedResponseStream(ResponseStream[UpdateT, FinalT]): - """ResponseStream whose content is sealed once its gate has run. - - Created by :meth:`ResponseStream.buffered_and_gated`. Hooks registered before - the gate runs are applied to the buffered content ahead of the gate; once the - gate has run, registering transform or result hooks raises so nothing can - rewrite content past the gate. (Cleanup hooks remain allowed: they cannot - influence content.) - """ - - _gate_sealed: bool = False - - def with_transform_hook( - self, - hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None], - ) -> ResponseStream[UpdateT, FinalT]: - """Register a transform hook; rejected once the stream's gate has run.""" - if self._gate_sealed: - raise RuntimeError( - "Cannot register a transform hook on a gated ResponseStream after its gate has " - "run: content is sealed by the gate's verdict." - ) - return super().with_transform_hook(hook) - - def with_result_hook( - self, - hook: Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None], - ) -> ResponseStream[UpdateT, FinalT]: - """Register a result hook; rejected once the stream's gate has run.""" - if self._gate_sealed: - raise RuntimeError( - "Cannot register a result hook on a gated ResponseStream after its gate has " - "run: content is sealed by the gate's verdict." - ) - return super().with_result_hook(hook) - - @classmethod - def create_buffered_and_gated( - cls, - consume: Callable[[], Awaitable[tuple[Sequence[UpdateT], FinalT]]], - gate: Callable[[list[UpdateT], FinalT], Awaitable[tuple[FinalT, bool]]], - rederive: Callable[[FinalT], Sequence[UpdateT]], - ) -> _GatedResponseStream[UpdateT, FinalT]: - """Build the gated stream for :meth:`ResponseStream.buffered_and_gated`.""" - holder: dict[str, Any] = {} - - async def _materialize() -> AsyncGenerator[UpdateT]: - stream = cast(_GatedResponseStream[UpdateT, FinalT], holder["stream"]) - updates, final = await consume() - # Drain the hooks registered on the gated stream so far and apply them to - # the buffered content before the gate (contract step 2). Draining also - # means _record_update applies nothing during replay. - transform_hooks = list(stream._transform_hooks) - stream._transform_hooks.clear() - result_hooks = list(stream._result_hooks) - stream._result_hooks.clear() - cleanup_hooks = list(stream._cleanup_hooks) - stream._cleanup_hooks.clear() - hooked_updates: list[UpdateT] = [] - for update in updates: - hooked_update = update - for hook in transform_hooks: - hooked = hook(hooked_update) - if isawaitable(hooked): - hooked = await hooked - if hooked is not None: - hooked_update = cast(UpdateT, hooked) - hooked_updates.append(hooked_update) - for result_hook in result_hooks: - hooked_final = result_hook(final) - if isawaitable(hooked_final): - hooked_final = await hooked_final - if hooked_final is not None: - final = cast(FinalT, hooked_final) - for cleanup_hook in cleanup_hooks: - cleanup_result = cleanup_hook() - if isawaitable(cleanup_result): - await cleanup_result - gated_final, gate_transformed = await gate(hooked_updates, final) - holder["final"] = gated_final - stream._gate_sealed = True - # No-divergence rule (contract step 4), owned here: hooks or a gate - # transform mean the buffered updates may no longer match the verdicted - # result, so the released updates are re-derived from it. - hooks_applied = bool(transform_hooks or result_hooks) - released = rederive(gated_final) if (hooks_applied or gate_transformed) else hooked_updates - for update in released: - yield update - - def _finalizer(_: Sequence[UpdateT]) -> FinalT: - return cast(FinalT, holder["final"]) - - stream: _GatedResponseStream[UpdateT, FinalT] = cls(_materialize(), finalizer=_finalizer) - holder["stream"] = stream - return stream - - # region ChatOptions diff --git a/python/packages/core/tests/core/test_agent_hooks.py b/python/packages/core/tests/core/test_agent_hooks.py index 5a2c66a55e0..a3cbc8c5827 100644 --- a/python/packages/core/tests/core/test_agent_hooks.py +++ b/python/packages/core/tests/core/test_agent_hooks.py @@ -2522,7 +2522,7 @@ async def test_tool_nested_run_inside_drained_attempt_persists_inline() -> None: @requires_sdk -async def test_stream_hooks_cannot_rewrite_egress_after_the_verdict(chat_client_base: MockBaseChatClient) -> None: +async def test_stream_hooks_are_covered_by_the_output_verdict(chat_client_base: MockBaseChatClient) -> None: seen_at_output: list[Any] = [] class RecordingOutputGuard: @@ -2532,7 +2532,7 @@ def intercept(self, context: dict[str, Any]) -> Any: return ALLOW class HookInjector(AgentMiddleware): - """Previously: rewrote egressed updates AFTER the output verdict (fail-open).""" + """Rewrites the stream inside the agent-hooks enforcement boundary.""" async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None: def sneak(update: AgentResponseUpdate) -> AgentResponseUpdate: @@ -2555,11 +2555,12 @@ def sneak(update: AgentResponseUpdate) -> AgentResponseUpdate: updates.append(update.text) final = await stream.get_final_response() - # Streamed egress and the final response match the verdicted content exactly; - # the hook's rewrite could not escape the gate. - assert seen_at_output == ["update - hi"] - assert "".join(updates) == "update - hi" - assert final.text == "update - hi" + # The transform runs before the output verdict, and streamed egress and the + # final response both match the content the verdict inspected. + expected = "update - hi INJECTED-AFTER-VERDICT" + assert seen_at_output == [expected] + assert "".join(updates) == expected + assert final.text == expected @requires_sdk @@ -2648,7 +2649,9 @@ def rederive(final: str) -> list[str]: order.append(f"rederive({final})") return list(final) - stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=gate, rederive=rederive) + with pytest.warns(DeprecationWarning, match="buffered_and_gated"): + stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=gate, rederive=rederive) + assert type(stream) is ResponseStream def transform(update: str) -> str: order.append(f"transform({update})") @@ -2693,7 +2696,12 @@ async def consume() -> tuple[list[str], str]: async def transforming_gate(updates: list[str], final: str) -> tuple[str, bool]: return "XY", True - stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=transforming_gate, rederive=rederive) + with pytest.warns(DeprecationWarning, match="buffered_and_gated"): + stream = cast("Any", ResponseStream).buffered_and_gated( + consume=consume, + gate=transforming_gate, + rederive=rederive, + ) assert [update async for update in stream] == ["X", "Y"] assert await stream.get_final_response() == "XY" assert rederived == ["XY"] @@ -2702,7 +2710,12 @@ async def passthrough_gate(updates: list[str], final: str) -> tuple[str, bool]: return final, False rederived.clear() - stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=passthrough_gate, rederive=rederive) + with pytest.warns(DeprecationWarning, match="buffered_and_gated"): + stream = cast("Any", ResponseStream).buffered_and_gated( + consume=consume, + gate=passthrough_gate, + rederive=rederive, + ) assert [update async for update in stream] == ["a", "b"] assert rederived == [] @@ -2720,7 +2733,8 @@ async def gate(updates: list[str], final: str) -> tuple[str, bool]: def rederive(final: str) -> list[str]: raise RuntimeError("rederive failed") - stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=gate, rederive=rederive) + with pytest.warns(DeprecationWarning, match="buffered_and_gated"): + stream = cast("Any", ResponseStream).buffered_and_gated(consume=consume, gate=gate, rederive=rederive) released: list[str] = [] with pytest.raises(RuntimeError, match="rederive failed"): async for update in stream: diff --git a/python/packages/core/tests/core/test_types.py b/python/packages/core/tests/core/test_types.py index 27f5cf2efb9..40bf49c5ca4 100644 --- a/python/packages/core/tests/core/test_types.py +++ b/python/packages/core/tests/core/test_types.py @@ -4671,6 +4671,372 @@ async def async_hook(response: ChatResponse) -> ChatResponse: assert final.text == "async_update_0update_1" # ty: ignore[unresolved-attribute] +class TestResponseStreamGatesAndBuffering: + """Tests for blocking gates and buffered update release.""" + + def test_constructor_rejects_conflicting_or_invalid_pipeline_configuration(self) -> None: + """Constructor aliases and gate phases fail explicitly when ambiguous.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + + with pytest.raises(ValueError, match="transform_hooks and update_transforms"): + ResponseStream( + updates(), + transform_hooks=[str.upper], + update_transforms=[str.lower], + ) + with pytest.raises(ValueError, match="result_hooks and result_transforms"): + ResponseStream( + updates(), + result_hooks=[lambda value: value], + result_transforms=[lambda value: value], + ) + with pytest.raises(ValueError, match="result_to_updates"): + ResponseStream(updates(), result_to_updates=list) + with pytest.raises(ValueError, match="Unknown update_gates"): + ResponseStream( + updates(), + update_gates=cast(Any, {"invalid": []}), + ) + + async def test_constructor_runs_multiple_gates_around_transforms(self) -> None: + """Constructor configuration preserves phase and registration order.""" + order: list[str] = [] + + async def updates() -> AsyncIterable[str]: + yield "a" + + def update_gate(name: str) -> Callable[[str], None]: + def gate(value: str) -> None: + order.append(f"{name}({value})") + + return gate + + def result_gate(name: str) -> Callable[[str], None]: + def gate(value: str) -> None: + order.append(f"{name}({value})") + + return gate + + def update_transform(value: str) -> str: + order.append(f"update_transform({value})") + return value + "!" + + def finalizer(values: Sequence[str]) -> str: + order.append("finalizer") + return "".join(values) + + def result_transform(value: str) -> str: + order.append(f"result_transform({value})") + return value + "?" + + stream = ResponseStream[str, str]( + updates(), + finalizer=finalizer, + update_gates=cast( + Any, + { + "before_transform": [update_gate("update_before_1"), update_gate("update_before_2")], + "after_transform": [update_gate("update_after")], + }, + ), + result_gates=cast( + Any, + { + "before_transform": [result_gate("result_before")], + "after_transform": [result_gate("result_after")], + }, + ), + update_transforms=[update_transform], + result_transforms=[result_transform], + ) + + assert [update async for update in stream] == ["a!"] + assert await stream.get_final_response() == "a?" + assert order == [ + "update_before_1(a)", + "update_before_2(a)", + "update_transform(a)", + "update_after(a!)", + "finalizer", + "result_before(a)", + "result_transform(a)", + "result_after(a?)", + ] + + async def test_gate_returning_content_is_rejected(self) -> None: + """Gates block by raising and cannot replace content.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + + stream: ResponseStream[str, Sequence[str]] = ResponseStream( + updates(), + update_gates=cast(Any, {"before_transform": [lambda value: value]}), + ) + + with pytest.raises(TypeError, match="must return None or raise"): + await anext(stream) + assert stream.updates == [] + + async def test_async_update_and_result_gates_are_awaited(self) -> None: + """Async gates run at both content levels.""" + seen: list[str] = [] + + async def updates() -> AsyncIterable[str]: + yield "a" + + async def update_gate(value: str) -> None: + await asyncio.sleep(0) + seen.append(f"update:{value}") + + async def result_gate(value: str) -> None: + await asyncio.sleep(0) + seen.append(f"result:{value}") + + stream = ( + ResponseStream[str, str](updates(), finalizer=lambda values: "".join(values)) + .with_update_gate(update_gate) + .with_result_gate(result_gate) + ) + + assert [update async for update in stream] == ["a"] + assert await stream.get_final_response() == "a" + assert seen == ["update:a", "result:a"] + + async def test_live_after_update_gate_does_not_expose_blocked_update(self) -> None: + """An update rejected after transformation is not reported as released.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + + def block(_: str) -> None: + raise ValueError("blocked update") + + stream: ResponseStream[str, Sequence[str]] = ResponseStream( + updates(), + update_transforms=[str.upper], + update_gates=cast(Any, {"after_transform": [block]}), + ) + + with pytest.raises(ValueError, match="blocked update"): + await anext(stream) + assert stream.updates == [] + + async def test_buffered_after_update_gate_failure_releases_and_exposes_nothing(self) -> None: + """Buffered update gates validate every release update before any becomes visible.""" + cleanup_calls = 0 + + async def updates() -> AsyncIterable[str]: + yield "a" + yield "b" + + def block_second(value: str) -> None: + if value == "B": + raise ValueError("blocked update") + + def cleanup() -> None: + nonlocal cleanup_calls + cleanup_calls += 1 + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + update_transforms=[str.upper], + update_gates=cast(Any, {"after_transform": [block_second]}), + cleanup_hooks=[cleanup], + stream_updates=False, + ) + + released: list[str] = [] + with pytest.raises(ValueError, match="blocked update"): + async for update in stream: + released.append(update) + + assert released == [] + assert stream.updates == [] + assert cleanup_calls == 1 + + async def test_buffered_result_replacement_rederives_release_updates(self) -> None: + """A buffered result replacement becomes the authoritative released representation.""" + seen_before: list[str] = [] + seen_after: list[str] = [] + gated_updates: list[str] = [] + + async def updates() -> AsyncIterable[str]: + yield "a" + yield "b" + + def result_to_updates(value: str) -> Sequence[str]: + return list(value) + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + update_gates=cast(Any, {"after_transform": [lambda value: gated_updates.append(value)]}), + update_transforms=[lambda value: value + "!"], + result_gates=cast( + Any, + { + "before_transform": [lambda value: seen_before.append(value)], + "after_transform": [lambda value: seen_after.append(value)], + }, + ), + result_transforms=[lambda _: "XY"], + stream_updates=False, + result_to_updates=result_to_updates, + ) + + assert [update async for update in stream] == ["X", "Y"] + assert await stream.get_final_response() == "XY" + assert stream.updates == ["X", "Y"] + assert seen_before == ["ab"] + assert seen_after == ["XY"] + assert gated_updates == ["a!", "b!", "X", "Y"] + + async def test_buffered_result_replacement_requires_result_to_updates(self) -> None: + """Buffered replacement cannot silently replay stale updates.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + result_transforms=[lambda _: "replacement"], + stream_updates=False, + ) + + with pytest.raises(RuntimeError, match="result_to_updates"): + await anext(stream) + assert stream.updates == [] + + async def test_buffered_result_gate_failure_runs_cleanup_and_releases_nothing(self) -> None: + """Result gates remain fail-closed after source cleanup.""" + cleanup_calls = 0 + + async def updates() -> AsyncIterable[str]: + yield "a" + + def block(_: str) -> None: + raise ValueError("blocked result") + + def cleanup() -> None: + nonlocal cleanup_calls + cleanup_calls += 1 + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + result_gates=cast(Any, {"after_transform": [block]}), + cleanup_hooks=[cleanup], + stream_updates=False, + ) + + with pytest.raises(ValueError, match="blocked result"): + await stream.get_final_response() + assert stream.updates == [] + assert cleanup_calls == 1 + + async def test_live_result_replacement_does_not_rewrite_emitted_updates(self) -> None: + """Live streaming retains already-emitted updates while replacing the final result.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + yield "b" + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + result_transforms=[lambda _: "XY"], + ) + + assert [update async for update in stream] == ["a", "b"] + assert await stream.get_final_response() == "XY" + assert stream.updates == ["a", "b"] + + async def test_buffering_helpers_freeze_content_configuration_on_first_pull(self) -> None: + """Fluent helpers configure an existing stream until buffered consumption starts.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + + stream: ResponseStream[str, str] = ( + ResponseStream(updates(), finalizer=lambda values: "".join(values)) + .with_update_gate(lambda _: None, phase="before_transform") + .with_update_transform(str.upper) + .with_result_gate(lambda _: None) + .buffer_updates() + ) + + assert await anext(stream) == "A" + with pytest.raises(RuntimeError, match="sealed"): + stream.with_result_transform(lambda value: value) + with pytest.raises(RuntimeError, match="consumption has started"): + stream.with_update_gate(lambda _: None) + + async def test_ordinary_stream_still_allows_late_transform_hooks(self) -> None: + """Streams without gates or buffering retain their established late-hook behavior.""" + + async def updates() -> AsyncIterable[str]: + yield "a" + yield "b" + + stream: ResponseStream[str, Sequence[str]] = ResponseStream(updates()) + + assert await anext(stream) == "a" + stream.with_transform_hook(str.upper) + assert await anext(stream) == "B" + + async def test_buffered_map_preserves_inner_result_hooks(self) -> None: + """Buffered mapped streams still finalize their inner stream exactly once.""" + inner_result_calls = 0 + + def inner_result_hook(response: ChatResponse) -> ChatResponse: + nonlocal inner_result_calls + inner_result_calls += 1 + return response + + inner = ResponseStream( + _generate_updates(2), + finalizer=_combine_updates, + result_hooks=[inner_result_hook], + ) + outer = inner.map(lambda update: update, _combine_updates).buffer_updates() + + assert [update.text async for update in outer] == ["update_0", "update_1"] + assert (await outer.get_final_response()).text == "update_0update_1" + assert inner_result_calls == 1 + + async def test_buffered_flat_map_replays_every_outer_update(self) -> None: + """Buffered flat-map streams retain zero-to-many outer update behavior.""" + inner = ResponseStream(_generate_updates(2), finalizer=_combine_updates) + outer = inner.flat_map(lambda update: [update, update], _combine_updates).buffer_updates() + + assert [update.text async for update in outer] == [ + "update_0", + "update_0", + "update_1", + "update_1", + ] + assert (await outer.get_final_response()).text == "update_0update_0update_1update_1" + + async def test_buffered_from_awaitable_resolves_source_once(self) -> None: + """Buffered wrappers do not await an inner stream factory more than once.""" + resolutions = 0 + + async def create_stream() -> ResponseStream[ChatResponseUpdate, ChatResponse]: + nonlocal resolutions + resolutions += 1 + return ResponseStream(_generate_updates(1), finalizer=_combine_updates) + + stream = ResponseStream.from_awaitable(create_stream()).buffer_updates() + + assert (await stream.get_final_response()).text == "update_0" + assert resolutions == 1 + + class TestResponseStreamFinalizer: """Tests for the finalizer.""" diff --git a/python/samples/02-agents/middleware/README.md b/python/samples/02-agents/middleware/README.md index 205dd01504a..4b7d0bbdfb0 100644 --- a/python/samples/02-agents/middleware/README.md +++ b/python/samples/02-agents/middleware/README.md @@ -19,11 +19,11 @@ This folder contains focused middleware samples for `Agent`, chat clients, tools | [`function_based_middleware.py`](./function_based_middleware.py) | Shows function-based agent and function middleware. | | [`middleware_termination.py`](./middleware_termination.py) | Demonstrates stopping a middleware pipeline early. | | [`message_injection_middleware.py`](./message_injection_middleware.py) | Demonstrates `MessageInjectionMiddleware` with a real Foundry chat client: enqueueing a follow-up message into the active session while a long-running async tool is awaiting. | -| [`override_result_with_middleware.py`](./override_result_with_middleware.py) | Shows how middleware can replace regular and streaming results, then post-process the final response. | +| [`override_result_with_middleware.py`](./override_result_with_middleware.py) | Shows how middleware registers result transforms and buffered re-derivation on its context before execution, then post-processes regular and streaming responses. | | [`runtime_context_delegation.py`](./runtime_context_delegation.py) | Demonstrates delegating arguments with runtime context data. | | [`session_behavior_middleware.py`](./session_behavior_middleware.py) | Shows how middleware interacts with session-backed runs. | | [`shared_state_middleware.py`](./shared_state_middleware.py) | Demonstrates sharing mutable state across middleware invocations. | -| [`usage_tracking_middleware.py`](./usage_tracking_middleware.py) | Demonstrates one chat middleware function that tracks per-call usage in non-streaming and streaming tool-loop runs. | +| [`usage_tracking_middleware.py`](./usage_tracking_middleware.py) | Demonstrates one chat middleware function that registers stream transforms on `ChatContext` before execution to track per-call usage in non-streaming and streaming tool-loop runs. | ## Running the usage tracking sample diff --git a/python/samples/02-agents/middleware/override_result_with_middleware.py b/python/samples/02-agents/middleware/override_result_with_middleware.py index 5e886290cf9..238a669896a 100644 --- a/python/samples/02-agents/middleware/override_result_with_middleware.py +++ b/python/samples/02-agents/middleware/override_result_with_middleware.py @@ -2,7 +2,7 @@ import asyncio import re -from collections.abc import AsyncIterable, Awaitable, Callable +from collections.abc import Awaitable, Callable, Sequence from random import randint from typing import Annotated @@ -16,7 +16,6 @@ ChatResponseUpdate, Content, Message, - ResponseStream, tool, ) from agent_framework.openai import OpenAIChatClient @@ -36,11 +35,11 @@ - Replacing function outputs with custom messages or transformed data - Using middleware for result filtering, formatting, or enhancement - Detecting streaming vs non-streaming execution using context.stream -- Overriding streaming results with custom async generators +- Overriding streaming results with result transforms and buffered re-derivation The weather override middleware lets the original weather function execute normally, then replaces its result with a custom "perfect weather" message. For streaming responses, -it creates a custom async generator that yields the override message in chunks. +it buffers the existing ResponseStream and derives released updates from the replacement result. """ @@ -56,73 +55,65 @@ def get_weather( return f"The weather in {location} is {conditions[randint(0, 3)]} with a high of {randint(10, 30)}°C." +def _chat_response_to_updates(response: ChatResponse) -> Sequence[ChatResponseUpdate]: + """Convert a replacement chat response into buffered release updates.""" + return [ChatResponseUpdate(contents=list(message.contents), role="assistant") for message in response.messages] + + async def weather_override_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None: """Chat middleware that overrides weather results for both streaming and non-streaming cases.""" + chunks = [ + "due to special atmospheric conditions, ", + "all locations are experiencing perfect weather today! ", + "Temperature is a comfortable 22°C with gentle breezes. ", + "Perfect day for outdoor activities!", + ] + + if context.stream: + replacement = ChatResponse( + messages=[ + Message(role="assistant", contents=[f"Weather Advisory: [{i}] {chunk_text}"]) + for i, chunk_text in enumerate(chunks) + ] + ) + context.stream_result_transforms.append(lambda _: replacement) + context.stream_buffer_updates = True + context.stream_result_to_updates = _chat_response_to_updates + await call_next() + return - # Let the original agent execution complete first + # Non-streaming middleware updates the concrete result after it exists. await call_next() - - # Check if there's a result to override (agent called weather function) if context.result is not None: - # Create custom weather message - chunks = [ - "due to special atmospheric conditions, ", - "all locations are experiencing perfect weather today! ", - "Temperature is a comfortable 22°C with gentle breezes. ", - "Perfect day for outdoor activities!", - ] - - if context.stream and isinstance(context.result, ResponseStream): - - async def _override_stream() -> AsyncIterable[ChatResponseUpdate]: - for i, chunk_text in enumerate(chunks): - yield ChatResponseUpdate( - contents=[Content.from_text(text=f"Weather Advisory: [{i}] {chunk_text}")], - role="assistant", - ) - - context.result = ResponseStream(_override_stream(), finalizer=ChatResponse.from_updates) - else: - # For non-streaming: just replace with a new message - current_text = context.result.text if isinstance(context.result, ChatResponse) else "" - custom_message = f"Weather Advisory: [0] {''.join(chunks)} Original message was: {current_text}" - context.result = ChatResponse(messages=[Message(role="assistant", contents=[custom_message])]) + current_text = context.result.text if isinstance(context.result, ChatResponse) else "" + custom_message = f"Weather Advisory: [0] {''.join(chunks)} Original message was: {current_text}" + context.result = ChatResponse(messages=[Message(role="assistant", contents=[custom_message])]) async def validate_weather_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None: """Chat middleware that simulates result validation for both streaming and non-streaming cases.""" - await call_next() - validation_note = "Validation: weather data verified." - if context.result is None: - return + if context.stream: - if context.stream and isinstance(context.result, ResponseStream): - result_stream = context.result + def _append_validation(response: ChatResponse) -> ChatResponse: + response.messages.append(Message(role="assistant", contents=[validation_note])) + return response - async def _validated_stream() -> AsyncIterable[ChatResponseUpdate]: - async for update in result_stream: - yield update - yield ChatResponseUpdate( - contents=[Content.from_text(text=validation_note)], - role="assistant", - ) + context.stream_result_transforms.append(_append_validation) + context.stream_buffer_updates = True + context.stream_result_to_updates = _chat_response_to_updates + await call_next() + return - context.result = ResponseStream(_validated_stream(), finalizer=ChatResponse.from_updates) - elif isinstance(context.result, ChatResponse): + await call_next() + if isinstance(context.result, ChatResponse): context.result.messages.append(Message(role="assistant", contents=[validation_note])) async def agent_cleanup_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None: """Agent middleware that validates chat middleware effects and cleans the result.""" - await call_next() - - if context.result is None: - return - validation_note = "Validation: weather data verified." - state = {"found_prefix": False, "found_validation": False} def _sanitize(response: AgentResponse) -> AgentResponse: @@ -170,34 +161,37 @@ def _sanitize(response: AgentResponse) -> AgentResponse: response.messages = cleaned_messages return response - if context.stream and isinstance(context.result, ResponseStream): - - def _clean_update(update: AgentResponseUpdate) -> AgentResponseUpdate: - cleaned_contents: list[Content] = [] + def _clean_update(update: AgentResponseUpdate) -> AgentResponseUpdate: + cleaned_contents: list[Content] = [] - for content in update.contents or []: - if not content.text: - cleaned_contents.append(content) - continue - text = content.text - if "Weather Advisory:" in text: - state["found_prefix"] = True - text = text.replace("Weather Advisory:", "") - if validation_note in text: - state["found_validation"] = True - text = text.replace(validation_note, "").strip() - if not text: - continue - text = re.sub(r"\[\d+\]\s*", "", text) - content.text = text + for content in update.contents or []: + if not content.text: cleaned_contents.append(content) + continue + text = content.text + if "Weather Advisory:" in text: + state["found_prefix"] = True + text = text.replace("Weather Advisory:", "") + if validation_note in text: + state["found_validation"] = True + text = text.replace(validation_note, "").strip() + if not text: + continue + text = re.sub(r"\[\d+\]\s*", "", text) + content.text = text + cleaned_contents.append(content) - update.contents = cleaned_contents - return update + update.contents = cleaned_contents + return update - context.result.with_transform_hook(_clean_update) - context.result.with_result_hook(_sanitize) - elif isinstance(context.result, AgentResponse): + if context.stream: + context.stream_update_transforms.append(_clean_update) + context.stream_result_transforms.append(_sanitize) + await call_next() + return + + await call_next() + if isinstance(context.result, AgentResponse): context.result = _sanitize(context.result) diff --git a/python/samples/02-agents/middleware/usage_tracking_middleware.py b/python/samples/02-agents/middleware/usage_tracking_middleware.py index 056d518fd45..42f0130fec7 100644 --- a/python/samples/02-agents/middleware/usage_tracking_middleware.py +++ b/python/samples/02-agents/middleware/usage_tracking_middleware.py @@ -90,8 +90,8 @@ def capture_final_usage(result: ChatResponse) -> ChatResponse: print(f"\n[Streaming model call #{call_number}] Final usage: {result.usage_details}") return result - context.stream_transform_hooks.append(capture_usage_update) - context.stream_result_hooks.append(capture_final_usage) + context.stream_update_transforms.append(capture_usage_update) + context.stream_result_transforms.append(capture_final_usage) await call_next() return diff --git a/python/samples/02-agents/response_stream.py b/python/samples/02-agents/response_stream.py index 365e4c2bcda..25b09ffa090 100644 --- a/python/samples/02-agents/response_stream.py +++ b/python/samples/02-agents/response_stream.py @@ -21,40 +21,78 @@ - How do you process updates as they arrive? - How do you also get a final, complete response? - How do you ensure the underlying stream is only consumed once? -- How do you add custom logic (hooks) at different stages? +- How do you transform or block content at defined stages? +- How do you hold all updates until a final result is approved? ResponseStream solves all these problems by wrapping an async iterable and providing: - Multiple consumption patterns (iteration OR direct finalization) -- Hook points for transformation, cleanup, finalization, and result processing -- The `wrap()` API to layer behavior without double-consuming the stream +- Ordered gates (checks that return `None` to allow a value or raise to block it) and transforms +- Live or buffered update release +- Cleanup, finalization, and mapping without double-consuming the stream -=== The Four Hook Types === +=== The Content Pipeline === -ResponseStream provides four ways to inject custom logic. All can be passed via constructor -or added later via fluent methods: +Each update follows this pipeline: -1. **Transform Hooks** (`transform_hooks=[]` or `.with_transform_hook()`) - - Called for EACH update as it's yielded during iteration - - Can transform updates before they're returned to the consumer - - Multiple hooks are called in order, each receiving the previous hook's output - - Only triggered during iteration (not when calling get_final_response directly) +```text +source update + -> before-transform update gates + -> update transforms + -> after-transform update gates + -> stream immediately or buffer +``` + +Updates stream immediately by default. Set `stream_updates=False` in the +constructor, or call `.buffer_updates()` on an existing stream, to hold every +update until finalization and all configured gates have passed. -2. **Cleanup Hooks** (`cleanup_hooks=[]` or `.with_cleanup_hook()`) - - Called ONCE when iteration completes (stream fully consumed), BEFORE finalizer - - Used for cleanup: closing connections, releasing resources, logging - - Cannot modify the stream or response - - Triggered regardless of how the stream ends (normal completion or exception) +The final result follows a matching pipeline: -3. **Finalizer** (`finalizer=` constructor parameter) - - Called ONCE when `get_final_response()` is invoked - - Receives the list of collected updates and converts to the final type - - There is only ONE finalizer per stream (set at construction) +```text +finalizer + -> before-transform result gates + -> result transforms + -> after-transform result gates +``` -4. **Result Hooks** (`result_hooks=[]` or `.with_result_hook()`) - - Called ONCE after the finalizer produces its result - - Transform the final response before returning - - Multiple result hooks are called in order, each receiving the previous result - - Can return None to keep the previous value unchanged +The complete lifecycle is: + +1. **Pull a source update.** +2. **Run before-transform update gates** from + `update_gates={"before_transform": [...]}`. +3. **Apply update transforms** from `update_transforms=[]` or + `.with_update_transform()`, in registration order. A transform returns a + replacement update, or `None` to keep the current update. +4. **Run after-transform update gates** from + `update_gates={"after_transform": [...]}`. A gate can block a transformed + update even when buffered mode will hold rather than immediately emit it. +5. **Emit or hold the transformed update.** + - Live mode emits it immediately after step 4 passes. + - Buffered mode holds it until the final result has passed its complete pipeline. +6. **Repeat steps 1-5** until the source is exhausted. +7. **Run cleanup hooks** from `cleanup_hooks=[]` or `.with_cleanup_hook()`. +8. **Run the finalizer** supplied through `finalizer=` over the collected source + updates. +9. **Run before-transform result gates** from + `result_gates={"before_transform": [...]}`. +10. **Apply result transforms** from `result_transforms=[]` or + `.with_result_transform()`, in registration order. A transform returns a + replacement result, or `None` to keep the current result. +11. **Run after-transform result gates** from + `result_gates={"after_transform": [...]}`. +12. **Gate and emit buffered updates.** If a result transform returned a + replacement, `result_to_updates` creates the updates to release. ResponseStream + runs every after-transform update gate against those replacement updates before + emitting the first one. + +The older `transform_hooks`, `result_hooks`, `.with_transform_hook()`, and +`.with_result_hook()` names remain compatible aliases. + +Middleware should register stream transforms, gates, buffering, and conversion +on its `AgentContext` or `ChatContext` before calling `call_next()`. The +middleware pipeline applies that configuration to the eventual ResponseStream +after the chain unwinds. Use the fluent ResponseStream helpers when code outside +middleware already owns a concrete stream. === Two Consumption Patterns === @@ -67,7 +105,7 @@ - Transform hooks are called for each yielded item - Cleanup hooks are called after the last item - The stream collects all updates internally for later finalization -- Does not run the finalizer automatically +- The stream finalizes automatically when iteration reaches the end **Pattern 2: Direct Finalization** ```python @@ -75,7 +113,7 @@ ``` - If the stream hasn't been iterated, it auto-iterates (consuming all updates) - The finalizer converts collected updates to a final response -- Result hooks transform the response +- Update and result pipelines run normally - You get the complete response without ever seeing individual updates ** Pattern 3: Combined Usage ** @@ -92,14 +130,14 @@ final = await response_stream.get_final_response() # Get the aggregated result ``` -=== Chaining with .map() and .with_finalizer() === +=== Chaining with .map(), .flat_map(), and .with_finalizer() === When building a Agent on top of a ChatClient, we face a challenge: - The ChatClient returns a ResponseStream[ChatResponseUpdate, ChatResponse] - The Agent needs to return a ResponseStream[AgentResponseUpdate, AgentResponse] - We can't iterate the ChatClient's stream twice! -The `.map()` and `.with_finalizer()` methods solve this by creating new ResponseStreams that: +The mapping and finalizer methods solve this by creating new ResponseStreams that: - Delegate iteration to the inner stream (only consuming it once) - Maintain their OWN separate transform hooks, result hooks, and cleanup hooks - Allow type-safe transformation of updates and final responses @@ -184,31 +222,30 @@ def combine_updates(updates: Sequence[ChatResponseUpdate]) -> ChatResponse: print(f"Number of updates collected internally: {len(stream2.updates)}") # ========================================================================= - # Example 3: Transform hooks - transform updates during iteration + # Example 3: Update transforms - transform updates during iteration # ========================================================================= - print("\n=== Example 3: Transform Hooks ===\n") + print("\n=== Example 3: Update Transforms ===\n") update_count = {"value": 0} - def counting_hook(update: ChatResponseUpdate) -> ChatResponseUpdate: - """Hook that counts and annotates each update.""" + def counting_transform(update: ChatResponseUpdate) -> ChatResponseUpdate: + """Transform that counts each update without replacing it.""" update_count["value"] += 1 - # Return the update (or a modified version) return update - def uppercase_hook(update: ChatResponseUpdate) -> ChatResponseUpdate: - """Hook that converts text to uppercase.""" + def uppercase_transform(update: ChatResponseUpdate) -> ChatResponseUpdate: + """Transform that converts text to uppercase.""" if update.text: return ChatResponseUpdate( contents=[Content.from_text(update.text.upper())], role=None, response_id=update.response_id ) return update - # Pass transform_hooks directly to constructor + # Pass update transforms directly to the constructor. stream3: ResponseStream[ChatResponseUpdate, ChatResponse] = ResponseStream( generate_updates(), finalizer=combine_updates, - transform_hooks=[counting_hook, uppercase_hook], # First counts, then uppercases + update_transforms=[counting_transform, uppercase_transform], ) print("Iterating with hooks applied:") @@ -242,18 +279,18 @@ async def cleanup_hook() -> None: print(f"Cleanup was performed: {cleanup_performed['value']}") # ========================================================================= - # Example 5: Result hooks - transform the final response + # Example 5: Result transforms - transform the final response # ========================================================================= - print("\n=== Example 5: Result Hooks ===\n") + print("\n=== Example 5: Result Transforms ===\n") - def add_metadata_hook(response: ChatResponse) -> ChatResponse: - """Result hook that adds metadata to the response.""" + def add_metadata_transform(response: ChatResponse) -> ChatResponse: + """Result transform that adds metadata to the response.""" response.additional_properties["processed"] = True response.additional_properties["word_count"] = len((response.text or "").split()) return response - def wrap_in_quotes_hook(response: ChatResponse) -> ChatResponse: - """Result hook that wraps the response text in quotes.""" + def wrap_in_quotes_transform(response: ChatResponse) -> ChatResponse: + """Result transform that wraps the response text in quotes.""" if response.text: return ChatResponse( messages=[Message(contents=[f'"{response.text}"'], role="assistant")], @@ -261,11 +298,11 @@ def wrap_in_quotes_hook(response: ChatResponse) -> ChatResponse: ) return response - # Finalizer converts updates to response, then result hooks transform it + # The finalizer creates a response, then result transforms run in order. stream5: ResponseStream[ChatResponseUpdate, ChatResponse] = ResponseStream( generate_updates(), finalizer=combine_updates, - result_hooks=[add_metadata_hook, wrap_in_quotes_hook], # First adds metadata, then wraps in quotes + result_transforms=[add_metadata_transform, wrap_in_quotes_transform], ) final5 = await stream5.get_final_response() @@ -273,9 +310,100 @@ def wrap_in_quotes_hook(response: ChatResponse) -> ChatResponse: print(f"Metadata: {final5.additional_properties}") # ========================================================================= - # Example 6: The wrap() API - layering without double-consumption + # Example 6: Gates before and after transforms + # ========================================================================= + print("\n=== Example 6: Gates Around Transforms ===\n") + + async def generate_policy_updates() -> AsyncIterable[ChatResponseUpdate]: + """Produce content that must be transformed before egress.""" + for text in ("Public content. ", "Internal secret."): + await asyncio.sleep(0.05) + yield ChatResponseUpdate(contents=[Content.from_text(text)], role="assistant") + + def inspect_source_update(update: ChatResponseUpdate) -> None: + """A before-transform gate can inspect the provider's original update.""" + print(f" [Before gate] Saw: '{update.text}'") + + def redact_update(update: ChatResponseUpdate) -> ChatResponseUpdate: + """Transforms own all content replacement.""" + text = (update.text or "").replace("Internal secret", "[redacted]") + return ChatResponseUpdate(contents=[Content.from_text(text)], role="assistant") + + def require_safe_update(update: ChatResponseUpdate) -> None: + """An after-transform gate blocks if unsafe content would egress.""" + if "secret" in (update.text or "").lower(): + raise RuntimeError("Unsafe update was not redacted.") + + def redact_result(response: ChatResponse) -> ChatResponse: + """Apply the corresponding replacement to the finalized response.""" + text = (response.text or "").replace("Internal secret", "[redacted]") + return ChatResponse(messages=[Message(role="assistant", contents=[text])]) + + def require_safe_result(response: ChatResponse) -> None: + """Validate the final result after its transforms.""" + if "secret" in (response.text or "").lower(): + raise RuntimeError("Unsafe final result was not redacted.") + + gated_stream: ResponseStream[ChatResponseUpdate, ChatResponse] = ResponseStream( + generate_policy_updates(), + finalizer=combine_updates, + update_gates={ + "before_transform": [inspect_source_update], + "after_transform": [require_safe_update], + }, + update_transforms=[redact_update], + result_gates={"after_transform": [require_safe_result]}, + result_transforms=[redact_result], + ) + + print("Released updates:") + async for update in gated_stream: + print(f" -> '{update.text}'") + print(f"Final result: '{(await gated_stream.get_final_response()).text}'") + + # ========================================================================= + # Example 7: Buffer updates so a final-result replacement controls egress + # ========================================================================= + print("\n=== Example 7: Buffered Final Replacement ===\n") + + def log_original_result(response: ChatResponse) -> None: + """A before-transform result gate sees the original finalized response.""" + print(f" [Before result gate] Original: '{response.text}'") + + def replace_result(_: ChatResponse) -> ChatResponse: + """Return a completely different final response.""" + return ChatResponse(messages=[Message(role="assistant", contents=["Approved replacement response."])]) + + def response_to_updates(response: ChatResponse) -> Sequence[ChatResponseUpdate]: + """Convert a replacement final result back into updates for buffered release.""" + return [ChatResponseUpdate(contents=list(message.contents), role="assistant") for message in response.messages] + + buffered_stream: ResponseStream[ChatResponseUpdate, ChatResponse] = ResponseStream( + generate_policy_updates(), + finalizer=combine_updates, + ) + + # Fluent methods are convenient when middleware or another layer receives an + # existing ResponseStream rather than constructing it itself. + ( + buffered_stream + .with_result_gate(log_original_result, phase="before_transform") + .with_result_transform(replace_result) + .with_result_gate(require_safe_result, phase="after_transform") + .with_update_transform(redact_update) + .with_update_gate(require_safe_update, phase="after_transform") + .buffer_updates(result_to_updates=response_to_updates) + ) + + print("The source is fully consumed and the replacement is approved before the first update is released:") + async for update in buffered_stream: + print(f" -> '{update.text}'") + print(f"Final replacement: '{(await buffered_stream.get_final_response()).text}'") + + # ========================================================================= + # Example 8: Mapping - layering without double-consumption # ========================================================================= - print("\n=== Example 6: wrap() API for Layering ===\n") + print("\n=== Example 8: Mapping for Layering ===\n") # Simulate what ChatClient returns inner_stream = ResponseStream(generate_updates(), finalizer=combine_updates) @@ -315,9 +443,9 @@ def to_agent_response(updates: Sequence[ChatResponseUpdate]) -> ChatResponse: print(f"Inner stream consumed: {inner_stream._consumed}") # ========================================================================= - # Example 7: Combining all patterns + # Example 9: Combining lifecycle patterns # ========================================================================= - print("\n=== Example 7: Full Integration ===\n") + print("\n=== Example 9: Lifecycle Integration ===\n") stats = {"updates": 0, "characters": 0} @@ -332,16 +460,16 @@ def log_cleanup() -> None: print(f" [Cleanup] Stream complete: {stats['updates']} updates, {stats['characters']} chars") def add_stats_to_response(response: ChatResponse) -> ChatResponse: - """Result hook to include the statistics in the final response.""" + """Result transform that includes statistics in the final response.""" response.additional_properties["stats"] = stats.copy() return response - # All hooks can be passed via constructor + # Transforms and cleanup hooks can be assembled together in the constructor. full_stream: ResponseStream[ChatResponseUpdate, ChatResponse] = ResponseStream( generate_updates(), finalizer=combine_updates, - transform_hooks=[track_stats], - result_hooks=[add_stats_to_response], + update_transforms=[track_stats], + result_transforms=[add_stats_to_response], cleanup_hooks=[log_cleanup], ) @@ -356,3 +484,16 @@ def add_stats_to_response(response: ChatResponse) -> ChatResponse: if __name__ == "__main__": asyncio.run(main()) + +# Expected output includes: +# === Example 6: Gates Around Transforms === +# [Before gate] Saw: 'Public content. ' +# -> 'Public content. ' +# [Before gate] Saw: 'Internal secret.' +# -> '[redacted].' +# Final result: 'Public content. [redacted].' +# +# === Example 7: Buffered Final Replacement === +# [Before result gate] Original: 'Public content. Internal secret.' +# -> 'Approved replacement response.' +# Final replacement: 'Approved replacement response.' From 0aeadeef2f2dc6db4352fbc84a7350aa6c56ce58 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Tue, 29 Sep 2026 15:18:11 +0200 Subject: [PATCH 2/3] Python: Address ResponseStream review feedback --- .../core/agent_framework/_agent_hooks.py | 75 +++--- .../core/agent_framework/_middleware.py | 244 ++++++++++++++++-- .../packages/core/agent_framework/_types.py | 139 ++++++++-- .../core/tests/core/test_agent_hooks.py | 54 ++++ .../core/tests/core/test_middleware.py | 64 +++++ python/packages/core/tests/core/test_types.py | 71 +++++ python/samples/02-agents/response_stream.py | 5 +- 7 files changed, 565 insertions(+), 87 deletions(-) diff --git a/python/packages/core/agent_framework/_agent_hooks.py b/python/packages/core/agent_framework/_agent_hooks.py index a547d147b76..e2d419fa9e4 100644 --- a/python/packages/core/agent_framework/_agent_hooks.py +++ b/python/packages/core/agent_framework/_agent_hooks.py @@ -1296,42 +1296,28 @@ async def _consumption_scope() -> AsyncGenerator[None]: _RUN_STATE.reset(run_token) async def _transform(final: AgentResponse[Any]) -> AgentResponse[Any] | None: - from agent_hooks import InterceptionBlocked + if not isinstance(final, AgentResponse): + raise MiddlewareException( + f"agent-hooks cannot guard a streamed run result of type {type(final).__name__}; " + "the output interception point was not emitted." + ) + transformed = await self._emit_output(state, final) + return final if transformed else None - try: - if not isinstance(final, AgentResponse): - raise MiddlewareException( - f"agent-hooks cannot guard a streamed run result of type {type(final).__name__}; " - "the output interception point was not emitted." - ) - transformed = await self._emit_output(state, final) - return final if transformed else None - except InterceptionBlocked: - gate_handle.drop() - await self._emit_shutdown(state, "error") - raise - except asyncio.CancelledError: - await self._emit_shutdown(state, "cancelled") - raise - except BaseException: - await self._emit_shutdown(state, "error") - raise + async def _release() -> None: + await gate_handle.flush() + await self._emit_shutdown(state, "completed") - async def _release(_: AgentResponse[Any]) -> None: - try: - await gate_handle.flush() - await self._emit_shutdown(state, "completed") - except asyncio.CancelledError: - await self._emit_shutdown(state, "cancelled") - raise - except BaseException: - await self._emit_shutdown(state, "error") - raise + async def _release_error(exc: BaseException) -> None: + gate_handle.drop() + await self._emit_shutdown(state, "cancelled" if isinstance(exc, asyncio.CancelledError) else "error") - context.stream_result_transforms.append(_transform) - context.stream_result_gates_after.append(_release) + context._stream_terminal_result_transforms.append(_transform) # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_to_updates = _agent_updates_from_response # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_is_authoritative = True # pyright: ignore[reportPrivateUsage] + context._stream_release_hooks.append(_release) # pyright: ignore[reportPrivateUsage] + context._stream_release_error_hooks.append(_release_error) # pyright: ignore[reportPrivateUsage] context.stream_buffer_updates = True - context.stream_result_to_updates = _agent_updates_from_response context.stream_consumption_context_manager_factories.append(_consumption_scope) return inner @@ -1442,29 +1428,26 @@ async def _consumption_scope() -> AsyncGenerator[None]: yield async def _transform(final: ChatResponse[Any]) -> ChatResponse[Any] | None: - from agent_hooks import InterceptionBlocked - if not isinstance(final, ChatResponse): raise MiddlewareException( f"agent-hooks cannot guard a streamed chat result of type {type(final).__name__}; " "the post_model_call interception point was not emitted." ) - try: - changed = await self._emit_post_model_call(state, model_id, final) - return final if changed else None - except InterceptionBlocked: - # §6.1: the deferred per-service-call persistence for the denied - # response is dropped, never executed. - gate_handle.drop() - raise + changed = await self._emit_post_model_call(state, model_id, final) + return final if changed else None - async def _release(_: ChatResponse[Any]) -> None: + async def _release() -> None: await gate_handle.flush() - context.stream_result_transforms.append(_transform) - context.stream_result_gates_after.append(_release) + def _release_error(_: BaseException) -> None: + gate_handle.drop() + + context._stream_terminal_result_transforms.append(_transform) # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_to_updates = _chat_updates_from_response # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_is_authoritative = True # pyright: ignore[reportPrivateUsage] + context._stream_release_hooks.append(_release) # pyright: ignore[reportPrivateUsage] + context._stream_release_error_hooks.append(_release_error) # pyright: ignore[reportPrivateUsage] context.stream_buffer_updates = True - context.stream_result_to_updates = _chat_updates_from_response context.stream_consumption_context_manager_factories.append(_consumption_scope) return inner diff --git a/python/packages/core/agent_framework/_middleware.py b/python/packages/core/agent_framework/_middleware.py index 1b473b9a515..82ac1c65247 100644 --- a/python/packages/core/agent_framework/_middleware.py +++ b/python/packages/core/agent_framework/_middleware.py @@ -329,11 +329,8 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self.stream_update_transforms = list(stream_transform_hooks or []) - self.stream_result_transforms = list(stream_result_hooks or []) - # Compatibility aliases for the original middleware context field names. - self.stream_transform_hooks = self.stream_update_transforms - self.stream_result_hooks = self.stream_result_transforms + self._stream_update_transforms = list(stream_transform_hooks or []) + self._stream_result_transforms = list(stream_result_hooks or []) self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) self.stream_update_gates_before: list[Callable[[AgentResponseUpdate], object]] = [] self.stream_update_gates_after: list[Callable[[AgentResponseUpdate], object]] = [] @@ -342,12 +339,101 @@ def __init__( self.stream_buffer_updates = False self.stream_result_to_updates: Callable[[AgentResponse[Any]], Sequence[AgentResponseUpdate]] | None = None self.stream_consumption_context_manager_factories: list[Callable[[], Any]] = [] + self._stream_terminal_result_transforms: list[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ] = [] + self._stream_terminal_result_gates: list[Callable[[AgentResponse[Any]], object]] = [] + self._stream_terminal_result_to_updates: ( + Callable[[AgentResponse[Any]], Sequence[AgentResponseUpdate]] | None + ) = None + self._stream_terminal_result_is_authoritative = False + self._stream_release_hooks: list[Callable[[], Awaitable[None] | None]] = [] + self._stream_release_error_hooks: list[Callable[[BaseException], Awaitable[None] | None]] = [] # Set by egress-enforcement middleware (agent-hooks): the run-persistence gate # covering this pipeline's run. The final handler offers it for adoption by # the run it starts (see _sessions._offer_run_persistence_gate_claim), so the # gate binds to that run's identity and never to middleware-initiated runs. self._run_persistence_gate: _RunPersistenceGate | None = None + @property + def stream_update_transforms( + self, + ) -> list[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]]: + """Transforms applied to agent updates after the middleware chain unwinds.""" + return self._stream_update_transforms + + @stream_update_transforms.setter + def stream_update_transforms( + self, + transforms: Sequence[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]], + ) -> None: + self._stream_update_transforms = list(transforms) + + @property + def stream_transform_hooks( + self, + ) -> list[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]]: + """Compatibility alias for :attr:`stream_update_transforms`.""" + return self._stream_update_transforms + + @stream_transform_hooks.setter + def stream_transform_hooks( + self, + hooks: Sequence[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]], + ) -> None: + self._stream_update_transforms = list(hooks) + + @property + def stream_result_transforms( + self, + ) -> list[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ]: + """Transforms applied to the finalized agent response after middleware unwinds.""" + return self._stream_result_transforms + + @stream_result_transforms.setter + def stream_result_transforms( + self, + transforms: Sequence[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ], + ) -> None: + self._stream_result_transforms = list(transforms) + + @property + def stream_result_hooks( + self, + ) -> list[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ]: + """Compatibility alias for :attr:`stream_result_transforms`.""" + return self._stream_result_transforms + + @stream_result_hooks.setter + def stream_result_hooks( + self, + hooks: Sequence[ + Callable[ + [AgentResponse[Any]], + AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, + ] + ], + ) -> None: + self._stream_result_transforms = list(hooks) + def _resolve_run_start_tools(self) -> list[ToolTypes]: """Resolve the run-start tool list for this invocation, normalized. @@ -669,11 +755,8 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self.stream_update_transforms = list(stream_transform_hooks or []) - self.stream_result_transforms = list(stream_result_hooks or []) - # Compatibility aliases for the original middleware context field names. - self.stream_transform_hooks = self.stream_update_transforms - self.stream_result_hooks = self.stream_result_transforms + self._stream_update_transforms = list(stream_transform_hooks or []) + self._stream_result_transforms = list(stream_result_hooks or []) self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) self.stream_update_gates_before: list[Callable[[ChatResponseUpdate], object]] = [] self.stream_update_gates_after: list[Callable[[ChatResponseUpdate], object]] = [] @@ -682,9 +765,98 @@ def __init__( self.stream_buffer_updates = False self.stream_result_to_updates: Callable[[ChatResponse[Any]], Sequence[ChatResponseUpdate]] | None = None self.stream_consumption_context_manager_factories: list[Callable[[], Any]] = [] + self._stream_terminal_result_transforms: list[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ] = [] + self._stream_terminal_result_gates: list[Callable[[ChatResponse[Any]], object]] = [] + self._stream_terminal_result_to_updates: Callable[[ChatResponse[Any]], Sequence[ChatResponseUpdate]] | None = ( + None + ) + self._stream_terminal_result_is_authoritative = False + self._stream_release_hooks: list[Callable[[], Awaitable[None] | None]] = [] + self._stream_release_error_hooks: list[Callable[[BaseException], Awaitable[None] | None]] = [] self._message_replacements: list[tuple[Message, tuple[Message, ...]]] = [] self._fallback_reconciliation_messages: list[Message] | None = None + @property + def stream_update_transforms( + self, + ) -> list[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]]: + """Transforms applied to chat updates after the middleware chain unwinds.""" + return self._stream_update_transforms + + @stream_update_transforms.setter + def stream_update_transforms( + self, + transforms: Sequence[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]], + ) -> None: + self._stream_update_transforms = list(transforms) + + @property + def stream_transform_hooks( + self, + ) -> list[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]]: + """Compatibility alias for :attr:`stream_update_transforms`.""" + return self._stream_update_transforms + + @stream_transform_hooks.setter + def stream_transform_hooks( + self, + hooks: Sequence[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]], + ) -> None: + self._stream_update_transforms = list(hooks) + + @property + def stream_result_transforms( + self, + ) -> list[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ]: + """Transforms applied to the finalized chat response after middleware unwinds.""" + return self._stream_result_transforms + + @stream_result_transforms.setter + def stream_result_transforms( + self, + transforms: Sequence[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ], + ) -> None: + self._stream_result_transforms = list(transforms) + + @property + def stream_result_hooks( + self, + ) -> list[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ]: + """Compatibility alias for :attr:`stream_result_transforms`.""" + return self._stream_result_transforms + + @stream_result_hooks.setter + def stream_result_hooks( + self, + hooks: Sequence[ + Callable[ + [ChatResponse[Any]], + ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, + ] + ], + ) -> None: + self._stream_result_transforms = list(hooks) + def record_message_replacement( self, replacement: Message, @@ -1334,18 +1506,32 @@ async def current_handler() -> None: context.result.buffer_updates(result_to_updates=context.stream_result_to_updates) for factory in context.stream_consumption_context_manager_factories: context.result.with_consumption_context_manager(factory) - for hook in context.stream_update_transforms: + for hook in reversed(context.stream_update_transforms): context.result.with_transform_hook(hook) - for result_hook in context.stream_result_transforms: + for result_hook in reversed(context.stream_result_transforms): context.result.with_result_hook(result_hook) - for gate in context.stream_update_gates_before: + for gate in reversed(context.stream_update_gates_before): context.result.with_update_gate(gate, phase="before_transform") - for gate in context.stream_update_gates_after: + for gate in reversed(context.stream_update_gates_after): context.result.with_update_gate(gate, phase="after_transform") - for gate in context.stream_result_gates_before: + for gate in reversed(context.stream_result_gates_before): context.result.with_result_gate(gate, phase="before_transform") - for gate in context.stream_result_gates_after: + for gate in reversed(context.stream_result_gates_after): context.result.with_result_gate(gate, phase="after_transform") + for transform in context._stream_terminal_result_transforms: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_transform(transform) # pyright: ignore[reportPrivateUsage] + for gate in context._stream_terminal_result_gates: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_gate(gate) # pyright: ignore[reportPrivateUsage] + if context._stream_terminal_result_to_updates is not None: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_to_updates( # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_to_updates # pyright: ignore[reportPrivateUsage] + ) + if context._stream_terminal_result_is_authoritative: # pyright: ignore[reportPrivateUsage] + context.result._with_authoritative_terminal_result() # pyright: ignore[reportPrivateUsage] + for hook in context._stream_release_hooks: # pyright: ignore[reportPrivateUsage] + context.result._with_release_hook(hook) # pyright: ignore[reportPrivateUsage] + for hook in context._stream_release_error_hooks: # pyright: ignore[reportPrivateUsage] + context.result._with_release_error_hook(hook) # pyright: ignore[reportPrivateUsage] for cleanup_hook in context.stream_cleanup_hooks: context.result.with_cleanup_hook(cleanup_hook) return context.result @@ -1530,18 +1716,32 @@ async def current_handler() -> None: context.result.buffer_updates(result_to_updates=context.stream_result_to_updates) for factory in context.stream_consumption_context_manager_factories: context.result.with_consumption_context_manager(factory) - for hook in context.stream_update_transforms: + for hook in reversed(context.stream_update_transforms): context.result.with_transform_hook(hook) - for result_hook in context.stream_result_transforms: + for result_hook in reversed(context.stream_result_transforms): context.result.with_result_hook(result_hook) - for gate in context.stream_update_gates_before: + for gate in reversed(context.stream_update_gates_before): context.result.with_update_gate(gate, phase="before_transform") - for gate in context.stream_update_gates_after: + for gate in reversed(context.stream_update_gates_after): context.result.with_update_gate(gate, phase="after_transform") - for gate in context.stream_result_gates_before: + for gate in reversed(context.stream_result_gates_before): context.result.with_result_gate(gate, phase="before_transform") - for gate in context.stream_result_gates_after: + for gate in reversed(context.stream_result_gates_after): context.result.with_result_gate(gate, phase="after_transform") + for transform in context._stream_terminal_result_transforms: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_transform(transform) # pyright: ignore[reportPrivateUsage] + for gate in context._stream_terminal_result_gates: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_gate(gate) # pyright: ignore[reportPrivateUsage] + if context._stream_terminal_result_to_updates is not None: # pyright: ignore[reportPrivateUsage] + context.result._with_terminal_result_to_updates( # pyright: ignore[reportPrivateUsage] + context._stream_terminal_result_to_updates # pyright: ignore[reportPrivateUsage] + ) + if context._stream_terminal_result_is_authoritative: # pyright: ignore[reportPrivateUsage] + context.result._with_authoritative_terminal_result() # pyright: ignore[reportPrivateUsage] + for hook in context._stream_release_hooks: # pyright: ignore[reportPrivateUsage] + context.result._with_release_hook(hook) # pyright: ignore[reportPrivateUsage] + for hook in context._stream_release_error_hooks: # pyright: ignore[reportPrivateUsage] + context.result._with_release_error_hook(hook) # pyright: ignore[reportPrivateUsage] for cleanup_hook in context.stream_cleanup_hooks: context.result.with_cleanup_hook(cleanup_hook) return context.result diff --git a/python/packages/core/agent_framework/_types.py b/python/packages/core/agent_framework/_types.py index b4609360b9e..5ce2a9a740e 100644 --- a/python/packages/core/agent_framework/_types.py +++ b/python/packages/core/agent_framework/_types.py @@ -3452,17 +3452,23 @@ def __init__( ) self._cleanup_hooks: list[Callable[[], Awaitable[None] | None]] = list(cleanup_hooks or ()) self._cleanup_run: bool = False - self._stream_error: Exception | None = None + self._stream_error: BaseException | None = None self._stream_updates = stream_updates self._result_to_updates = result_to_updates + self._terminal_result_to_updates: Callable[[FinalT], Sequence[UpdateT]] | None = None + self._terminal_result_is_authoritative = False self._result_was_transformed = False + self._terminal_result_transforms: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] = [] + self._terminal_result_gates: list[Callable[[FinalT], object]] = [] + self._release_hooks: list[Callable[[], Awaitable[None] | None]] = [] + self._release_error_hooks: list[Callable[[BaseException], Awaitable[None] | None]] = [] self._released_updates: list[UpdateT] = [] self._buffered_candidate_updates: list[UpdateT] = [] self._buffered_output_updates: list[UpdateT] = [] self._buffered_output_index = 0 self._buffered_materialized = False self._buffered_materializing = False - self._buffered_materialization_error: Exception | None = None + self._buffered_materialization_error: BaseException | None = None self._return_final_after_buffer_release_error = False self._content_pipeline_started = False self._content_hooks_sealed = False @@ -3810,13 +3816,41 @@ async def _pull_next_update(self, *, run_after_gates: bool) -> UpdateT: update = await update return await self._record_update(update, run_after_gates=run_after_gates) - async def _handle_stream_error(self, exc: Exception) -> None: + async def _handle_stream_error(self, exc: BaseException) -> None: self._stream_error = exc try: await self._run_cleanup_hooks() finally: self._stream_error = None + async def _run_release_hooks(self) -> None: + for hook in self._release_hooks: + result = hook() + if isawaitable(result): + await result + + async def _run_release_error_hooks(self, exc: BaseException) -> None: + for hook in self._release_error_hooks: + result = hook(exc) + if isawaitable(result): + await result + + async def _abort_buffered_materialization(self, exc: BaseException) -> None: + self._stream_error = exc + try: + iterator = self._iterator + if iterator is not None: + if isinstance(iterator, ResponseStream): + await cast(ResponseStream[UpdateT, Any], iterator).close() + else: + close = getattr(iterator, "aclose", None) + if close is not None: + await close() + self._consumed = True + await self._run_cleanup_hooks() + finally: + self._stream_error = None + async def _prepare_final_result(self) -> None: if self._result_prepared: return @@ -3860,6 +3894,20 @@ async def _complete_final_result(self) -> None: self._final_result, target="result", ) + terminal_result, transformed = await self._apply_transforms( + cast( + Sequence[Callable[[Any], Any | Awaitable[Any | None] | None]], + self._terminal_result_transforms, + ), + self._final_result, + ) + self._final_result = terminal_result + self._result_was_transformed = self._result_was_transformed or transformed + await self._run_gates( + cast(Sequence[Callable[[Any], object]], self._terminal_result_gates), + self._final_result, + target="result", + ) self._finalized = True async def _finish_consumption(self) -> None: @@ -3878,6 +3926,7 @@ async def _ensure_buffered_materialized(self) -> None: raise RuntimeError("ResponseStream does not support concurrent buffered consumption.") self._buffered_materializing = True + release_phase_started = False try: async with contextlib.AsyncExitStack() as stack: for factory in self._consumption_context_manager_factories: @@ -3887,43 +3936,48 @@ async def _ensure_buffered_materialized(self) -> None: else: stack.enter_context(manager) - buffered_updates: list[UpdateT] = [] + self._buffered_candidate_updates = [] while True: try: - buffered_updates.append(await self._pull_next_update(run_after_gates=True)) + self._buffered_candidate_updates.append(await self._pull_next_update(run_after_gates=True)) except StopAsyncIteration: break - self._buffered_candidate_updates = buffered_updates if not self._consumed: self._consumed = True await self._run_cleanup_hooks() await self._prepare_final_result() + release_phase_started = True await self._complete_final_result() released_updates: list[UpdateT] - if self._result_to_updates is not None: - released_updates = list(self._result_to_updates(cast(FinalT, self._final_result))) + needs_rederive = self._result_was_transformed or self._terminal_result_is_authoritative + result_to_updates = self._terminal_result_to_updates or self._result_to_updates + if needs_rederive: + if result_to_updates is None: + raise RuntimeError( + "A buffered ResponseStream transform returned a replacement, but result_to_updates " + "was not configured." + ) + released_updates = list(result_to_updates(cast(FinalT, self._final_result))) for update in released_updates: await self._run_gates( cast(Sequence[Callable[[Any], object]], self._update_gates_after), update, target="update", ) - elif self._result_was_transformed: - raise RuntimeError( - "A buffered ResponseStream result transform returned a replacement, but result_to_updates " - "was not configured." - ) else: - released_updates = buffered_updates + released_updates = self._buffered_candidate_updates + await self._run_release_hooks() self._buffered_output_updates = released_updates self._buffered_materialized = True - except Exception as exc: + except BaseException as exc: try: - await self._handle_stream_error(exc) - except Exception as cleanup_exc: + if release_phase_started: + await self._run_release_error_hooks(exc) + await self._abort_buffered_materialization(exc) + except BaseException as cleanup_exc: self._buffered_materialization_error = cleanup_exc raise self._buffered_materialization_error = exc @@ -4157,6 +4211,57 @@ def with_consumption_context_manager( self._consumption_context_manager_factories.append(cm_factory) return self + def _with_terminal_result_transform( + self, + transform: Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a framework-owned transform after all composable result stages.""" + self._ensure_content_configuration_mutable() + self._terminal_result_transforms.append(transform) + return self + + def _with_terminal_result_gate( + self, + gate: Callable[[FinalT], object], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a framework-owned gate after terminal result transforms.""" + self._ensure_content_configuration_mutable() + self._terminal_result_gates.append(gate) + return self + + def _with_terminal_result_to_updates( + self, + result_to_updates: Callable[[FinalT], Sequence[UpdateT]], + ) -> ResponseStream[UpdateT, FinalT]: + """Bind the trusted converter used after terminal enforcement transforms content.""" + self._ensure_content_configuration_mutable() + self._terminal_result_to_updates = result_to_updates + return self + + def _with_authoritative_terminal_result(self) -> ResponseStream[UpdateT, FinalT]: + """Require buffered release updates to be derived from the terminal result.""" + self._ensure_content_configuration_mutable() + self._terminal_result_is_authoritative = True + return self + + def _with_release_hook( + self, + hook: Callable[[], Awaitable[None] | None], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a framework-owned hook after release updates pass every gate.""" + self._ensure_content_configuration_mutable() + self._release_hooks.append(hook) + return self + + def _with_release_error_hook( + self, + hook: Callable[[BaseException], Awaitable[None] | None], + ) -> ResponseStream[UpdateT, FinalT]: + """Register a framework-owned hook for terminal enforcement or release failures.""" + self._ensure_content_configuration_mutable() + self._release_error_hooks.append(hook) + return self + async def _run_cleanup_hooks(self) -> None: if self._cleanup_run: return diff --git a/python/packages/core/tests/core/test_agent_hooks.py b/python/packages/core/tests/core/test_agent_hooks.py index a3cbc8c5827..c87c424d028 100644 --- a/python/packages/core/tests/core/test_agent_hooks.py +++ b/python/packages/core/tests/core/test_agent_hooks.py @@ -964,6 +964,7 @@ async def test_streaming_output_deny_releases_nothing(chat_client_base: MockBase assert updates == [] # nothing egressed before the deny assert points(records)[-1] == "agent_shutdown" # the record trail is closed + assert points(records).count("agent_shutdown") == 1 @requires_sdk @@ -2563,6 +2564,59 @@ def sneak(update: AgentResponseUpdate) -> AgentResponseUpdate: assert final.text == expected +@requires_sdk +async def test_outer_result_transform_runs_before_output_verdict(chat_client_base: MockBaseChatClient) -> None: + """Terminal enforcement covers transforms registered after the bundle unwinds.""" + seen_at_output: list[Any] = [] + + class RecordingOutputGuard: + def intercept(self, context: dict[str, Any]) -> Any: + if context["interception_point"] == "output": + seen_at_output.append(context["target"]["content"]) + return ALLOW + + class OuterTransform(AgentMiddleware): + async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None: + await call_next() + context.stream_result_transforms.append( + lambda _: AgentResponse(messages=[Message(role="assistant", contents=["outer replacement"])]) + ) + + agent = Agent( + client=chat_client_base, + middleware=[OuterTransform(), create_agent_hooks_middleware([RecordingOutputGuard()])], + ) + + stream = agent.run("hi", stream=True) + assert "".join([update.text async for update in stream]) == "outer replacement" + assert (await stream.get_final_response()).text == "outer replacement" + assert seen_at_output == ["outer replacement"] + + +@requires_sdk +async def test_outer_middleware_cannot_replace_agent_hooks_release_converter( + chat_client_base: MockBaseChatClient, +) -> None: + """The trusted terminal converter owns post-verdict update derivation.""" + + class OuterConverter(AgentMiddleware): + async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None: + await call_next() + context.stream_result_to_updates = lambda _: [ + AgentResponseUpdate(contents=[Content.from_text("BYPASS")], role="assistant") + ] + + agent = Agent( + client=chat_client_base, + middleware=[OuterConverter(), create_agent_hooks_middleware([AllowGuard()])], + ) + + stream = agent.run("hi", stream=True) + released = "".join([update.text async for update in stream]) + assert released == "update - hi" + assert "BYPASS" not in released + + @requires_sdk async def test_as_tool_stream_callback_sees_nothing_on_deny() -> None: # The observe direction of the gate contract: `as_tool(stream_callback=...)` is a diff --git a/python/packages/core/tests/core/test_middleware.py b/python/packages/core/tests/core/test_middleware.py index c3e59903317..f6a28295c57 100644 --- a/python/packages/core/tests/core/test_middleware.py +++ b/python/packages/core/tests/core/test_middleware.py @@ -71,6 +71,22 @@ def test_init_with_session(self, mock_agent: SupportsAgentRun) -> None: assert context.stream is False assert context.metadata == {} + def test_stream_transform_alias_assignment_stays_synchronized(self, mock_agent: SupportsAgentRun) -> None: + """Legacy and preferred transform fields share one backing configuration.""" + context = AgentContext(agent=mock_agent, messages=[Message(role="user", contents=["test"])]) + + def legacy(update: AgentResponseUpdate) -> AgentResponseUpdate: + return update + + def preferred(response: AgentResponse[Any]) -> AgentResponse[Any]: + return response + + context.stream_transform_hooks = [legacy] + assert context.stream_update_transforms == [legacy] + + context.stream_result_transforms = [preferred] + assert context.stream_result_hooks == [preferred] + class TestFunctionInvocationContext: """Test cases for FunctionInvocationContext.""" @@ -131,6 +147,22 @@ def test_init_with_custom_values(self, mock_chat_client: Any) -> None: assert context.stream is True assert context.metadata == metadata + def test_stream_transform_alias_assignment_stays_synchronized(self, mock_chat_client: Any) -> None: + """Legacy and preferred transform fields share one backing configuration.""" + context = ChatContext(client=mock_chat_client, messages=[], options={}) + + def preferred(update: ChatResponseUpdate) -> ChatResponseUpdate: + return update + + def legacy(response: ChatResponse[Any]) -> ChatResponse[Any]: + return response + + context.stream_update_transforms = [preferred] + assert context.stream_transform_hooks == [preferred] + + context.stream_result_hooks = [legacy] + assert context.stream_result_transforms == [legacy] + def test_record_message_replacement(self, mock_chat_client: Any) -> None: """Replacement provenance is validated, deduplicated, and returned.""" source = Message(role="tool", contents=["source"]) @@ -685,6 +717,38 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: assert updates[1].text == "chunk2" assert execution_order == ["test_before", "test_after", "handler_start", "handler_end"] + async def test_stream_result_transforms_follow_middleware_unwind_order(self, mock_chat_client: Any) -> None: + """Inner middleware transforms run before outer post-processing transforms.""" + + class ResultTransformMiddleware(ChatMiddleware): + def __init__(self, suffix: str): + self.suffix = suffix + + async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None: + def append_suffix(response: ChatResponse) -> ChatResponse: + return ChatResponse( + messages=[Message(role="assistant", contents=[f"{response.text}{self.suffix}"])] + ) + + context.stream_result_transforms.append(append_suffix) + await call_next() + + pipeline = ChatMiddlewarePipeline( + ResultTransformMiddleware("-outer"), + ResultTransformMiddleware("-inner"), + ) + context = ChatContext(client=mock_chat_client, messages=[], options={}, stream=True) + + def final_handler(_: ChatContext) -> ResponseStream[ChatResponseUpdate, ChatResponse]: + async def updates() -> AsyncIterable[ChatResponseUpdate]: + yield ChatResponseUpdate(contents=[Content.from_text("base")], role="assistant") + + return ResponseStream(updates(), finalizer=ChatResponse.from_updates) + + stream = await pipeline.execute(context, final_handler) + assert isinstance(stream, ResponseStream) + assert (await stream.get_final_response()).text == "base-inner-outer" + async def test_execute_with_pre_next_termination(self, mock_chat_client: Any) -> None: """Test pipeline execution with termination before next().""" middleware = self.PreNextTerminateChatMiddleware() diff --git a/python/packages/core/tests/core/test_types.py b/python/packages/core/tests/core/test_types.py index 40bf49c5ca4..2c76441d58b 100644 --- a/python/packages/core/tests/core/test_types.py +++ b/python/packages/core/tests/core/test_types.py @@ -4895,6 +4895,38 @@ def result_to_updates(value: str) -> Sequence[str]: assert seen_after == ["XY"] assert gated_updates == ["a!", "b!", "X", "Y"] + async def test_buffered_converter_preserves_unchanged_updates(self) -> None: + """A configured converter is not used when no transform replaces content.""" + converter_calls = 0 + original = ChatResponseUpdate( + contents=[Content.from_text("a")], + role="assistant", + continuation_token={}, + additional_properties={"provider": "metadata"}, + ) + + async def updates() -> AsyncIterable[ChatResponseUpdate]: + yield original + + def result_to_updates(_: ChatResponse) -> Sequence[ChatResponseUpdate]: + nonlocal converter_calls + converter_calls += 1 + return [ChatResponseUpdate(contents=[Content.from_text("rebuilt")], role="assistant")] + + stream = ResponseStream( + updates(), + finalizer=ChatResponse.from_updates, + stream_updates=False, + result_to_updates=result_to_updates, + ) + + released = [update async for update in stream] + assert released == [original] + assert released[0] is original + assert released[0].continuation_token == {} + assert released[0].additional_properties == {"provider": "metadata"} + assert converter_calls == 0 + async def test_buffered_result_replacement_requires_result_to_updates(self) -> None: """Buffered replacement cannot silently replay stale updates.""" @@ -4939,6 +4971,45 @@ def cleanup() -> None: assert stream.updates == [] assert cleanup_calls == 1 + async def test_buffered_cancellation_is_terminal(self) -> None: + """Cancellation closes the source, runs cleanup once, and cannot be retried.""" + source_waiting = asyncio.Event() + source_closed = False + cleanup_calls = 0 + + async def updates() -> AsyncIterable[str]: + nonlocal source_closed + try: + yield "a" + source_waiting.set() + await asyncio.Event().wait() + finally: + source_closed = True + + def cleanup() -> None: + nonlocal cleanup_calls + cleanup_calls += 1 + + stream: ResponseStream[str, str] = ResponseStream( + updates(), + finalizer=lambda values: "".join(values), + cleanup_hooks=[cleanup], + stream_updates=False, + ) + + pull = asyncio.create_task(anext(stream)) + await source_waiting.wait() + pull.cancel() + with pytest.raises(asyncio.CancelledError): + await pull + + assert stream._buffered_candidate_updates == ["a"] + assert stream.updates == [] + assert source_closed is True + assert cleanup_calls == 1 + with pytest.raises(asyncio.CancelledError): + await anext(stream) + async def test_live_result_replacement_does_not_rewrite_emitted_updates(self) -> None: """Live streaming retains already-emitted updates while replacing the final result.""" diff --git a/python/samples/02-agents/response_stream.py b/python/samples/02-agents/response_stream.py index 25b09ffa090..917a4490233 100644 --- a/python/samples/02-agents/response_stream.py +++ b/python/samples/02-agents/response_stream.py @@ -91,8 +91,9 @@ Middleware should register stream transforms, gates, buffering, and conversion on its `AgentContext` or `ChatContext` before calling `call_next()`. The middleware pipeline applies that configuration to the eventual ResponseStream -after the chain unwinds. Use the fluent ResponseStream helpers when code outside -middleware already owns a concrete stream. +after the chain unwinds, in unwind order: inner middleware post-processing runs +before outer middleware post-processing. Use the fluent ResponseStream helpers +when code outside middleware already owns a concrete stream. === Two Consumption Patterns === From c7d34d0737d65180b36522a0e7d9224d40503f45 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Tue, 29 Sep 2026 15:20:24 +0200 Subject: [PATCH 3/3] Python: Preserve middleware hook assignment compatibility --- .../core/agent_framework/_middleware.py | 100 +++--------------- .../core/tests/core/test_middleware.py | 4 +- 2 files changed, 14 insertions(+), 90 deletions(-) diff --git a/python/packages/core/agent_framework/_middleware.py b/python/packages/core/agent_framework/_middleware.py index 82ac1c65247..8af5c885c4c 100644 --- a/python/packages/core/agent_framework/_middleware.py +++ b/python/packages/core/agent_framework/_middleware.py @@ -329,8 +329,8 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self._stream_update_transforms = list(stream_transform_hooks or []) - self._stream_result_transforms = list(stream_result_hooks or []) + self.stream_transform_hooks = list(stream_transform_hooks or []) + self.stream_result_hooks = list(stream_result_hooks or []) self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) self.stream_update_gates_before: list[Callable[[AgentResponseUpdate], object]] = [] self.stream_update_gates_after: list[Callable[[AgentResponseUpdate], object]] = [] @@ -363,28 +363,14 @@ def stream_update_transforms( self, ) -> list[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]]: """Transforms applied to agent updates after the middleware chain unwinds.""" - return self._stream_update_transforms + return self.stream_transform_hooks @stream_update_transforms.setter def stream_update_transforms( self, transforms: Sequence[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]], ) -> None: - self._stream_update_transforms = list(transforms) - - @property - def stream_transform_hooks( - self, - ) -> list[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]]: - """Compatibility alias for :attr:`stream_update_transforms`.""" - return self._stream_update_transforms - - @stream_transform_hooks.setter - def stream_transform_hooks( - self, - hooks: Sequence[Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]], - ) -> None: - self._stream_update_transforms = list(hooks) + self.stream_transform_hooks = list(transforms) @property def stream_result_transforms( @@ -396,7 +382,7 @@ def stream_result_transforms( ] ]: """Transforms applied to the finalized agent response after middleware unwinds.""" - return self._stream_result_transforms + return self.stream_result_hooks @stream_result_transforms.setter def stream_result_transforms( @@ -408,31 +394,7 @@ def stream_result_transforms( ] ], ) -> None: - self._stream_result_transforms = list(transforms) - - @property - def stream_result_hooks( - self, - ) -> list[ - Callable[ - [AgentResponse[Any]], - AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, - ] - ]: - """Compatibility alias for :attr:`stream_result_transforms`.""" - return self._stream_result_transforms - - @stream_result_hooks.setter - def stream_result_hooks( - self, - hooks: Sequence[ - Callable[ - [AgentResponse[Any]], - AgentResponse[Any] | Awaitable[AgentResponse[Any] | None] | None, - ] - ], - ) -> None: - self._stream_result_transforms = list(hooks) + self.stream_result_hooks = list(transforms) def _resolve_run_start_tools(self) -> list[ToolTypes]: """Resolve the run-start tool list for this invocation, normalized. @@ -755,8 +717,8 @@ def __init__( self.function_invocation_kwargs: dict[str, Any] = ( dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {} ) - self._stream_update_transforms = list(stream_transform_hooks or []) - self._stream_result_transforms = list(stream_result_hooks or []) + self.stream_transform_hooks = list(stream_transform_hooks or []) + self.stream_result_hooks = list(stream_result_hooks or []) self.stream_cleanup_hooks = list(stream_cleanup_hooks or []) self.stream_update_gates_before: list[Callable[[ChatResponseUpdate], object]] = [] self.stream_update_gates_after: list[Callable[[ChatResponseUpdate], object]] = [] @@ -786,28 +748,14 @@ def stream_update_transforms( self, ) -> list[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]]: """Transforms applied to chat updates after the middleware chain unwinds.""" - return self._stream_update_transforms + return self.stream_transform_hooks @stream_update_transforms.setter def stream_update_transforms( self, transforms: Sequence[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]], ) -> None: - self._stream_update_transforms = list(transforms) - - @property - def stream_transform_hooks( - self, - ) -> list[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]]: - """Compatibility alias for :attr:`stream_update_transforms`.""" - return self._stream_update_transforms - - @stream_transform_hooks.setter - def stream_transform_hooks( - self, - hooks: Sequence[Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]], - ) -> None: - self._stream_update_transforms = list(hooks) + self.stream_transform_hooks = list(transforms) @property def stream_result_transforms( @@ -819,7 +767,7 @@ def stream_result_transforms( ] ]: """Transforms applied to the finalized chat response after middleware unwinds.""" - return self._stream_result_transforms + return self.stream_result_hooks @stream_result_transforms.setter def stream_result_transforms( @@ -831,31 +779,7 @@ def stream_result_transforms( ] ], ) -> None: - self._stream_result_transforms = list(transforms) - - @property - def stream_result_hooks( - self, - ) -> list[ - Callable[ - [ChatResponse[Any]], - ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, - ] - ]: - """Compatibility alias for :attr:`stream_result_transforms`.""" - return self._stream_result_transforms - - @stream_result_hooks.setter - def stream_result_hooks( - self, - hooks: Sequence[ - Callable[ - [ChatResponse[Any]], - ChatResponse[Any] | Awaitable[ChatResponse[Any] | None] | None, - ] - ], - ) -> None: - self._stream_result_transforms = list(hooks) + self.stream_result_hooks = list(transforms) def record_message_replacement( self, diff --git a/python/packages/core/tests/core/test_middleware.py b/python/packages/core/tests/core/test_middleware.py index f6a28295c57..314860811ca 100644 --- a/python/packages/core/tests/core/test_middleware.py +++ b/python/packages/core/tests/core/test_middleware.py @@ -81,7 +81,7 @@ def legacy(update: AgentResponseUpdate) -> AgentResponseUpdate: def preferred(response: AgentResponse[Any]) -> AgentResponse[Any]: return response - context.stream_transform_hooks = [legacy] + context.stream_transform_hooks = cast(Any, [legacy]) assert context.stream_update_transforms == [legacy] context.stream_result_transforms = [preferred] @@ -160,7 +160,7 @@ def legacy(response: ChatResponse[Any]) -> ChatResponse[Any]: context.stream_update_transforms = [preferred] assert context.stream_transform_hooks == [preferred] - context.stream_result_hooks = [legacy] + context.stream_result_hooks = cast(Any, [legacy]) assert context.stream_result_transforms == [legacy] def test_record_message_replacement(self, mock_chat_client: Any) -> None: