Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 62 additions & 62 deletions python/packages/core/agent_framework/_agent_hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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``
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -1255,25 +1257,27 @@ async def _process_streaming(

def _gated_agent_stream(
self,
context: AgentContext,
state: _RunState,
inner: ResponseStream[AgentResponseUpdate, AgentResponse[Any]],
gate_handle: _RunPersistenceGate,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
"""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
Comment thread
eavanvalkenburg marked this conversation as resolved.
except asyncio.CancelledError:
await self._emit_shutdown(state, "cancelled")
raise
Expand All @@ -1290,40 +1294,32 @@ 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]:
from agent_hooks import InterceptionBlocked
async def _transform(final: AgentResponse[Any]) -> AgentResponse[Any] | None:
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)
await gate_handle.flush()
await self._emit_shutdown(state, "completed")
return final, transformed
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")

return cast(
"ResponseStream[AgentResponseUpdate, AgentResponse[Any]]",
cast(Any, ResponseStream).buffered_and_gated(
consume=_consume, gate=_gate, rederive=_agent_updates_from_response
),
)
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_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_consumption_context_manager_factories.append(_consumption_scope)
return inner


class _AgentHooksChatMiddleware(_AgentHooksMiddlewareBase, ChatMiddleware):
Expand Down Expand Up @@ -1367,7 +1363,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:
Expand Down Expand Up @@ -1406,6 +1406,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]],
Expand All @@ -1417,39 +1418,38 @@ 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

async def _gate(_updates: list[ChatResponseUpdate], final: ChatResponse[Any]) -> tuple[ChatResponse[Any], bool]:
from agent_hooks import InterceptionBlocked
yield

async def _transform(final: ChatResponse[Any]) -> ChatResponse[Any] | None:
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)
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() -> 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
),
)
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_consumption_context_manager_factories.append(_consumption_scope)
return inner


class _AgentHooksFunctionMiddleware(_AgentHooksMiddlewareBase, FunctionMiddleware):
Expand Down
2 changes: 1 addition & 1 deletion python/packages/core/agent_framework/_agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading
Loading