diff --git a/python/packages/ag-ui/AGENTS.md b/python/packages/ag-ui/AGENTS.md index 94a11e1ca2..8de91344a1 100644 --- a/python/packages/ag-ui/AGENTS.md +++ b/python/packages/ag-ui/AGENTS.md @@ -112,6 +112,27 @@ 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 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, 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 + 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. Classify hydration once and keep it on the direct + 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 `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 ee9c0af4ed..2a53d7f3ac 100644 --- a/python/packages/ag-ui/README.md +++ b/python/packages/ag-ui/README.md @@ -389,6 +389,14 @@ 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 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: ```json @@ -398,6 +406,42 @@ 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, + 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, 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). + +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 +512,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 cece3ed33d..ab44b2249d 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 @@ -103,6 +103,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 @@ -2390,6 +2391,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]] = [] @@ -2802,7 +2894,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 @@ -3099,6 +3195,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 @@ -3297,6 +3412,46 @@ async def _run_agent_stream( if resolved_approval_results or any(message.get("function_approvals") for message in snapshot_messages): _merge_resolved_approval_results_into_snapshot(snapshot_messages, messages) + retired_approval_interrupt_ids = { + *handled_resume_ids, + *( + str(interrupt["id"]) + for interrupt in _normalize_resume_interrupts(resume_payload) + if resolved_approval_results or validated_approved_responses or approval_snapshot_reconciliations + ), + *( + reconciliation.interrupt_id + for reconciliation in approval_snapshot_reconciliations + if reconciliation.retire_interrupt + ), + } + + def remaining_stored_interrupts(retired_interrupt_ids: set[str]) -> list[dict[str, Any]] | None: + stored_interrupts = snapshot_session.stored.interrupt if snapshot_session.stored is not None else None + if not stored_interrupts: + return None + remaining = [ + interrupt + for interrupt in stored_interrupts + if str(interrupt.get("id") or interrupt.get("interruptId")) not in retired_interrupt_ids + ] + return remaining or None + + def merge_snapshot_interrupts( + stored_interrupts: list[dict[str, Any]] | None, + current_interrupts: list[dict[str, Any]], + ) -> list[dict[str, Any]] | None: + combined_by_id: dict[str, dict[str, Any]] = {} + anonymous: list[dict[str, Any]] = [] + for interrupt in [*(stored_interrupts or []), *current_interrupts]: + interrupt_id = interrupt.get("id") or interrupt.get("interruptId") + if interrupt_id is None: + anonymous.append(interrupt) + else: + combined_by_id[str(interrupt_id)] = interrupt + combined = [*combined_by_id.values(), *anonymous] + return combined or None + if replacement_approval_requests: yield RunStartedEvent(run_id=run_id, thread_id=thread_id) for request in replacement_approval_requests: @@ -3311,18 +3466,14 @@ 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, + interrupt=merge_snapshot_interrupts( + remaining_stored_interrupts(retired_approval_interrupt_ids), + flow.interrupts, ), ) - _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 @@ -3336,8 +3487,31 @@ async def _run_agent_stream( flow.current_state.update(approved_state_updates) approved_state_snapshot_emitted = True + is_confirm_changes_response = _is_confirm_changes_response(messages) + preserved_interrupts: list[dict[str, Any]] | None = None + if ( + (resolved_approval_results or retired_approval_interrupt_ids) + and snapshot_session.enabled + and not config.use_service_session + and not is_confirm_changes_response + ): + 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) + preserved_interrupts = remaining_stored_interrupts(retired_approval_interrupt_ids) + 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=preserved_interrupts, + ) + # Handle confirm_changes response (state confirmation flow - emit confirmation and stop) - if _is_confirm_changes_response(messages): + if is_confirm_changes_response: + confirm_additional_properties = cast(dict[str, Any], messages[-1].additional_properties or {}) + confirm_interrupt_id = confirm_additional_properties.get("tool_call_id") + confirm_remaining_interrupts = remaining_stored_interrupts( + {str(confirm_interrupt_id)} if confirm_interrupt_id else set() + ) yield RunStartedEvent(run_id=run_id, thread_id=thread_id) # Emit approved state snapshot before confirmation message if approved_state_snapshot_emitted: @@ -3352,18 +3526,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, - ), + interrupt=confirm_remaining_interrupts, ) - _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 @@ -3381,6 +3548,21 @@ 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 ) + + def snapshot_interrupts() -> list[dict[str, Any]] | None: + return merge_snapshot_interrupts(preserved_interrupts, flow.interrupts) + + 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) + await save_thread_snapshot( + persisted_messages=safe_point_messages, + state=latest_state_snapshot, + interrupt=snapshot_interrupts(), + ) + initial_state_snapshot = flow.current_state if state_schema and flow.current_state else None # With both IDs supplied there is nothing to wait for, so start the run before @@ -3417,6 +3599,8 @@ async def _run_agent_stream( stream = await _normalize_response_stream(response_stream) async for update in _iterate_with_context(stream, telemetry_context): + result_safe_point = False + # Collect updates for structured output processing if response_format is not None: all_updates.append(update) @@ -3457,6 +3641,12 @@ 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", + "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 @@ -3566,6 +3756,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 result_safe_point + and not flow.waiting_for_approval + and not config.use_service_session + ): + await save_safe_point_snapshot() + # Stop if waiting for approval if flow.waiting_for_approval: break @@ -3596,8 +3794,11 @@ async def _run_agent_stream( 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 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_safe_point_snapshot() # If no updates at all, still emit RunStarted if not run_started_emitted: @@ -3791,16 +3992,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, - ), + interrupt=snapshot_interrupts(), ) - _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 3b2241b48d..2a4080cca7 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_endpoint.py @@ -4,9 +4,11 @@ from __future__ import annotations +import asyncio import copy import logging from collections.abc import AsyncGenerator, Sequence +from contextlib import suppress from inspect import isawaitable from typing import Any, cast @@ -15,10 +17,11 @@ 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 +from ._run_common import _is_snapshot_hydration_request from ._snapshots import ( _DEFAULT_STATE_INPUT_KEY, _SNAPSHOT_SCOPE_INPUT_KEY, @@ -30,6 +33,8 @@ logger = logging.getLogger(__name__) +_DETACHED_READER_START_TIMEOUT_SECONDS = 1.0 +_DETACHED_STREAM_QUEUE_SIZE = 16 _KEEPALIVE_COMMENT = "keepalive" @@ -78,6 +83,13 @@ def _validate_keepalive_seconds(keepalive_seconds: float | None) -> None: raise ValueError("keepalive_seconds must be positive or None.") +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 add_agent_framework_fastapi_endpoint( app: FastAPI, agent: SupportsAgentRun | AgentFrameworkAgent | Workflow | AgentFrameworkWorkflow, @@ -93,6 +105,9 @@ 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, + max_detached_runs: int = 32, + detached_run_timeout_seconds: float = 3600, ) -> None: """Add an AG-UI endpoint to a FastAPI app. @@ -128,8 +143,18 @@ 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. + 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): @@ -167,6 +192,27 @@ def add_agent_framework_fastapi_endpoint( snapshot_scope_resolver=snapshot_scope_resolver, ) + 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: + 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 +223,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 +251,30 @@ 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( + input_data, + 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: + 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) + if detached_runs and not snapshot_hydration_request: + 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: @@ -257,6 +329,112 @@ 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 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() + + 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 + 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) + if reader_started.is_set(): + 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 + 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 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: + completed = True + return + 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')}", + ) + retain_background_task(drain_task) + + stream = detached_event_generator() + else: + stream = event_generator() + headers = { "Cache-Control": "no-cache", "Connection": "keep-alive", @@ -267,14 +445,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/_run_common.py b/python/packages/ag-ui/agent_framework_ag_ui/_run_common.py index cdbf1d71c4..5c47a105b8 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/_snapshot_session.py b/python/packages/ag-ui/agent_framework_ag_ui/_snapshot_session.py index fa63dda39f..976a2cc916 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/agent_framework_ag_ui/_workflow.py b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py index 33d3bede30..93c703e484 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 c10e20fcfc..1a29960c1c 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, @@ -56,6 +56,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 @@ -130,6 +132,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 @@ -1203,8 +1256,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: @@ -1218,9 +1271,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), @@ -1243,6 +1308,28 @@ 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_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) + + 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__ @@ -1256,12 +1343,36 @@ 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 + 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: """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 + 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: """Importing endpoint helpers does not trigger sse-starlette's process-global transport hooks.""" import_check = ( @@ -1274,73 +1385,874 @@ def test_endpoint_module_import_does_not_import_sse_transport() -> None: check=False, ) - assert result.returncode == 0, result.stderr or result.stdout + assert result.returncode == 0, result.stderr or result.stdout + + +async def test_endpoint_keepalive_enabled_emits_static_comment_during_silent_gap(streaming_chat_client_stub): + """Enabled keepalive sends static SSE comments without changing AG-UI data frames.""" + app = FastAPI() + + async def stream_fn(messages: Any, options: Any, **kwargs: Any): + del messages, options, kwargs + await asyncio.sleep(0.05) + yield ChatResponseUpdate(contents=[Content.from_text(text="Done")]) + + agent = Agent(name="test", instructions="Test agent", client=streaming_chat_client_stub(stream_fn)) + + add_agent_framework_fastapi_endpoint(app, agent, path="/keepalive", keepalive_seconds=0.01) + + client = TestClient(app) + response = client.post("/keepalive", json={"messages": [{"role": "user", "content": "Hello"}]}) + + assert response.status_code == 200 + assert response.headers["content-type"] == "text/event-stream; charset=utf-8" + assert response.headers["cache-control"] == "no-cache" + assert response.headers["connection"] == "keep-alive" + assert response.headers["x-accel-buffering"] == "no" + + content = response.content.decode("utf-8") + comments = [line for line in content.splitlines() if line.startswith(":")] + assert comments + assert set(comments) == {": keepalive"} + assert "data: data:" not in content + + event_types = [event.get("type") for event in _decode_sse_events(response)] + assert "RUN_STARTED" in event_types + assert "TEXT_MESSAGE_CONTENT" in event_types + assert "RUN_FINISHED" in event_types + + +async def test_endpoint_keepalive_disabled_preserves_streaming_response_shape(streaming_chat_client_stub): + """Disabled keepalive keeps the original SSE data frames without transport comments.""" + app = FastAPI() + + async def stream_fn(messages: Any, options: Any, **kwargs: Any): + del messages, options, kwargs + await asyncio.sleep(0.05) + yield ChatResponseUpdate(contents=[Content.from_text(text="Done")]) + + agent = Agent(name="test", instructions="Test agent", client=streaming_chat_client_stub(stream_fn)) + + add_agent_framework_fastapi_endpoint(app, agent, path="/no-keepalive", keepalive_seconds=None) + + client = TestClient(app) + response = client.post("/no-keepalive", json={"messages": [{"role": "user", "content": "Hello"}]}) + + assert response.status_code == 200 + assert response.headers["content-type"] == "text/event-stream; charset=utf-8" + assert response.headers["cache-control"] == "no-cache" + assert response.headers["connection"] == "keep-alive" + assert response.headers["x-accel-buffering"] == "no" + + content = response.content.decode("utf-8") + assert ": keepalive" not in content + assert "data: data:" not in content + + event_types = [event.get("type") for event in _decode_sse_events(response)] + assert "RUN_STARTED" in event_types + assert "TEXT_MESSAGE_CONTENT" in event_types + 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, + 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_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_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: + """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_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_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_keepalive_enabled_emits_static_comment_during_silent_gap(streaming_chat_client_stub): - """Enabled keepalive sends static SSE comments without changing AG-UI data frames.""" - app = FastAPI() - async def stream_fn(messages: Any, options: Any, **kwargs: Any): +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 - await asyncio.sleep(0.05) - yield ChatResponseUpdate(contents=[Content.from_text(text="Done")]) + 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="test", instructions="Test agent", client=streaming_chat_client_stub(stream_fn)) + 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, + ) - add_agent_framework_fastapi_endpoint(app, agent, path="/keepalive", keepalive_seconds=0.01) + 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) - client = TestClient(app) - response = client.post("/keepalive", json={"messages": [{"role": "user", "content": "Hello"}]}) + retry_status, retry_body = await _post_asgi_request( + app, + "/detached-expiry", + { + "runId": "retry-run", + "threadId": "expiring-thread", + "messages": [{"role": "user", "content": "retry"}], + }, + ) - assert response.status_code == 200 - assert response.headers["content-type"] == "text/event-stream; charset=utf-8" - assert response.headers["cache-control"] == "no-cache" - assert response.headers["connection"] == "keep-alive" - assert response.headers["x-accel-buffering"] == "no" + assert retry_status == 200 + assert b"retry-complete" in retry_body - content = response.content.decode("utf-8") - comments = [line for line in content.splitlines() if line.startswith(":")] - assert comments - assert set(comments) == {": keepalive"} - assert "data: data:" not in content - event_types = [event.get("type") for event in _decode_sse_events(response)] - assert "RUN_STARTED" in event_types - assert "TEXT_MESSAGE_CONTENT" in event_types - assert "RUN_FINISHED" in event_types +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 -async def test_endpoint_keepalive_disabled_preserves_streaming_response_shape(streaming_chat_client_stub): - """Disabled keepalive keeps the original SSE data frames without transport comments.""" + 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, + ) - async def stream_fn(messages: Any, options: Any, **kwargs: Any): - del messages, options, kwargs - await asyncio.sleep(0.05) - yield ChatResponseUpdate(contents=[Content.from_text(text="Done")]) + 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", + ) - agent = Agent(name="test", instructions="Test agent", client=streaming_chat_client_stub(stream_fn)) + 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 - add_agent_framework_fastapi_endpoint(app, agent, path="/no-keepalive", keepalive_seconds=None) + 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 - client = TestClient(app) - response = client.post("/no-keepalive", json={"messages": [{"role": "user", "content": "Hello"}]}) + checkpoint_status, checkpoint_body = 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 == 200 + assert b'"type":"MESSAGES_SNAPSHOT"' in checkpoint_body - assert response.status_code == 200 - assert response.headers["content-type"] == "text/event-stream; charset=utf-8" - assert response.headers["cache-control"] == "no-cache" - assert response.headers["connection"] == "keep-alive" - assert response.headers["x-accel-buffering"] == "no" + 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 - content = response.content.decode("utf-8") - assert ": keepalive" not in content - assert "data: data:" not in content + 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) - event_types = [event.get("type") for event in _decode_sse_events(response)] - assert "RUN_STARTED" in event_types - assert "TEXT_MESSAGE_CONTENT" in event_types - assert "RUN_FINISHED" in event_types + assert final_status == 200 async def test_endpoint_keepalive_disabled_does_not_import_sse_transport(build_chat_client) -> None: @@ -1381,6 +2293,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() @@ -5645,6 +6579,14 @@ async def test_endpoint_fides_replacement_rotates_lifecycle_generation( snapshot = await snapshot_store.get(scope="tenant-a", thread_id=thread_id) assert snapshot is not None and snapshot.session_state is not None snapshot.session_state[security.source_id]["pending_policy_approvals"][original_id]["created_at"] = 0.0 + assert snapshot.interrupt is not None + snapshot.interrupt.append( + { + "id": "other-input", + "reason": "input_required", + "message": "Choose another value.", + } + ) await snapshot_store.save(scope="tenant-a", thread_id=thread_id, snapshot=snapshot) stale = client.post( @@ -5666,6 +6608,10 @@ async def test_endpoint_fides_replacement_rotates_lifecycle_generation( assert replacement_id != original_id assert replacement_interrupts[0]["toolCallId"] == "call-guarded" assert executed == [] + replaced_snapshot = await snapshot_store.get(scope="tenant-a", thread_id=thread_id) + assert replaced_snapshot is not None + assert replaced_snapshot.interrupt is not None + assert {interrupt["id"] for interrupt in replaced_snapshot.interrupt} == {replacement_id, "other-input"} approval_state = wrapped_agent._approval_state_store.get_tool_approval_state( approval_state_thread_id(scope="tenant-a", thread_id=thread_id) ) @@ -9651,8 +10597,12 @@ async def run(self, input_data: dict[str, Any]): assert response.status_code == 200 -async def test_agent_endpoint_confirm_changes_clears_persisted_interrupt(streaming_chat_client_stub): - """A confirm_changes response persists the completed turn and clears the stored interrupt.""" +@pytest.mark.parametrize("with_other_interrupt", [False, True], ids=["only-confirm", "preserve-other"]) +async def test_agent_endpoint_confirm_changes_clears_persisted_interrupt( + streaming_chat_client_stub, + with_other_interrupt: bool, +): + """A confirm response retires only its own persisted interrupt.""" app = FastAPI() call_count = 0 @@ -9698,6 +10648,17 @@ async def stream_fn(messages: Any, options: Any, **kwargs: Any): first_finished = [event for event in first_events if event.get("type") == "RUN_FINISHED"] first_interrupts = _run_finished_interrupts(first_finished[-1]) confirm_call_id = first_interrupts[0]["id"] + if with_other_interrupt: + stored_snapshot = await store.get(scope="tenant-a", thread_id="agent-thread") + assert stored_snapshot is not None and stored_snapshot.interrupt is not None + stored_snapshot.interrupt.append( + { + "id": "other-input", + "reason": "input_required", + "message": "Choose another value.", + } + ) + await store.save(scope="tenant-a", thread_id="agent-thread", snapshot=stored_snapshot) confirm_response = client.post( "/snapshots", @@ -9723,7 +10684,10 @@ async def stream_fn(messages: Any, options: Any, **kwargs: Any): assert hydrate_response.status_code == 200 assert call_count == 1 events = _decode_sse_events(hydrate_response) - assert "outcome" not in events[-1] + if with_other_interrupt: + assert [interrupt["id"] for interrupt in _run_finished_interrupts(events[-1])] == ["other-input"] + else: + assert "outcome" not in events[-1] messages = _latest_messages_snapshot(hydrate_response) assert any( message.get("role") == "assistant" and message.get("content") == "Changes confirmed and applied successfully!" @@ -10089,6 +11053,150 @@ def get_weather(city: str) -> str: assert ("tool", "function_result", "call_get_weather", None) in received +async def test_agent_endpoint_approved_result_is_persisted_before_provider_continuation_failure(): + """Hydration keeps an executed approval result when the provider continuation fails.""" + from agent_framework import AgentResponse, ResponseInvalidatedException, ResponseStream + + store = InMemoryAGUIThreadSnapshotStore() + client, agent, executed_cities = _build_weather_approval_endpoint(snapshot_store=store) + original_run = agent.run + + def failing_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates() -> AsyncIterator[AgentResponseUpdate]: + if False: # pragma: no cover + yield AgentResponseUpdate() + raise ResponseInvalidatedException("provider continuation failed") + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + agent.run = failing_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + resume_response = client.post( + "/approval", + json={ + "runId": "run-resume-failure", + "threadId": "thread-weather", + "messages": [], + "resume": [ + { + "interruptId": "call_get_weather", + "status": "resolved", + "payload": {"accepted": True}, + } + ], + }, + ) + + resume_events = _decode_sse_events(resume_response) + assert executed_cities == ["Seattle"] + assert [ + (event["toolCallId"], event["content"]) for event in resume_events if event.get("type") == "TOOL_CALL_RESULT" + ] == [("call_get_weather", "Sunny in Seattle")] + assert [event["code"] for event in resume_events if event.get("type") == "RUN_ERROR"] == [ + "ResponseInvalidatedException" + ] + + snapshot = await store.get(scope="tenant-a", thread_id="thread-weather") + assert snapshot is not None + assert snapshot.interrupt is None + assert "Sunny in Seattle" in str(snapshot.messages) + + +async def test_agent_endpoint_rejected_approval_is_persisted_before_provider_continuation_failure(): + """Hydration cannot replay a rejected approval when provider continuation fails.""" + from agent_framework import AgentResponse, ResponseInvalidatedException, ResponseStream + + store = InMemoryAGUIThreadSnapshotStore() + client, agent, executed_cities = _build_weather_approval_endpoint(snapshot_store=store) + original_run = agent.run + + def failing_run(*args: Any, **kwargs: Any) -> Any: + if not kwargs.get("stream", False): + return original_run(*args, **kwargs) + + async def updates() -> AsyncIterator[AgentResponseUpdate]: + if False: # pragma: no cover + yield AgentResponseUpdate() + raise ResponseInvalidatedException("provider continuation failed") + + return ResponseStream(updates(), finalizer=AgentResponse.from_updates) + + agent.run = failing_run # type: ignore[assignment, method-assign] # ty: ignore[invalid-assignment] + resume_response = client.post( + "/approval", + json={ + "runId": "run-reject-failure", + "threadId": "thread-weather", + "messages": [], + "resume": [ + { + "interruptId": "call_get_weather", + "status": "resolved", + "payload": {"approved": False}, + } + ], + }, + ) + + resume_events = _decode_sse_events(resume_response) + assert executed_cities == [] + assert [event["code"] for event in resume_events if event.get("type") == "RUN_ERROR"] == [ + "ResponseInvalidatedException" + ] + + snapshot = await store.get(scope="tenant-a", thread_id="thread-weather") + assert snapshot is not None + assert snapshot.interrupt is None + assert "Tool call invocation was rejected by user" in str(snapshot.messages) + + +async def test_agent_endpoint_approval_resume_preserves_other_stored_interrupts(): + """A successful approval continuation retains unrelated stored interrupts.""" + store = InMemoryAGUIThreadSnapshotStore() + client, _, executed_cities = _build_weather_approval_endpoint(snapshot_store=store) + snapshot = await store.get(scope="tenant-a", thread_id="thread-weather") + assert snapshot is not None + assert snapshot.interrupt is not None + snapshot.interrupt.append( + { + "id": "other-input", + "reason": "input_required", + "message": "Choose another value.", + } + ) + await store.save(scope="tenant-a", thread_id="thread-weather", snapshot=snapshot) + + resume_response = client.post( + "/approval", + json={ + "runId": "run-resume-preserve", + "threadId": "thread-weather", + "messages": [], + "resume": [ + { + "interruptId": "call_get_weather", + "status": "resolved", + "payload": {"accepted": True}, + } + ], + }, + ) + + assert resume_response.status_code == 200 + assert executed_cities == ["Seattle"] + final_snapshot = await store.get(scope="tenant-a", thread_id="thread-weather") + assert final_snapshot is not None + assert final_snapshot.interrupt == [ + { + "id": "other-input", + "reason": "input_required", + "message": "Choose another value.", + } + ] + + async def test_agent_endpoint_cancelled_approval_resume_clears_persisted_interrupt(): """Cancelling an approval resume cancels the whole approval set and clears the stored interrupt prompt.""" executed_cities: list[str] = [] 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 ccd53dbb5e..f4d7b02fde 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,8 +21,8 @@ ToolCallArgsEvent, ToolCallStartEvent, ) -from agent_framework import AgentResponseUpdate, Content, Message, ResponseStream -from agent_framework.exceptions import AgentInvalidResponseException +from agent_framework import AgentResponse, AgentResponseUpdate, Content, Message, ResponseStream +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 @@ -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,546 @@ async def test_service_session_rejects_disabled_provider_storage(): ] +@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 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_usage( + { + "input_token_count": 1, + "output_token_count": 1, + "total_token_count": 2, + } + ) + ], + role="assistant", + ) + if failure_mode == "stream": + raise invalidated + + stream = ResponseStream(updates(), finalizer=AgentResponse.from_updates) + if failure_mode == "result_hook": + + def invalidate_result(response: AgentResponse[Any]) -> None: + del response + raise invalidated + + stream.with_result_hook(invalidate_result) + return stream + + 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) + + 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"}], + } + ) + ] + + 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_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 + + 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_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 + + 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() + 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, + 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=[ + 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", + ) + ] + ) + 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" + 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(): + """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] 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 1841a0d37c..2e6b8ed56c 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."""