Server runtime (#342)

This commit is contained in:
Adrian Lyjak
2026-02-10 14:25:47 -05:00
committed by GitHub
parent e981f7312a
commit 528d5623a3
46 changed files with 5355 additions and 2437 deletions
+5
View File
@@ -0,0 +1,5 @@
---
"llama-agents-client": minor
---
Add SSE event streaming with sequence-based cursors and automatic reconnection on connection drop
+1 -1
View File
@@ -2,4 +2,4 @@
"llama-agents-server": minor
---
Test pre-release functioning
Refactor server internals from monolithic handler to composable runtime decorators (ServerRuntimeDecorator, PersistenceDecorator, IdleReleaseDecorator) enabling pluggable server runtimes
+5
View File
@@ -0,0 +1,5 @@
---
"llama-agents-server": minor
---
Add tick storage, event storage with SSE subscription, per-run state stores, and centralized handler status transitions to AbstractWorkflowStore and SQLite/memory implementations
@@ -8,6 +8,7 @@ from typing import (
Any,
AsyncGenerator,
AsyncIterator,
Literal,
overload,
)
@@ -202,55 +203,89 @@ class WorkflowClient:
self,
handler_id: str,
include_internal_events: bool = False,
lock_timeout: float = 1,
after_sequence: int | Literal["now"] = -1,
max_reconnect_attempts: int = 3,
) -> AsyncGenerator[EventEnvelopeWithMetadata, None]:
"""
Stream events as they are produced by the workflow.
Uses SSE (Server-Sent Events) mode and automatically reconnects from
the last received event on connection drops.
Args:
handler_id (str): ID of the handler running the workflow
include_internal_events (bool): Include internal workflow events. Defaults to False.
lock_timeout (float): Timeout (in seconds) for acquiring the lock to iterate over the events.
after_sequence (int | str): Sequence number to start streaming after. Defaults to -1 (all events). Use ``"now"`` to only receive new events.
max_reconnect_attempts (int): Maximum number of reconnect attempts on connection drop. Defaults to 3.
Returns:
AsyncGenerator[EventEnvelopeWithMetadata, None]: Generator for the events that are streamed as instances of `EventEnvelopeWithMetadata`.
"""
incl_inter = "true" if include_internal_events else "false"
url = f"/events/{handler_id}"
last_sequence = after_sequence
attempts = 0
async with self._get_client() as client:
try:
async with client.stream(
"GET",
url,
params={
"sse": "false",
"include_internal": incl_inter,
"acquire_timeout": lock_timeout,
},
headers={"Connection": "keep-alive"},
timeout=None,
) as response:
# Handle different response codes
if response.status_code == 404:
raise ValueError("Handler not found")
elif response.status_code == 204:
# Handler completed, no more events
return
while True:
async with self._get_client() as client:
try:
async with client.stream(
"GET",
url,
params={
"sse": "true",
"include_internal": incl_inter,
"after_sequence": str(last_sequence),
},
headers={"Connection": "keep-alive"},
timeout=None,
) as response:
if response.status_code == 404:
raise ValueError("Handler not found")
elif response.status_code == 204:
return
_raise_for_status_with_body(response)
_raise_for_status_with_body(response)
async for line in response.aiter_lines():
if line.strip(): # Skip empty lines
event = EventEnvelopeWithMetadata.model_validate_json(line)
yield event
# Reset attempts on successful connection
attempts = 0
except httpx.TimeoutException:
raise TimeoutError(
f"Timeout waiting for events from handler {handler_id}"
)
except httpx.RequestError as e:
raise ConnectionError(f"Failed to connect to event stream: {e}")
# Parse SSE stream: "id: N\ndata: {...}\n\n"
current_id: str | None = None
async for line in response.aiter_lines():
stripped = line.strip()
if not stripped:
# Empty line = end of SSE event
continue
if stripped.startswith("id:"):
current_id = stripped[3:].strip()
elif stripped.startswith("data:"):
data = stripped[5:].strip()
event = EventEnvelopeWithMetadata.model_validate_json(
data
)
if current_id is not None:
try:
last_sequence = int(current_id)
except ValueError:
pass
current_id = None
yield event
# Stream ended normally (server closed connection)
return
except httpx.TimeoutException:
raise TimeoutError(
f"Timeout waiting for events from handler {handler_id}"
)
except (httpx.RequestError, ConnectionError):
attempts += 1
if attempts > max_reconnect_attempts:
raise ConnectionError(
f"Failed to connect to event stream after {max_reconnect_attempts} attempts"
)
# Retry from last received sequence
async def send_event(
self,
@@ -2,7 +2,7 @@ from __future__ import annotations
from typing import Any, Literal
from pydantic import BaseModel, Field
from pydantic import BaseModel
from workflows.representation import WorkflowGraph
from .serializable_events import EventEnvelopeWithMetadata
@@ -35,13 +35,6 @@ class HandlersListResponse(BaseModel):
class HealthResponse(BaseModel):
status: Literal["healthy"]
loaded_workflows: int = Field(
description="Number of workflow handlers currently loaded in memory"
)
active_workflows: int = Field(
description="Number of workflow handlers that are active (not idle)"
)
idle_workflows: int = Field(description="Number of workflow handlers that are idle")
class WorkflowsListResponse(BaseModel):
@@ -35,13 +35,7 @@ class GreetingWorkflow(Workflow):
return OutputEvent(greeting=f"{ev.greeting}{'!' * ev.exclamation_marks}")
greeting_wf = GreetingWorkflow()
class CrashingWorkflow(Workflow):
@step
async def crashing_step(self, ev: StartEvent) -> StopEvent:
raise ValueError("Workflow crashed intentionally")
crashing_wf = CrashingWorkflow()
@@ -1,18 +1,24 @@
from __future__ import annotations
from contextlib import asynccontextmanager
from typing import AsyncIterator, Union
from unittest.mock import AsyncMock
import httpx
import pytest
from client_test_workflows import (
CrashingWorkflow,
GreetEvent,
GreetingWorkflow,
InputEvent,
OutputEvent,
crashing_wf,
greeting_wf,
)
from httpx import ASGITransport, AsyncClient
from llama_agents.client import WorkflowClient
from llama_agents.client.protocol.serializable_events import (
EventEnvelopeWithMetadata,
)
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from llama_agents.server import MemoryWorkflowStore
from llama_agents.server.server import WorkflowServer
@@ -20,8 +26,8 @@ from llama_agents.server.server import WorkflowServer
def server() -> WorkflowServer:
# Use MemoryWorkflowStore so get_handlers() can retrieve from persistence
ws = WorkflowServer(workflow_store=MemoryWorkflowStore())
ws.add_workflow(name="greeting", workflow=greeting_wf)
ws.add_workflow(name="crashing", workflow=crashing_wf)
ws.add_workflow(name="greeting", workflow=GreetingWorkflow())
ws.add_workflow(name="crashing", workflow=CrashingWorkflow())
return ws
@@ -260,3 +266,121 @@ async def test_error_message_format(client: WorkflowClient) -> None:
"404 Not Found for POST http://test/workflows/nonexistent_workflow/run. Response: Workflow not found"
== error_message
) # Status code
def _envelope(msg: str) -> EventEnvelopeWithMetadata:
return EventEnvelopeWithMetadata(
value={"msg": msg}, qualified_name=None, type="TestEvent", types=None
)
# Each "connection" in a script is a list of SSE events to yield, optionally
# ending with an exception to simulate a disconnect. A bare exception means
# the connection fails before yielding any data.
ConnectionScript = Union[
list[Union[tuple[int, EventEnvelopeWithMetadata], Exception]], Exception
]
class FakeStreamClient:
"""Mock httpx client that replays a scripted sequence of SSE connections."""
def __init__(self, script: list[ConnectionScript]) -> None:
self._script = list(script)
self.captured_params: list[dict[str, str]] = []
self._call = 0
@asynccontextmanager
async def stream(
self,
method: str,
url: str,
params: dict[str, str] | None = None,
**kwargs: object,
) -> AsyncIterator[AsyncMock]:
self.captured_params.append(params or {})
assert self._call < len(self._script), "More connections than scripted"
entry = self._script[self._call]
self._call += 1
if isinstance(entry, Exception):
raise entry
events = entry
tail_error: Exception | None = None
# If the last element is an exception, pop it as a mid-stream error
if events and isinstance(events[-1], Exception):
tail_error = events[-1] # type: ignore[assignment]
events = events[:-1] # type: ignore[assignment]
resp = AsyncMock()
resp.status_code = 200
async def aiter_lines() -> AsyncIterator[str]:
for seq, env in events: # type: ignore[union-attr]
yield f"id: {seq}"
yield f"data: {env.model_dump_json()}"
yield ""
if tail_error is not None:
raise tail_error
resp.aiter_lines = aiter_lines
yield resp
async def _collect(
script: list[ConnectionScript], **kwargs: object
) -> list[EventEnvelopeWithMetadata]:
fake = FakeStreamClient(script)
wf_client = WorkflowClient(httpx_client=fake) # type: ignore[arg-type]
events = [
e
async for e in wf_client.get_workflow_events(handler_id="h", **kwargs) # type: ignore[arg-type]
]
return events
@pytest.mark.asyncio
async def test_reconnect_resumes_from_last_sequence() -> None:
e1, e2, e3 = _envelope("first"), _envelope("second"), _envelope("third")
fake = FakeStreamClient(
[
[(0, e1), httpx.RemoteProtocolError("reset")],
[(1, e2), (2, e3)],
]
)
wf_client = WorkflowClient(httpx_client=fake) # type: ignore[arg-type]
events = [e async for e in wf_client.get_workflow_events(handler_id="h")]
assert [e.value["msg"] for e in events] == ["first", "second", "third"]
assert fake.captured_params[0]["after_sequence"] == "-1"
assert fake.captured_params[1]["after_sequence"] == "0"
@pytest.mark.asyncio
async def test_reconnect_exceeds_max_attempts_raises() -> None:
with pytest.raises(ConnectionError, match="after 2 attempts"):
await _collect(
[httpx.ConnectError("refused")] * 3,
max_reconnect_attempts=2,
)
@pytest.mark.asyncio
async def test_reconnect_resets_attempts_on_success() -> None:
e1, e2 = _envelope("a"), _envelope("b")
events = await _collect(
[
[(0, e1), httpx.ReadError("broken")],
httpx.ReadError("broken again"),
[(1, e2)],
],
max_reconnect_attempts=2,
)
assert [e.value["msg"] for e in events] == ["a", "b"]
@pytest.mark.asyncio
async def test_timeout_exception_not_retried() -> None:
with pytest.raises(TimeoutError, match="Timeout"):
await _collect([httpx.ReadTimeout("timed out")])
+5 -2
View File
@@ -11,7 +11,8 @@ dev = [
"pytest-cov>=7.0.0",
"pytest-timeout>=2.4.0",
"pytest-xdist>=3.8.0",
"time-machine>=2.19.0,<3.0.0"
"time-machine>=2.19.0,<3.0.0",
"llama-agents-integration-tests"
]
[project]
@@ -24,7 +25,8 @@ dependencies = [
"llama-index-workflows>=2.12.0,<3.0.0",
"llama-agents-client>=0.1.0,<0.2.0",
"starlette>=0.39.0",
"uvicorn>=0.32.0"
"uvicorn>=0.32.0",
"httpx>=0.27.0"
]
[project.scripts]
@@ -49,3 +51,4 @@ module-name = "llama_agents.server"
[tool.uv.sources]
llama-index-workflows = {workspace = true}
llama-agents-client = {workspace = true}
llama-agents-integration-tests = {workspace = true}
@@ -6,13 +6,15 @@ from ._store.abstract_workflow_store import (
HandlerQuery,
PersistentHandler,
)
from ._store.memory_workflow_store import MemoryWorkflowStore
from ._store.sqlite.sqlite_workflow_store import SqliteWorkflowStore
from .server import WorkflowServer
__all__ = [
"WorkflowServer",
"AbstractWorkflowStore",
"HandlerQuery",
"PersistentHandler",
"WorkflowServer",
"MemoryWorkflowStore",
"SqliteWorkflowStore",
]
@@ -18,7 +18,6 @@ from llama_agents.client.protocol import (
WorkflowEventsListResponse,
WorkflowGraphResponse,
WorkflowSchemaResponse,
is_status_completed,
)
from llama_agents.client.protocol.serializable_events import (
EventEnvelope,
@@ -35,15 +34,21 @@ from starlette.routing import Route
from starlette.schemas import SchemaGenerator
from starlette.staticfiles import StaticFiles
from workflows import Context, Workflow
from workflows.events import InternalDispatchEvent, StartEvent
from workflows.events import Event, InternalDispatchEvent, StartEvent
from workflows.representation import get_workflow_representation
from workflows.utils import _nanoid as nanoid
from ._handler import NoLockAvailable, _NamedWorkflow, _WorkflowHandler
from ._service import _WorkflowService
from ._service import (
EventSendError,
HandlerCompletedError,
HandlerNotFoundError,
_WorkflowService,
)
from ._store.abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
Status,
is_terminal_status,
)
logger = logging.getLogger()
@@ -61,6 +66,7 @@ class _WorkflowAPI:
assets_path: Path = _DEFAULT_ASSETS_PATH,
) -> None:
self._service = service
self._additional_events: dict[str, list[type[Event]]] = {}
middleware = middleware or [
Middleware(
@@ -90,6 +96,19 @@ class _WorkflowAPI:
"/", app=StaticFiles(directory=assets_path, html=True), name="ui"
)
def register_additional_events(self, name: str, events: list[type[Event]]) -> None:
self._additional_events[name] = events
def get_workflow_events(self, workflow_name: str) -> list[type[Event]]:
workflow = self._service.get_workflow(workflow_name)
if workflow is None:
return []
return workflow.events + (self._additional_events.get(workflow_name) or [])
def event_registry(self, workflow_name: str) -> dict[str, type[Event]]:
"""Return a name→type mapping of events for the given workflow."""
return {e.__name__: e for e in self.get_workflow_events(workflow_name)}
def _routes(self) -> list[Route]:
return [
Route("/workflows", self._list_workflows, methods=["GET"]),
@@ -236,28 +255,11 @@ class _WorkflowAPI:
status:
type: string
example: healthy
loaded_workflows:
type: integer
description: Number of workflow handlers currently loaded in memory
active_workflows:
type: integer
description: Number of workflow handlers that are active (not idle)
idle_workflows:
type: integer
description: Number of workflow handlers that are idle
required: [status, loaded_workflows, active_workflows, idle_workflows]
required: [status]
"""
loaded = len(self._service._handlers)
idle = sum(
1 for h in self._service._handlers.values() if h.idle_since is not None
)
active = loaded - idle
return JSONResponse(
HealthResponse(
status="healthy",
loaded_workflows=loaded,
active_workflows=active,
idle_workflows=idle,
).model_dump()
)
@@ -280,7 +282,7 @@ class _WorkflowAPI:
type: string
required: [workflows]
"""
workflow_names = list(self._service._workflows.keys())
workflow_names = self._service.get_workflow_names()
return JSONResponse({"workflows": workflow_names})
async def _list_workflow_events(self, request: Request) -> JSONResponse:
@@ -314,12 +316,10 @@ class _WorkflowAPI:
raise HTTPException(status_code=400, detail="name param is required")
name = request.path_params["name"]
if name not in self._service._workflows:
if self._service.get_workflow(name) is None:
raise HTTPException(status_code=404, detail=f"Workflow '{name}' not found")
events = self._service._workflows[name].events + (
self._service._additional_events.get(name, []) or []
)
events = self.get_workflow_events(name)
return JSONResponse(
WorkflowEventsListResponse(
@@ -377,46 +377,36 @@ class _WorkflowAPI:
"""
workflow = self._extract_workflow(request)
context, start_event, handler_id = await self._extract_run_params(
request, workflow.workflow, workflow.name
request, workflow, workflow.workflow_name
)
if start_event is not None:
input_ev = workflow.workflow.start_event_class.model_validate(start_event)
input_ev = workflow.start_event_class.model_validate(start_event)
else:
input_ev = None
try:
wrapper = await self._service.start_workflow(
workflow=_NamedWorkflow(name=workflow.name, workflow=workflow.workflow),
started = await self._service.start_workflow(
workflow=workflow,
handler_id=handler_id,
context=context,
start_event=input_ev,
)
handler = wrapper.run_handler
try:
await handler
status = 200
except Exception as e:
status = 500
logger.error(f"Error running workflow: {e}", exc_info=True)
if wrapper.task is not None:
try:
await wrapper.task
except Exception:
pass
# explicitly close handlers from this synchronous api so they don't linger with events
# that no-one is listening for
await self._service.close_handler(wrapper)
return JSONResponse(
wrapper.to_response_model().model_dump(), status_code=status
)
except Exception as e:
status = 500
logger.error(f"Error running workflow: {e}", exc_info=True)
raise HTTPException(
detail=f"Error running workflow: {e}", status_code=status
)
raise HTTPException(detail=f"Error running workflow: {e}", status_code=500)
try:
handler_data = await self._service.await_workflow(started)
status = 200 if handler_data.status == "completed" else 500
except Exception as e:
logger.error(f"Error running workflow: {e}", exc_info=True)
handler_data = await self._service.load_handler(handler_id)
status = 500
return JSONResponse(
handler_data.model_dump() if handler_data else {}, status_code=status
)
async def _get_events_schema(self, request: Request) -> JSONResponse:
"""
@@ -453,14 +443,14 @@ class _WorkflowAPI:
"""
workflow = self._extract_workflow(request)
try:
start_event_schema = workflow.workflow.start_event_class.model_json_schema()
start_event_schema = workflow.start_event_class.model_json_schema()
except Exception as e:
raise HTTPException(
detail=f"Error getting schema of start event for workflow: {e}",
status_code=500,
)
try:
stop_event_schema = workflow.workflow.stop_event_class.model_json_schema()
stop_event_schema = workflow.stop_event_class.model_json_schema()
except Exception as e:
raise HTTPException(
detail=f"Error getting schema of stop event for workflow: {e}",
@@ -506,7 +496,7 @@ class _WorkflowAPI:
"""
workflow = self._extract_workflow(request)
try:
workflow_graph = get_workflow_representation(workflow.workflow)
workflow_graph = get_workflow_representation(workflow)
except Exception as e:
raise HTTPException(
detail=f"Error while getting JSON workflow representation: {e}",
@@ -561,47 +551,105 @@ class _WorkflowAPI:
"""
workflow = self._extract_workflow(request)
context, start_event, handler_id = await self._extract_run_params(
request, workflow.workflow, workflow.name
request, workflow, workflow.workflow_name
)
if start_event is not None:
input_ev = workflow.workflow.start_event_class.model_validate(start_event)
input_ev = workflow.start_event_class.model_validate(start_event)
else:
input_ev = None
try:
wrapper = await self._service.start_workflow(
workflow=_NamedWorkflow(name=workflow.name, workflow=workflow.workflow),
handler_data = await self._service.start_workflow(
workflow=workflow,
handler_id=handler_id,
context=context,
start_event=input_ev,
)
except Exception as e:
raise HTTPException(
detail=f"Initial persistence failed: {e}", status_code=500
)
return JSONResponse(wrapper.to_response_model().model_dump())
return JSONResponse(handler_data.model_dump())
async def _load_handler(self, handler_id: str) -> HandlerData:
wrapper = self._service._handlers.get(handler_id)
if wrapper is None:
found = await self._service._workflow_store.query(
HandlerQuery(handler_id_in=[handler_id])
)
if not found:
raise HTTPException(detail="Handler not found", status_code=404)
existing = found[0]
return _WorkflowHandler.handler_data_from_persistent(existing)
else:
if wrapper.run_handler.done() and wrapper.task is not None:
try:
await wrapper.task # make sure its fully done
except Exception:
# failed workflows raise their exception here
pass # failed workflows raise their exception here
handler_data = await self._service.load_handler(handler_id)
if handler_data is None:
raise HTTPException(detail="Handler not found", status_code=404)
return handler_data
return wrapper.to_response_model()
async def _resolve_event_stream(
self,
handler_id: str,
*,
after_sequence: int | None,
include_internal: bool,
include_qualified_name: bool,
) -> AsyncGenerator[tuple[int, EventEnvelopeWithMetadata], None] | None:
"""Resolve a handler to an event stream.
Args:
handler_id: The handler to stream events for.
after_sequence: Resume after this sequence number. None means "now"
(skip historical events).
include_internal: Whether to include internal dispatch events.
include_qualified_name: Whether to include qualified_name in envelopes.
Returns:
An async generator of (sequence, envelope) tuples, or None if the
handler is completed and all events have been consumed.
Raises:
HTTPException: 404 if handler not found or has no run.
"""
store = self._service.store
# Resolve handler_id → run_id via persistence
found = await store.query(HandlerQuery(handler_id_in=[handler_id]))
if not found:
raise HTTPException(detail="Handler not found", status_code=404)
persistent = found[0]
run_id = persistent.run_id
if run_id is None:
raise HTTPException(detail="Handler has no associated run", status_code=404)
# Resolve "now" cursor to current max sequence
if after_sequence is None:
all_current = await store.query_events(run_id)
after_sequence = all_current[-1].sequence if all_current else -1
# Check if already fully consumed
remaining_events = await store.query_events(
run_id, after_sequence=after_sequence
)
if not remaining_events:
all_events = await store.query_events(run_id)
run_is_complete = is_terminal_status(persistent.status) or (
bool(all_events)
and AbstractWorkflowStore._is_terminal_event(all_events[-1])
)
if run_is_complete:
return None
_INTERNAL_EVENT_TYPE = InternalDispatchEvent.__name__
async def event_gen() -> AsyncGenerator[
tuple[int, EventEnvelopeWithMetadata], None
]:
async for stored_event in store.subscribe_events(
run_id,
after_sequence=after_sequence, # type: ignore[arg-type]
):
envelope = stored_event.event
types = (envelope.types or []) + [envelope.type]
if not include_internal and _INTERNAL_EVENT_TYPE in types:
continue
if not include_qualified_name:
envelope = envelope.model_copy(update={"qualified_name": None})
yield stored_event.sequence, envelope
return event_gen()
async def _get_workflow_result(self, request: Request) -> JSONResponse:
"""
@@ -646,7 +694,7 @@ class _WorkflowAPI:
handler_data = await self._load_handler(handler_id)
status = (
202
if handler_data.status in "running"
if handler_data.status == "running"
else 200
if handler_data.status == "completed"
else 500
@@ -706,7 +754,7 @@ class _WorkflowAPI:
handler_data = await self._load_handler(handler_id)
status = (
202
if handler_data.status in "running"
if handler_data.status == "running"
else 200
if handler_data.status == "completed"
else 500
@@ -720,8 +768,10 @@ class _WorkflowAPI:
description: |
Streams events produced by a workflow execution. Events are emitted as
newline-delimited JSON by default, or as Server-Sent Events when `sse=true`.
Event data is returned as an envelope that preserves backward-compatible fields
and adds metadata for type-safety on the client:
Multiple clients can stream the same handler concurrently. Disconnected
clients can resume from their last-seen position via `after_sequence`.
Event data is returned as an envelope:
{
"value": <pydantic serialized value>,
"types": [<class names from MRO excluding the event class and base Event>],
@@ -729,9 +779,6 @@ class _WorkflowAPI:
"qualified_name": <python module path + class name>,
}
Event queue is mutable. Elements are added to the queue by the workflow handler, and removed by any consumer of the queue.
The queue is protected by a lock that is acquired by the consumer, so only one consumer of the queue at a time is allowed.
parameters:
- in: path
name: handler_id
@@ -754,12 +801,20 @@ class _WorkflowAPI:
default: false
description: If true, include internal workflow events (e.g., step state changes).
- in: query
name: acquire_timeout
name: after_sequence
required: false
schema:
type: number
default: 1
description: Timeout for acquiring the lock to iterate over the events.
oneOf:
- type: integer
- type: string
enum: [now]
default: now
description: >
Resume streaming after this event sequence number. Use -1
to start from the beginning, or "now" (default) to skip historical events and
only receive events appended after the request is made.
In SSE mode, the Last-Event-ID request header takes priority over
this parameter.
- in: query
name: include_qualified_name
required: false
@@ -791,11 +846,12 @@ class _WorkflowAPI:
type: string
description: The qualified name of the event.
required: [value, type]
204:
description: Handler completed and all events already consumed
404:
description: Handler not found
"""
handler_id = request.path_params["handler_id"]
timeout = request.query_params.get("acquire_timeout", "1").lower()
include_internal = (
request.query_params.get("include_internal", "false").lower() == "true"
)
@@ -803,64 +859,49 @@ class _WorkflowAPI:
request.query_params.get("include_qualified_name", "true").lower() == "true"
)
sse = request.query_params.get("sse", "true").lower() == "true"
try:
timeout = float(timeout)
except ValueError:
raise HTTPException(
detail=f"Invalid acquire_timeout: '{timeout}'", status_code=400
)
handler = self._service._handlers.get(handler_id)
if handler is None:
# Try to reload from persistence (for released idle workflows)
after_sequence_str = request.query_params.get("after_sequence", "now")
after_sequence_is_now = after_sequence_str.lower() == "now"
if after_sequence_is_now:
after_sequence: int | None = None # resolved by helper
else:
try:
handler, persisted = await self._service.try_reload_handler(handler_id)
except Exception as e:
after_sequence = int(after_sequence_str)
except ValueError:
raise HTTPException(
detail=f"Failed to reload handler: {e}", status_code=500
detail=f"Invalid after_sequence: '{after_sequence_str}'",
status_code=400,
)
if handler is None:
if persisted:
status = persisted.status
if status in {"completed", "failed", "cancelled"}:
raise HTTPException(
detail="Handler is completed", status_code=204
)
raise HTTPException(detail="Handler not found", status_code=404)
if handler.queue.empty() and handler.task is not None and handler.task.done():
# https://html.spec.whatwg.org/multipage/server-sent-events.html
# Clients will reconnect if the connection is closed; a client can
# be told to stop reconnecting using the HTTP 204 No Content response code.
# SSE Last-Event-ID header overrides after_sequence
if sse:
last_event_id = request.headers.get("last-event-id")
if last_event_id is not None:
try:
after_sequence = int(last_event_id)
except ValueError:
pass # Ignore non-integer Last-Event-ID
gen = await self._resolve_event_stream(
handler_id,
after_sequence=after_sequence,
include_internal=include_internal,
include_qualified_name=include_qualified_name,
)
if gen is None:
raise HTTPException(detail="Handler is completed", status_code=204)
# Get raw_event query parameter
media_type = "text/event-stream" if sse else "application/x-ndjson"
try:
generator = await handler.acquire_events_stream(timeout=timeout)
except NoLockAvailable as e:
raise HTTPException(
detail=f"No lock available to acquire after {timeout}s timeout",
status_code=409,
) from e
async def event_stream(handler: _WorkflowHandler) -> AsyncGenerator[str, None]:
async for event in generator:
if not include_internal and isinstance(event, InternalDispatchEvent):
continue
envelope = EventEnvelopeWithMetadata.from_event(
event, include_qualified_name=include_qualified_name
)
async def format_stream() -> AsyncGenerator[str, None]:
async for sequence, envelope in gen:
payload = envelope.model_dump_json()
if sse:
# emit as untyped data. Difficult to subscribe to dynamic event types with SSE.
yield f"data: {payload}\n\n"
yield f"id: {sequence}\ndata: {payload}\n\n"
else:
yield f"{payload}\n"
await asyncio.sleep(0)
return StreamingResponse(event_stream(handler), media_type=media_type)
return StreamingResponse(format_stream(), media_type=media_type)
async def _get_handlers(self, request: Request) -> JSONResponse:
"""
@@ -931,7 +972,7 @@ class _WorkflowAPI:
if status_values is not None
else None
)
persistent_handlers = await self._service._workflow_store.query(
persistent_handlers = await self._service.query_handlers(
HandlerQuery(status_in=status_in, workflow_name_in=workflow_name_in)
)
items = [
@@ -1012,84 +1053,49 @@ class _WorkflowAPI:
"""
handler_id = request.path_params["handler_id"]
# Check if handler exists
wrapper = self._service._handlers.get(handler_id)
if wrapper is not None and is_status_completed(wrapper.status):
raise HTTPException(detail="Workflow already completed", status_code=409)
if wrapper is None:
# Try to reload from persistence (for released idle workflows)
try:
wrapper, persisted = await self._service.try_reload_handler(handler_id)
except Exception as e:
raise HTTPException(
detail=f"Failed to reload handler: {e}", status_code=500
)
if wrapper is None:
# Check if it exists but is completed
if persisted and is_status_completed(persisted.status):
raise HTTPException(
detail="Workflow already completed", status_code=409
)
elif persisted is None:
raise HTTPException(detail="Handler not found", status_code=404)
else:
# Shouldn't really happen
raise HTTPException(
detail=f"Failed to resume incomplete handler with status {persisted.status}",
status_code=500,
)
# Immediately mark active to cancel the idle timer before it can fire.
# This prevents a race where the timer releases the handler before we
# finish processing the event.
wrapper.mark_active()
handler = wrapper.run_handler
# Get the context
ctx = handler.ctx
if ctx is None:
raise HTTPException(detail="Context not available", status_code=500)
# Parse request body
try:
body = await request.json()
event_str = body.get("event")
step = body.get("step")
if not event_str:
raise HTTPException(detail="Event data is required", status_code=400)
# Deserialize the event
try:
event = EventEnvelope.parse(
event_str, self._service.event_registry(wrapper.workflow_name)
)
except EventValidationError as e:
raise HTTPException(detail=str(e), status_code=400)
except Exception as e:
raise HTTPException(
detail=f"Failed to deserialize event: {e}", status_code=400
)
# Send the event to the context
try:
ctx.send_event(event, step=step)
except Exception as e:
raise HTTPException(
detail=f"Failed to send event: {e}", status_code=400
)
return JSONResponse(SendEventResponse(status="sent").model_dump())
except HTTPException:
raise
except Exception as e:
raise HTTPException(
detail=f"Error processing request: {e}", status_code=500
)
event_data = body.get("event")
step = body.get("step")
if not event_data:
raise HTTPException(detail="Event data is required", status_code=400)
try:
handler_data = await self._service.resolve_handler(handler_id)
except HandlerNotFoundError:
raise HTTPException(detail="Handler not found", status_code=404)
except HandlerCompletedError:
raise HTTPException(detail="Workflow already completed", status_code=409)
try:
event = EventEnvelope.parse(
event_data, self.event_registry(handler_data.workflow_name)
)
except EventValidationError as e:
raise HTTPException(detail=str(e), status_code=400)
except Exception as e:
raise HTTPException(
detail=f"Failed to deserialize event: {e}", status_code=400
)
try:
await self._service.send_event(handler_id, event, step=step)
except HandlerNotFoundError:
raise HTTPException(detail="Handler not found", status_code=404)
except HandlerCompletedError:
raise HTTPException(detail="Workflow already completed", status_code=409)
except EventSendError as e:
raise HTTPException(detail=str(e), status_code=500)
return JSONResponse(SendEventResponse(status="sent").model_dump())
async def _cancel_handler(self, request: Request) -> JSONResponse:
"""
---
@@ -1130,40 +1136,25 @@ class _WorkflowAPI:
# Simple boolean parsing aligned with other APIs (e.g., `sse`): only "true" enables
purge = request.query_params.get("purge", "false").lower() == "true"
wrapper = self._service._handlers.get(handler_id)
if wrapper is None and not purge:
result = await self._service.cancel_handler(handler_id, purge=purge)
if result is None:
raise HTTPException(detail="Handler not found", status_code=404)
# Close the handler if it exists (this will cancel and trigger auto-checkpoint)
if wrapper is not None:
await self._service.close_handler(wrapper)
# Handle persistence
if purge:
n_deleted = await self._service._workflow_store.delete(
HandlerQuery(handler_id_in=[handler_id])
)
if n_deleted == 0:
raise HTTPException(detail="Handler not found", status_code=404)
return JSONResponse(
CancelHandlerResponse(
status="deleted" if purge else "cancelled"
).model_dump()
)
return JSONResponse(CancelHandlerResponse(status=result).model_dump())
#
# Private methods
#
def _extract_workflow(self, request: Request) -> _NamedWorkflow:
def _extract_workflow(self, request: Request) -> Workflow:
if "name" not in request.path_params:
raise HTTPException(detail="'name' parameter missing", status_code=400)
name = request.path_params["name"]
if name not in self._service._workflows:
workflow = self._service.get_workflow(name)
if workflow is None:
raise HTTPException(detail="Workflow not found", status_code=404)
return _NamedWorkflow(name=name, workflow=self._service._workflows[name])
return workflow
async def _extract_run_params(
self, request: Request, workflow: Workflow, workflow_name: str
@@ -1184,7 +1175,7 @@ class _WorkflowAPI:
try:
start_event = EventEnvelope.parse(
start_event_data,
self._service.event_registry(workflow_name),
self.event_registry(workflow_name),
explicit_event=workflow.start_event_class,
)
@@ -1205,18 +1196,6 @@ class _WorkflowAPI:
context = None
if context_data:
context = Context.from_dict(workflow=workflow, data=context_data)
elif handler_id:
persisted_handlers = await self._service._workflow_store.query(
HandlerQuery(
handler_id_in=[handler_id],
workflow_name_in=[workflow_name],
status_in=["completed"],
)
)
if len(persisted_handlers) == 0:
raise HTTPException(detail="Handler not found", status_code=404)
context = Context.from_dict(workflow, persisted_handlers[0].ctx)
handler_id = handler_id or nanoid()
return (context, start_event, handler_id)
@@ -1,399 +0,0 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
import logging
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import AsyncGenerator, Awaitable, Callable
from llama_agents.client.protocol import HandlerData
from llama_agents.client.protocol.serializable_events import (
EventEnvelopeWithMetadata,
)
from llama_index_instrumentation.dispatcher import instrument_tags
from workflows.errors import WorkflowRuntimeError
from workflows.events import (
Event,
StepState,
StepStateChanged,
StopEvent,
UnhandledEvent,
WorkflowIdleEvent,
)
from workflows.handler import WorkflowHandler
from workflows.workflow import Workflow
from ._store.abstract_workflow_store import (
AbstractWorkflowStore,
PersistentHandler,
Status,
)
logger = logging.getLogger()
@dataclass
class _WorkflowHandler:
"""A wrapper around a handler: WorkflowHandler. Necessary to monitor and dispatch events from the handler's stream_events"""
run_handler: WorkflowHandler
queue: asyncio.Queue[Event]
task: asyncio.Task[None] | None
# only one consumer of the queue at a time allowed
consumer_mutex: asyncio.Lock
# metadata
handler_id: str
workflow_name: str
started_at: datetime
updated_at: datetime
completed_at: datetime | None
# Dependencies for persistence
_workflow_store: AbstractWorkflowStore
_persistence_backoff: list[float]
_on_finish: Callable[[], Awaitable[None]] | None = None
idle_since: datetime | None = None
# Idle release support
_idle_release_timeout: timedelta | None = None
_on_idle_release: Callable[[_WorkflowHandler], Awaitable[None]] | None = None
_idle_release_timer: asyncio.Task[None] | None = None
_skip_checkpoint: bool = False # Set to prevent checkpointing stale handlers
def _as_persistent(self) -> PersistentHandler:
"""Persist the current handler state immediately to the workflow store."""
self.updated_at = datetime.now(timezone.utc)
if self.status in ("completed", "failed", "cancelled"):
self.completed_at = self.updated_at
persistent = PersistentHandler(
handler_id=self.handler_id,
workflow_name=self.workflow_name,
status=self.status,
run_id=self.run_handler.run_id,
error=self.error,
result=self.result,
started_at=self.started_at,
updated_at=self.updated_at,
completed_at=self.completed_at,
idle_since=self.idle_since,
ctx=self.run_handler.ctx.to_dict() if self.run_handler.ctx else {},
)
return persistent
async def persist(self, persistent: PersistentHandler) -> None:
await self._workflow_store.update(persistent)
async def checkpoint(self) -> None:
"""Persist with retry/backoff; cancel handler when retries exhausted."""
if self._skip_checkpoint:
logger.debug(f"Skipping checkpoint for handler {self.handler_id}")
return
backoffs = list(self._persistence_backoff)
try:
persistent = self._as_persistent()
except Exception as e:
logger.error(
f"Failed to checkpoint handler {self.handler_id} to persistent state. Is there non-serializable state in an event or the state store? {e}",
exc_info=True,
)
raise
while True:
try:
await self.persist(persistent)
return
except Exception as e:
backoff = backoffs.pop(0) if backoffs else None
if backoff is None:
logger.error(
f"Failed to checkpoint handler {self.handler_id} after final attempt. Failing the handler.",
exc_info=True,
)
# Cancel the underlying workflow; do not re-raise here to allow callers to decide behavior
try:
self.run_handler.cancel()
except Exception:
pass
raise
logger.error(
f"Failed to checkpoint handler {self.handler_id}. Retrying in {backoff} seconds: {e}"
)
await asyncio.sleep(backoff)
def to_response_model(self) -> HandlerData:
"""Convert runtime handler to API response model."""
return HandlerData(
handler_id=self.handler_id,
workflow_name=self.workflow_name,
run_id=self.run_handler.run_id,
status=self.status,
started_at=self.started_at.isoformat(),
updated_at=self.updated_at.isoformat(),
completed_at=self.completed_at.isoformat()
if self.completed_at is not None
else None,
error=self.error,
result=EventEnvelopeWithMetadata.from_event(self.result)
if self.result is not None
else None,
)
@staticmethod
def handler_data_from_persistent(persistent: PersistentHandler) -> HandlerData:
return HandlerData(
handler_id=persistent.handler_id,
workflow_name=persistent.workflow_name,
run_id=persistent.run_id,
status=persistent.status,
started_at=persistent.started_at.isoformat()
if persistent.started_at is not None
else datetime.now(timezone.utc).isoformat(),
updated_at=persistent.updated_at.isoformat()
if persistent.updated_at is not None
else None,
completed_at=persistent.completed_at.isoformat()
if persistent.completed_at is not None
else None,
error=persistent.error,
result=EventEnvelopeWithMetadata.from_event(persistent.result)
if persistent.result is not None
else None,
)
@property
def status(self) -> Status:
"""Get the current status by inspecting the handler state."""
if not self.run_handler.done():
return "running"
# done - check if cancelled first
if self.run_handler.cancelled():
return "cancelled"
# then check for exception
exc = self.run_handler.exception()
if exc is not None:
return "failed"
return "completed"
@property
def error(self) -> str | None:
if not self.run_handler.done():
return None
try:
exc = self.run_handler.exception()
except asyncio.CancelledError:
return None
return str(exc) if exc is not None else None
@property
def result(self) -> StopEvent | None:
if not self.run_handler.done():
return None
try:
return self.run_handler.get_stop_event()
except asyncio.CancelledError:
return None
except Exception:
return None
def _start_idle_release_timer(self) -> None:
"""Start a timer to release this handler after the idle timeout."""
if self._idle_release_timeout is None or self._on_idle_release is None:
return
# Cancel any existing timer first
self._cancel_idle_release_timer()
timeout_seconds = self._idle_release_timeout.total_seconds()
async def release_after_timeout() -> None:
try:
await asyncio.sleep(timeout_seconds)
# Only release if still idle and no active stream consumers
if self.idle_since is not None:
if self.consumer_mutex.locked():
# Mutex is locked - reschedule to try again later
self._start_idle_release_timer()
return
if self._on_idle_release is not None:
await self._on_idle_release(self)
except asyncio.CancelledError:
pass # Timer was cancelled, nothing to do
self._idle_release_timer = asyncio.create_task(release_after_timeout())
def _cancel_idle_release_timer(self, skip_checkpoint: bool = False) -> None:
"""Cancel any pending idle release timer."""
if skip_checkpoint:
self._skip_checkpoint = True
if self._idle_release_timer is not None:
self._idle_release_timer.cancel()
self._idle_release_timer = None
def mark_idle(self, idle_since: datetime | None = None) -> None:
self.idle_since = idle_since or datetime.now(timezone.utc)
self._start_idle_release_timer()
def mark_active(self) -> None:
"""Mark this handler as active (not idle).
Call this when an event is being sent to prevent premature release.
"""
if self.idle_since is not None:
self.idle_since = None
self._cancel_idle_release_timer()
def start_streaming(self, on_finish: Callable[[], Awaitable[None]]) -> None:
"""Start streaming events from the handler and managing state."""
self.task = asyncio.create_task(self._stream_events(on_finish=on_finish))
async def _stream_events(self, on_finish: Callable[[], Awaitable[None]]) -> None:
"""Internal method that streams events, updates status, and persists state."""
with instrument_tags({"handler_id": self.handler_id}):
await self.checkpoint()
self._on_finish = on_finish
try:
async for event in self.run_handler.stream_events(expose_internal=True):
# Track idle state transitions and manage release timer
if isinstance(event, WorkflowIdleEvent):
self.mark_idle()
elif isinstance(event, UnhandledEvent):
self.mark_idle()
elif (
isinstance(event, StepStateChanged)
and event.step_state == StepState.RUNNING
):
self.mark_active()
if ( # Watch for a specific internal event that signals the step is complete
isinstance(event, StepStateChanged)
and event.step_state == StepState.NOT_RUNNING
):
state = (
self.run_handler.ctx.to_dict()
if self.run_handler.ctx
else None
)
if state is None:
logger.warning(
f"Context state is None for handler {self.handler_id}. This is not expected."
)
continue
await self.checkpoint()
self.queue.put_nowait(event)
except WorkflowRuntimeError:
# Stream was already consumed - this can happen during handler
# cancellation when run_handler is cancelled before this task.
# This is benign; we'll proceed to cleanup.
pass
# Workflow is completing - cancel any pending release timer
self._cancel_idle_release_timer()
# done when stream events are complete
try:
await self.run_handler
except asyncio.CancelledError:
# Handler was cancelled - status will be automatically detected via handler.cancelled()
logger.info(f"Workflow run {self.handler_id} was cancelled")
# Don't re-raise, just let the task complete
except Exception as e:
logger.error(
f"Workflow run {self.handler_id} failed! {e}", exc_info=True
)
await self.checkpoint()
async def acquire_events_stream(
self, timeout: float = 1
) -> AsyncGenerator[Event, None]:
"""
Acquires the lock to iterate over the events, and returns generator of events.
"""
try:
await asyncio.wait_for(self.consumer_mutex.acquire(), timeout=timeout)
except asyncio.TimeoutError:
raise NoLockAvailable(
f"No lock available to acquire after {timeout}s timeout"
)
return self._iter_events(timeout=timeout)
async def _iter_events(self, timeout: float = 1) -> AsyncGenerator[Event, None]:
"""
Converts the queue to an async generator while the workflow is still running, and there are still events.
For better or worse, multiple consumers will compete for events
"""
queue_get_task: asyncio.Task[Event] | None = None
try:
while not self.queue.empty() or (
self.task is not None and not self.task.done()
):
available_events = []
while not self.queue.empty():
available_events.append(self.queue.get_nowait())
for event in available_events:
yield event
queue_get_task = asyncio.create_task(self.queue.get())
task_waitable = self.task
done, pending = await asyncio.wait(
{queue_get_task, task_waitable}
if task_waitable is not None
else {queue_get_task},
return_when=asyncio.FIRST_COMPLETED,
)
if queue_get_task in done:
yield await queue_get_task
queue_get_task = None
else: # otherwise task completed, so nothing else will be published to the queue
queue_get_task.cancel()
queue_get_task = None
break
finally:
# Cancel any pending queue.get() task to prevent orphaned tasks from
# consuming events after the consumer disconnects.
if queue_get_task is not None:
if not queue_get_task.done():
queue_get_task.cancel()
try:
await queue_get_task
except asyncio.CancelledError:
pass
if self._on_finish is not None and self.run_handler.done():
# clean up the resources if the stream has been consumed
await self._on_finish()
self.consumer_mutex.release()
async def cancel_handlers_and_tasks(self) -> None:
"""Cancel the handler and release it from the store."""
if not self.run_handler.done():
try:
self.run_handler.cancel()
except Exception:
pass
try:
await self.run_handler.cancel_run()
except Exception:
pass
if self.task and not self.task.done():
self.task.cancel()
try:
await self.task
except asyncio.CancelledError:
pass
class NoLockAvailable(Exception):
"""Raised when no lock is available to acquire after a timeout"""
pass
@dataclass
class _NamedWorkflow:
name: str
workflow: Workflow
@@ -2,6 +2,8 @@
# Copyright (c) 2026 LlamaIndex Inc.
"""Keyed lock utility for per-key mutual exclusion with automatic cleanup."""
from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager
from typing import AsyncIterator
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
@@ -0,0 +1,242 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""IdleReleaseDecorator and supporting adapters.
Wraps a PersistenceDecorator to add idle detection, memory release, and
reload-on-demand for idle workflow handlers.
"""
from __future__ import annotations
import asyncio
import logging
from collections.abc import Coroutine
from datetime import datetime, timezone
from typing import Any
from typing_extensions import override
from workflows.context.serializers import BaseSerializer
from workflows.events import (
Event,
StartEvent,
WorkflowIdleEvent,
)
from workflows.runtime.types.internal_state import BrokerState
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
V2RuntimeCompatibilityShim,
)
from workflows.runtime.types.ticks import WorkflowTick
from workflows.workflow import Workflow
from .._keyed_lock import KeyedLock
from .._store.abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
)
from .persistence_runtime import PersistenceDecorator
from .runtime_decorators import (
BaseExternalRunAdapterDecorator,
BaseInternalRunAdapterDecorator,
BaseRuntimeDecorator,
)
logger = logging.getLogger(__name__)
class _IdleReleaseInternalRunAdapter(BaseInternalRunAdapterDecorator):
"""Internal adapter that detects idle events and schedules release."""
def __init__(
self,
decorated: InternalRunAdapter,
runtime: IdleReleaseDecorator,
store: AbstractWorkflowStore,
) -> None:
super().__init__(decorated)
self._runtime = runtime
self._store = store
@override
async def write_to_event_stream(self, event: Event) -> None:
if isinstance(event, WorkflowIdleEvent):
idle_since = datetime.now(timezone.utc)
await self._store.update_handler_status(
self.run_id, status="running", idle_since=idle_since
)
await super().write_to_event_stream(event)
if isinstance(event, WorkflowIdleEvent):
self._runtime._spawn_task(self._runtime._deferred_release(self.run_id))
class IdleReleaseExternalRunAdapter(BaseExternalRunAdapterDecorator):
"""Proxy adapter that adds reload-on-demand for idle-released handlers.
The inner adapter is resolved lazily via a property because
``get_external_adapter`` is sync but reload is async — the inner run
may not exist yet when this adapter is constructed.
"""
def __init__(self, runtime: IdleReleaseDecorator, run_id: str) -> None:
# Intentionally skip super().__init__ — _decorated is a lazy property.
self._runtime = runtime
self._run_id = run_id
@property # type: ignore[override]
def _decorated(self) -> ExternalRunAdapter:
return self._runtime._decorated.get_external_adapter(self._run_id)
@_decorated.setter
def _decorated(self, value: ExternalRunAdapter) -> None:
pass
@property
def run_id(self) -> str:
return self._run_id
@override
async def send_event(self, tick: WorkflowTick) -> None:
async with self._runtime._reload_lock(self.run_id):
if self.run_id not in self._runtime._active_run_ids:
await self._runtime._ensure_active_run_locked(self.run_id)
else:
await self._runtime._store.update_handler_status(
self.run_id, idle_since=None
)
await self._decorated.send_event(tick)
class IdleReleaseDecorator(BaseRuntimeDecorator):
"""Runtime decorator for idle detection, memory release, and reload-on-demand.
Must wrap a PersistenceDecorator (or compatible runtime) to access
context_from_ticks for reloading released handlers.
"""
def __init__(
self,
decorated: PersistenceDecorator,
store: AbstractWorkflowStore,
idle_timeout: float = 60.0,
) -> None:
super().__init__(decorated)
self._store = store
self._persistence: PersistenceDecorator = decorated
self._reload_lock = KeyedLock()
self._active_run_ids: set[str] = set()
self._background_tasks: set[asyncio.Task[None]] = set()
self.stop_task: asyncio.Task[None] | None = None
self._idle_timeout = idle_timeout
def _spawn_task(self, coro: Coroutine[Any, Any, None]) -> asyncio.Task[None]:
task = asyncio.create_task(coro)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
return task
@override
def run_workflow(
self,
run_id: str,
workflow: Workflow,
init_state: BrokerState,
start_event: StartEvent | None = None,
serialized_state: dict[str, Any] | None = None,
serializer: BaseSerializer | None = None,
) -> ExternalRunAdapter:
self._active_run_ids.add(run_id)
return super().run_workflow(
run_id,
workflow,
init_state,
start_event=start_event,
serialized_state=serialized_state,
serializer=serializer,
)
@override
def get_internal_adapter(self, workflow: Workflow) -> InternalRunAdapter:
inner_adapter = self._decorated.get_internal_adapter(workflow)
return _IdleReleaseInternalRunAdapter(inner_adapter, self, self._store)
@override
def get_external_adapter(self, run_id: str) -> ExternalRunAdapter:
return IdleReleaseExternalRunAdapter(self, run_id)
async def _deferred_release(self, run_id: str) -> None:
"""Wait for idle_timeout then release the handler if still idle."""
await asyncio.sleep(self._idle_timeout)
await self._release_idle_handler(run_id)
async def _release_idle_handler(self, run_id: str) -> None:
"""Release an idle handler from memory."""
async with self._reload_lock(run_id):
handlers = await self._store.query(HandlerQuery(run_id_in=[run_id]))
if len(handlers) != 1 or handlers[0].idle_since is None:
return
elapsed = (
datetime.now(timezone.utc) - handlers[0].idle_since
).total_seconds()
if elapsed < self._idle_timeout:
return
if run_id not in self._active_run_ids:
return
self._active_run_ids.discard(run_id)
self._abort_inner_run(run_id)
logger.info(f"Released idle handler [run_id={run_id}] from memory")
def _abort_inner_run(self, run_id: str) -> None:
"""Cancel the inner runtime's control loop task for a run."""
try:
inner_adapter = self._decorated.get_external_adapter(run_id)
except Exception:
return
if isinstance(inner_adapter, V2RuntimeCompatibilityShim):
inner_adapter.abort()
else:
raise ValueError(f"Inner adapter {inner_adapter} does not support abort")
async def _ensure_active_run(self, run_id: str) -> None:
if run_id in self._active_run_ids:
return
async with self._reload_lock(run_id):
await self._ensure_active_run_locked(run_id)
async def _ensure_active_run_locked(self, run_id: str) -> None:
if run_id in self._active_run_ids:
return
handlers = await self._store.query(HandlerQuery(run_id_in=[run_id]))
if len(handlers) != 1:
raise ValueError(
f"Expected 1 handler for run {run_id}, got {len(handlers)}"
)
handler = handlers[0]
workflow = self._persistence.get_tracked_workflow(handler.workflow_name)
if workflow is None:
raise ValueError(f"Workflow {handler.workflow_name} not found")
context = await self._persistence.context_from_ticks(workflow, run_id)
workflow.run(ctx=context, run_id=run_id)
self._active_run_ids.add(run_id)
await self._store.update_handler_status(run_id, idle_since=None)
logger.info(
f"Reloaded workflow [handler_id={handler.handler_id}, run_id={run_id}] from persistence"
)
@override
def destroy(self) -> None:
super().destroy()
if self.stop_task is not None:
try:
self.stop_task.cancel()
except Exception:
pass
self.stop_task = self._spawn_task(self._on_server_stop())
async def _on_server_stop(self) -> None:
"""Cancel all active runs."""
run_ids = list(self._active_run_ids)
logger.info(f"Shutting down. Cancelling {len(run_ids)} handlers.")
for run_id in run_ids:
self._abort_inner_run(run_id)
self._active_run_ids.clear()
@@ -0,0 +1,267 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""PersistenceDecorator and _PersistenceInternalRunAdapter.
Wraps a basic runtime to add tick persistence and auto-restart on server
start. Does NOT handle idle detection or reload-on-demand — those live in
IdleReleaseDecorator.
"""
from __future__ import annotations
import asyncio
import logging
import sqlite3
from collections.abc import Coroutine
from datetime import datetime, timezone
from typing import Any
from typing_extensions import override
from workflows import Context
from workflows.context.context_types import SerializedContext
from workflows.context.serializers import BaseSerializer, JsonSerializer
from workflows.events import StartEvent
from workflows.runtime.control_loop import rebuild_state_from_ticks
from workflows.runtime.types.internal_state import BrokerState
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
Runtime,
)
from workflows.runtime.types.ticks import WorkflowTick, WorkflowTickAdapter
from workflows.workflow import Workflow
from .._store.abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
PersistentHandler,
as_legacy_context_store,
)
from .._store.sqlite.sqlite_state_store import SqliteStateStore
from .runtime_decorators import (
BaseInternalRunAdapterDecorator,
BaseRuntimeDecorator,
)
logger = logging.getLogger(__name__)
class _PersistenceInternalRunAdapter(BaseInternalRunAdapterDecorator):
"""Internal adapter that persists ticks to the workflow store."""
def __init__(
self,
decorated: InternalRunAdapter,
store: AbstractWorkflowStore,
) -> None:
super().__init__(decorated)
self._store = store
@override
async def on_tick(self, tick: WorkflowTick) -> None:
await super().on_tick(tick)
tick_data = WorkflowTickAdapter.dump_python(tick, mode="json")
try:
await self._store.append_tick(self.run_id, tick_data)
except Exception:
logger.exception(
"Failed to persist tick for run %s",
self.run_id,
)
class PersistenceDecorator(BaseRuntimeDecorator):
"""Runtime decorator for tick persistence and auto-restart.
Manages workflow tracking, tick persistence via internal adapter,
and resuming previously running workflows on server start.
"""
def __init__(
self,
decorated: Runtime,
store: AbstractWorkflowStore,
) -> None:
super().__init__(decorated)
self._store = store
self._workflows_by_name: dict[str, Workflow] = {}
self._active_run_ids: set[str] = set()
self._background_tasks: set[asyncio.Task[None]] = set()
self.resume_task: asyncio.Task[None] | None = None
def _spawn_task(self, coro: Coroutine[Any, Any, None]) -> asyncio.Task[None]:
task = asyncio.create_task(coro)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
return task
@override
def run_workflow(
self,
run_id: str,
workflow: Workflow,
init_state: BrokerState,
start_event: StartEvent | None = None,
serialized_state: dict[str, Any] | None = None,
serializer: BaseSerializer | None = None,
) -> ExternalRunAdapter:
self._active_run_ids.add(run_id)
return super().run_workflow(
run_id,
workflow,
init_state,
start_event=start_event,
serialized_state=serialized_state,
serializer=serializer,
)
@override
def get_internal_adapter(self, workflow: Workflow) -> InternalRunAdapter:
inner_adapter = self._decorated.get_internal_adapter(workflow)
return _PersistenceInternalRunAdapter(inner_adapter, self._store)
@override
def track_workflow(self, workflow: Workflow) -> None:
self._workflows_by_name[workflow.workflow_name] = workflow
super().track_workflow(workflow)
@override
def untrack_workflow(self, workflow: Workflow) -> None:
self._workflows_by_name.pop(workflow.workflow_name, None)
super().untrack_workflow(workflow)
def get_tracked_workflow(self, name: str) -> Workflow | None:
"""Look up a tracked workflow by name (used by IdleReleaseDecorator)."""
return self._workflows_by_name.get(name)
@override
def launch(self) -> None:
super().launch()
self.resume_task = self._spawn_task(
self._on_server_start(self._workflows_by_name)
)
async def _on_server_start(self, registered_workflows: dict[str, Workflow]) -> None:
"""Resume previously running (non-idle) workflows from persistence."""
handlers = await self._store.query(
HandlerQuery(
status_in=["running"],
workflow_name_in=list(registered_workflows.keys()),
is_idle=False,
)
)
for persistent in handlers:
workflow = registered_workflows.get(persistent.workflow_name)
if workflow is None:
continue
if persistent.run_id is None:
logger.error(f"Run ID is required for handler {persistent.handler_id}")
continue
run_id = persistent.run_id
if run_id in self._active_run_ids:
continue
try:
context = await self.context_from_ticks(workflow, run_id)
workflow.run(ctx=context, run_id=run_id)
except Exception as e:
logger.error(
f"Failed to resume handler {persistent.handler_id} for workflow {persistent.workflow_name}: {e}"
)
try:
now = datetime.now(timezone.utc)
await self._store.update(
PersistentHandler(
handler_id=persistent.handler_id,
workflow_name=persistent.workflow_name,
status="failed",
run_id=persistent.run_id,
error=str(e),
result=None,
started_at=persistent.started_at,
updated_at=now,
completed_at=now,
)
)
except Exception:
pass
continue
@override
def destroy(self) -> None:
super().destroy()
if self.resume_task is not None:
try:
self.resume_task.cancel()
except Exception:
pass
async def context_from_ticks(
self, workflow: Workflow, run_id: str
) -> Context | None:
"""Rebuild a Context from persisted ticks (and legacy ctx if available)."""
stored_ticks = await self._store.get_ticks(run_id)
serializer = JsonSerializer()
legacy_ctx = self._get_legacy_ctx(run_id)
if not stored_ticks and not legacy_ctx:
return None
if legacy_ctx:
self._seed_legacy_state(run_id, legacy_ctx)
parsed = SerializedContext.from_dict_auto(legacy_ctx)
init_state = BrokerState.from_serialized(parsed, workflow, serializer)
else:
init_state = BrokerState.from_workflow(workflow)
if stored_ticks:
ticks = [
WorkflowTickAdapter.validate_python(st.tick_data) for st in stored_ticks
]
init_state = rebuild_state_from_ticks(init_state, ticks)
serialized = init_state.to_serialized(serializer)
return Context.from_dict(
workflow=workflow, data=serialized.model_dump(), serializer=serializer
)
def _get_legacy_ctx(self, run_id: str) -> dict[str, Any] | None:
legacy_store = as_legacy_context_store(self._store)
if legacy_store is None:
return None
try:
return legacy_store.get_legacy_ctx(run_id)
except Exception:
logger.warning(
"Failed to read legacy ctx for run %s", run_id, exc_info=True
)
return None
def _seed_legacy_state(self, run_id: str, legacy_ctx: dict[str, Any]) -> None:
try:
parsed = SerializedContext.from_dict_auto(legacy_ctx)
except Exception:
logger.warning(
"Failed to parse legacy ctx for state migration, run %s", run_id
)
return
state_data = parsed.state
if not state_data:
return
state_store = self._store.create_state_store(run_id)
if not isinstance(state_store, SqliteStateStore):
return
conn = sqlite3.connect(state_store._db_path)
try:
row = conn.execute(
"SELECT 1 FROM state WHERE run_id = ?", (run_id,)
).fetchone()
if row is not None:
return
finally:
conn.close()
state_store._write_in_memory_state(state_data)
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""
Base decorator classes for Runtime, InternalRunAdapter, and ExternalRunAdapter.
These provide a simple forwarding pattern: accept an inner instance, delegate
every method to it. Subclasses override only the methods they need to customise.
"""
from __future__ import annotations
import asyncio
import logging
from contextlib import contextmanager
from typing import Any, AsyncGenerator, Generator
from workflows.context.serializers import BaseSerializer
from workflows.context.state_store import StateStore
from workflows.events import (
Event,
StartEvent,
StopEvent,
)
from workflows.runtime.types.internal_state import BrokerState
from workflows.runtime.types.named_task import NamedTask
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
RegisteredWorkflow,
Runtime,
WaitResult,
)
from workflows.runtime.types.ticks import WorkflowTick
from workflows.workflow import Workflow
logger = logging.getLogger(__name__)
class BaseRuntimeDecorator(Runtime):
"""Decorator base for :class:`Runtime`.
Wraps an inner runtime and forwards every call to it. Subclasses can
override individual methods to add behaviour (logging, metrics, auth,
etc.) without re-implementing the full interface.
"""
def __init__(self, decorated: Runtime) -> None:
super().__init__()
self._decorated = decorated
def register(self, workflow: Workflow) -> RegisteredWorkflow:
return self._decorated.register(workflow)
def run_workflow(
self,
run_id: str,
workflow: Workflow,
init_state: BrokerState,
start_event: StartEvent | None = None,
serialized_state: dict[str, Any] | None = None,
serializer: BaseSerializer | None = None,
) -> ExternalRunAdapter:
return self._decorated.run_workflow(
run_id,
workflow,
init_state,
start_event=start_event,
serialized_state=serialized_state,
serializer=serializer,
)
def get_internal_adapter(self, workflow: Workflow) -> InternalRunAdapter:
return self._decorated.get_internal_adapter(workflow)
def get_external_adapter(self, run_id: str) -> ExternalRunAdapter:
return self._decorated.get_external_adapter(run_id)
def launch(self) -> None:
super().launch()
self._decorated.launch()
def destroy(self) -> None:
self._decorated.destroy()
def track_workflow(self, workflow: Workflow) -> None:
self._pending.add(workflow)
self._decorated.track_workflow(workflow)
def untrack_workflow(self, workflow: Workflow) -> None:
self._pending.discard(workflow)
self._decorated.untrack_workflow(workflow)
def get_registered(self, workflow: Workflow) -> RegisteredWorkflow | None:
return self._decorated.get_registered(workflow)
@contextmanager
def registering(self) -> Generator[Runtime, None, None]:
with self._decorated.registering() as rt:
yield rt
class BaseInternalRunAdapterDecorator(InternalRunAdapter):
"""Decorator base for :class:`InternalRunAdapter`.
Wraps an inner adapter and forwards every call to it. Subclasses can
override individual methods to intercept or augment behaviour.
"""
def __init__(self, decorated: InternalRunAdapter) -> None:
self._decorated = decorated
@property
def run_id(self) -> str:
return self._decorated.run_id
async def write_to_event_stream(self, event: Event) -> None:
await self._decorated.write_to_event_stream(event)
async def get_now(self) -> float:
return await self._decorated.get_now()
async def send_event(self, tick: WorkflowTick) -> None:
await self._decorated.send_event(tick)
async def wait_receive(
self,
timeout_seconds: float | None = None,
) -> WaitResult:
return await self._decorated.wait_receive(timeout_seconds)
async def close(self) -> None:
await self._decorated.close()
def get_state_store(self) -> StateStore[Any] | None:
return self._decorated.get_state_store()
async def finalize_step(self) -> None:
await self._decorated.finalize_step()
async def on_tick(self, tick: WorkflowTick) -> None:
await self._decorated.on_tick(tick)
async def wait_for_next_task(
self,
task_set: list[NamedTask],
timeout: float | None = None,
) -> asyncio.Task[Any] | None:
return await self._decorated.wait_for_next_task(task_set, timeout)
class BaseExternalRunAdapterDecorator(ExternalRunAdapter):
"""Decorator base for :class:`ExternalRunAdapter`.
Wraps an inner adapter and forwards every call to it. Subclasses can
override individual methods to intercept or augment behaviour.
"""
def __init__(self, decorated: ExternalRunAdapter) -> None:
self._decorated = decorated
@property
def run_id(self) -> str:
return self._decorated.run_id
async def send_event(self, tick: WorkflowTick) -> None:
await self._decorated.send_event(tick)
def stream_published_events(self) -> AsyncGenerator[Event, None]:
return self._decorated.stream_published_events()
async def close(self) -> None:
await self._decorated.close()
async def get_result(self) -> StopEvent:
return await self._decorated.get_result()
async def cancel(self) -> None:
await self._decorated.cancel()
def get_state_store(self) -> StateStore[Any] | None:
return self._decorated.get_state_store()
@@ -0,0 +1,270 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""
Server runtime decorator: the main required runtime decorator for workflows
served by the WorkflowServer. Handles event recording, handler persistence,
and status updates.
"""
from __future__ import annotations
import asyncio
import logging
from datetime import datetime, timezone
from typing import Any, Awaitable, Callable
from llama_agents.client.protocol.serializable_events import (
EventEnvelopeWithMetadata,
)
from llama_agents.server._runtime.runtime_decorators import (
BaseInternalRunAdapterDecorator,
)
from typing_extensions import override
from workflows.context.serializers import BaseSerializer
from workflows.context.state_store import (
InMemoryStateStore,
StateStore,
infer_state_type,
)
from workflows.events import (
Event,
StartEvent,
StopEvent,
WorkflowCancelledEvent,
WorkflowFailedEvent,
WorkflowTimedOutEvent,
)
from workflows.handler import WorkflowHandler
from workflows.runtime.types.internal_state import BrokerState
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
Runtime,
)
from workflows.workflow import Workflow
from .._store.abstract_workflow_store import (
AbstractWorkflowStore,
PersistentHandler,
Status,
)
from .runtime_decorators import BaseRuntimeDecorator
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# _ServerInternalRunAdapter
# ---------------------------------------------------------------------------
class _ServerInternalRunAdapter(BaseInternalRunAdapterDecorator):
"""Internal adapter that records every emitted event to the workflow store.
Handles event recording and terminal-event status updates.
"""
def __init__(
self,
decorated: InternalRunAdapter,
runtime: ServerRuntimeDecorator,
*,
state_type: type[Any] | None = None,
) -> None:
super().__init__(decorated)
self._runtime = runtime
self._store = runtime._store
self._state_type = state_type
self._state_store: StateStore[Any] | None = None
@override
def get_state_store(self) -> StateStore[Any]:
if self._state_store is not None:
return self._state_store
store = self._store.create_state_store(self.run_id, self._state_type)
# Seed with initial context state if provided at run start
initial = self._runtime._initial_state.pop(self.run_id, None)
if initial is not None and isinstance(store, InMemoryStateStore):
store._state = initial
self._state_store = store
return store
@override
async def write_to_event_stream(self, event: Event) -> None:
"""
Monitors for writes to the event stream that indicate a workflow has terminated.
"""
if isinstance(event, WorkflowFailedEvent):
await self._runtime._handle_status_update(
run_id=self.run_id,
status="failed",
error=event.exception_message,
)
elif isinstance(event, WorkflowTimedOutEvent):
await self._runtime._handle_status_update(
run_id=self.run_id,
status="failed",
error=f"Workflow timed out after {event.timeout}s",
)
elif isinstance(event, WorkflowCancelledEvent):
await self._runtime._handle_status_update(
run_id=self.run_id, status="cancelled"
)
elif isinstance(event, StopEvent):
await self._runtime._handle_status_update(
run_id=self.run_id,
status="completed",
result=event,
)
envelope = EventEnvelopeWithMetadata.from_event(event)
await self._store.append_event(self.run_id, envelope)
# Forward to inner adapter (e.g. _DurableInternalRunAdapter for idle detection)
await super().write_to_event_stream(event)
# ---------------------------------------------------------------------------
# ServerRuntimeDecorator -- adapter wrapping, handler persistence,
# status updates, and workflow registry
# ---------------------------------------------------------------------------
class ServerRuntimeDecorator(BaseRuntimeDecorator):
"""
Runtime decorator that wraps the main runtime to also record events to a configured
workflow store, for integration with the WorkflowService for querying
workflow run state.
"""
def __init__(
self,
decorated: Runtime,
store: AbstractWorkflowStore,
*,
persistence_backoff: list[float] | None = None,
) -> None:
super().__init__(decorated)
self._store: AbstractWorkflowStore = store
self._registered_workflows: dict[str, Workflow] = {}
self._initial_state: dict[str, Any] = {}
self._persistence_backoff = (
list(persistence_backoff) if persistence_backoff is not None else [0.5, 3]
)
async def _retry_store_write(self, coro_fn: Callable[[], Awaitable[None]]) -> None:
"""Wrap a store write with retry/backoff."""
backoffs = list(self._persistence_backoff)
while True:
try:
await coro_fn()
return
except Exception as e:
backoff = backoffs.pop(0) if backoffs else None
if backoff is None:
logger.error(
"Store write failed after final attempt",
exc_info=True,
)
raise
logger.error(f"Store write failed, retrying in {backoff}s: {e}")
await asyncio.sleep(backoff)
# ------------------------------------------------------------------
# Workflow registration
# ------------------------------------------------------------------
@override
def track_workflow(self, workflow: Workflow) -> None:
# Keep a strong reference — the base WorkflowSet uses weak refs,
# so without this the workflow can be GC'd before launch().
self._registered_workflows[workflow.workflow_name] = workflow
super().track_workflow(workflow)
@override
def untrack_workflow(self, workflow: Workflow) -> None:
self._registered_workflows.pop(workflow.workflow_name, None)
super().untrack_workflow(workflow)
def get_workflow(self, name: str) -> Workflow | None:
return self._registered_workflows.get(name)
def get_workflow_names(self) -> list[str]:
return list(self._registered_workflows.keys())
# ------------------------------------------------------------------
# Adapter wiring
# ------------------------------------------------------------------
async def _handle_status_update(
self,
run_id: str,
status: Status,
result: StopEvent | None = None,
error: str | None = None,
) -> None:
"""Callback for adapter terminal-event status updates."""
await self._retry_store_write(
lambda: self._store.update_handler_status(
run_id, status=status, result=result, error=error
)
)
@override
def run_workflow(
self,
run_id: str,
workflow: Workflow,
init_state: BrokerState,
start_event: StartEvent | None = None,
serialized_state: dict[str, Any] | None = None,
serializer: BaseSerializer | None = None,
) -> ExternalRunAdapter:
if serialized_state and serializer:
try:
seed_store = InMemoryStateStore.from_dict(serialized_state, serializer)
self._initial_state[run_id] = seed_store._state
except Exception:
pass
return super().run_workflow(
run_id,
workflow,
init_state,
start_event=start_event,
serialized_state=serialized_state,
serializer=serializer,
)
def get_internal_adapter(self, workflow: Workflow) -> InternalRunAdapter:
"""Wraps the inner runtime's adapter in _ServerInternalRunAdapter."""
inner_adapter = self._decorated.get_internal_adapter(workflow)
state_type = infer_state_type(workflow)
return _ServerInternalRunAdapter(inner_adapter, self, state_type=state_type)
# ------------------------------------------------------------------
# Handler persistence
# ------------------------------------------------------------------
async def run_workflow_handler(
self,
handler_id: str,
workflow_name: str,
handler: WorkflowHandler,
) -> WorkflowHandler:
"""Persist initial handler record to store, then notify decorator chain."""
started_at = datetime.now(timezone.utc)
await self._retry_store_write(
lambda: self._store.update(
PersistentHandler(
handler_id=handler_id,
workflow_name=workflow_name,
status="running",
run_id=handler.run_id,
started_at=started_at,
updated_at=started_at,
)
)
)
return handler
@@ -1,280 +1,272 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""
Application-level orchestration layer for workflow handler lifecycle.
_WorkflowService is a plain class (not a Runtime subclass) that provides
the public interface consumed by _api.py. It delegates to the decorated
runtime for persistence and adapter wiring, and to the store for queries.
"""
from __future__ import annotations
import asyncio
import logging
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from typing import Literal
from llama_agents.client.protocol import HandlerData
from llama_agents.client.protocol.serializable_events import (
EventEnvelopeWithMetadata,
)
from llama_agents.server._runtime.server_runtime import ServerRuntimeDecorator
from llama_index_instrumentation.dispatcher import instrument_tags
from workflows import Context, Workflow
from workflows import Context
from workflows.events import Event, StartEvent
from workflows.handler import WorkflowHandler
from workflows.utils import _nanoid as nanoid
from workflows.workflow import Workflow
from ._handler import _NamedWorkflow, _WorkflowHandler
from ._keyed_lock import KeyedLock
from ._store.abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
PersistentHandler,
is_terminal_status,
)
from ._store.memory_workflow_store import MemoryWorkflowStore
logger = logging.getLogger()
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Exceptions
# ---------------------------------------------------------------------------
class HandlerNotFoundError(Exception):
pass
class HandlerCompletedError(Exception):
pass
class EventSendError(Exception):
pass
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def handler_data_from_persistent(persistent: PersistentHandler) -> HandlerData:
return HandlerData(
handler_id=persistent.handler_id,
workflow_name=persistent.workflow_name,
run_id=persistent.run_id,
status=persistent.status,
started_at=persistent.started_at.isoformat()
if persistent.started_at is not None
else datetime.now(timezone.utc).isoformat(),
updated_at=persistent.updated_at.isoformat()
if persistent.updated_at is not None
else None,
completed_at=persistent.completed_at.isoformat()
if persistent.completed_at is not None
else None,
error=persistent.error,
result=EventEnvelopeWithMetadata.from_event(persistent.result)
if persistent.result is not None
else None,
)
# ---------------------------------------------------------------------------
# _WorkflowService
# ---------------------------------------------------------------------------
class _WorkflowService:
"""Handler lifecycle, persistence, and event registry management.
"""Application-level service facade for workflow handler lifecycle.
This layer owns the _handlers dict, _workflows dict, _reload_lock,
and all lifecycle methods. It has no knowledge of HTTP.
This is NOT a Runtime. It holds references to the decorated runtime
(for running workflows and getting adapters) and the store (for queries).
"""
def __init__(
self,
*,
workflow_store: AbstractWorkflowStore | None = None,
persistence_backoff: list[float] = [0.5, 3],
idle_release_timeout: timedelta | None = timedelta(seconds=10),
runtime: ServerRuntimeDecorator,
store: AbstractWorkflowStore,
) -> None:
self._workflows: dict[str, Workflow] = {}
self._additional_events: dict[str, list[type[Event]] | None] = {}
self._handlers: dict[str, _WorkflowHandler] = {}
self._workflow_store = (
workflow_store if workflow_store is not None else MemoryWorkflowStore()
)
self._persistence_backoff = list(persistence_backoff)
self._idle_release_timeout = idle_release_timeout
self._reload_lock = KeyedLock()
self._runtime: ServerRuntimeDecorator = runtime
self._store = store
def add_workflow(
# ------------------------------------------------------------------
# Workflow registration
# ------------------------------------------------------------------
def get_workflow(self, name: str) -> Workflow | None:
return self._runtime.get_workflow(name)
def get_workflow_names(self) -> list[str]:
return self._runtime.get_workflow_names()
# ------------------------------------------------------------------
# Store access
# ------------------------------------------------------------------
@property
def store(self) -> AbstractWorkflowStore:
return self._store
async def query_handlers(self, query: HandlerQuery) -> list[PersistentHandler]:
return await self._store.query(query)
# ------------------------------------------------------------------
# Handler lifecycle
# ------------------------------------------------------------------
async def load_handler(self, handler_id: str) -> HandlerData | None:
found = await self._store.query(HandlerQuery(handler_id_in=[handler_id]))
if not found:
return None
return handler_data_from_persistent(found[0])
async def resolve_handler(self, handler_id: str) -> HandlerData:
handler_data = await self.load_handler(handler_id)
if handler_data is None:
raise HandlerNotFoundError()
if is_terminal_status(handler_data.status):
raise HandlerCompletedError()
return handler_data
async def send_event(
self,
name: str,
workflow: Workflow,
additional_events: list[type[Event]] | None = None,
handler_id: str,
event: Event,
step: str | None = None,
) -> None:
self._workflows[name] = workflow
if additional_events is not None:
self._additional_events[name] = additional_events
"""Send a parsed event to a running handler."""
handler_data = await self.resolve_handler(handler_id)
async def start(self) -> None:
"""Resume previously running (non-idle) workflows from persistence."""
handlers = await self._workflow_store.query(
HandlerQuery(
status_in=["running"],
workflow_name_in=list(self._workflows.keys()),
is_idle=False,
workflow = self._runtime.get_workflow(handler_data.workflow_name)
if workflow is None:
raise EventSendError(
f"Workflow {handler_data.workflow_name} not registered"
)
)
for persistent in handlers:
workflow = self._workflows[persistent.workflow_name]
try:
await self.start_workflow(
workflow=_NamedWorkflow(
name=persistent.workflow_name, workflow=workflow
),
handler_id=persistent.handler_id,
context=Context.from_dict(workflow=workflow, data=persistent.ctx),
)
except Exception as e:
logger.error(
f"Failed to resume handler {persistent.handler_id} for workflow {persistent.workflow_name}: {e}"
)
try:
now = datetime.now(timezone.utc)
await self._workflow_store.update(
PersistentHandler(
handler_id=persistent.handler_id,
workflow_name=persistent.workflow_name,
status="failed",
run_id=persistent.run_id,
error=str(e),
result=None,
started_at=persistent.started_at,
updated_at=now,
completed_at=now,
ctx=persistent.ctx,
)
)
except Exception:
pass
continue
if handler_data.run_id is None:
raise EventSendError(f"Handler {handler_id} has no run ID")
async def stop(self) -> None:
logger.info(
f"Shutting down Workflow server. Cancelling {len(self._handlers)} handlers."
)
await asyncio.gather(
*[self.close_handler(handler) for handler in list(self._handlers.values())]
)
self._handlers.clear()
try:
handler = WorkflowHandler(
workflow, self._runtime.get_external_adapter(handler_data.run_id)
)
await handler.send_event(event, step=step)
except Exception as e:
raise EventSendError(f"Failed to send event: {e}") from e
async def cancel_handler(
self, handler_id: str, purge: bool = False
) -> Literal["cancelled", "deleted"] | None:
found = await self._store.query(HandlerQuery(handler_id_in=[handler_id]))
if not found:
return None
persisted = handler_data_from_persistent(found[0])
if not purge and (
persisted.run_id is None or is_terminal_status(persisted.status)
):
return None
is_terminal = is_terminal_status(persisted.status)
if not is_terminal and persisted.run_id is not None:
handler = self._workflow_run_handler(
persisted.workflow_name, persisted.run_id
)
await self._cancel_run(handler)
if purge:
n_deleted = await self._store.delete(
HandlerQuery(handler_id_in=[handler_id])
)
if n_deleted == 0:
return None
return "deleted" if purge else "cancelled"
async def start_workflow(
self,
workflow: _NamedWorkflow,
workflow: Workflow,
handler_id: str,
start_event: StartEvent | None = None,
context: Context | None = None,
idle_since: datetime | None = None,
) -> _WorkflowHandler:
"""Start a workflow and return a wrapper for the handler."""
) -> HandlerData:
with instrument_tags({"handler_id": handler_id}):
handler = workflow.workflow.run(
handler = workflow.run(
ctx=context,
start_event=start_event,
)
wrapper = await self.run_workflow_handler(
handler_id, workflow.name, handler, idle_since=idle_since
await self._runtime.run_workflow_handler(
handler_id, workflow.workflow_name, handler
)
return wrapper
handler_data = await self.load_handler(handler_id)
if handler_data is None:
raise RuntimeError(f"Handler {handler_id} not found after creation")
return handler_data
async def run_workflow_handler(
self,
handler_id: str,
workflow_name: str,
handler: WorkflowHandler,
idle_since: datetime | None = None,
) -> _WorkflowHandler:
"""Create a wrapper for the handler and start streaming events."""
queue: asyncio.Queue[Event] = asyncio.Queue()
started_at = datetime.now(timezone.utc)
async def await_workflow(self, handler: HandlerData) -> HandlerData:
if handler.run_id is None:
raise HandlerNotFoundError("Handler exists, but has no run ID")
run = self._workflow_run_handler(handler.workflow_name, handler.run_id)
wrapper = _WorkflowHandler(
run_handler=handler,
queue=queue,
task=None,
consumer_mutex=asyncio.Lock(),
handler_id=handler_id,
workflow_name=workflow_name,
started_at=started_at,
updated_at=started_at,
completed_at=None,
_workflow_store=self._workflow_store,
_persistence_backoff=self._persistence_backoff,
_idle_release_timeout=self._idle_release_timeout,
_on_idle_release=self.release_handler,
try:
await run
except Exception:
pass
handler_data = await self.load_handler(handler.handler_id)
if handler_data is None:
raise HandlerNotFoundError()
return handler_data
# ------------------------------------------------------------------
# Start / stop
# ------------------------------------------------------------------
async def start(self) -> None:
"""Launch runtimes and register tracked workflows."""
self._runtime.launch()
async def stop(self) -> None:
"""Stop active runs and destroy the runtime."""
self._runtime.destroy()
# ------------------------------------------------------------------
# Private helpers
# ------------------------------------------------------------------
def _workflow_run_handler(self, workflow_name: str, run_id: str) -> WorkflowHandler:
workflow = self._runtime.get_workflow(workflow_name)
if workflow is None:
raise HandlerNotFoundError(f"Workflow {workflow_name} not registered")
return WorkflowHandler(
workflow=workflow,
external_adapter=workflow._runtime.get_external_adapter(run_id),
)
wrapper.idle_since = idle_since
# Initial checkpoint before registration; fail fast if persistence is unavailable
await wrapper.checkpoint()
# Now register and start streaming
self._handlers[handler_id] = wrapper
async def on_finish() -> None:
self._handlers.pop(handler_id, None)
wrapper.start_streaming(on_finish=on_finish)
return wrapper
async def close_handler(self, handler: _WorkflowHandler) -> None:
"""Close and cleanup a handler."""
await handler.cancel_handlers_and_tasks()
self._handlers.pop(handler.handler_id, None)
async def release_handler(self, wrapper: _WorkflowHandler) -> None:
"""Release an idle handler from memory, keeping it in persistence."""
handler_id = wrapper.handler_id
async with self._reload_lock(handler_id):
current = self._handlers.get(handler_id)
if current is not None and current is not wrapper:
logger.debug(
f"Skipping release checkpoint for {handler_id}: "
"handler was already reloaded"
)
wrapper._cancel_idle_release_timer(skip_checkpoint=True)
await wrapper.cancel_handlers_and_tasks()
return
self._handlers.pop(handler_id, None)
wrapper._cancel_idle_release_timer()
async def _cancel_run(self, run: WorkflowHandler) -> None:
"""Gracefully cancel the workflow run, then kill tasks."""
if not run.done():
try:
await wrapper.checkpoint()
finally:
await wrapper.cancel_handlers_and_tasks()
logger.info(f"Released idle workflow {handler_id} from memory")
async def try_reload_handler(
self, handler_id: str
) -> tuple[_WorkflowHandler | None, PersistentHandler | None]:
"""Attempt to reload a released handler from persistence.
Uses per-handler locking to prevent concurrent reloads from creating
duplicate workflow instances.
Returns (wrapper, persistent_data). The persistent data is returned
so callers can inspect it without re-querying the store.
"""
async with self._reload_lock(handler_id):
if handler_id in self._handlers:
return self._handlers[handler_id], None
found = await self._workflow_store.query(
HandlerQuery(handler_id_in=[handler_id])
)
if not found:
return None, None
handler_data = found[0]
if handler_data.status != "running":
return None, handler_data
workflow = self._workflows.get(handler_data.workflow_name)
if workflow is None:
logger.warning(
f"Cannot reload {handler_id}: workflow {handler_data.workflow_name} not registered"
)
return None, handler_data
await run.cancel_run()
except Exception:
pass
try:
context = Context.from_dict(workflow=workflow, data=handler_data.ctx)
wrapper = await self.start_workflow(
workflow=_NamedWorkflow(
name=handler_data.workflow_name, workflow=workflow
),
handler_id=handler_id,
context=context,
idle_since=handler_data.idle_since,
)
await run
except (asyncio.CancelledError, Exception):
pass
await self._kill_run(run)
if wrapper.idle_since is not None:
wrapper._start_idle_release_timer()
logger.info(f"Reloaded workflow {handler_id} from persistence")
return wrapper, handler_data
except Exception as e:
logger.error(f"Failed to reload handler {handler_id}: {e}")
raise
def event_registry(self, workflow_name: str) -> dict[str, type[Event]]:
items = {e.__name__: e for e in self._workflows[workflow_name].events}
items.update(
{
e.__name__: e
for e in self._additional_events.get(workflow_name, None) or []
}
)
return items
def prepare_run_params(
self,
workflow_name: str,
start_event_data: dict | None,
context_data: dict | None,
handler_id: str | None,
run_kwargs: dict | None,
) -> tuple[str, dict | None]:
"""Prepare and validate run parameters from already-parsed request data.
Returns (handler_id, start_event_data) where handler_id may be generated
if not provided.
"""
if run_kwargs and start_event_data is None:
start_event_data = run_kwargs
handler_id = handler_id or nanoid()
return handler_id, start_event_data
async def _kill_run(self, run: WorkflowHandler) -> None:
"""Force-kill the handler without graceful cancellation."""
if not run.done():
try:
run.cancel()
except Exception:
pass
@@ -1,2 +0,0 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
@@ -1,25 +1,52 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
import logging
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator
from dataclasses import dataclass
from datetime import datetime
from typing import Any, List, Literal
from datetime import datetime, timezone
from enum import Enum
from typing import Any, List, Literal, Protocol, runtime_checkable
from llama_agents.client.protocol.serializable_events import (
EventEnvelopeWithMetadata,
)
from pydantic import (
BaseModel,
field_serializer,
field_validator,
)
from workflows.context import JsonSerializer
from workflows.context.state_store import StateStore
from workflows.events import StopEvent
logger = logging.getLogger(__name__)
Status = Literal["running", "completed", "failed", "cancelled"]
TERMINAL_STATUSES: frozenset[Status] = frozenset(("completed", "failed", "cancelled"))
def is_terminal_status(status: Status) -> bool:
return status in TERMINAL_STATUSES
class _Unset(Enum):
UNSET = "UNSET"
_UNSET = _Unset.UNSET
@dataclass()
class HandlerQuery:
# Matches if any of the handler_ids match
handler_id_in: List[str] | None = None
# Matches if any of the run_ids match
run_id_in: List[str] | None = None
# Matches if any of the workflow_names match
workflow_name_in: List[str] | None = None
# Matches if the status flag matches
@@ -39,7 +66,6 @@ class PersistentHandler(BaseModel):
updated_at: datetime | None = None
completed_at: datetime | None = None
idle_since: datetime | None = None
ctx: dict[str, Any] = {}
@field_validator("result", mode="before")
@classmethod
@@ -65,7 +91,29 @@ class PersistentHandler(BaseModel):
return result
class StoredTick(BaseModel):
run_id: str
sequence: int
timestamp: datetime
tick_data: dict[str, Any]
class StoredEvent(BaseModel):
run_id: str
sequence: int
timestamp: datetime
event: EventEnvelopeWithMetadata
class AbstractWorkflowStore(ABC):
poll_interval: float = 0.1
@abstractmethod
def create_state_store(
self, run_id: str, state_type: type[Any] | None = None
) -> StateStore[Any]:
"""Create a persistent state store for the given run. see e.g. InMemoryStateStore for a reference implementation."""
@abstractmethod
async def query(self, query: HandlerQuery) -> List[PersistentHandler]: ...
@@ -74,3 +122,99 @@ class AbstractWorkflowStore(ABC):
@abstractmethod
async def delete(self, query: HandlerQuery) -> int: ...
@abstractmethod
async def append_event(
self, run_id: str, event: EventEnvelopeWithMetadata
) -> None: ...
@abstractmethod
async def query_events(
self, run_id: str, after_sequence: int | None = None, limit: int | None = None
) -> list[StoredEvent]: ...
@abstractmethod
async def append_tick(self, run_id: str, tick_data: dict[str, Any]) -> None: ...
@abstractmethod
async def get_ticks(self, run_id: str) -> list[StoredTick]: ...
async def update_handler_status(
self,
run_id: str,
*,
status: Status | None = None,
result: StopEvent | None = None,
error: str | None = None,
idle_since: datetime | None | _Unset = _UNSET,
) -> None:
"""Update status and related fields for an existing handler.
Loads the handler by run_id, updates status/timestamps/provided fields,
and writes back. If the handler is not found, logs a warning and returns.
"""
found = await self.query(HandlerQuery(run_id_in=[run_id]))
if not found:
logger.warning("update_handler_status: run %s not found, skipping", run_id)
return
handler = found[0]
now = datetime.now(timezone.utc)
if status is not None:
handler.status = status
handler.updated_at = now
if status in ("completed", "failed", "cancelled"):
handler.completed_at = now
if result is not None:
handler.result = result
if error is not None:
handler.error = error
if not isinstance(idle_since, _Unset):
handler.idle_since = idle_since
await self.update(handler)
@staticmethod
def _is_terminal_event(event: StoredEvent) -> bool:
"""Check if a stored event is terminal (StopEvent or subclass, etc.)."""
types = (event.event.types or []) + [event.event.type]
return StopEvent.__name__ in types
async def subscribe_events(
self, run_id: str, after_sequence: int = -1
) -> AsyncIterator[StoredEvent]:
"""Stream events starting after *after_sequence*, yielding in real time.
The default implementation polls via :meth:`query_events`.
:class:`MemoryWorkflowStore` overrides this with condition-based
notification so there is no polling.
The iterator terminates once a terminal event
(``StopEvent``, ``WorkflowFailedEvent``, ``WorkflowCancelledEvent``)
is yielded.
"""
cursor = after_sequence
while True:
events = await self.query_events(run_id, after_sequence=cursor)
for event in events:
yield event
cursor = event.sequence
if self._is_terminal_event(event):
return
if not events:
await asyncio.sleep(self.poll_interval)
@runtime_checkable
class LegacyContextStore(Protocol):
"""Opt-in protocol for stores that can provide old serialized context data from the ctx column."""
def get_legacy_ctx(self, run_id: str) -> dict[str, Any] | None:
"""Return the old serialized context dict for a run, or None if not available."""
...
def as_legacy_context_store(store: AbstractWorkflowStore) -> LegacyContextStore | None:
"""Return the store as a LegacyContextStore if it supports it, else None."""
if isinstance(store, LegacyContextStore):
return store
return None
@@ -1,9 +1,20 @@
from typing import Dict, List
from __future__ import annotations
import asyncio
import weakref
from collections.abc import AsyncIterator
from datetime import datetime, timezone
from typing import Any, Dict, List
from llama_agents.client.protocol.serializable_events import EventEnvelopeWithMetadata
from workflows.context.state_store import DictState, InMemoryStateStore
from .abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
PersistentHandler,
StoredEvent,
StoredTick,
)
@@ -15,6 +26,12 @@ def _matches_query(handler: PersistentHandler, query: HandlerQuery) -> bool:
if handler.handler_id not in query.handler_id_in:
return False
if query.run_id_in is not None:
if len(query.run_id_in) == 0:
return False
if handler.run_id not in query.run_id_in:
return False
if query.workflow_name_in is not None:
if len(query.workflow_name_in) == 0:
return False
@@ -38,6 +55,21 @@ def _matches_query(handler: PersistentHandler, query: HandlerQuery) -> bool:
class MemoryWorkflowStore(AbstractWorkflowStore):
def __init__(self) -> None:
self.handlers: Dict[str, PersistentHandler] = {}
self.events: Dict[str, List[StoredEvent]] = {}
self.ticks: Dict[str, List[StoredTick]] = {}
self.state_stores: Dict[str, InMemoryStateStore[Any]] = {}
self._conditions: weakref.WeakValueDictionary[str, asyncio.Condition] = (
weakref.WeakValueDictionary()
)
def create_state_store(
self, run_id: str, state_type: type[Any] | None = None
) -> InMemoryStateStore[Any]:
if run_id not in self.state_stores:
self.state_stores[run_id] = InMemoryStateStore(
state_type() if state_type else DictState()
)
return self.state_stores[run_id]
async def query(self, query: HandlerQuery) -> List[PersistentHandler]:
return [
@@ -58,3 +90,102 @@ class MemoryWorkflowStore(AbstractWorkflowStore):
for handler_id in to_delete:
del self.handlers[handler_id]
return len(to_delete)
def _get_or_create_condition(self, run_id: str) -> asyncio.Condition:
"""Get or create a condition for a run_id.
The caller is responsible for holding a strong reference to the
returned Condition for as long as it needs notifications.
"""
cond = self._conditions.get(run_id)
if cond is None:
cond = asyncio.Condition()
self._conditions[run_id] = cond
return cond
async def append_event(self, run_id: str, event: EventEnvelopeWithMetadata) -> None:
if run_id not in self.events:
self.events[run_id] = []
existing = self.events[run_id]
next_seq = (existing[-1].sequence + 1) if existing else 0
stored = StoredEvent(
run_id=run_id,
sequence=next_seq,
timestamp=datetime.now(timezone.utc),
event=event,
)
existing.append(stored)
condition = self._conditions.get(run_id)
if condition is not None:
async with condition:
condition.notify_all()
async def query_events(
self,
run_id: str,
after_sequence: int | None = None,
limit: int | None = None,
) -> List[StoredEvent]:
events = self.events.get(run_id, [])
if after_sequence is not None:
events = [e for e in events if e.sequence > after_sequence]
if limit is not None:
events = events[:limit]
return events
async def append_tick(self, run_id: str, tick_data: dict[str, Any]) -> None:
if run_id not in self.ticks:
self.ticks[run_id] = []
existing = self.ticks[run_id]
next_seq = (existing[-1].sequence + 1) if existing else 0
stored = StoredTick(
run_id=run_id,
sequence=next_seq,
timestamp=datetime.now(timezone.utc),
tick_data=tick_data,
)
existing.append(stored)
async def get_ticks(self, run_id: str) -> list[StoredTick]:
return list(self.ticks.get(run_id, []))
async def subscribe_events(
self, run_id: str, after_sequence: int = -1
) -> AsyncIterator[StoredEvent]:
"""Condition-based subscription — no polling.
Uses list-index cursoring rather than sequence-field cursoring to
handle duplicate sequence numbers (which occur when multiple internal
adapters share the same run_id).
"""
condition = self._get_or_create_condition(run_id)
# Determine starting index: skip events with sequence <= after_sequence
all_events = self.events.get(run_id, [])
if after_sequence >= 0:
cursor = 0
for i, e in enumerate(all_events):
if e.sequence <= after_sequence:
cursor = i + 1
# cursor is now the index of the first event to yield
else:
cursor = 0
while True:
all_events = self.events.get(run_id, [])
batch = all_events[cursor:]
for event in batch:
yield event
cursor += 1
if self._is_terminal_event(event):
return
# Before waiting, check if the run already has a terminal event
# that we've already passed (e.g. cursor is beyond all events but
# the last event was terminal). This prevents hanging when a late
# subscriber joins after the run is fully complete.
if all_events and self._is_terminal_event(all_events[-1]):
return
# No new events — wait for the producer to notify
async with condition:
all_events = self.events.get(run_id, [])
if len(all_events) <= cursor:
await condition.wait()
@@ -0,0 +1,33 @@
PRAGMA user_version = 4;
CREATE TABLE IF NOT EXISTS ticks (
id INTEGER PRIMARY KEY AUTOINCREMENT,
run_id TEXT NOT NULL,
sequence INTEGER NOT NULL,
timestamp TEXT NOT NULL,
tick_data TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_ticks_run_id ON ticks (run_id);
CREATE INDEX IF NOT EXISTS idx_ticks_run_id_sequence ON ticks (run_id, sequence);
CREATE TABLE IF NOT EXISTS state (
run_id TEXT PRIMARY KEY,
state_json TEXT NOT NULL DEFAULT '{}',
state_type TEXT NOT NULL DEFAULT 'DictState',
state_module TEXT NOT NULL DEFAULT 'workflows.context.state_store',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS events (
id INTEGER PRIMARY KEY AUTOINCREMENT,
run_id TEXT NOT NULL,
sequence INTEGER NOT NULL,
timestamp TEXT NOT NULL,
event_json TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_events_run_id_sequence ON events (run_id, sequence);
CREATE INDEX IF NOT EXISTS idx_handlers_run_id ON handlers (run_id);
@@ -0,0 +1,257 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
import functools
import json
import logging
import sqlite3
import uuid
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from typing import Any, AsyncGenerator, Generic, Literal, Type
from pydantic import BaseModel
from typing_extensions import TypeVar
from workflows.context.serializers import BaseSerializer, JsonSerializer
from workflows.context.state_store import (
DictState,
create_cleared_state,
deserialize_dict_state_data,
deserialize_state_from_dict,
get_by_path,
merge_state,
parse_in_memory_state,
serialize_dict_state_data,
set_by_path,
)
logger = logging.getLogger(__name__)
MODEL_T = TypeVar("MODEL_T", bound=BaseModel, default=DictState) # type: ignore[reportGeneralTypeIssues]
class SqliteSerializedState(BaseModel):
"""Serialized state referencing a sqlite database row."""
store_type: Literal["sqlite"] = "sqlite"
run_id: str
def _utc_now() -> datetime:
return datetime.now(timezone.utc)
class SqliteStateStore(Generic[MODEL_T]):
"""Sqlite-backed StateStore implementation.
Every get() reads from the database, every set() writes through.
No in-memory cache — the database is the source of truth.
"""
state_type: Type[MODEL_T]
def __init__(
self,
db_path: str,
run_id: str,
state_type: Type[MODEL_T] | None = None,
serializer: BaseSerializer | None = None,
) -> None:
self._db_path = db_path
self._run_id = run_id
self.state_type = state_type or DictState # type: ignore[assignment]
self._serializer = serializer or JsonSerializer()
@property
def run_id(self) -> str:
return self._run_id
@functools.cached_property
def _lock(self) -> asyncio.Lock:
"""Lazy lock initialization for Python 3.14+ compatibility."""
return asyncio.Lock()
def _connect(self) -> sqlite3.Connection:
return sqlite3.connect(self._db_path)
def _write_in_memory_state(self, serialized_state: dict[str, Any]) -> None:
"""Migrate InMemory-format state into the database."""
state = deserialize_state_from_dict(serialized_state, self._serializer)
self._save_state(state) # type: ignore[arg-type]
def _serialize_state(self, state: MODEL_T) -> str:
"""Serialize state model to JSON string."""
if isinstance(state, DictState):
return json.dumps(serialize_dict_state_data(state, self._serializer))
return self._serializer.serialize(state)
def _deserialize_state(self, state_json: str) -> MODEL_T:
"""Deserialize state from JSON string."""
if issubclass(self.state_type, DictState):
data = json.loads(state_json)
return deserialize_dict_state_data(data, self._serializer) # type: ignore[return-value]
return self._serializer.deserialize(state_json)
def _create_default_state(self) -> MODEL_T:
return self.state_type()
def _load_state(self) -> MODEL_T:
"""Load state from database. Creates default if row doesn't exist."""
conn = self._connect()
try:
cursor = conn.cursor()
cursor.execute(
"SELECT state_json FROM state WHERE run_id = ?",
(self._run_id,),
)
row = cursor.fetchone()
if row is None:
state = self._create_default_state()
self._save_state(state, conn)
conn.commit()
return state
return self._deserialize_state(row[0])
finally:
conn.close()
def _save_state(
self, state: MODEL_T, conn: sqlite3.Connection | None = None
) -> None:
"""Save state to database."""
should_close = conn is None
if conn is None:
conn = self._connect()
try:
now = _utc_now().isoformat()
state_json = self._serialize_state(state)
conn.execute(
"""
INSERT INTO state (run_id, state_json, state_type, state_module, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?)
ON CONFLICT(run_id) DO UPDATE SET
state_json = excluded.state_json,
state_type = excluded.state_type,
state_module = excluded.state_module,
updated_at = excluded.updated_at
""",
(
self._run_id,
state_json,
type(state).__name__,
type(state).__module__,
now,
now,
),
)
if should_close:
conn.commit()
finally:
if should_close:
conn.close()
async def get_state(self) -> MODEL_T:
"""Return a copy of the current state model."""
state = self._load_state()
return state.model_copy()
async def set_state(self, state: MODEL_T) -> None:
"""Replace or merge into the current state model."""
conn = self._connect()
try:
cursor = conn.cursor()
cursor.execute(
"SELECT state_json FROM state WHERE run_id = ?",
(self._run_id,),
)
row = cursor.fetchone()
if row is None:
self._save_state(state, conn)
conn.commit()
return
current_state = self._deserialize_state(row[0])
merged = merge_state(current_state, state)
self._save_state(merged, conn) # type: ignore[arg-type]
conn.commit()
finally:
conn.close()
async def get(self, path: str, default: Any = ...) -> Any:
"""Get a nested value using dot-separated paths."""
state = self._load_state()
return get_by_path(state, path, default)
async def set(self, path: str, value: Any) -> None:
"""Set a nested value using dot-separated paths."""
async with self.edit_state() as state:
set_by_path(state, path, value)
async def clear(self) -> None:
"""Reset the state to its type defaults."""
await self.set_state(create_cleared_state(self.state_type))
@asynccontextmanager
async def edit_state(self) -> AsyncGenerator[MODEL_T, None]:
"""Edit state transactionally under a lock."""
async with self._lock:
state = self._load_state()
yield state
self._save_state(state)
def to_dict(self, serializer: BaseSerializer) -> dict[str, Any]:
"""Serialize state store metadata for persistence.
Returns metadata only — actual state lives in the database.
"""
payload = SqliteSerializedState(run_id=self._run_id)
return payload.model_dump()
@classmethod
def from_dict(
cls,
serialized_state: dict[str, Any],
serializer: BaseSerializer,
db_path: str | None = None,
state_type: type[BaseModel] | None = None,
run_id: str | None = None,
) -> SqliteStateStore[Any]:
"""Restore a state store from serialized payload.
Handles both InMemorySerializedState (migrates data to DB on first use)
and SqliteSerializedState (reconnects to existing row).
"""
if not serialized_state:
raise ValueError("Cannot restore SqliteStateStore from empty dict")
store_type = serialized_state.get("store_type")
if store_type == "sqlite":
parsed = SqliteSerializedState.model_validate(serialized_state)
effective_run_id = run_id or parsed.run_id
if db_path is None:
raise ValueError("db_path is required for SqliteStateStore.from_dict()")
return cls(
db_path=db_path,
run_id=effective_run_id,
state_type=state_type, # type: ignore[arg-type]
serializer=serializer,
)
# InMemory format — migrate data to DB immediately
parse_in_memory_state(serialized_state)
effective_run_id = run_id or str(uuid.uuid4())
if db_path is None:
raise ValueError("db_path is required for SqliteStateStore.from_dict()")
store = cls(
db_path=db_path,
run_id=effective_run_id,
state_type=state_type, # type: ignore[arg-type]
serializer=serializer,
)
store._write_in_memory_state(serialized_state)
return store
@@ -1,23 +1,57 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
import json
import sqlite3
import weakref
from collections.abc import AsyncIterator
from datetime import datetime
from typing import List, Optional, Sequence, Tuple
from typing import Any, List, Sequence
from llama_agents.client.protocol.serializable_events import EventEnvelopeWithMetadata
from workflows.context import JsonSerializer
from ..abstract_workflow_store import (
AbstractWorkflowStore,
HandlerQuery,
PersistentHandler,
StoredEvent,
StoredTick,
)
from .migrate import run_migrations
from .sqlite_state_store import SqliteStateStore
class SqliteWorkflowStore(AbstractWorkflowStore):
def __init__(self, db_path: str) -> None:
def __init__(self, db_path: str, poll_interval: float = 1.0) -> None:
self.db_path = db_path
self.poll_interval = poll_interval
self._conditions: weakref.WeakValueDictionary[str, asyncio.Condition] = (
weakref.WeakValueDictionary()
)
self._init_db()
def create_state_store(
self, run_id: str, state_type: type[Any] | None = None
) -> SqliteStateStore[Any]:
return SqliteStateStore(
db_path=self.db_path, run_id=run_id, state_type=state_type
)
def _get_or_create_condition(self, run_id: str) -> asyncio.Condition:
"""Get or create a condition for a run_id.
The caller is responsible for holding a strong reference to the
returned Condition for as long as it needs notifications.
"""
cond = self._conditions.get(run_id)
if cond is None:
cond = asyncio.Condition()
self._conditions[run_id] = cond
return cond
def _init_db(self) -> None:
conn = sqlite3.connect(self.db_path)
try:
@@ -33,7 +67,7 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
clauses, params = filter_spec
sql = """SELECT handler_id, workflow_name, status, run_id, error, result,
started_at, updated_at, completed_at, idle_since, ctx FROM handlers"""
started_at, updated_at, completed_at, idle_since FROM handlers"""
if clauses:
sql = f"{sql} WHERE {' AND '.join(clauses)}"
conn = sqlite3.connect(self.db_path)
@@ -53,8 +87,8 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
cursor.execute(
"""
INSERT INTO handlers (handler_id, workflow_name, status, run_id, error, result,
started_at, updated_at, completed_at, idle_since, ctx)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
started_at, updated_at, completed_at, idle_since)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(handler_id) DO UPDATE SET
workflow_name = excluded.workflow_name,
status = excluded.status,
@@ -64,8 +98,7 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
started_at = excluded.started_at,
updated_at = excluded.updated_at,
completed_at = excluded.completed_at,
idle_since = excluded.idle_since,
ctx = excluded.ctx
idle_since = excluded.idle_since
""",
(
handler.handler_id,
@@ -80,7 +113,6 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
handler.updated_at.isoformat() if handler.updated_at else None,
handler.completed_at.isoformat() if handler.completed_at else None,
handler.idle_since.isoformat() if handler.idle_since else None,
json.dumps(handler.ctx),
),
)
conn.commit()
@@ -107,11 +139,142 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
return int(deleted)
def _build_filters(
self, query: HandlerQuery
) -> Optional[Tuple[List[str], List[str]]]:
clauses: List[str] = []
params: List[str] = []
async def append_event(self, run_id: str, event: EventEnvelopeWithMetadata) -> None:
conn = sqlite3.connect(self.db_path)
try:
conn.execute(
"""INSERT INTO events (run_id, sequence, timestamp, event_json)
VALUES (?, COALESCE((SELECT MAX(sequence) FROM events WHERE run_id = ?), -1) + 1, CURRENT_TIMESTAMP, ?)""",
(
run_id,
run_id,
event.model_dump_json(),
),
)
conn.commit()
finally:
conn.close()
condition = self._conditions.get(run_id)
if condition is not None:
async with condition:
condition.notify_all()
async def query_events(
self,
run_id: str,
after_sequence: int | None = None,
limit: int | None = None,
) -> list[StoredEvent]:
conn = sqlite3.connect(self.db_path)
try:
sql = "SELECT run_id, sequence, timestamp, event_json FROM events WHERE run_id = ?"
params: list[Any] = [run_id]
if after_sequence is not None:
sql += " AND sequence > ?"
params.append(after_sequence)
sql += " ORDER BY sequence"
if limit is not None:
sql += " LIMIT ?"
params.append(limit)
cursor = conn.cursor()
cursor.execute(sql, params)
rows = cursor.fetchall()
finally:
conn.close()
return [
StoredEvent(
run_id=row[0],
sequence=row[1],
timestamp=datetime.fromisoformat(row[2]),
event=EventEnvelopeWithMetadata.model_validate_json(row[3]),
)
for row in rows
]
async def subscribe_events(
self, run_id: str, after_sequence: int = -1
) -> AsyncIterator[StoredEvent]:
condition = self._get_or_create_condition(run_id)
cursor = after_sequence
while True:
events = await self.query_events(run_id, after_sequence=cursor)
for event in events:
yield event
cursor = event.sequence
if self._is_terminal_event(event):
return
if not events:
# Wait for notification or poll timeout
async with condition:
try:
await asyncio.wait_for(
condition.wait(), timeout=self.poll_interval
)
except TimeoutError:
pass
async def append_tick(self, run_id: str, tick_data: dict[str, Any]) -> None:
conn = sqlite3.connect(self.db_path)
try:
conn.execute(
"""INSERT INTO ticks (run_id, sequence, timestamp, tick_data)
VALUES (?, COALESCE((SELECT MAX(sequence) FROM ticks WHERE run_id = ?), -1) + 1, CURRENT_TIMESTAMP, ?)""",
(
run_id,
run_id,
json.dumps(tick_data),
),
)
conn.commit()
finally:
conn.close()
async def get_ticks(self, run_id: str) -> List[StoredTick]:
conn = sqlite3.connect(self.db_path)
try:
cursor = conn.cursor()
cursor.execute(
"SELECT run_id, sequence, timestamp, tick_data FROM ticks WHERE run_id = ? ORDER BY sequence",
(run_id,),
)
rows = cursor.fetchall()
finally:
conn.close()
return [
StoredTick(
run_id=row[0],
sequence=row[1],
timestamp=datetime.fromisoformat(row[2]),
tick_data=json.loads(row[3]),
)
for row in rows
]
def get_legacy_ctx(self, run_id: str) -> dict[str, Any] | None:
"""Read the old ctx column for a run_id, if present."""
conn = sqlite3.connect(self.db_path)
try:
cursor = conn.cursor()
cursor.execute(
"SELECT ctx FROM handlers WHERE run_id = ?",
(run_id,),
)
row = cursor.fetchone()
if row is None or row[0] is None:
return None
try:
data = json.loads(row[0])
if not isinstance(data, dict) or not data:
return None
return data
except (json.JSONDecodeError, TypeError):
return None
finally:
conn.close()
def _build_filters(self, query: HandlerQuery) -> tuple[list[str], list[str]] | None:
clauses: list[str] = []
params: list[str] = []
def add_in_clause(column: str, values: Sequence[str]) -> None:
placeholders = ",".join(["?"] * len(values))
@@ -128,6 +291,11 @@ class SqliteWorkflowStore(AbstractWorkflowStore):
return None
add_in_clause("handler_id", query.handler_id_in)
if query.run_id_in is not None:
if len(query.run_id_in) == 0:
return None
add_in_clause("run_id", query.run_id_in)
if query.status_in is not None:
if len(query.status_in) == 0:
return None
@@ -157,5 +325,4 @@ def _row_to_persistent_handler(row: tuple) -> PersistentHandler:
updated_at=datetime.fromisoformat(row[7]) if row[7] else None,
completed_at=datetime.fromisoformat(row[8]) if row[8] else None,
idle_since=datetime.fromisoformat(row[9]) if row[9] else None,
ctx=json.loads(row[10]),
)
@@ -5,17 +5,22 @@ from __future__ import annotations
import json
import logging
from contextlib import asynccontextmanager
from datetime import timedelta
from typing import Any, AsyncGenerator
import uvicorn
from llama_agents.server._runtime.idle_release_runtime import IdleReleaseDecorator
from llama_agents.server._runtime.persistence_runtime import PersistenceDecorator
from llama_agents.server._runtime.server_runtime import ServerRuntimeDecorator
from starlette.middleware import Middleware
from workflows import Workflow
from workflows.events import Event
from workflows.plugins.basic import basic_runtime
from workflows.runtime.types.plugin import Runtime
from ._api import _WorkflowAPI
from ._service import _WorkflowService
from ._store.abstract_workflow_store import AbstractWorkflowStore
from ._store.memory_workflow_store import MemoryWorkflowStore
logger = logging.getLogger()
@@ -28,26 +33,62 @@ class WorkflowServer:
workflow_store: AbstractWorkflowStore | None = None,
# retry/backoff seconds for persisting the handler state in the store after failures. Configurable mainly for testing.
persistence_backoff: list[float] = [0.5, 3],
# Release idle workflows from memory after this timeout (None = disabled)
idle_release_timeout: timedelta | None = timedelta(seconds=10),
runtime: Runtime | None = None,
idle_timeout: float = 60.0,
):
self._service = _WorkflowService(
workflow_store=workflow_store,
persistence_backoff=persistence_backoff,
idle_release_timeout=idle_release_timeout,
self._workflow_store = (
workflow_store if workflow_store is not None else MemoryWorkflowStore()
)
inner: Runtime = (
runtime
if runtime is not None
else IdleReleaseDecorator(
PersistenceDecorator(basic_runtime, store=self._workflow_store),
store=self._workflow_store,
idle_timeout=idle_timeout,
)
)
self._runtime: ServerRuntimeDecorator = ServerRuntimeDecorator(
inner,
store=self._workflow_store,
persistence_backoff=list(persistence_backoff),
)
self._service = _WorkflowService(
runtime=self._runtime, store=self._workflow_store
)
self._api = _WorkflowAPI(self._service, middleware=middleware)
self.app = self._api.app
# ------------------------------------------------------------------
# Workflow registration
# ------------------------------------------------------------------
def add_workflow(
self,
name: str,
workflow: Workflow,
additional_events: list[type[Event]] | None = None,
) -> None:
self._service.add_workflow(name, workflow, additional_events)
workflow._switch_workflow_name(name)
workflow._switch_runtime(self._runtime)
async def start(self) -> "WorkflowServer":
if additional_events is not None:
self._api.register_additional_events(name, additional_events)
def get_workflows(self) -> dict[str, Workflow]:
"""Return registered workflows as a dict by name. Only available after start()."""
return {
n: wf
for n in self._service.get_workflow_names()
if (wf := self._service.get_workflow(n)) is not None
}
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
async def start(self) -> WorkflowServer:
"""Resumes previously running workflows, if they were not complete at last shutdown.
Idle workflows are not resumed - they remain released and will be
@@ -57,7 +98,7 @@ class WorkflowServer:
return self
@asynccontextmanager
async def contextmanager(self) -> AsyncGenerator["WorkflowServer", None]:
async def contextmanager(self) -> AsyncGenerator[WorkflowServer, None]:
"""Use this server as a context manager to start and stop it"""
await self.start()
try:
@@ -68,6 +109,10 @@ class WorkflowServer:
async def stop(self) -> None:
await self._service.stop()
# ------------------------------------------------------------------
# Serve
# ------------------------------------------------------------------
async def serve(
self,
host: str = "localhost",
@@ -6,12 +6,13 @@ import asyncio
import socket
import time
from contextlib import asynccontextmanager
from pathlib import Path
from typing import AsyncGenerator, Awaitable, Callable, TypeVar
import httpx
import pytest
import uvicorn
from llama_agents.server import WorkflowServer
from llama_agents.server import MemoryWorkflowStore, SqliteWorkflowStore, WorkflowServer
from workflows import Context, Workflow, step
from workflows.events import (
Event,
@@ -193,6 +194,16 @@ class StructuredStartWorkflow(Workflow):
return StopEvent(result=ev.message)
@pytest.fixture
def memory_store() -> MemoryWorkflowStore:
return MemoryWorkflowStore()
@pytest.fixture
def sqlite_store(tmp_path: Path) -> SqliteWorkflowStore:
return SqliteWorkflowStore(str(tmp_path / "test.db"))
@pytest.fixture
def simple_test_workflow() -> Workflow:
return SimpleTestWorkflow()
File diff suppressed because it is too large Load Diff
@@ -1,64 +1,16 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
import json
from datetime import datetime, timezone
from typing import AsyncGenerator
from unittest.mock import MagicMock
import pytest
from llama_agents.client.protocol import HandlerData
from llama_agents.server._handler import _WorkflowHandler
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from workflows.events import Event, StopEvent
from workflows.handler import WorkflowHandler
from workflows.runtime.types.plugin import ExternalRunAdapter
from workflows.runtime.types.ticks import WorkflowTick
from workflows.workflow import Workflow
class MockRunAdapter(ExternalRunAdapter):
"""Minimal mock adapter for testing."""
def __init__(self, run_id: str) -> None:
self._run_id = run_id
self._result: asyncio.Future[StopEvent] = asyncio.Future()
@property
def run_id(self) -> str:
return self._run_id
@property
def is_running(self) -> bool:
return not self._result.done()
async def get_result(self) -> StopEvent:
return await self._result
def get_result_or_none(self) -> StopEvent | None:
if self._result.done() and not self._result.cancelled():
return self._result.result()
return None
async def stream_published_events(self) -> AsyncGenerator[Event, None]:
result = await self._result
yield result
def abort(self) -> None:
if not self._result.done():
self._result.cancel()
def set_result(self, result: StopEvent) -> None:
if not self._result.done():
self._result.set_result(result)
async def send_event(self, tick: WorkflowTick) -> None:
pass
async def close(self) -> None:
pass
from llama_agents.server import PersistentHandler
from llama_agents.server._service import handler_data_from_persistent
from workflows.events import StopEvent
class MyStopEvent(StopEvent):
@@ -66,45 +18,22 @@ class MyStopEvent(StopEvent):
@pytest.mark.asyncio
async def test_workflow_handler_to_dict_json_roundtrip() -> None:
workflow = MagicMock(spec=Workflow)
workflow.workflow_name = "TestWorkflow"
adapter = MockRunAdapter(run_id="test-run-id")
# Set the result on the adapter
stop_event = MyStopEvent(message="ok")
adapter.set_result(stop_event)
handler: WorkflowHandler = WorkflowHandler(
workflow=workflow, external_adapter=adapter
)
# Wait for the result task to complete
await asyncio.sleep(0)
queue: asyncio.Queue[Event] = asyncio.Queue()
async def noop() -> None:
return None
task: asyncio.Task[None] = asyncio.create_task(noop())
# Ensure the task is completed so the wrapper mimics a finished handler
await asyncio.sleep(0)
async def test_handler_data_from_persistent_json_roundtrip() -> None:
now = datetime.now(timezone.utc)
wrapper = _WorkflowHandler(
_workflow_store=MemoryWorkflowStore(),
_persistence_backoff=[0.0, 0.0],
run_handler=handler,
queue=queue,
task=task,
consumer_mutex=asyncio.Lock(),
stop_event = MyStopEvent(message="ok")
persistent = PersistentHandler(
handler_id="handler-1",
workflow_name="wf",
run_id="test-run-id",
status="completed",
started_at=now,
updated_at=now,
completed_at=now,
result=stop_event,
)
response_model = wrapper.to_response_model()
response_model = handler_data_from_persistent(persistent)
# JSON serialization should not error
s = json.dumps(response_model.model_dump())
reparsed_dict = json.loads(s)
@@ -1,788 +0,0 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""Tests for idle workflow release and reload functionality."""
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Any, Optional
import pytest
import time_machine
from httpx import ASGITransport, AsyncClient
from llama_agents.server._handler import _WorkflowHandler
from llama_agents.server._store.abstract_workflow_store import (
HandlerQuery,
PersistentHandler,
Status,
)
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from llama_agents.server.server import WorkflowServer
from server_test_fixtures import async_yield, wait_for_passing
from workflows import Context, Workflow, step
from workflows.context.context_types import SerializedContext
from workflows.events import HumanResponseEvent, StartEvent, StopEvent
class WaitableExternalEvent(HumanResponseEvent):
"""Event sent from external sources."""
response: str
class WaitingWorkflow(Workflow):
"""Workflow that uses ctx.wait_for_event() to properly become idle."""
@step
async def start_and_wait(self, ctx: Context, ev: StartEvent) -> StopEvent:
# Use ctx.wait_for_event() to create a waiter - this makes the workflow idle
external = await ctx.wait_for_event(WaitableExternalEvent)
return StopEvent(result=f"received: {external.response}")
def get_handler_in_memory(server: WorkflowServer, handler_id: str) -> _WorkflowHandler:
wrapper = server._service._handlers.get(handler_id)
assert wrapper is not None, f"Handler {handler_id} not found in memory"
return wrapper
def assert_handler_in_memory(server: WorkflowServer, handler_id: str) -> None:
get_handler_in_memory(server, handler_id)
def assert_handler_not_in_memory(server: WorkflowServer, handler_id: str) -> None:
assert handler_id not in server._service._handlers, (
f"Handler {handler_id} still in memory"
)
def make_server(
memory_store: MemoryWorkflowStore,
waiting_workflow: Workflow,
idle_release_timeout: Optional[timedelta],
*,
persistence_backoff: Optional[list[float]] = None,
) -> WorkflowServer:
if persistence_backoff is None:
server = WorkflowServer(
workflow_store=memory_store,
idle_release_timeout=idle_release_timeout,
)
else:
server = WorkflowServer(
workflow_store=memory_store,
idle_release_timeout=idle_release_timeout,
persistence_backoff=persistence_backoff,
)
server.add_workflow(
"test", waiting_workflow, additional_events=[WaitableExternalEvent]
)
return server
async def start_waiting_handler(
server: WorkflowServer, handler_id: str
) -> _WorkflowHandler:
handler = server._service._workflows["test"].run()
await server._service.run_workflow_handler(handler_id, "test", handler)
await async_yield(20)
wrapper = get_handler_in_memory(server, handler_id)
assert wrapper.idle_since is not None
return wrapper
async def advance_time(traveller: Any, delta: timedelta, iterations: int = 10) -> None:
traveller.shift(delta)
await async_yield(iterations)
async def seed_persistent_handler(
store: MemoryWorkflowStore,
handler_id: str,
*,
idle_since: Optional[datetime],
ctx: dict[str, object],
status: Status = "running",
) -> None:
await store.update(
PersistentHandler(
handler_id=handler_id,
workflow_name="test",
status=status,
idle_since=idle_since,
ctx=ctx,
)
)
@pytest.fixture
def memory_store() -> MemoryWorkflowStore:
return MemoryWorkflowStore()
@pytest.fixture
def waiting_workflow() -> Workflow:
return WaitingWorkflow()
@pytest.mark.asyncio
async def test_is_idle_query_filter_memory_store() -> None:
"""Test that is_idle filter works in MemoryWorkflowStore."""
store = MemoryWorkflowStore()
now = datetime.now(timezone.utc)
# Handler that is idle (has idle_since set)
await seed_persistent_handler(
store,
"idle-1",
idle_since=now - timedelta(minutes=5),
ctx={},
)
# Handler that is not idle (no idle_since)
await seed_persistent_handler(store, "active-1", idle_since=None, ctx={})
# Another idle handler
await seed_persistent_handler(
store,
"idle-2",
idle_since=now - timedelta(seconds=10),
ctx={},
)
# Query for idle handlers
idle_results = await store.query(HandlerQuery(is_idle=True))
assert len(idle_results) == 2
assert {r.handler_id for r in idle_results} == {"idle-1", "idle-2"}
# Query for non-idle handlers
active_results = await store.query(HandlerQuery(is_idle=False))
assert len(active_results) == 1
assert active_results[0].handler_id == "active-1"
# Query without filter returns all
all_results = await store.query(HandlerQuery())
assert len(all_results) == 3
@pytest.mark.asyncio
async def test_workflow_becomes_idle_and_is_released(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that a workflow becomes idle after a step and can be released."""
idle_timeout = timedelta(milliseconds=50)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
# Start a workflow
handler_id = "release-test-1"
await start_waiting_handler(server, handler_id)
# Advance time past the idle timeout to trigger the timer
await advance_time(traveller, timedelta(milliseconds=100))
# The handler should be released from memory
assert_handler_not_in_memory(server, handler_id)
# But should still exist in the store with status "running"
persisted = await memory_store.query(
HandlerQuery(handler_id_in=[handler_id])
)
assert len(persisted) == 1
assert persisted[0].status == "running"
assert persisted[0].idle_since is not None
@pytest.mark.asyncio
async def test_released_workflow_is_reloaded_on_event(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that a released workflow is reloaded when an event is sent."""
idle_timeout = timedelta(milliseconds=50)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
# Start a workflow
handler_id = "reload-test-1"
await start_waiting_handler(server, handler_id)
# Advance time past the idle timeout to trigger the timer
await advance_time(traveller, timedelta(milliseconds=100))
# Handler should be released
assert_handler_not_in_memory(server, handler_id)
# Now reload by using _try_reload_handler
reloaded, persisted = await server._service.try_reload_handler(handler_id)
assert reloaded is not None
assert_handler_in_memory(server, handler_id)
assert persisted is not None
assert persisted.status == "running"
# Send the event to complete the workflow
ctx = reloaded.run_handler.ctx
assert ctx is not None
ctx.send_event(WaitableExternalEvent(response="hello"))
result = await reloaded.run_handler
assert result == "received: hello"
@pytest.mark.asyncio
async def test_idle_release_restores_idle_since_on_reload(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that idle_since is preserved when reloading via _try_reload_handler."""
server = make_server(
memory_store,
waiting_workflow,
timedelta(minutes=5), # Long timeout
)
# Seed a persisted handler with idle_since
idle_time = datetime.now(timezone.utc) - timedelta(minutes=2)
ctx = SerializedContext().model_dump(mode="python")
await seed_persistent_handler(
memory_store,
"idle-restore-1",
idle_since=idle_time,
ctx=ctx,
)
async with server.contextmanager():
# Idle handler should NOT be in memory on startup
assert_handler_not_in_memory(server, "idle-restore-1")
# Reload it (simulating an event arriving)
wrapper, persisted = await server._service.try_reload_handler("idle-restore-1")
assert wrapper is not None
assert wrapper.idle_since == idle_time
assert persisted is not None
assert persisted.status == "running"
@pytest.mark.asyncio
async def test_reloaded_idle_workflow_is_released_again(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that a reloaded workflow that stays idle gets released again."""
# Use very short timeout - this test needs real time for timers
idle_timeout = timedelta(milliseconds=1)
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
# Start a workflow
handler_id = "reload-idle-test-1"
handler = server._service._workflows["test"].run()
wrapper = await server._service.run_workflow_handler(
handler_id, "test", handler
)
async def wrapper_is_idle() -> None:
assert wrapper.idle_since is not None
await wait_for_passing(wrapper_is_idle, interval=0.01, max_duration=1.5)
# Reload the workflow (simulating an event arriving) once the handler
# is released from memory.
async def reload_from_store() -> tuple[_WorkflowHandler, PersistentHandler]:
reloaded, persisted = await server._service.try_reload_handler(handler_id)
assert reloaded is not None
assert persisted is not None
assert reloaded is not wrapper
return reloaded, persisted
reloaded, persisted = await wait_for_passing(
reload_from_store, interval=0.01, max_duration=1.5
)
assert persisted is not None
assert persisted.status == "running"
# The reloaded handler should have idle_since restored
assert reloaded.idle_since is not None
# Wait for the reloaded handler to be released again by observing
# that a subsequent reload returns a new handler instance.
async def reload_after_release() -> _WorkflowHandler:
reloaded_again, persisted_again = await server._service.try_reload_handler(
handler_id
)
assert reloaded_again is not None
assert persisted_again is not None
assert reloaded_again is not reloaded
return reloaded_again
reloaded_again = await wait_for_passing(
reload_after_release, interval=0.01, max_duration=1.5
)
await server._service.close_handler(reloaded_again)
@pytest.mark.asyncio
async def test_idle_handlers_not_resumed_on_server_start(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that idle handlers are not loaded into memory on server start."""
# Seed the store with an idle handler
idle_time = datetime.now(timezone.utc) - timedelta(minutes=2)
ctx = SerializedContext().model_dump(mode="python")
await seed_persistent_handler(
memory_store,
"idle-on-start-1",
idle_since=idle_time,
ctx=ctx,
)
# Also seed an active (non-idle) handler
await seed_persistent_handler(
memory_store,
"active-on-start-1",
idle_since=None, # Not idle
ctx=ctx,
)
server = make_server(
memory_store,
waiting_workflow,
timedelta(minutes=5),
)
async with server.contextmanager():
# The idle handler should NOT be in memory
assert_handler_not_in_memory(server, "idle-on-start-1")
# The active handler SHOULD be in memory
assert_handler_in_memory(server, "active-on-start-1")
# The idle handler should still exist in the store
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["idle-on-start-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "running"
@pytest.mark.asyncio
async def test_idle_release_disabled_when_timeout_none(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Test that no timer is started when idle_release_timeout is None."""
server = make_server(
memory_store,
waiting_workflow,
None, # Disabled
)
async with server.contextmanager():
# Start a workflow
handler_id = "no-release-test"
wrapper = await start_waiting_handler(server, handler_id)
# Timer should not be set since idle_release_timeout is None
assert wrapper._idle_release_timer is None
@pytest.mark.asyncio
async def test_idle_release_cancels_runtime(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Idle release should stop the workflow runtime.
When a workflow is released from memory, the underlying WorkflowHandler
and its context/broker should be cancelled. Currently, _release_handler
only cancels the server's stream task and removes from _handlers, but
the actual workflow runtime keeps running in the background.
This test captures a reference to the run_handler before release, then
verifies it should be done/cancelled after release.
"""
idle_timeout = timedelta(milliseconds=50)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
# Start a workflow
handler_id = "runtime-leak-test"
wrapper = await start_waiting_handler(server, handler_id)
# Capture reference to the run_handler before release
run_handler = wrapper.run_handler
# Advance time past the idle timeout to trigger the timer
await advance_time(traveller, timedelta(milliseconds=100))
# Handler should be released
assert_handler_not_in_memory(server, handler_id)
assert run_handler.done(), (
"run_handler is still running after release workflow runtime was not stopped, causing memory leak"
)
@pytest.mark.asyncio
async def test_idle_release_waits_for_stream_consumer_then_releases(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Idle release should wait for active consumers and then reschedule.
The idle release timer checks if consumer_mutex is locked and skips release
if so. However, it does NOT reschedule another timer attempt. This means
if a client briefly holds a stream during the timeout window and then
disconnects, the handler will never be released.
This test holds the mutex during the timer, releases it, then verifies
the handler should eventually be released (but won't be due to the issue).
"""
idle_timeout = timedelta(milliseconds=50)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
# Start workflow
handler_id = "mutex-reschedule-test"
wrapper = await start_waiting_handler(server, handler_id)
# Grab the mutex before timer fires, then hold it past the timeout
async with wrapper.consumer_mutex:
# Advance time past the timeout to trigger the timer
await advance_time(traveller, timedelta(milliseconds=100))
# Handler should still be present (timer was blocked by mutex)
assert_handler_in_memory(server, handler_id)
# Now mutex is released - handler should eventually be released
# because the timer reschedules when mutex was locked.
# Advance time to let the rescheduled timer fire.
await advance_time(traveller, timedelta(milliseconds=100))
# Handler should be released because timer was rescheduled
assert_handler_not_in_memory(server, handler_id)
@pytest.mark.asyncio
async def test_mark_active_prevents_release_during_event_post(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Marking active should protect from release during event post.
When an event is posted to wake a workflow, the idle_since is only cleared
when StepStateChanged(RUNNING) is observed in _stream_events. However,
the idle timer can fire in the window between event posting and the
running event being processed, causing the handler to be released while
the workflow is actively processing.
DESIRED IMPLEMENTATION:
We want to support short/immediate idle timeouts (e.g., 0 timeout to release
workflows immediately when idle). However, "revived" workflows (ones that
just received an event but haven't processed it yet) need protection. The
fix should use a minimum grace period for revived workflows - when an event
is sent to an idle workflow, the timer should be rescheduled with at least
a minimum timeout (e.g., 1 second) regardless of the configured timeout.
This prevents immediate release while still allowing fast release of truly
idle workflows.
The fix implemented: The server's _post_event endpoint calls mark_active()
before sending an event, which clears idle_since and cancels the timer.
This test verifies that mark_active() properly protects the handler.
"""
idle_timeout = timedelta(milliseconds=50)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
handler_id = "race-test"
wrapper = await start_waiting_handler(server, handler_id)
# Mark active before sending event (this is what _post_event does)
# This clears idle_since and cancels the timer
wrapper.mark_active()
# Send event to wake the workflow
ctx = wrapper.run_handler.ctx
assert ctx is not None
ctx.send_event(WaitableExternalEvent(response="wake-up"))
# Advance time past the original idle timeout
# The workflow should be protected from release because mark_active was called
await advance_time(traveller, timedelta(milliseconds=100))
# Handler should NOT be released because mark_active()
# cleared idle_since and cancelled the timer before the event was sent
assert_handler_in_memory(server, handler_id)
@pytest.mark.asyncio
async def test_try_reload_is_singleton_under_concurrency(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Concurrent reload requests should only create one handler.
_try_reload_handler has no locking mechanism, so two parallel requests
can both reload and start the same workflow. This creates split-brain
scenarios where multiple workflow instances process events independently.
This test fires multiple reload requests in parallel and checks that
only one succeeds in creating a handler.
"""
server = make_server(
memory_store,
waiting_workflow,
timedelta(minutes=5),
)
# Seed a persisted idle handler
idle_time = datetime.now(timezone.utc) - timedelta(minutes=2)
ctx = SerializedContext().model_dump(mode="python")
handler_id = "concurrent-reload-test"
await seed_persistent_handler(
memory_store,
handler_id,
idle_since=idle_time,
ctx=ctx,
)
async with server.contextmanager():
assert_handler_not_in_memory(server, handler_id)
# Track how many workflow instances were created
instances_created: list[object] = []
async def reload_and_track() -> None:
wrapper, _ = await server._service.try_reload_handler(handler_id)
if wrapper is not None:
# Track the actual run_handler object identity
instances_created.append(id(wrapper.run_handler))
# Fire 5 concurrent reload requests
await asyncio.gather(*[reload_and_track() for _ in range(5)])
# We should have exactly one handler in memory
assert_handler_in_memory(server, handler_id)
# Multiple unique workflow instances may have been created
# (even though only one ends up in _handlers, the others are leaked)
unique_instances = set(instances_created)
assert len(unique_instances) == 1, (
f"{len(unique_instances)} different workflow instances were created "
f"by concurrent reloads. Only one should be created. "
f"The extra instances are leaked and may process events independently."
)
@pytest.mark.asyncio
async def test_release_handler_cancels_runtime_on_checkpoint_failure(
memory_store: MemoryWorkflowStore,
waiting_workflow: Workflow,
) -> None:
"""Release should cancel runtime even if checkpoint fails.
If checkpoint() raises after the handler is already removed from _handlers,
the runtime cancel path never runs, leaking a live workflow with no
in-memory reference.
The current _release_handler structure is:
1. Pop from _handlers (handler removed from memory)
2. Call checkpoint() (if this fails, exception propagates)
3. Cancel runtime (never reached if step 2 fails)
This test triggers a realistic failure by storing a non-serializable object
(a lambda) in the context state. When the idle release timer fires and
tries to checkpoint, serialization fails and the runtime leaks.
The fix should ensure that if ANY part of release fails, the workflow
runtime is still properly cancelled (e.g., use try/finally).
"""
idle_timeout = timedelta(milliseconds=100)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False) as traveller:
server = make_server(
memory_store,
waiting_workflow,
idle_timeout,
persistence_backoff=[], # No retries - fail immediately
)
await server.start()
run_handler = None
try:
handler_id = "checkpoint-fail-test"
wrapper = await start_waiting_handler(server, handler_id)
# Capture reference to verify runtime is stopped
run_handler = wrapper.run_handler
ctx = run_handler.ctx
assert ctx is not None
# Store a non-serializable object in context store directly
# (bypassing the Context.store property which requires internal context)
# This will cause checkpoint() to fail when it tries to serialize
state_store = ctx._face._external_adapter.get_state_store() # type: ignore[union-attr]
assert state_store is not None
await state_store.set("bad_data", lambda x: x)
# Advance time past the idle timeout to trigger the timer
await advance_time(traveller, timedelta(milliseconds=200), iterations=20)
# Verify handler was removed from _handlers (release started)
assert_handler_not_in_memory(server, handler_id)
# The run_handler should be done/cancelled even if checkpoint failed
# but it's still running because the cancel code was never reached
assert run_handler.done(), (
"run_handler is still running after release attempt failed. "
"_release_handler pops from _handlers before checkpoint(), so if "
"checkpoint() raises (e.g., due to non-serializable state), the "
"cancel path is never reached and the workflow runtime leaks."
)
finally:
# Clean up the leaked runtime manually (since the issue prevents normal cleanup)
if run_handler is not None and not run_handler.done():
run_handler.cancel()
try:
await run_handler.cancel_run()
except Exception:
pass
await server.stop()
@pytest.mark.asyncio
async def test_send_event_clears_idle_state_before_processing(
memory_store: MemoryWorkflowStore, waiting_workflow: Workflow
) -> None:
"""Verify that send_event() clears idle state immediately to prevent race conditions.
mark_active() is called BEFORE processing the event to prevent a race where
the idle release timer fires while we're still handling the request. This means
idle_since is cleared even if send_event() subsequently fails.
"""
idle_timeout = timedelta(milliseconds=100)
with time_machine.travel("2026-01-07T12:00:00Z", tick=False):
server = make_server(memory_store, waiting_workflow, idle_timeout)
async with server.contextmanager():
transport = ASGITransport(app=server.app)
async with AsyncClient(
transport=transport, base_url="http://test"
) as client:
# Start a workflow via HTTP
response = await client.post("/workflows/test/run-nowait", json={})
assert response.status_code == 200
handler_id = response.json()["handler_id"]
# Wait for workflow to become idle (just needs event loop iterations)
await async_yield(20)
wrapper = get_handler_in_memory(server, handler_id)
assert wrapper is not None and wrapper.idle_since is not None
# Post an event with a bad step name via HTTP - this will fail
response = await client.post(
f"/events/{handler_id}",
json={
"event": {
"type": "WaitableExternalEvent",
"value": {"response": "test"},
},
"step": "nonexistent_step", # This step doesn't exist
},
)
# The endpoint returns 400 for bad step
assert response.status_code == 400
# mark_active() is called BEFORE send_event() to prevent race conditions.
# Even though send_event() failed, idle state is cleared.
assert wrapper.idle_since is None, (
"idle_since should be cleared before send_event() is attempted"
)
assert wrapper._idle_release_timer is None, "timer should be cancelled"
@pytest.mark.asyncio
async def test_release_skips_checkpoint_if_handler_was_reloaded(
memory_store: MemoryWorkflowStore,
waiting_workflow: Workflow,
) -> None:
"""Verify that _release_handler skips checkpoint if handler was reloaded.
When release and reload race, _release_handler should detect if a new
handler instance is now in _handlers and skip the checkpoint to avoid
overwriting newer state.
This test verifies:
1. Start handler, wait for idle
2. Reload the handler (simulating it was released earlier)
3. Update the reloaded handler's state
4. Call _release_handler with the OLD wrapper
5. Verify the old release doesn't overwrite the new state
"""
# time_machine with tick=False for fast test execution
with time_machine.travel("2026-01-07T12:27:00.000-08:00", tick=False):
server = WorkflowServer(
workflow_store=memory_store,
idle_release_timeout=None, # Disable auto-release for manual control
)
server.add_workflow(
"test", waiting_workflow, additional_events=[WaitableExternalEvent]
)
async with server.contextmanager():
# Start workflow and wait for it to become idle
handler_id = "race-test-handler"
await start_waiting_handler(server, handler_id)
old_wrapper = get_handler_in_memory(server, handler_id)
original_idle_since = old_wrapper.idle_since
assert original_idle_since is not None
# Checkpoint the old state (simulating what release would have done)
await old_wrapper.checkpoint()
# Simulate the scenario where handler was released and then reloaded:
# Remove old handler from memory
server._service._handlers.pop(handler_id, None)
# Reload the handler (this gets the persisted state)
reloaded, persisted = await server._service.try_reload_handler(handler_id)
assert reloaded is not None
assert reloaded is not old_wrapper # Different instance
assert persisted is not None
assert persisted.status == "running"
# The reloaded handler receives an event and becomes active
reloaded.mark_active()
assert reloaded.idle_since is None
# Checkpoint the new state
await reloaded.checkpoint()
# Verify store has the new state (idle_since=None)
stored = await memory_store.query(HandlerQuery(handler_id_in=[handler_id]))
assert stored[0].idle_since is None, "Store should have new state"
# NOW call _release_handler with the OLD wrapper
# This simulates a delayed release that happens after reload
# _release_handler should detect that a different handler is in _handlers
# and skip the checkpoint
await server._service.release_handler(old_wrapper)
# Verify the store still has the correct (new) state
stored = await memory_store.query(HandlerQuery(handler_id_in=[handler_id]))
assert len(stored) == 1
# The old release should have skipped the checkpoint
# because it detected a different handler instance in _handlers
assert stored[0].idle_since is None, (
f"Release should have skipped checkpoint since handler was reloaded. "
f"Expected idle_since=None (from reload), "
f"but got idle_since={stored[0].idle_since}."
)
@@ -4,12 +4,9 @@
from __future__ import annotations
from datetime import timedelta
import pytest
from llama_agents.client.client import WorkflowClient
from llama_agents.server import WorkflowServer
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from llama_agents.server import MemoryWorkflowStore, WorkflowServer
from server_test_fixtures import (
live_server, # type: ignore[import]
wait_for_passing, # type: ignore[import]
@@ -34,9 +31,6 @@ async def test_fast_idle_timeout_does_not_drop_valid_event() -> None:
def make_server() -> WorkflowServer:
server = WorkflowServer(
workflow_store=MemoryWorkflowStore(),
# 50ms is still fast enough to test idle release behavior while
# avoiding race conditions in test infrastructure (especially with xdist)
idle_release_timeout=timedelta(milliseconds=50),
)
server.add_workflow(
"waiting",
@@ -1,12 +1,25 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
from datetime import datetime, timezone
import pytest
from llama_agents.server._store.abstract_workflow_store import (
from llama_agents.client.protocol.serializable_events import EventEnvelopeWithMetadata
from llama_agents.server import (
AbstractWorkflowStore,
HandlerQuery,
MemoryWorkflowStore,
PersistentHandler,
)
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from workflows.events import StopEvent
from llama_agents.server._store.abstract_workflow_store import Status, StoredEvent
from workflows.events import (
Event,
StopEvent,
WorkflowCancelledEvent,
WorkflowFailedEvent,
)
@pytest.mark.asyncio
@@ -17,7 +30,6 @@ async def test_update_and_query_returns_inserted_handler() -> None:
handler_id="h1",
workflow_name="wf_a",
status="running",
ctx={"state": {"x": 1, "y": [1, 2, 3]}},
)
await store.update(handler)
@@ -32,7 +44,6 @@ async def test_update_and_query_returns_inserted_handler() -> None:
assert found.handler_id == "h1"
assert found.workflow_name == "wf_a"
assert found.status == "running"
assert found.ctx == {"state": {"x": 1, "y": [1, 2, 3]}}
@pytest.mark.asyncio
@@ -45,17 +56,15 @@ async def test_update_on_conflict_overwrites_existing_row() -> None:
handler_id="h2",
workflow_name="wf_b",
status="running",
ctx={"k": "v1"},
)
)
# Update same handler_id (completed) with new ctx
# Update same handler_id (completed)
await store.update(
PersistentHandler(
handler_id="h2",
workflow_name="wf_b",
status="completed",
ctx={"k": "v2", "n": 42},
)
)
@@ -65,7 +74,7 @@ async def test_update_on_conflict_overwrites_existing_row() -> None:
)
assert result_in_progress == []
# Should be returned for status=completed with latest ctx
# Should be returned for status=completed with latest values
result_completed = await store.query(
HandlerQuery(workflow_name_in=["wf_b"], status_in=["completed"])
)
@@ -74,7 +83,6 @@ async def test_update_on_conflict_overwrites_existing_row() -> None:
assert found.handler_id == "h2"
assert found.workflow_name == "wf_b"
assert found.status == "completed"
assert found.ctx == {"k": "v2", "n": 42}
@pytest.mark.asyncio
@@ -86,7 +94,6 @@ async def test_delete_filters_by_query() -> None:
handler_id="delete-me",
workflow_name="wf_delete",
status="completed",
ctx={"val": 1},
)
)
await store.update(
@@ -94,7 +101,6 @@ async def test_delete_filters_by_query() -> None:
handler_id="keep-me",
workflow_name="wf_keep",
status="running",
ctx={"val": 2},
)
)
@@ -115,7 +121,6 @@ async def test_delete_noop_on_empty_filter() -> None:
handler_id="delete-me",
workflow_name="wf_delete",
status="completed",
ctx={},
)
)
@@ -138,7 +143,6 @@ async def test_query_filters_by_handler_id_and_empty_lists() -> None:
handler_id=hid,
workflow_name=wf,
status="running",
ctx={"seed": hid},
)
)
@@ -169,7 +173,6 @@ async def test_query_filters_by_multiple_statuses() -> None:
handler_id="h1",
workflow_name="wf",
status="running",
ctx={},
)
)
await store.update(
@@ -177,7 +180,6 @@ async def test_query_filters_by_multiple_statuses() -> None:
handler_id="h2",
workflow_name="wf",
status="completed",
ctx={},
)
)
await store.update(
@@ -185,7 +187,6 @@ async def test_query_filters_by_multiple_statuses() -> None:
handler_id="h3",
workflow_name="wf",
status="failed",
ctx={},
)
)
await store.update(
@@ -193,7 +194,6 @@ async def test_query_filters_by_multiple_statuses() -> None:
handler_id="h4",
workflow_name="wf",
status="cancelled",
ctx={},
)
)
@@ -216,7 +216,6 @@ async def test_query_filters_by_workflow_name() -> None:
handler_id="h1",
workflow_name="wf_a",
status="running",
ctx={},
)
)
await store.update(
@@ -224,7 +223,6 @@ async def test_query_filters_by_workflow_name() -> None:
handler_id="h2",
workflow_name="wf_b",
status="running",
ctx={},
)
)
await store.update(
@@ -232,7 +230,6 @@ async def test_query_filters_by_workflow_name() -> None:
handler_id="h3",
workflow_name="wf_a",
status="completed",
ctx={},
)
)
@@ -257,7 +254,6 @@ async def test_query_combines_multiple_filters() -> None:
handler_id="h1",
workflow_name="wf_a",
status="running",
ctx={},
)
)
await store.update(
@@ -265,7 +261,6 @@ async def test_query_combines_multiple_filters() -> None:
handler_id="h2",
workflow_name="wf_a",
status="completed",
ctx={},
)
)
await store.update(
@@ -273,7 +268,6 @@ async def test_query_combines_multiple_filters() -> None:
handler_id="h3",
workflow_name="wf_b",
status="running",
ctx={},
)
)
await store.update(
@@ -281,7 +275,6 @@ async def test_query_combines_multiple_filters() -> None:
handler_id="h4",
workflow_name="wf_b",
status="completed",
ctx={},
)
)
@@ -315,7 +308,6 @@ async def test_delete_removes_multiple_matching_handlers() -> None:
handler_id=f"h{i}",
workflow_name="wf",
status="completed" if i % 2 == 0 else "running",
ctx={},
)
)
@@ -343,7 +335,6 @@ async def test_store_handles_all_datetime_fields() -> None:
started_at=now,
updated_at=now,
completed_at=now,
ctx={"data": "value"},
)
await store.update(handler)
@@ -367,7 +358,6 @@ async def test_store_handles_error_field() -> None:
workflow_name="wf",
status="failed",
error="Something went wrong",
ctx={},
)
await store.update(handler)
@@ -388,3 +378,195 @@ async def test_empty_store_returns_empty_results() -> None:
# Delete from empty store
deleted = await store.delete(HandlerQuery(handler_id_in=["nonexistent"]))
assert deleted == 0
@pytest.mark.asyncio
async def test_update_handler_status_with_nonexistent_run_id() -> None:
store = MemoryWorkflowStore()
# Should not raise when run_id does not exist
await store.update_handler_status("nonexistent-run-id", status="completed")
@pytest.mark.asyncio
async def test_update_handler_status_sets_status_and_completed_at() -> None:
store = MemoryWorkflowStore()
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
)
)
await store.update_handler_status("run-1", status="completed")
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
assert len(result) == 1
handler = result[0]
assert handler.status == "completed"
assert handler.updated_at is not None
assert handler.completed_at is not None
@pytest.mark.asyncio
async def test_update_handler_status_with_result() -> None:
store = MemoryWorkflowStore()
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
)
)
stop = StopEvent(result={"answer": 42})
await store.update_handler_status("run-1", status="completed", result=stop)
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
handler = result[0]
assert handler.status == "completed"
assert handler.result == stop
@pytest.mark.asyncio
async def test_update_handler_status_with_error() -> None:
store = MemoryWorkflowStore()
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
)
)
await store.update_handler_status("run-1", status="failed", error="boom")
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
handler = result[0]
assert handler.status == "failed"
assert handler.error == "boom"
assert handler.completed_at is not None
@pytest.mark.asyncio
async def test_update_handler_status_idle_since_explicit_none_clears() -> None:
store = MemoryWorkflowStore()
now = datetime.now(timezone.utc)
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
idle_since=now,
)
)
# Passing idle_since=None explicitly should clear it
await store.update_handler_status("run-1", idle_since=None)
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
assert result[0].idle_since is None
@pytest.mark.asyncio
async def test_update_handler_status_idle_since_unset_preserves() -> None:
store = MemoryWorkflowStore()
now = datetime.now(timezone.utc)
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
idle_since=now,
)
)
# Not passing idle_since at all should preserve the existing value
await store.update_handler_status("run-1", status="running")
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
assert result[0].idle_since == now
@pytest.mark.asyncio
async def test_update_handler_status_non_terminal_does_not_set_completed_at() -> None:
store = MemoryWorkflowStore()
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
)
)
# Update status to "running" (non-terminal) should not set completed_at
await store.update_handler_status("run-1", status="running")
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
assert result[0].completed_at is None
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_status", ["completed", "failed", "cancelled"])
async def test_update_handler_status_terminal_sets_completed_at(
terminal_status: Status,
) -> None:
store = MemoryWorkflowStore()
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="wf",
status="running",
run_id="run-1",
)
)
await store.update_handler_status("run-1", status=terminal_status)
result = await store.query(HandlerQuery(run_id_in=["run-1"]))
assert result[0].completed_at is not None
def _make_stored_event(event: Event, run_id: str = "run-1") -> StoredEvent:
return StoredEvent(
run_id=run_id,
sequence=0,
timestamp=datetime.now(timezone.utc),
event=EventEnvelopeWithMetadata.from_event(event),
)
def test_is_terminal_event_stop_event() -> None:
stored = _make_stored_event(StopEvent(result="done"))
assert AbstractWorkflowStore._is_terminal_event(stored) is True
def test_is_terminal_event_regular_event() -> None:
stored = _make_stored_event(Event())
assert AbstractWorkflowStore._is_terminal_event(stored) is False
def test_is_terminal_event_workflow_failed_event() -> None:
# WorkflowFailedEvent extends StopEvent, so it should be terminal
event = WorkflowFailedEvent(
step_name="my_step",
exception_type="ValueError",
exception_message="bad value",
traceback="",
attempts=1,
elapsed_seconds=0.1,
)
stored = _make_stored_event(event)
assert AbstractWorkflowStore._is_terminal_event(stored) is True
def test_is_terminal_event_workflow_cancelled_event() -> None:
# WorkflowCancelledEvent extends StopEvent, so it should be terminal
stored = _make_stored_event(WorkflowCancelledEvent())
assert AbstractWorkflowStore._is_terminal_event(stored) is True
@@ -2,7 +2,7 @@ from __future__ import annotations
from typing import Any, cast
from llama_agents.server._store.abstract_workflow_store import PersistentHandler
from llama_agents.server import PersistentHandler
from workflows.events import StopEvent
@@ -0,0 +1,174 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""Tests for the base runtime decorator forwarding classes."""
from __future__ import annotations
from typing import Any, AsyncGenerator
from unittest.mock import MagicMock
from llama_agents.server._runtime.runtime_decorators import (
BaseExternalRunAdapterDecorator,
BaseInternalRunAdapterDecorator,
BaseRuntimeDecorator,
)
from workflows.context.state_store import StateStore
from workflows.events import (
Event,
StopEvent,
)
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
RegisteredWorkflow,
Runtime,
WaitResult,
WaitResultTimeout,
)
from workflows.runtime.types.ticks import WorkflowTick
# -- Stubs -----------------------------------------------------------------
class StubInternalAdapter(InternalRunAdapter):
def __init__(self) -> None:
self.closed = False
@property
def run_id(self) -> str:
return "r1"
async def write_to_event_stream(self, event: Event) -> None:
pass
async def get_now(self) -> float:
return 1.0
async def send_event(self, tick: WorkflowTick) -> None:
pass
async def wait_receive(self, timeout_seconds: float | None = None) -> WaitResult:
return WaitResultTimeout()
async def close(self) -> None:
self.closed = True
def get_state_store(self) -> StateStore[Any] | None:
return None
class StubExternalAdapter(ExternalRunAdapter):
def __init__(self) -> None:
self.closed = False
@property
def run_id(self) -> str:
return "r1"
async def send_event(self, tick: WorkflowTick) -> None:
pass
async def stream_published_events(self) -> AsyncGenerator[Event, None]:
yield StopEvent(result="done")
async def close(self) -> None:
self.closed = True
async def get_result(self) -> StopEvent:
return StopEvent(result="done")
def get_state_store(self) -> StateStore[Any] | None:
return None
class StubRuntime(Runtime):
def __init__(self) -> None:
super().__init__()
self.launched = False
def register(self, workflow: Any) -> RegisteredWorkflow:
return RegisteredWorkflow(
workflow=workflow, workflow_run_fn=MagicMock(), steps={}
)
def run_workflow(
self,
run_id: str,
workflow: Any,
init_state: Any,
start_event: Any = None,
serialized_state: dict[str, Any] | None = None,
serializer: Any = None,
) -> ExternalRunAdapter:
return StubExternalAdapter()
def get_internal_adapter(self, workflow: Any) -> InternalRunAdapter:
return StubInternalAdapter()
def get_external_adapter(self, run_id: str) -> ExternalRunAdapter:
return StubExternalAdapter()
def launch(self) -> None:
self.launched = True
def destroy(self) -> None:
pass
# -- Tests -----------------------------------------------------------------
def test_runtime_decorator_forwards() -> None:
inner = StubRuntime()
dec = BaseRuntimeDecorator(inner)
dec.launch()
assert inner.launched
async def test_internal_adapter_decorator_forwards() -> None:
inner = StubInternalAdapter()
dec = BaseInternalRunAdapterDecorator(inner)
assert dec.run_id == "r1"
assert await dec.get_now() == 1.0
await dec.close()
assert inner.closed
async def test_external_adapter_decorator_forwards() -> None:
inner = StubExternalAdapter()
dec = BaseExternalRunAdapterDecorator(inner)
assert dec.run_id == "r1"
result = await dec.get_result()
assert result.result == "done"
await dec.close()
assert inner.closed
async def test_subclass_can_override_selectively() -> None:
"""Override one method; the rest still forward."""
class Custom(BaseInternalRunAdapterDecorator):
async def get_now(self) -> float:
return 42.0
inner = StubInternalAdapter()
dec = Custom(inner)
assert await dec.get_now() == 42.0
assert dec.run_id == "r1" # still forwarded
def test_runtime_decorator_forwards_untrack() -> None:
from workflows import Workflow, step
from workflows.events import StartEvent
class SimpleWorkflow(Workflow):
@step
async def start(self, ev: StartEvent) -> StopEvent:
return StopEvent(result="done")
inner = StubRuntime()
dec = BaseRuntimeDecorator(inner)
wf = SimpleWorkflow(runtime=dec)
assert wf in dec._pending
dec.untrack_workflow(wf)
assert wf not in dec._pending
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
from typing import Any
from unittest.mock import AsyncMock, Mock, patch
@@ -10,13 +12,6 @@ from starlette.middleware import Middleware
from workflows.workflow import Workflow
def test_init() -> None:
server = WorkflowServer()
assert len(server.app.user_middleware) == 1
assert server._service._workflows == {}
assert server._service._handlers == {}
def test_init_custom_middleware() -> None:
custom_middleware = [Mock(spec=Middleware)]
server = WorkflowServer(middleware=custom_middleware) # type: ignore
@@ -26,8 +21,8 @@ def test_init_custom_middleware() -> None:
def test_add_workflow(simple_test_workflow: Workflow) -> None:
server = WorkflowServer()
server.add_workflow("test", simple_test_workflow)
assert "test" in server._service._workflows
assert server._service._workflows["test"] == simple_test_workflow
assert "test" in server.get_workflows()
assert server.get_workflows()["test"] == simple_test_workflow
@pytest.mark.asyncio
@@ -68,7 +63,7 @@ def test_extract_workflow_success(simple_test_workflow: Workflow) -> None:
mock_request = Mock()
mock_request.path_params = {"name": "test"}
assert server._api._extract_workflow(mock_request).workflow == simple_test_workflow
assert server._api._extract_workflow(mock_request) is simple_test_workflow
def test_extract_workflow_missing_name() -> None:
@@ -8,18 +8,17 @@ import json
from collections import Counter
from contextlib import asynccontextmanager
from datetime import datetime
from types import SimpleNamespace
from typing import Any, AsyncGenerator, AsyncIterator
import pytest
import pytest_asyncio
from httpx import ASGITransport, AsyncClient, Response
from llama_agents.server import WorkflowServer
from llama_agents.server._store.abstract_workflow_store import (
from llama_agents.server import (
HandlerQuery,
MemoryWorkflowStore,
PersistentHandler,
WorkflowServer,
)
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from llama_index_instrumentation.dispatcher import active_instrument_tags
from server_test_fixtures import (
ExternalEvent, # type: ignore[import]
@@ -62,7 +61,7 @@ def server(
interactive_workflow: Workflow,
) -> WorkflowServer:
# Use MemoryWorkflowStore so get_handlers() can retrieve from persistence
server = WorkflowServer(workflow_store=MemoryWorkflowStore())
server = WorkflowServer(workflow_store=MemoryWorkflowStore(), idle_timeout=0.01)
server.add_workflow("test", simple_test_workflow)
server.add_workflow("error", error_workflow)
server.add_workflow("streaming", streaming_workflow)
@@ -92,7 +91,7 @@ async def server_with_persisted_handlers(
for handler in persisted_handlers:
await store.update(handler)
server_with_store = WorkflowServer(workflow_store=store)
server_with_store = WorkflowServer(workflow_store=store, idle_timeout=0.01)
server_with_store.add_workflow("interactive", interactive_workflow)
async with server_with_store.contextmanager():
@@ -170,9 +169,6 @@ async def test_health_check(client: AsyncClient) -> None:
assert response.status_code == 200
data = response.json()
assert data["status"] == "healthy"
assert data["loaded_workflows"] == 0
assert data["active_workflows"] == 0
assert data["idle_workflows"] == 0
@pytest.mark.asyncio
@@ -430,18 +426,18 @@ async def test_stream_events_success(client: AsyncClient) -> None:
data = response.json()
handler_id = data["handler_id"]
# Stream events
response = await client.get(f"/events/{handler_id}")
# Stream events (after_sequence=-1 to get all from beginning)
response = await client.get(f"/events/{handler_id}?after_sequence=-1")
assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
# Collect streamed events
events: list[dict[str, Any]] = []
async for line in response.aiter_lines():
if line.strip():
line = line.strip()
if line.startswith("data: "):
event_data = json.loads(line.removeprefix("data: "))
assert isinstance(event_data, dict)
# Filter out empty events
if event_data:
events.append(event_data)
@@ -466,8 +462,8 @@ async def test_stream_events_sse(client: AsyncClient) -> None:
data = response.json()
handler_id = data["handler_id"]
# Stream events in SSE format
response = await client.get(f"/events/{handler_id}?sse=true")
# Stream events in SSE format (after_sequence=-1 to get all from beginning)
response = await client.get(f"/events/{handler_id}?sse=true&after_sequence=-1")
assert response.status_code == 200
assert response.headers["content-type"].startswith("text/event-stream")
@@ -504,8 +500,8 @@ async def test_stream_events_sse(client: AsyncClient) -> None:
assert event["data"]["value"]["message"] == f"event_{i}"
assert event["data"]["value"]["sequence"] == i
# stream completed
response = await client.get(f"/events/{handler_id}?sse=true")
# reconnect with after_sequence beyond last event returns 204
response = await client.get(f"/events/{handler_id}?sse=true&after_sequence=999999")
assert response.status_code == 204
@@ -517,46 +513,41 @@ async def test_stream_events_not_found(client: AsyncClient) -> None:
@pytest.mark.asyncio
async def test_stream_events_single_consumer(client: AsyncClient) -> None:
"""Test that the consumer lock mechanism works with acquire_timeout."""
# Start a streaming workflow that completes quickly
handler_response = await client.post("/workflows/interactive/run-nowait", json={})
handler_id = handler_response.json()["handler_id"]
# send 2 simultaneous requests
a = asyncio.create_task(
client.send(
client.build_request("GET", f"/events/{handler_id}?acquire_timeout=0.01"),
stream=True,
)
)
b = asyncio.create_task(
client.send(
client.build_request("GET", f"/events/{handler_id}?acquire_timeout=0.01"),
stream=True,
)
async def test_stream_events_multiple_consumers(client: AsyncClient) -> None:
"""Multiple concurrent consumers can stream the same handler's events."""
# Start a streaming workflow
response = await client.post(
"/workflows/streaming/run-nowait", json={"kwargs": {"count": 2}}
)
handler_id = response.json()["handler_id"]
# wait for one to be rejected
done, pending = await asyncio.wait({a, b}, return_when=asyncio.FIRST_COMPLETED)
# Two concurrent stream requests (after_sequence=-1 to get all from beginning)
a = asyncio.create_task(client.get(f"/events/{handler_id}?after_sequence=-1"))
b = asyncio.create_task(client.get(f"/events/{handler_id}?after_sequence=-1"))
# Assert that the done request got a 409 response
assert len(done) == 1
done_response = list(done)[0].result()
assert done_response.status_code == 409
response_a, response_b = await asyncio.gather(a, b)
# Send an ExternalEvent to complete the workflow
send_response = await client.post(
f"/events/{handler_id}",
json={
"event": JsonSerializer().serialize(ExternalEvent(response="test-response"))
},
)
assert send_response.status_code == 200
assert response_a.status_code == 200
assert response_b.status_code == 200
# Wait for the pending response and stream it
pending_response = await list(pending)[0]
assert pending_response.status_code == 200
# Both consumers should receive the same events
def parse_events(text: str) -> list[dict[str, Any]]:
events = []
for line in text.strip().split("\n"):
line = line.strip()
if line.startswith("data: "):
data = json.loads(line.removeprefix("data: "))
if data:
events.append(data)
return events
events_a = parse_events(response_a.text)
events_b = parse_events(response_b.text)
# Both should have the same event types
types_a = [e["type"] for e in events_a]
types_b = [e["type"] for e in events_b]
assert types_a == types_b
@pytest.mark.asyncio
@@ -572,15 +563,16 @@ async def test_stream_events_no_events_default_hides_internal(
data = response.json()
handler_id = data["handler_id"]
# Stream without include_internal
response = await client.get(f"/events/{handler_id}")
# Stream without include_internal (after_sequence=-1 to get all from beginning)
response = await client.get(f"/events/{handler_id}?after_sequence=-1")
assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
# Collect events
events = []
async for line in response.aiter_lines():
if line.strip():
line = line.strip()
if line.startswith("data: "):
event_data = json.loads(line.removeprefix("data: "))
if event_data:
events.append(event_data)
@@ -603,15 +595,18 @@ async def test_stream_events_include_internal_true(client: AsyncClient) -> None:
data = response.json()
handler_id = data["handler_id"]
# Stream with include_internal=true
response = await client.get(f"/events/{handler_id}?include_internal=true")
# Stream with include_internal=true (after_sequence=-1 to get all from beginning)
response = await client.get(
f"/events/{handler_id}?include_internal=true&after_sequence=-1"
)
assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
# Collect events
events = []
async for line in response.aiter_lines():
if line.strip():
line = line.strip()
if line.startswith("data: "):
event_data = json.loads(line.removeprefix("data: "))
if event_data:
events.append(event_data)
@@ -776,14 +771,12 @@ async def test_get_handlers_filters_status_and_workflow_name(
# Seed persistence with mixed handlers
persisted = [
PersistentHandler(
handler_id="h1", workflow_name="interactive", status="running", ctx={}
handler_id="h1", workflow_name="interactive", status="running"
),
PersistentHandler(
handler_id="h2", workflow_name="interactive", status="completed", ctx={}
),
PersistentHandler(
handler_id="h3", workflow_name="other", status="failed", ctx={}
handler_id="h2", workflow_name="interactive", status="completed"
),
PersistentHandler(handler_id="h3", workflow_name="other", status="failed"),
]
async with server_with_persisted_handlers(
@@ -814,13 +807,13 @@ async def test_get_handlers_filters_multiple_status_params(
) -> None:
persisted = [
PersistentHandler(
handler_id="ha", workflow_name="interactive", status="completed", ctx={}
handler_id="ha", workflow_name="interactive", status="completed"
),
PersistentHandler(
handler_id="hb", workflow_name="interactive", status="failed", ctx={}
handler_id="hb", workflow_name="interactive", status="failed"
),
PersistentHandler(
handler_id="hc", workflow_name="interactive", status="running", ctx={}
handler_id="hc", workflow_name="interactive", status="running"
),
]
@@ -985,29 +978,6 @@ async def test_post_event_invalid_event_data(client: AsyncClient) -> None:
assert "Failed to deserialize event" in response.text
@pytest.mark.asyncio
async def test_post_event_context_not_available(
client: AsyncClient, server: WorkflowServer
) -> None:
# Dumb test for code coverage. Inject a dummy handler with no context to trigger 500 path
wrapper = SimpleNamespace(
run_handler=SimpleNamespace(done=lambda: False, ctx=None),
workflow_name="test",
status="running",
mark_active=lambda: None,
)
handler_id = "noctx-1"
server._service._handlers[handler_id] = wrapper # type: ignore[assignment]
try:
response = await client.post(f"/events/{handler_id}", json={"event": "{}"})
assert response.status_code == 500
assert "Context not available" in response.text
finally:
server._service._handlers.pop(handler_id, None)
@pytest.mark.asyncio
async def test_post_event_body_parsing_error(client: AsyncClient) -> None:
# Start interactive workflow which waits for an event (keeps running)
@@ -1125,7 +1095,6 @@ async def test_delete_persisted_handler_removes_from_store(
handler_id="persist-only",
workflow_name="interactive",
status="completed",
ctx={},
)
],
) as (_server, client, store):
@@ -1152,7 +1121,6 @@ async def test_stop_only_persisted_handler_without_removal_returns_not_found(
handler_id="store-only",
workflow_name="interactive",
status="completed",
ctx={},
)
],
) as (_server, client, store):
@@ -1208,9 +1176,9 @@ async def test_stream_events_after_completion_should_return_unconsumed_events(
await wait_for_passing(_wait_done)
# Now fetch events AFTER completion. Expect the unconsumed events to still be retrievable.
# Use NDJSON for easier parsing.
resp = await client.get(f"/events/{handler_id}?sse=false")
# Now fetch events AFTER completion. Expect all events to be retrievable.
# Use NDJSON for easier parsing. after_sequence=-1 to get all from beginning.
resp = await client.get(f"/events/{handler_id}?sse=false&after_sequence=-1")
assert resp.status_code == 200
assert resp.headers["content-type"].startswith("application/x-ndjson")
@@ -1224,6 +1192,143 @@ async def test_stream_events_after_completion_should_return_unconsumed_events(
assert len(lines) == 4
@pytest.mark.asyncio
async def test_stream_events_sse_includes_id_field(client: AsyncClient) -> None:
"""SSE events include an id: field with the event sequence number."""
response = await client.post(
"/workflows/streaming/run-nowait", json={"kwargs": {"count": 2}}
)
handler_id = response.json()["handler_id"]
response = await client.get(f"/events/{handler_id}?sse=true&after_sequence=-1")
assert response.status_code == 200
# Parse raw SSE frames and extract id fields
ids: list[int] = []
for line in response.text.strip().split("\n"):
line = line.strip()
if line.startswith("id: "):
ids.append(int(line.removeprefix("id: ")))
# Every SSE event should have an id
assert len(ids) >= 2
# Ids should be monotonically increasing
assert ids == sorted(ids)
@pytest.mark.asyncio
async def test_stream_events_last_event_id_header(client: AsyncClient) -> None:
"""SSE Last-Event-ID header takes priority over after_sequence query param."""
response = await client.post(
"/workflows/streaming/run-nowait", json={"kwargs": {"count": 3}}
)
handler_id = response.json()["handler_id"]
# First, stream all events to get the sequence numbers
response = await client.get(f"/events/{handler_id}?sse=true&after_sequence=-1")
assert response.status_code == 200
ids: list[int] = []
for line in response.text.strip().split("\n"):
line = line.strip()
if line.startswith("id: "):
ids.append(int(line.removeprefix("id: ")))
assert len(ids) >= 3
# Reconnect with Last-Event-ID header set to skip past all events.
# The query param says after_sequence=-1 (from beginning), but the header
# should override it.
response = await client.get(
f"/events/{handler_id}?sse=true&after_sequence=-1",
headers={"last-event-id": str(ids[-1])},
)
# Should get 204 because Last-Event-ID is past the last event and the run
# is complete.
assert response.status_code == 204
@pytest.mark.asyncio
async def test_stream_events_after_sequence_now(client: AsyncClient) -> None:
"""after_sequence=now skips historical events, only receives new ones."""
# Start a streaming workflow
response = await client.post(
"/workflows/streaming/run-nowait", json={"kwargs": {"count": 3}}
)
handler_id = response.json()["handler_id"]
# Wait for completion so all events are stored
async def _wait_done() -> None:
r = await client.get(f"/handlers/{handler_id}")
if r.status_code != 200:
raise AssertionError("not done")
await wait_for_passing(_wait_done)
# Now request with after_sequence=now. Since the workflow is already complete,
# "now" resolves to the last sequence, and there are no remaining events, so
# we should get 204.
response = await client.get(f"/events/{handler_id}?after_sequence=now")
assert response.status_code == 204
@pytest.mark.asyncio
async def test_stream_events_after_sequence_now_receives_future_events(
interactive_workflow: Workflow,
) -> None:
"""after_sequence=now on a running workflow receives only events appended after the request."""
async with server_with_persisted_handlers(interactive_workflow) as (
_server,
client,
store,
):
# Start the interactive workflow (it waits for an external event)
start_resp = await client.post("/workflows/interactive/run-nowait", json={})
handler_id = start_resp.json()["handler_id"]
# Wait until some events are stored (at least the internal dispatch)
await asyncio.sleep(0.1)
# Count events currently in the store
found = await store.query(HandlerQuery(handler_id_in=[handler_id]))
run_id = found[0].run_id
assert run_id is not None
events_before = await store.query_events(run_id)
assert len(events_before) > 0
# Start streaming with after_sequence=now — should skip all existing events
stream_task = asyncio.create_task(
client.get(f"/events/{handler_id}?sse=false&after_sequence=now")
)
# Give the streaming request time to start
await asyncio.sleep(0.05)
# Send an external event to progress the workflow
serializer = JsonSerializer()
event = ExternalEvent(response="after-now")
event_str = serializer.serialize(event)
await client.post(f"/events/{handler_id}", json={"event": event_str})
response = await stream_task
assert response.status_code == 200
# Parse NDJSON lines
events = []
for line in response.text.strip().split("\n"):
line = line.strip()
if line:
events.append(json.loads(line))
# The events should include at minimum the StopEvent from completion.
# They should NOT include any of the events that existed before "now".
event_types = [e["type"] for e in events]
assert "StopEvent" in event_types
# Verify we got fewer events than the total stored — the historical ones were skipped
all_events = await store.query_events(run_id)
assert len(events) < len(all_events)
@pytest.mark.asyncio
async def test_instrument_tags_contains_handler_id_in_server_context() -> None:
seen_handler_id: dict[str, str | None] = {"handler_id": None}
@@ -1236,7 +1341,7 @@ async def test_instrument_tags_contains_handler_id_in_server_context() -> None:
seen_handler_id["handler_id"] = hid
return StopEvent()
server = WorkflowServer(workflow_store=MemoryWorkflowStore())
server = WorkflowServer(workflow_store=MemoryWorkflowStore(), idle_timeout=0.01)
server.add_workflow("tags", TagReadingWorkflow())
async with server.contextmanager():
@@ -1258,17 +1363,3 @@ async def test_instrument_tags_contains_handler_id_in_server_context() -> None:
assert data["status"] == "completed"
assert seen_handler_id["handler_id"] is not None
assert seen_handler_id["handler_id"] == handler_id
@pytest.mark.asyncio
async def test_run_sync_removes_handler_even_with_unconsumed_events(
client: AsyncClient, server: WorkflowServer
) -> None:
# Run a streaming workflow synchronously; it emits user events but we don't consume them here.
resp = await client.post("/workflows/streaming/run", json={"kwargs": {"count": 2}})
assert resp.status_code == 200
data = resp.json()
assert data["status"] == "completed"
# The synchronous run path should clean up the handler from memory even if events remain
assert len(server._service._handlers) == 0
@@ -1,416 +0,0 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
import asyncio
from typing import AsyncGenerator
import pytest
from httpx import ASGITransport, AsyncClient
from llama_agents.server import WorkflowServer
from llama_agents.server._store.abstract_workflow_store import (
HandlerQuery,
PersistentHandler,
)
from llama_agents.server._store.memory_workflow_store import MemoryWorkflowStore
from server_test_fixtures import ( # type: ignore[import]
ExternalEvent,
RequestedExternalEvent,
wait_for_passing, # type: ignore[import]
)
from workflows.context.context_types import SerializedContext
from workflows.events import Event, InternalDispatchEvent, StopEvent
from workflows.workflow import Workflow
@pytest.fixture
def memory_store() -> MemoryWorkflowStore:
return MemoryWorkflowStore()
@pytest.fixture
async def server_with_store(
memory_store: MemoryWorkflowStore, interactive_workflow: Workflow
) -> AsyncGenerator[WorkflowServer, None]:
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("test", interactive_workflow)
async with server.contextmanager():
yield server
@pytest.fixture
async def server_with_store_and_simple_workflow(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> AsyncGenerator[WorkflowServer, None]:
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("test", simple_test_workflow)
async with server.contextmanager():
yield server
@pytest.mark.asyncio
async def test_store_is_updated_on_step_completion(
server_with_store: WorkflowServer, memory_store: MemoryWorkflowStore
) -> None:
server = server_with_store
# Start a workflow through internal runner to exercise persistence updates
handler_id = "persist-1"
handler = server._service._workflows["test"].run()
await server._service.run_workflow_handler(handler_id, "test", handler)
handler = server._service._handlers[handler_id]
# wait for first step to complete
async def get_non_internal_event() -> Event:
item = await server._service._handlers[handler_id].queue.get()
if isinstance(item, InternalDispatchEvent):
raise ValueError("Internal event received. Try again")
return item
item = await wait_for_passing(get_non_internal_event)
assert isinstance(item, RequestedExternalEvent)
# much sure its stored and running
persistent_list = await memory_store.query(HandlerQuery(handler_id_in=[handler_id]))
assert persistent_list
persistent = persistent_list[0]
assert persistent.workflow_name == "test"
assert persistent.status == "running"
assert isinstance(persistent.ctx, dict)
# now, validate that the workflow completes when responding
handler = server._service._handlers[handler_id].run_handler
ctx = handler.ctx
assert ctx is not None
ctx.send_event(ExternalEvent(response="pong"))
result = await handler
# wait for event loop to resolve all tasks
task = server._service._handlers[handler_id].task
assert task is not None
await task
await asyncio.sleep(0) # let even loop resolve other waiters on the internal
assert result == "received: pong"
updated = memory_store.handlers[handler_id]
assert updated.status == "completed"
@pytest.mark.asyncio
async def test_resume_active_handlers_across_server_restart(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
# Seed the store with a valid serialized Context using public API
handler_id = "resume-1"
initial_ctx = SerializedContext().model_dump(mode="python")
await memory_store.update(
PersistentHandler(
handler_id=handler_id,
workflow_name="test",
status="running",
ctx=initial_ctx,
)
)
# Second server: same store and workflow, explicitly initialize active handlers
server2 = WorkflowServer(workflow_store=memory_store)
server2.add_workflow("test", simple_test_workflow)
async with server2.contextmanager(): # start and stop it
# The handler should be registered under the same id
assert handler_id in server2._service._handlers
# Await its completion through internal result future
result = await server2._service._handlers[handler_id].run_handler
assert result == "processed: default"
@pytest.mark.asyncio
async def test_startup_marks_invalid_persisted_context_as_failed(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""Server should not crash on invalid persisted context; it should mark it failed."""
# Seed an invalid context payload that will fail Context.from_dict
# Make the context structurally valid but with an invalid streaming_queue JSON
invalid_ctx = {
"state": {},
"streaming_queue": "[]",
"queues": {"process": "not-deserializable-as-a-queue"},
"event_buffers": {},
"in_progress": {},
"accepted_events": [],
"broker_log": [],
"is_running": True,
"waiting_ids": [],
}
handler_id = "bad-ctx-1"
await memory_store.update(
PersistentHandler(
handler_id=handler_id,
workflow_name="test",
status="running",
ctx=invalid_ctx,
)
)
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("test", simple_test_workflow)
async with server.contextmanager():
# Invalid handler should not be registered
assert handler_id not in server._service._handlers
# After startup attempt, it should be marked as failed in the store
persisted = memory_store.handlers[handler_id]
assert persisted.status == "failed"
@pytest.mark.asyncio
async def test_store_is_updated_on_workflow_failure(
memory_store: MemoryWorkflowStore, error_workflow: Workflow
) -> None:
# Build a server with a failing workflow and the in-memory store
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("error", error_workflow)
async with server.contextmanager():
# Start a workflow through internal runner to exercise persistence updates
handler_id = "fail-1"
handler = server._service._workflows["error"].run()
await server._service.run_workflow_handler(handler_id, "error", handler)
# Await the failure of the handler itself
with pytest.raises(ValueError, match="Test error"):
await handler
# Ensure the background streaming task has completed and persisted status
task = server._service._handlers[handler_id].task
assert task is not None
await task
await asyncio.sleep(0)
# Verify store captured the failed status and has a context snapshot
persistent_list = await memory_store.query(
HandlerQuery(handler_id_in=[handler_id])
)
assert persistent_list
persistent = persistent_list[0]
assert persistent.workflow_name == "error"
assert persistent.status == "failed"
assert isinstance(persistent.ctx, dict)
@pytest.mark.asyncio
async def test_startup_ignores_unregistered_workflows(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
# Unknown workflow entry
await memory_store.update(
PersistentHandler(
handler_id="unknown-1", workflow_name="unknown", status="running", ctx={}
)
)
# Known workflow entry to be resumed
await memory_store.update(
PersistentHandler(
handler_id="known-1",
workflow_name="test",
status="running",
ctx=SerializedContext().model_dump(mode="python"),
)
)
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("test", simple_test_workflow)
async with server.contextmanager(): # start and stop it
assert "unknown-1" not in server._service._handlers
assert "known-1" in server._service._handlers
# Await completion of the resumed known handler
result = await server._service._handlers["known-1"].run_handler
assert result == "processed: default"
def patch_store_update_to_fail(
monkeypatch: pytest.MonkeyPatch, store: MemoryWorkflowStore, fail_count: int
) -> dict[str, int]:
"""Monkeypatch `store.update` to fail `fail_count` times, then delegate to original.
Returns a dict with a mutable `count` key to inspect attempt count.
"""
attempts: dict[str, int] = {"count": 0}
original_update = store.update
async def wrapped(handler: PersistentHandler) -> None:
attempts["count"] += 1
if attempts["count"] <= fail_count:
raise Exception(
f"Simulated persistence failure (attempt {attempts['count']})"
)
await original_update(handler)
monkeypatch.setattr(store, "update", wrapped)
return attempts
@pytest.mark.asyncio
async def test_persistence_retries_on_failure(
simple_test_workflow: Workflow, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Test that persistence operations are retried according to backoff configuration."""
# Create a store that fails twice then succeeds
store = MemoryWorkflowStore()
attempts = patch_store_update_to_fail(monkeypatch, store, fail_count=2)
# Configure server with custom backoff (shorter for testing)
server = WorkflowServer(
workflow_store=store,
persistence_backoff=[0.0, 0.0], # Very short backoffs for testing
)
server.add_workflow("test", simple_test_workflow)
async with server.contextmanager():
# Start a workflow to trigger persistence
handler_id = "retry-test"
handler = server._service._workflows["test"].run()
await server._service.run_workflow_handler(handler_id, "test", handler)
# Wait for workflow completion
result = await server._service._handlers[handler_id].run_handler
assert result == "processed: default"
# Wait for background streaming task to complete, ignoring its expected exception
task = server._service._handlers[handler_id].task
try:
assert task is not None
await task
except Exception:
pass
await asyncio.sleep(0)
# Verify that retries occurred and eventually succeeded
# There can be multiple checkpoints (initial running, step completion, final completion)
# so attempts will be >= initial + retries
assert attempts["count"] >= 3 # Initial + 2 retries
persistent_list = await store.query(HandlerQuery(handler_id_in=[handler_id]))
assert persistent_list
persistent = persistent_list[0]
assert persistent.status == "completed"
@pytest.mark.asyncio
async def test_workflow_cancelled_after_all_retries_fail(
streaming_workflow: Workflow, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Test that workflow is cancelled when all persistence retries are exhausted."""
# Create a store that always fails
store = MemoryWorkflowStore()
attempts = patch_store_update_to_fail(
monkeypatch, store, fail_count=10
) # Fail more than backoff attempts
# Configure server with custom backoff (shorter for testing)
server = WorkflowServer(
workflow_store=store,
persistence_backoff=[0.0, 0.0], # 2 retries, very short backoffs
)
server.add_workflow("test", streaming_workflow)
async with server.contextmanager():
# Start a workflow to trigger persistence
handler_id = "cancel-test"
handler = server._service._workflows["test"].run()
# Should raise HTTPException if persistence fails on initial checkpoint
with pytest.raises(Exception):
await server._service.run_workflow_handler(handler_id, "test", handler)
# Verify retry attempts and no registration/persistence on failure
assert attempts["count"] == 3 # Initial + 2 retries
assert handler_id not in server._service._handlers
persistent_list = await store.query(HandlerQuery(handler_id_in=[handler_id]))
assert not persistent_list
@pytest.mark.asyncio
async def test_resume_across_runs(
memory_store: MemoryWorkflowStore, cumulative_workflow: Workflow
) -> None:
"""Test that workflow context accumulates data across multiple runs using handler_id continuation."""
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("cumulative", cumulative_workflow)
async with server.contextmanager():
transport = ASGITransport(app=server.app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
# First run - should start with count=0, increment by 5
response = await client.post(
"/workflows/cumulative/run", json={"start_event": {"increment": 5}}
)
assert response.status_code == 200
resp_data = response.json()
assert resp_data["result"]["value"]["result"] == "count: 5, runs: 1"
# Get the handler id for that run
handler_id = resp_data["handler_id"]
# Wait for the handler to be fully persisted as completed
await asyncio.sleep(0.1)
# Verify it's persisted in the store as completed
persisted_list = await memory_store.query(
HandlerQuery(handler_id_in=[handler_id])
)
assert persisted_list
persisted = persisted_list[0]
assert persisted.status == "completed"
# Second run - should start with count=5, increment by 3
response2 = await client.post(
"/workflows/cumulative/run",
json={"handler_id": handler_id, "start_event": {"increment": 3}},
)
assert response2.status_code == 200
resp_data2 = response2.json()
assert resp_data2["result"]["value"]["result"] == "count: 8, runs: 2"
# Verify the handler id is the same
assert resp_data2["handler_id"] == handler_id
# Wait for the handler to be fully persisted as completed
await asyncio.sleep(0.1)
# Verify memory store has only one handler
assert len(memory_store.handlers) == 1
assert memory_store.handlers[handler_id].status == "completed"
@pytest.mark.asyncio
async def test_result_for_completed_persisted_handler_without_runtime_registration(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""A completed handler persisted in the store but not registered in memory should still return its result."""
handler_id = "store-completed-1"
# Seed a completed handler directly in the store (server won't load completed handlers at startup)
await memory_store.update(
PersistentHandler(
handler_id=handler_id,
workflow_name="test",
status="completed",
result=StopEvent(result="processed: default"),
ctx=SerializedContext().model_dump(mode="python"),
)
)
# Start a server with the same store and workflow registered
server = WorkflowServer(workflow_store=memory_store)
server.add_workflow("test", simple_test_workflow)
async with server.contextmanager():
# Ensure the handler is not registered in runtime memory
assert handler_id not in server._service._handlers
# But the API should still return the persisted result
transport = ASGITransport(app=server.app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.get(f"/handlers/{handler_id}")
assert response.status_code == 200
data = response.json()
assert data["status"] == "completed"
assert data["result"]["value"]["result"] == "processed: default"
@@ -0,0 +1,393 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
"""Tests for ServerRuntimeDecorator and _ServerInternalRunAdapter."""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any, AsyncGenerator
from unittest.mock import MagicMock
import pytest
from llama_agents.server import (
HandlerQuery,
MemoryWorkflowStore,
PersistentHandler,
WorkflowServer,
)
from llama_agents.server._runtime.idle_release_runtime import IdleReleaseDecorator
from llama_agents.server._runtime.runtime_decorators import BaseRuntimeDecorator
from llama_agents.server._runtime.server_runtime import (
ServerRuntimeDecorator,
_ServerInternalRunAdapter,
)
from workflows import Workflow, step
from workflows.context.state_store import StateStore
from workflows.events import (
Event,
StartEvent,
StopEvent,
WorkflowCancelledEvent,
WorkflowFailedEvent,
WorkflowTimedOutEvent,
)
from workflows.runtime.types.plugin import (
ExternalRunAdapter,
InternalRunAdapter,
RegisteredWorkflow,
Runtime,
WaitResult,
WaitResultTimeout,
)
from workflows.runtime.types.ticks import WorkflowTick
# -- Stubs -----------------------------------------------------------------
class StubInternalAdapter(InternalRunAdapter):
def __init__(self) -> None:
self.closed = False
@property
def run_id(self) -> str:
return "r1"
async def write_to_event_stream(self, event: Event) -> None:
pass
async def get_now(self) -> float:
return 1.0
async def send_event(self, tick: WorkflowTick) -> None:
pass
async def wait_receive(self, timeout_seconds: float | None = None) -> WaitResult:
return WaitResultTimeout()
async def close(self) -> None:
self.closed = True
def get_state_store(self) -> StateStore[Any] | None:
return None
class StubExternalAdapter(ExternalRunAdapter):
def __init__(self) -> None:
self.closed = False
@property
def run_id(self) -> str:
return "r1"
async def send_event(self, tick: WorkflowTick) -> None:
pass
async def stream_published_events(self) -> AsyncGenerator[Event, None]:
yield StopEvent(result="done")
async def close(self) -> None:
self.closed = True
async def get_result(self) -> StopEvent:
return StopEvent(result="done")
def get_state_store(self) -> StateStore[Any] | None:
return None
class StubRuntime(Runtime):
def __init__(self) -> None:
super().__init__()
self.launched = False
def register(self, workflow: Any) -> RegisteredWorkflow:
return RegisteredWorkflow(
workflow=workflow, workflow_run_fn=MagicMock(), steps={}
)
def run_workflow(
self,
run_id: str,
workflow: Any,
init_state: Any,
start_event: Any = None,
serialized_state: dict[str, Any] | None = None,
serializer: Any = None,
) -> ExternalRunAdapter:
return StubExternalAdapter()
def get_internal_adapter(self, workflow: Any) -> InternalRunAdapter:
return StubInternalAdapter()
def get_external_adapter(self, run_id: str) -> ExternalRunAdapter:
return StubExternalAdapter()
def launch(self) -> None:
self.launched = True
def destroy(self) -> None:
pass
class SimpleWorkflow(Workflow):
@step
async def start(self, ev: StartEvent) -> StopEvent:
return StopEvent(result="done")
# -- Tests -----------------------------------------------------------------
def test_add_workflow_sets_workflow_name() -> None:
server = WorkflowServer()
wf = SimpleWorkflow(runtime=StubRuntime())
server.add_workflow("greeting", wf)
assert wf.workflow_name == "greeting"
def test_add_workflow_wraps_runtime_with_decorator() -> None:
server = WorkflowServer()
wf = SimpleWorkflow(runtime=StubRuntime())
server.add_workflow("greeting", wf)
assert isinstance(wf.runtime, BaseRuntimeDecorator)
def test_add_workflow_no_double_wrap() -> None:
server = WorkflowServer()
wf = SimpleWorkflow(runtime=StubRuntime())
server.add_workflow("greeting", wf)
server.add_workflow("greeting", wf)
assert isinstance(wf.runtime, ServerRuntimeDecorator)
# Inner should be IdleReleaseDecorator, not another ServerRuntimeDecorator
assert isinstance(wf.runtime._decorated, IdleReleaseDecorator)
assert not isinstance(wf.runtime._decorated, ServerRuntimeDecorator)
def test_server_runtime_decorator_wraps_internal_adapter() -> None:
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
wf = SimpleWorkflow(runtime=decorator)
adapter = decorator.get_internal_adapter(wf)
assert isinstance(adapter, _ServerInternalRunAdapter)
async def test_server_internal_adapter_records_events_to_store() -> None:
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
wf = SimpleWorkflow(runtime=decorator)
adapter = decorator.get_internal_adapter(wf)
await adapter.write_to_event_stream(StopEvent(result="hello"))
await adapter.write_to_event_stream(StopEvent(result="world"))
events = await store.query_events(adapter.run_id)
assert len(events) == 2
assert events[0].sequence == 0
assert events[1].sequence == 1
assert events[0].event.type == "StopEvent"
assert events[1].event.type == "StopEvent"
async def test_server_internal_adapter_forwards_to_inner() -> None:
"""The adapter should forward write_to_event_stream to the inner adapter.
Events are recorded to the store AND forwarded so that inner decorators
(e.g. _DurableInternalRunAdapter) can process them for idle detection.
"""
class RecordingAdapter(StubInternalAdapter):
def __init__(self) -> None:
super().__init__()
self.recorded_events: list[Event] = []
async def write_to_event_stream(self, event: Event) -> None:
self.recorded_events.append(event)
inner = RecordingAdapter()
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
adapter = _ServerInternalRunAdapter(inner, decorator)
stop = StopEvent(result="forwarded")
await adapter.write_to_event_stream(stop)
# Event is forwarded to inner adapter for decorator chain processing
assert len(inner.recorded_events) == 1
# And also recorded in the store
events = await store.query_events(adapter.run_id)
assert len(events) == 1
def test_add_workflow_uses_server_runtime_decorator() -> None:
server = WorkflowServer()
wf = SimpleWorkflow(runtime=StubRuntime())
server.add_workflow("test", wf)
assert isinstance(wf.runtime, ServerRuntimeDecorator)
async def test_concurrent_runs_get_independent_sequences() -> None:
"""Two adapters from the same decorator should have independent sequences."""
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
adapter_a = _ServerInternalRunAdapter(StubInternalAdapterWithId("run-a"), decorator)
adapter_b = _ServerInternalRunAdapter(StubInternalAdapterWithId("run-b"), decorator)
# Interleave writes from both adapters
await adapter_a.write_to_event_stream(StopEvent(result="a1"))
await adapter_b.write_to_event_stream(StopEvent(result="b1"))
await adapter_a.write_to_event_stream(StopEvent(result="a2"))
await adapter_b.write_to_event_stream(StopEvent(result="b2"))
await adapter_b.write_to_event_stream(StopEvent(result="b3"))
events_a = await store.query_events("run-a")
events_b = await store.query_events("run-b")
assert len(events_a) == 2
assert [e.sequence for e in events_a] == [0, 1]
assert len(events_b) == 3
assert [e.sequence for e in events_b] == [0, 1, 2]
class StubInternalAdapterWithId(StubInternalAdapter):
def __init__(self, run_id: str) -> None:
super().__init__()
self._run_id = run_id
@property
def run_id(self) -> str:
return self._run_id
@pytest.mark.parametrize(
"event, expected_status, expected_error, expected_has_result",
[
pytest.param(
StopEvent(result="done"),
"completed",
None,
True,
id="stop-event",
),
pytest.param(
WorkflowFailedEvent(
step_name="s",
exception_type="E",
exception_message="boom",
traceback="",
attempts=1,
elapsed_seconds=0.0,
),
"failed",
"boom",
False,
id="failed-event",
),
pytest.param(
WorkflowTimedOutEvent(
timeout=10.0,
active_steps=["s"],
),
"failed",
"Workflow timed out after 10.0s",
False,
id="timed-out-event",
),
pytest.param(
WorkflowCancelledEvent(),
"cancelled",
None,
False,
id="cancelled-event",
),
],
)
async def test_terminal_event_status_transitions(
event: Event,
expected_status: str,
expected_error: str | None,
expected_has_result: bool,
) -> None:
"""Writing a terminal event updates handler status in the store."""
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
decorator._persistence_backoff = [0, 0]
run_id = "run-terminal"
# Seed a handler record so update_handler_status can find it
await store.update(
PersistentHandler(
handler_id="h1",
workflow_name="test",
status="running",
run_id=run_id,
started_at=datetime.now(timezone.utc),
)
)
inner = StubInternalAdapterWithId(run_id)
adapter = _ServerInternalRunAdapter(inner, decorator)
await adapter.write_to_event_stream(event)
found = await store.query(HandlerQuery(run_id_in=[run_id]))
assert len(found) == 1
handler = found[0]
assert handler.status == expected_status
assert handler.error == expected_error
if expected_has_result:
assert handler.result is not None
else:
assert handler.result is None
async def test_run_workflow_handler_persists_initial_record() -> None:
"""run_workflow_handler creates a running handler record in the store."""
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
decorator._persistence_backoff = [0, 0]
mock_handler = MagicMock(run_id="test-run")
await decorator.run_workflow_handler("h-init", "my_workflow", mock_handler)
found = await store.query(HandlerQuery(handler_id_in=["h-init"]))
assert len(found) == 1
handler = found[0]
assert handler.handler_id == "h-init"
assert handler.workflow_name == "my_workflow"
assert handler.status == "running"
assert handler.run_id == "test-run"
assert handler.started_at is not None
async def test_retry_store_write_succeeds_after_failures() -> None:
"""_retry_store_write retries and eventually succeeds."""
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
decorator._persistence_backoff = [0, 0]
call_count = 0
async def flaky() -> None:
nonlocal call_count
call_count += 1
if call_count < 3:
raise RuntimeError("transient")
await decorator._retry_store_write(flaky)
assert call_count == 3
async def test_retry_store_write_raises_after_exhaustion() -> None:
"""_retry_store_write raises when all retries are exhausted."""
store = MemoryWorkflowStore()
decorator = ServerRuntimeDecorator(StubRuntime(), store=store)
decorator._persistence_backoff = [0, 0]
async def always_fail() -> None:
raise RuntimeError("permanent")
with pytest.raises(RuntimeError, match="permanent"):
await decorator._retry_store_write(always_fail)
@@ -0,0 +1,328 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import sqlite3
from pathlib import Path
import pytest
from llama_agents.server import SqliteWorkflowStore
from llama_agents.server._store.sqlite.migrate import run_migrations
from llama_agents.server._store.sqlite.sqlite_state_store import (
SqliteStateStore,
)
from pydantic import BaseModel
from workflows.context.serializers import JsonSerializer
from workflows.context.state_store import DictState, InMemoryStateStore
# -- Typed state models for testing --
class CounterState(BaseModel):
count: int = 0
label: str = "default"
class ExtendedCounterState(CounterState):
extra: str = "extra_default"
# -- Fixtures --
@pytest.fixture
def db_path(tmp_path: Path) -> str:
path = str(tmp_path / "test_state.db")
conn = sqlite3.connect(path)
try:
run_migrations(conn)
conn.commit()
finally:
conn.close()
return path
@pytest.fixture
def store(db_path: str) -> SqliteStateStore[DictState]:
return SqliteStateStore(db_path=db_path, run_id="run-1")
@pytest.fixture
def typed_store(db_path: str) -> SqliteStateStore[CounterState]:
return SqliteStateStore(
db_path=db_path,
run_id="run-typed",
state_type=CounterState,
)
# -- Basic get/set tests --
@pytest.mark.asyncio
async def test_get_returns_default_dict_state(
store: SqliteStateStore[DictState],
) -> None:
state = await store.get_state()
assert isinstance(state, DictState)
assert dict(state) == {}
@pytest.mark.asyncio
async def test_set_and_get_path(store: SqliteStateStore[DictState]) -> None:
await store.set("foo", 42)
value = await store.get("foo")
assert value == 42
@pytest.mark.asyncio
async def test_set_nested_path(store: SqliteStateStore[DictState]) -> None:
await store.set("a.b.c", "deep")
value = await store.get("a.b.c")
assert value == "deep"
@pytest.mark.asyncio
async def test_get_missing_path_raises(store: SqliteStateStore[DictState]) -> None:
with pytest.raises(ValueError, match="not found"):
await store.get("nonexistent")
@pytest.mark.asyncio
async def test_get_missing_path_returns_default(
store: SqliteStateStore[DictState],
) -> None:
value = await store.get("nonexistent", default="fallback")
assert value == "fallback"
@pytest.mark.asyncio
async def test_set_empty_path_raises(store: SqliteStateStore[DictState]) -> None:
with pytest.raises(ValueError, match="cannot be empty"):
await store.set("", 42)
# -- get_state / set_state --
@pytest.mark.asyncio
async def test_set_state_replaces_dict_state(
store: SqliteStateStore[DictState],
) -> None:
await store.set("x", 1)
new_state = DictState(y=2)
await store.set_state(new_state)
state = await store.get_state()
assert "y" in state
assert "x" not in state
@pytest.mark.asyncio
async def test_typed_state_get_returns_default(
typed_store: SqliteStateStore[CounterState],
) -> None:
state = await typed_store.get_state()
assert isinstance(state, CounterState)
assert state.count == 0
assert state.label == "default"
@pytest.mark.asyncio
async def test_typed_state_set_and_get(
typed_store: SqliteStateStore[CounterState],
) -> None:
await typed_store.set_state(CounterState(count=5, label="updated"))
state = await typed_store.get_state()
assert state.count == 5
assert state.label == "updated"
@pytest.mark.asyncio
async def test_set_state_parent_type_merge(db_path: str) -> None:
"""Setting a parent type state merges fields, preserving child-specific fields."""
store: SqliteStateStore[ExtendedCounterState] = SqliteStateStore(
db_path=db_path,
run_id="run-merge",
state_type=ExtendedCounterState,
)
await store.set_state(ExtendedCounterState(count=1, label="init", extra="mine"))
# Set parent type — should merge
parent = CounterState(count=10, label="merged")
await store.set_state(parent) # type: ignore[arg-type]
state = await store.get_state()
assert state.count == 10
assert state.label == "merged"
assert state.extra == "mine" # child field preserved
# -- edit_state --
@pytest.mark.asyncio
async def test_edit_state_dict(store: SqliteStateStore[DictState]) -> None:
await store.set("counter", 0)
async with store.edit_state() as state:
state["counter"] = state["counter"] + 1
value = await store.get("counter")
assert value == 1
@pytest.mark.asyncio
async def test_edit_state_typed(typed_store: SqliteStateStore[CounterState]) -> None:
async with typed_store.edit_state() as state:
state.count += 10
result = await typed_store.get_state()
assert result.count == 10
# -- clear --
@pytest.mark.asyncio
async def test_clear_resets_state(store: SqliteStateStore[DictState]) -> None:
await store.set("x", 99)
await store.clear()
state = await store.get_state()
assert dict(state) == {}
@pytest.mark.asyncio
async def test_clear_resets_typed_state(
typed_store: SqliteStateStore[CounterState],
) -> None:
await typed_store.set_state(CounterState(count=100, label="dirty"))
await typed_store.clear()
state = await typed_store.get_state()
assert state.count == 0
assert state.label == "default"
# -- Persistence across instances --
@pytest.mark.asyncio
async def test_state_persists_across_instances(db_path: str) -> None:
"""State set by one store instance is readable by a new instance pointing at the same DB."""
store1: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-persist"
)
await store1.set("key", "value")
store2: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-persist"
)
value = await store2.get("key")
assert value == "value"
@pytest.mark.asyncio
async def test_typed_state_persists_across_instances(db_path: str) -> None:
store1: SqliteStateStore[CounterState] = SqliteStateStore(
db_path=db_path,
run_id="run-typed-persist",
state_type=CounterState,
)
await store1.set_state(CounterState(count=42, label="persisted"))
store2: SqliteStateStore[CounterState] = SqliteStateStore(
db_path=db_path,
run_id="run-typed-persist",
state_type=CounterState,
)
state = await store2.get_state()
assert state.count == 42
assert state.label == "persisted"
@pytest.mark.asyncio
async def test_different_run_ids_are_isolated(db_path: str) -> None:
store_a: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-a"
)
store_b: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-b"
)
await store_a.set("x", "from-a")
await store_b.set("x", "from-b")
assert await store_a.get("x") == "from-a"
assert await store_b.get("x") == "from-b"
# -- to_dict / from_dict --
@pytest.mark.asyncio
async def test_to_dict_returns_metadata_only(
store: SqliteStateStore[DictState],
) -> None:
await store.set("key", "value")
serializer = JsonSerializer()
d = store.to_dict(serializer)
assert d["store_type"] == "sqlite"
assert d["run_id"] == "run-1"
assert "state_data" not in d
@pytest.mark.asyncio
async def test_from_dict_sqlite_format(db_path: str) -> None:
"""from_dict with sqlite format reconnects to existing row."""
store1: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-fromdict"
)
await store1.set("saved", True)
serializer = JsonSerializer()
payload = store1.to_dict(serializer)
store2 = SqliteStateStore.from_dict(
payload, serializer, db_path=db_path, state_type=DictState
)
value = await store2.get("saved")
assert value is True
@pytest.mark.asyncio
async def test_from_dict_in_memory_format_migrates(db_path: str) -> None:
"""from_dict with InMemorySerializedState format stores data on first DB access."""
serializer = JsonSerializer()
in_memory_store = InMemoryStateStore(DictState(migrated_key="migrated_value"))
payload = in_memory_store.to_dict(serializer)
store = SqliteStateStore.from_dict(
payload,
serializer,
db_path=db_path,
state_type=DictState,
run_id="run-migrate",
)
value = await store.get("migrated_key")
assert value == "migrated_value"
@pytest.mark.asyncio
async def test_from_dict_empty_raises() -> None:
with pytest.raises(ValueError, match="Cannot restore"):
SqliteStateStore.from_dict({}, JsonSerializer())
# -- Migration applies cleanly --
@pytest.mark.asyncio
async def test_migration_applies_on_existing_db(tmp_path: Path) -> None:
"""Verify the state table can be created on a DB that already has other tables."""
db_path = str(tmp_path / "existing.db")
# Create DB with workflow store tables first
SqliteWorkflowStore(db_path)
# Now create state store on same DB — should work
store: SqliteStateStore[DictState] = SqliteStateStore(
db_path=db_path, run_id="run-coexist"
)
await store.set("coexist", True)
value = await store.get("coexist")
assert value is True
@@ -1,11 +1,11 @@
from __future__ import annotations
from pathlib import Path
import pytest
from llama_agents.server._store.abstract_workflow_store import (
from llama_agents.server import (
HandlerQuery,
PersistentHandler,
)
from llama_agents.server._store.sqlite.sqlite_workflow_store import (
SqliteWorkflowStore,
)
from workflows.events import StopEvent
@@ -20,7 +20,6 @@ async def test_update_and_query_returns_inserted_handler(tmp_path: Path) -> None
handler_id="h1",
workflow_name="wf_a",
status="running",
ctx={"state": {"x": 1, "y": [1, 2, 3]}},
)
await store.update(handler)
@@ -35,7 +34,6 @@ async def test_update_and_query_returns_inserted_handler(tmp_path: Path) -> None
assert found.handler_id == "h1"
assert found.workflow_name == "wf_a"
assert found.status == "running"
assert found.ctx == {"state": {"x": 1, "y": [1, 2, 3]}}
@pytest.mark.asyncio
@@ -49,17 +47,15 @@ async def test_update_on_conflict_overwrites_existing_row(tmp_path: Path) -> Non
handler_id="h2",
workflow_name="wf_b",
status="running",
ctx={"k": "v1"},
)
)
# Update same handler_id (completed) with new ctx
# Update same handler_id (completed)
await store.update(
PersistentHandler(
handler_id="h2",
workflow_name="wf_b",
status="completed",
ctx={"k": "v2", "n": 42},
)
)
@@ -69,7 +65,7 @@ async def test_update_on_conflict_overwrites_existing_row(tmp_path: Path) -> Non
)
assert result_in_progress == []
# Should be returned for completed=True with latest ctx
# Should be returned for completed=True with latest values
result_completed = await store.query(
HandlerQuery(workflow_name_in=["wf_b"], status_in=["completed"])
)
@@ -78,7 +74,6 @@ async def test_update_on_conflict_overwrites_existing_row(tmp_path: Path) -> Non
assert found.handler_id == "h2"
assert found.workflow_name == "wf_b"
assert found.status == "completed"
assert found.ctx == {"k": "v2", "n": 42}
@pytest.mark.asyncio
@@ -91,7 +86,6 @@ async def test_delete_filters_by_query(tmp_path: Path) -> None:
handler_id="delete-me",
workflow_name="wf_delete",
status="completed",
ctx={"val": 1},
)
)
await store.update(
@@ -99,7 +93,6 @@ async def test_delete_filters_by_query(tmp_path: Path) -> None:
handler_id="keep-me",
workflow_name="wf_keep",
status="running",
ctx={"val": 2},
)
)
@@ -121,7 +114,6 @@ async def test_delete_noop_on_empty_filter(tmp_path: Path) -> None:
handler_id="delete-me",
workflow_name="wf_delete",
status="completed",
ctx={},
)
)
@@ -145,7 +137,6 @@ async def test_query_filters_by_handler_id_and_empty_lists(tmp_path: Path) -> No
handler_id=hid,
workflow_name=wf,
status="running",
ctx={"seed": hid},
)
)
@@ -193,7 +184,6 @@ async def test_update_pydantic_result_serialization(
workflow_name="wf_pyd",
status="completed",
result=event,
ctx={"state": {"ok": True}},
)
# This would raise TypeError if the store used json.dumps(handler.result)
@@ -0,0 +1,288 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
from datetime import datetime, timezone
import pytest
from llama_agents.server import (
HandlerQuery,
MemoryWorkflowStore,
PersistentHandler,
WorkflowServer,
)
from llama_agents.server._service import EventSendError, HandlerCompletedError
from server_test_fixtures import ( # type: ignore[import]
ErrorWorkflow,
ExternalEvent,
wait_for_passing,
)
from workflows import Workflow
@pytest.mark.asyncio
async def test_cancel_running_handler(
memory_store: MemoryWorkflowStore, interactive_workflow: Workflow
) -> None:
"""Start an interactive workflow, cancel it, and verify status becomes cancelled."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow(
"interactive", interactive_workflow, additional_events=[ExternalEvent]
)
async with server.contextmanager():
handler_data = await server._service.start_workflow(
interactive_workflow, "cancel-test-1"
)
assert handler_data.run_id is not None
result = await server._service.cancel_handler("cancel-test-1")
assert result == "cancelled"
async def status_is_cancelled() -> None:
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["cancel-test-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "cancelled"
await wait_for_passing(status_is_cancelled, max_duration=2.0, interval=0.01)
@pytest.mark.asyncio
async def test_cancel_handler_with_purge(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""Start and complete a workflow, then purge it from the store."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("simple", simple_test_workflow)
async with server.contextmanager():
await server._service.start_workflow(simple_test_workflow, "purge-test-1")
# Wait for completion
async def handler_completed() -> None:
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["purge-test-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "completed"
await wait_for_passing(handler_completed, max_duration=2.0, interval=0.01)
result = await server._service.cancel_handler("purge-test-1", purge=True)
assert result == "deleted"
# Handler should be gone from store
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["purge-test-1"])
)
assert len(persisted) == 0
@pytest.mark.asyncio
async def test_cancel_handler_not_found(memory_store: MemoryWorkflowStore) -> None:
"""Cancelling a nonexistent handler returns None."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
async with server.contextmanager():
result = await server._service.cancel_handler("nonexistent")
assert result is None
@pytest.mark.asyncio
async def test_send_event_workflow_not_registered(
memory_store: MemoryWorkflowStore,
) -> None:
"""Sending an event to a handler whose workflow is not registered raises EventSendError."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
# Seed store with a handler for an unregistered workflow
await memory_store.update(
PersistentHandler(
handler_id="orphan-handler",
workflow_name="unregistered",
status="running",
run_id="some-run-id",
started_at=datetime.now(timezone.utc),
)
)
async with server.contextmanager():
with pytest.raises(EventSendError, match="not registered"):
await server._service.send_event(
"orphan-handler", ExternalEvent(response="hello")
)
@pytest.mark.asyncio
async def test_send_event_no_run_id(
memory_store: MemoryWorkflowStore, interactive_workflow: Workflow
) -> None:
"""Sending an event to a handler with no run_id raises EventSendError."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow(
"interactive", interactive_workflow, additional_events=[ExternalEvent]
)
# Seed store with a handler that has no run_id
await memory_store.update(
PersistentHandler(
handler_id="no-run-handler",
workflow_name="interactive",
status="running",
run_id=None,
started_at=datetime.now(timezone.utc),
)
)
async with server.contextmanager():
with pytest.raises(EventSendError, match="no run ID"):
await server._service.send_event(
"no-run-handler", ExternalEvent(response="hello")
)
@pytest.mark.asyncio
async def test_start_workflow_happy_path(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""start_workflow returns HandlerData with correct initial fields."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("simple", simple_test_workflow)
async with server.contextmanager():
handler_data = await server._service.start_workflow(
simple_test_workflow, "start-hp-1"
)
assert handler_data.handler_id == "start-hp-1"
assert handler_data.workflow_name == "simple"
assert handler_data.run_id is not None
assert handler_data.status == "running"
@pytest.mark.asyncio
async def test_await_workflow_happy_path(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""await_workflow returns completed HandlerData."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("simple", simple_test_workflow)
async with server.contextmanager():
handler_data = await server._service.start_workflow(
simple_test_workflow, "await-hp-1"
)
result = await server._service.await_workflow(handler_data)
assert result.status == "completed"
assert result.handler_id == "await-hp-1"
@pytest.mark.asyncio
async def test_await_workflow_error_returns_failed(
memory_store: MemoryWorkflowStore,
) -> None:
"""await_workflow on an ErrorWorkflow returns failed status, not an exception."""
error_wf = ErrorWorkflow()
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("error", error_wf)
async with server.contextmanager():
handler_data = await server._service.start_workflow(error_wf, "await-err-1")
result = await server._service.await_workflow(handler_data)
assert result.status == "failed"
@pytest.mark.asyncio
async def test_resolve_handler_raises_on_completed(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""resolve_handler raises HandlerCompletedError for a terminal handler."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("simple", simple_test_workflow)
async with server.contextmanager():
await server._service.start_workflow(simple_test_workflow, "resolve-done-1")
async def handler_completed() -> None:
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["resolve-done-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "completed"
await wait_for_passing(handler_completed, max_duration=2.0, interval=0.01)
with pytest.raises(HandlerCompletedError):
await server._service.resolve_handler("resolve-done-1")
@pytest.mark.asyncio
async def test_send_event_happy_path(
memory_store: MemoryWorkflowStore, interactive_workflow: Workflow
) -> None:
"""send_event delivers an event and the workflow completes."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow(
"interactive", interactive_workflow, additional_events=[ExternalEvent]
)
async with server.contextmanager():
handler_data = await server._service.start_workflow(
interactive_workflow, "send-hp-1"
)
# Wait for the workflow to emit the InputRequiredEvent (meaning it's waiting)
run_id = handler_data.run_id
assert run_id is not None
async def handler_emitted_input_required() -> None:
events = await memory_store.query_events(run_id)
assert any(e.event.type == "RequestedExternalEvent" for e in events)
await wait_for_passing(
handler_emitted_input_required, max_duration=2.0, interval=0.01
)
await server._service.send_event("send-hp-1", ExternalEvent(response="pong"))
async def handler_completed() -> None:
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["send-hp-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "completed"
await wait_for_passing(handler_completed, max_duration=2.0, interval=0.01)
@pytest.mark.asyncio
async def test_cancel_terminal_handler_without_purge(
memory_store: MemoryWorkflowStore, simple_test_workflow: Workflow
) -> None:
"""cancel_handler on an already-completed handler without purge returns None."""
server = WorkflowServer(workflow_store=memory_store, idle_timeout=0.01)
server.add_workflow("simple", simple_test_workflow)
async with server.contextmanager():
await server._service.start_workflow(simple_test_workflow, "cancel-term-1")
async def handler_completed() -> None:
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["cancel-term-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "completed"
await wait_for_passing(handler_completed, max_duration=2.0, interval=0.01)
result = await server._service.cancel_handler("cancel-term-1", purge=False)
assert result is None
# Handler should still exist unchanged
persisted = await memory_store.query(
HandlerQuery(handler_id_in=["cancel-term-1"])
)
assert len(persisted) == 1
assert persisted[0].status == "completed"
@@ -0,0 +1,257 @@
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 LlamaIndex Inc.
from __future__ import annotations
import asyncio
from pathlib import Path
from typing import Optional
import pytest
from llama_agents.client.protocol.serializable_events import EventEnvelopeWithMetadata
from llama_agents.server import (
AbstractWorkflowStore,
MemoryWorkflowStore,
SqliteWorkflowStore,
)
from llama_agents.server._store.abstract_workflow_store import StoredEvent
from workflows.events import (
Event,
StopEvent,
WorkflowCancelledEvent,
WorkflowFailedEvent,
)
def make_envelope(
event: Event | None = None,
seq_label: int = 0,
) -> EventEnvelopeWithMetadata:
"""Create an EventEnvelopeWithMetadata by serializing a real Event."""
if event is None:
event = Event(data=f"seq-{seq_label}")
return EventEnvelopeWithMetadata.from_event(event, include_qualified_name=False)
def _stop(seq_label: int = 0) -> StopEvent:
return StopEvent(data=f"seq-{seq_label}")
def _failed() -> WorkflowFailedEvent:
return WorkflowFailedEvent(
step_name="test_step",
exception_type="ValueError",
exception_message="boom",
traceback="",
attempts=1,
elapsed_seconds=0.0,
)
def _cancelled() -> WorkflowCancelledEvent:
return WorkflowCancelledEvent()
async def _subscribe_and_collect(
store: AbstractWorkflowStore,
run_id: str,
after_sequence: Optional[int] = None,
) -> tuple[list[StoredEvent], asyncio.Task[None]]:
"""Subscribe to events, returning the collected list and the consumer task."""
collected: list[StoredEvent] = []
started = asyncio.Event()
async def consumer() -> None:
started.set()
kwargs = {} if after_sequence is None else {"after_sequence": after_sequence}
async for event in store.subscribe_events(run_id, **kwargs):
collected.append(event)
task = asyncio.create_task(consumer())
await started.wait()
await asyncio.sleep(0.01)
return collected, task
@pytest.fixture(params=["memory", "sqlite"])
def store(request: pytest.FixtureRequest, tmp_path: Path) -> AbstractWorkflowStore:
if request.param == "memory":
return MemoryWorkflowStore()
else:
return SqliteWorkflowStore(str(tmp_path / "test.sqlite"), poll_interval=0.05)
@pytest.mark.asyncio
async def test_append_single_event_and_query_it_back(
store: AbstractWorkflowStore,
) -> None:
await store.append_event("run-1", make_envelope(seq_label=0))
result = await store.query_events("run-1")
assert len(result) == 1
assert result[0].run_id == "run-1"
assert result[0].sequence == 0
assert result[0].event.type == "Event"
@pytest.mark.asyncio
async def test_append_multiple_events_and_query_all(
store: AbstractWorkflowStore,
) -> None:
for i in range(5):
await store.append_event("run-1", make_envelope(seq_label=i))
result = await store.query_events("run-1")
assert len(result) == 5
assert [e.sequence for e in result] == [0, 1, 2, 3, 4]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"seed_count, after_sequence, limit, expected_sequences",
[
pytest.param(5, 2, None, [3, 4], id="after_sequence_only"),
pytest.param(5, None, 3, [0, 1, 2], id="limit_only"),
pytest.param(10, 3, 2, [4, 5], id="after_sequence_and_limit"),
],
)
async def test_query_events_with_filters(
store: AbstractWorkflowStore,
seed_count: int,
after_sequence: Optional[int],
limit: Optional[int],
expected_sequences: list[int],
) -> None:
for i in range(seed_count):
await store.append_event("run-1", make_envelope(seq_label=i))
result = await store.query_events(
"run-1", after_sequence=after_sequence, limit=limit
)
assert len(result) == len(expected_sequences)
assert [e.sequence for e in result] == expected_sequences
@pytest.mark.asyncio
async def test_query_events_for_nonexistent_run_id_returns_empty(
store: AbstractWorkflowStore,
) -> None:
result = await store.query_events("nonexistent-run")
assert result == []
@pytest.mark.asyncio
async def test_events_from_different_run_ids_are_isolated(
store: AbstractWorkflowStore,
) -> None:
for i in range(3):
await store.append_event("run-a", make_envelope(seq_label=i))
for i in range(2):
await store.append_event("run-b", make_envelope(seq_label=i))
result_a = await store.query_events("run-a")
result_b = await store.query_events("run-b")
assert len(result_a) == 3
assert all(e.run_id == "run-a" for e in result_a)
assert len(result_b) == 2
assert all(e.run_id == "run-b" for e in result_b)
@pytest.mark.asyncio
async def test_subscribe_events_receives_appended_events(
store: AbstractWorkflowStore,
) -> None:
"""Appending events wakes a waiting subscriber."""
collected, task = await _subscribe_and_collect(store, "run-1")
await store.append_event("run-1", make_envelope(seq_label=0))
await store.append_event("run-1", make_envelope(seq_label=1))
await store.append_event("run-1", make_envelope(event=_stop(seq_label=2)))
await asyncio.wait_for(task, timeout=5.0)
assert len(collected) == 3
assert [e.sequence for e in collected] == [0, 1, 2]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"terminal_event, expected_type",
[
pytest.param(_failed(), "WorkflowFailedEvent", id="failed"),
pytest.param(_cancelled(), "WorkflowCancelledEvent", id="cancelled"),
],
)
async def test_subscribe_events_terminates_on_terminal_event(
store: AbstractWorkflowStore,
terminal_event: Event,
expected_type: str,
) -> None:
"""Subscriber terminates after receiving a terminal event."""
collected, task = await _subscribe_and_collect(store, "run-1")
if expected_type == "WorkflowFailedEvent":
await store.append_event("run-1", make_envelope())
await store.append_event("run-1", make_envelope(event=terminal_event))
await asyncio.wait_for(task, timeout=5.0)
assert collected[-1].event.type == expected_type
@pytest.mark.asyncio
async def test_subscribe_events_multiple_concurrent_subscribers(
store: AbstractWorkflowStore,
) -> None:
"""Multiple concurrent subscribers on the same run each receive all events."""
collected_a, task_a = await _subscribe_and_collect(store, "run-1")
collected_b, task_b = await _subscribe_and_collect(store, "run-1")
await store.append_event("run-1", make_envelope(seq_label=0))
await store.append_event("run-1", make_envelope(seq_label=1))
await store.append_event("run-1", make_envelope(event=_stop(seq_label=2)))
await asyncio.wait_for(asyncio.gather(task_a, task_b), timeout=5.0)
assert len(collected_a) == 3
assert len(collected_b) == 3
assert [e.sequence for e in collected_a] == [0, 1, 2]
assert [e.sequence for e in collected_b] == [0, 1, 2]
@pytest.mark.asyncio
async def test_subscribe_events_with_after_sequence(
store: AbstractWorkflowStore,
) -> None:
"""Subscriber can resume from a specific sequence position."""
await store.append_event("run-1", make_envelope(seq_label=0))
await store.append_event("run-1", make_envelope(seq_label=1))
await store.append_event("run-1", make_envelope(seq_label=2))
collected, task = await _subscribe_and_collect(store, "run-1", after_sequence=1)
await store.append_event("run-1", make_envelope(event=_stop(seq_label=3)))
await asyncio.wait_for(task, timeout=5.0)
# Should have events 2, 3 (skipping 0 and 1)
assert len(collected) == 2
assert [e.sequence for e in collected] == [2, 3]
@pytest.mark.asyncio
async def test_subscribe_events_already_terminated(
store: AbstractWorkflowStore,
) -> None:
"""If events already contain a terminal event, subscriber terminates immediately."""
await store.append_event("run-1", make_envelope(seq_label=0))
await store.append_event("run-1", make_envelope(event=_stop(seq_label=1)))
collected: list[StoredEvent] = []
async for event in store.subscribe_events("run-1"):
collected.append(event)
assert len(collected) == 2
assert collected[-1].event.type == "StopEvent"
@@ -52,14 +52,6 @@ T = TypeVar("T", bound=Event)
EventBuffer = dict[str, list[Event]]
# Only warn once about unserializable keys
class UnserializableKeyWarning(Warning):
pass
warnings.simplefilter("once", UnserializableKeyWarning)
# TODO(v3) remove this class, and replace with direct references to the pre/internal/external contexts
class Context(Generic[MODEL_T]):
"""
@@ -121,10 +113,6 @@ class Context(Generic[MODEL_T]):
- [InMemoryStateStore][workflows.context.state_store.InMemoryStateStore]
"""
# These keys are set by pre-built workflows and
# are known to be unserializable in some cases.
known_unserializable_keys = ("memory",)
# Current face - context is in exactly one state at a time
_face: (
PreContext[MODEL_T] | ExternalContext[MODEL_T, Any] | InternalContext[MODEL_T]
@@ -32,6 +32,9 @@ if TYPE_CHECKING:
MAX_DEPTH = 1000
# Keys set by pre-built workflows that are known to be unserializable in some cases.
KNOWN_UNSERIALIZABLE_KEYS: tuple[str, ...] = ("memory",)
class InMemorySerializedState(BaseModel):
"""Serialized state containing actual data (from InMemoryStateStore)."""
@@ -81,7 +84,7 @@ def parse_in_memory_state(
def serialize_dict_state_data(
state: DictState,
serializer: BaseSerializer,
known_unserializable_keys: tuple[str, ...] = (),
known_unserializable_keys: tuple[str, ...] = KNOWN_UNSERIALIZABLE_KEYS,
) -> dict[str, Any]:
"""Serialize DictState items to {"_data": {...}} format.
@@ -115,7 +118,7 @@ def serialize_dict_state_data(
def create_in_memory_payload(
state: BaseModel,
serializer: BaseSerializer,
known_unserializable_keys: tuple[str, ...] = (),
known_unserializable_keys: tuple[str, ...] = KNOWN_UNSERIALIZABLE_KEYS,
) -> InMemorySerializedState:
"""Create InMemorySerializedState from any state model.
@@ -447,10 +450,6 @@ class InMemoryStateStore(Generic[MODEL_T]):
- [Context.store][workflows.context.context.Context.store]
"""
# These keys are set by pre-built workflows and
# are known to be unserializable in some cases.
known_unserializable_keys = ("memory",)
state_type: Type[MODEL_T]
def __init__(self, initial_state: MODEL_T):
@@ -510,9 +509,7 @@ class InMemoryStateStore(Generic[MODEL_T]):
dict[str, Any]: A payload suitable for
[from_dict][workflows.context.state_store.InMemoryStateStore.from_dict].
"""
payload = create_in_memory_payload(
self._state, serializer, self.known_unserializable_keys
)
payload = create_in_memory_payload(self._state, serializer)
return payload.model_dump()
@classmethod
@@ -3,7 +3,7 @@
"""Re-export sqlite components from the optional llama-agents-server package."""
try:
from llama_agents.server._store.sqlite.sqlite_workflow_store import (
from llama_agents.server import (
SqliteWorkflowStore,
)
except ImportError as e:
Generated
+4
View File
@@ -1731,6 +1731,7 @@ name = "llama-agents-server"
version = "0.2.0rc0"
source = { editable = "packages/llama-agents-server" }
dependencies = [
{ name = "httpx" },
{ name = "llama-agents-client" },
{ name = "llama-index-workflows" },
{ name = "starlette" },
@@ -1740,6 +1741,7 @@ dependencies = [
[package.dev-dependencies]
dev = [
{ name = "hatch" },
{ name = "llama-agents-integration-tests" },
{ name = "pytest" },
{ name = "pytest-asyncio" },
{ name = "pytest-cov" },
@@ -1751,6 +1753,7 @@ dev = [
[package.metadata]
requires-dist = [
{ name = "httpx", specifier = ">=0.27.0" },
{ name = "llama-agents-client", editable = "packages/llama-agents-client" },
{ name = "llama-index-workflows", editable = "packages/llama-index-workflows" },
{ name = "starlette", specifier = ">=0.39.0" },
@@ -1760,6 +1763,7 @@ requires-dist = [
[package.metadata.requires-dev]
dev = [
{ name = "hatch", specifier = ">=1.14.1" },
{ name = "llama-agents-integration-tests", editable = "packages/llama-agents-integration-tests" },
{ name = "pytest", specifier = ">=8.4.2" },
{ name = "pytest-asyncio", specifier = ">=1.0.0" },
{ name = "pytest-cov", specifier = ">=7.0.0" },