mirror of
https://github.com/run-llama/workflows-py.git
synced 2026-08-26 21:41:14 -04:00
Server runtime (#342)
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
---
|
||||
"llama-agents-client": minor
|
||||
---
|
||||
|
||||
Add SSE event streaming with sequence-based cursors and automatic reconnection on connection drop
|
||||
@@ -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
|
||||
|
||||
@@ -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")])
|
||||
|
||||
@@ -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.
|
||||
+242
@@ -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.
|
||||
|
||||
+147
-3
@@ -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
|
||||
|
||||
+132
-1
@@ -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()
|
||||
|
||||
+33
@@ -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);
|
||||
+257
@@ -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
|
||||
+181
-14
@@ -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:
|
||||
|
||||
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user