From a70fed415d728b009a7aa29563fa075158410f41 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Wed, 23 Sep 2026 11:17:31 +0200 Subject: [PATCH 1/5] Python: make AG-UI runs disconnect durable --- python/packages/ag-ui/AGENTS.md | 14 + python/packages/ag-ui/README.md | 44 +- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 78 ++- .../ag-ui/agent_framework_ag_ui/_endpoint.py | 116 ++++- .../_snapshot_session.py | 12 +- .../ag-ui/tests/ag_ui/test_endpoint.py | 484 +++++++++++++++++- python/packages/ag-ui/tests/ag_ui/test_run.py | 280 +++++++++- 7 files changed, 977 insertions(+), 51 deletions(-) diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index 94a11e1ca29..297e665290e 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -112,6 +112,20 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. `state` entry, so conversation continuation is unreachable from client input by construction; keep it that way. - `confirm_changes` snapshot cleanup resolves the synthetic confirmation back to its original `function_call_id`; it must never concatenate unrelated tool results or record accepted changes without a matching real result. +- Configured stateless agent Thread Snapshot stores persist after complete model-roundtrip, tool-result, and approval + safe points, then persist the terminal state. Save only after the complete update and approval lifecycle side effects + are applied; do not write every text delta or schedule unordered background snapshot writes. Service-session and + workflow Thread Snapshots retain terminal-save cadence so replayable messages cannot advance without matching provider + continuation state; workflow checkpoints own incremental workflow runtime state. +- Disconnect-safe execution is endpoint-owned and opt-in through + `add_agent_framework_fastapi_endpoint(detached_runs=True)`. Keep its producer queue bounded, retain and observe + producer/drainer tasks, and let the producer own final snapshot/checkpoint/approval persistence after the SSE reader + disconnects. This mode discards unread events; it is not a resumable event log. +- While a detached mutation is active, the endpoint may serve an empty Snapshot Hydrate Request but must reject another + mutation for the same `(Snapshot Scope, threadId)` with HTTP 409. The guard is process-local and does not replace + cross-replica coordination. +- Detached runs may outlive FastAPI request-scoped disposable resources. Resolve authorization and Snapshot Scope before + spawning the producer, and do not rely on request-owned clients or sessions remaining open after disconnect. - SSE keepalive is endpoint-owned transport behavior configured through `add_agent_framework_fastapi_endpoint(keepalive_seconds=...)`. It emits SSE comments only; do not add `PING`, `HEARTBEAT`, or `KEEPALIVE` AG-UI events, and do not add runner-level keepalive settings. diff --git a/python/packages/ag-ui/README.md b/python/packages/ag-ui/README.md index ee9c0af4edf..4eaacff3d7b 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -389,6 +389,12 @@ add_agent_framework_fastapi_endpoint( ) ``` +Configured stateless agent snapshot stores are updated after completed model roundtrips, tool-result batches, and +approval safe points, then written once more with the terminal run state. This limits progress loss during long agent +runs without persisting every streaming text delta. Service-session snapshots retain terminal-save cadence so +replayable messages cannot advance without their matching provider continuation state. Workflow Thread Snapshots also +keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow runtime state. + A frontend can then hydrate the latest stored snapshot for the scoped thread: ```json @@ -398,6 +404,36 @@ A frontend can then hydrate the latest stored snapshot for the scoped thread: } ``` +### Disconnect-safe runs + +By default, AG-UI execution remains attached to the SSE response: disconnecting the client cancels the response +generator and can stop the run. Set `detached_runs=True` when a finite agent or workflow run must continue through +snapshot, checkpoint, and approval-state finalization after the HTTP reader disconnects: + +```python +add_agent_framework_fastapi_endpoint( + app, + agent, + "/", + snapshot_store=snapshot_store, + snapshot_scope_resolver=resolve_snapshot_scope, + detached_runs=True, +) +``` + +Detached execution uses a bounded endpoint-owned producer queue. While a detached mutating request is active, another +mutating request for the same `(Snapshot Scope, threadId)` returns HTTP 409; an empty snapshot Hydrate Request remains +allowed and returns the latest committed safe point. Equal Thread ids in different Snapshot Scopes remain independent. + +This option does not add resumable event replay. Events emitted while no client is attached are consumed and discarded, +so a reconnecting client recovers from Thread Snapshots rather than resuming the original SSE position. Applications +that require replay of every in-flight event must provide their own authenticated event log and resume route with +retention and cross-replica semantics appropriate to their deployment. + +FastAPI dependencies and other request-scoped resources may be released after the disconnected response ends. Resolve +authorization, Snapshot Scope, and other durable values before the run starts; detached tools and providers must not +retain request-owned clients or sessions that are expected to close with the HTTP request. + Endpoint configuration requires `snapshot_scope_resolver` whenever a snapshot store is configured, including when the store is already set on a pre-wrapped `AgentFrameworkAgent` or `AgentFrameworkWorkflow`. The resolver returns the application-defined Snapshot Scope used with the AG-UI Thread id as the storage key. The endpoint also derives @@ -468,9 +504,11 @@ encryption, integrity protection, access control, retention, audit, data residen custom stores remain source-compatible because `session_state` is optional, but they provide Session State Continuity only when they round-trip that field unchanged with the rest of the snapshot. -The supported consistency model is one active run per `(Snapshot Scope, threadId)`. Concurrent writes to the same -scoped thread remain last-writer-wins. Applications that require stronger consistency must serialize those runs using -coordination appropriate to their deployment; a process-local lock does not provide distributed consistency. +The supported consistency model is one active run per `(Snapshot Scope, threadId)`. With `detached_runs=True`, one +registered endpoint enforces that rule in-process by rejecting concurrent mutations while allowing hydration. +Coordination is not shared across endpoint registrations, workers, or replicas; applications that require distributed +serialization must provide it using infrastructure appropriate to their deployment. Without detached execution, +concurrent writes to the same scoped thread remain last-writer-wins. ## Architecture diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 9fbb15f8d65..4a3e939fd6c 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -3070,6 +3070,25 @@ async def _run_agent_stream( _restore_tool_approval_state(session, approval_state_store, approval_thread_id) approval_middleware_pipeline = _approval_observer_middleware_pipeline(agent, session) + async def save_thread_snapshot( + *, + persisted_messages: list[dict[str, Any]], + state: dict[str, Any] | None, + interrupt: list[dict[str, Any]] | None, + ) -> None: + _save_tool_approval_state(session, approval_state_store, approval_thread_id) + await snapshot_session.save( + messages=_bound_host_payload_history(_persistable_host_payload_history(persisted_messages)), + state=state, + interrupt=interrupt, + session_state=_safe_serialize_session_continuation_state( + session, + agent, + shared_state_keys=set(flow.current_state).difference(protected_session_state_keys), + include_service_session_id=config.use_service_session, + ), + ) + authenticated_cancellations = [ _approval_observer_response(occurrence, cancelled=True) for interrupt_id in cancelled_resume_ids @@ -3282,18 +3301,11 @@ async def _run_agent_stream( persisted_messages = snapshot_messages if resume_payload is not None and not seeded_resume_from_snapshot and snapshot_seed_messages is None: persisted_messages = snapshot_session.resume_seeded_messages(persisted_messages) - await snapshot_session.save( - messages=_bound_host_payload_history(_persistable_host_payload_history(persisted_messages)), + await save_thread_snapshot( + persisted_messages=persisted_messages, state=cast(dict[str, Any], make_json_safe(flow.current_state)) if flow.current_state else None, interrupt=flow.interrupts or None, - session_state=_safe_serialize_session_continuation_state( - session, - agent, - shared_state_keys=set(flow.current_state).difference(protected_session_state_keys), - include_service_session_id=config.use_service_session, - ), ) - _save_tool_approval_state(session, approval_state_store, approval_thread_id) yield _build_run_finished_event(run_id=run_id, thread_id=thread_id, interrupts=flow.interrupts) return @@ -3323,18 +3335,11 @@ async def _run_agent_stream( # Generic resume requests carry only the synthesized response, so prepend # stored history unless this run already seeded raw messages from it. persisted_messages = snapshot_session.resume_seeded_messages(persisted_messages) - await snapshot_session.save( - messages=_bound_host_payload_history(_persistable_host_payload_history(persisted_messages)), + await save_thread_snapshot( + persisted_messages=persisted_messages, state=cast(dict[str, Any], make_json_safe(flow.current_state)) if flow.current_state else None, interrupt=None, - session_state=_safe_serialize_session_continuation_state( - session, - agent, - shared_state_keys=set(flow.current_state).difference(protected_session_state_keys), - include_service_session_id=config.use_service_session, - ), ) - _save_tool_approval_state(session, approval_state_store, approval_thread_id) yield _build_run_finished_event(run_id=run_id, thread_id=thread_id) return @@ -3350,6 +3355,18 @@ async def _run_agent_stream( latest_state_snapshot: dict[str, Any] | None = ( cast(dict[str, Any], make_json_safe(flow.current_state)) if flow.current_state else None ) + + async def save_flow_snapshot() -> None: + safe_point_event = _build_messages_snapshot(flow, snapshot_messages) + safe_point_messages = _event_messages_to_snapshot_dicts(list(safe_point_event.messages)) + if resume_payload is not None and not seeded_resume_from_snapshot and snapshot_seed_messages is None: + safe_point_messages = snapshot_session.resume_seeded_messages(safe_point_messages) + await save_thread_snapshot( + persisted_messages=safe_point_messages, + state=latest_state_snapshot, + interrupt=flow.interrupts or None, + ) + # Agent middleware can defer the inner run until streaming begins, so the # telemetry override must cover construction, stream resolution, and every pull. # Drive the A2UI runner when one is active (see the gate above); the original agent @@ -3369,6 +3386,8 @@ async def _run_agent_stream( stream = await _normalize_response_stream(response_stream) async for update in _iterate_with_context(stream, telemetry_context): + save_safe_point = update.finish_reason is not None + # Collect updates for structured output processing if response_format is not None: all_updates.append(update) @@ -3416,6 +3435,8 @@ async def _run_agent_stream( # Emit events for each content item for content in update.contents: content_type = getattr(content, "type", None) + if content_type in {"function_result", "function_approval_request"}: + save_safe_point = True logger.debug(f"Processing content type={content_type}, message_id={flow.message_id}") forwarded_reapproval_handled = False native_approval_result = False @@ -3525,6 +3546,14 @@ async def _run_agent_stream( if native_approval_result: native_approval_flow_result_ids.update(id(result) for result in flow.tool_results[result_offset:]) + if ( + snapshot_session.enabled + and save_safe_point + and not flow.waiting_for_approval + and not config.use_service_session + ): + await save_flow_snapshot() + # Stop if waiting for approval if flow.waiting_for_approval: break @@ -3557,6 +3586,8 @@ async def _run_agent_stream( if flow.waiting_for_approval and isinstance(stream, ResponseStream): await stream.get_final_response() + if snapshot_session.enabled: + await save_flow_snapshot() # If no updates at all, still emit RunStarted if not run_started_emitted: @@ -3758,16 +3789,9 @@ async def _run_agent_stream( # Generic resume requests carry only the synthesized response, so prepend # stored history unless this run already seeded raw messages from it. persisted_messages = snapshot_session.resume_seeded_messages(persisted_messages) - await snapshot_session.save( - messages=_bound_host_payload_history(_persistable_host_payload_history(persisted_messages)), + await save_thread_snapshot( + persisted_messages=persisted_messages, state=latest_state_snapshot, interrupt=flow.interrupts or None, - session_state=_safe_serialize_session_continuation_state( - session, - agent, - shared_state_keys=set(flow.current_state).difference(protected_session_state_keys), - include_service_session_id=config.use_service_session, - ), ) - _save_tool_approval_state(session, approval_state_store, approval_thread_id) yield _build_run_finished_event(run_id=run_id, thread_id=thread_id, interrupts=flow.interrupts) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py index 3b2241b48d0..fd8f5585bd3 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -4,6 +4,7 @@ from __future__ import annotations +import asyncio import copy import logging from collections.abc import AsyncGenerator, Sequence @@ -15,7 +16,7 @@ from agent_framework import CheckpointStorage, SupportsAgentRun, Workflow from fastapi import FastAPI, HTTPException from fastapi.params import Depends -from fastapi.responses import Response, StreamingResponse +from fastapi.responses import JSONResponse, Response, StreamingResponse from ._agent import AgentFrameworkAgent from ._approval_state import _APPROVAL_SCOPE_INPUT_KEY @@ -30,6 +31,7 @@ logger = logging.getLogger(__name__) +_DETACHED_STREAM_QUEUE_SIZE = 16 _KEEPALIVE_COMMENT = "keepalive" @@ -78,6 +80,14 @@ def _validate_keepalive_seconds(keepalive_seconds: float | None) -> None: raise ValueError("keepalive_seconds must be positive or None.") +def _is_snapshot_hydration_request(request: AGUIRequest, *, snapshot_persistence_active: bool) -> bool: + """Return whether a request only replays the latest stored snapshot.""" + if not snapshot_persistence_active or request.messages or request.resume is not None: + return False + forwarded_props = request.forwarded_props or {} + return not (forwarded_props.get("checkpoint_id") or forwarded_props.get("checkpointId")) + + def add_agent_framework_fastapi_endpoint( app: FastAPI, agent: SupportsAgentRun | AgentFrameworkAgent | Workflow | AgentFrameworkWorkflow, @@ -93,6 +103,7 @@ def add_agent_framework_fastapi_endpoint( checkpoint_storage: CheckpointStorage | None = None, keepalive_seconds: float | None = 15, a2ui_config: dict[str, Any] | None = None, + detached_runs: bool = False, ) -> None: """Add an AG-UI endpoint to a FastAPI app. @@ -128,6 +139,11 @@ def add_agent_framework_fastapi_endpoint( the surface-generation tool (``forwardedProps.injectA2UITool``). Keys: ``inject_a2ui_tool`` (backend opt-in override), ``default_catalog_id``, ``catalog``, ``guidelines``, ``recovery``, ``default_surface_id``. + detached_runs: Whether agent/workflow execution continues after the SSE client disconnects. Defaults to False. + When enabled, the endpoint owns a bounded background producer and completes runner persistence even if the + HTTP reader is cancelled. Request-scoped disposable resources may be released after disconnect, so detached + work must use values resolved before streaming rather than retaining request-owned clients or sessions. + This option does not provide resumable event replay. """ _validate_keepalive_seconds(keepalive_seconds) @@ -167,6 +183,26 @@ def add_agent_framework_fastapi_endpoint( snapshot_scope_resolver=snapshot_scope_resolver, ) + background_tasks: set[asyncio.Task[Any]] = set() + active_runs: dict[tuple[str | None, str], asyncio.Task[None]] = {} + + def retain_background_task(task: asyncio.Task[Any]) -> None: + background_tasks.add(task) + + def task_done(completed: asyncio.Task[Any]) -> None: + background_tasks.discard(completed) + if completed.cancelled(): + return + exception = completed.exception() + if exception is not None: + logger.error( + "[%s] Detached stream task failed", + path, + exc_info=(type(exception), exception, exception.__traceback__), + ) + + task.add_done_callback(task_done) + @app.post(path, tags=tags or ["AG-UI"], dependencies=dependencies, response_model=None) # type: ignore[arg-type] async def agent_endpoint(request_body: AGUIRequest) -> Response: """Handle AG-UI agent requests. @@ -177,12 +213,14 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: try: input_data = request_body.model_dump(exclude_none=True) snapshot_persistence_active = _get_snapshot_store(protocol_runner) is not None + snapshot_scope: str | None = None if snapshot_scope_resolver is not None: - snapshot_scope = snapshot_scope_resolver(request_body) - if isawaitable(snapshot_scope): - snapshot_scope = await snapshot_scope - if not isinstance(snapshot_scope, str) or not snapshot_scope: + resolved_scope = snapshot_scope_resolver(request_body) + if isawaitable(resolved_scope): + resolved_scope = await resolved_scope + if not isinstance(resolved_scope, str) or not resolved_scope: raise ValueError("snapshot_scope_resolver must return a non-empty string.") + snapshot_scope = resolved_scope input_data[_APPROVAL_SCOPE_INPUT_KEY] = snapshot_scope input_data[_SNAPSHOT_SCOPE_INPUT_KEY] = snapshot_scope if default_state: @@ -203,6 +241,23 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: logger.info(f"Received request at {path}: {input_data.get('run_id', 'no-run-id')}") keepalive_enabled = keepalive_seconds is not None + active_run_key: tuple[str | None, str] | None = None + if ( + detached_runs + and request_body.thread_id is not None + and not _is_snapshot_hydration_request( + request_body, + snapshot_persistence_active=snapshot_persistence_active, + ) + ): + active_run_key = (snapshot_scope, request_body.thread_id) + active_task = active_runs.get(active_run_key) + if active_task is not None and not active_task.done(): + return JSONResponse( + status_code=409, + content={"detail": "An AG-UI run is already active for this scoped thread."}, + ) + active_runs.pop(active_run_key, None) def prepare_frame(encoded: str) -> str | bytes: if keepalive_enabled: @@ -257,6 +312,53 @@ async def event_generator() -> AsyncGenerator[str | bytes]: except Exception: logger.exception("[%s] Failed to encode RUN_ERROR event", path) + async def drain_detached_stream(queue: asyncio.Queue[str | bytes | None]) -> None: + while await queue.get() is not None: + pass + + stream: AsyncGenerator[str | bytes] + if detached_runs: + queue: asyncio.Queue[str | bytes | None] = asyncio.Queue(maxsize=_DETACHED_STREAM_QUEUE_SIZE) + + async def produce_events() -> None: + try: + async for frame in event_generator(): + await queue.put(frame) + finally: + current_task = asyncio.current_task() + if active_run_key is not None and active_runs.get(active_run_key) is current_task: + active_runs.pop(active_run_key, None) + await queue.put(None) + + producer_task = asyncio.create_task( + produce_events(), + name=f"ag-ui-run-{input_data.get('run_id', 'generated')}", + ) + if active_run_key is not None: + active_runs[active_run_key] = producer_task + retain_background_task(producer_task) + + async def detached_event_generator() -> AsyncGenerator[str | bytes]: + completed = False + try: + while True: + item = await queue.get() + if item is None: + completed = True + return + yield item + finally: + if not completed and not producer_task.done(): + drain_task = asyncio.create_task( + drain_detached_stream(queue), + name=f"ag-ui-drain-{input_data.get('run_id', 'generated')}", + ) + retain_background_task(drain_task) + + stream = detached_event_generator() + else: + stream = event_generator() + headers = { "Cache-Control": "no-cache", "Connection": "keep-alive", @@ -267,14 +369,14 @@ async def event_generator() -> AsyncGenerator[str | bytes]: from sse_starlette.sse import EventSourceResponse return EventSourceResponse( - event_generator(), + stream, ping=cast(int, keepalive_seconds), ping_message_factory=lambda: ServerSentEvent(comment=_KEEPALIVE_COMMENT), headers=headers, media_type="text/event-stream", ) return StreamingResponse( - event_generator(), + stream, media_type="text/event-stream", headers=headers, ) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_snapshot_session.py b/python/packages/ag-ui/agent_framework_ag_ui/_snapshot_session.py index fa63dda39f1..976a2cc9168 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_snapshot_session.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_snapshot_session.py @@ -4,8 +4,8 @@ A ThreadSnapshotSession is opened once per run and owns every interaction with the AG-UI Thread Snapshot store: the load-once read, hydration replay, -the effective-state overlay, resume message seeding, and the save whose -storage failures must never surface on an already-streamed run. +the effective-state overlay, resume message seeding, and ordered safe-point +and terminal saves whose storage failures must never replace streamed events. """ from __future__ import annotations @@ -155,11 +155,11 @@ async def save( interrupt: list[dict[str, Any]] | None, session_state: dict[str, Any] | None, ) -> None: - """Commit the latest thread snapshot in one write when persistence is configured. + """Commit the latest thread snapshot when persistence is configured. - The run has already streamed by the time this is called, so a store - failure is logged and swallowed; the previous snapshot stays - authoritative for hydration. + Calls may follow already-emitted stream events, so a store failure is + logged and swallowed rather than changing the stream into a late + ``RUN_ERROR``. The previous snapshot stays authoritative for hydration. """ if self._store is None or self._scope is None: return diff --git a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py index 69bf104ec8a..e6a4af2d667 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py +++ b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py @@ -8,7 +8,7 @@ import subprocess import sys from collections import Counter -from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Sequence from contextvars import ContextVar from dataclasses import dataclass from inspect import signature @@ -17,7 +17,7 @@ from unittest.mock import AsyncMock, Mock import pytest -from ag_ui.core import MessagesSnapshotEvent, RunStartedEvent, StateSnapshotEvent +from ag_ui.core import BaseEvent, MessagesSnapshotEvent, RunFinishedEvent, RunStartedEvent, StateSnapshotEvent from agent_framework import ( Agent, AgentContext, @@ -129,6 +129,57 @@ async def send(message: ASGIMessage) -> None: await asyncio.wait_for(app(scope, cast(Receive, receive), cast(Send, send)), timeout=5) +async def _post_asgi_request( + app: FastAPI, + path: str, + payload: dict[str, Any], +) -> tuple[int, bytes]: + """Run one complete ASGI request and return its status and response body.""" + request_sent = False + response_complete = asyncio.Event() + status_code = 0 + chunks: list[bytes] = [] + body = json.dumps(payload).encode() + + async def receive() -> ASGIMessage: + nonlocal request_sent + if not request_sent: + request_sent = True + return {"type": "http.request", "body": body, "more_body": False} + await response_complete.wait() + return {"type": "http.disconnect"} + + async def send(message: ASGIMessage) -> None: + nonlocal status_code + if message["type"] == "http.response.start": + status_code = int(message["status"]) + return + if message["type"] != "http.response.body": + return + chunk = message.get("body", b"") + if isinstance(chunk, bytes): + chunks.append(chunk) + if not message.get("more_body", False): + response_complete.set() + + scope: Scope = { + "type": "http", + "asgi": {"version": "3.0"}, + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": path, + "raw_path": path.encode(), + "query_string": b"", + "root_path": "", + "headers": [(b"content-type", b"application/json"), (b"host", b"testserver")], + "client": ("testclient", 50000), + "server": ("testserver", 80), + } + await asyncio.wait_for(app(scope, cast(Receive, receive), cast(Send, send)), timeout=5) + return status_code, b"".join(chunks) + + def _run_finished_interrupts(event: dict[str, Any]) -> list[dict[str, Any]]: """Return canonical interrupts from an SSE RUN_FINISHED event.""" assert "interrupt" not in event @@ -1202,8 +1253,8 @@ def event_types_of(response: Any) -> list[str]: assert "TEXT_MESSAGE_CONTENT" in resumed_types -async def test_add_endpoint_accepts_keepalive_option_for_supported_runners(build_chat_client): - """Keepalive configuration is accepted at the endpoint seam for every supported runner shape.""" +async def test_add_endpoint_accepts_transport_options_for_supported_runners(build_chat_client): + """Endpoint transport configuration is accepted for every supported runner shape.""" @executor(id="start") async def start(message: Any, ctx: WorkflowContext[Any, Any]) -> None: @@ -1217,9 +1268,21 @@ async def start(message: Any, ctx: WorkflowContext[Any, Any]) -> None: name="wrapped", ) - add_agent_framework_fastapi_endpoint(app, raw_agent, path="/raw-agent", keepalive_seconds=0.5) + add_agent_framework_fastapi_endpoint( + app, + raw_agent, + path="/raw-agent", + keepalive_seconds=0.5, + detached_runs=True, + ) add_agent_framework_fastapi_endpoint(app, wrapped_agent, path="/wrapped-agent", keepalive_seconds=None) - add_agent_framework_fastapi_endpoint(app, workflow, path="/raw-workflow", keepalive_seconds=1.0) + add_agent_framework_fastapi_endpoint( + app, + workflow, + path="/raw-workflow", + keepalive_seconds=1.0, + detached_runs=True, + ) add_agent_framework_fastapi_endpoint( app, AgentFrameworkWorkflow(workflow=workflow), @@ -1242,6 +1305,20 @@ def test_add_endpoint_keepalive_default_is_enabled() -> None: assert parameter.default == 15 +def test_add_endpoint_detached_runs_default_is_disabled() -> None: + """Detached execution is opt-in so current disconnect cancellation semantics remain the default.""" + parameter = signature(add_agent_framework_fastapi_endpoint).parameters["detached_runs"] + + assert parameter.default is False + + +def test_add_endpoint_detached_runs_preserves_existing_positional_parameter_order() -> None: + """The new option is appended after existing parameters so positional a2ui_config callers do not shift.""" + parameters = list(signature(add_agent_framework_fastapi_endpoint).parameters) + + assert parameters.index("detached_runs") > parameters.index("a2ui_config") + + def test_add_endpoint_docstring_describes_keepalive_transport_behavior() -> None: """The public endpoint docs describe keepalive as transport comments, not AG-UI events.""" docstring = add_agent_framework_fastapi_endpoint.__doc__ @@ -1255,12 +1332,30 @@ def test_add_endpoint_docstring_describes_keepalive_transport_behavior() -> None assert "do not change AG-UI events" in normalized_docstring +def test_add_endpoint_docstring_describes_detached_run_behavior() -> None: + """The public endpoint docs distinguish detached completion from resumable replay.""" + docstring = add_agent_framework_fastapi_endpoint.__doc__ + + assert docstring is not None + normalized_docstring = " ".join(docstring.split()) + assert "detached_runs" in normalized_docstring + assert "Defaults to False" in normalized_docstring + assert "continues after the SSE client disconnects" in normalized_docstring + assert "does not provide resumable event replay" in normalized_docstring + + def test_keepalive_option_is_endpoint_owned() -> None: """Keepalive is endpoint transport configuration, not runner configuration.""" assert "keepalive_seconds" not in signature(AgentFrameworkAgent).parameters assert "keepalive_seconds" not in signature(AgentFrameworkWorkflow).parameters +def test_detached_runs_option_is_endpoint_owned() -> None: + """Detached execution is endpoint transport configuration, not runner configuration.""" + assert "detached_runs" not in signature(AgentFrameworkAgent).parameters + assert "detached_runs" not in signature(AgentFrameworkWorkflow).parameters + + def test_endpoint_module_import_does_not_import_sse_transport() -> None: """Importing endpoint helpers does not trigger sse-starlette's process-global transport hooks.""" import_check = ( @@ -1342,6 +1437,383 @@ async def stream_fn(messages: Any, options: Any, **kwargs: Any): assert "RUN_FINISHED" in event_types +@pytest.mark.parametrize("keepalive_seconds", [None, 0.01]) +async def test_endpoint_detached_run_completes_and_saves_after_client_disconnect( + streaming_chat_client_stub: Any, + keepalive_seconds: float | None, +) -> None: + """A disconnected reader cannot strand a bounded detached producer or skip final persistence.""" + release = asyncio.Event() + source_completed = asyncio.Event() + observed_scopes: list[str] = [] + request_scope: ContextVar[str] = ContextVar("detached-request-scope") + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del messages, options, kwargs + observed_scopes.append(request_scope.get()) + yield ChatResponseUpdate(contents=[Content.from_text(text="first")], role="assistant") + await release.wait() + for index in range(32): + yield ChatResponseUpdate( + contents=[Content.from_text(text=f"-chunk-{index}")], + role="assistant", + finish_reason="stop" if index == 31 else None, + ) + source_completed.set() + + def resolve_scope(request: AGUIRequest) -> str: + scope = str((request.state or {})["scope"]) + request_scope.set(scope) + return scope + + store = InMemoryAGUIThreadSnapshotStore() + agent = Agent( + name="detached", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached", + snapshot_store=store, + snapshot_scope_resolver=resolve_scope, + keepalive_seconds=keepalive_seconds, + detached_runs=True, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached", + { + "runId": "run-detached", + "threadId": "thread-detached", + "messages": [{"role": "user", "content": "Start"}], + "state": {"scope": "tenant-a"}, + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + assert not source_completed.is_set() + + release.set() + await asyncio.wait_for(source_completed.wait(), timeout=5) + + snapshot = None + for _ in range(100): + snapshot = await store.get(scope="tenant-a", thread_id="thread-detached") + if snapshot is not None and "chunk-31" in json.dumps(snapshot.messages): + break + await asyncio.sleep(0.01) + + assert snapshot is not None + assert "chunk-31" in json.dumps(snapshot.messages) + assert observed_scopes == ["tenant-a"] + + +async def test_endpoint_disconnect_still_cancels_run_by_default( + streaming_chat_client_stub: Any, +) -> None: + """Without the opt-in, disconnect retains the existing cancellation behavior.""" + release = asyncio.Event() + source_closed = asyncio.Event() + source_completed = False + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + nonlocal source_completed + del messages, options, kwargs + try: + yield ChatResponseUpdate(contents=[Content.from_text(text="first")], role="assistant") + await release.wait() + source_completed = True + yield ChatResponseUpdate( + contents=[Content.from_text(text="-finished")], + role="assistant", + finish_reason="stop", + ) + finally: + source_closed.set() + + store = InMemoryAGUIThreadSnapshotStore() + agent = Agent( + name="cancel-on-disconnect", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/cancel-on-disconnect", + snapshot_store=store, + snapshot_scope_resolver=lambda _request: "tenant-a", + keepalive_seconds=None, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/cancel-on-disconnect", + { + "runId": "run-cancelled", + "threadId": "thread-cancelled", + "messages": [{"role": "user", "content": "Start"}], + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + await asyncio.wait_for(source_closed.wait(), timeout=5) + + assert not source_completed + assert await store.get(scope="tenant-a", thread_id="thread-cancelled") is None + + +async def test_endpoint_detached_workflow_runner_completes_after_disconnect() -> None: + """The endpoint-owned producer keeps workflow runners alive as well as agent runners.""" + release = asyncio.Event() + completed = asyncio.Event() + + class BlockingWorkflowRunner(AgentFrameworkWorkflow): + def __init__(self) -> None: + self.snapshot_store = None + self.checkpoint_storage = None + + async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: + run_id = str(input_data["run_id"]) + thread_id = str(input_data["thread_id"]) + yield RunStartedEvent(run_id=run_id, thread_id=thread_id) + yield StateSnapshotEvent(snapshot={"started": True}) + await release.wait() + completed.set() + yield RunFinishedEvent(run_id=run_id, thread_id=thread_id) + + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + BlockingWorkflowRunner(), + path="/detached-workflow", + keepalive_seconds=None, + detached_runs=True, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached-workflow", + { + "runId": "workflow-run", + "threadId": "workflow-thread", + "messages": [{"role": "user", "content": "Start"}], + }, + event_type="STATE_SNAPSHOT", + ) + assert not completed.is_set() + + release.set() + await asyncio.wait_for(completed.wait(), timeout=5) + + +async def test_endpoint_detached_connected_failure_emits_run_error_and_completes( + streaming_chat_client_stub: Any, +) -> None: + """Detached producer failures retain the connected stream's RUN_ERROR contract.""" + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del messages, options, kwargs + if False: # pragma: no cover + yield ChatResponseUpdate() + raise RuntimeError("detached failure") + + agent = Agent( + name="detached-error", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached-error", + keepalive_seconds=None, + detached_runs=True, + ) + + status_code, body = await _post_asgi_request( + app, + "/detached-error", + { + "runId": "run-error", + "threadId": "thread-error", + "messages": [{"role": "user", "content": "Start"}], + }, + ) + + assert status_code == 200 + events = [json.loads(line[6:]) for line in body.decode().splitlines() if line.startswith("data: ")] + run_errors = [event for event in events if event.get("type") == "RUN_ERROR"] + assert len(run_errors) == 1 + assert run_errors[0]["code"] == "RuntimeError" + + +async def test_endpoint_detached_run_guards_active_scoped_thread_mutations( + streaming_chat_client_stub: Any, +) -> None: + """Active detached runs allow hydration but reject same-scope mutations until final persistence.""" + release = asyncio.Event() + completed_scopes: set[str] = set() + request_scope: ContextVar[str] = ContextVar("detached-guard-scope") + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del messages, options, kwargs + scope = request_scope.get() + yield ChatResponseUpdate(contents=[Content.from_text(text=f"started-{scope}")], role="assistant") + await release.wait() + yield ChatResponseUpdate( + contents=[Content.from_text(text=f"-finished-{scope}")], + role="assistant", + finish_reason="stop", + ) + completed_scopes.add(scope) + + def resolve_scope(request: AGUIRequest) -> str: + scope = str((request.state or {})["scope"]) + request_scope.set(scope) + return scope + + store = InMemoryAGUIThreadSnapshotStore() + await store.save( + scope="tenant-a", + thread_id="shared-thread", + snapshot=AGUIThreadSnapshot( + messages=[{"id": "stored-user", "role": "user", "content": "Stored"}], + ), + ) + agent = Agent( + name="detached-guard", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached-guard", + snapshot_store=store, + snapshot_scope_resolver=resolve_scope, + keepalive_seconds=None, + detached_runs=True, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached-guard", + { + "runId": "run-tenant-a", + "threadId": "shared-thread", + "messages": [{"role": "user", "content": "Start tenant A"}], + "state": {"scope": "tenant-a"}, + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + + hydration_status, hydration_body = await _post_asgi_request( + app, + "/detached-guard", + { + "runId": "hydrate-tenant-a", + "threadId": "shared-thread", + "messages": [], + "state": {"scope": "tenant-a"}, + }, + ) + assert hydration_status == 200 + assert b'"type":"MESSAGES_SNAPSHOT"' in hydration_body + + mutation_status, mutation_body = await _post_asgi_request( + app, + "/detached-guard", + { + "runId": "conflict-tenant-a", + "threadId": "shared-thread", + "messages": [{"role": "user", "content": "Race tenant A"}], + "state": {"scope": "tenant-a"}, + }, + ) + assert mutation_status == 409 + assert b"already active" in mutation_body + + checkpoint_status, _ = await _post_asgi_request( + app, + "/detached-guard", + { + "runId": "checkpoint-tenant-a", + "threadId": "shared-thread", + "messages": [], + "state": {"scope": "tenant-a"}, + "forwardedProps": {"checkpoint_id": "checkpoint-1"}, + }, + ) + assert checkpoint_status == 409 + + await _post_until_sse_event_then_disconnect( + app, + "/detached-guard", + { + "runId": "run-tenant-b", + "threadId": "shared-thread", + "messages": [{"role": "user", "content": "Start tenant B"}], + "state": {"scope": "tenant-b"}, + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + + release.set() + for _ in range(100): + tenant_a_snapshot = await store.get(scope="tenant-a", thread_id="shared-thread") + tenant_b_snapshot = await store.get(scope="tenant-b", thread_id="shared-thread") + if ( + tenant_a_snapshot is not None + and tenant_b_snapshot is not None + and "finished-tenant-a" in json.dumps(tenant_a_snapshot.messages) + and "finished-tenant-b" in json.dumps(tenant_b_snapshot.messages) + ): + break + await asyncio.sleep(0.01) + + assert completed_scopes == {"tenant-a", "tenant-b"} + + final_status = 409 + for _ in range(100): + final_status, _ = await _post_asgi_request( + app, + "/detached-guard", + { + "runId": "after-completion", + "threadId": "shared-thread", + "messages": [{"role": "user", "content": "Continue"}], + "state": {"scope": "tenant-a"}, + }, + ) + if final_status != 409: + break + await asyncio.sleep(0.01) + + assert final_status == 200 + + async def test_endpoint_keepalive_disabled_does_not_import_sse_transport(build_chat_client) -> None: """Disabled keepalive avoids importing sse-starlette's transport module.""" saved_sse_modules = { diff --git a/python/packages/ag-ui/tests/ag_ui/test_run.py b/python/packages/ag-ui/tests/ag_ui/test_run.py index ccd53dbb5e9..bf5836a8c71 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_run.py +++ b/python/packages/ag-ui/tests/ag_ui/test_run.py @@ -2,6 +2,7 @@ """Tests for _agent_run.py helper functions and FlowState.""" +import asyncio from typing import Any, cast import pytest @@ -20,7 +21,7 @@ ToolCallArgsEvent, ToolCallStartEvent, ) -from agent_framework import AgentResponseUpdate, Content, Message, ResponseStream +from agent_framework import AgentResponse, AgentResponseUpdate, Content, Message, ResponseStream from agent_framework.exceptions import AgentInvalidResponseException from conftest import StubAgent # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports] @@ -43,7 +44,7 @@ ApprovalStatus, ResumeDecision, ) -from agent_framework_ag_ui._approval_state import InMemoryAGUIApprovalStateStore +from agent_framework_ag_ui._approval_state import InMemoryAGUIApprovalStateStore, approval_state_thread_id from agent_framework_ag_ui._run_common import ( FlowState, _build_run_finished_event, @@ -2899,6 +2900,281 @@ async def test_service_session_rejects_disabled_provider_storage(): ] +async def test_snapshot_is_saved_at_model_roundtrip_safe_point_before_run_completion(): + """A configured snapshot store receives completed model output before the overall run ends.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + release = asyncio.Event() + safe_point_saved = asyncio.Event() + + class RecordingStore(InMemoryAGUIThreadSnapshotStore): + async def save(self, **kwargs: Any) -> None: + await super().save(**kwargs) + snapshot = kwargs["snapshot"] + if "round-one" in str(snapshot.messages): + safe_point_saved.set() + + stub = StubAgent() + original_run = stub.run + + def blocking_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates(): + yield AgentResponseUpdate( + contents=[Content.from_text(text="round-one")], + role="assistant", + finish_reason="stop", + ) + await release.wait() + yield AgentResponseUpdate( + contents=[Content.from_text(text="-round-two")], + role="assistant", + finish_reason="stop", + ) + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + stub.run = blocking_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = RecordingStore() + agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) + payload = { + "thread_id": "incremental-thread", + "run_id": "incremental-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + + async def collect_events() -> list[Any]: + return [event async for event in agent.run(payload)] + + run_task = asyncio.create_task(collect_events()) + await asyncio.wait_for(safe_point_saved.wait(), timeout=5) + + safe_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") + assert safe_snapshot is not None + assert "round-one" in str(safe_snapshot.messages) + assert "round-two" not in str(safe_snapshot.messages) + + release.set() + await asyncio.wait_for(run_task, timeout=5) + + final_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") + assert final_snapshot is not None + assert "round-one-round-two" in str(final_snapshot.messages) + + +async def test_snapshot_is_saved_after_tool_result_without_finish_reason(): + """A completed tool-result batch is durable even when its update has no finish reason.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + release = asyncio.Event() + tool_result_saved = asyncio.Event() + + class RecordingStore(InMemoryAGUIThreadSnapshotStore): + async def save(self, **kwargs: Any) -> None: + await super().save(**kwargs) + snapshot = kwargs["snapshot"] + if "tool-output" in str(snapshot.messages): + tool_result_saved.set() + + stub = StubAgent() + original_run = stub.run + + def blocking_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates(): + yield AgentResponseUpdate( + contents=[Content.from_function_call(call_id="call-1", name="lookup", arguments={})], + role="assistant", + finish_reason="tool_calls", + ) + yield AgentResponseUpdate( + contents=[Content.from_function_result(call_id="call-1", result="tool-output")], + role="tool", + ) + await release.wait() + yield AgentResponseUpdate( + contents=[Content.from_text(text="done")], + role="assistant", + finish_reason="stop", + ) + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + stub.run = blocking_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = RecordingStore() + agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) + payload = { + "thread_id": "tool-safe-point-thread", + "run_id": "tool-safe-point-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + + async def collect_events() -> list[Any]: + return [event async for event in agent.run(payload)] + + run_task = asyncio.create_task(collect_events()) + await asyncio.wait_for(tool_result_saved.wait(), timeout=5) + + safe_snapshot = await store.get(scope="tenant-a", thread_id="tool-safe-point-thread") + assert safe_snapshot is not None + assert "tool-output" in str(safe_snapshot.messages) + assert "done" not in str(safe_snapshot.messages) + + release.set() + await asyncio.wait_for(run_task, timeout=5) + + final_snapshot = await store.get(scope="tenant-a", thread_id="tool-safe-point-thread") + assert final_snapshot is not None + assert "tool-output" in str(final_snapshot.messages) + assert "done" in str(final_snapshot.messages) + + +async def test_service_session_snapshot_waits_for_terminal_continuation_state(): + """Service-session messages are not advanced before their provider continuation is final.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + continuation_set = asyncio.Event() + emit_next_update = asyncio.Event() + second_update_processed = asyncio.Event() + finish_run = asyncio.Event() + + stub = StubAgent() + original_run = stub.run + + def blocking_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + session = kwargs["session"] + + async def updates(): + yield AgentResponseUpdate( + contents=[Content.from_text(text="round-one")], + role="assistant", + finish_reason="stop", + ) + session.service_session_id = "provider-continuation" + continuation_set.set() + await emit_next_update.wait() + yield AgentResponseUpdate( + contents=[Content.from_text(text="-round-two")], + role="assistant", + finish_reason="stop", + ) + second_update_processed.set() + await finish_run.wait() + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + stub.run = blocking_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkAgent( + agent=stub, + use_service_session=True, + snapshot_store=store, + ) + payload = { + "thread_id": "service-safe-point-thread", + "run_id": "service-safe-point-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + + async def collect_events() -> list[Any]: + return [event async for event in agent.run(payload)] + + run_task = asyncio.create_task(collect_events()) + await asyncio.wait_for(continuation_set.wait(), timeout=5) + + assert await store.get(scope="tenant-a", thread_id="service-safe-point-thread") is None + + emit_next_update.set() + await asyncio.wait_for(second_update_processed.wait(), timeout=5) + + assert await store.get(scope="tenant-a", thread_id="service-safe-point-thread") is None + + finish_run.set() + await asyncio.wait_for(run_task, timeout=5) + + final_snapshot = await store.get(scope="tenant-a", thread_id="service-safe-point-thread") + assert final_snapshot is not None + assert "round-one-round-two" in str(final_snapshot.messages) + assert final_snapshot.session_state == { + "__ag_ui_provider_service_session_id": "provider-continuation", + } + + +async def test_interrupt_snapshot_is_saved_after_approval_lifecycle_registration(): + """An incremental interrupt snapshot never gets ahead of its server-side approval authority.""" + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + + approval_saved = asyncio.Event() + state_store = InMemoryAGUIApprovalStateStore() + scoped_thread_id = approval_state_thread_id(scope="tenant-a", thread_id="approval-thread") + + class RecordingStore(InMemoryAGUIThreadSnapshotStore): + async def save(self, **kwargs: Any) -> None: + snapshot = kwargs["snapshot"] + for interrupt in snapshot.interrupt or []: + occurrence = state_store.lifecycle.occurrence_for_alias( + thread_id=scoped_thread_id, + interrupt_id=str(interrupt["id"]), + ) + assert occurrence is not None + approval_saved.set() + await super().save(**kwargs) + + function_call = Content.from_function_call( + call_id="provider-call", + name="write_doc", + arguments={"content": "draft"}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=function_call, + ) + stub = StubAgent( + updates=[ + AgentResponseUpdate( + contents=[approval_request], + role="assistant", + finish_reason="tool_calls", + ) + ] + ) + store = RecordingStore() + config = AgentConfig(snapshot_store=store) + payload = { + "thread_id": "approval-thread", + "run_id": "approval-run", + "__ag_ui_snapshot_scope": "tenant-a", + "__ag_ui_approval_scope": "tenant-a", + "messages": [{"role": "user", "content": "Write"}], + } + + _ = [ + event + async for event in run_agent_stream( + payload, + stub, + config, + approval_state_store=state_store, + ) + ] + + assert approval_saved.is_set() + snapshot = await store.get(scope="tenant-a", thread_id="approval-thread") + assert snapshot is not None + assert snapshot.interrupt is not None + assert snapshot.interrupt[0]["id"] == "approval-occurrence" + + async def test_stateless_snapshot_excludes_only_provider_service_session_state(): """Stateless runs restore unrelated private state but not provider-owned continuation.""" from conftest import StubAgent # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports] From 8084ec67a1264435de662092eddc59ed1c62ec19 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Wed, 23 Sep 2026 11:36:26 +0200 Subject: [PATCH 2/5] Python: harden detached AG-UI run lifecycle --- python/packages/ag-ui/AGENTS.md | 15 +- python/packages/ag-ui/README.md | 16 +- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 30 +- .../ag-ui/agent_framework_ag_ui/_endpoint.py | 110 ++++++- .../ag-ui/tests/ag_ui/test_endpoint.py | 268 ++++++++++++++++++ python/packages/ag-ui/tests/ag_ui/test_run.py | 199 ++++++++++++- 6 files changed, 612 insertions(+), 26 deletions(-) diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index 297e665290e..85c542b8d37 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -112,15 +112,18 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. `state` entry, so conversation continuation is unreachable from client input by construction; keep it that way. - `confirm_changes` snapshot cleanup resolves the synthetic confirmation back to its original `function_call_id`; it must never concatenate unrelated tool results or record accepted changes without a matching real result. -- Configured stateless agent Thread Snapshot stores persist after complete model-roundtrip, tool-result, and approval - safe points, then persist the terminal state. Save only after the complete update and approval lifecycle side effects - are applied; do not write every text delta or schedule unordered background snapshot writes. Service-session and - workflow Thread Snapshots retain terminal-save cadence so replayable messages cannot advance without matching provider - continuation state; workflow checkpoints own incremental workflow runtime state. +- Configured stateless agent Thread Snapshot stores persist after finalized model-roundtrip, function/MCP tool-result, + and approval safe points, then persist the terminal state. A `finish_reason` is only a pending boundary: pull the + stream again so the inner finalizer and provider/context side effects complete before saving the preceding turn. Save + only after the complete update and approval lifecycle side effects are applied; do not write every text delta or + schedule unordered background snapshot writes. Service-session and workflow Thread Snapshots retain terminal-save + cadence so replayable messages cannot advance without matching provider continuation state; workflow checkpoints own + incremental workflow runtime state. - Disconnect-safe execution is endpoint-owned and opt-in through `add_agent_framework_fastapi_endpoint(detached_runs=True)`. Keep its producer queue bounded, retain and observe producer/drainer tasks, and let the producer own final snapshot/checkpoint/approval persistence after the SSE reader - disconnects. This mode discards unread events; it is not a resumable event log. + disconnects. Bound endpoint-wide producer admission and cancel abandoned producers after the configured timeout. + This mode discards unread events; it is not a resumable event log. - While a detached mutation is active, the endpoint may serve an empty Snapshot Hydrate Request but must reject another mutation for the same `(Snapshot Scope, threadId)` with HTTP 409. The guard is process-local and does not replace cross-replica coordination. diff --git a/python/packages/ag-ui/README.md b/python/packages/ag-ui/README.md index 4eaacff3d7b..60524f8152b 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -389,11 +389,12 @@ add_agent_framework_fastapi_endpoint( ) ``` -Configured stateless agent snapshot stores are updated after completed model roundtrips, tool-result batches, and -approval safe points, then written once more with the terminal run state. This limits progress loss during long agent -runs without persisting every streaming text delta. Service-session snapshots retain terminal-save cadence so -replayable messages cannot advance without their matching provider continuation state. Workflow Thread Snapshots also -keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow runtime state. +Configured stateless agent snapshot stores are updated after finalized model roundtrips, function/MCP tool-result +batches, and approval safe points, then written once more with the terminal run state. This limits progress loss during +long agent runs without persisting every streaming text delta. Service-session snapshots retain terminal-save cadence +so replayable messages cannot advance without their matching provider continuation state. Workflow Thread Snapshots +also keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow runtime +state. A frontend can then hydrate the latest stored snapshot for the scoped thread: @@ -418,12 +419,17 @@ add_agent_framework_fastapi_endpoint( snapshot_store=snapshot_store, snapshot_scope_resolver=resolve_snapshot_scope, detached_runs=True, + max_detached_runs=32, + detached_run_timeout_seconds=3600, ) ``` Detached execution uses a bounded endpoint-owned producer queue. While a detached mutating request is active, another mutating request for the same `(Snapshot Scope, threadId)` returns HTTP 409; an empty snapshot Hydrate Request remains allowed and returns the latest committed safe point. Equal Thread ids in different Snapshot Scopes remain independent. +Each endpoint registration retains at most `max_detached_runs` producers (32 by default); requests beyond that limit +receive HTTP 503. After a reader disconnects or never starts, `detached_run_timeout_seconds` cancels a stalled producer +and releases its capacity (one hour by default). This option does not add resumable event replay. Events emitted while no client is attached are consumed and discarded, so a reconnecting client recovers from Thread Snapshots rather than resuming the original SSE position. Applications diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 4a3e939fd6c..1b046fb3b97 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -3375,6 +3375,7 @@ async def save_flow_snapshot() -> None: telemetry_context = partial(_use_telemetry_conversation_id, telemetry_conversation_id) stream_completed = False native_approval_flow_result_ids: set[int] = set() + pending_model_safe_point = False try: with telemetry_context(): for queued_executions in forwarded_executions.values(): @@ -3386,7 +3387,14 @@ async def save_flow_snapshot() -> None: stream = await _normalize_response_stream(response_stream) async for update in _iterate_with_context(stream, telemetry_context): - save_safe_point = update.finish_reason is not None + if snapshot_session.enabled and pending_model_safe_point and not config.use_service_session: + # Pulling the next update finalizes the preceding model turn and + # completes provider/context side effects before this snapshot. + await save_flow_snapshot() + pending_model_safe_point = False + + model_safe_point = update.finish_reason is not None + result_safe_point = False # Collect updates for structured output processing if response_format is not None: @@ -3435,8 +3443,12 @@ async def save_flow_snapshot() -> None: # Emit events for each content item for content in update.contents: content_type = getattr(content, "type", None) - if content_type in {"function_result", "function_approval_request"}: - save_safe_point = True + if content_type in { + "function_result", + "function_approval_request", + "mcp_server_tool_result", + }: + result_safe_point = True logger.debug(f"Processing content type={content_type}, message_id={flow.message_id}") forwarded_reapproval_handled = False native_approval_result = False @@ -3548,11 +3560,14 @@ async def save_flow_snapshot() -> None: if ( snapshot_session.enabled - and save_safe_point + and result_safe_point + and not model_safe_point and not flow.waiting_for_approval and not config.use_service_session ): await save_flow_snapshot() + if model_safe_point and not flow.waiting_for_approval and not config.use_service_session: + pending_model_safe_point = True # Stop if waiting for approval if flow.waiting_for_approval: @@ -3584,9 +3599,10 @@ async def save_flow_snapshot() -> None: approval_state_store.lifecycle.recover_unfinished(intent) forwarded_executions.clear() - if flow.waiting_for_approval and isinstance(stream, ResponseStream): - await stream.get_final_response() - if snapshot_session.enabled: + if flow.waiting_for_approval: + if isinstance(stream, ResponseStream): + await stream.get_final_response() + if snapshot_session.enabled and not config.use_service_session: await save_flow_snapshot() # If no updates at all, still emit RunStarted diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py index fd8f5585bd3..f05ebfc9205 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -8,6 +8,7 @@ import copy import logging from collections.abc import AsyncGenerator, Sequence +from contextlib import suppress from inspect import isawaitable from typing import Any, cast @@ -20,6 +21,7 @@ from ._agent import AgentFrameworkAgent from ._approval_state import _APPROVAL_SCOPE_INPUT_KEY +from ._run_common import _extract_resume_payload from ._snapshots import ( _DEFAULT_STATE_INPUT_KEY, _SNAPSHOT_SCOPE_INPUT_KEY, @@ -31,6 +33,7 @@ logger = logging.getLogger(__name__) +_DETACHED_READER_START_TIMEOUT_SECONDS = 1.0 _DETACHED_STREAM_QUEUE_SIZE = 16 _KEEPALIVE_COMMENT = "keepalive" @@ -80,9 +83,21 @@ def _validate_keepalive_seconds(keepalive_seconds: float | None) -> None: raise ValueError("keepalive_seconds must be positive or None.") -def _is_snapshot_hydration_request(request: AGUIRequest, *, snapshot_persistence_active: bool) -> bool: +def _validate_detached_run_options(max_detached_runs: int, detached_run_timeout_seconds: float) -> None: + if max_detached_runs < 1: + raise ValueError("max_detached_runs must be greater than 0.") + if detached_run_timeout_seconds <= 0: + raise ValueError("detached_run_timeout_seconds must be positive.") + + +def _is_snapshot_hydration_request( + request: AGUIRequest, + input_data: dict[str, Any], + *, + snapshot_persistence_active: bool, +) -> bool: """Return whether a request only replays the latest stored snapshot.""" - if not snapshot_persistence_active or request.messages or request.resume is not None: + if not snapshot_persistence_active or request.messages or _extract_resume_payload(input_data) is not None: return False forwarded_props = request.forwarded_props or {} return not (forwarded_props.get("checkpoint_id") or forwarded_props.get("checkpointId")) @@ -104,6 +119,8 @@ def add_agent_framework_fastapi_endpoint( keepalive_seconds: float | None = 15, a2ui_config: dict[str, Any] | None = None, detached_runs: bool = False, + max_detached_runs: int = 32, + detached_run_timeout_seconds: float = 3600, ) -> None: """Add an AG-UI endpoint to a FastAPI app. @@ -144,8 +161,13 @@ def add_agent_framework_fastapi_endpoint( HTTP reader is cancelled. Request-scoped disposable resources may be released after disconnect, so detached work must use values resolved before streaming rather than retaining request-owned clients or sessions. This option does not provide resumable event replay. + max_detached_runs: Maximum number of detached producers retained by this endpoint registration. Defaults to 32. + Additional requests receive HTTP 503 until capacity is released. + detached_run_timeout_seconds: Maximum time a producer may remain active after its SSE reader disconnects or + never starts. Defaults to 3600 seconds. Expired producers are cancelled and release endpoint capacity. """ _validate_keepalive_seconds(keepalive_seconds) + _validate_detached_run_options(max_detached_runs, detached_run_timeout_seconds) protocol_runner: AgentFrameworkAgent | AgentFrameworkWorkflow if isinstance(agent, AgentFrameworkWorkflow): @@ -184,6 +206,7 @@ def add_agent_framework_fastapi_endpoint( ) background_tasks: set[asyncio.Task[Any]] = set() + producer_tasks: set[asyncio.Task[None]] = set() active_runs: dict[tuple[str | None, str], asyncio.Task[None]] = {} def retain_background_task(task: asyncio.Task[Any]) -> None: @@ -247,6 +270,7 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: and request_body.thread_id is not None and not _is_snapshot_hydration_request( request_body, + input_data, snapshot_persistence_active=snapshot_persistence_active, ) ): @@ -258,6 +282,15 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: content={"detail": "An AG-UI run is already active for this scoped thread."}, ) active_runs.pop(active_run_key, None) + if detached_runs: + for completed_task in tuple(producer_tasks): + if completed_task.done(): + producer_tasks.discard(completed_task) + if len(producer_tasks) >= max_detached_runs: + return JSONResponse( + status_code=503, + content={"detail": "AG-UI detached run capacity is exhausted."}, + ) def prepare_frame(encoded: str) -> str | bytes: if keepalive_enabled: @@ -319,16 +352,35 @@ async def drain_detached_stream(queue: asyncio.Queue[str | bytes | None]) -> Non stream: AsyncGenerator[str | bytes] if detached_runs: queue: asyncio.Queue[str | bytes | None] = asyncio.Queue(maxsize=_DETACHED_STREAM_QUEUE_SIZE) + reader_started = asyncio.Event() + reader_abandoned = asyncio.Event() async def produce_events() -> None: try: async for frame in event_generator(): + if reader_abandoned.is_set(): + continue + if not reader_started.is_set(): + try: + queue.put_nowait(frame) + except asyncio.QueueFull: + reader_abandoned.set() + while True: + try: + queue.get_nowait() + except asyncio.QueueEmpty: + break + continue await queue.put(frame) + except asyncio.CancelledError: + reader_abandoned.set() + raise finally: current_task = asyncio.current_task() if active_run_key is not None and active_runs.get(active_run_key) is current_task: active_runs.pop(active_run_key, None) - await queue.put(None) + if reader_started.is_set() or not reader_abandoned.is_set(): + await queue.put(None) producer_task = asyncio.create_task( produce_events(), @@ -336,11 +388,62 @@ async def produce_events() -> None: ) if active_run_key is not None: active_runs[active_run_key] = producer_task + producer_tasks.add(producer_task) + producer_task.add_done_callback(producer_tasks.discard) retain_background_task(producer_task) + async def expire_abandoned_run() -> None: + if not reader_started.is_set(): + try: + await asyncio.wait_for( + reader_started.wait(), + timeout=_DETACHED_READER_START_TIMEOUT_SECONDS, + ) + except asyncio.TimeoutError: + reader_abandoned.set() + if producer_task.done(): + return + if not reader_abandoned.is_set(): + abandoned_wait = asyncio.create_task(reader_abandoned.wait()) + done, _ = await asyncio.wait( + {producer_task, abandoned_wait}, + return_when=asyncio.FIRST_COMPLETED, + ) + if producer_task in done: + abandoned_wait.cancel() + with suppress(asyncio.CancelledError): + await abandoned_wait + return + if producer_task.done(): + return + try: + await asyncio.wait_for( + asyncio.shield(producer_task), + timeout=detached_run_timeout_seconds, + ) + except asyncio.TimeoutError: + logger.error( + "[%s] Detached run exceeded %.1f seconds after reader disconnect; cancelling", + path, + detached_run_timeout_seconds, + ) + reader_abandoned.set() + producer_task.cancel() + with suppress(asyncio.CancelledError): + await producer_task + + expiry_task = asyncio.create_task( + expire_abandoned_run(), + name=f"ag-ui-expiry-{input_data.get('run_id', 'generated')}", + ) + retain_background_task(expiry_task) + async def detached_event_generator() -> AsyncGenerator[str | bytes]: completed = False try: + reader_started.set() + if reader_abandoned.is_set(): + return while True: item = await queue.get() if item is None: @@ -349,6 +452,7 @@ async def detached_event_generator() -> AsyncGenerator[str | bytes]: yield item finally: if not completed and not producer_task.done(): + reader_abandoned.set() drain_task = asyncio.create_task( drain_detached_stream(queue), name=f"ag-ui-drain-{input_data.get('run_id', 'generated')}", diff --git a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py index e6a4af2d667..ef724bc929c 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py +++ b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py @@ -55,6 +55,8 @@ from conftest import StubAgent # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports] from fastapi import FastAPI, Header, HTTPException from fastapi.params import Depends +from fastapi.responses import StreamingResponse +from fastapi.routing import APIRoute from fastapi.testclient import TestClient from starlette.types import Message as ASGIMessage from starlette.types import Receive, Scope, Send @@ -1312,6 +1314,14 @@ def test_add_endpoint_detached_runs_default_is_disabled() -> None: assert parameter.default is False +def test_add_endpoint_detached_run_limits_have_bounded_defaults() -> None: + """Detached producers have finite admission and post-disconnect lifetime defaults.""" + parameters = signature(add_agent_framework_fastapi_endpoint).parameters + + assert parameters["max_detached_runs"].default == 32 + assert parameters["detached_run_timeout_seconds"].default == 3600 + + def test_add_endpoint_detached_runs_preserves_existing_positional_parameter_order() -> None: """The new option is appended after existing parameters so positional a2ui_config callers do not shift.""" parameters = list(signature(add_agent_framework_fastapi_endpoint).parameters) @@ -1342,6 +1352,8 @@ def test_add_endpoint_docstring_describes_detached_run_behavior() -> None: assert "Defaults to False" in normalized_docstring assert "continues after the SSE client disconnects" in normalized_docstring assert "does not provide resumable event replay" in normalized_docstring + assert "Maximum number of detached producers" in normalized_docstring + assert "Maximum time a producer may remain active" in normalized_docstring def test_keepalive_option_is_endpoint_owned() -> None: @@ -1354,6 +1366,10 @@ def test_detached_runs_option_is_endpoint_owned() -> None: """Detached execution is endpoint transport configuration, not runner configuration.""" assert "detached_runs" not in signature(AgentFrameworkAgent).parameters assert "detached_runs" not in signature(AgentFrameworkWorkflow).parameters + assert "max_detached_runs" not in signature(AgentFrameworkAgent).parameters + assert "max_detached_runs" not in signature(AgentFrameworkWorkflow).parameters + assert "detached_run_timeout_seconds" not in signature(AgentFrameworkAgent).parameters + assert "detached_run_timeout_seconds" not in signature(AgentFrameworkWorkflow).parameters def test_endpoint_module_import_does_not_import_sse_transport() -> None: @@ -1515,6 +1531,68 @@ def resolve_scope(request: AGUIRequest) -> str: assert observed_scopes == ["tenant-a"] +async def test_endpoint_detached_run_completes_when_response_stream_never_starts( + streaming_chat_client_stub: Any, +) -> None: + """An unstarted response iterator cannot strand a producer behind its bounded queue.""" + source_completed = asyncio.Event() + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del messages, options, kwargs + for index in range(32): + yield ChatResponseUpdate( + contents=[Content.from_text(text=f"chunk-{index}")], + role="assistant", + finish_reason="stop" if index == 31 else None, + ) + source_completed.set() + + store = InMemoryAGUIThreadSnapshotStore() + agent = Agent( + name="unstarted-reader", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/unstarted-reader", + snapshot_store=store, + snapshot_scope_resolver=lambda _request: "tenant-a", + keepalive_seconds=None, + detached_runs=True, + ) + route = next(route for route in app.routes if getattr(route, "path", None) == "/unstarted-reader") + assert isinstance(route, APIRoute) + + response = await route.endpoint( + AGUIRequest.model_validate( + { + "runId": "unstarted-run", + "threadId": "unstarted-thread", + "messages": [{"role": "user", "content": "Start"}], + } + ) + ) + assert isinstance(response, StreamingResponse) + + await asyncio.wait_for(source_completed.wait(), timeout=5) + snapshot = None + for _ in range(100): + snapshot = await store.get(scope="tenant-a", thread_id="unstarted-thread") + if snapshot is not None and "chunk-31" in json.dumps(snapshot.messages): + break + await asyncio.sleep(0.01) + + assert snapshot is not None + assert "chunk-31" in json.dumps(snapshot.messages) + + async def test_endpoint_disconnect_still_cancels_run_by_default( streaming_chat_client_stub: Any, ) -> None: @@ -1664,6 +1742,157 @@ async def stream_fn( assert run_errors[0]["code"] == "RuntimeError" +async def test_endpoint_detached_run_admission_is_bounded( + streaming_chat_client_stub: Any, +) -> None: + """A caller cannot retain more detached producers than the endpoint limit.""" + release = asyncio.Event() + first_started = asyncio.Event() + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del options, kwargs + text = messages[-1].text + yield ChatResponseUpdate(contents=[Content.from_text(text=f"started-{text}")], role="assistant") + if text == "first": + first_started.set() + await release.wait() + yield ChatResponseUpdate( + contents=[Content.from_text(text="-done")], + role="assistant", + finish_reason="stop", + ) + + agent = Agent( + name="detached-capacity", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached-capacity", + keepalive_seconds=None, + detached_runs=True, + max_detached_runs=1, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached-capacity", + { + "runId": "first-run", + "threadId": "first-thread", + "messages": [{"role": "user", "content": "first"}], + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + await asyncio.wait_for(first_started.wait(), timeout=5) + + rejected_status, rejected_body = await _post_asgi_request( + app, + "/detached-capacity", + { + "runId": "second-run", + "threadId": "second-thread", + "messages": [{"role": "user", "content": "second"}], + }, + ) + assert rejected_status == 503 + assert b"capacity is exhausted" in rejected_body + + release.set() + accepted_status = 503 + for _ in range(100): + accepted_status, _ = await _post_asgi_request( + app, + "/detached-capacity", + { + "runId": "second-run", + "threadId": "second-thread", + "messages": [{"role": "user", "content": "second"}], + }, + ) + if accepted_status != 503: + break + await asyncio.sleep(0.01) + + assert accepted_status == 200 + + +async def test_endpoint_detached_run_expires_after_reader_disconnect( + streaming_chat_client_stub: Any, +) -> None: + """A stalled abandoned run is cancelled and releases its scoped-thread slot.""" + first_closed = asyncio.Event() + run_count = 0 + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + nonlocal run_count + del messages, options, kwargs + run_count += 1 + if run_count == 1: + try: + yield ChatResponseUpdate(contents=[Content.from_text(text="started")], role="assistant") + await asyncio.Event().wait() + finally: + first_closed.set() + return + yield ChatResponseUpdate( + contents=[Content.from_text(text="retry-complete")], + role="assistant", + finish_reason="stop", + ) + + agent = Agent( + name="detached-expiry", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached-expiry", + keepalive_seconds=None, + detached_runs=True, + detached_run_timeout_seconds=0.05, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached-expiry", + { + "runId": "expiring-run", + "threadId": "expiring-thread", + "messages": [{"role": "user", "content": "first"}], + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + await asyncio.wait_for(first_closed.wait(), timeout=5) + + retry_status, retry_body = await _post_asgi_request( + app, + "/detached-expiry", + { + "runId": "retry-run", + "threadId": "expiring-thread", + "messages": [{"role": "user", "content": "retry"}], + }, + ) + + assert retry_status == 200 + assert b"retry-complete" in retry_body + + async def test_endpoint_detached_run_guards_active_scoped_thread_mutations( streaming_chat_client_stub: Any, ) -> None: @@ -1768,6 +1997,23 @@ def resolve_scope(request: AGUIRequest) -> str: ) assert checkpoint_status == 409 + for forwarded_props in ( + {"resume": [{"interruptId": "approval-1", "status": "resolved", "payload": {"approved": True}}]}, + {"command": {"resume": [{"interruptId": "approval-1", "status": "resolved", "payload": {"approved": True}}]}}, + ): + forwarded_resume_status, _ = await _post_asgi_request( + app, + "/detached-guard", + { + "runId": "forwarded-resume-tenant-a", + "threadId": "shared-thread", + "messages": [], + "state": {"scope": "tenant-a"}, + "forwardedProps": forwarded_props, + }, + ) + assert forwarded_resume_status == 409 + await _post_until_sse_event_then_disconnect( app, "/detached-guard", @@ -1852,6 +2098,28 @@ def test_add_endpoint_rejects_non_positive_keepalive_interval(build_chat_client, add_agent_framework_fastapi_endpoint(app, agent, path="/invalid", keepalive_seconds=keepalive_seconds) +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"max_detached_runs": 0}, "max_detached_runs must be greater than 0"), + ({"max_detached_runs": -1}, "max_detached_runs must be greater than 0"), + ({"detached_run_timeout_seconds": 0}, "detached_run_timeout_seconds must be positive"), + ({"detached_run_timeout_seconds": -1}, "detached_run_timeout_seconds must be positive"), + ], +) +def test_add_endpoint_rejects_invalid_detached_run_limits( + build_chat_client: Any, + kwargs: dict[str, Any], + message: str, +) -> None: + """Detached admission and expiry configuration must remain finite and positive.""" + app = FastAPI() + agent = Agent(name="test", instructions="Test agent", client=build_chat_client()) + + with pytest.raises(ValueError, match=message): + add_agent_framework_fastapi_endpoint(app, agent, path="/invalid-detached", **kwargs) + + async def test_endpoint_with_state_schema(build_chat_client): """Test endpoint with state_schema parameter.""" app = FastAPI() diff --git a/python/packages/ag-ui/tests/ag_ui/test_run.py b/python/packages/ag-ui/tests/ag_ui/test_run.py index bf5836a8c71..f875eba5f46 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_run.py +++ b/python/packages/ag-ui/tests/ag_ui/test_run.py @@ -2901,11 +2901,13 @@ async def test_service_session_rejects_disabled_provider_storage(): async def test_snapshot_is_saved_at_model_roundtrip_safe_point_before_run_completion(): - """A configured snapshot store receives completed model output before the overall run ends.""" + """Model output is saved only after the stream advances through turn finalization.""" from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore - release = asyncio.Event() + emit_next_update = asyncio.Event() safe_point_saved = asyncio.Event() + finish_run = asyncio.Event() + turn_finalized = asyncio.Event() class RecordingStore(InMemoryAGUIThreadSnapshotStore): async def save(self, **kwargs: Any) -> None: @@ -2920,6 +2922,7 @@ async def save(self, **kwargs: Any) -> None: def blocking_run(*args: Any, **kwargs: Any) -> Any: if not kwargs.get("stream", False): return original_run(*args, **kwargs) + session = kwargs["session"] async def updates(): yield AgentResponseUpdate( @@ -2927,12 +2930,14 @@ async def updates(): role="assistant", finish_reason="stop", ) - await release.wait() + session.state["turn_finalized"] = True + turn_finalized.set() + await emit_next_update.wait() yield AgentResponseUpdate( contents=[Content.from_text(text="-round-two")], role="assistant", - finish_reason="stop", ) + await finish_run.wait() return ResponseStream(updates(), finalizer=AgentResponse.from_updates) @@ -2950,14 +2955,19 @@ async def collect_events() -> list[Any]: return [event async for event in agent.run(payload)] run_task = asyncio.create_task(collect_events()) + await asyncio.wait_for(turn_finalized.wait(), timeout=5) + assert await store.get(scope="tenant-a", thread_id="incremental-thread") is None + + emit_next_update.set() await asyncio.wait_for(safe_point_saved.wait(), timeout=5) safe_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") assert safe_snapshot is not None assert "round-one" in str(safe_snapshot.messages) assert "round-two" not in str(safe_snapshot.messages) + assert safe_snapshot.session_state == {"turn_finalized": True} - release.set() + finish_run.set() await asyncio.wait_for(run_task, timeout=5) final_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") @@ -3035,6 +3045,78 @@ async def collect_events() -> list[Any]: assert "done" in str(final_snapshot.messages) +async def test_snapshot_is_saved_after_mcp_tool_result_without_finish_reason(): + """An MCP tool-result batch is durable even when its update has no finish reason.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + release = asyncio.Event() + tool_result_saved = asyncio.Event() + + class RecordingStore(InMemoryAGUIThreadSnapshotStore): + async def save(self, **kwargs: Any) -> None: + await super().save(**kwargs) + snapshot = kwargs["snapshot"] + if "mcp-output" in str(snapshot.messages): + tool_result_saved.set() + + stub = StubAgent() + original_run = stub.run + + def blocking_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates(): + yield AgentResponseUpdate( + contents=[ + Content.from_mcp_server_tool_call( + call_id="mcp-1", + tool_name="lookup", + server_name="server", + arguments={}, + ) + ], + role="assistant", + finish_reason="tool_calls", + ) + yield AgentResponseUpdate( + contents=[Content.from_mcp_server_tool_result(call_id="mcp-1", output="mcp-output")], + role="tool", + ) + await release.wait() + yield AgentResponseUpdate( + contents=[Content.from_text(text="done")], + role="assistant", + finish_reason="stop", + ) + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + stub.run = blocking_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = RecordingStore() + agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) + payload = { + "thread_id": "mcp-safe-point-thread", + "run_id": "mcp-safe-point-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + + async def collect_events() -> list[Any]: + return [event async for event in agent.run(payload)] + + run_task = asyncio.create_task(collect_events()) + await asyncio.wait_for(tool_result_saved.wait(), timeout=5) + + safe_snapshot = await store.get(scope="tenant-a", thread_id="mcp-safe-point-thread") + assert safe_snapshot is not None + assert "mcp-output" in str(safe_snapshot.messages) + assert "done" not in str(safe_snapshot.messages) + + release.set() + await asyncio.wait_for(run_task, timeout=5) + + async def test_service_session_snapshot_waits_for_terminal_continuation_state(): """Service-session messages are not advanced before their provider continuation is final.""" from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore @@ -3175,6 +3257,113 @@ async def save(self, **kwargs: Any) -> None: assert snapshot.interrupt[0]["id"] == "approval-occurrence" +async def test_plain_async_iterable_persists_waiting_approval_snapshot(): + """A non-ResponseStream runner still persists its approval interrupt before finishing.""" + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + + approval_saved = asyncio.Event() + + class RecordingStore(InMemoryAGUIThreadSnapshotStore): + async def save(self, **kwargs: Any) -> None: + await super().save(**kwargs) + if kwargs["snapshot"].interrupt: + approval_saved.set() + + function_call = Content.from_function_call( + call_id="plain-call", + name="write_doc", + arguments={"content": "draft"}, + id="plain-approval", + ) + approval_request = Content.from_function_approval_request( + id="plain-approval", + function_call=function_call, + ) + stub = StubAgent() + original_run = stub.run + + def plain_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates(): + yield AgentResponseUpdate( + contents=[approval_request], + role="assistant", + finish_reason="tool_calls", + ) + + return updates() + + stub.run = plain_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = RecordingStore() + config = AgentConfig(snapshot_store=store) + payload = { + "thread_id": "plain-approval-thread", + "run_id": "plain-approval-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Write"}], + } + + _ = [event async for event in run_agent_stream(payload, stub, config)] + + assert approval_saved.is_set() + + +async def test_service_session_approval_snapshot_is_terminal_only(): + """Approval interruption does not add an intermediate service-session snapshot write.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + class CountingStore(InMemoryAGUIThreadSnapshotStore): + def __init__(self) -> None: + super().__init__() + self.save_count = 0 + + async def save(self, **kwargs: Any) -> None: + self.save_count += 1 + await super().save(**kwargs) + + function_call = Content.from_function_call( + call_id="service-call", + name="write_doc", + arguments={"content": "draft"}, + id="service-approval", + ) + approval_request = Content.from_function_approval_request( + id="service-approval", + function_call=function_call, + ) + store = CountingStore() + agent = AgentFrameworkAgent( + agent=StubAgent( + updates=[ + AgentResponseUpdate( + contents=[approval_request], + role="assistant", + finish_reason="tool_calls", + ) + ] + ), + use_service_session=True, + service_session_id_from_thread_id=True, + snapshot_store=store, + ) + + _ = [ + event + async for event in agent.run( + { + "thread_id": "service-approval-thread", + "run_id": "service-approval-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Write"}], + } + ) + ] + + assert store.save_count == 1 + + async def test_stateless_snapshot_excludes_only_provider_service_session_state(): """Stateless runs restore unrelated private state but not provider-owned continuation.""" from conftest import StubAgent # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports] From ea839d799357fcaeb483752991c84f8e0341bac3 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Thu, 24 Sep 2026 12:18:05 +0200 Subject: [PATCH 3/5] Python: keep AG-UI hydration outside detached capacity --- python/packages/ag-ui/AGENTS.md | 5 +- python/packages/ag-ui/README.md | 3 +- .../ag-ui/agent_framework_ag_ui/_endpoint.py | 19 ++--- .../ag-ui/tests/ag_ui/test_endpoint.py | 76 +++++++++++++++++++ 4 files changed, 89 insertions(+), 14 deletions(-) diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index 85c542b8d37..d7fac45e1dd 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -125,8 +125,9 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. disconnects. Bound endpoint-wide producer admission and cancel abandoned producers after the configured timeout. This mode discards unread events; it is not a resumable event log. - While a detached mutation is active, the endpoint may serve an empty Snapshot Hydrate Request but must reject another - mutation for the same `(Snapshot Scope, threadId)` with HTTP 409. The guard is process-local and does not replace - cross-replica coordination. + mutation for the same `(Snapshot Scope, threadId)` with HTTP 409. Classify hydration once and keep it on the direct + response path so it bypasses detached admission and producer wrapping. The guard is process-local and does not + replace cross-replica coordination. - Detached runs may outlive FastAPI request-scoped disposable resources. Resolve authorization and Snapshot Scope before spawning the producer, and do not rely on request-owned clients or sessions remaining open after disconnect. - SSE keepalive is endpoint-owned transport behavior configured through diff --git a/python/packages/ag-ui/README.md b/python/packages/ag-ui/README.md index 60524f8152b..32ba741b43e 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -426,7 +426,8 @@ add_agent_framework_fastapi_endpoint( Detached execution uses a bounded endpoint-owned producer queue. While a detached mutating request is active, another mutating request for the same `(Snapshot Scope, threadId)` returns HTTP 409; an empty snapshot Hydrate Request remains -allowed and returns the latest committed safe point. Equal Thread ids in different Snapshot Scopes remain independent. +allowed, bypasses detached producer capacity, and returns the latest committed safe point. Equal Thread ids in +different Snapshot Scopes remain independent. Each endpoint registration retains at most `max_detached_runs` producers (32 by default); requests beyond that limit receive HTTP 503. After a reader disconnects or never starts, `detached_run_timeout_seconds` cancels a stalled producer and releases its capacity (one hour by default). diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py index f05ebfc9205..74e1b08ca94 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -264,16 +264,13 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: logger.info(f"Received request at {path}: {input_data.get('run_id', 'no-run-id')}") keepalive_enabled = keepalive_seconds is not None + snapshot_hydration_request = _is_snapshot_hydration_request( + request_body, + input_data, + snapshot_persistence_active=snapshot_persistence_active, + ) active_run_key: tuple[str | None, str] | None = None - if ( - detached_runs - and request_body.thread_id is not None - and not _is_snapshot_hydration_request( - request_body, - input_data, - snapshot_persistence_active=snapshot_persistence_active, - ) - ): + if detached_runs and request_body.thread_id is not None and not snapshot_hydration_request: active_run_key = (snapshot_scope, request_body.thread_id) active_task = active_runs.get(active_run_key) if active_task is not None and not active_task.done(): @@ -282,7 +279,7 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: content={"detail": "An AG-UI run is already active for this scoped thread."}, ) active_runs.pop(active_run_key, None) - if detached_runs: + if detached_runs and not snapshot_hydration_request: for completed_task in tuple(producer_tasks): if completed_task.done(): producer_tasks.discard(completed_task) @@ -350,7 +347,7 @@ async def drain_detached_stream(queue: asyncio.Queue[str | bytes | None]) -> Non pass stream: AsyncGenerator[str | bytes] - if detached_runs: + if detached_runs and not snapshot_hydration_request: queue: asyncio.Queue[str | bytes | None] = asyncio.Queue(maxsize=_DETACHED_STREAM_QUEUE_SIZE) reader_started = asyncio.Event() reader_abandoned = asyncio.Event() diff --git a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py index ef724bc929c..0597b1ffe62 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py +++ b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py @@ -1824,6 +1824,82 @@ async def stream_fn( assert accepted_status == 200 +async def test_endpoint_snapshot_hydration_bypasses_detached_run_capacity( + streaming_chat_client_stub: Any, +) -> None: + """Hydration remains available while a disconnected run occupies the only detached slot.""" + release = asyncio.Event() + completed = asyncio.Event() + + async def stream_fn( + messages: list[Message], + options: dict[str, Any], + **kwargs: Any, + ) -> AsyncIterator[ChatResponseUpdate]: + del messages, options, kwargs + yield ChatResponseUpdate(contents=[Content.from_text(text="started")], role="assistant") + await release.wait() + yield ChatResponseUpdate( + contents=[Content.from_text(text="-finished")], + role="assistant", + finish_reason="stop", + ) + completed.set() + + store = InMemoryAGUIThreadSnapshotStore() + await store.save( + scope="tenant-a", + thread_id="capacity-thread", + snapshot=AGUIThreadSnapshot( + messages=[{"id": "stored-user", "role": "user", "content": "Stored safe point"}], + ), + ) + agent = Agent( + name="detached-hydration-capacity", + instructions="Test agent", + client=streaming_chat_client_stub(stream_fn), + ) + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + agent, + path="/detached-hydration-capacity", + snapshot_store=store, + snapshot_scope_resolver=lambda _request: "tenant-a", + keepalive_seconds=None, + detached_runs=True, + max_detached_runs=1, + ) + + await _post_until_sse_event_then_disconnect( + app, + "/detached-hydration-capacity", + { + "runId": "active-run", + "threadId": "capacity-thread", + "messages": [{"role": "user", "content": "Start"}], + }, + event_type="TEXT_MESSAGE_CONTENT", + ) + + hydration_status, hydration_body = await _post_asgi_request( + app, + "/detached-hydration-capacity", + { + "runId": "hydrate-run", + "threadId": "capacity-thread", + "messages": [], + }, + ) + + assert hydration_status == 200 + assert b'"type":"MESSAGES_SNAPSHOT"' in hydration_body + assert b"Stored safe point" in hydration_body + + release.set() + await asyncio.wait_for(completed.wait(), timeout=5) + + async def test_endpoint_detached_run_expires_after_reader_disconnect( streaming_chat_client_stub: Any, ) -> None: From 901d4a352caac3c20baa33f6e5ef25872b0e7272 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Fri, 25 Sep 2026 12:13:07 +0200 Subject: [PATCH 4/5] Python: finalize AG-UI durability boundaries --- python/packages/ag-ui/AGENTS.md | 14 +-- python/packages/ag-ui/README.md | 12 +- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 11 -- .../ag-ui/agent_framework_ag_ui/_endpoint.py | 28 ++--- .../ag-ui/tests/ag_ui/test_endpoint.py | 117 ++++++++++++++++++ python/packages/ag-ui/tests/ag_ui/test_run.py | 110 ++++++++-------- 6 files changed, 194 insertions(+), 98 deletions(-) diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index d7fac45e1dd..9c12210e41d 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -112,13 +112,13 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. `state` entry, so conversation continuation is unreachable from client input by construction; keep it that way. - `confirm_changes` snapshot cleanup resolves the synthetic confirmation back to its original `function_call_id`; it must never concatenate unrelated tool results or record accepted changes without a matching real result. -- Configured stateless agent Thread Snapshot stores persist after finalized model-roundtrip, function/MCP tool-result, - and approval safe points, then persist the terminal state. A `finish_reason` is only a pending boundary: pull the - stream again so the inner finalizer and provider/context side effects complete before saving the preceding turn. Save - only after the complete update and approval lifecycle side effects are applied; do not write every text delta or - schedule unordered background snapshot writes. Service-session and workflow Thread Snapshots retain terminal-save - cadence so replayable messages cannot advance without matching provider continuation state; workflow checkpoints own - incremental workflow runtime state. +- Configured stateless agent Thread Snapshot stores persist after function/MCP tool-result and approval safe points, + then persist the terminal state. Never treat a yielded `finish_reason` or a following metadata update as successful + model-turn finalization: inner finalizers and result hooks may still reject that output. Save only after the complete + result/approval update and lifecycle side effects are applied; do not write every text delta or schedule unordered + background snapshot writes. Service-session and workflow Thread Snapshots retain terminal-save cadence so replayable + messages cannot advance without matching provider continuation state; workflow checkpoints own incremental workflow + runtime state. - Disconnect-safe execution is endpoint-owned and opt-in through `add_agent_framework_fastapi_endpoint(detached_runs=True)`. Keep its producer queue bounded, retain and observe producer/drainer tasks, and let the producer own final snapshot/checkpoint/approval persistence after the SSE reader diff --git a/python/packages/ag-ui/README.md b/python/packages/ag-ui/README.md index 32ba741b43e..9e40a114dd6 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -389,12 +389,12 @@ add_agent_framework_fastapi_endpoint( ) ``` -Configured stateless agent snapshot stores are updated after finalized model roundtrips, function/MCP tool-result -batches, and approval safe points, then written once more with the terminal run state. This limits progress loss during -long agent runs without persisting every streaming text delta. Service-session snapshots retain terminal-save cadence -so replayable messages cannot advance without their matching provider continuation state. Workflow Thread Snapshots -also keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow runtime -state. +Configured stateless agent snapshot stores are updated after function/MCP tool-result batches and approval safe points, +then written once more with the terminal run state. These boundaries capture completed model/tool rounds without +persisting model output whose stream finalizer may still reject it. Service-session snapshots retain terminal-save +cadence so replayable messages cannot advance without their matching provider continuation state. Workflow Thread +Snapshots also keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow +runtime state. A frontend can then hydrate the latest stored snapshot for the scoped thread: diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 1b046fb3b97..0fcc41713e7 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -3375,7 +3375,6 @@ async def save_flow_snapshot() -> None: telemetry_context = partial(_use_telemetry_conversation_id, telemetry_conversation_id) stream_completed = False native_approval_flow_result_ids: set[int] = set() - pending_model_safe_point = False try: with telemetry_context(): for queued_executions in forwarded_executions.values(): @@ -3387,13 +3386,6 @@ async def save_flow_snapshot() -> None: stream = await _normalize_response_stream(response_stream) async for update in _iterate_with_context(stream, telemetry_context): - if snapshot_session.enabled and pending_model_safe_point and not config.use_service_session: - # Pulling the next update finalizes the preceding model turn and - # completes provider/context side effects before this snapshot. - await save_flow_snapshot() - pending_model_safe_point = False - - model_safe_point = update.finish_reason is not None result_safe_point = False # Collect updates for structured output processing @@ -3561,13 +3553,10 @@ async def save_flow_snapshot() -> None: if ( snapshot_session.enabled and result_safe_point - and not model_safe_point and not flow.waiting_for_approval and not config.use_service_session ): await save_flow_snapshot() - if model_safe_point and not flow.waiting_for_approval and not config.use_service_session: - pending_model_safe_point = True # Stop if waiting for approval if flow.waiting_for_approval: diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py index 74e1b08ca94..b0e8348da79 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -354,20 +354,16 @@ async def drain_detached_stream(queue: asyncio.Queue[str | bytes | None]) -> Non async def produce_events() -> None: try: + try: + await asyncio.wait_for( + reader_started.wait(), + timeout=_DETACHED_READER_START_TIMEOUT_SECONDS, + ) + except asyncio.TimeoutError: + reader_abandoned.set() async for frame in event_generator(): if reader_abandoned.is_set(): continue - if not reader_started.is_set(): - try: - queue.put_nowait(frame) - except asyncio.QueueFull: - reader_abandoned.set() - while True: - try: - queue.get_nowait() - except asyncio.QueueEmpty: - break - continue await queue.put(frame) except asyncio.CancelledError: reader_abandoned.set() @@ -376,7 +372,7 @@ async def produce_events() -> None: current_task = asyncio.current_task() if active_run_key is not None and active_runs.get(active_run_key) is current_task: active_runs.pop(active_run_key, None) - if reader_started.is_set() or not reader_abandoned.is_set(): + if reader_started.is_set(): await queue.put(None) producer_task = asyncio.create_task( @@ -390,14 +386,6 @@ async def produce_events() -> None: retain_background_task(producer_task) async def expire_abandoned_run() -> None: - if not reader_started.is_set(): - try: - await asyncio.wait_for( - reader_started.wait(), - timeout=_DETACHED_READER_START_TIMEOUT_SECONDS, - ) - except asyncio.TimeoutError: - reader_abandoned.set() if producer_task.done(): return if not reader_abandoned.is_set(): diff --git a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py index 0597b1ffe62..7e13fbc7fc6 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py +++ b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py @@ -1453,6 +1453,52 @@ async def stream_fn(messages: Any, options: Any, **kwargs: Any): assert "RUN_FINISHED" in event_types +@pytest.mark.parametrize("keepalive_seconds", [None, 15]) +async def test_endpoint_detached_fast_connected_producer_preserves_all_events( + keepalive_seconds: float | None, +) -> None: + """A connected reader applies queue backpressure instead of being mistaken for a disconnect.""" + + class FastWorkflowRunner(AgentFrameworkWorkflow): + def __init__(self) -> None: + self.snapshot_store = None + self.checkpoint_storage = None + + async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: + run_id = str(input_data["run_id"]) + thread_id = str(input_data["thread_id"]) + yield RunStartedEvent(run_id=run_id, thread_id=thread_id) + for index in range(32): + yield StateSnapshotEvent(snapshot={"index": index}) + yield RunFinishedEvent(run_id=run_id, thread_id=thread_id) + + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + FastWorkflowRunner(), + path="/detached-fast-connected", + keepalive_seconds=keepalive_seconds, + detached_runs=True, + ) + + status_code, body = await _post_asgi_request( + app, + "/detached-fast-connected", + { + "runId": "fast-run", + "threadId": "fast-thread", + "messages": [{"role": "user", "content": "Start"}], + }, + ) + + assert status_code == 200 + events = [json.loads(line[6:]) for line in body.decode().splitlines() if line.startswith("data: ")] + assert len(events) == 34 + assert events[0]["type"] == "RUN_STARTED" + assert [event["snapshot"]["index"] for event in events[1:-1]] == list(range(32)) + assert events[-1]["type"] == "RUN_FINISHED" + + @pytest.mark.parametrize("keepalive_seconds", [None, 0.01]) async def test_endpoint_detached_run_completes_and_saves_after_client_disconnect( streaming_chat_client_stub: Any, @@ -1593,6 +1639,77 @@ async def stream_fn( assert "chunk-31" in json.dumps(snapshot.messages) +async def test_endpoint_exactly_full_unstarted_stream_releases_detached_capacity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Completion cannot block on a full queue when the response body never starts.""" + import agent_framework_ag_ui._endpoint as endpoint_module + + monkeypatch.setattr(endpoint_module, "_DETACHED_READER_START_TIMEOUT_SECONDS", 0.01) + first_completed = asyncio.Event() + + class ExactCapacityWorkflowRunner(AgentFrameworkWorkflow): + def __init__(self) -> None: + self.snapshot_store = None + self.checkpoint_storage = None + self.run_count = 0 + + async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: + self.run_count += 1 + if self.run_count == 1: + for index in range(16): + yield StateSnapshotEvent(snapshot={"index": index}) + first_completed.set() + return + run_id = str(input_data["run_id"]) + thread_id = str(input_data["thread_id"]) + yield RunStartedEvent(run_id=run_id, thread_id=thread_id) + yield RunFinishedEvent(run_id=run_id, thread_id=thread_id) + + app = FastAPI() + add_agent_framework_fastapi_endpoint( + app, + ExactCapacityWorkflowRunner(), + path="/exact-capacity", + keepalive_seconds=None, + detached_runs=True, + max_detached_runs=1, + ) + route = next(route for route in app.routes if getattr(route, "path", None) == "/exact-capacity") + assert isinstance(route, APIRoute) + + response = await route.endpoint( + AGUIRequest.model_validate( + { + "runId": "unstarted-run", + "threadId": "unstarted-thread", + "messages": [{"role": "user", "content": "Start"}], + } + ) + ) + assert isinstance(response, StreamingResponse) + await asyncio.wait_for(first_completed.wait(), timeout=5) + + retry_status = 503 + retry_body = b"" + for _ in range(100): + retry_status, retry_body = await _post_asgi_request( + app, + "/exact-capacity", + { + "runId": "retry-run", + "threadId": "retry-thread", + "messages": [{"role": "user", "content": "Retry"}], + }, + ) + if retry_status != 503: + break + await asyncio.sleep(0.01) + + assert retry_status == 200 + assert b'"type":"RUN_FINISHED"' in retry_body + + async def test_endpoint_disconnect_still_cancels_run_by_default( streaming_chat_client_stub: Any, ) -> None: diff --git a/python/packages/ag-ui/tests/ag_ui/test_run.py b/python/packages/ag-ui/tests/ag_ui/test_run.py index f875eba5f46..8060abc778d 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_run.py +++ b/python/packages/ag-ui/tests/ag_ui/test_run.py @@ -22,7 +22,7 @@ ToolCallStartEvent, ) from agent_framework import AgentResponse, AgentResponseUpdate, Content, Message, ResponseStream -from agent_framework.exceptions import AgentInvalidResponseException +from agent_framework.exceptions import AgentInvalidResponseException, ResponseInvalidatedException from conftest import StubAgent # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports] from agent_framework_ag_ui._agent import AgentConfig @@ -2900,79 +2900,81 @@ async def test_service_session_rejects_disabled_provider_storage(): ] -async def test_snapshot_is_saved_at_model_roundtrip_safe_point_before_run_completion(): - """Model output is saved only after the stream advances through turn finalization.""" - from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore - - emit_next_update = asyncio.Event() - safe_point_saved = asyncio.Event() - finish_run = asyncio.Event() - turn_finalized = asyncio.Event() - - class RecordingStore(InMemoryAGUIThreadSnapshotStore): - async def save(self, **kwargs: Any) -> None: - await super().save(**kwargs) - snapshot = kwargs["snapshot"] - if "round-one" in str(snapshot.messages): - safe_point_saved.set() +@pytest.mark.parametrize("failure_mode", ["stream", "result_hook"]) +async def test_finish_reason_and_trailing_usage_do_not_persist_invalidated_model_output( + failure_mode: str, +) -> None: + """A yielded finish reason is not durable until the run reaches a successful safe point.""" + from agent_framework_ag_ui import AgentFrameworkAgent, AGUIThreadSnapshot, InMemoryAGUIThreadSnapshotStore + invalidated = ResponseInvalidatedException("provider invalidated partial response output") + function_call = Content.from_function_call(call_id="c1", name="lookup", arguments={}) stub = StubAgent() original_run = stub.run - def blocking_run(*args: Any, **kwargs: Any) -> Any: + def invalidating_run(*args: Any, **kwargs: Any) -> Any: if not kwargs.get("stream", False): return original_run(*args, **kwargs) - session = kwargs["session"] async def updates(): yield AgentResponseUpdate( - contents=[Content.from_text(text="round-one")], + contents=[function_call], role="assistant", - finish_reason="stop", + finish_reason="tool_calls", ) - session.state["turn_finalized"] = True - turn_finalized.set() - await emit_next_update.wait() yield AgentResponseUpdate( - contents=[Content.from_text(text="-round-two")], + contents=[ + Content.from_usage( + { + "input_token_count": 1, + "output_token_count": 1, + "total_token_count": 2, + } + ) + ], role="assistant", ) - await finish_run.wait() + if failure_mode == "stream": + raise invalidated - return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + stream = ResponseStream(updates(), finalizer=AgentResponse.from_updates) + if failure_mode == "result_hook": - stub.run = blocking_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] - store = RecordingStore() - agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) - payload = { - "thread_id": "incremental-thread", - "run_id": "incremental-run", - "__ag_ui_snapshot_scope": "tenant-a", - "messages": [{"role": "user", "content": "Start"}], - } + def invalidate_result(response: AgentResponse[Any]) -> None: + del response + raise invalidated - async def collect_events() -> list[Any]: - return [event async for event in agent.run(payload)] + stream.with_result_hook(invalidate_result) + return stream - run_task = asyncio.create_task(collect_events()) - await asyncio.wait_for(turn_finalized.wait(), timeout=5) - assert await store.get(scope="tenant-a", thread_id="incremental-thread") is None - - emit_next_update.set() - await asyncio.wait_for(safe_point_saved.wait(), timeout=5) - - safe_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") - assert safe_snapshot is not None - assert "round-one" in str(safe_snapshot.messages) - assert "round-two" not in str(safe_snapshot.messages) - assert safe_snapshot.session_state == {"turn_finalized": True} + stub.run = invalidating_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = InMemoryAGUIThreadSnapshotStore() + await store.save( + scope="tenant-a", + thread_id="invalidated-thread", + snapshot=AGUIThreadSnapshot( + messages=[{"id": "baseline", "role": "user", "content": "Persisted baseline"}], + ), + ) + agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) - finish_run.set() - await asyncio.wait_for(run_task, timeout=5) + with pytest.raises(ResponseInvalidatedException, match="provider invalidated"): + _ = [ + event + async for event in agent.run( + { + "thread_id": "invalidated-thread", + "run_id": f"invalidated-{failure_mode}", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + ) + ] - final_snapshot = await store.get(scope="tenant-a", thread_id="incremental-thread") - assert final_snapshot is not None - assert "round-one-round-two" in str(final_snapshot.messages) + snapshot = await store.get(scope="tenant-a", thread_id="invalidated-thread") + assert snapshot is not None + assert snapshot.messages == [{"id": "baseline", "role": "user", "content": "Persisted baseline"}] + assert "c1" not in str(snapshot.messages) async def test_snapshot_is_saved_after_tool_result_without_finish_reason(): From 0d82e97f4ae585a5eb46a050a520ab6b535f3790 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Mon, 28 Sep 2026 14:59:20 +0200 Subject: [PATCH 5/5] Python: align AG-UI hydration and snapshot safety --- python/packages/ag-ui/AGENTS.md | 15 ++- python/packages/ag-ui/README.md | 9 +- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 106 +++++++++++++++++- .../ag-ui/agent_framework_ag_ui/_endpoint.py | 19 +--- .../agent_framework_ag_ui/_run_common.py | 17 +++ .../ag-ui/agent_framework_ag_ui/_workflow.py | 7 +- .../ag-ui/tests/ag_ui/test_endpoint.py | 5 +- python/packages/ag-ui/tests/ag_ui/test_run.py | 76 ++++++++++++- .../ag-ui/tests/ag_ui/test_run_common.py | 37 ++++++ 9 files changed, 256 insertions(+), 35 deletions(-) diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index 9c12210e41d..8de91344a1a 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -115,10 +115,11 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. - Configured stateless agent Thread Snapshot stores persist after function/MCP tool-result and approval safe points, then persist the terminal state. Never treat a yielded `finish_reason` or a following metadata update as successful model-turn finalization: inner finalizers and result hooks may still reject that output. Save only after the complete - result/approval update and lifecycle side effects are applied; do not write every text delta or schedule unordered - background snapshot writes. Service-session and workflow Thread Snapshots retain terminal-save cadence so replayable - messages cannot advance without matching provider continuation state; workflow checkpoints own incremental workflow - runtime state. + result/approval update and lifecycle side effects are applied, and project only completed call/result groups and + current approval controls into the intermediate snapshot. Sibling text, reasoning, and unrelated pending calls remain + terminal-only. Do not write every text delta or schedule unordered background snapshot writes. Service-session and + workflow Thread Snapshots retain terminal-save cadence so replayable messages cannot advance without matching + provider continuation state; workflow checkpoints own incremental workflow runtime state. - Disconnect-safe execution is endpoint-owned and opt-in through `add_agent_framework_fastapi_endpoint(detached_runs=True)`. Keep its producer queue bounded, retain and observe producer/drainer tasks, and let the producer own final snapshot/checkpoint/approval persistence after the SSE reader @@ -126,8 +127,10 @@ AG-UI protocol integration for building agent UIs with the AG-UI standard. This mode discards unread events; it is not a resumable event log. - While a detached mutation is active, the endpoint may serve an empty Snapshot Hydrate Request but must reject another mutation for the same `(Snapshot Scope, threadId)` with HTTP 409. Classify hydration once and keep it on the direct - response path so it bypasses detached admission and producer wrapping. The guard is process-local and does not - replace cross-replica coordination. + response path so it bypasses detached admission and producer wrapping. Endpoint admission, agent hydration, and + workflow hydration must all use `_run_common._is_snapshot_hydration_request`; pass workflow checkpoint capability + explicitly so checkpoint resumes never become hydration. The guard is process-local and does not replace cross-replica + coordination. - Detached runs may outlive FastAPI request-scoped disposable resources. Resolve authorization and Snapshot Scope before spawning the producer, and do not rely on request-owned clients or sessions remaining open after disconnect. - SSE keepalive is endpoint-owned transport behavior configured through diff --git a/python/packages/ag-ui/README.md b/python/packages/ag-ui/README.md index 9e40a114dd6..2a53d7f3ac2 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -391,10 +391,11 @@ add_agent_framework_fastapi_endpoint( Configured stateless agent snapshot stores are updated after function/MCP tool-result batches and approval safe points, then written once more with the terminal run state. These boundaries capture completed model/tool rounds without -persisting model output whose stream finalizer may still reject it. Service-session snapshots retain terminal-save -cadence so replayable messages cannot advance without their matching provider continuation state. Workflow Thread -Snapshots also keep their terminal-save cadence; workflow checkpointing remains the mechanism for incremental workflow -runtime state. +persisting sibling text or reasoning whose stream finalizer may still reject it. Intermediate snapshots project only +completed call/result groups and current approval controls; the terminal snapshot retains the complete finalized +output. Service-session snapshots retain terminal-save cadence so replayable messages cannot advance without their +matching provider continuation state. Workflow Thread Snapshots also keep their terminal-save cadence; workflow +checkpointing remains the mechanism for incremental workflow runtime state. A frontend can then hydrate the latest stored snapshot for the scoped thread: diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 0fcc41713e7..3130c49fb89 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -102,6 +102,7 @@ _extract_resume_payload, # type: ignore _extract_tool_result_display, # type: ignore _has_only_tool_calls, # type: ignore + _is_snapshot_hydration_request, # type: ignore _iterate_with_context, # type: ignore _normalize_resume_interrupts, # type: ignore _new_tool_call_segment_id, # type: ignore @@ -2363,6 +2364,97 @@ def _build_messages_snapshot( return MessagesSnapshotEvent(messages=_project_host_payload_history(bounded_messages)) # type: ignore[arg-type] +def _safe_point_tool_call_ids(flow: FlowState) -> set[str]: + """Return tool calls whose snapshot content is safe before stream finalization.""" + safe_ids = {str(result["toolCallId"]) for result in flow.tool_results if result.get("toolCallId") is not None} + interrupt_call_ids = { + str(tool_call_id) + for interrupt in flow.interrupts + if (tool_call_id := interrupt.get("toolCallId") or interrupt.get("id")) is not None + } + safe_ids.update(interrupt_call_ids) + for call in flow.pending_tool_calls: + call_id = call.get("id") + function = call.get("function") + if call_id is None or not isinstance(function, Mapping) or function.get("name") != "confirm_changes": + continue + arguments = function.get("arguments") + try: + parsed_arguments = json.loads(arguments) if isinstance(arguments, str) else arguments + except json.JSONDecodeError: + parsed_arguments = None + if not isinstance(parsed_arguments, Mapping): + continue + original_call_id = parsed_arguments.get("function_call_id") + if original_call_id is None or str(original_call_id) in interrupt_call_ids: + safe_ids.add(str(call_id)) + return safe_ids + + +def _build_safe_point_messages_snapshot( + flow: FlowState, + snapshot_messages: list[dict[str, Any]], +) -> MessagesSnapshotEvent: + """Build an intermediate snapshot without unfinalized text or reasoning.""" + all_messages = list(snapshot_messages) + safe_call_ids = _safe_point_tool_call_ids(flow) + results_by_call_id: dict[str, list[dict[str, Any]]] = {} + for result in flow.tool_results: + call_id = result.get("toolCallId") + if call_id is not None: + results_by_call_id.setdefault(str(call_id), []).append(result) + + emitted_call_ids: set[str] = set() + for segment in flow.snapshot_segments: + if segment.get("kind") != "tool_calls": + continue + calls = [ + flow.tool_calls_by_id[call_id] + for call_id in segment.get("call_ids", []) + if call_id in safe_call_ids and call_id in flow.tool_calls_by_id + ] + if not calls: + continue + all_messages.append( + { + "id": str(segment.get("id") or generate_event_id()), + "role": "assistant", + "tool_calls": [call.copy() for call in calls], + } + ) + for call in calls: + call_id = str(call["id"]) + emitted_call_ids.add(call_id) + all_messages.extend(results_by_call_id.get(call_id, [])) + + leftover_calls = [ + call + for call in flow.pending_tool_calls + if (call_id := call.get("id")) is not None + and str(call_id) in safe_call_ids + and str(call_id) not in emitted_call_ids + ] + if leftover_calls: + all_messages.append( + { + "id": generate_event_id(), + "role": "assistant", + "tool_calls": [call.copy() for call in leftover_calls], + } + ) + for call in leftover_calls: + call_id = str(call["id"]) + emitted_call_ids.add(call_id) + all_messages.extend(results_by_call_id.get(call_id, [])) + + for call_id, results in results_by_call_id.items(): + if call_id not in emitted_call_ids: + all_messages.extend(results) + + bounded_messages = _bound_host_payload_history(_persistable_host_payload_history(all_messages)) + return MessagesSnapshotEvent(messages=_project_host_payload_history(bounded_messages)) # type: ignore[arg-type] + + def _text_events_to_snapshot_messages(events: list[BaseEvent]) -> list[dict[str, Any]]: """Convert streamed text-message events into snapshot message dictionaries.""" messages: list[dict[str, Any]] = [] @@ -2773,7 +2865,11 @@ async def _run_agent_stream( await snapshot_session.clear_interrupts(interrupt_ids=retired_interrupt_ids) stored_snapshot = snapshot_session.stored stored_pending_approval_interrupt_ids.difference_update(retired_interrupt_ids) - if snapshot_session.enabled and not raw_messages and resume_payload is None: + if _is_snapshot_hydration_request( + input_data, + snapshot_enabled=snapshot_session.enabled, + supports_checkpoint_resume=False, + ): async for event in snapshot_session.hydrate_events(run_id=run_id): yield event return @@ -3356,8 +3452,8 @@ async def save_thread_snapshot( cast(dict[str, Any], make_json_safe(flow.current_state)) if flow.current_state else None ) - async def save_flow_snapshot() -> None: - safe_point_event = _build_messages_snapshot(flow, snapshot_messages) + async def save_safe_point_snapshot() -> None: + safe_point_event = _build_safe_point_messages_snapshot(flow, snapshot_messages) safe_point_messages = _event_messages_to_snapshot_dicts(list(safe_point_event.messages)) if resume_payload is not None and not seeded_resume_from_snapshot and snapshot_seed_messages is None: safe_point_messages = snapshot_session.resume_seeded_messages(safe_point_messages) @@ -3556,7 +3652,7 @@ async def save_flow_snapshot() -> None: and not flow.waiting_for_approval and not config.use_service_session ): - await save_flow_snapshot() + await save_safe_point_snapshot() # Stop if waiting for approval if flow.waiting_for_approval: @@ -3592,7 +3688,7 @@ async def save_flow_snapshot() -> None: if isinstance(stream, ResponseStream): await stream.get_final_response() if snapshot_session.enabled and not config.use_service_session: - await save_flow_snapshot() + await save_safe_point_snapshot() # If no updates at all, still emit RunStarted if not run_started_emitted: diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py index b0e8348da79..2a4080cca73 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -21,7 +21,7 @@ from ._agent import AgentFrameworkAgent from ._approval_state import _APPROVAL_SCOPE_INPUT_KEY -from ._run_common import _extract_resume_payload +from ._run_common import _is_snapshot_hydration_request from ._snapshots import ( _DEFAULT_STATE_INPUT_KEY, _SNAPSHOT_SCOPE_INPUT_KEY, @@ -90,19 +90,6 @@ def _validate_detached_run_options(max_detached_runs: int, detached_run_timeout_ raise ValueError("detached_run_timeout_seconds must be positive.") -def _is_snapshot_hydration_request( - request: AGUIRequest, - input_data: dict[str, Any], - *, - snapshot_persistence_active: bool, -) -> bool: - """Return whether a request only replays the latest stored snapshot.""" - if not snapshot_persistence_active or request.messages or _extract_resume_payload(input_data) is not None: - return False - forwarded_props = request.forwarded_props or {} - return not (forwarded_props.get("checkpoint_id") or forwarded_props.get("checkpointId")) - - def add_agent_framework_fastapi_endpoint( app: FastAPI, agent: SupportsAgentRun | AgentFrameworkAgent | Workflow | AgentFrameworkWorkflow, @@ -265,9 +252,9 @@ async def agent_endpoint(request_body: AGUIRequest) -> Response: keepalive_enabled = keepalive_seconds is not None snapshot_hydration_request = _is_snapshot_hydration_request( - request_body, input_data, - snapshot_persistence_active=snapshot_persistence_active, + snapshot_enabled=snapshot_persistence_active, + supports_checkpoint_resume=isinstance(protocol_runner, AgentFrameworkWorkflow), ) active_run_key: tuple[str | None, str] | None = None if detached_runs and request_body.thread_id is not None and not snapshot_hydration_request: diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_run_common.py b/python/packages/ag-ui/agent_framework_ag_ui/_run_common.py index cdbf1d71c47..5c47a105b8e 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_run_common.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_run_common.py @@ -174,6 +174,23 @@ def _extract_resume_payload(input_data: dict[str, Any]) -> Any: return forwarded_props_dict.get("resume") +def _is_snapshot_hydration_request( + input_data: dict[str, Any], + *, + snapshot_enabled: bool, + supports_checkpoint_resume: bool, +) -> bool: + """Return whether a request only replays the latest stored snapshot.""" + if not snapshot_enabled or input_data.get("messages") or _extract_resume_payload(input_data) is not None: + return False + if not supports_checkpoint_resume: + return True + forwarded_props = input_data.get("forwarded_props") or input_data.get("forwardedProps") + if not isinstance(forwarded_props, Mapping): + return True + return not (forwarded_props.get("checkpoint_id") or forwarded_props.get("checkpointId")) + + def _strict_resume_entries(resume_payload: Any) -> tuple[list[dict[str, Any]], str | None]: """Parse resume entries for pending interrupt contract validation.""" if isinstance(resume_payload, list): diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py index 33d3bede30e..93c703e484a 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py @@ -37,6 +37,7 @@ from ._run_common import ( _cancelled_resume_interrupt_ids, _extract_resume_payload, + _is_snapshot_hydration_request, _normalize_resume_interrupts, _reconstruct_messages_from_thread_snapshot, ) @@ -619,7 +620,11 @@ async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: # A checkpoint resume legitimately carries no new messages; it must reach the # core workflow's restore path rather than replaying a stored thread snapshot. - if checkpoint_id is None and snapshot_session.enabled and not raw_messages and resume_payload is None: + if _is_snapshot_hydration_request( + input_data, + snapshot_enabled=snapshot_session.enabled, + supports_checkpoint_resume=True, + ): async for event in snapshot_session.hydrate_events(run_id=run_id): yield event return diff --git a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py index 7e13fbc7fc6..f84ed233c51 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_endpoint.py +++ b/python/packages/ag-ui/tests/ag_ui/test_endpoint.py @@ -2177,7 +2177,7 @@ def resolve_scope(request: AGUIRequest) -> str: assert mutation_status == 409 assert b"already active" in mutation_body - checkpoint_status, _ = await _post_asgi_request( + checkpoint_status, checkpoint_body = await _post_asgi_request( app, "/detached-guard", { @@ -2188,7 +2188,8 @@ def resolve_scope(request: AGUIRequest) -> str: "forwardedProps": {"checkpoint_id": "checkpoint-1"}, }, ) - assert checkpoint_status == 409 + assert checkpoint_status == 200 + assert b'"type":"MESSAGES_SNAPSHOT"' in checkpoint_body for forwarded_props in ( {"resume": [{"interruptId": "approval-1", "status": "resolved", "payload": {"approved": True}}]}, diff --git a/python/packages/ag-ui/tests/ag_ui/test_run.py b/python/packages/ag-ui/tests/ag_ui/test_run.py index 8060abc778d..f4d7b02fdea 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_run.py +++ b/python/packages/ag-ui/tests/ag_ui/test_run.py @@ -2977,6 +2977,68 @@ def invalidate_result(response: AgentResponse[Any]) -> None: assert "c1" not in str(snapshot.messages) +async def test_tool_result_safe_point_excludes_unfinalized_sibling_output() -> None: + """Intermediate result snapshots retain completed groups but omit sibling model output.""" + from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore + + invalidated = ResponseInvalidatedException("provider invalidated result update") + function_call = Content.from_function_call(call_id="safe-call", name="lookup", arguments={}) + stub = StubAgent() + original_run = stub.run + + def invalidating_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates(): + yield AgentResponseUpdate( + contents=[function_call], + role="assistant", + finish_reason="tool_calls", + ) + yield AgentResponseUpdate( + contents=[ + Content.from_text(text="unfinalized sibling text"), + Content.from_text_reasoning(id="unsafe-reasoning", text="unfinalized reasoning"), + Content.from_function_result(call_id="safe-call", result="completed tool output"), + ], + role="tool", + ) + + stream = ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + def invalidate_result(response: AgentResponse[Any]) -> None: + del response + raise invalidated + + return stream.with_result_hook(invalidate_result) + + stub.run = invalidating_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkAgent(agent=stub, snapshot_store=store) + + with pytest.raises(ResponseInvalidatedException, match="provider invalidated"): + _ = [ + event + async for event in agent.run( + { + "thread_id": "safe-projection-thread", + "run_id": "safe-projection-run", + "__ag_ui_snapshot_scope": "tenant-a", + "messages": [{"role": "user", "content": "Start"}], + } + ) + ] + + snapshot = await store.get(scope="tenant-a", thread_id="safe-projection-thread") + assert snapshot is not None + serialized = str(snapshot.messages) + assert "safe-call" in serialized + assert "completed tool output" in serialized + assert "unfinalized sibling text" not in serialized + assert "unfinalized reasoning" not in serialized + + async def test_snapshot_is_saved_after_tool_result_without_finish_reason(): """A completed tool-result batch is durable even when its update has no finish reason.""" from agent_framework_ag_ui import AgentFrameworkAgent, InMemoryAGUIThreadSnapshotStore @@ -3198,12 +3260,14 @@ async def test_interrupt_snapshot_is_saved_after_approval_lifecycle_registration from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore approval_saved = asyncio.Event() + saved_snapshots: list[Any] = [] state_store = InMemoryAGUIApprovalStateStore() scoped_thread_id = approval_state_thread_id(scope="tenant-a", thread_id="approval-thread") class RecordingStore(InMemoryAGUIThreadSnapshotStore): async def save(self, **kwargs: Any) -> None: snapshot = kwargs["snapshot"] + saved_snapshots.append(snapshot) for interrupt in snapshot.interrupt or []: occurrence = state_store.lifecycle.occurrence_for_alias( thread_id=scoped_thread_id, @@ -3226,7 +3290,11 @@ async def save(self, **kwargs: Any) -> None: stub = StubAgent( updates=[ AgentResponseUpdate( - contents=[approval_request], + contents=[ + Content.from_text(text="unfinalized approval sibling"), + Content.from_text_reasoning(id="approval-reasoning", text="unfinalized approval reasoning"), + approval_request, + ], role="assistant", finish_reason="tool_calls", ) @@ -3257,6 +3325,12 @@ async def save(self, **kwargs: Any) -> None: assert snapshot is not None assert snapshot.interrupt is not None assert snapshot.interrupt[0]["id"] == "approval-occurrence" + assert len(saved_snapshots) >= 2 + intermediate_snapshot = saved_snapshots[0] + assert "unfinalized approval sibling" not in str(intermediate_snapshot.messages) + assert "unfinalized approval reasoning" not in str(intermediate_snapshot.messages) + assert "unfinalized approval sibling" in str(snapshot.messages) + assert "unfinalized approval reasoning" in str(snapshot.messages) async def test_plain_async_iterable_persists_waiting_approval_snapshot(): diff --git a/python/packages/ag-ui/tests/ag_ui/test_run_common.py b/python/packages/ag-ui/tests/ag_ui/test_run_common.py index 1841a0d37cc..2e6b8ed56ce 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_run_common.py +++ b/python/packages/ag-ui/tests/ag_ui/test_run_common.py @@ -33,6 +33,7 @@ _emit_tool_result, _extract_resume_payload, _extract_tool_result_state, + _is_snapshot_hydration_request, _normalize_resume_interrupts, _reconstruct_messages_from_thread_snapshot, _strict_resume_entries, @@ -125,6 +126,42 @@ def test_canonical_resume_entry_uses_interrupt_id_and_payload(self): assert result == [{"id": "req_1", "value": {"approved": True}, "status": "resolved"}] +@pytest.mark.parametrize( + ("input_data", "snapshot_enabled", "supports_checkpoint_resume", "expected"), + [ + ({"messages": []}, True, False, True), + ({"messages": []}, True, True, True), + ({"messages": [{"role": "user", "content": "continue"}]}, True, False, False), + ({"messages": [], "resume": [{"interruptId": "i1"}]}, True, False, False), + ({"messages": [], "forwardedProps": {"resume": [{"interruptId": "i1"}]}}, True, False, False), + ( + {"messages": [], "forwardedProps": {"command": {"resume": [{"interruptId": "i1"}]}}}, + True, + False, + False, + ), + ({"messages": [], "forwardedProps": {"checkpoint_id": "cp1"}}, True, False, True), + ({"messages": [], "forwardedProps": {"checkpoint_id": "cp1"}}, True, True, False), + ({"messages": []}, False, False, False), + ], +) +def test_snapshot_hydration_classification_is_shared( + input_data: dict[str, Any], + snapshot_enabled: bool, + supports_checkpoint_resume: bool, + expected: bool, +) -> None: + """Hydration admission matches runner resume capabilities.""" + assert ( + _is_snapshot_hydration_request( + input_data, + snapshot_enabled=snapshot_enabled, + supports_checkpoint_resume=supports_checkpoint_resume, + ) + is expected + ) + + class TestStrictResumeEntries: """Tests for strict canonical resume-entry parsing."""