From 21b2607e7cf0437daf76d2df7214f5cfde37f5f5 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:08:20 +0200 Subject: [PATCH 01/28] feat(durable): add execution contracts and configuration --- pyproject.toml | 23 ++ src/hayhooks/durable/mode.py | 37 +++ src/hayhooks/durable/models.py | 356 ++++++++++++++++++++++++++++ src/hayhooks/settings.py | 64 ++++- tests/durable_contract.py | 73 ++++++ tests/test_optional_dependencies.py | 48 ++++ tests/test_settings.py | 63 +++++ 7 files changed, 663 insertions(+), 1 deletion(-) create mode 100644 src/hayhooks/durable/mode.py create mode 100644 src/hayhooks/durable/models.py create mode 100644 tests/durable_contract.py create mode 100644 tests/test_optional_dependencies.py diff --git a/pyproject.toml b/pyproject.toml index 7a8e0dad..174d5fd7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,6 +46,11 @@ mcp = [ ] a2a = [ "a2a-sdk[http-server]>=1.1.0,<2.0", + "redis>=5,<9", +] +durable = [ + "haystack-ai>=3,<4", + "redis>=5,<9", ] chainlit = [ "chainlit>=2.0.0", @@ -138,6 +143,24 @@ all = "pytest -vv {args:tests}" all-cov = "all --cov=hayhooks" types = "ty check {args:src/hayhooks}" +[tool.hatch.envs.test-v3] +features = ["durable", "a2a", "mcp"] +extra-dependencies = [ + "qdrant-haystack", + "trafilatura", + "pytest", + "pytest-asyncio", + "pytest-cov", + "pytest-mock", + "ty", +] + +[tool.hatch.envs.test-v3.scripts] +unit = "pytest -vv -m 'not integration' {args:tests}" +integration = "pytest -vv -m integration {args:tests}" +all = "pytest -vv {args:tests}" +types = "ty check {args:src/hayhooks}" + [tool.ty.environment] python-version = "3.10" diff --git a/src/hayhooks/durable/mode.py b/src/hayhooks/durable/mode.py new file mode 100644 index 00000000..8e5cd2ba --- /dev/null +++ b/src/hayhooks/durable/mode.py @@ -0,0 +1,37 @@ +"""One authoritative classification of durable wrapper authoring modes.""" + +from __future__ import annotations + +from enum import Enum + +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper + + +class DurableAuthoringMode(str, Enum): + """How a wrapper participates in the durable runtime.""" + + NONE = "none" + WRAPPER = "wrapper" + MANAGED_AGENT = "managed_agent" + + +def _durable_method_implementations(wrapper: BasePipelineWrapper) -> tuple[bool, bool]: + wrapper_type = type(wrapper) + return ( + bool(getattr(wrapper, "_is_run_durable_implemented", False)) + or wrapper_type.run_durable is not BasePipelineWrapper.run_durable, + bool(getattr(wrapper, "_is_run_durable_async_implemented", False)) + or wrapper_type.run_durable_async is not BasePipelineWrapper.run_durable_async, + ) + + +def durable_authoring_mode(wrapper: BasePipelineWrapper) -> DurableAuthoringMode: + """Classify a wrapper once, with explicit wrapper methods taking precedence.""" + if any(_durable_method_implementations(wrapper)): + return DurableAuthoringMode.WRAPPER + if getattr(wrapper, "durable", False) and wrapper.pipeline is not None: + return DurableAuthoringMode.MANAGED_AGENT + return DurableAuthoringMode.NONE + + +__all__ = ["DurableAuthoringMode", "durable_authoring_mode"] diff --git a/src/hayhooks/durable/models.py b/src/hayhooks/durable/models.py new file mode 100644 index 00000000..80824f54 --- /dev/null +++ b/src/hayhooks/durable/models.py @@ -0,0 +1,356 @@ +""" +Persisted durable-execution model and JSON-safe value helpers. + +This module is transport- and storage-neutral. REST, A2A, Redis, and +Haystack-specific code build on these types rather than extending them. +""" + +from __future__ import annotations + +import json +import re +from collections import deque +from collections.abc import Mapping +from dataclasses import dataclass, field +from datetime import datetime, timezone +from enum import Enum +from typing import Any, TypeAlias, cast + +from hayhooks.durable.engine import ExecutionLeaseLostError as _ExecutionLeaseLostError +from hayhooks.durable.engine import ExecutionStatus, normalize_cancellation_reason + +JsonScalar: TypeAlias = str | int | float | bool | None +JsonValue: TypeAlias = JsonScalar | list["JsonValue"] | dict[str, "JsonValue"] +ExecutionLeaseLostError = _ExecutionLeaseLostError +DEFAULT_MAX_RECORD_BYTES = 1_000_000 +DEFAULT_MAX_PROGRESS_EVENTS = 100 +DEFAULT_MAX_PROGRESS_BYTES = 8_192 +_SENSITIVE_NAME = r"(?:api[_ -]?key|access[_ -]?token|authorization|bearer|password|secret)" + + +def utc_now() -> datetime: + """Return a timezone-aware UTC timestamp.""" + return datetime.now(timezone.utc) + + +def _as_utc(value: datetime) -> datetime: + return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc) + + +def json_safe(value: Any) -> JsonValue: + """Convert common application values to a JSON-compatible value.""" + if value is None or isinstance(value, str | int | float | bool): + return value + if isinstance(value, Enum): + return json_safe(value.value) + if isinstance(value, datetime): + return value.isoformat() + if isinstance(value, Mapping): + return {str(key): json_safe(item) for key, item in value.items()} + if isinstance(value, list | tuple | set | frozenset | deque): + return [json_safe(item) for item in value] + converter = getattr(value, "to_dict", None) + if callable(converter): + return json_safe(converter()) + msg = f"{type(value).__name__} is not JSON serializable" + raise TypeError(msg) + + +def validate_json(value: Any, *, limit: int | None, label: str) -> JsonValue: + """Normalize a value and, when requested, bound its encoded size.""" + try: + safe = json_safe(value) + encoded = json.dumps(safe, ensure_ascii=False, separators=(",", ":"), allow_nan=False) + except (TypeError, ValueError) as error: + msg = f"{label} must be JSON serializable" + raise ValueError(msg) from error + if limit is not None and len(encoded.encode("utf-8")) > limit: + msg = f"{label} exceeds the {limit}-byte durable execution limit" + raise ExecutionRecordSizeError(msg) + return safe + + +def _sanitize_error_message(error: BaseException) -> str: + message = str(error) + patterns = ( + (r"(?i)(authorization\s*:\s*)bearer\s+[A-Za-z0-9._~+/=-]+", r"\1"), + (rf"(?i)([?&]{_SENSITIVE_NAME}=)[^&#\s]+", r"\1"), + (rf'(?i)(["\']{_SENSITIVE_NAME}["\']\s*:\s*)["\'][^"\']*["\']', r'\1""'), + (rf"(?i)({_SENSITIVE_NAME})\s*[:=]\s*[^\s,;&]+", r"\1="), + ) + for pattern, replacement in patterns: + message = re.sub(pattern, replacement, message) + return message[:2_000] + + +class ExecutionStoreError(RuntimeError): + """A storage or lease-heartbeat operation failed transiently.""" + + +class ExecutionAdmissionError(RuntimeError): + """A durable store rejected new work because a configured capacity limit is full.""" + + retry_after_seconds = 1 + + def __init__(self, policy: str) -> None: + self.policy = policy + super().__init__(f"Durable execution admission rejected by {policy} capacity limit") + + +class ExecutionRecordSizeError(ValueError): + """An execution record cannot fit within its configured persistence limit.""" + + +class ExecutionCanceledError(RuntimeError): + """Raised at a cooperative cancellation boundary.""" + + +class ExecutionSuspendedError(RuntimeError): + """Internal signal used after a claim atomically enters ``waiting``.""" + + +class RetryableExecutionError(RuntimeError): + """Ask the manager to persist retry metadata and redeliver later.""" + + def __init__(self, message: str, *, delay: float = 0.0) -> None: + super().__init__(message) + self.delay = max(0.0, delay) + + +class ExecutionKind(str, Enum): + """The Haystack object adapter selected for an execution.""" + + PIPELINE = "pipeline" + AGENT = "agent" + + +@dataclass +class ExecutionError: + """Sanitized persisted failure metadata.""" + + type: str + message: str + retryable: bool = False + code: str | None = None + + def to_dict(self) -> dict[str, JsonValue]: + return {"type": self.type, "message": self.message, "retryable": self.retryable, "code": self.code} + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> ExecutionError: + return cls( + type=str(value.get("type", "ExecutionError")), + message=str(value.get("message", ""))[:2_000], + retryable=bool(value.get("retryable", False)), + code=str(value["code"]) if value.get("code") is not None else None, + ) + + @classmethod + def from_exception(cls, error: BaseException, *, retryable: bool = False) -> ExecutionError: + code = getattr(error, "code", None) + return cls( + type=type(error).__name__, + message=_sanitize_error_message(error), + retryable=retryable, + code=str(code) if code is not None else None, + ) + + +@dataclass +class ExecutionProgressEvent: + """A bounded, safe event visible through REST and A2A projections.""" + + sequence: int + message: str + timestamp: datetime = field(default_factory=utc_now) + kind: str = "progress" + metadata: dict[str, JsonValue] = field(default_factory=dict) + + def __post_init__(self) -> None: + if self.sequence < 1: + msg = "progress event sequence numbers must start at one" + raise ValueError(msg) + self.timestamp = _as_utc(self.timestamp) + self.kind = str(self.kind)[:128] + self.message = str(self.message)[:2_000] + self.metadata = cast( + dict[str, JsonValue], + validate_json(self.metadata, limit=None, label="progress metadata"), + ) + validate_json(self.to_dict(), limit=DEFAULT_MAX_PROGRESS_BYTES, label="progress event") + + def to_dict(self) -> dict[str, JsonValue]: + return { + "sequence": self.sequence, + "kind": self.kind, + "message": self.message, + "timestamp": self.timestamp.isoformat(), + "metadata": self.metadata, + } + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> ExecutionProgressEvent: + timestamp = value.get("timestamp") + return cls( + sequence=int(value["sequence"]), + kind=str(value.get("kind", "progress")), + message=str(value.get("message", "")), + timestamp=datetime.fromisoformat(str(timestamp)) if timestamp else utc_now(), + metadata=dict(cast(Mapping[str, Any], value.get("metadata", {}))), + ) + + +@dataclass +class ExecutionCheckpoint: + """Private recovery data, discriminated by the selected adapter.""" + + kind: ExecutionKind + data: dict[str, JsonValue] + + def __post_init__(self) -> None: + self.kind = ExecutionKind(self.kind) + self.data = cast(dict[str, JsonValue], validate_json(self.data, limit=None, label="checkpoint")) + + def to_dict(self) -> dict[str, JsonValue]: + return {"kind": self.kind.value, "data": self.data} + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> ExecutionCheckpoint: + return cls(kind=ExecutionKind(str(value["kind"])), data=dict(cast(Mapping[str, Any], value["data"]))) + + +@dataclass +class ExecutionRecord: + """Hayhooks' versioned, private durable execution envelope.""" + + execution_id: str + execution_kind: ExecutionKind + deployment_name: str + definition_revision: str + validated_input: dict[str, JsonValue] + operation_fingerprint: str = "" + owner_id: str | None = None + status: ExecutionStatus = ExecutionStatus.QUEUED + sequence: int = 0 + attempt: int = 0 + checkpoint: ExecutionCheckpoint | None = None + application_state: dict[str, JsonValue] = field(default_factory=dict) + wait: dict[str, JsonValue] | None = None + progress: list[ExecutionProgressEvent] = field(default_factory=list) + result: JsonValue | None = None + error: ExecutionError | None = None + last_retry_error: ExecutionError | None = None + retry_at: datetime | None = None + cancel_requested_at: datetime | None = None + cancel_reason: str | None = None + created_at: datetime = field(default_factory=utc_now) + updated_at: datetime = field(default_factory=utc_now) + max_progress_events: int = field(default=DEFAULT_MAX_PROGRESS_EVENTS, repr=False, compare=False) + max_record_bytes: int = field(default=DEFAULT_MAX_RECORD_BYTES, repr=False, compare=False) + + def __post_init__(self) -> None: + if not self.execution_id or not self.deployment_name or not self.definition_revision: + msg = "execution_id, deployment_name, and definition_revision must be non-empty" + raise ValueError(msg) + self.execution_kind = ExecutionKind(self.execution_kind) + self.status = ExecutionStatus(self.status) + self.created_at = _as_utc(self.created_at) + self.updated_at = _as_utc(self.updated_at) + if self.cancel_requested_at is not None: + self.cancel_requested_at = _as_utc(self.cancel_requested_at) + self.cancel_reason = normalize_cancellation_reason(self.cancel_reason) + self.validated_input = cast( + dict[str, JsonValue], + validate_json(self.validated_input, limit=self.max_record_bytes, label="validated input"), + ) + self.application_state = cast( + dict[str, JsonValue], + validate_json(self.application_state, limit=self.max_record_bytes, label="application state"), + ) + if self.wait is not None: + self.wait = cast(dict[str, JsonValue], validate_json(self.wait, limit=self.max_record_bytes, label="wait")) + if self.result is not None: + self.result = validate_json(self.result, limit=self.max_record_bytes, label="result") + if self.checkpoint is not None and not isinstance(self.checkpoint, ExecutionCheckpoint): + self.checkpoint = ExecutionCheckpoint.from_dict(cast(Mapping[str, Any], self.checkpoint)) + if self.checkpoint is not None: + self.checkpoint.data = cast( + dict[str, JsonValue], + validate_json(self.checkpoint.data, limit=self.max_record_bytes, label="checkpoint"), + ) + if self.error is not None and not isinstance(self.error, ExecutionError): + self.error = ExecutionError.from_dict(cast(Mapping[str, Any], self.error)) + if self.last_retry_error is not None and not isinstance(self.last_retry_error, ExecutionError): + self.last_retry_error = ExecutionError.from_dict(cast(Mapping[str, Any], self.last_retry_error)) + self.progress = [ + event + if isinstance(event, ExecutionProgressEvent) + else ExecutionProgressEvent.from_dict(cast(Mapping[str, Any], event)) + for event in self.progress + ] + self._trim_progress() + + @property + def terminal(self) -> bool: + return self.status.terminal + + def touch(self) -> None: + self.sequence += 1 + self.updated_at = utc_now() + + def append_progress( + self, message: str, *, kind: str = "progress", metadata: Mapping[str, Any] | None = None + ) -> ExecutionProgressEvent: + event = ExecutionProgressEvent( + sequence=self.progress[-1].sequence + 1 if self.progress else 1, + message=message, + kind=kind, + metadata=dict(metadata or {}), + ) + self.progress.append(event) + self._trim_progress() + self.touch() + return event + + def mark_canceled(self) -> None: + """Mark this execution as canceled unless it is already terminal.""" + if not self.terminal or self.status == ExecutionStatus.CANCELED: + if self.cancel_requested_at is None: + self.cancel_requested_at = utc_now() + self.status = ExecutionStatus.CANCELED + self.error = None + self.result = None + self.retry_at = None + self.wait = None + self.touch() + + def mark_failed(self, error: BaseException | ExecutionError) -> None: + self.status = ExecutionStatus.FAILED + self.error = error if isinstance(error, ExecutionError) else ExecutionError.from_exception(error) + self.wait = None + self.touch() + + def safe_view(self, *, links: Mapping[str, str] | None = None) -> dict[str, JsonValue]: + """Return only the public projection; never expose inputs/checkpoints.""" + waiting = None + if self.wait is not None: + waiting = {key: self.wait[key] for key in ("kind", "message", "expected_input_schema") if key in self.wait} + return { + "execution_id": self.execution_id, + "status": self.status.value, + "attempt": self.attempt, + "sequence": self.sequence, + "progress": [event.to_dict() for event in self.progress], + "result": self.result, + "error": self.error.to_dict() if self.error else None, + "waiting": waiting, + "cancellation_requested_at": (self.cancel_requested_at.isoformat() if self.cancel_requested_at else None), + "created_at": self.created_at.isoformat(), + "updated_at": self.updated_at.isoformat(), + "links": dict(links or {}), + } + + def _trim_progress(self) -> None: + self.max_progress_events = max(1, self.max_progress_events) + if len(self.progress) > self.max_progress_events: + self.progress = self.progress[-self.max_progress_events :] diff --git a/src/hayhooks/settings.py b/src/hayhooks/settings.py index 72eece7d..bb57f477 100644 --- a/src/hayhooks/settings.py +++ b/src/hayhooks/settings.py @@ -3,8 +3,9 @@ from typing import Literal from dotenv import find_dotenv, load_dotenv -from pydantic import Field +from pydantic import Field, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict +from typing_extensions import Self from hayhooks.server.logger import log @@ -81,6 +82,67 @@ class AppSettings(BaseSettings): # tools, e.g. the a2a-inspector, still speak 0.3 during the 1.0 transition) a2a_v0_3_compat: bool = True + # Built-in A2A task-store backend. ``auto`` selects Redis only when durable + # A2A execution also uses Redis; an explicit memory choice is never changed. + a2a_task_store: Literal["auto", "memory", "redis"] = "auto" + + # Connection settings for the built-in Redis task store. + a2a_redis_url: str = "redis://localhost:6379/0" + a2a_redis_key_prefix: str = "hayhooks:a2a" + a2a_redis_socket_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + a2a_redis_socket_connect_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + a2a_redis_health_check_interval: int = Field(default=30, ge=0, le=3_600) + a2a_terminal_task_ttl_seconds: int = Field(default=604_800, ge=1) + a2a_task_snapshot_cache_size: int = Field(default=1_024, ge=1, le=1_000_000) + a2a_list_scan_batch_size: int = Field(default=500, ge=1, le=10_000) + + # Durable executions use Redis by default. Memory is an explicit volatile + # development/test choice and is never selected after a Redis failure. + durable_store: Literal["memory", "redis"] = "redis" + durable_redis_url: str = "redis://localhost:6379/0" + durable_redis_key_prefix: str = "hayhooks:durable" + durable_redis_socket_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + durable_redis_socket_connect_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + durable_redis_health_check_interval: int = Field(default=30, ge=0, le=3_600) + + # Retention and safety limits are application settings, not wrapper API. + durable_terminal_ttl_seconds: int = Field(default=604_800, ge=1) + durable_max_progress_events: int = Field(default=100, ge=1, le=10_000) + durable_max_record_bytes: int = Field(default=1_000_000, ge=1_024) + # Zero disables the deployment-wide queued/running/waiting admission cap. + durable_max_nonterminal_executions: int = Field(default=0, ge=0) + durable_shutdown_grace_period: float = Field(default=5.0, ge=0.0) + durable_max_attempts: int = Field(default=3, ge=1, le=1_000) + durable_retry_base_delay: float = Field(default=1.0, ge=0.0, le=86_400.0) + durable_retry_max_delay: float = Field(default=60.0, ge=0.0, le=604_800.0) + # Shared worker and lease-maintenance polling per deployment/replica at concurrency 1: + # 0.25 s -> ~16 idle Redis commands/s and ~125 ms average pickup latency. + # 0.50 s -> ~8 idle Redis commands/s and ~250 ms average pickup latency. + # 1.00 s -> ~4 idle Redis commands/s and ~500 ms average pickup latency (default). + durable_poll_interval: float = Field(default=1.0, ge=0.05, le=60.0) + # When configured, a trusted reverse proxy must strip any client-supplied + # value and inject the authenticated owner. Empty means bearer-ID mode. + durable_trusted_owner_header: str = "" + + durable_lease_duration_ms: int = Field(default=30_000, ge=1, le=86_400_000) + durable_lease_commit_safety_ms: int = Field(default=1_500, ge=0, le=86_400_000) + + # Keep the default conservative until Agent tools and shared components are + # proven concurrency-safe. + durable_execution_concurrency: int = Field(default=1, ge=1, le=128) + + @model_validator(mode="after") + def _validate_durable_lease_margin(self) -> Self: + if self.durable_lease_commit_safety_ms >= self.durable_lease_duration_ms: + msg = "durable_lease_commit_safety_ms must be smaller than durable_lease_duration_ms" + raise ValueError(msg) + if self.durable_lease_duration_ms - self.durable_lease_commit_safety_ms <= max( + 10, self.durable_lease_duration_ms / 3 + ): + msg = "durable lease duration minus commit safety must exceed the heartbeat interval" + raise ValueError(msg) + return self + # Disable SSL verification when making requests from the CLI disable_ssl: bool = False diff --git a/tests/durable_contract.py b/tests/durable_contract.py new file mode 100644 index 00000000..3b5c7d48 --- /dev/null +++ b/tests/durable_contract.py @@ -0,0 +1,73 @@ +"""Shared observable contract for the memory and Redis durable stores.""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from hayhooks.durable.engine import ( + Checkpoint, + Claim, + Complete, + ExecutionStatus, + RequestCancellation, + Resume, + Suspend, + initial_control, +) +from hayhooks.durable.redis import ExecutionIdempotencyConflictError + + +def control(run_id: str = "run_1", *, idempotency_digest: str = "a" * 64, binding_digest: str = "b" * 64): + return initial_control( + run_id=run_id, + idempotency_digest=idempotency_digest, + idempotency_binding_digest=binding_digest, + deployment="integration", + definition_revision="rev-1", + owner_id="owner", + kind="pipeline", + now_ms=0, + ) + + +async def assert_store_contract(store: Any) -> None: + """Exercise public store behavior without inspecting backend internals.""" + accepted = await store.submit(control(), b"{}", binding_digest="b" * 64) + replay = await store.submit(control("run_2"), b"{}", binding_digest="b" * 64) + assert accepted.created and not replay.created + assert replay.control.run_id == "run_1" + with pytest.raises(ExecutionIdempotencyConflictError): + await store.submit(control("run_3", binding_digest="e" * 64), b"{}", binding_digest="e" * 64) + + run_id = await store.read_candidate() + assert run_id is not None + claimed = await store.transition( + run_id, + Claim("worker", 0, 1_000, 3, "rev-1"), + candidate=True, + ) + assert claimed.next_control.status is ExecutionStatus.RUNNING + checkpointed = await store.transition(run_id, Checkpoint(1, "worker", 0, 1_000, b"checkpoint")) + suspended = await store.transition( + run_id, + Suspend(checkpointed.next_control.fence, "worker", 0, b"checkpoint-2", b"wait"), + ) + assert suspended.next_control.status is ExecutionStatus.WAITING + resumed = await store.transition(run_id, Resume(0, "rev-1", b"checkpoint-3")) + assert resumed.next_control.status is ExecutionStatus.QUEUED + + candidate = await store.read_candidate() + assert candidate == run_id + claimed_again = await store.transition( + run_id, + Claim("worker", 0, 1_000, 3, "rev-1"), + candidate=True, + ) + await store.transition(run_id, RequestCancellation(0, "stop")) + terminal = await store.transition( + run_id, + Complete(claimed_again.next_control.fence, "worker", 0, b"ignored-result"), + ) + assert terminal.next_control.status is ExecutionStatus.CANCELED diff --git a/tests/test_optional_dependencies.py b/tests/test_optional_dependencies.py new file mode 100644 index 00000000..321986d2 --- /dev/null +++ b/tests/test_optional_dependencies.py @@ -0,0 +1,48 @@ +import subprocess +import sys +import textwrap + + +def test_base_import_and_cli_do_not_require_a2a_or_redis(): + script = textwrap.dedent( + """ + import importlib.abc + import sys + + class OptionalDependencyBlocker(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path=None, target=None): + if ( + fullname == "a2a" + or fullname.startswith("a2a.") + or fullname == "redis" + or fullname.startswith("redis.") + ): + raise ModuleNotFoundError(f"blocked optional dependency: {fullname}") + return None + + sys.meta_path.insert(0, OptionalDependencyBlocker()) + + import hayhooks + from hayhooks.cli import hayhooks_cli + from hayhooks.durable import ( + DurableRuntime, + ExecutionStore, + ExecutionStoreProvider, + InMemoryExecutionStoreProvider, + RedisExecutionStoreProvider, + durable_runtime, + ) + from hayhooks.server.app import create_app + + assert hayhooks.__name__ == "hayhooks" + assert callable(hayhooks_cli) + assert callable(create_app) + assert DurableRuntime is not None + assert ExecutionStore is not None + assert ExecutionStoreProvider is not None + assert InMemoryExecutionStoreProvider is not None + assert RedisExecutionStoreProvider is not None + assert durable_runtime is not None + """ + ) + subprocess.run([sys.executable, "-c", script], check=True, capture_output=True, text=True) diff --git a/tests/test_settings.py b/tests/test_settings.py index 8e7343fa..fe329090 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -50,6 +50,69 @@ def test_env_var_prefix(monkeypatch): assert settings.port == 5678 +def test_durable_redis_settings_defaults(monkeypatch): + names = ( + "HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS", + "HAYHOOKS_DURABLE_POLL_INTERVAL", + "HAYHOOKS_DURABLE_LEASE_DURATION_MS", + "HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS", + "HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT", + "HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT", + "HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL", + ) + for name in names: + monkeypatch.delenv(name, raising=False) + + settings = AppSettings() + + assert settings.durable_lease_duration_ms == 30_000 + assert settings.durable_lease_commit_safety_ms == 1_500 + assert settings.durable_poll_interval == 1.0 + assert settings.durable_max_nonterminal_executions == 0 + assert settings.durable_redis_socket_timeout == 5.0 + assert settings.durable_redis_socket_connect_timeout == 5.0 + assert settings.durable_redis_health_check_interval == 30 + + +def test_durable_redis_settings_from_environment(monkeypatch): + monkeypatch.setenv("HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS", "250") + monkeypatch.setenv("HAYHOOKS_DURABLE_POLL_INTERVAL", "0.5") + monkeypatch.setenv("HAYHOOKS_DURABLE_LEASE_DURATION_MS", "45000") + monkeypatch.setenv("HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS", "2000") + monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT", "3.5") + monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT", "2.5") + monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL", "20") + + settings = AppSettings() + + assert settings.durable_lease_duration_ms == 45_000 + assert settings.durable_lease_commit_safety_ms == 2_000 + assert settings.durable_poll_interval == 0.5 + assert settings.durable_max_nonterminal_executions == 250 + assert settings.durable_redis_socket_timeout == 3.5 + assert settings.durable_redis_socket_connect_timeout == 2.5 + assert settings.durable_redis_health_check_interval == 20 + + +def test_durable_lease_safety_margin_must_leave_time_for_a_commit() -> None: + with pytest.raises(ValueError, match="durable_lease_commit_safety_ms"): + AppSettings(durable_lease_duration_ms=1_000, durable_lease_commit_safety_ms=1_000) + with pytest.raises(ValueError, match="heartbeat interval"): + AppSettings(durable_lease_duration_ms=1, durable_lease_commit_safety_ms=0) + with pytest.raises(ValueError, match="heartbeat interval"): + AppSettings(durable_lease_duration_ms=1_000, durable_lease_commit_safety_ms=700) + + +def test_a2a_task_store_bounds_from_environment(monkeypatch): + monkeypatch.setenv("HAYHOOKS_A2A_TASK_SNAPSHOT_CACHE_SIZE", "64") + monkeypatch.setenv("HAYHOOKS_A2A_LIST_SCAN_BATCH_SIZE", "25") + + configured = AppSettings() + + assert configured.a2a_task_snapshot_cache_size == 64 + assert configured.a2a_list_scan_batch_size == 25 + + def test_cors(): default_settings = AppSettings() assert default_settings.cors_allow_origins == ["*"] From d0acca7337a275b8518d7ab4526d335c933f2bcb Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:08:45 +0200 Subject: [PATCH 02/28] feat(durable): implement portable execution stores --- src/hayhooks/durable/adapters.py | 386 ++++++++++++++++ src/hayhooks/durable/backend.py | 208 +++++++++ src/hayhooks/durable/redis.py | 461 +++++++++++++++++++ src/hayhooks/durable/reference.py | 181 ++++++++ src/hayhooks/durable/store.py | 706 ++++++++++++++++++++++++++++++ tests/conftest.py | 20 + tests/test_durable_redis_codec.py | 113 +++++ tests/test_durable_reference.py | 93 ++++ tests/test_durable_store.py | 495 +++++++++++++++++++++ 9 files changed, 2663 insertions(+) create mode 100644 src/hayhooks/durable/adapters.py create mode 100644 src/hayhooks/durable/backend.py create mode 100644 src/hayhooks/durable/redis.py create mode 100644 src/hayhooks/durable/reference.py create mode 100644 src/hayhooks/durable/store.py create mode 100644 tests/test_durable_redis_codec.py create mode 100644 tests/test_durable_reference.py create mode 100644 tests/test_durable_store.py diff --git a/src/hayhooks/durable/adapters.py b/src/hayhooks/durable/adapters.py new file mode 100644 index 00000000..41f820e3 --- /dev/null +++ b/src/hayhooks/durable/adapters.py @@ -0,0 +1,386 @@ +""" +Haystack 3 adapters used by :class:`hayhooks.durable.context.DurableContext`. + +The adapters use only Haystack's public PipelineSnapshot, Agent, State, and +hook APIs. The imports intentionally stay lazy so the base Hayhooks install +continues to support Haystack 2 for non-durable deployments. +""" + +from __future__ import annotations + +import asyncio +import importlib +from collections.abc import Mapping +from typing import Any, cast + +from haystack.lazy_imports import LazyImport + +from hayhooks.durable.context import DurableContext +from hayhooks.durable.models import ExecutionCheckpoint, ExecutionKind, RetryableExecutionError, validate_json + +_HAYSTACK_V3_ERROR = ( + "Durable execution requires Haystack 3. Install `hayhooks[durable]` in the durable server environment." +) +_AGENT_CHECKPOINT_PHASE = "_hayhooks_agent_checkpoint_phase" +_AGENT_FINAL_PHASE = "after_run" +_AGENT_INTERNAL_STATE_KEYS = frozenset(("continue_run", "tools", "hook_context")) + + +async def _run_fenced_thread(function: Any, /, *args: Any, **kwargs: Any) -> Any: + """Keep the caller's durable claim alive until non-cancellable thread work exits.""" + task = asyncio.create_task(asyncio.to_thread(function, *args, **kwargs)) + try: + return await asyncio.shield(task) + except asyncio.CancelledError: + return await task + + +# Keep every Haystack 3-only symbol behind Haystack's supported optional-import +# boundary. This module is imported by the base package in Haystack 2 +# environments, where Agent, hooks, and snapshots intentionally do not exist. +with LazyImport(_HAYSTACK_V3_ERROR) as haystack_v3_import: + from haystack import Pipeline + from haystack.components.agents import Agent + from haystack.components.agents.state import State + from haystack.core.errors import BreakpointException, PipelineRuntimeError + from haystack.dataclasses import ChatMessage + from haystack.dataclasses.breakpoints import Breakpoint, PipelineSnapshot + + FunctionHook = importlib.import_module("haystack.hooks.from_function").FunctionHook + + +def require_haystack_v3() -> None: + """Fail durable deployment explicitly when the optional v3 extra is missing.""" + try: + haystack_v3_import.check() + import haystack + except ImportError as error: # pragma: no cover - dependency failure + raise RuntimeError(_HAYSTACK_V3_ERROR) from error + major = str(getattr(haystack, "__version__", "0")).split(".", maxsplit=1)[0] + if major != "3": + raise RuntimeError(_HAYSTACK_V3_ERROR) + + +class HaystackDurableAdapter: + """Bind a validated Haystack 3 Pipeline or Agent to execution contexts.""" + + def __init__(self, pipeline: Any, kind: ExecutionKind) -> None: + require_haystack_v3() + self.pipeline = pipeline + self.kind = kind + if kind is ExecutionKind.PIPELINE: + self._validate_pipeline() + else: + self._validate_agent() + self._install_agent_checkpoint_hooks() + + def _validate_pipeline(self) -> None: + haystack_v3_import.check() + if not isinstance(self.pipeline, Pipeline): + msg = "run_durable Pipeline wrappers must set self.pipeline to a Haystack 3 Pipeline" + raise TypeError(msg) + + def _validate_agent(self) -> None: + haystack_v3_import.check() + if not isinstance(self.pipeline, Agent): + msg = "durable Agent wrappers must set self.pipeline to a Haystack 3 Agent" + raise TypeError(msg) + + async def run_pipeline_async( + self, context: DurableContext, data: dict[str, Any], *, checkpoint_at: list[str] + ) -> dict[str, Any]: + return cast( + dict[str, Any], + await _run_fenced_thread(self.run_pipeline, context, data, checkpoint_at=checkpoint_at), + ) + + def run_pipeline( + self, context: DurableContext, data: dict[str, Any], *, checkpoint_at: list[str] + ) -> dict[str, Any]: + if self.kind is not ExecutionKind.PIPELINE: + msg = "run_pipeline is available only when self.pipeline is a Haystack Pipeline" + raise TypeError(msg) + snapshot = None + if context.record.checkpoint is not None: + checkpoint = context.record.checkpoint + if checkpoint.kind is not ExecutionKind.PIPELINE: + msg = "The persisted checkpoint is not a PipelineSnapshot" + raise TypeError(msg) + snapshot = PipelineSnapshot.from_dict(cast(dict[str, Any], checkpoint.data["snapshot"])) + boundaries = list(checkpoint_at) + if snapshot is not None: + break_point: Any = snapshot.break_point + completed_visits = snapshot.pipeline_state.component_visits + boundaries = [ + name + for name in boundaries + if completed_visits.get(name, 0) == 0 + and not (name == break_point.component_name and break_point.visit_count == 0) + ] + + next_data = data if snapshot is None else {} + try: + while boundaries: + component_name = boundaries.pop(0) + break_point = Breakpoint(component_name=component_name) + try: + return cast( + dict[str, Any], + self.pipeline.run(data=next_data, pipeline_snapshot=snapshot, break_point=break_point), + ) + except BreakpointException as error: + if error.pipeline_snapshot is None: + msg = "Haystack breakpoint did not expose a PipelineSnapshot" + raise RetryableExecutionError(msg) from error + snapshot = error.pipeline_snapshot + context.record.append_progress( + f"Checkpoint saved before pipeline component '{component_name}'", + kind="checkpoint", + ) + context._sync_await(context.checkpoint(_pipeline_checkpoint(context, snapshot))) + next_data = {} + return cast(dict[str, Any], self.pipeline.run(data=next_data, pipeline_snapshot=snapshot)) + except PipelineRuntimeError as error: + if error.pipeline_snapshot is not None: + context._sync_await(context.checkpoint(_pipeline_checkpoint(context, error.pipeline_snapshot))) + raise + + async def run_agent_async(self, context: DurableContext, *, messages: list[Any], **kwargs: Any) -> dict[str, Any]: + if self.kind is not ExecutionKind.AGENT: + msg = "run_agent_async is available only when self.pipeline is a Haystack Agent" + raise TypeError(msg) + final_result = _final_agent_result(context) + if final_result is not None: + return final_result + method = getattr(self.pipeline, "run_async", None) + if callable(method): + return cast(dict[str, Any], await method(messages=messages, **kwargs)) + return cast( + dict[str, Any], + await _run_fenced_thread(self.run_agent, context, messages=messages, **kwargs), + ) + + def run_agent(self, context: DurableContext, *, messages: list[Any], **kwargs: Any) -> dict[str, Any]: + if self.kind is not ExecutionKind.AGENT: + msg = "run_agent is available only when self.pipeline is a Haystack Agent" + raise TypeError(msg) + final_result = _final_agent_result(context) + if final_result is not None: + return final_result + return cast(dict[str, Any], self.pipeline.run(messages=messages, **kwargs)) + + def _install_agent_checkpoint_hooks(self) -> None: # noqa: C901, PLR0915 + """Install once; hooks select the active execution through ContextVar.""" + if getattr(self.pipeline, "_hayhooks_durable_hooks_installed", False): + return + + def restore_before_run(state: State) -> None: + if context := _current_durable_context(): + _restore_agent_state(context, state) + + async def restore_before_run_async(state: State) -> None: + if context := _current_durable_context(): + _restore_agent_state(context, state) + + def check_cancelled_before_llm(state: State) -> None: + del state + if context := _current_durable_context(): + context.check_cancelled_sync() + + async def check_cancelled_before_llm_async(state: State) -> None: + del state + if context := _current_durable_context(): + await context.check_cancelled() + + def checkpoint_after_tool(state: State) -> None: + context = _current_durable_context() + if context is None: + return + if not _agent_exits_after_tools(state, self.pipeline.exit_conditions): + context._sync_await(_checkpoint_agent_state(context, state)) + context.check_cancelled_sync() + + async def checkpoint_after_tool_async(state: State) -> None: + context = _current_durable_context() + if context is None: + return + if not _agent_exits_after_tools(state, self.pipeline.exit_conditions): + await _checkpoint_agent_state(context, state) + await context.check_cancelled() + + def checkpoint_on_exit(state: State) -> None: + context = _current_durable_context() + if context is not None and state.data["continue_run"]: + context._sync_await(_checkpoint_agent_state(context, state)) + + async def checkpoint_on_exit_async(state: State) -> None: + context = _current_durable_context() + if context is not None and state.data["continue_run"]: + await _checkpoint_agent_state(context, state) + + def checkpoint_after_run(state: State) -> None: + if context := _current_durable_context(): + context._sync_await(_checkpoint_agent_state(context, state, final=True)) + + async def checkpoint_after_run_async(state: State) -> None: + if context := _current_durable_context(): + await _checkpoint_agent_state(context, state, final=True) + + # The module uses postponed annotations, while Haystack validates hook + # signatures with ``inspect.signature`` rather than resolving hints. + for function in ( + restore_before_run, + restore_before_run_async, + check_cancelled_before_llm, + check_cancelled_before_llm_async, + checkpoint_after_tool, + checkpoint_after_tool_async, + checkpoint_on_exit, + checkpoint_on_exit_async, + checkpoint_after_run, + checkpoint_after_run_async, + ): + function.__annotations__["state"] = State + + hooks = dict(getattr(self.pipeline, "hooks", {}) or {}) + hooks["before_run"] = [ + FunctionHook(function=restore_before_run, async_function=restore_before_run_async), + *hooks.get("before_run", []), + ] + hooks["before_llm"] = [ + FunctionHook(function=check_cancelled_before_llm, async_function=check_cancelled_before_llm_async), + *hooks.get("before_llm", []), + ] + hooks["after_tool"] = [ + *hooks.get("after_tool", []), + FunctionHook(function=checkpoint_after_tool, async_function=checkpoint_after_tool_async), + ] + hooks["on_exit"] = [ + *hooks.get("on_exit", []), + FunctionHook(function=checkpoint_on_exit, async_function=checkpoint_on_exit_async), + ] + hooks["after_run"] = [ + *hooks.get("after_run", []), + FunctionHook(function=checkpoint_after_run, async_function=checkpoint_after_run_async), + ] + self.pipeline.hooks = hooks + self.pipeline._hayhooks_durable_hooks_installed = True + + +def _current_durable_context() -> DurableContext | None: + from hayhooks.durable.context import get_current_durable_context + + return get_current_durable_context() + + +def _pipeline_checkpoint(context: DurableContext, snapshot: Any) -> ExecutionCheckpoint: + return ExecutionCheckpoint( + ExecutionKind.PIPELINE, + {"snapshot": validate_json(snapshot.to_dict(), limit=context.record.max_record_bytes, label="snapshot")}, + ) + + +def _checkpoint_data(state: Any, context: DurableContext, *, final: bool = False) -> dict[str, Any]: + """Exclude live resources from State's otherwise public serialization.""" + payload = _without_live_agent_resources(state.to_dict()) + if final: + payload[_AGENT_CHECKPOINT_PHASE] = _AGENT_FINAL_PHASE + return cast( + dict[str, Any], + validate_json(payload, limit=context.record.max_record_bytes, label="Agent state"), + ) + + +def _without_live_agent_resources(payload: Mapping[str, Any]) -> dict[str, Any]: + """Remove per-run Agent resources from a serialized state checkpoint.""" + cleaned = dict(payload) + data = dict(cast(Mapping[str, Any], cleaned.get("data", {}))) + schema = dict(cast(Mapping[str, Any], cleaned.get("schema", {}))) + serialization_schema = dict(data.get("serialization_schema", {})) + properties = dict(serialization_schema.get("properties", {})) + serialized_data = dict(data.get("serialized_data", {})) + for key in ("tools", "hook_context"): + schema.pop(key, None) + properties.pop(key, None) + serialized_data.pop(key, None) + serialization_schema["properties"] = properties + data["serialization_schema"] = serialization_schema + data["serialized_data"] = serialized_data + cleaned["schema"] = schema + cleaned["data"] = data + return cleaned + + +async def _checkpoint_agent_state(context: DurableContext, state: Any, *, final: bool = False) -> None: + context.record.append_progress( + "Agent final checkpoint saved" if final else "Agent step checkpoint saved", kind="checkpoint" + ) + await context.checkpoint(ExecutionCheckpoint(ExecutionKind.AGENT, _checkpoint_data(state, context, final=final))) + + +def _agent_exits_after_tools(state: State, exit_conditions: list[str]) -> bool: + """Return whether Haystack will stop after the current tool-result messages.""" + if exit_conditions == ["text"]: + return False + matched = False + for message in reversed(state.data.get("messages", [])): + result = message.tool_call_result + if result is None: + break + if result.origin.tool_name not in exit_conditions: + continue + if result.error: + return False + matched = True + return matched + + +def _final_agent_result(context: DurableContext) -> dict[str, Any] | None: + """Return a checkpointed terminal Agent result without re-entering the Agent loop.""" + checkpoint = context.record.checkpoint + if ( + checkpoint is None + or checkpoint.kind is not ExecutionKind.AGENT + or checkpoint.data.get(_AGENT_CHECKPOINT_PHASE) != _AGENT_FINAL_PHASE + ): + return None + state = State.from_dict(_without_live_agent_resources(checkpoint.data)) + result = {key: value for key, value in state.data.items() if key not in _AGENT_INTERNAL_STATE_KEYS} + if messages := result.get("messages"): + result["last_message"] = messages[-1] + return result + + +def _restore_agent_state(context: DurableContext, state: Any) -> None: + """Restore a recovered State, retaining fresh per-run live resources.""" + checkpoint = context.record.checkpoint + if checkpoint is None or checkpoint.kind is not ExecutionKind.AGENT: + return + restored = State.from_dict(_without_live_agent_resources(checkpoint.data)) + live_tools = state.data.get("tools") + live_hook_context = state.data.get("hook_context") + state.data.clear() + state.data.update(restored.data) + if live_tools is not None: + state.data["tools"] = live_tools + if live_hook_context is not None: + state.data["hook_context"] = live_hook_context + resume = context.take_resume_input() + if isinstance(resume, dict) and isinstance(resume.get("messages"), list): + state.data.setdefault("messages", []).extend( + ChatMessage.from_dict(message) for message in resume["messages"] if isinstance(message, dict) + ) + + +def execution_kind(pipeline: Any) -> ExecutionKind: + """Classify a real Haystack 3 Pipeline or Agent behind the lazy boundary.""" + require_haystack_v3() + if isinstance(pipeline, Pipeline): + return ExecutionKind.PIPELINE + if isinstance(pipeline, Agent): + return ExecutionKind.AGENT + msg = "Durable wrappers must set self.pipeline to a real Haystack 3 Pipeline or Agent" + raise TypeError(msg) + + +__all__ = ["HaystackDurableAdapter", "execution_kind", "require_haystack_v3"] diff --git a/src/hayhooks/durable/backend.py b/src/hayhooks/durable/backend.py new file mode 100644 index 00000000..26d74bae --- /dev/null +++ b/src/hayhooks/durable/backend.py @@ -0,0 +1,208 @@ +"""Backend-neutral durable-store policy and contract.""" +# ruff: noqa: EM101, EM102 + +from __future__ import annotations + +import json +from collections.abc import Callable, Mapping +from dataclasses import dataclass, replace +from typing import Any, Protocol + +from hayhooks.durable.engine import ( + Checkpoint, + Complete, + ExecutionCommand, + ExecutionControl, + ExecutionPayloadSizeError, + Fail, + Heartbeat, + PayloadKind, + RequestCancellation, + Resume, + ScheduleRetry, + Suspend, + TransitionPlan, + validate_run_id, +) +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError + +MAINTENANCE_BATCH_SIZE = 100 +DEFAULT_TRANSACTION_MAX_RETRIES = 8 +DEFAULT_TRANSACTION_BACKOFF_MAX_MS = 25 + + +class ExecutionStoreCorruptionError(ExecutionStoreError): + """Persisted backend state cannot be safely decoded as durable control data.""" + + +class ExecutionContentionError(ExecutionStoreError): + """A bounded optimistic transaction could not obtain a stable snapshot.""" + + +class ExecutionIdempotencyConflictError(RuntimeError): + """A logical idempotency key was reused for a different request binding.""" + + +@dataclass(frozen=True, slots=True) +class SubmissionResult: + created: bool + control: ExecutionControl + + +@dataclass(frozen=True, slots=True) +class ExecutionStoreConfig: + key_prefix: str = "hayhooks:durable" + transaction_max_retries: int = DEFAULT_TRANSACTION_MAX_RETRIES + transaction_backoff_max_ms: int = DEFAULT_TRANSACTION_BACKOFF_MAX_MS + lease_commit_safety_ms: int = 50 + terminal_ttl_seconds: int = 604_800 + max_nonterminal_executions: int = 0 + max_input_bytes: int = 256_000 + max_checkpoint_bytes: int = 512_000 + max_result_bytes: int = 512_000 + max_error_bytes: int = 64_000 + max_wait_bytes: int = 64_000 + max_progress_events: int = 100 + max_progress_event_bytes: int = 8_192 + + def __post_init__(self) -> None: + for name in ( + "transaction_max_retries", + "terminal_ttl_seconds", + "max_input_bytes", + "max_checkpoint_bytes", + "max_result_bytes", + "max_error_bytes", + "max_wait_bytes", + "max_progress_events", + "max_progress_event_bytes", + ): + if getattr(self, name) < 1: + raise ValueError(f"{name} must be positive") + if min(self.transaction_backoff_max_ms, self.lease_commit_safety_ms, self.max_nonterminal_executions) < 0: + raise ValueError("durable limits cannot be negative") + + +class ExecutionBackend(Protocol): + """Internal operations needed by the durable adapter.""" + + config: ExecutionStoreConfig + deployment: str + + async def initialize(self) -> None: ... + + async def submit( + self, control: ExecutionControl, input_payload: bytes, *, binding_digest: str + ) -> SubmissionResult: ... + + async def get(self, run_id: str) -> ExecutionControl | None: ... + + async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> dict[PayloadKind, bytes | None]: ... + + async def read_progress(self, run_id: str) -> list[bytes]: ... + + async def transition( + self, run_id: str, command: ExecutionCommand, *, candidate: bool = False + ) -> TransitionPlan: ... + + async def read_candidate(self) -> str | None: ... + + async def maintain(self, command_factory: Callable[[int, int], ExecutionCommand]) -> int: ... + + async def operational_counts(self) -> dict[str, int]: ... + + +def parse_idempotency_binding(value: str) -> tuple[str, str]: + """Decode the execution ID and request digest stored for an idempotency key.""" + run_id, separator, binding = value.partition("|") + validate_run_id(run_id) + if not separator or not binding: + raise ValueError("idempotency binding is invalid") + return run_id, binding + + +def parse_lease_member(value: str) -> tuple[str, int]: + """Decode the execution ID and fence stored in the lease-expiry index.""" + run_id, separator, raw_fence = value.rpartition("|") + validate_run_id(run_id) + if not separator: + raise ValueError("lease index member is invalid") + fence = int(raw_fence) + if fence < 0: + raise ValueError("lease index member has a negative fence") + return run_id, fence + + +def bind_command(command: ExecutionCommand, *, now_ms: int, lease_commit_safety_ms: int) -> ExecutionCommand: + """Apply the backend clock and lease safety policy before reduction.""" + if isinstance(command, (Heartbeat, Checkpoint, ScheduleRetry, Suspend, Complete, Fail)): + return replace(command, now_ms=now_ms, lease_commit_safety_ms=lease_commit_safety_ms) + return replace(command, now_ms=now_ms) + + +def check_admission(raw: Mapping[Any, Any], config: ExecutionStoreConfig) -> None: + """Validate the optional single-deployment nonterminal cap.""" + values = {key.decode("utf-8") if isinstance(key, bytes) else str(key): int(value) for key, value in raw.items()} + nonterminal = values.get("nonterminal", 0) + if nonterminal < 0: + raise ExecutionStoreCorruptionError("capacity contains a negative counter") + if config.max_nonterminal_executions and nonterminal >= config.max_nonterminal_executions: + raise ExecutionAdmissionError("deployment_nonterminal") + + +def validate_command_payloads(command: ExecutionCommand, config: ExecutionStoreConfig) -> None: + """Reject command payloads before either backend attempts a transition.""" + checks: tuple[tuple[str, bytes | None, int], ...] = () + if isinstance(command, Checkpoint): + checks = ( + ("checkpoint", command.payload, config.max_checkpoint_bytes), + *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), + ) + elif isinstance(command, Suspend): + checks = ( + ("checkpoint", command.checkpoint, config.max_checkpoint_bytes), + ("wait", command.wait, config.max_wait_bytes), + *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), + ) + elif isinstance(command, Resume): + checks = ( + ("checkpoint", command.checkpoint, config.max_checkpoint_bytes), + *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), + ) + elif isinstance(command, RequestCancellation): + checks = tuple(("progress", event, config.max_progress_event_bytes) for event in command.progress_events) + elif isinstance(command, Complete): + checks = ( + ("result", command.result, config.max_result_bytes), + *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), + ) + elif isinstance(command, Fail): + checks = ( + ("error", command.error, config.max_error_bytes), + *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), + ) + elif isinstance(command, ScheduleRetry): + checks = (("error", command.error, config.max_error_bytes),) + for label, payload, limit in checks: + if payload is not None and len(payload) > limit: + raise ExecutionPayloadSizeError(f"{label} payload exceeds its configured byte limit") + + +def bind_progress_sequences(plan: TransitionPlan, config: ExecutionStoreConfig) -> TransitionPlan: + """Persist reducer-assigned progress sequence numbers inside event payloads.""" + if not plan.progress_events: + return plan + events = [] + for event in plan.progress_events: + try: + value = json.loads(event.data) + except (UnicodeDecodeError, json.JSONDecodeError) as error: + raise ExecutionPayloadSizeError("progress payload is not valid JSON") from error + if not isinstance(value, dict): + raise ExecutionPayloadSizeError("progress payload must be a JSON object") + value["sequence"] = event.sequence + encoded = json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode() + if len(encoded) > config.max_progress_event_bytes: + raise ExecutionPayloadSizeError("progress payload exceeds its configured byte limit") + events.append(replace(event, data=encoded)) + return replace(plan, progress_events=tuple(events)) diff --git a/src/hayhooks/durable/redis.py b/src/hayhooks/durable/redis.py new file mode 100644 index 00000000..41d90ff3 --- /dev/null +++ b/src/hayhooks/durable/redis.py @@ -0,0 +1,461 @@ +"""Redis codecs and optimistic transactions for the lean durable engine.""" +# ruff: noqa: EM101, EM102 + +from __future__ import annotations + +import asyncio +import hashlib +import random +import re +from collections.abc import Callable, Mapping +from dataclasses import asdict, replace +from typing import Any, cast + +from hayhooks.durable.backend import ( + DEFAULT_TRANSACTION_BACKOFF_MAX_MS, + DEFAULT_TRANSACTION_MAX_RETRIES, + MAINTENANCE_BATCH_SIZE, + ExecutionContentionError, + ExecutionIdempotencyConflictError, + ExecutionStoreConfig, + ExecutionStoreCorruptionError, + SubmissionResult, + bind_command, + bind_progress_sequences, + check_admission, + parse_idempotency_binding, + parse_lease_member, + validate_command_payloads, +) +from hayhooks.durable.engine import ( + MAX_CONTROL_SCALAR_BYTES, + ExecutionCommand, + ExecutionControl, + ExecutionNotFoundError, + ExecutionPayloadSizeError, + ExecutionStatus, + InvalidExecutionTransitionError, + PayloadKind, + TransitionPlan, + decide, + submission_plan, + validate_run_id, +) + +_DIGEST_PREFIX = "hayhooks-durable-v2" + + +class RedisKeys: + """Build private keys without exposing deployment or idempotency values.""" + + def __init__(self, key_prefix: str, deployment: str) -> None: + if not key_prefix.rstrip(":"): + raise ValueError("key_prefix cannot be empty") + self.deployment_digest = digest("deployment", deployment) + prefix = key_prefix.rstrip(":") + self.base = f"{prefix}:{{{self.deployment_digest}}}" + + @property + def runnable(self) -> str: + return f"{self.base}:runnable" + + @property + def lease_expiry(self) -> str: + return f"{self.base}:lease-expiry" + + @property + def capacity(self) -> str: + return f"{self.base}:capacity" + + def control(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:control" + + def input(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:input" + + def checkpoint(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:checkpoint" + + def result(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:result" + + def error(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:error" + + def progress(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:progress" + + def wait(self, run_id: str) -> str: + return f"{self._execution_base(run_id)}:wait" + + def idempotency(self, idempotency_digest: str) -> str: + if not re.fullmatch(r"[a-f0-9]{64}", idempotency_digest): + raise ValueError("idempotency digest must be a sha256 hex value") + return f"{self.base}:idem:{idempotency_digest}" + + @staticmethod + def lease_member(run_id: str, fence: int) -> str: + validate_run_id(run_id) + if fence < 0: + raise ValueError("fence cannot be negative") + return f"{run_id}|{fence}" + + def _execution_base(self, run_id: str) -> str: + validate_run_id(run_id) + return f"{self.base}:exec:{run_id}" + + +def digest(domain: str, value: str) -> str: + """Return the stable domain-separated digest used for isolated key material.""" + return hashlib.sha256(f"{_DIGEST_PREFIX}:{domain}:".encode() + value.encode()).hexdigest() + + +def encode_control(control: ExecutionControl) -> dict[str, str]: + encoded: dict[str, str] = {} + for field_name, value in asdict(control).items(): + if value is None: + continue + encoded[field_name] = value.value if isinstance(value, ExecutionStatus) else str(value) + return encoded + + +def decode_control(values: Mapping[str | bytes, str | bytes | int]) -> ExecutionControl: + decoded = {_text(key): _text(value) for key, value in values.items()} + required = { + "run_id", + "idempotency_digest", + "idempotency_binding_digest", + "deployment", + "definition_revision", + "kind", + "status", + "version", + "fence", + "run_attempt", + "application_retry_count", + "progress_sequence", + "created_at_ms", + "updated_at_ms", + } + missing = required.difference(decoded) + if missing: + raise ExecutionStoreCorruptionError(f"control Hash is missing required fields: {', '.join(sorted(missing))}") + try: + status = ExecutionStatus(decoded["status"]) + except ValueError as error: + raise ExecutionStoreCorruptionError("control Hash has an unknown status") from error + for name in ( + "run_id", + "idempotency_digest", + "idempotency_binding_digest", + "deployment", + "definition_revision", + "kind", + ): + if not decoded[name] or len(decoded[name].encode()) > MAX_CONTROL_SCALAR_BYTES: + raise ExecutionStoreCorruptionError(f"control Hash has invalid {name}") + for name in ("owner_id", "lease_owner", "cancel_reason"): + if name in decoded and len(decoded[name].encode()) > MAX_CONTROL_SCALAR_BYTES: + raise ExecutionStoreCorruptionError(f"control Hash has oversized {name}") + integers = ( + "version", + "fence", + "run_attempt", + "application_retry_count", + "progress_sequence", + "created_at_ms", + "updated_at_ms", + ) + numeric: dict[str, int | None] = {name: _nonnegative_int(decoded[name], name) for name in integers} + for name in ("available_at_ms", "lease_expires_at_ms", "cancel_requested_at_ms"): + numeric[name] = _nonnegative_int(decoded[name], name) if name in decoded else None + try: + return ExecutionControl( + run_id=decoded["run_id"], + idempotency_digest=decoded["idempotency_digest"], + idempotency_binding_digest=decoded["idempotency_binding_digest"], + deployment=decoded["deployment"], + definition_revision=decoded["definition_revision"], + owner_id=decoded.get("owner_id"), + kind=decoded["kind"], + status=status, + lease_owner=decoded.get("lease_owner"), + cancel_reason=decoded.get("cancel_reason"), + **cast(Any, numeric), + ) + except (TypeError, ValueError, ExecutionPayloadSizeError) as error: + raise ExecutionStoreCorruptionError("control Hash violates durable invariants") from error + + +async def redis_time_ms(pipe: Any) -> int: + seconds, micros = await pipe.time() + return int(seconds) * 1_000 + int(micros) // 1_000 + + +class RedisExecutionStore: + """Redis implementation using one due-time ZSET and one lease ZSET.""" + + def __init__(self, redis: Any, *, deployment: str, config: ExecutionStoreConfig | None = None) -> None: + self.redis = redis + self.config = config or ExecutionStoreConfig() + self.keys = RedisKeys(self.config.key_prefix, deployment) + self.deployment = deployment + + async def initialize(self) -> None: + try: + info = await self.redis.info("server") + version = tuple(int(piece) for piece in _text(info["redis_version"]).split(".")[:2]) + except Exception as error: + raise RuntimeError("unable to validate Redis server capabilities") from error + if version < (6, 2): + raise RuntimeError("durable Redis requires Redis 6.2 or later") + + async def submit( # noqa: C901 + self, control: ExecutionControl, input_payload: bytes, *, binding_digest: str + ) -> SubmissionResult: + if control.deployment != self.deployment: + raise ValueError("control deployment does not match this store") + if len(input_payload) > self.config.max_input_bytes: + raise ExecutionPayloadSizeError("input payload exceeds configured size") + idem_key = self.keys.idempotency(control.idempotency_digest) + control_key = self.keys.control(control.run_id) + for attempt in range(self.config.transaction_max_retries): + async with self.redis.pipeline(transaction=True) as pipe: + try: + watch_keys = [idem_key, control_key] + if self.config.max_nonterminal_executions: + watch_keys.append(self.keys.capacity) + await pipe.watch(*watch_keys) + existing = await pipe.get(idem_key) + if existing is not None: + try: + existing_run, existing_binding = parse_idempotency_binding(_text(existing)) + except ValueError as error: + raise ExecutionStoreCorruptionError("idempotency binding has an invalid format") from error + mapped_control_key = self.keys.control(existing_run) + await pipe.watch(mapped_control_key) + raw_control = await pipe.hgetall(mapped_control_key) + if existing_binding != binding_digest: + raise ExecutionIdempotencyConflictError("idempotency key is bound to another request") + if raw_control: + return SubmissionResult(created=False, control=decode_control(raw_control)) + if self.config.max_nonterminal_executions: + check_admission(await pipe.hgetall(self.keys.capacity), self.config) + now_ms = await redis_time_ms(pipe) + candidate = replace(control, created_at_ms=now_ms, updated_at_ms=now_ms) + plan = submission_plan(candidate, input_payload) + pipe.multi() + if existing is not None: + pipe.delete(idem_key) + pipe.set(idem_key, f"{candidate.run_id}|{binding_digest}") + self._apply_plan(pipe, candidate, plan, new_submission=True) + await pipe.execute() + return SubmissionResult(created=True, control=candidate) + except redis_watch_error(): + await self._backoff(attempt) + raise ExecutionContentionError("submission transaction retry budget exhausted") + + async def get(self, run_id: str) -> ExecutionControl | None: + values = await self.redis.hgetall(self.keys.control(run_id)) + return decode_control(values) if values else None + + async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> dict[PayloadKind, bytes | None]: + if not kinds: + return {} + async with self.redis.pipeline(transaction=False) as pipe: + for kind in kinds: + pipe.get(self._payload_key(run_id, kind)) + values = await pipe.execute() + return {kind: bytes(value) if value is not None else None for kind, value in zip(kinds, values, strict=True)} + + async def read_progress(self, run_id: str) -> list[bytes]: + return [bytes(value) for value in await self.redis.lrange(self.keys.progress(run_id), 0, -1)] + + async def transition(self, run_id: str, command: ExecutionCommand, *, candidate: bool = False) -> TransitionPlan: + validate_command_payloads(command, self.config) + control_key = self.keys.control(run_id) + for attempt in range(self.config.transaction_max_retries): + async with self.redis.pipeline(transaction=True) as pipe: + try: + watch_keys = [control_key] + if candidate: + watch_keys.append(self.keys.runnable) + await pipe.watch(*watch_keys) + current_values = await pipe.hgetall(control_key) + if not current_values: + if candidate: + pipe.multi() + pipe.zrem(self.keys.runnable, run_id) + await pipe.execute() + raise ExecutionNotFoundError(f"execution '{run_id}' was not found") + current = decode_control(current_values) + try: + plan = bind_progress_sequences( + decide( + current, + bind_command( + command, + now_ms=await redis_time_ms(pipe), + lease_commit_safety_ms=self.config.lease_commit_safety_ms, + ), + ), + self.config, + ) + except InvalidExecutionTransitionError: + if not candidate: + raise + pipe.multi() + pipe.zrem(self.keys.runnable, run_id) + if current.status is ExecutionStatus.QUEUED: + pipe.zadd(self.keys.runnable, {run_id: _runnable_score(current)}) + await pipe.execute() + return TransitionPlan(current) + pipe.multi() + self._apply_plan(pipe, current, plan) + await pipe.execute() + return plan + except redis_watch_error(): + await self._backoff(attempt) + raise ExecutionContentionError("execution transition retry budget exhausted") + + async def read_candidate(self) -> str | None: + now_ms = await self._time_ms() + members = await self.redis.zrangebyscore(self.keys.runnable, "-inf", now_ms, start=0, num=1) + return _text(members[0]) if members else None + + async def maintain(self, command_factory: Callable[[int, int], ExecutionCommand]) -> int: + now_ms = await self._time_ms() + entries = await self.redis.zrangebyscore( + self.keys.lease_expiry, + "-inf", + now_ms, + start=0, + num=MAINTENANCE_BATCH_SIZE, + withscores=True, + ) + recovered = 0 + for member, deadline in entries: + try: + run_id, fence = parse_lease_member(_text(member)) + await self.transition(run_id, command_factory(fence, int(deadline))) + recovered += 1 + except (ExecutionNotFoundError, ValueError): + await self.redis.zrem(self.keys.lease_expiry, member) + return recovered + + async def operational_counts(self) -> dict[str, int]: + async with self.redis.pipeline(transaction=False) as pipe: + pipe.hget(self.keys.capacity, "nonterminal") + pipe.zcard(self.keys.runnable) + pipe.zcard(self.keys.lease_expiry) + nonterminal, runnable, leases = await pipe.execute() + return { + "nonterminal": int(_text(nonterminal)) if nonterminal is not None else 0, + "runnable": int(runnable), + "lease_expiry": int(leases), + } + + def _apply_plan( # noqa: C901 + self, pipe: Any, current: ExecutionControl, plan: TransitionPlan, *, new_submission: bool = False + ) -> None: + next_control = plan.next_control + current_fields = encode_control(current) + next_fields = encode_control(next_control) + pipe.hset(self.keys.control(next_control.run_id), mapping=next_fields) + removed_fields = tuple(set(current_fields).difference(next_fields)) + if removed_fields: + pipe.hdel(self.keys.control(next_control.run_id), *removed_fields) + for write in plan.payload_writes: + pipe.set(self._payload_key(next_control.run_id, write.kind), write.data) + for kind in plan.payload_deletes: + pipe.delete(self._payload_key(next_control.run_id, kind)) + for event in plan.progress_events: + pipe.rpush(self.keys.progress(next_control.run_id), event.data) + pipe.ltrim(self.keys.progress(next_control.run_id), -self.config.max_progress_events, -1) + + pipe.zrem(self.keys.runnable, next_control.run_id) + if next_control.status is ExecutionStatus.QUEUED: + pipe.zadd(self.keys.runnable, {next_control.run_id: _runnable_score(next_control)}) + + if plan.lease_index_update is not None: + member = RedisKeys.lease_member(next_control.run_id, plan.lease_index_update.fence) + if plan.lease_index_update.deadline_ms is None: + pipe.zrem(self.keys.lease_expiry, member) + else: + pipe.zadd(self.keys.lease_expiry, {member: plan.lease_index_update.deadline_ms}) + + if new_submission: + pipe.hincrby(self.keys.capacity, "nonterminal", 1) + elif not current.terminal and next_control.terminal: + pipe.hincrby(self.keys.capacity, "nonterminal", -1) + for key in self._execution_keys(next_control.run_id): + pipe.expire(key, self.config.terminal_ttl_seconds) + pipe.expire(self.keys.idempotency(next_control.idempotency_digest), self.config.terminal_ttl_seconds) + + async def _backoff(self, attempt: int) -> None: + await redis_transaction_backoff( + attempt, + max_retries=self.config.transaction_max_retries, + max_backoff_ms=self.config.transaction_backoff_max_ms, + ) + + async def _time_ms(self) -> int: + return await redis_time_ms(self.redis) + + def _execution_keys(self, run_id: str) -> tuple[str, ...]: + return ( + self.keys.control(run_id), + self.keys.input(run_id), + self.keys.checkpoint(run_id), + self.keys.result(run_id), + self.keys.error(run_id), + self.keys.progress(run_id), + self.keys.wait(run_id), + ) + + def _payload_key(self, run_id: str, kind: PayloadKind) -> str: + return { + PayloadKind.INPUT: self.keys.input, + PayloadKind.CHECKPOINT: self.keys.checkpoint, + PayloadKind.RESULT: self.keys.result, + PayloadKind.ERROR: self.keys.error, + PayloadKind.WAIT: self.keys.wait, + }[kind](run_id) + + +def _runnable_score(control: ExecutionControl) -> int: + return control.available_at_ms if control.available_at_ms is not None else control.updated_at_ms + + +def _text(value: str | bytes | int) -> str: + return value.decode() if isinstance(value, bytes) else str(value) + + +def _nonnegative_int(value: str, name: str) -> int: + try: + parsed = int(value) + except ValueError as error: + raise ExecutionStoreCorruptionError(f"control Hash has non-integer {name}") from error + if parsed < 0: + raise ExecutionStoreCorruptionError(f"control Hash has negative {name}") + return parsed + + +async def redis_transaction_backoff( + attempt: int, + *, + max_retries: int = DEFAULT_TRANSACTION_MAX_RETRIES, + max_backoff_ms: int = DEFAULT_TRANSACTION_BACKOFF_MAX_MS, +) -> None: + """Apply the bounded jitter shared by Redis optimistic transactions.""" + if attempt + 1 < max_retries and max_backoff_ms: + await asyncio.sleep(random.uniform(0, max_backoff_ms) / 1_000) # noqa: S311 + + +def redis_watch_error() -> type[Exception]: + """Load redis-py's optimistic-transaction conflict lazily.""" + try: + from redis.exceptions import WatchError + except ImportError as error: # pragma: no cover + raise RuntimeError("Redis durable storage requires the redis package") from error + return WatchError diff --git a/src/hayhooks/durable/reference.py b/src/hayhooks/durable/reference.py new file mode 100644 index 00000000..e847a1dc --- /dev/null +++ b/src/hayhooks/durable/reference.py @@ -0,0 +1,181 @@ +"""Deterministic in-memory implementation of the lean durable backend.""" +# ruff: noqa: EM101, EM102 + +from __future__ import annotations + +import time +from collections.abc import Callable + +from hayhooks.durable.backend import ( + MAINTENANCE_BATCH_SIZE, + ExecutionIdempotencyConflictError, + ExecutionStoreConfig, + SubmissionResult, + bind_command, + bind_progress_sequences, + check_admission, + parse_idempotency_binding, + parse_lease_member, + validate_command_payloads, +) +from hayhooks.durable.engine import ( + ExecutionCommand, + ExecutionControl, + ExecutionNotFoundError, + ExecutionStatus, + InvalidExecutionTransitionError, + PayloadKind, + TransitionPlan, + decide, + submission_plan, +) + + +class InMemoryExecutionStore: + """Single-process reference model used by the shared backend contract.""" + + def __init__(self, *, deployment: str, config: ExecutionStoreConfig | None = None) -> None: + self.deployment = deployment + self.config = config or ExecutionStoreConfig() + self._controls: dict[str, ExecutionControl] = {} + self._payloads: dict[str, dict[PayloadKind, bytes]] = {} + self._progress: dict[str, list[bytes]] = {} + self._runnable: dict[str, int] = {} + self._lease_expiry: dict[str, int] = {} + self._capacity = {"nonterminal": 0} + self._idempotency: dict[str, str] = {} + self._terminal_cleanup: dict[str, tuple[int, str, str]] = {} + + @staticmethod + def _now_ms() -> int: + return round(time.time() * 1_000) + + async def initialize(self) -> None: + return None + + async def submit(self, control: ExecutionControl, input_payload: bytes, *, binding_digest: str) -> SubmissionResult: + if control.deployment != self.deployment: + raise ValueError("control deployment does not match this store") + if len(input_payload) > self.config.max_input_bytes: + raise ValueError("input payload exceeds configured size") + existing = self._idempotency.get(control.idempotency_digest) + if existing is not None: + existing_run, existing_binding = parse_idempotency_binding(existing) + existing_control = self._controls.get(existing_run) + if existing_binding != binding_digest: + raise ExecutionIdempotencyConflictError("idempotency key is bound to another request") + if existing_control is not None: + return SubmissionResult(created=False, control=existing_control) + check_admission(self._capacity, self.config) + self._idempotency[control.idempotency_digest] = f"{control.run_id}|{binding_digest}" + self._apply_plan(control, submission_plan(control, input_payload)) + self._capacity["nonterminal"] += 1 + return SubmissionResult(created=True, control=control) + + async def get(self, run_id: str) -> ExecutionControl | None: + return self._controls.get(run_id) + + async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> dict[PayloadKind, bytes | None]: + return {kind: self._payloads.get(run_id, {}).get(kind) for kind in kinds} + + async def read_progress(self, run_id: str) -> list[bytes]: + return list(self._progress.get(run_id, ())) + + async def transition(self, run_id: str, command: ExecutionCommand, *, candidate: bool = False) -> TransitionPlan: + current = self._controls.get(run_id) + if current is None: + raise ExecutionNotFoundError(f"execution '{run_id}' was not found") + command = bind_command( + command, now_ms=self._now_ms(), lease_commit_safety_ms=self.config.lease_commit_safety_ms + ) + validate_command_payloads(command, self.config) + try: + plan = bind_progress_sequences(decide(current, command), self.config) + except InvalidExecutionTransitionError: + if not candidate: + raise + self._runnable.pop(run_id, None) + if current.status is ExecutionStatus.QUEUED: + self._runnable[run_id] = _runnable_score(current) + return TransitionPlan(current) + self._apply_plan(current, plan) + return plan + + async def read_candidate(self) -> str | None: + now_ms = self._now_ms() + due = ((score, run_id) for run_id, score in self._runnable.items() if score <= now_ms) + try: + return min(due)[1] + except ValueError: + return None + + async def maintain(self, command_factory: Callable[[int, int], ExecutionCommand]) -> int: + now_ms = self._now_ms() + recovered = 0 + for member, deadline in sorted(self._lease_expiry.items(), key=lambda item: item[1])[:MAINTENANCE_BATCH_SIZE]: + if deadline > now_ms: + break + run_id, fence = parse_lease_member(member) + try: + await self.transition(run_id, command_factory(fence, deadline)) + recovered += 1 + except ExecutionNotFoundError: + self._lease_expiry.pop(member, None) + self._cleanup_terminal(now_ms) + return recovered + + async def operational_counts(self) -> dict[str, int]: + return { + "nonterminal": self._capacity["nonterminal"], + "runnable": len(self._runnable), + "lease_expiry": len(self._lease_expiry), + } + + def _cleanup_terminal(self, now_ms: int) -> None: + for run_id, (expires_at, idem_digest, idem_value) in tuple(self._terminal_cleanup.items()): + if expires_at > now_ms: + continue + self._terminal_cleanup.pop(run_id) + if self._idempotency.get(idem_digest) == idem_value: + self._idempotency.pop(idem_digest) + self._controls.pop(run_id, None) + self._payloads.pop(run_id, None) + self._progress.pop(run_id, None) + + def _apply_plan(self, current: ExecutionControl, plan: TransitionPlan) -> None: + next_control = plan.next_control + self._controls[next_control.run_id] = next_control + payloads = self._payloads.setdefault(next_control.run_id, {}) + for write in plan.payload_writes: + payloads[write.kind] = write.data + for kind in plan.payload_deletes: + payloads.pop(kind, None) + if plan.progress_events: + progress = self._progress.setdefault(next_control.run_id, []) + progress.extend(event.data for event in plan.progress_events) + del progress[: -self.config.max_progress_events] + + self._runnable.pop(next_control.run_id, None) + if next_control.status is ExecutionStatus.QUEUED: + self._runnable[next_control.run_id] = _runnable_score(next_control) + + if plan.lease_index_update is not None: + member = f"{next_control.run_id}|{plan.lease_index_update.fence}" + if plan.lease_index_update.deadline_ms is None: + self._lease_expiry.pop(member, None) + else: + self._lease_expiry[member] = plan.lease_index_update.deadline_ms + + if not current.terminal and next_control.terminal: + self._capacity["nonterminal"] -= 1 + if self._capacity["nonterminal"] < 0: + raise RuntimeError("reference capacity counter underflow") + self._terminal_cleanup[next_control.run_id] = ( + next_control.updated_at_ms + self.config.terminal_ttl_seconds * 1_000, + next_control.idempotency_digest, + f"{next_control.run_id}|{next_control.idempotency_binding_digest}", + ) + + +def _runnable_score(control: ExecutionControl) -> int: + return control.available_at_ms if control.available_at_ms is not None else control.updated_at_ms diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py new file mode 100644 index 00000000..1e0f1307 --- /dev/null +++ b/src/hayhooks/durable/store.py @@ -0,0 +1,706 @@ +"""Application-facing durable store backed by the simplified engine.""" + +from __future__ import annotations + +import asyncio +import json +import time +from collections.abc import Awaitable, Mapping +from contextlib import suppress +from datetime import datetime, timezone +from typing import Any, cast + +from hayhooks.durable.backend import ExecutionBackend, ExecutionContentionError, ExecutionStoreConfig, SubmissionResult +from hayhooks.durable.context import RESUME_INPUT_KEY +from hayhooks.durable.engine import ( + Checkpoint, + Claim, + Complete, + ExecutionControl, + ExecutionLeaseLostError, + ExecutionNotFoundError, + Fail, + Heartbeat, + InvalidExecutionTransitionError, + PayloadKind, + RecoverExpiredLease, + RequestCancellation, + Resume, + ScheduleRetry, + Suspend, + initial_control, + normalize_cancellation_reason, +) +from hayhooks.durable.engine import ExecutionStatus as EngineStatus +from hayhooks.durable.models import ( + DEFAULT_MAX_PROGRESS_BYTES, + ExecutionError, + ExecutionKind, + ExecutionRecord, + ExecutionRecordSizeError, + ExecutionStatus, + ExecutionStoreError, + JsonValue, +) +from hayhooks.durable.redis import RedisExecutionStore, digest +from hayhooks.durable.reference import InMemoryExecutionStore +from hayhooks.server.logger import log +from hayhooks.settings import AppSettings, settings + +_RECORD_PAYLOADS = ( + PayloadKind.INPUT, + PayloadKind.CHECKPOINT, + PayloadKind.RESULT, + PayloadKind.ERROR, + PayloadKind.WAIT, +) + + +def _encode(value: Any, *, limit: int, label: str) -> bytes: + try: + encoded = json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode("utf-8") + except (TypeError, ValueError) as error: + msg = f"{label} is not JSON serializable" + raise ExecutionRecordSizeError(msg) from error + if len(encoded) > limit: + msg = f"{label} exceeds the {limit}-byte durable execution limit" + raise ExecutionRecordSizeError(msg) + return encoded + + +def _decode(payload: bytes, *, label: str) -> Any: + try: + return json.loads(payload) + except (UnicodeDecodeError, json.JSONDecodeError) as error: + msg = f"durable {label} payload is invalid" + raise RuntimeError(msg) from error + + +def _datetime(ms: int) -> datetime: + return datetime.fromtimestamp(ms / 1_000, tz=timezone.utc) + + +def _error_from_payload(payload: bytes | None) -> ExecutionError | None: + if payload is None: + return None + # The engine may terminalize an incompatible revision before application + # code has had a chance to build an ``ExecutionError``. Keep that + # engine-owned reason visible through the public record rather than + # treating an otherwise healthy terminal record as corrupt. + try: + decoded = _decode(payload, label="error") + except RuntimeError: + try: + decoded = payload.decode("utf-8") + except UnicodeDecodeError as error: + msg = "durable error payload is invalid" + raise RuntimeError(msg) from error + if isinstance(decoded, Mapping): + return ExecutionError.from_dict(cast(Mapping[str, Any], decoded)) + return ExecutionError(type="DurableExecutionError", message=str(decoded)) + + +class ExecutionClaim: + """Application-facing fenced claim backed by a durable control fence.""" + + def __init__( + self, + store: ExecutionStore, + control: Any, + record: ExecutionRecord, + worker_id: str, + confirmed_at: float, + ) -> None: + self.store = store + self.control = control + self._record = record + self.worker_id = worker_id + self._heartbeat: asyncio.Task[None] | None = None + self._finished = False + self._lost = False + self._lost_event = asyncio.Event() + self._confirmed_until = confirmed_at + self.store.lease_safe_duration + self._persisted_progress_sequence = control.progress_sequence + + @property + def record(self) -> ExecutionRecord: + return self._record + + @property + def lost_event(self) -> asyncio.Event: + return self._lost_event + + async def __aenter__(self) -> ExecutionClaim: + await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) + self._heartbeat = asyncio.create_task( + self._heartbeat_loop(), + name=f"durable-heartbeat:{self.record.execution_id}", + ) + return self + + async def __aexit__(self, exc_type: Any, exc: Any, traceback: Any) -> None: + if self._heartbeat is not None: + self._heartbeat.cancel() + with suppress(asyncio.CancelledError): + await self._heartbeat + + async def checkpoint(self) -> None: + self._ensure_owned() + await self._transition( + Checkpoint( + self.control.fence, + self.worker_id, + 0, + self.store.lease_duration_ms, + self.store._snapshot(self.record), + self._new_progress(), + ) + ) + + async def cancellation_requested(self) -> bool: + control = await self.store._core_call( + "read cancellation state", + self.store.core.get(self.record.execution_id), + ) + if control is None: + self._mark_lost() + msg = f"Execution '{self.record.execution_id}' no longer exists" + raise ExecutionLeaseLostError(msg) + if control.cancel_requested_at_ms is not None: + self.record.cancel_requested_at = _datetime(control.cancel_requested_at_ms) + self.record.cancel_reason = control.cancel_reason + return True + return False + + async def complete(self) -> None: + self._ensure_owned() + if self.record.status is ExecutionStatus.CANCELED and self.control.cancel_requested_at_ms is None: + cancellation = await self.store._core_call( + "request cancellation", + self.store.core.transition( + self.record.execution_id, + RequestCancellation(0, self.record.cancel_reason), + ), + ) + self._sync(cancellation.next_control) + if self.record.status is ExecutionStatus.FAILED: + error = self.record.error or ExecutionError(type="ExecutionError", message="Execution failed") + await self._transition( + Fail( + self.control.fence, + self.worker_id, + 0, + _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="error"), + self._new_progress(), + ) + ) + else: + await self._transition( + Complete( + self.control.fence, + self.worker_id, + 0, + _encode(self.record.result, limit=self.store.config.max_result_bytes, label="result"), + self._new_progress(), + ) + ) + self._finished = True + + async def suspend(self) -> None: + self._ensure_owned() + await self._transition( + Suspend( + self.control.fence, + self.worker_id, + 0, + self.store._snapshot(self.record), + _encode(self.record.wait, limit=self.store.config.max_wait_bytes, label="wait"), + self._new_progress(), + ) + ) + self._finished = True + + async def retry(self, error: ExecutionError, *, delay: float) -> None: + self._ensure_owned() + await self._transition( + ScheduleRetry( + self.control.fence, + self.worker_id, + 0, + max(0, round(delay * 1_000)), + self.store.max_application_retries, + _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="retry error"), + ) + ) + self._finished = True + + async def _transition(self, command: Any) -> Any: + confirmed_at = time.monotonic() + try: + plan = await self.store._core_call( + "persist execution transition", + self.store.core.transition(self.record.execution_id, command), + ) + except ExecutionLeaseLostError: + self._mark_lost() + raise + self._sync(plan.next_control, confirmed_at=confirmed_at) + return plan + + async def _heartbeat_loop(self) -> None: + while not self._finished and not self._lost: + await asyncio.sleep(self.store.heartbeat_interval) + try: + await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) + except ExecutionLeaseLostError: + return + except Exception: + if time.monotonic() >= self._confirmed_until: + self._mark_lost() + return + continue + + def _new_progress(self) -> tuple[bytes, ...]: + events = tuple(event for event in self.record.progress if event.sequence > self._persisted_progress_sequence) + return tuple( + _encode(event.to_dict(), limit=self.store.config.max_progress_event_bytes, label="progress") + for event in events + ) + + def _sync(self, control: Any, *, confirmed_at: float | None = None) -> None: + self.control = control + self.record.attempt = control.run_attempt + self.record.sequence = control.version + self.record.status = control.status + self._persisted_progress_sequence = control.progress_sequence + if confirmed_at is not None: + self._confirmed_until = confirmed_at + self.store.lease_safe_duration + + def _ensure_owned(self) -> None: + if self._lost: + msg = f"Execution lease for '{self.record.execution_id}' was lost" + raise ExecutionLeaseLostError(msg) + + def _mark_lost(self) -> None: + if not self._lost: + self._lost = True + self._lost_event.set() + + +class ExecutionStore: + """Public durable store contract backed by control records and payloads.""" + + def __init__( # noqa: PLR0913 + self, + core: ExecutionBackend, + *, + definition_revision: str | None = None, + lease_duration_ms: int = 30_000, + max_run_attempts: int = 3, + max_progress_events: int = 100, + max_record_bytes: int = 1_000_000, + ) -> None: + self.core = core + self.config = core.config + self.definition_revision = definition_revision + self.lease_duration_ms = lease_duration_ms + self.max_run_attempts = max_run_attempts + self.max_application_retries = max(0, max_run_attempts - 1) + self.max_progress_events = max_progress_events + self.max_record_bytes = max_record_bytes + if self.config.lease_commit_safety_ms >= self.lease_duration_ms: + msg = "lease_commit_safety_ms must be smaller than lease_duration_ms" + raise ValueError(msg) + self.heartbeat_interval = max(0.01, self.lease_duration_ms / 3_000) + if self.lease_duration_ms / 1_000 - self.config.lease_commit_safety_ms / 1_000 <= self.heartbeat_interval: + msg = "lease duration minus commit safety must exceed the heartbeat interval" + raise ValueError(msg) + self.lease_safe_duration = max( + 0.01, + self.lease_duration_ms / 1_000 - self.config.lease_commit_safety_ms / 1_000, + ) + + async def initialize(self) -> None: + """Initialize and validate the backing execution store.""" + await self._core_call("initialize durable store", self.core.initialize()) + + async def submit(self, record: ExecutionRecord) -> bool: + created, _ = await self.submit_with_record(record) + return created + + async def submit_with_record(self, record: ExecutionRecord) -> tuple[bool, ExecutionRecord]: + """Atomically submit a record or return its idempotent predecessor.""" + input_payload = self._input(record) + binding_digest = digest("binding", record.operation_fingerprint) + control = initial_control( + run_id=record.execution_id, + idempotency_digest=digest("idempotency", record.execution_id), + idempotency_binding_digest=binding_digest, + deployment=record.deployment_name, + definition_revision=record.definition_revision, + owner_id=record.owner_id, + kind=record.execution_kind.value, + now_ms=round(time.time() * 1_000), + ) + result: SubmissionResult = await self._core_call( + "submit execution", + self.core.submit(control, input_payload, binding_digest=binding_digest), + ) + view = await self._read_view(result.control.run_id) + if view is None: + msg = f"Submitted execution '{result.control.run_id}' is not available" + raise ExecutionStoreError(msg) + return result.created, view[1] + + async def get(self, execution_id: str) -> ExecutionRecord | None: + view = await self._read_view(execution_id) + return view[1] if view is not None else None + + async def claim_next(self, worker_name: str) -> ExecutionClaim | None: + """Claim one due execution with a new ownership fence, if available.""" + run_id = await self._core_call("read runnable candidate", self.core.read_candidate()) + if run_id is None: + return None + try: + confirmed_at = time.monotonic() + plan = await self._core_call( + "claim execution", + self.core.transition( + run_id, + Claim( + worker_name, 0, self.lease_duration_ms, self.max_run_attempts, self.definition_revision or "" + ), + candidate=True, + ), + ) + except (ExecutionLeaseLostError, ExecutionNotFoundError, InvalidExecutionTransitionError): + return None + if plan.next_control.status is not EngineStatus.RUNNING: + return None + view = await self._read_view(run_id) + if view is None: + return None + current, record = view + if ( + current.status is not EngineStatus.RUNNING + or current.fence != plan.next_control.fence + or current.lease_owner != worker_name + ): + return None + return ExecutionClaim(self, current, record, worker_name, confirmed_at) + + async def request_cancel(self, execution_id: str, reason: str | None = None) -> bool: + """Persist a cancellation request, returning whether it was accepted.""" + control = await self._core_call("read execution for cancellation", self.core.get(execution_id)) + if control is None: + return False + if control.terminal: + return control.status is EngineStatus.CANCELED + event = _encode( + { + "sequence": control.progress_sequence + 1, + "message": "Cancellation requested", + "timestamp": datetime.now(timezone.utc).isoformat(), + "kind": "cancellation_requested", + "metadata": {}, + }, + limit=self.config.max_progress_event_bytes, + label="progress", + ) + await self._core_call( + "request cancellation", + self.core.transition( + execution_id, + RequestCancellation(0, normalize_cancellation_reason(reason), (event,)), + ), + ) + return True + + async def resume(self, execution_id: str, update: JsonValue | None = None) -> bool: + """Resume a waiting execution with an optional JSON-safe application update.""" + view = await self._read_view(execution_id) + if view is None: + return False + control, record = view + if control.status is not EngineStatus.WAITING: + return False + if update is not None: + record.application_state[RESUME_INPUT_KEY] = update + record.wait = None + event = record.append_progress("Execution resumed", kind="resumed") + try: + plan = await self._core_call( + "resume execution", + self.core.transition( + execution_id, + Resume( + 0, + self.definition_revision or control.definition_revision, + self._snapshot(record), + (_encode(event.to_dict(), limit=self.config.max_progress_event_bytes, label="progress"),), + ), + ), + ) + except InvalidExecutionTransitionError: + return False + return plan.next_control.status is EngineStatus.QUEUED + + def set_definition_revision(self, definition_revision: str) -> None: + """Set the revision accepted by future claims and resumes.""" + self.definition_revision = definition_revision + + async def maintain(self) -> None: + """Recover expired leases and repair their derived indexes.""" + recover = lambda fence, deadline: RecoverExpiredLease( # noqa: E731 + 0, + fence, + deadline, + self.max_run_attempts, + self.definition_revision or "", + ) + recovered = await self._core_call("maintain durable indexes", self.core.maintain(recover)) + if recovered: + log.bind(deployment=self.core.deployment, recovered=recovered).debug( + "Recovered expired durable execution leases" + ) + + async def operational_counts(self) -> dict[str, int]: + """Return authoritative nonterminal, runnable, and lease-index counts.""" + return await self._core_call("read durable operational counts", self.core.operational_counts()) + + async def _core_call(self, operation: str, awaitable: Awaitable[Any]) -> Any: + """Expose backend outages through the public retryable-store contract.""" + try: + return await awaitable + except Exception as error: + if _is_redis_error(error): + msg = f"durable Redis store failed while attempting to {operation}" + raise ExecutionStoreError(msg) from error + raise + + async def _read_view( + self, execution_id: str, *, control: ExecutionControl | None = None + ) -> tuple[ExecutionControl, ExecutionRecord] | None: + """Read one control/payload view that was stable for its construction.""" + for _ in range(self.config.transaction_max_retries): + control = control or await self._core_call("read execution control", self.core.get(execution_id)) + if control is None: + return None + payloads, progress = await asyncio.gather( + self._core_call("read execution payloads", self.core.read_payloads(control.run_id, _RECORD_PAYLOADS)), + self._core_call("read execution progress", self.core.read_progress(control.run_id)), + ) + current = await self._core_call("recheck execution control", self.core.get(execution_id)) + if current is None: + return None + if current.version != control.version: + control = None + continue + + # A stable control must retain the payload that explains its + # terminal or waiting state; otherwise the view is corrupt. + missing = PayloadKind.INPUT.value if payloads[PayloadKind.INPUT] is None else None + if current.status is EngineStatus.COMPLETED: + missing = missing or (PayloadKind.RESULT.value if payloads[PayloadKind.RESULT] is None else None) + elif current.status is EngineStatus.FAILED: + missing = missing or (PayloadKind.ERROR.value if payloads[PayloadKind.ERROR] is None else None) + elif current.status is EngineStatus.WAITING: + missing = missing or (PayloadKind.WAIT.value if payloads[PayloadKind.WAIT] is None else None) + if missing is not None: + msg = f"Execution '{execution_id}' is missing its required {missing} payload" + raise RuntimeError(msg) + return current, self._record(current, payloads, progress) + msg = "execution changed while its durable view was being read" + raise ExecutionContentionError(msg) + + def _record( + self, + control: ExecutionControl, + payloads: Mapping[PayloadKind, bytes | None], + progress_payloads: list[bytes], + ) -> ExecutionRecord: + """Translate one stable engine view into the established public record.""" + input_payload = payloads[PayloadKind.INPUT] + if input_payload is None: + msg = f"Execution '{control.run_id}' is missing its immutable input payload" + raise RuntimeError(msg) + input_data = _decode(input_payload, label="input") + checkpoint_payload = payloads[PayloadKind.CHECKPOINT] + snapshot = _decode(checkpoint_payload, label="checkpoint") if checkpoint_payload is not None else {} + progress = [_decode(event, label="progress") for event in progress_payloads] + result_payload = payloads[PayloadKind.RESULT] + result = _decode(result_payload, label="result") if result_payload is not None else None + error = _error_from_payload(payloads[PayloadKind.ERROR]) + wait_payload = payloads[PayloadKind.WAIT] + wait = _decode(wait_payload, label="wait") if wait_payload is not None else None + return ExecutionRecord( + execution_id=control.run_id, + execution_kind=ExecutionKind(control.kind), + deployment_name=control.deployment, + definition_revision=control.definition_revision, + validated_input=input_data["validated_input"], + operation_fingerprint=input_data["operation_fingerprint"], + owner_id=control.owner_id, + status=ExecutionStatus(control.status.value), + sequence=control.version, + attempt=control.run_attempt, + checkpoint=snapshot.get("checkpoint"), + application_state=snapshot.get("application_state", {}), + wait=wait, + progress=progress, + result=result if control.status is EngineStatus.COMPLETED else None, + error=error if control.status is EngineStatus.FAILED else None, + last_retry_error=error if not control.terminal else None, + retry_at=_datetime(control.available_at_ms) if control.available_at_ms is not None else None, + cancel_requested_at=( + _datetime(control.cancel_requested_at_ms) if control.cancel_requested_at_ms is not None else None + ), + cancel_reason=control.cancel_reason, + created_at=_datetime(control.created_at_ms), + updated_at=_datetime(control.updated_at_ms), + max_progress_events=self.max_progress_events, + max_record_bytes=self.max_record_bytes, + ) + + def _input(self, record: ExecutionRecord) -> bytes: + return _encode( + { + "validated_input": record.validated_input, + "operation_fingerprint": record.operation_fingerprint, + }, + limit=self.config.max_input_bytes, + label="validated input", + ) + + def _snapshot(self, record: ExecutionRecord) -> bytes: + return _encode( + { + "checkpoint": record.checkpoint.to_dict() if record.checkpoint else None, + "application_state": record.application_state, + }, + limit=self.config.max_checkpoint_bytes, + label="checkpoint", + ) + + +class RedisExecutionStoreProvider: + """Application-owned Redis client and deployment stores.""" + + def __init__( # noqa: PLR0913 - mirrors the configurable Redis task-store provider + self, + redis_url: str | None = None, + *, + redis: Any | None = None, + key_prefix: str | None = None, + close_redis: bool = True, + app_settings: AppSettings | None = None, + socket_timeout: float | None = None, + socket_connect_timeout: float | None = None, + health_check_interval: int | None = None, + ) -> None: + source_settings = app_settings if app_settings is not None else settings + self.app_settings = source_settings.model_copy(deep=True) + self.config = _config(app_settings=self.app_settings, key_prefix=key_prefix) + self.close_redis = close_redis + self.socket_timeout = ( + socket_timeout if socket_timeout is not None else self.app_settings.durable_redis_socket_timeout + ) + self.socket_connect_timeout = ( + socket_connect_timeout + if socket_connect_timeout is not None + else self.app_settings.durable_redis_socket_connect_timeout + ) + self.health_check_interval = ( + health_check_interval + if health_check_interval is not None + else self.app_settings.durable_redis_health_check_interval + ) + if redis is None: + try: + from redis.asyncio import Redis + except ImportError as error: # pragma: no cover - optional dependency guard + msg = 'Durable Redis storage requires `pip install "hayhooks[durable]`.' + raise ImportError(msg) from error + redis = Redis.from_url( + redis_url or self.app_settings.durable_redis_url, + decode_responses=False, + socket_timeout=self.socket_timeout, + socket_connect_timeout=self.socket_connect_timeout, + health_check_interval=self.health_check_interval, + ) + self.redis = redis + self.cores: dict[str, RedisExecutionStore] = {} + + def create_execution_store(self, deployment_name: str) -> ExecutionStore: + core = self.cores.get(deployment_name) + if core is None: + core = RedisExecutionStore(self.redis, deployment=deployment_name, config=self.config) + self.cores[deployment_name] = core + # A candidate deployment must not mutate the active deployment's accepted + # definition revision while it is still preparing or rolling back. + return _execution_store(core, app_settings=self.app_settings) + + async def close(self) -> None: + if self.close_redis: + await self.redis.aclose() + + +class InMemoryExecutionStoreProvider: + """Volatile reference backend for local development and tests.""" + + def __init__(self, *, app_settings: AppSettings | None = None) -> None: + source_settings = app_settings if app_settings is not None else settings + self.app_settings = source_settings.model_copy(deep=True) + self.config = _config(app_settings=self.app_settings) + self.cores: dict[str, InMemoryExecutionStore] = {} + + def create_execution_store(self, deployment_name: str) -> ExecutionStore: + core = self.cores.get(deployment_name) + if core is None: + core = InMemoryExecutionStore(deployment=deployment_name, config=self.config) + self.cores[deployment_name] = core + return _execution_store(core, app_settings=self.app_settings) + + async def close(self) -> None: + return None + + +def _execution_store(core: ExecutionBackend, *, app_settings: AppSettings) -> ExecutionStore: + """Build the same public adapter for both built-in backend implementations.""" + return ExecutionStore( + core, + lease_duration_ms=app_settings.durable_lease_duration_ms, + max_run_attempts=app_settings.durable_max_attempts, + max_progress_events=app_settings.durable_max_progress_events, + max_record_bytes=app_settings.durable_max_record_bytes, + ) + + +def _config(*, app_settings: AppSettings, key_prefix: str | None = None) -> ExecutionStoreConfig: + max_record = app_settings.durable_max_record_bytes + progress_bytes = DEFAULT_MAX_PROGRESS_BYTES + return ExecutionStoreConfig( + key_prefix=key_prefix or app_settings.durable_redis_key_prefix, + lease_commit_safety_ms=app_settings.durable_lease_commit_safety_ms, + terminal_ttl_seconds=app_settings.durable_terminal_ttl_seconds, + max_nonterminal_executions=app_settings.durable_max_nonterminal_executions, + max_input_bytes=max_record, + max_checkpoint_bytes=max_record, + max_result_bytes=max_record, + max_error_bytes=max_record, + max_wait_bytes=max_record, + max_progress_events=app_settings.durable_max_progress_events, + max_progress_event_bytes=progress_bytes, + ) + + +def _is_redis_error(error: BaseException) -> bool: + """Avoid exposing Redis client exceptions directly through HTTP routes.""" + try: + from redis.exceptions import RedisError + except ImportError: # pragma: no cover - Redis is an optional dependency + return False + return isinstance(error, RedisError) + + +__all__ = ["ExecutionClaim", "ExecutionStore"] diff --git a/tests/conftest.py b/tests/conftest.py index 5cde7614..1e719ece 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,6 @@ +import os import shutil +import uuid from collections.abc import Iterator from contextlib import contextmanager from contextvars import ContextVar @@ -10,6 +12,7 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from haystack.tracing import Span, Tracer, disable_tracing, enable_tracing +from redis.asyncio import Redis from hayhooks.server.app import create_app from hayhooks.server.logger import log @@ -18,6 +21,23 @@ from hayhooks.settings import settings +@pytest.fixture +async def isolated_redis(): + redis_url = os.getenv("HAYHOOKS_TEST_REDIS_URL") + if not redis_url: + pytest.skip("set HAYHOOKS_TEST_REDIS_URL to run the real-Redis suite") + redis = Redis.from_url(redis_url, decode_responses=False) + await redis.ping() + prefix = f"hayhooks:test:{uuid.uuid4().hex}" + try: + yield redis, prefix + finally: + keys = [key async for key in redis.scan_iter(match=f"{prefix}:*")] + if keys: + await redis.delete(*keys) + await redis.aclose() + + class _RecordedSpan(Span): def __init__( self, operation_name: str, tags: dict[str, Any], trace_id: int, span_id: int, parent_span_id: int | None diff --git a/tests/test_durable_redis_codec.py b/tests/test_durable_redis_codec.py new file mode 100644 index 00000000..6aebb382 --- /dev/null +++ b/tests/test_durable_redis_codec.py @@ -0,0 +1,113 @@ +"""Redis layout tests that do not duplicate reducer decisions.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, Mock + +import pytest + +from hayhooks.durable.backend import MAINTENANCE_BATCH_SIZE, ExecutionStoreCorruptionError +from hayhooks.durable.engine import MAX_CONTROL_SCALAR_BYTES, initial_control +from hayhooks.durable.redis import RedisExecutionStore, RedisKeys, decode_control, digest, encode_control + + +def control(**changes: object): + values = { + "run_id": "run-1", + "idempotency_digest": "a" * 64, + "idempotency_binding_digest": "d" * 64, + "deployment": "deployment", + "definition_revision": "rev-1", + "owner_id": "owner", + "kind": "pipeline", + "now_ms": 100, + } + values.update(changes) + return initial_control(**values) + + +def test_key_builder_is_namespaced_hash_tagged_and_private() -> None: + keys = RedisKeys("tenant:durable", "an unsafe deployment/name") + idem = digest("idempotency", "raw-client-key") + + all_keys = ( + keys.runnable, + keys.lease_expiry, + keys.capacity, + keys.control("run_1"), + keys.idempotency(idem), + ) + assert all("{" + keys.deployment_digest + "}" in key for key in all_keys) + assert "unsafe" not in " ".join(all_keys) + assert "raw-client-key" not in keys.idempotency(idem) + assert RedisKeys.lease_member("run_1", 7) == "run_1|7" + + +def test_key_builder_rejects_user_influenced_unsafe_components() -> None: + keys = RedisKeys("tenant:durable", "deployment") + with pytest.raises(ValueError): + keys.control("run:injection") + with pytest.raises(ValueError): + keys.idempotency("not-a-digest") + + +def test_control_hash_round_trip_omits_absent_optionals() -> None: + original = control(owner_id=None) + encoded = encode_control(original) + restored = decode_control(encoded) + + assert restored == original + assert "lease_owner" not in encoded + assert "available_at_ms" not in encoded + + +@pytest.mark.parametrize( + "field", + [ + "run_id", + "idempotency_digest", + "idempotency_binding_digest", + "deployment", + "definition_revision", + "owner_id", + "kind", + ], +) +def test_control_rejects_scalars_redis_cannot_decode(field) -> None: + with pytest.raises(ValueError, match=field): + control(**{field: "x" * (MAX_CONTROL_SCALAR_BYTES + 1)}) + + +@pytest.mark.parametrize( + "mutate", + [ + lambda values: values.pop("version"), + lambda values: values.__setitem__("status", "unexpected"), + lambda values: values.__setitem__("fence", "-1"), + lambda values: values.__setitem__("run_id", "x" * 5_000), + lambda values: values.__setitem__("lease_owner", "worker"), + ], +) +def test_control_hash_corruption_is_rejected(mutate) -> None: + values = encode_control(control()) + mutate(values) + with pytest.raises(ExecutionStoreCorruptionError): + decode_control(values) + + +async def test_maintenance_reads_a_fixed_batch_of_due_leases() -> None: + redis = AsyncMock() + redis.time.return_value = (123, 456_000) + store = RedisExecutionStore(redis, deployment="deployment") + + await store.maintain(Mock()) + + redis.time.assert_awaited_once() + redis.zrangebyscore.assert_awaited_once_with( + store.keys.lease_expiry, + "-inf", + 123_456, + start=0, + num=MAINTENANCE_BATCH_SIZE, + withscores=True, + ) diff --git a/tests/test_durable_reference.py b/tests/test_durable_reference.py new file mode 100644 index 00000000..28efe872 --- /dev/null +++ b/tests/test_durable_reference.py @@ -0,0 +1,93 @@ +"""Shared contract and lean-index checks for the reference backend.""" + +from __future__ import annotations + +import pytest + +from hayhooks.durable.backend import ExecutionStoreConfig +from hayhooks.durable.engine import ( + Checkpoint, + Claim, + ExecutionPayloadSizeError, + RecoverExpiredLease, + RequestCancellation, + Suspend, +) +from hayhooks.durable.reference import InMemoryExecutionStore +from tests.durable_contract import assert_store_contract, control + + +async def test_reference_store_matches_contract() -> None: + store = InMemoryExecutionStore( + deployment="integration", + config=ExecutionStoreConfig( + max_input_bytes=64, + max_checkpoint_bytes=64, + max_result_bytes=64, + max_error_bytes=64, + max_wait_bytes=64, + max_progress_events=2, + max_progress_event_bytes=32, + ), + ) + await store.initialize() + await assert_store_contract(store) + + +async def test_reference_rejects_oversized_payload_before_transition() -> None: + store = InMemoryExecutionStore( + deployment="integration", + config=ExecutionStoreConfig( + max_input_bytes=64, + max_checkpoint_bytes=8, + max_result_bytes=64, + max_error_bytes=64, + max_wait_bytes=64, + max_progress_events=2, + max_progress_event_bytes=32, + ), + ) + await store.submit(control(), b"{}", binding_digest="b" * 64) + run_id = await store.read_candidate() + assert run_id is not None + claim = await store.transition(run_id, Claim("worker", 0, 1_000, 3, "rev-1"), candidate=True) + with pytest.raises(ExecutionPayloadSizeError): + await store.transition(run_id, Checkpoint(claim.next_control.fence, "worker", 0, 1_000, b"too-large")) + with pytest.raises(ExecutionPayloadSizeError): + await store.transition(run_id, Suspend(claim.next_control.fence, "worker", 0, b"ok", b"x" * 65)) + assert await store.get(run_id) == claim.next_control + + +async def test_maintenance_repairs_only_the_stale_lease_member() -> None: + store = InMemoryExecutionStore(deployment="integration") + await store.submit(control(), b"{}", binding_digest="b" * 64) + store._lease_expiry["run_1|7"] = 0 + + def recover(fence: int, deadline: int) -> RecoverExpiredLease: + return RecoverExpiredLease(0, fence, deadline, 3, "rev-1") + + await store.maintain(recover) + assert "run_1|7" not in store._lease_expiry + + run_id = await store.read_candidate() + assert run_id is not None + claim = await store.transition(run_id, Claim("worker", 0, 1_000_000, 3, "rev-1"), candidate=True) + live_member = f"{run_id}|{claim.next_control.fence}" + live_deadline = store._lease_expiry[live_member] + store._lease_expiry[f"{run_id}|0"] = 0 + await store.maintain(recover) + assert store._lease_expiry == {live_member: live_deadline} + + +async def test_candidate_read_is_non_destructive_and_cancel_removes_runnable() -> None: + store = InMemoryExecutionStore(deployment="integration") + first = control("run_a", idempotency_digest="a" * 64, binding_digest="b" * 64) + second = control("run_b", idempotency_digest="c" * 64, binding_digest="d" * 64) + await store.submit(first, b"{}", binding_digest="b" * 64) + await store.submit(second, b"{}", binding_digest="d" * 64) + assert await store.read_candidate() == "run_a" + assert await store.read_candidate() == "run_a" + await store.transition("run_a", RequestCancellation(0, "cancel")) + await store.transition("run_b", RequestCancellation(0, "cancel")) + assert await store.read_candidate() is None + assert await store.operational_counts() == {"nonterminal": 0, "runnable": 0, "lease_expiry": 0} diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py new file mode 100644 index 00000000..d1ae29e2 --- /dev/null +++ b/tests/test_durable_store.py @@ -0,0 +1,495 @@ +"""Application-contract tests for the durable store adapter.""" + +from __future__ import annotations + +import asyncio +import time +from dataclasses import replace +from unittest.mock import AsyncMock + +import pytest + +from hayhooks.durable.backend import ExecutionStoreConfig +from hayhooks.durable.context import RESUME_INPUT_KEY +from hayhooks.durable.engine import Claim, Heartbeat, RequestCancellation, Resume +from hayhooks.durable.manager import DurableExecutionManager +from hayhooks.durable.models import ( + ExecutionCheckpoint, + ExecutionKind, + ExecutionProgressEvent, + ExecutionRecord, + ExecutionRecordSizeError, + ExecutionStatus, + ExecutionStoreError, + RetryableExecutionError, +) +from hayhooks.durable.reference import InMemoryExecutionStore +from hayhooks.durable.runtime import DurableRuntime +from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider +from hayhooks.settings import AppSettings, settings + + +def _config() -> ExecutionStoreConfig: + return ExecutionStoreConfig( + max_input_bytes=512, + max_checkpoint_bytes=512, + max_result_bytes=512, + max_error_bytes=512, + max_wait_bytes=512, + max_progress_events=2, + max_progress_event_bytes=256, + ) + + +def _store(*, config: ExecutionStoreConfig | None = None, **options) -> ExecutionStore: + return ExecutionStore( + InMemoryExecutionStore(deployment="deployment", config=config or _config()), + definition_revision="rev-1", + **options, + ) + + +def _record() -> ExecutionRecord: + return ExecutionRecord( + execution_id="run_1", + execution_kind=ExecutionKind.PIPELINE, + deployment_name="deployment", + definition_revision="rev-1", + validated_input={"question": "hello"}, + operation_fingerprint="request-fingerprint", + owner_id="owner", + max_progress_events=2, + max_record_bytes=512, + ) + + +def test_builtin_providers_snapshot_explicit_durable_settings() -> None: + app_settings = AppSettings( + durable_redis_key_prefix="portable:durable", + durable_lease_duration_ms=45_000, + durable_lease_commit_safety_ms=2_000, + durable_terminal_ttl_seconds=123, + durable_max_nonterminal_executions=12, + durable_max_attempts=7, + durable_max_progress_events=17, + durable_max_record_bytes=32_768, + ) + memory_store = InMemoryExecutionStoreProvider(app_settings=app_settings).create_execution_store("portable") + redis_provider = RedisExecutionStoreProvider( + redis=AsyncMock(), + app_settings=app_settings, + socket_timeout=1.5, + socket_connect_timeout=2.5, + health_check_interval=0, + ) + redis_store = redis_provider.create_execution_store("portable") + + for store in (memory_store, redis_store): + assert store.lease_duration_ms == 45_000 + assert store.max_run_attempts == 7 + assert store.max_progress_events == 17 + assert store.max_record_bytes == 32_768 + assert store.config.key_prefix == "portable:durable" + assert store.config.lease_commit_safety_ms == 2_000 + assert store.config.terminal_ttl_seconds == 123 + assert store.config.max_nonterminal_executions == 12 + + assert redis_provider.socket_timeout == 1.5 + assert redis_provider.socket_connect_timeout == 2.5 + assert redis_provider.health_check_interval == 0 + + +def test_runtime_uses_its_explicit_settings_for_the_default_provider() -> None: + app_settings = AppSettings(durable_store="memory", durable_lease_duration_ms=45_000) + runtime = DurableRuntime(app_settings=app_settings) + + provider = runtime._provider() + + assert isinstance(provider, InMemoryExecutionStoreProvider) + assert provider.app_settings.durable_lease_duration_ms == 45_000 + + +async def test_runtime_uses_implicit_provider_settings_until_close(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "durable_store", "memory") + original_attempts = settings.durable_max_attempts + runtime = DurableRuntime() + provider = runtime._provider() + + monkeypatch.setattr(settings, "durable_max_attempts", original_attempts + 1) + + assert runtime.app_settings.durable_max_attempts == provider.app_settings.durable_max_attempts == original_attempts + await runtime.close() + assert runtime.app_settings.durable_max_attempts == original_attempts + 1 + + +def test_runtime_uses_supplied_builtin_provider_as_its_settings_source() -> None: + app_settings = AppSettings(durable_store="memory", durable_lease_duration_ms=45_000, durable_max_attempts=7) + provider = InMemoryExecutionStoreProvider(app_settings=app_settings) + runtime = DurableRuntime(provider) + + store = provider.create_execution_store("portable") + + assert runtime.app_settings.durable_lease_duration_ms == store.lease_duration_ms == 45_000 + assert runtime.app_settings.durable_max_attempts == store.max_run_attempts == 7 + + +def test_runtime_rejects_conflicting_builtin_provider_settings() -> None: + provider = InMemoryExecutionStoreProvider(app_settings=AppSettings(durable_lease_duration_ms=45_000)) + + with pytest.raises(ValueError, match="settings must match"): + DurableRuntime(provider, app_settings=AppSettings(durable_lease_duration_ms=60_000)) + + +async def test_store_preserves_public_checkpoint_progress_wait_resume_and_result_contract() -> None: + store = _store( + lease_duration_ms=10_000, + max_run_attempts=3, + max_progress_events=2, + max_record_bytes=512, + ) + await store.initialize() + + created, submitted = await store.submit_with_record(_record()) + replayed, replay = await store.submit_with_record(_record()) + assert created and not replayed + assert replay.execution_id == submitted.execution_id + + claim = await store.claim_next("worker") + assert claim is not None + async with claim: + claim.record.application_state["step"] = "checkpointed" + claim.record.checkpoint = ExecutionCheckpoint(ExecutionKind.PIPELINE, {"component": "search"}) + claim.record.append_progress("checkpoint saved", kind="checkpoint") + await claim.checkpoint() + claim.record.wait = {"kind": "approval", "message": "continue?"} + claim.record.status = ExecutionStatus.WAITING + claim.record.append_progress("waiting", kind="waiting") + await claim.suspend() + + waiting = await store.get("run_1") + assert waiting is not None + assert waiting.status is ExecutionStatus.WAITING + assert waiting.application_state == {"step": "checkpointed"} + assert [event.kind for event in waiting.progress] == ["checkpoint", "waiting"] + + assert await store.resume("run_1", {"approved": True}) + resumed = await store.claim_next("worker") + assert resumed is not None + async with resumed: + assert resumed.record.application_state.pop(RESUME_INPUT_KEY) == {"approved": True} + resumed.record.result = {"answer": "done"} + resumed.record.status = ExecutionStatus.COMPLETED + await resumed.complete() + + completed = await store.get("run_1") + assert completed is not None + assert completed.status is ExecutionStatus.COMPLETED + assert completed.result == {"answer": "done"} + assert completed.wait is None + assert [event.kind for event in completed.progress] == ["waiting", "resumed"] + + +async def test_cancel_and_checkpoint_assign_distinct_persisted_progress_sequences( + monkeypatch: pytest.MonkeyPatch, +) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + entered = asyncio.Event() + release = asyncio.Event() + original_transition = core.transition + + async def gated_transition(run_id, command, *, candidate=False): + if isinstance(command, RequestCancellation): + entered.set() + await release.wait() + return await original_transition(run_id, command, candidate=candidate) + + monkeypatch.setattr(core, "transition", gated_transition) + claim = await store.claim_next("worker") + assert claim is not None + async with claim: + cancellation = asyncio.create_task(store.request_cancel("run_1")) + await entered.wait() + claim.record.append_progress("checkpoint", kind="checkpoint") + await claim.checkpoint() + release.set() + assert await cancellation + + record = await store.get("run_1") + assert record is not None + assert [event.sequence for event in record.progress] == [1, 2] + + +async def test_losing_resume_race_returns_false(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + async with claim: + claim.record.status = ExecutionStatus.WAITING + claim.record.wait = {"kind": "approval"} + await claim.suspend() + + entered = asyncio.Event() + release = asyncio.Event() + original_transition = core.transition + + async def gated_transition(run_id, command, *, candidate=False): + if isinstance(command, Resume): + entered.set() + await release.wait() + return await original_transition(run_id, command, candidate=candidate) + + monkeypatch.setattr(core, "transition", gated_transition) + resumed = asyncio.create_task(store.resume("run_1", {"approved": True})) + await entered.wait() + await store.request_cancel("run_1") + release.set() + + assert not await resumed + record = await store.get("run_1") + assert record is not None and record.status is ExecutionStatus.CANCELED + + +async def test_retry_exhaustion_persists_its_progress_event() -> None: + store = _store(max_run_attempts=1) + + async def runner(context): + await context.retry("again", delay=0) + + manager = DurableExecutionManager( + "deployment", store, runner, adapter=object(), poll_interval=0.001, max_attempts=1 + ) + await manager.start() + try: + await store.submit(_record()) + for _ in range(100): + record = await store.get("run_1") + if record is not None and record.terminal: + break + await asyncio.sleep(0.001) + else: + pytest.fail("retry exhaustion did not become terminal") + finally: + await manager.close() + + assert record is not None + assert [event.kind for event in record.progress] == ["retry_exhausted"] + + +async def test_store_runs_through_the_existing_durable_manager_contract() -> None: + store = _store( + lease_duration_ms=10_000, + max_run_attempts=3, + max_progress_events=2, + max_record_bytes=512, + ) + + async def runner(context): + context.record.application_state["phase"] = "running" + await context.report_progress("started") + return {"answer": "done"} + + manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) + await manager.start() + try: + await store.submit(_record()) + for _ in range(100): + record = await store.get("run_1") + if record and record.terminal: + break + await asyncio.sleep(0.001) + else: + pytest.fail("manager did not complete the submitted execution") + assert record.status is ExecutionStatus.COMPLETED + assert record.result == {"answer": "done"} + assert record.application_state == {"phase": "running"} + assert [event.message for event in record.progress] == ["started"] + finally: + await manager.close() + + +async def test_canceled_runner_restarts_its_worker_slot() -> None: + store = _store( + config=replace(_config(), lease_commit_safety_ms=1), + lease_duration_ms=50, + ) + calls = 0 + + async def runner(_context): + nonlocal calls + calls += 1 + if calls == 1: + raise asyncio.CancelledError + return {"answer": "done"} + + manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) + await manager.start() + try: + await store.submit(_record()) + for _ in range(200): + record = await store.get("run_1") + if record is not None and record.terminal: + break + await asyncio.sleep(0.005) + else: + pytest.fail("canceled runner did not recover") + finally: + await manager.close() + + assert calls == 2 + assert record is not None + assert record.status is ExecutionStatus.COMPLETED + + +async def test_oversized_result_fails_without_replaying_the_runner() -> None: + store = _store(max_record_bytes=512) + calls = 0 + + async def runner(_context): + nonlocal calls + calls += 1 + return {"answer": "x" * 512} + + manager = DurableExecutionManager("deployment", store, runner, adapter=object()) + await store.initialize() + assert await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + + await manager._process_claim(claim) + + record = await store.get("run_1") + assert calls == 1 + assert record is not None + assert record.status is ExecutionStatus.FAILED + assert record.error is not None + assert record.error.code == "record_too_large" + + +def test_progress_event_enforces_the_full_serialized_size() -> None: + with pytest.raises(ExecutionRecordSizeError, match="progress event"): + ExecutionProgressEvent(sequence=1, message="", kind="", metadata={"data": "x" * 8_150}) + + +async def test_in_memory_store_uses_real_time_for_delayed_retries() -> None: + store = _store( + config=replace(_config(), lease_commit_safety_ms=1), + lease_duration_ms=10_000, + max_run_attempts=3, + max_progress_events=2, + max_record_bytes=512, + ) + attempts = 0 + + async def runner(_context): + nonlocal attempts + attempts += 1 + if attempts == 1: + msg = "try again" + raise RetryableExecutionError(msg, delay=0.01) + return {"answer": "done"} + + manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) + await manager.start() + try: + await store.submit(_record()) + for _ in range(100): + record = await store.get("run_1") + if record is not None and record.terminal: + break + await asyncio.sleep(0.005) + else: + pytest.fail("in-memory durable retry did not become due") + finally: + await manager.close() + + assert attempts == 2 + assert record.status is ExecutionStatus.COMPLETED + + +def test_store_rejects_a_lease_safety_margin_equal_to_the_lease() -> None: + with pytest.raises(ValueError, match="lease_commit_safety_ms"): + _store( + config=replace(_config(), lease_commit_safety_ms=100), + lease_duration_ms=100, + ) + + +def test_store_requires_the_safe_lease_window_to_cover_one_heartbeat() -> None: + with pytest.raises(ValueError, match="heartbeat interval"): + _store( + config=replace(_config(), lease_commit_safety_ms=0), + lease_duration_ms=1, + ) + + +async def test_local_lease_deadlines_exclude_backend_round_trips(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1", lease_duration_ms=10_000) + await store.submit(_record()) + original_transition = core.transition + started: dict[type, float] = {} + + async def delayed_transition(run_id, command, *, candidate=False): + if isinstance(command, Claim | Heartbeat): + started[type(command)] = time.monotonic() + await asyncio.sleep(0.05) + return await original_transition(run_id, command, candidate=candidate) + + monkeypatch.setattr(core, "transition", delayed_transition) + claim = await store.claim_next("worker") + assert claim is not None + assert claim._confirmed_until == pytest.approx(started[Claim] + store.lease_safe_duration, abs=0.02) + + async with claim: + assert claim._confirmed_until == pytest.approx(started[Heartbeat] + store.lease_safe_duration, abs=0.02) + + +async def test_store_normalizes_redis_client_failures(monkeypatch: pytest.MonkeyPatch) -> None: + from redis.exceptions import ConnectionError as RedisConnectionError + + store = _store() + + async def unavailable(_execution_id: str): + msg = "offline" + raise RedisConnectionError(msg) + + monkeypatch.setattr(store.core, "get", unavailable) + with pytest.raises(ExecutionStoreError, match="read execution"): + await store.get("run_1") + + +async def test_store_retries_a_view_changed_while_reading_payloads(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.initialize() + await store.submit(_record()) + original_read_progress = core.read_progress + changed = False + + async def read_progress_and_cancel(run_id: str) -> list[bytes]: + nonlocal changed + progress = await original_read_progress(run_id) + if not changed: + changed = True + await core.transition(run_id, RequestCancellation(0, "cancel")) + return progress + + monkeypatch.setattr(core, "read_progress", read_progress_and_cancel) + record = await store.get("run_1") + + assert changed + assert record is not None + assert record.status is ExecutionStatus.CANCELED + + +async def test_concurrent_adapter_claims_return_only_the_fenced_owner() -> None: + store = _store() + await store.submit(_record()) + claims = await asyncio.gather(*(store.claim_next(f"worker-{index}") for index in range(3))) + owners = [claim for claim in claims if claim is not None] + assert len(owners) == 1 + assert owners[0].control.lease_owner == owners[0].worker_id From 5b7b2fe484c07da810eaf4782fb4e7a1317d9f18 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:08:57 +0200 Subject: [PATCH 03/28] feat(durable): add recoverable execution engine --- src/hayhooks/durable/context.py | 173 +++ src/hayhooks/durable/engine.py | 595 +++++++++ src/hayhooks/durable/manager.py | 507 +++++++ tests/test_durable_engine.py | 151 +++ tests/test_durable_execution.py | 1171 +++++++++++++++++ tests/test_durable_process_recovery.py | 327 +++++ .../pipeline_wrapper.py | 74 ++ tests/test_redis_execution_integration.py | 235 ++++ 8 files changed, 3233 insertions(+) create mode 100644 src/hayhooks/durable/context.py create mode 100644 src/hayhooks/durable/engine.py create mode 100644 src/hayhooks/durable/manager.py create mode 100644 tests/test_durable_engine.py create mode 100644 tests/test_durable_execution.py create mode 100644 tests/test_durable_process_recovery.py create mode 100644 tests/test_files/durable_process_recovery/pipeline_wrapper.py create mode 100644 tests/test_redis_execution_integration.py diff --git a/src/hayhooks/durable/context.py b/src/hayhooks/durable/context.py new file mode 100644 index 00000000..50d07df4 --- /dev/null +++ b/src/hayhooks/durable/context.py @@ -0,0 +1,173 @@ +"""Store contracts and the context exposed to durable application code.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Coroutine, Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from typing import TYPE_CHECKING, Any, cast + +from hayhooks.durable.models import ( + ExecutionCanceledError, + ExecutionCheckpoint, + ExecutionProgressEvent, + ExecutionStatus, + ExecutionSuspendedError, + JsonValue, + RetryableExecutionError, + validate_json, +) + +if TYPE_CHECKING: + from hayhooks.durable.adapters import HaystackDurableAdapter + from hayhooks.durable.store import ExecutionClaim + +RESUME_INPUT_KEY = "__hayhooks_resume_input" + + +_active_context: ContextVar[DurableContext | None] = ContextVar("hayhooks_durable_context", default=None) + + +def get_current_durable_context() -> DurableContext | None: + """Return the context active in a durable wrapper, component, hook, or tool.""" + return _active_context.get() + + +@contextmanager +def execution_context_scope(context: DurableContext): + token = _active_context.set(context) + try: + yield + finally: + _active_context.reset(token) + + +class DurableContext: + """Execution controls and adapters bound to one claimed record.""" + + def __init__( + self, + claim: ExecutionClaim, + adapter: HaystackDurableAdapter, + *, + event_loop: asyncio.AbstractEventLoop | None = None, + ) -> None: + self.claim = claim + self.record = claim.record + self.adapter = adapter + self._event_loop = event_loop or asyncio.get_running_loop() + + @property + def execution_id(self) -> str: + return self.record.execution_id + + @property + def attempt(self) -> int: + return self.record.attempt + + @property + def state(self) -> dict[str, JsonValue]: + return self.record.application_state + + @property + def resume_input(self) -> JsonValue | None: + """Return the most recently persisted resume payload without consuming it.""" + return self.record.application_state.get(RESUME_INPUT_KEY) + + def take_resume_input(self) -> JsonValue | None: + """Consume the persisted resume payload exactly once within this attempt.""" + return self.record.application_state.pop(RESUME_INPUT_KEY, None) + + async def checkpoint(self, checkpoint: ExecutionCheckpoint | None = None) -> None: + if checkpoint is not None: + if checkpoint.kind is not self.record.execution_kind: + msg = ( + f"{checkpoint.kind.value} checkpoint cannot be used for " + f"{self.record.execution_kind.value} execution" + ) + raise ValueError(msg) + checkpoint.data = cast( + dict[str, JsonValue], + validate_json(checkpoint.data, limit=self.record.max_record_bytes, label="checkpoint"), + ) + self.record.checkpoint = checkpoint + await self.claim.checkpoint() + + async def report_progress( + self, message: str, *, kind: str = "progress", metadata: Mapping[str, Any] | None = None + ) -> ExecutionProgressEvent: + event = self.record.append_progress(message, kind=kind, metadata=metadata) + await self.claim.checkpoint() + return event + + def report_progress_sync( + self, message: str, *, kind: str = "progress", metadata: Mapping[str, Any] | None = None + ) -> ExecutionProgressEvent: + return self._sync_await(self.report_progress(message, kind=kind, metadata=metadata)) + + async def check_cancelled(self) -> None: + if await self.claim.cancellation_requested(): + msg = "Durable execution cancellation was requested" + raise ExecutionCanceledError(msg) + + def check_cancelled_sync(self) -> None: + self._sync_await(self.check_cancelled()) + + async def retry(self, message: str, *, delay: float | None = None) -> None: + """Request a bounded, durable retry from application code.""" + raise RetryableExecutionError(message, delay=delay or 0.0) + + def retry_sync(self, message: str, *, delay: float | None = None) -> None: + """Synchronous counterpart to :meth:`retry`.""" + raise RetryableExecutionError(message, delay=delay or 0.0) + + async def suspend(self, wait: Mapping[str, Any], *, update: Mapping[str, Any] | None = None) -> None: + """Atomically checkpoint and move this execution to durable ``waiting``.""" + self.record.wait = cast( + dict[str, JsonValue], validate_json(dict(wait), limit=self.record.max_record_bytes, label="wait") + ) + if update is not None: + self.record.application_state.update( + cast( + dict[str, JsonValue], + validate_json(dict(update), limit=self.record.max_record_bytes, label="wait update"), + ) + ) + self.record.status = ExecutionStatus.WAITING + self.record.append_progress("Execution is waiting for resume", kind="waiting") + await self.claim.suspend() + raise ExecutionSuspendedError() + + def suspend_sync(self, wait: Mapping[str, Any], *, update: Mapping[str, Any] | None = None) -> None: + self._sync_await(self.suspend(wait, update=update)) + + async def run_pipeline_async( + self, data: Mapping[str, Any], *, checkpoint_at: list[str] | None = None + ) -> dict[str, Any]: + return await self.adapter.run_pipeline_async(self, dict(data), checkpoint_at=checkpoint_at or []) + + def run_pipeline(self, data: Mapping[str, Any], *, checkpoint_at: list[str] | None = None) -> dict[str, Any]: + return self.adapter.run_pipeline(self, dict(data), checkpoint_at=checkpoint_at or []) + + async def run_agent_async(self, *, messages: list[Any], **kwargs: Any) -> dict[str, Any]: + return await self.adapter.run_agent_async(self, messages=messages, **kwargs) + + def run_agent(self, *, messages: list[Any], **kwargs: Any) -> dict[str, Any]: + return self.adapter.run_agent(self, messages=messages, **kwargs) + + def _sync_await(self, awaitable: Awaitable[Any]) -> Any: + """Bridge a synchronous wrapper/component thread to its manager loop.""" + try: + running = asyncio.get_running_loop() + except RuntimeError: + running = None + if running is self._event_loop: + msg = "A synchronous durable context method cannot run on the server event loop" + raise RuntimeError(msg) + + async def resolve() -> Any: + return await awaitable + + future = asyncio.run_coroutine_threadsafe(cast(Coroutine[Any, Any, Any], resolve()), self._event_loop) + return future.result() diff --git a/src/hayhooks/durable/engine.py b/src/hayhooks/durable/engine.py new file mode 100644 index 00000000..34b9d5d0 --- /dev/null +++ b/src/hayhooks/durable/engine.py @@ -0,0 +1,595 @@ +""" +Storage-neutral durable execution state machine. + +The reducer is the only place that decides an execution lifecycle. Storage +only persists its effects and derives its indexes from the old and new control. +""" +# ruff: noqa: EM101, EM102, PLR0911, PLR0913 + +from __future__ import annotations + +import re +from dataclasses import dataclass, replace +from enum import Enum + +RUN_ID_PATTERN = r"[A-Za-z0-9_-]{1,128}" +MAX_CONTROL_SCALAR_BYTES = 4_096 +MAX_CANCELLATION_REASON_LENGTH = 2_000 +_RUN_ID_RE = re.compile(f"{RUN_ID_PATTERN}\\Z") + + +def validate_run_id(run_id: str) -> None: + """Reject execution IDs that cannot be embedded safely in backend keys.""" + if not _RUN_ID_RE.fullmatch(run_id): + raise ValueError("run_id has an invalid key-safe format") + + +def normalize_cancellation_reason(reason: str | None) -> str | None: + """Bound cancellation text by both characters and persisted UTF-8 bytes.""" + if not reason: + return None + encoded = str(reason)[:MAX_CANCELLATION_REASON_LENGTH].encode()[:MAX_CONTROL_SCALAR_BYTES] + return encoded.decode(errors="ignore") + + +class ExecutionStatus(str, Enum): + QUEUED = "queued" + RUNNING = "running" + WAITING = "waiting" + COMPLETED = "completed" + FAILED = "failed" + CANCELED = "canceled" + + @property + def terminal(self) -> bool: + return self in {self.COMPLETED, self.FAILED, self.CANCELED} + + +class PayloadKind(str, Enum): + INPUT = "input" + CHECKPOINT = "checkpoint" + RESULT = "result" + ERROR = "error" + WAIT = "wait" + + +class ExecutionNotFoundError(RuntimeError): + pass + + +class ExecutionLeaseLostError(RuntimeError): + pass + + +class InvalidExecutionTransitionError(RuntimeError): + pass + + +class ExecutionPayloadSizeError(RuntimeError): + pass + + +@dataclass(frozen=True, slots=True) +class ExecutionControl: + """ + The compact, authoritative execution state. + + Times come from Redis ``TIME`` (or the reference-store equivalent), so a + worker never decides leases from its own wall clock. + """ + + run_id: str + idempotency_digest: str + idempotency_binding_digest: str + deployment: str + definition_revision: str + owner_id: str | None + kind: str + status: ExecutionStatus = ExecutionStatus.QUEUED + version: int = 1 + fence: int = 0 + run_attempt: int = 0 + application_retry_count: int = 0 + available_at_ms: int | None = None + lease_owner: str | None = None + lease_expires_at_ms: int | None = None + cancel_requested_at_ms: int | None = None + cancel_reason: str | None = None + progress_sequence: int = 0 + created_at_ms: int = 0 + updated_at_ms: int = 0 + + def __post_init__(self) -> None: + for name in ( + "run_id", + "idempotency_digest", + "idempotency_binding_digest", + "deployment", + "definition_revision", + "kind", + ): + value = getattr(self, name) + if not value or len(value.encode()) > MAX_CONTROL_SCALAR_BYTES: + raise ValueError(f"{name} must be non-empty and at most {MAX_CONTROL_SCALAR_BYTES} bytes") + for name in ("owner_id", "lease_owner", "cancel_reason"): + value = getattr(self, name) + if value is not None and len(value.encode()) > MAX_CONTROL_SCALAR_BYTES: + raise ValueError(f"{name} must be at most {MAX_CONTROL_SCALAR_BYTES} bytes") + validate_run_id(self.run_id) + for name in ( + "version", + "fence", + "run_attempt", + "application_retry_count", + "progress_sequence", + "created_at_ms", + "updated_at_ms", + ): + if getattr(self, name) < 0: + raise ValueError(f"{name} cannot be negative") + if self.status is ExecutionStatus.RUNNING: + if not self.lease_owner or self.lease_expires_at_ms is None: + raise ValueError("running controls require lease_owner and lease_expires_at_ms") + elif self.lease_owner is not None or self.lease_expires_at_ms is not None: + raise ValueError("only running controls may hold a lease") + + @property + def terminal(self) -> bool: + return self.status.terminal + + +@dataclass(frozen=True, slots=True) +class PayloadWrite: + kind: PayloadKind + data: bytes + + +@dataclass(frozen=True, slots=True) +class ProgressEvent: + sequence: int + data: bytes + + +@dataclass(frozen=True, slots=True) +class LeaseIndexUpdate: + deadline_ms: int | None + fence: int + + +@dataclass(frozen=True, slots=True) +class TransitionPlan: + next_control: ExecutionControl + payload_writes: tuple[PayloadWrite, ...] = () + payload_deletes: tuple[PayloadKind, ...] = () + progress_events: tuple[ProgressEvent, ...] = () + lease_index_update: LeaseIndexUpdate | None = None + + +@dataclass(frozen=True, slots=True) +class Claim: + worker_id: str + now_ms: int + lease_duration_ms: int + max_run_attempts: int + worker_revision: str + + +@dataclass(frozen=True, slots=True) +class Heartbeat: + fence: int + worker_id: str + now_ms: int + lease_duration_ms: int + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class Checkpoint: + fence: int + worker_id: str + now_ms: int + lease_duration_ms: int + payload: bytes + progress_events: tuple[bytes, ...] = () + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class RequestCancellation: + now_ms: int + reason: str | None = None + progress_events: tuple[bytes, ...] = () + + +@dataclass(frozen=True, slots=True) +class ScheduleRetry: + fence: int + worker_id: str + now_ms: int + delay_ms: int + max_application_retries: int + error: bytes = b"" + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class Suspend: + fence: int + worker_id: str + now_ms: int + checkpoint: bytes + wait: bytes + progress_events: tuple[bytes, ...] = () + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class Resume: + now_ms: int + worker_revision: str + checkpoint: bytes | None = None + progress_events: tuple[bytes, ...] = () + + +@dataclass(frozen=True, slots=True) +class Complete: + fence: int + worker_id: str + now_ms: int + result: bytes + progress_events: tuple[bytes, ...] = () + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class Fail: + fence: int + worker_id: str + now_ms: int + error: bytes + progress_events: tuple[bytes, ...] = () + lease_commit_safety_ms: int = 0 + + +@dataclass(frozen=True, slots=True) +class RecoverExpiredLease: + now_ms: int + indexed_fence: int + indexed_deadline_ms: int + max_run_attempts: int + worker_revision: str + + +ExecutionCommand = ( + Claim + | Heartbeat + | Checkpoint + | RequestCancellation + | ScheduleRetry + | Suspend + | Resume + | Complete + | Fail + | RecoverExpiredLease +) + + +def initial_control( + *, + run_id: str, + idempotency_digest: str, + idempotency_binding_digest: str, + deployment: str, + definition_revision: str, + owner_id: str | None, + kind: str, + now_ms: int, +) -> ExecutionControl: + """Create the payload-free, version-one control record for a new queued execution.""" + return ExecutionControl( + run_id=run_id, + idempotency_digest=idempotency_digest, + idempotency_binding_digest=idempotency_binding_digest, + deployment=deployment, + definition_revision=definition_revision, + owner_id=owner_id, + kind=kind, + created_at_ms=now_ms, + updated_at_ms=now_ms, + ) + + +def submission_plan(control: ExecutionControl, input_payload: bytes) -> TransitionPlan: + """Validate a new control and return the initial input-persistence plan.""" + if control.status is not ExecutionStatus.QUEUED or control.version != 1: + raise InvalidExecutionTransitionError("only a new queued control can be submitted") + return TransitionPlan(control, payload_writes=(PayloadWrite(PayloadKind.INPUT, input_payload),)) + + +def decide(control: ExecutionControl, command: ExecutionCommand) -> TransitionPlan: # noqa: C901 + """ + Reduce one command into the next control and its atomic persistence effects. + + The reducer performs no I/O and never mutates ``control``. Invalid lifecycle + or lease transitions raise ``InvalidExecutionTransitionError`` or + ``ExecutionLeaseLostError`` before any effects can be persisted. + """ + if isinstance(command, Claim): + return _claim(control, command) + if isinstance(command, Heartbeat): + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + return TransitionPlan( + replace(control, lease_expires_at_ms=command.now_ms + command.lease_duration_ms), + lease_index_update=LeaseIndexUpdate(command.now_ms + command.lease_duration_ms, control.fence), + ) + if isinstance(command, Checkpoint): + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + next_control = _business( + control, + command.now_ms, + progress_sequence=control.progress_sequence + len(command.progress_events), + lease_expires_at_ms=command.now_ms + command.lease_duration_ms, + ) + return TransitionPlan( + next_control, + payload_writes=(PayloadWrite(PayloadKind.CHECKPOINT, command.payload),), + progress_events=_progress_events(control.progress_sequence, command.progress_events), + lease_index_update=LeaseIndexUpdate(next_control.lease_expires_at_ms, control.fence), + ) + if isinstance(command, RequestCancellation): + return _cancel(control, command) + if isinstance(command, ScheduleRetry): + return _retry(control, command) + if isinstance(command, Suspend): + return _suspend(control, command) + if isinstance(command, Resume): + return _resume(control, command) + if isinstance(command, Complete): + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + return _terminal_or_canceled( + control, + command.now_ms, + ExecutionStatus.COMPLETED, + PayloadKind.RESULT, + command.result, + command.progress_events, + ) + if isinstance(command, Fail): + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + return _terminal_or_canceled( + control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, command.error, command.progress_events + ) + if isinstance(command, RecoverExpiredLease): + return _recover(control, command) + raise TypeError(f"unsupported execution command {type(command).__name__}") + + +def _claim(control: ExecutionControl, command: Claim) -> TransitionPlan: + if control.status is not ExecutionStatus.QUEUED: + raise InvalidExecutionTransitionError("execution is not queued") + if control.available_at_ms is not None and control.available_at_ms > command.now_ms: + raise InvalidExecutionTransitionError("queued execution is not due") + if control.definition_revision != command.worker_revision: + return _terminal( + control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"definition revision is incompatible" + ) + if control.cancel_requested_at_ms is not None: + return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) + if control.run_attempt >= command.max_run_attempts: + return _terminal(control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"run attempts exhausted") + deadline = command.now_ms + command.lease_duration_ms + next_control = _business( + control, + command.now_ms, + status=ExecutionStatus.RUNNING, + fence=control.fence + 1, + run_attempt=control.run_attempt + 1, + available_at_ms=None, + lease_owner=command.worker_id, + lease_expires_at_ms=deadline, + ) + return TransitionPlan(next_control, lease_index_update=LeaseIndexUpdate(deadline, next_control.fence)) + + +def _cancel(control: ExecutionControl, command: RequestCancellation) -> TransitionPlan: + if control.terminal or control.cancel_requested_at_ms is not None: + return TransitionPlan(control) + if control.status is ExecutionStatus.RUNNING: + return TransitionPlan( + _business( + control, + command.now_ms, + cancel_requested_at_ms=command.now_ms, + cancel_reason=normalize_cancellation_reason(command.reason), + progress_sequence=control.progress_sequence + len(command.progress_events), + ), + progress_events=_progress_events(control.progress_sequence, command.progress_events), + ) + return _terminal( + _business( + control, + command.now_ms, + cancel_requested_at_ms=command.now_ms, + cancel_reason=normalize_cancellation_reason(command.reason), + progress_sequence=control.progress_sequence + len(command.progress_events), + ), + command.now_ms, + ExecutionStatus.CANCELED, + None, + None, + increment_version=False, + progress_events=_progress_events(control.progress_sequence, command.progress_events), + ) + + +def _retry(control: ExecutionControl, command: ScheduleRetry) -> TransitionPlan: + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + if control.cancel_requested_at_ms is not None: + return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) + if control.application_retry_count >= command.max_application_retries: + return _terminal( + control, + command.now_ms, + ExecutionStatus.FAILED, + PayloadKind.ERROR, + command.error or b"application retries exhausted", + ) + due = command.now_ms + max(0, command.delay_ms) + next_control = _business( + control, + command.now_ms, + status=ExecutionStatus.QUEUED, + application_retry_count=control.application_retry_count + 1, + available_at_ms=due, + lease_owner=None, + lease_expires_at_ms=None, + ) + return TransitionPlan( + next_control, + payload_writes=((PayloadWrite(PayloadKind.ERROR, command.error),) if command.error else ()), + payload_deletes=((PayloadKind.ERROR,) if not command.error else ()), + lease_index_update=LeaseIndexUpdate(None, control.fence), + ) + + +def _suspend(control: ExecutionControl, command: Suspend) -> TransitionPlan: + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + if control.cancel_requested_at_ms is not None: + return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) + next_control = _business( + control, + command.now_ms, + status=ExecutionStatus.WAITING, + progress_sequence=control.progress_sequence + len(command.progress_events), + lease_owner=None, + lease_expires_at_ms=None, + ) + return TransitionPlan( + next_control, + payload_writes=( + PayloadWrite(PayloadKind.CHECKPOINT, command.checkpoint), + PayloadWrite(PayloadKind.WAIT, command.wait), + ), + progress_events=_progress_events(control.progress_sequence, command.progress_events), + lease_index_update=LeaseIndexUpdate(None, control.fence), + ) + + +def _resume(control: ExecutionControl, command: Resume) -> TransitionPlan: + if control.status is not ExecutionStatus.WAITING: + raise InvalidExecutionTransitionError("only waiting executions can resume") + if control.definition_revision != command.worker_revision: + return _terminal( + control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"definition revision is incompatible" + ) + if control.cancel_requested_at_ms is not None: + return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) + writes = (PayloadWrite(PayloadKind.CHECKPOINT, command.checkpoint),) if command.checkpoint is not None else () + next_control = _business( + control, + command.now_ms, + status=ExecutionStatus.QUEUED, + progress_sequence=control.progress_sequence + len(command.progress_events), + ) + return TransitionPlan( + next_control, + payload_writes=writes, + payload_deletes=(PayloadKind.WAIT,), + progress_events=_progress_events(control.progress_sequence, command.progress_events), + ) + + +def _recover(control: ExecutionControl, command: RecoverExpiredLease) -> TransitionPlan: + if control.status is not ExecutionStatus.RUNNING or control.fence != command.indexed_fence: + return TransitionPlan(control, lease_index_update=LeaseIndexUpdate(None, command.indexed_fence)) + assert control.lease_expires_at_ms is not None + if control.lease_expires_at_ms != command.indexed_deadline_ms: + return TransitionPlan(control, lease_index_update=LeaseIndexUpdate(control.lease_expires_at_ms, control.fence)) + if control.lease_expires_at_ms > command.now_ms: + return TransitionPlan(control) + if control.cancel_requested_at_ms is not None: + return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) + if control.definition_revision != command.worker_revision: + return _terminal( + control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"definition revision is incompatible" + ) + if control.run_attempt >= command.max_run_attempts: + return _terminal(control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"run attempts exhausted") + next_control = _business( + control, command.now_ms, status=ExecutionStatus.QUEUED, lease_owner=None, lease_expires_at_ms=None + ) + return TransitionPlan(next_control, lease_index_update=LeaseIndexUpdate(None, control.fence)) + + +def _terminal_or_canceled( + control: ExecutionControl, + now_ms: int, + status: ExecutionStatus, + payload_kind: PayloadKind, + payload: bytes, + progress_values: tuple[bytes, ...] = (), +) -> TransitionPlan: + progress_events = _progress_events(control.progress_sequence, progress_values) + if control.cancel_requested_at_ms is not None: + return _terminal(control, now_ms, ExecutionStatus.CANCELED, None, None, progress_events=progress_events) + return _terminal(control, now_ms, status, payload_kind, payload, progress_events=progress_events) + + +def _terminal( + control: ExecutionControl, + now_ms: int, + status: ExecutionStatus, + payload_kind: PayloadKind | None, + payload: bytes | None, + *, + increment_version: bool = True, + progress_events: tuple[ProgressEvent, ...] = (), +) -> TransitionPlan: + if control.terminal: + raise InvalidExecutionTransitionError("terminal execution cannot transition") + kwargs: dict[str, object] = { + "status": status, + "available_at_ms": None, + "lease_owner": None, + "lease_expires_at_ms": None, + } + if progress_events: + kwargs["progress_sequence"] = progress_events[-1].sequence + next_control = ( + _business(control, now_ms, **kwargs) if increment_version else replace(control, updated_at_ms=now_ms, **kwargs) + ) + writes = (PayloadWrite(payload_kind, payload or b""),) if payload_kind is not None else () + deletes: tuple[PayloadKind, ...] = () + if payload_kind is PayloadKind.RESULT: + deletes = (PayloadKind.ERROR,) + elif payload_kind is PayloadKind.ERROR: + deletes = (PayloadKind.RESULT,) + elif status is ExecutionStatus.CANCELED: + deletes = (PayloadKind.RESULT, PayloadKind.ERROR) + return TransitionPlan( + next_control, + payload_writes=writes, + payload_deletes=(*deletes, PayloadKind.WAIT), + progress_events=progress_events, + lease_index_update=LeaseIndexUpdate(None, control.fence), + ) + + +def _owned(control: ExecutionControl, fence: int, worker_id: str, now_ms: int, safety_margin_ms: int) -> None: + if ( + control.status is not ExecutionStatus.RUNNING + or control.fence != fence + or control.lease_owner != worker_id + or control.lease_expires_at_ms is None + or safety_margin_ms < 0 + or now_ms >= control.lease_expires_at_ms - safety_margin_ms + ): + raise ExecutionLeaseLostError("execution is no longer owned by this worker fence") + + +def _business(control: ExecutionControl, now_ms: int, **changes: object) -> ExecutionControl: + return replace(control, version=control.version + 1, updated_at_ms=now_ms, **changes) + + +def _progress_events(sequence: int, values: tuple[bytes, ...]) -> tuple[ProgressEvent, ...]: + return tuple(ProgressEvent(sequence + index, value) for index, value in enumerate(values, start=1)) diff --git a/src/hayhooks/durable/manager.py b/src/hayhooks/durable/manager.py new file mode 100644 index 00000000..bae7de29 --- /dev/null +++ b/src/hayhooks/durable/manager.py @@ -0,0 +1,507 @@ +"""Durable worker lifecycle and execution state transitions.""" + +from __future__ import annotations + +import asyncio +import random +import socket +import uuid +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import asynccontextmanager, suppress +from typing import TypeAlias, cast + +from hayhooks.durable.adapters import HaystackDurableAdapter +from hayhooks.durable.context import DurableContext, execution_context_scope +from hayhooks.durable.models import ( + ExecutionCanceledError, + ExecutionError, + ExecutionLeaseLostError, + ExecutionRecord, + ExecutionRecordSizeError, + ExecutionStatus, + ExecutionStoreError, + ExecutionSuspendedError, + JsonValue, + RetryableExecutionError, + validate_json, +) +from hayhooks.durable.store import ExecutionClaim, ExecutionStore +from hayhooks.server.logger import log + +RecordRunner: TypeAlias = Callable[[DurableContext], Awaitable[JsonValue]] + + +class SubmissionGate: + """Atomically admit submissions or close and wait for admitted work.""" + + def __init__(self) -> None: + self._condition = asyncio.Condition() + self._open = False + self._active = 0 + + @property + def open(self) -> bool: + return self._open + + def activate(self) -> None: + self._open = True + + @asynccontextmanager + async def admit(self) -> AsyncIterator[None]: + async with self._condition: + if not self._open: + msg = "Durable deployment is not accepting submissions" + raise RuntimeError(msg) + self._active += 1 + try: + yield + finally: + async with self._condition: + self._active -= 1 + if not self._active: + self._condition.notify_all() + + async def close_and_wait(self) -> None: + async with self._condition: + self._open = False + await self._condition.wait_for(lambda: self._active == 0) + + +class DurableExecutionManager: + """Bounded worker manager shared by REST and A2A adapters.""" + + def __init__( # noqa: PLR0913 + self, + name: str, + store: ExecutionStore, + runner: RecordRunner, + adapter: HaystackDurableAdapter, + *, + concurrency: int = 1, + poll_interval: float = 1.0, + shutdown_grace_period: float = 5.0, + max_attempts: int = 3, + retry_base_delay: float = 1.0, + retry_max_delay: float = 60.0, + ) -> None: + if concurrency < 1: + msg = "durable execution concurrency must be at least one" + raise ValueError(msg) + self.name = name + self.store = store + self.runner = runner + self.adapter = adapter + self.concurrency = concurrency + self.poll_interval = poll_interval + self.shutdown_grace_period = shutdown_grace_period + self.max_attempts = max(1, max_attempts) + self.retry_base_delay = max(0.0, retry_base_delay) + self.retry_max_delay = max(self.retry_base_delay, retry_max_delay) + self._workers: list[asyncio.Task[None]] = [] + self._draining_workers: set[asyncio.Task[None]] = set() + self._draining_runs: set[asyncio.Future[JsonValue]] = set() + self._maintenance_task: asyncio.Task[None] | None = None + self._maintenance_error_streak = 0 + self._prepared = False + self._started = False + self._accepting_claims = False + self._worker_generation = 0 + self._submission_gate = SubmissionGate() + + async def start(self) -> None: + """Prepare storage and activate this manager's workers.""" + if self._started: + self._accepting_claims = True + return + await self.prepare() + self.activate() + + async def prepare(self) -> None: + """Initialize storage while keeping workers and submissions disabled.""" + if self._prepared: + return + await self.store.initialize() + self._prepared = True + log.bind(deployment=self.name).debug("Prepared durable execution store") + + def activate(self) -> None: + """Start workers for an initialized deployment without an await gap.""" + if self._started: + self._accepting_claims = True + return + if not self._prepared: + msg = "durable execution manager must be prepared before activation" + raise RuntimeError(msg) + self._started = True + self._accepting_claims = True + self._submission_gate.activate() + self._worker_generation += 1 + generation = self._worker_generation + identity = f"{socket.gethostname()}-{uuid.uuid4().hex[:8]}" + self._workers = [self._start_worker(identity, slot, generation) for slot in range(self.concurrency)] + self._start_maintenance(generation) + log.bind(deployment=self.name, workers=self.concurrency).debug("Activated durable execution workers") + + def deactivate(self) -> None: + self._accepting_claims = False + self._worker_generation += 1 + + async def quiesce(self) -> None: + """Close admission before workers stop or stranded work is counted.""" + await self._submission_gate.close_and_wait() + self.deactivate() + log.bind(deployment=self.name).debug("Quiesced durable execution manager") + + @property + def started(self) -> bool: + return self._started + + @property + def accepting(self) -> bool: + return self._started and self._accepting_claims and self._submission_gate.open + + @property + def health(self) -> dict[str, JsonValue]: + """Return a payload-safe worker projection for readiness and diagnostics.""" + running = sum(not worker.done() for worker in self._workers) + maintenance_running = self._maintenance_task is not None and not self._maintenance_task.done() + maintenance_healthy = self._maintenance_task is None or ( + maintenance_running and self._maintenance_error_streak == 0 + ) + return { + "healthy": not self._prepared + or (self._started and self._accepting_claims and running == self.concurrency and maintenance_healthy), + "configured_slots": self.concurrency, + "running_slots": running, + "draining_slots": sum(not worker.done() for worker in self._draining_workers), + "draining_runs": sum(not runner.done() for runner in self._draining_runs), + "maintenance_running": maintenance_running, + "accepting": self.accepting, + } + + async def health_snapshot(self) -> dict[str, JsonValue]: + """Add storage-level queue/state counts to the local readiness view.""" + health = self.health + try: + health["counts"] = cast(JsonValue, await self.store.operational_counts()) + except Exception as error: + health["healthy"] = False + health["operational_error"] = type(error).__name__ + return health + + @property + def draining(self) -> bool: + return any(not worker.done() for worker in self._draining_workers) or any( + not runner.done() for runner in self._draining_runs + ) + + async def wait_drained(self) -> None: + """Wait for detached application work before closing shared storage.""" + pending = [worker for worker in self._draining_workers if not worker.done()] + pending_runs = [runner for runner in self._draining_runs if not runner.done()] + if pending or pending_runs: + await asyncio.gather(*pending, *pending_runs, return_exceptions=True) + + async def close(self) -> None: + """Stop claims and retain lease-lost work until application code exits.""" + await self.quiesce() + if self._maintenance_task is not None: + self._maintenance_task.cancel() + with suppress(asyncio.CancelledError): + await self._maintenance_task + self._maintenance_task = None + if not self._workers: + self._started = False + log.bind(deployment=self.name).debug("Closed durable execution manager") + return + done, pending = await asyncio.wait(self._workers, timeout=self.shutdown_grace_period) + if pending: + self._draining_workers.update(pending) + for worker in pending: + worker.add_done_callback(self._draining_workers.discard) + log.warning( + "{} | {} durable worker slot(s) exceeded the {:.2f}s shutdown grace period; " + "claims remain fenced and heartbeating until application work exits", + self.name, + len(pending), + self.shutdown_grace_period, + ) + self._log_worker_failures(done) + self._workers = [] + self._started = False + log.bind(deployment=self.name, draining=len(pending)).debug("Closed durable execution manager") + + async def submit_with_record(self, record: ExecutionRecord) -> tuple[bool, ExecutionRecord]: + """Admit one submission while the deployment gate remains open.""" + async with self._submission_gate.admit(): + return await self.store.submit_with_record(record) + + def _start_worker(self, identity: str, slot: int, generation: int) -> asyncio.Task[None]: + worker_name = f"{self.name}:{identity}:{slot}" + worker = asyncio.create_task( + self._worker(worker_name, generation), + name=f"durable:{self.name}:{slot}", + ) + worker.add_done_callback(lambda completed: self._worker_done(identity, slot, generation, completed)) + return worker + + def _start_maintenance(self, generation: int) -> None: + self._maintenance_task = asyncio.create_task( + self._maintenance_loop(self.store.maintain, self.poll_interval, generation), + name=f"durable-maintenance:{self.name}", + ) + + async def _maintenance_loop( + self, + maintain: Callable[[], Awaitable[None]], + interval: float, + generation: int, + ) -> None: + """Supervise bounded store maintenance independently from worker claims.""" + while self._accepting_claims and generation == self._worker_generation: + try: + await maintain() + except asyncio.CancelledError: + raise + except Exception as error: + self._maintenance_error_streak += 1 + log.opt(exception=error).warning("{} | durable maintenance failed; retrying", self.name) + else: + self._maintenance_error_streak = 0 + await asyncio.sleep(interval) + + def _worker_done( + self, + identity: str, + slot: int, + generation: int, + worker: asyncio.Task[None], + ) -> None: + if self._accepting_claims and generation == self._worker_generation: + error = None if worker.cancelled() else worker.exception() + if error is not None: + log.opt(exception=error).error( + "{} | durable worker slot {} stopped unexpectedly; restarting", + self.name, + slot, + ) + else: + log.error("{} | durable worker slot {} stopped unexpectedly; restarting", self.name, slot) + replacement = self._start_worker(identity, slot, generation) + try: + index = self._workers.index(worker) + except ValueError: + replacement.cancel() + else: + self._workers[index] = replacement + + def _log_worker_failures(self, workers: set[asyncio.Task[None]]) -> None: + for worker in workers: + if worker.cancelled(): + continue + error = worker.exception() + if error is not None: + log.opt(exception=error).warning( + "{} | durable worker '{}' ended with an error during shutdown", + self.name, + worker.get_name(), + ) + + async def _worker(self, worker_name: str, generation: int) -> None: + consecutive_store_errors = 0 + while self._accepting_claims and generation == self._worker_generation: + try: + claim = await self.store.claim_next(worker_name) + except asyncio.CancelledError: + raise + except Exception as error: + consecutive_store_errors += 1 + await self._backoff_store_error(error, consecutive_store_errors, operation="claim") + continue + if claim is None: + await asyncio.sleep(self.poll_interval) + continue + if not self._accepting_claims or generation != self._worker_generation: + return + attempt_log = log.bind( + deployment=self.name, + execution_id=claim.record.execution_id, + attempt=claim.record.attempt, + ) + attempt_log.debug("Claimed durable execution") + try: + await self._process_claim(claim) + consecutive_store_errors = 0 + attempt_log.bind(status=claim.record.status.value).debug("Finished durable execution attempt") + except ExecutionLeaseLostError: + consecutive_store_errors = 0 + log.warning("{} | lost durable execution claim {}", self.name, claim.record.execution_id) + except asyncio.CancelledError: + raise + except Exception as error: + consecutive_store_errors += 1 + await self._backoff_store_error(error, consecutive_store_errors, operation="transition") + + async def _process_claim(self, claim: ExecutionClaim) -> None: # noqa: C901, PLR0912 - explicit attempt outcomes + """Run one fenced claim and leave every terminal decision to the store.""" + async with claim: + # A reclaimed delivery has already incremented ``attempt`` in the + # store. Do not invoke application code again after the total + # execution-attempt budget has been consumed by prior crashes, + # lease losses, or explicit retries. + if claim.record.attempt > self.max_attempts: + await self._terminalize_attempt_exhaustion(claim) + return + + context = DurableContext(claim, self.adapter) + try: + if await claim.cancellation_requested(): + await self._complete_canceled(claim) + return + + with execution_context_scope(context): + result = await self._run_with_lease_guard(claim, context) + except ExecutionSuspendedError: + pass + except ExecutionCanceledError: + await self._complete_canceled(claim) + except RetryableExecutionError as error: + if claim.record.attempt >= self.max_attempts: + await self._terminalize_attempt_exhaustion(claim) + return + + exponent = min(max(0, claim.record.attempt - 1), 30) + delay = error.delay if error.delay > 0 else self.retry_base_delay * (2**exponent) + await claim.retry( + ExecutionError.from_exception(error, retryable=True), + delay=min(max(0.0, delay), self.retry_max_delay), + ) + except ExecutionRecordSizeError: + await self._fail_oversized_record(claim) + except (ExecutionLeaseLostError, ExecutionStoreError, asyncio.CancelledError): + raise + except Exception as error: + claim.record.mark_failed(error) + await claim.complete() + else: + try: + # A cancellation accepted after the runner returns wins over + # the result, so clients never observe a completed canceled run. + if await claim.cancellation_requested(): + claim.record.mark_canceled() + else: + claim.record.result = validate_json(result, limit=claim.record.max_record_bytes, label="result") + claim.record.error = None + claim.record.status = ExecutionStatus.COMPLETED + claim.record.wait = None + claim.record.retry_at = None + + claim.record.touch() + await claim.complete() + except ExecutionRecordSizeError: + await self._fail_oversized_record(claim) + + async def _complete_canceled(self, claim: ExecutionClaim) -> None: + """Terminalize a cooperative cancellation with the current ownership fence.""" + claim.record.mark_canceled() + await claim.complete() + + async def _terminalize_attempt_exhaustion(self, claim: ExecutionClaim) -> None: + """Fail an execution without allowing an attempt beyond the configured bound.""" + exhausted = ExecutionError( + type="RetryExhausted", + message=f"Execution exhausted its {self.max_attempts} permitted attempts", + retryable=False, + code="retry_exhausted", + ) + claim.record.mark_failed(exhausted) + claim.record.append_progress("Execution retry limit reached", kind="retry_exhausted") + await claim.complete() + + async def _run_with_lease_guard(self, claim: ExecutionClaim, context: DurableContext) -> JsonValue: + """ + Stop cooperative work when a store reports definitive lease loss. + + Concrete stores expose an event that flips once the fenced lease is + definitively lost. + """ + lost_event = claim.lost_event + runner = asyncio.ensure_future(self.runner(context)) + loss_waiter = asyncio.create_task( + lost_event.wait(), + name=f"durable-lease-watch:{self.name}:{claim.record.execution_id}", + ) + try: + done, _ = await asyncio.wait({runner, loss_waiter}, return_when=asyncio.FIRST_COMPLETED) + if runner in done and not lost_event.is_set(): + return runner.result() + if not runner.done(): + runner.cancel() + self._track_draining_run(runner, claim.record.execution_id) + else: + # Retrieve a finished task's result so an application exception + # is not reported as an unhandled background-task failure. + with suppress(asyncio.CancelledError, Exception): + runner.result() + msg = f"Execution lease for '{claim.record.execution_id}' was lost" + raise ExecutionLeaseLostError(msg) + except asyncio.CancelledError: + if not runner.done(): + runner.cancel() + self._track_draining_run(runner, claim.record.execution_id) + raise + finally: + loss_waiter.cancel() + with suppress(asyncio.CancelledError): + await loss_waiter + + def _track_draining_run(self, runner: asyncio.Future[JsonValue], execution_id: str) -> None: + """Retain non-cooperative work until it exits after a lost lease.""" + if runner.done() or runner in self._draining_runs: + return + self._draining_runs.add(runner) + + def completed(done: asyncio.Future[JsonValue]) -> None: + self._draining_runs.discard(done) + if done.cancelled(): + return + with suppress(Exception): + done.result() + log.warning("{} | lease-lost durable work drained for execution {}", self.name, execution_id) + + runner.add_done_callback(completed) + + async def _fail_oversized_record(self, claim: ExecutionClaim) -> None: + """Persist a small terminal record after application state exceeded its bound.""" + record = claim.record + record.validated_input = {} + record.checkpoint = None + record.application_state = {} + record.wait = None + record.progress = [] + record.result = None + record.error = None + record.last_retry_error = None + record.retry_at = None + record.mark_failed( + ExecutionError( + type="ExecutionRecordTooLarge", + message=f"Execution exceeded its {record.max_record_bytes}-byte durable record limit", + retryable=False, + code="record_too_large", + ) + ) + await claim.complete() + + async def _backoff_store_error(self, error: BaseException, failures: int, *, operation: str) -> None: + exponent = min(max(0, failures - 1), 10) + ceiling = min(max(self.poll_interval, 0.01) * (2**exponent), 5.0) + delay = random.uniform(ceiling / 2, ceiling) # noqa: S311 - jitter is not security-sensitive + log.opt(exception=error).warning( + "{} | durable worker store {} failed; retrying in {:.2f}s: {}", + self.name, + operation, + delay, + error, + ) + await asyncio.sleep(delay) diff --git a/tests/test_durable_engine.py b/tests/test_durable_engine.py new file mode 100644 index 00000000..70418120 --- /dev/null +++ b/tests/test_durable_engine.py @@ -0,0 +1,151 @@ +"""Reducer lifecycle invariants for the lean durable engine.""" + +from __future__ import annotations + +from dataclasses import replace + +import pytest + +from hayhooks.durable.engine import ( + Checkpoint, + Claim, + Complete, + ExecutionLeaseLostError, + ExecutionStatus, + Fail, + Heartbeat, + InvalidExecutionTransitionError, + PayloadKind, + RecoverExpiredLease, + RequestCancellation, + Resume, + ScheduleRetry, + Suspend, + decide, + initial_control, + submission_plan, +) + + +def control(**changes: object): + defaults = { + "run_id": "run-1", + "idempotency_digest": "idem", + "idempotency_binding_digest": "binding", + "deployment": "deployment", + "definition_revision": "rev-1", + "owner_id": "owner", + "kind": "pipeline", + "now_ms": 100, + } + defaults.update(changes) + return initial_control(**defaults) + + +def claim(current, *, now_ms: int = 200): + return decide(current, Claim("worker-a", now_ms, 500, 3, "rev-1")) + + +def test_submission_is_an_initial_control_and_input_only() -> None: + plan = submission_plan(control(), b"input") + assert plan.next_control.version == 1 + assert plan.payload_writes[0].kind is PayloadKind.INPUT + assert plan.lease_index_update is None + + +def test_claim_fences_and_heartbeat_renews_without_business_version() -> None: + claimed = claim(control()).next_control + assert (claimed.status, claimed.version, claimed.fence, claimed.run_attempt) == ( + ExecutionStatus.RUNNING, + 2, + 1, + 1, + ) + heartbeat = decide(claimed, Heartbeat(1, "worker-a", 300, 500)).next_control + assert heartbeat.version == claimed.version + assert heartbeat.lease_expires_at_ms == 800 + with pytest.raises(ExecutionLeaseLostError): + decide(heartbeat, Heartbeat(0, "worker-a", 301, 500)) + with pytest.raises(ExecutionLeaseLostError): + decide(heartbeat, Heartbeat(1, "worker-a", 750, 500, 50)) + + +@pytest.mark.parametrize( + "command", + [ + lambda c: Checkpoint(c.fence, "stale", 300, 500, b"checkpoint"), + lambda c: Complete(c.fence + 1, "worker-a", 300, b"result"), + lambda c: ScheduleRetry(c.fence + 1, "worker-a", 300, 100, 2, b"error"), + lambda c: Suspend(c.fence + 1, "worker-a", 300, b"checkpoint", b"wait"), + ], +) +def test_owned_transitions_reject_stale_fences(command) -> None: + with pytest.raises(ExecutionLeaseLostError): + decide(claim(control()).next_control, command(claim(control()).next_control)) + + +def test_cancellation_wins_completion_and_retry() -> None: + claimed = claim(control()).next_control + canceled = decide(claimed, RequestCancellation(250, "💥" * 2_000)).next_control + assert canceled.cancel_reason == "💥" * 1_024 + assert len(canceled.cancel_reason.encode()) == 4_096 + terminal = decide(canceled, Complete(1, "worker-a", 300, b"result")) + assert terminal.next_control.status is ExecutionStatus.CANCELED + assert not terminal.payload_writes + retried = decide(canceled, ScheduleRetry(1, "worker-a", 300, 100, 2, b"error")) + assert retried.next_control.status is ExecutionStatus.CANCELED + +def test_retry_and_lease_recovery_requeue_without_resetting_retry_count() -> None: + claimed = claim(control()).next_control + retry = decide(claimed, ScheduleRetry(1, "worker-a", 300, 100, 2, b"retry")) + queued = retry.next_control + assert (queued.status, queued.available_at_ms, queued.run_attempt, queued.application_retry_count) == ( + ExecutionStatus.QUEUED, + 400, + 1, + 1, + ) + next_claim = claim(replace(queued, available_at_ms=None), now_ms=400).next_control + recovered = decide( + next_claim, + RecoverExpiredLease(1_000, next_claim.fence, next_claim.lease_expires_at_ms or 0, 3, "rev-1"), + ) + assert recovered.next_control.status is ExecutionStatus.QUEUED + assert recovered.next_control.application_retry_count == 1 + assert recovered.next_control.run_attempt == 2 + + +def test_wait_resume_and_progress_preserve_checkpoint_boundary() -> None: + claimed = claim(control()).next_control + waiting = decide( + claimed, + Suspend(1, "worker-a", 300, b"checkpoint", b"wait", progress_events=(b"waiting",)), + ) + assert waiting.next_control.status is ExecutionStatus.WAITING + assert waiting.lease_index_update and waiting.lease_index_update.deadline_ms is None + resumed = decide(waiting.next_control, Resume(400, "rev-1", b"checkpoint", (b"resumed",))) + assert resumed.next_control.status is ExecutionStatus.QUEUED + assert [event.sequence for event in (*waiting.progress_events, *resumed.progress_events)] == [1, 2] + + +def test_terminal_state_is_irreversible_and_payloads_are_exclusive() -> None: + completed = decide(claim(control()).next_control, Complete(1, "worker-a", 300, b"result")) + assert completed.payload_deletes == (PayloadKind.ERROR, PayloadKind.WAIT) + assert decide(completed.next_control, RequestCancellation(400, "late")).next_control == completed.next_control + with pytest.raises(InvalidExecutionTransitionError): + decide(completed.next_control, Claim("worker-a", 400, 500, 3, "rev-1")) + failed = decide(claim(control()).next_control, Fail(1, "worker-a", 300, b"error")) + assert failed.payload_deletes == (PayloadKind.RESULT, PayloadKind.WAIT) + + +def test_revision_mismatch_never_grants_a_fence() -> None: + plan = decide(control(), Claim("worker-a", 200, 500, 3, "rev-2")) + assert plan.next_control.status is ExecutionStatus.FAILED + assert plan.next_control.fence == 0 + + +def test_stale_lease_index_is_removed_without_changing_control() -> None: + current = control() + plan = decide(current, RecoverExpiredLease(200, 1, 100, 3, "rev-1")) + assert plan.next_control == current + assert plan.lease_index_update and plan.lease_index_update.deadline_ms is None diff --git a/tests/test_durable_execution.py b/tests/test_durable_execution.py new file mode 100644 index 00000000..d42112b7 --- /dev/null +++ b/tests/test_durable_execution.py @@ -0,0 +1,1171 @@ +import asyncio +import importlib.metadata +import json +import threading +import time +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient +from haystack import Pipeline, component +from haystack.components.agents import Agent +from haystack.components.agents.state import State +from haystack.core.errors import PipelineRuntimeError +from haystack.dataclasses import ChatMessage, ToolCall +from haystack.dataclasses.breakpoints import PipelineSnapshot +from haystack.tools import Tool +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper, DurableContext +from hayhooks.durable.adapters import ( + HaystackDurableAdapter, + _agent_exits_after_tools, + _checkpoint_agent_state, + _checkpoint_data, + _restore_agent_state, +) +from hayhooks.durable.context import execution_context_scope +from hayhooks.durable.models import ExecutionCheckpoint, ExecutionKind, ExecutionRecord, ExecutionStatus +from hayhooks.durable.runtime import ( + DurableDeployment, + DurableRuntime, + _canonical_json, + _operation_fingerprint, + durable_runtime, +) +from hayhooks.durable.store import InMemoryExecutionStoreProvider +from hayhooks.server.a2a.app import create_a2a_app +from hayhooks.server.app import create_app +from hayhooks.server.logger import log +from hayhooks.server.pipelines.registry import registry +from hayhooks.server.utils.deploy_utils import add_pipeline_api_route +from hayhooks.server.utils.mcp_utils import create_mcp_server, create_starlette_app +from hayhooks.server.utils.module_loader import ( + _set_method_implementation_flags, + create_pipeline_wrapper_instance, + load_pipeline_module, + unload_pipeline_modules, +) +from hayhooks.settings import settings + +pytestmark = pytest.mark.skipif( + not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" +) + +_DURABLE_EXECUTION_EXAMPLE = Path("examples/durable_execution/pipelines/durable_job") +_DURABLE_A2A_EXAMPLE = Path("examples/a2a_long_running/pipelines/long_running_agent") + + +class Request(BaseModel): + value: int + + +class Result(BaseModel): + value: int + + +class ResumeInput(BaseModel): + approved: bool + + +def test_set_valued_input_has_a_stable_idempotency_fingerprint() -> None: + class SetRequest(BaseModel): + tags: set[str] + + canonical = _canonical_json(SetRequest(tags={"zeta", "alpha"}).model_dump(mode="python")) + assert canonical == {"tags": ["alpha", "zeta"]} + assert _operation_fingerprint("job", "revision", canonical, owner_id=None) == _operation_fingerprint( + "job", "revision", {"tags": ["alpha", "zeta"]}, owner_id=None + ) + assert _operation_fingerprint("job", "revision", canonical, owner_id=None) != _operation_fingerprint( + "job", "revision", {"tags": ["zeta", "alpha"]}, owner_id=None + ) + + +async def test_portable_runtime_starts_its_own_deployments() -> None: + wrapper = Wrapper() + wrapper.setup() + runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + deployment = runtime.deployment("portable", wrapper) + + try: + await runtime.start() + assert runtime.started + assert deployment.manager.started + assert deployment.manager.accepting + finally: + await runtime.close() + + +def _create_mcp_app(): + return create_starlette_app(create_mcp_server()) + + +@pytest.mark.parametrize( + "app_factory", + [create_app, create_a2a_app, _create_mcp_app], + ids=["rest", "a2a", "mcp"], +) +async def test_app_lifespans_close_durable_runtime_when_start_fails(monkeypatch, app_factory) -> None: + closed = 0 + + async def fail_start() -> None: + msg = "durable startup failed" + raise RuntimeError(msg) + + async def close() -> None: + nonlocal closed + closed += 1 + + monkeypatch.setattr(settings, "pipelines_dir", "") + monkeypatch.setattr(durable_runtime, "start", fail_start) + monkeypatch.setattr(durable_runtime, "close", close) + app = app_factory() + + with pytest.raises(Exception) as exc_info: + async with app.router.lifespan_context(app): + pass + + assert "durable startup failed" in repr(exc_info.value) + assert closed == 1 + + +def _checkpoint_test_tool(value: str) -> str: + return value + + +class _Claim: + def __init__(self, record: ExecutionRecord) -> None: + self.record = record + self.checkpoints = 0 + + async def checkpoint(self) -> None: + self.checkpoints += 1 + + async def cancellation_requested(self) -> bool: + return False + + +def _agent_record(execution_id: str, deployment_name: str = "agent") -> ExecutionRecord: + return ExecutionRecord( + execution_id=execution_id, + execution_kind=ExecutionKind.AGENT, + deployment_name=deployment_name, + definition_revision="revision", + validated_input={"messages": []}, + ) + + +class Wrapper(BasePipelineWrapper): + durable_revision = "test-wrapper" + + def setup(self) -> None: + self.pipeline = Pipeline() + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + await context.report_progress("working") + return Result(value=request.value + 1) + + +class InvalidResultWrapper(Wrapper): + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + return {"value": "not-an-integer"} # type: ignore[return-value] + + +class SyncWrapper(BasePipelineWrapper): + durable_revision = "sync-wrapper" + + def setup(self) -> None: + self.pipeline = Pipeline() + + def run_durable(self, context: DurableContext, request: Request) -> Result: + context.report_progress_sync("working in a worker thread") + context.check_cancelled_sync() + return Result(value=request.value + 2) + + +class BlockingAdmissionWrapper(BasePipelineWrapper): + durable_revision = "blocking-admission" + + def setup(self) -> None: + self.pipeline = Pipeline() + self.started = threading.Event() + self.release = threading.Event() + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + self.started.set() + await asyncio.to_thread(self.release.wait) + return Result(value=request.value + 1) + + +async def test_durable_lifecycle_logs_identifiers_without_payload(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_poll_interval", 0.01) + wrapper = Wrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("logged-job", wrapper, InMemoryExecutionStoreProvider()) + records = [] + sink = log.add(lambda message: records.append(message.record), level="DEBUG") + try: + await deployment.start() + await deployment.submit({"value": 8675309}, execution_id="logged-execution") + expected = { + "Accepted durable execution submission", + "Claimed durable execution", + "Finished durable execution attempt", + } + for _ in range(200): + if expected <= {record["message"] for record in records}: + break + await asyncio.sleep(0.005) + else: + pytest.fail("durable lifecycle logs were not emitted") + finally: + await deployment.close() + log.remove(sink) + + lifecycle = [record for record in records if record["message"] in expected] + assert all(record["extra"]["deployment"] == "logged-job" for record in lifecycle) + assert all(record["extra"]["execution_id"] == "logged-execution" for record in lifecycle) + assert "8675309" not in str(lifecycle) + + +async def test_quiesce_waits_for_an_admitted_submission_and_rejects_later_ones(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_poll_interval", 0.01) + monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) + wrapper = Wrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("admission-gate", wrapper, InMemoryExecutionStoreProvider()) + await deployment.start() + entered = asyncio.Event() + release = asyncio.Event() + submit_with_record = deployment.store.submit_with_record + + async def paused_submit(record): + entered.set() + await release.wait() + return await submit_with_record(record) + + monkeypatch.setattr(deployment.store, "submit_with_record", paused_submit) + submission = asyncio.create_task(deployment.submit({"value": 1})) + await entered.wait() + quiescing = asyncio.create_task(deployment.quiesce()) + await asyncio.sleep(0) + assert not quiescing.done() + + release.set() + await submission + await quiescing + assert (await deployment.store.operational_counts())["nonterminal"] == 1 + with pytest.raises(RuntimeError, match="not accepting submissions"): + await deployment.submit({"value": 2}) + await deployment.close() + + +async def test_deployment_claims_reject_incompatible_work_without_a_revision_scan() -> None: + provider = InMemoryExecutionStoreProvider() + store = provider.create_execution_store("revision-safe") + await store.initialize() + + for execution_id, revision in (("waiting-old", "old"), ("waiting-current", "current")): + assert await store.submit( + ExecutionRecord( + execution_id=execution_id, + execution_kind=ExecutionKind.PIPELINE, + deployment_name="revision-safe", + definition_revision=revision, + validated_input={"value": 1}, + ) + ) + store.set_definition_revision(revision) + waiting = await store.claim_next("worker") + assert waiting is not None + async with waiting: + waiting.record.status = ExecutionStatus.WAITING + waiting.record.wait = {"kind": "input"} + await waiting.suspend() + assert await store.submit( + ExecutionRecord( + execution_id="queued-old", + execution_kind=ExecutionKind.PIPELINE, + deployment_name="revision-safe", + definition_revision="old", + validated_input={"value": 1}, + ) + ) + + wrapper = Wrapper() + wrapper.durable_revision = "current" + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("revision-safe", wrapper, provider) + await deployment.start() + try: + for _ in range(100): + queued_old = await store.get("queued-old") + if queued_old is not None and queued_old.terminal: + break + await asyncio.sleep(0.001) + else: + pytest.fail("incompatible queued work was not rejected by its first claim") + waiting_old = await store.get("waiting-old") + waiting_current = await store.get("waiting-current") + finally: + await deployment.close() + + assert queued_old is not None + assert queued_old.status is ExecutionStatus.FAILED + assert queued_old.error is not None + assert waiting_old is not None + assert waiting_old.status is ExecutionStatus.WAITING + assert waiting_current is not None + assert waiting_current.status is ExecutionStatus.WAITING + + +@component +class _CheckpointIncrement: + def __init__(self, *, fail_once: bool = False) -> None: + self.fail_once = fail_once + self.calls = 0 + + @component.output_types(value=int) + def run(self, value: int) -> dict[str, int]: + self.calls += 1 + if self.fail_once and self.calls == 1: + msg = "interrupted component" + raise RuntimeError(msg) + return {"value": value + 1} + + +class CheckpointPipelineWrapper(BasePipelineWrapper): + durable_revision = "checkpoint-pipeline" + + def setup(self) -> None: + self.first = _CheckpointIncrement() + self.second = _CheckpointIncrement(fail_once=True) + self.pipeline = Pipeline() + self.pipeline.add_component("first", self.first) + self.pipeline.add_component("second", self.second) + self.pipeline.connect("first.value", "second.value") + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + try: + result = await context.run_pipeline_async( + {"first": {"value": request.value}}, + checkpoint_at=["first", "second"], + ) + except PipelineRuntimeError: + await context.retry("retry interrupted pipeline", delay=0) + return Result(value=result["second"]["value"]) + + +class WaitingWrapper(Wrapper): + durable_resume_model = ResumeInput + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + if context.resume_input is None: + await context.suspend( + { + "kind": "approval", + "message": "Approve this job", + "expected_input_schema": ResumeInput.model_json_schema(), + "private_tool_arguments": {"secret": True}, + } + ) + resume = ResumeInput.model_validate(context.take_resume_input()) + return Result(value=request.value if resume.approved else -1) + + +@component +class FakeChatGenerator: + @component.output_types(replies=list[ChatMessage]) + def run(self, messages: list[ChatMessage], tools=None): + return {"replies": [ChatMessage.from_assistant("done")]} + + +class AgentRequest(BaseModel): + message: str + + +class AgentWrapper(BasePipelineWrapper): + durable_revision = "test-agent" + + def setup(self) -> None: + self.pipeline = Agent(chat_generator=FakeChatGenerator(), tools=[]) + + async def run_durable_async(self, context: DurableContext, request: AgentRequest) -> dict: + return await context.run_agent_async(messages=[ChatMessage.from_user(request.message)]) + + +class BuiltinAgentWrapper(BasePipelineWrapper): + durable = True + durable_revision = "builtin-agent" + + def setup(self) -> None: + self.pipeline = Agent(chat_generator=FakeChatGenerator(), tools=[]) + + +@pytest.fixture(autouse=True) +def clean_registry(): + registry.clear() + yield + registry.clear() + + +def _durable_app(monkeypatch, name, wrapper): + monkeypatch.setattr(settings, "durable_store", "memory") + wrapper.setup() + _set_method_implementation_flags(wrapper) + registry.add(name, wrapper) + app = create_app() + add_pipeline_api_route(app, name, wrapper) + return app + + +def _wait_for_status(client, url, status, message): + for _ in range(100): + record = client.get(url) + if record.json()["status"] == status: + return record + time.sleep(0.01) + pytest.fail(message) + + +def test_durable_rest_submission_is_direct_typed_and_idempotent(monkeypatch) -> None: + app = _durable_app(monkeypatch, "job", Wrapper()) + + with TestClient(app) as client: + submitted = client.post("/job/run-durable", json={"value": 4}, headers={"Idempotency-Key": "same"}) + duplicate = client.post("/job/run-durable", json={"value": 4}, headers={"Idempotency-Key": "same"}) + assert submitted.status_code == 202 + assert duplicate.headers["Idempotent-Replay"] == "true" + body = submitted.json() + assert set(body) == { + "execution_id", + "status", + "attempt", + "sequence", + "progress", + "result", + "error", + "waiting", + "cancellation_requested_at", + "created_at", + "updated_at", + "links", + } + inspected = _wait_for_status(client, body["links"]["self"], "completed", "durable execution did not complete") + + assert inspected.json()["result"] == {"value": 5} + assert submitted.headers["Location"] == body["links"]["self"] + + +def test_durable_rest_admission_is_retryable_and_preserves_owner_isolation(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_trusted_owner_header", "X-Authenticated-Owner") + monkeypatch.setattr(settings, "durable_max_nonterminal_executions", 1) + wrapper = BlockingAdmissionWrapper() + app = _durable_app(monkeypatch, "admission", wrapper) + + try: + with TestClient(app) as client: + alice_headers = {"Idempotency-Key": "alice-first", "X-Authenticated-Owner": "alice"} + first = client.post("/admission/run-durable", json={"value": 1}, headers=alice_headers) + assert first.status_code == 202 + assert wrapper.started.wait(timeout=1) + + replay = client.post("/admission/run-durable", json={"value": 1}, headers=alice_headers) + assert replay.status_code == 202 + assert replay.headers["Idempotent-Replay"] == "true" + + globally_limited = client.post( + "/admission/run-durable", + json={"value": 2}, + headers={"Idempotency-Key": "alice-second", "X-Authenticated-Owner": "alice"}, + ) + assert globally_limited.status_code == 503 + assert globally_limited.headers["Retry-After"] == "1" + assert "deployment_nonterminal" in globally_limited.json()["detail"] + + bob = client.post( + "/admission/run-durable", + json={"value": 3}, + headers={"Idempotency-Key": "bob-first", "X-Authenticated-Owner": "bob"}, + ) + assert bob.status_code == 503 + wrapper.release.set() + finally: + wrapper.release.set() + + +def test_durable_rest_can_inspect_and_cancel_an_execution_from_an_old_revision(monkeypatch) -> None: + wrapper = BlockingAdmissionWrapper() + app = _durable_app(monkeypatch, "rolling", wrapper) + + try: + with TestClient(app) as client: + submitted = client.post("/rolling/run-durable", json={"value": 1}) + assert wrapper.started.wait(timeout=5) + deployment = durable_runtime.current_deployment("rolling") + assert deployment is not None + deployment.revision = "replacement" + + links = submitted.json()["links"] + assert client.get(links["self"]).status_code == 200 + canceled = client.post(links["cancel"]) + assert canceled.status_code == 202 + assert canceled.json()["cancellation_requested_at"] is not None + wrapper.release.set() + finally: + wrapper.release.set() + + +def test_durable_result_annotation_is_validated_before_completion(monkeypatch) -> None: + app = _durable_app(monkeypatch, "invalid-result", InvalidResultWrapper()) + + with TestClient(app) as client: + submitted = client.post("/invalid-result/run-durable", json={"value": 4}) + url = submitted.json()["links"]["self"] + inspected = _wait_for_status(client, url, "failed", "invalid durable result did not become a terminal failure") + + assert inspected.json()["result"] is None + assert inspected.json()["error"] == { + "type": "ValueError", + "message": "Durable method result does not match its declared return annotation", + "retryable": False, + "code": None, + } + + +def test_durable_rest_rejects_mismatched_idempotency_payload(monkeypatch) -> None: + app = _durable_app(monkeypatch, "job", Wrapper()) + + with TestClient(app) as client: + first = client.post("/job/run-durable", json={"value": 4}, headers={"Idempotency-Key": "same"}) + conflict = client.post("/job/run-durable", json={"value": 5}, headers={"Idempotency-Key": "same"}) + + assert first.status_code == 202 + assert conflict.status_code == 409 + + +def test_durable_rest_uses_the_execution_id_key_grammar_at_every_boundary(monkeypatch) -> None: + app = _durable_app(monkeypatch, "job", Wrapper()) + + with TestClient(app) as client: + accepted = client.post("/job/run-durable", json={"value": 4}, headers={"Idempotency-Key": "a" * 128}) + rejected = [ + client.post("/job/run-durable", json={"value": 4}, headers={"Idempotency-Key": key}) + for key in ("part/child", "part.child", "part~child", "a" * 129) + ] + invalid_paths = [ + client.get("/job/executions/part.child"), + client.post("/job/executions/part.child/cancel"), + client.post("/job/executions/part.child/resume"), + ] + + assert accepted.status_code == 202 + assert all(response.status_code == 422 for response in rejected + invalid_paths) + + +def test_durable_rest_bounds_owner_scoped_idempotency_keys(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_trusted_owner_header", "X-Authenticated-Owner") + app = _durable_app(monkeypatch, "job", Wrapper()) + + with TestClient(app) as client: + accepted = client.post( + "/job/run-durable", + json={"value": 4}, + headers={"Idempotency-Key": "a" * 63, "X-Authenticated-Owner": "owner"}, + ) + rejected = client.post( + "/job/run-durable", + json={"value": 4}, + headers={"Idempotency-Key": "a" * 64, "X-Authenticated-Owner": "owner"}, + ) + + assert accepted.status_code == 202 + assert rejected.status_code == 422 + assert "63" in rejected.json()["detail"] + + +def test_durable_rest_maps_oversized_validated_request_to_422(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_max_record_bytes", 5) + app = _durable_app(monkeypatch, "job", Wrapper()) + + with TestClient(app) as client: + response = client.post("/job/run-durable", json={"value": 123}) + + assert response.status_code == 422 + assert "durable execution limit" in response.json()["detail"] + + +def test_durable_waiting_resume_is_typed_private_and_revision_safe(monkeypatch) -> None: + app = _durable_app(monkeypatch, "approval", WaitingWrapper()) + + with TestClient(app) as client: + submitted = client.post("/approval/run-durable", json={"value": 7}) + url = submitted.json()["links"]["self"] + waiting = _wait_for_status(client, url, "waiting", "execution did not wait") + assert waiting.json()["waiting"] == { + "kind": "approval", + "message": "Approve this job", + "expected_input_schema": ResumeInput.model_json_schema(), + } + deployment = durable_runtime.current_deployment("approval") + assert deployment is not None + revision = deployment.revision + deployment.revision = "replacement" + missing = client.post(f"{url}/resume") + assert missing.status_code == 422 + invalid = client.post(f"{url}/resume", json={"approved": "not-a-bool"}) + assert invalid.status_code == 422 + resumed = client.post(f"{url}/resume", json={"approved": True}) + assert resumed.status_code == 202 + deployment.revision = revision + completed = _wait_for_status(client, url, "completed", "resumed execution did not complete") + + assert completed.json()["result"] == {"value": 7} + openapi = app.openapi() + resume_schema = openapi["paths"]["/approval/executions/{execution_id}/resume"]["post"]["requestBody"]["content"][ + "application/json" + ]["schema"] + assert "ResumeInput" in str(resume_schema) + + +def test_durable_rest_enforces_configured_trusted_owner_header(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_trusted_owner_header", "X-Authenticated-Owner") + app = _durable_app(monkeypatch, "owned", Wrapper()) + + with TestClient(app) as client: + assert client.post("/owned/run-durable", json={"value": 1}).status_code == 401 + submitted = client.post( + "/owned/run-durable", + json={"value": 1}, + headers={"X-Authenticated-Owner": "alice"}, + ) + assert submitted.status_code == 202 + url = submitted.json()["links"]["self"] + assert client.get(url, headers={"X-Authenticated-Owner": "bob"}).status_code == 404 + assert client.get(url, headers={"X-Authenticated-Owner": "alice"}).status_code == 200 + oversized = client.post( + "/owned/run-durable", + json={"value": 1}, + headers={"X-Authenticated-Owner": "x" * 513}, + ) + assert oversized.status_code == 400 + assert "exceeds 512 characters" in oversized.json()["detail"] + + +def test_durable_deployment_requires_an_explicit_revision() -> None: + class MissingRevisionWrapper(BasePipelineWrapper): + def setup(self) -> None: + self.pipeline = Pipeline() + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + return Result(value=request.value) + + wrapper = MissingRevisionWrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + with pytest.raises(Exception, match="non-empty durable_revision"): + DurableDeployment("missing-revision", wrapper, InMemoryExecutionStoreProvider()) + + +def test_sync_durable_wrapper_uses_context_sync_controls(monkeypatch) -> None: + app = _durable_app(monkeypatch, "sync-job", SyncWrapper()) + + with TestClient(app) as client: + submitted = client.post("/sync-job/run-durable", json={"value": 4}) + url = submitted.json()["links"]["self"] + inspected = _wait_for_status(client, url, "completed", "sync durable execution did not complete") + + assert inspected.json()["result"] == {"value": 6} + assert inspected.json()["progress"][0]["message"] == "working in a worker thread" + + +async def test_sync_work_retains_claim_after_shutdown_grace_until_thread_exits(monkeypatch) -> None: + started = threading.Event() + release = threading.Event() + + class BlockingWrapper(BasePipelineWrapper): + durable_revision = "blocking-wrapper" + + def setup(self) -> None: + self.pipeline = Pipeline() + + def run_durable(self, context: DurableContext, request: Request) -> Result: + started.set() + assert release.wait(timeout=5) + return Result(value=request.value) + + monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) + provider = InMemoryExecutionStoreProvider() + wrapper = BlockingWrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("blocking", wrapper, provider) + await deployment.start() + _, submitted = await deployment.submit({"value": 9}) + assert await asyncio.to_thread(started.wait, 1) + + await deployment.close() + assert deployment.manager.draining + assert await deployment.store.claim_next("replacement") is None + + release.set() + await deployment.manager.wait_drained() + completed = await deployment.store.get(submitted.execution_id) + assert completed is not None + assert completed.status.value == "completed" + + +async def test_async_pipeline_thread_fallback_retains_cancellation_fence(monkeypatch) -> None: + started = threading.Event() + release = threading.Event() + adapter = HaystackDurableAdapter(Pipeline(), ExecutionKind.PIPELINE) + + def blocking_run(_context, _data, *, checkpoint_at): + assert checkpoint_at == ["component"] + started.set() + assert release.wait(timeout=5) + return {"done": True} + + monkeypatch.setattr(adapter, "run_pipeline", blocking_run) + task = asyncio.create_task( + adapter.run_pipeline_async(object(), {}, checkpoint_at=["component"]), # type: ignore[arg-type] + ) + assert await asyncio.to_thread(started.wait, 1) + + task.cancel() + await asyncio.sleep(0.01) + assert not task.done() + release.set() + assert await task == {"done": True} + + +async def test_pipeline_snapshot_round_trip_skips_completed_components_after_retry(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_retry_base_delay", 0) + monkeypatch.setattr(settings, "durable_retry_max_delay", 0) + provider = InMemoryExecutionStoreProvider() + wrapper = CheckpointPipelineWrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("checkpoint-pipeline", wrapper, provider) + await deployment.start() + try: + _, submitted = await deployment.submit({"value": 3}) + for _ in range(200): + completed = await deployment.store.get(submitted.execution_id) + if completed is not None and completed.terminal: + break + await asyncio.sleep(0.005) + else: + pytest.fail("checkpointed Pipeline did not finish its retry") + finally: + await deployment.close() + + assert completed is not None + assert completed.status.value == "completed" + assert completed.result == {"value": 5} + assert completed.attempt == 2 + assert completed.checkpoint is not None + PipelineSnapshot.from_dict(completed.checkpoint.data["snapshot"]) + assert wrapper.first.calls == 1 + assert wrapper.second.calls == 2 + checkpoint_events = [event for event in completed.progress if event.kind == "checkpoint"] + assert len(checkpoint_events) >= 2 + + +async def test_agent_state_checkpoint_restores_custom_state_and_typed_resume_messages() -> None: + schema = { + "counter": {"type": int}, + "tools": {"type": list}, + "hook_context": {"type": dict}, + } + record = _agent_record("agent-state") + checkpoint_state = State( + schema=schema, + data={ + "counter": 7, + "messages": [ChatMessage.from_user("before restart")], + "tools": ["old tool"], + "hook_context": {"request": "old"}, + }, + ) + checkpoint_context = DurableContext(_Claim(record), adapter=object()) + checkpoint_data = _checkpoint_data(checkpoint_state, checkpoint_context) + checkpoint_payload = checkpoint_data["data"] + assert "tools" not in checkpoint_payload["serialized_data"] + assert "hook_context" not in checkpoint_payload["serialized_data"] + assert "tools" not in checkpoint_payload["serialization_schema"]["properties"] + assert "hook_context" not in checkpoint_payload["serialization_schema"]["properties"] + store = InMemoryExecutionStoreProvider().create_execution_store("agent") + await store.initialize() + assert await store.submit(record) + store.set_definition_revision("revision") + waiting_claim = await store.claim_next("before-resume") + assert waiting_claim is not None + async with waiting_claim: + waiting_claim.record.checkpoint = ExecutionCheckpoint(ExecutionKind.AGENT, checkpoint_data) + await waiting_claim.checkpoint() + waiting_claim.record.status = ExecutionStatus.WAITING + await waiting_claim.suspend() + assert await store.resume( + record.execution_id, + {"messages": [ChatMessage.from_user("after restart").to_dict()]}, + ) + resumed_claim = await store.claim_next("after-resume") + assert resumed_claim is not None + context = DurableContext(resumed_claim, adapter=object()) + + restored_state = State( + schema=schema, + data={"counter": 0, "tools": ["live tool"], "hook_context": {"request": "live"}}, + ) + _restore_agent_state(context, restored_state) + + assert restored_state.data["counter"] == 7 + assert restored_state.data["tools"] == ["live tool"] + assert restored_state.data["hook_context"] == {"request": "live"} + assert [message.text for message in restored_state.data["messages"]] == ["before restart", "after restart"] + assert context.resume_input is None + + +async def test_agent_checkpoint_excludes_custom_tool_deserialization() -> None: + tool = Tool( + name="custom_tool", + description="A tool defined by the deployed wrapper", + parameters={"type": "object", "properties": {"value": {"type": "string"}}}, + function=_checkpoint_test_tool, + ) + record = _agent_record("agent-tool-checkpoint") + checkpoint_context = DurableContext(_Claim(record), adapter=object()) + checkpoint_state = State( + schema={"tools": {"type": list}, "hook_context": {"type": dict}}, + data={"tools": [tool], "hook_context": {"request": "old"}}, + ) + checkpoint_data = _checkpoint_data(checkpoint_state, checkpoint_context) + checkpoint_payload = checkpoint_data["data"] + assert "tools" not in checkpoint_payload["serialized_data"] + assert "tools" not in checkpoint_payload["serialization_schema"]["properties"] + record.checkpoint = ExecutionCheckpoint(ExecutionKind.AGENT, checkpoint_state.to_dict()) + + restored_state = State( + schema={"tools": {"type": list}, "hook_context": {"type": dict}}, + data={"tools": [tool], "hook_context": {"request": "live"}}, + ) + _restore_agent_state(DurableContext(_Claim(record), adapter=object()), restored_state) + + assert restored_state.data["tools"] == [tool] + assert restored_state.data["hook_context"] == {"request": "live"} + + +async def test_agent_checkpoint_and_progress_share_one_store_write() -> None: + record = _agent_record("coalesced-agent-checkpoint") + claim = _Claim(record) + context = DurableContext(claim, adapter=object()) + state = State(schema={"counter": {"type": int}}, data={"counter": 1}) + + await _checkpoint_agent_state(context, state) + + assert claim.checkpoints == 1 + assert record.checkpoint is not None + assert [event.kind for event in record.progress] == ["checkpoint"] + + +def test_agent_tool_exit_uses_only_the_current_tool_result_batch() -> None: + failed_then_retried = State( + schema={}, + data={ + "messages": [ + ChatMessage.from_tool("failed", origin=ToolCall(tool_name="finish", arguments={}), error=True), + ChatMessage.from_assistant("retrying"), + ChatMessage.from_tool("done", origin=ToolCall(tool_name="finish", arguments={})), + ] + }, + ) + current_non_exit_tool = State( + schema={}, + data={ + "messages": [ + ChatMessage.from_tool("done", origin=ToolCall(tool_name="finish", arguments={})), + ChatMessage.from_assistant("continuing"), + ChatMessage.from_tool("working", origin=ToolCall(tool_name="search", arguments={})), + ] + }, + ) + + assert _agent_exits_after_tools(failed_then_retried, ["finish"]) + assert not _agent_exits_after_tools(current_non_exit_tool, ["finish"]) + + +async def test_agent_after_run_saves_a_final_checkpoint() -> None: + record = _agent_record("agent-after-run") + adapter = HaystackDurableAdapter(Agent(chat_generator=FakeChatGenerator(), tools=[]), ExecutionKind.AGENT) + claim = _Claim(record) + context = DurableContext(claim, adapter=adapter) + + with execution_context_scope(context): + result = await adapter.run_agent_async(context, messages=[ChatMessage.from_user("question")]) + + assert result["last_message"].text == "done" + assert record.checkpoint is not None + assert record.checkpoint.data["_hayhooks_agent_checkpoint_phase"] == "after_run" + assert claim.checkpoints == 1 + + +async def test_durable_agent_hooks_ignore_ordinary_sync_and_async_runs() -> None: + agent = Agent(chat_generator=FakeChatGenerator(), tools=[]) + HaystackDurableAdapter(agent, ExecutionKind.AGENT) + hook_counts = {name: len(hooks) for name, hooks in agent.hooks.items()} + HaystackDurableAdapter(agent, ExecutionKind.AGENT) + + sync_result = agent.run(messages=[ChatMessage.from_user("sync")]) + async_result = await agent.run_async(messages=[ChatMessage.from_user("async")]) + + assert {name: len(hooks) for name, hooks in agent.hooks.items()} == hook_counts + assert sync_result["last_message"].text == "done" + assert async_result["last_message"].text == "done" + + +async def test_agent_after_run_checkpoint_recovers_without_another_llm_call() -> None: + record = _agent_record("agent-final-state") + adapter = HaystackDurableAdapter(Agent(chat_generator=FakeChatGenerator(), tools=[]), ExecutionKind.AGENT) + claim = _Claim(record) + context = DurableContext(claim, adapter=adapter) + final_state = State( + schema={"counter": {"type": int}}, + data={ + "counter": 2, + "messages": [ChatMessage.from_user("question"), ChatMessage.from_assistant("already complete")], + }, + ) + + record.checkpoint = ExecutionCheckpoint(ExecutionKind.AGENT, _checkpoint_data(final_state, context, final=True)) + result = await adapter.run_agent_async(context, messages=[ChatMessage.from_user("must not run")]) + + assert result["counter"] == 2 + assert result["last_message"].text == "already complete" + assert claim.checkpoints == 0 + + +async def test_agent_on_exit_continuation_checkpoints_application_state() -> None: + from haystack.hooks import hook + + @component + class FailAfterExit: + @component.output_types(replies=list[ChatMessage]) + def run(self, messages: list[ChatMessage], tools=None) -> dict: + del messages, tools + if hasattr(self, "called"): + msg = "interrupted after continuation" + raise RuntimeError(msg) + self.called = True + return {"replies": [ChatMessage.from_assistant("first exit")]} + + @hook + def continue_once(state: State) -> None: + state.set("marker", "saved") + state.set("continue_run", True) + + record = _agent_record("agent-on-exit") + adapter = HaystackDurableAdapter( + Agent( + chat_generator=FailAfterExit(), + tools=[], + state_schema={"marker": {"type": str}}, + hooks={"on_exit": [continue_once]}, + ), + ExecutionKind.AGENT, + ) + claim = _Claim(record) + context = DurableContext(claim, adapter=adapter) + + with execution_context_scope(context), pytest.raises(RuntimeError, match="interrupted after continuation"): + await adapter.run_agent_async(context, messages=[ChatMessage.from_user("question")]) + + assert record.checkpoint is not None + data = record.checkpoint.data["data"]["serialized_data"] + assert data["marker"] == "saved" + assert data["continue_run"] is True + assert claim.checkpoints == 1 + + +async def test_builtin_agent_leaves_resume_input_for_checkpoint_restoration() -> None: + class RestoringAdapter: + restored_messages: list[str] + + async def run_agent_async(self, context, *, messages, **_kwargs): + state = State( + schema={"messages": {"type": list[ChatMessage]}}, + data={"messages": messages}, + ) + _restore_agent_state(context, state) + self.restored_messages = [message.text for message in state.data["messages"]] + return {"messages": [message.to_dict() for message in state.data["messages"]]} + + wrapper = BuiltinAgentWrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("builtin-agent", wrapper, InMemoryExecutionStoreProvider()) + checkpoint_state = State( + schema={"messages": {"type": list[ChatMessage]}}, + data={"messages": [ChatMessage.from_user("before restart")]}, + ) + record = ExecutionRecord( + execution_id="agent-resume", + execution_kind=ExecutionKind.AGENT, + deployment_name="builtin-agent", + definition_revision=deployment.revision, + validated_input={"messages": [ChatMessage.from_user("initial").to_dict()]}, + checkpoint=ExecutionCheckpoint( + ExecutionKind.AGENT, + _checkpoint_data( + checkpoint_state, + DurableContext( + _Claim( + ExecutionRecord( + execution_id="checkpoint", + execution_kind=ExecutionKind.AGENT, + deployment_name="builtin-agent", + definition_revision=deployment.revision, + validated_input={"messages": []}, + ) + ), + adapter=object(), + ), + ), + ), + application_state={ + "__hayhooks_resume_input": { + "messages": [ChatMessage.from_user("after restart").to_dict()], + } + }, + ) + adapter = RestoringAdapter() + context = DurableContext(_Claim(record), adapter=adapter) + + await deployment._run(context) + + assert adapter.restored_messages == ["before restart", "after restart"] + assert context.resume_input is None + + +def test_durable_agent_uses_native_run_and_public_hooks(monkeypatch) -> None: + app = _durable_app(monkeypatch, "agent", AgentWrapper()) + + with TestClient(app) as client: + submitted = client.post("/agent/run-durable", json={"message": "hello"}) + url = submitted.json()["links"]["self"] + inspected = _wait_for_status(client, url, "completed", "durable Agent did not complete") + + assert inspected.json()["result"]["last_message"]["content"][0]["text"] == "done" + assert inspected.json()["progress"][0]["kind"] == "checkpoint" + + +@pytest.mark.parametrize( + ("module_name", "source", "kind"), + [ + ("durable_execution_example", _DURABLE_EXECUTION_EXAMPLE, ExecutionKind.PIPELINE), + ("durable_a2a_example", _DURABLE_A2A_EXAMPLE, ExecutionKind.AGENT), + ], +) +def test_durable_examples_load(monkeypatch, module_name, source, kind) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "test-key") + module = load_pipeline_module(module_name, source) + try: + wrapper = create_pipeline_wrapper_instance(module) + deployment = DurableDeployment(module_name, wrapper, InMemoryExecutionStoreProvider()) + assert deployment.kind is kind + assert deployment.revision + finally: + unload_pipeline_modules(module_name) + + +async def test_durable_execution_example_completes_retry_approval_and_real_pipeline(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_retry_base_delay", 0) + monkeypatch.setattr(settings, "durable_retry_max_delay", 0) + module_name = "durable_execution_end_to_end_example" + module = load_pipeline_module(module_name, _DURABLE_EXECUTION_EXAMPLE) + deployment = None + try: + wrapper = create_pipeline_wrapper_instance(module) + deployment = DurableDeployment("durable-execution-example", wrapper, InMemoryExecutionStoreProvider()) + await deployment.start() + _, submitted = await deployment.submit( + { + "documents": [{"document_id": "guide", "content": "durable document preparation"}], + "fail_first_attempt": True, + "require_approval": True, + "demo_delay_seconds": 0.01, + } + ) + + for _ in range(200): + waiting = await deployment.get(submitted.execution_id) + if waiting.status is ExecutionStatus.WAITING: + break + await asyncio.sleep(0.005) + else: + pytest.fail("durable execution example did not reach approval") + + assert await deployment.resume(submitted.execution_id, {"approved": True}) + for _ in range(200): + completed = await deployment.get(submitted.execution_id) + if completed.terminal: + break + await asyncio.sleep(0.005) + else: + pytest.fail("durable execution example did not complete") + + assert completed.status is ExecutionStatus.COMPLETED + assert completed.attempt == 3 + assert completed.result["document_count"] == 1 + assert completed.result["chunk_count"] == 1 + assert completed.checkpoint is not None + assert {event.kind for event in completed.progress} >= { + "accepted", + "retry_demo", + "waiting", + "checkpoint", + "demo_delay", + "completed", + } + finally: + if deployment is not None: + await deployment.close() + unload_pipeline_modules(module_name) + + +async def test_durable_a2a_example_tool_replays_its_external_effect_idempotently(monkeypatch, tmp_path) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "test-key") + monkeypatch.setenv("HAYHOOKS_EXAMPLE_INDEX_DB", str(tmp_path / "indexing-effects.sqlite3")) + monkeypatch.setenv("HAYHOOKS_EXAMPLE_TOOL_DELAY_SECONDS", "0") + module_name = "durable_a2a_tool_example" + module = load_pipeline_module(module_name, _DURABLE_A2A_EXAMPLE) + try: + record = _agent_record("a2a-tool-replay", "long-running-agent") + claim = _Claim(record) + context = DurableContext(claim, adapter=object()) + + with execution_context_scope(context): + first = await asyncio.to_thread( + module.prepare_document_for_indexing.invoke, + document_id="guide", + content="one two three", + ) + replay = await asyncio.to_thread( + module.prepare_document_for_indexing.invoke, + document_id="guide", + content="one two three", + ) + + assert json.loads(first)["side_effect_applied"] is True + assert json.loads(replay)["side_effect_applied"] is False + assert claim.checkpoints == 2 + assert [event.kind for event in record.progress] == [ + "side_effect_committed", + "side_effect_committed", + ] + finally: + unload_pipeline_modules(module_name) diff --git a/tests/test_durable_process_recovery.py b/tests/test_durable_process_recovery.py new file mode 100644 index 00000000..9115e03b --- /dev/null +++ b/tests/test_durable_process_recovery.py @@ -0,0 +1,327 @@ +"""One real-process crash/restart check for Redis-backed durable execution.""" + +from __future__ import annotations + +import importlib.metadata +import json +import os +import shutil +import signal +import sqlite3 +import subprocess +import sys +import time +import uuid +from pathlib import Path + +import pytest +import requests +from redis import Redis + +from hayhooks.server.a2a.redis_task_store import RedisTaskStore + +pytestmark = [ + pytest.mark.integration, + pytest.mark.skipif( + not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" + ), +] + +_REDIS_URL_ENV = "HAYHOOKS_TEST_REDIS_URL" +_PROCESS_RECOVERY_ENV = "HAYHOOKS_TEST_PROCESS_RECOVERY" +_FIXTURE_DIR = Path(__file__).parent / "test_files/durable_process_recovery" +_A2A_FIXTURE_DIR = Path(__file__).parent / "test_files/durable_a2a_process_recovery" +_CRASH_AFTER_A2A_SUBMIT_ENV = "HAYHOOKS_TEST_CRASH_AFTER_A2A_SUBMIT" + + +def _start_server( + port: int, + environment: dict[str, str], + factory: str = "hayhooks.cli.base:get_app", +) -> subprocess.Popen[str]: + return subprocess.Popen( # noqa: S603 + [ + sys.executable, + "-m", + "uvicorn", + factory, + "--factory", + "--host", + "127.0.0.1", + "--port", + str(port), + "--log-level", + "warning", + "--no-access-log", + ], + cwd=Path.cwd(), + env=environment, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + ) + + +def _stop_server(server: subprocess.Popen[str]) -> None: + if server.poll() is None: + server.terminate() + try: + server.wait(timeout=3) + except subprocess.TimeoutExpired: + server.kill() + server.wait(timeout=3) + + +def _server_error(server: subprocess.Popen[str]) -> str: + output = server.stdout.read() if server.stdout is not None else "" + return f"durable test server exited with {server.returncode}:\n{output}" + + +def _wait_for_server(server: subprocess.Popen[str], base_url: str) -> None: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + if server.poll() is not None: + pytest.fail(_server_error(server)) + try: + if requests.get(f"{base_url}/status", timeout=0.25).status_code == 200: + return + except requests.RequestException: + pass + time.sleep(0.05) + pytest.fail("durable test server did not become ready") + + +def _wait_for_file(path: Path) -> None: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + if path.exists(): + return + time.sleep(0.05) + pytest.fail("durable test wrapper did not reach its crash window") + + +def _wait_for_completion(execution_url: str) -> dict: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + response = requests.get(execution_url, timeout=0.5) + response.raise_for_status() + execution = response.json() + if execution["status"] == "completed": + return execution + if execution["status"] in {"failed", "canceled"}: + pytest.fail(f"durable execution ended as {execution['status']}: {execution}") + time.sleep(0.05) + pytest.fail("durable execution did not recover to completion") + + +def _cleanup_redis(redis_url: str, prefix: str) -> None: + redis = Redis.from_url(redis_url) + try: + keys = list(redis.scan_iter(match=f"{prefix}:*")) + if keys: + redis.delete(*keys) + finally: + redis.close() + + +def create_a2a_recovery_app(): + """Build the process-test A2A app after loading its pipeline fixture.""" + from hayhooks.durable.runtime import DurableDeployment + from hayhooks.server.a2a.app import create_a2a_app + from hayhooks.server.utils.deploy_utils import deploy_pipelines + + deploy_pipelines() + if os.getenv(_CRASH_AFTER_A2A_SUBMIT_ENV) == "1": + submit = DurableDeployment.submit + + async def submit_then_crash(self, *args, **kwargs): + result = await submit(self, *args, **kwargs) + os.kill(os.getpid(), signal.SIGKILL) + return result + + DurableDeployment.submit = submit_then_crash + return create_a2a_app() + + +def _a2a_rpc(base_url: str, method: str, params: dict, request_id: str) -> dict: + response = requests.post( + f"{base_url}/durable_agent/", + json={"jsonrpc": "2.0", "id": request_id, "method": method, "params": params}, + headers={"A2A-Version": "1.0"}, + timeout=2, + ) + response.raise_for_status() + payload = response.json() + assert "error" not in payload, payload + return payload["result"] + + +def _wait_for_task_state(base_url: str, task_id: str, state: str) -> dict: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + task = _a2a_rpc(base_url, "GetTask", {"id": task_id}, f"get-{task_id}") + if task["status"]["state"] == state: + return task + time.sleep(0.05) + pytest.fail(f"A2A task '{task_id}' did not reach {state}") + + +def _a2a_task(result: dict) -> dict: + return result.get("task", result) + + +def test_redis_durable_execution_survives_a_process_kill_and_restart(tmp_path: Path, unused_tcp_port: int) -> None: + if os.getenv(_PROCESS_RECOVERY_ENV) != "1": + pytest.skip(f"set {_PROCESS_RECOVERY_ENV}=1 to run the process-recovery smoke test") + redis_url = os.getenv(_REDIS_URL_ENV) + if not redis_url: + pytest.skip(f"set {_REDIS_URL_ENV} to run the process-recovery smoke test") + + pipelines_dir = tmp_path / "pipelines" + shutil.copytree(_FIXTURE_DIR, pipelines_dir / "recovery_job") + database_path = tmp_path / "effects.sqlite3" + ready_file = tmp_path / "ready" + prefix = f"hayhooks:test:process-recovery:{uuid.uuid4().hex}" + environment = os.environ | { + "HAYHOOKS_PIPELINES_DIR": str(pipelines_dir), + "HAYHOOKS_DURABLE_STORE": "redis", + "HAYHOOKS_DURABLE_REDIS_URL": redis_url, + "HAYHOOKS_DURABLE_REDIS_KEY_PREFIX": prefix, + "HAYHOOKS_DURABLE_LEASE_DURATION_MS": "250", + "HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS": "25", + "HAYHOOKS_DURABLE_MAX_ATTEMPTS": "2", + } + base_url = f"http://127.0.0.1:{unused_tcp_port}" + request_body = {"database_path": str(database_path), "ready_file": str(ready_file)} + headers = {"Idempotency-Key": "process-recovery"} + servers: list[subprocess.Popen[str]] = [] + + try: + server = _start_server(unused_tcp_port, environment) + servers.append(server) + _wait_for_server(server, base_url) + + submitted = requests.post(f"{base_url}/recovery_job/run-durable", json=request_body, headers=headers, timeout=1) + assert submitted.status_code == 202, submitted.text + execution = submitted.json() + execution_url = f"{base_url}{execution['links']['self']}" + _wait_for_file(ready_file) + + server.kill() + server.wait(timeout=3) + + server = _start_server(unused_tcp_port, environment) + servers.append(server) + _wait_for_server(server, base_url) + completed = _wait_for_completion(execution_url) + + assert completed["attempt"] == 2 + assert completed["result"] == {"attempt": 2, "effect_applied": False} + with sqlite3.connect(database_path) as connection: + assert connection.execute("SELECT COUNT(*) FROM checkpoint_runs").fetchone() == (1,) + assert connection.execute("SELECT COUNT(*) FROM effects").fetchone() == (1,) + + replay = requests.post(f"{base_url}/recovery_job/run-durable", json=request_body, headers=headers, timeout=1) + assert replay.status_code == 200, replay.text + assert replay.headers["Idempotent-Replay"] == "true" + assert replay.json()["execution_id"] == execution["execution_id"] + finally: + for server in servers: + _stop_server(server) + _cleanup_redis(redis_url, prefix) + + +def test_redis_durable_a2a_tasks_read_through_after_restart(tmp_path: Path, unused_tcp_port: int) -> None: + if os.getenv(_PROCESS_RECOVERY_ENV) != "1": + pytest.skip(f"set {_PROCESS_RECOVERY_ENV}=1 to run the process-recovery smoke test") + redis_url = os.getenv(_REDIS_URL_ENV) + if not redis_url: + pytest.skip(f"set {_REDIS_URL_ENV} to run the process-recovery smoke test") + + pipelines_dir = tmp_path / "pipelines" + shutil.copytree(_A2A_FIXTURE_DIR, pipelines_dir / "durable_agent") + prefix = f"hayhooks:test:a2a-process-recovery:{uuid.uuid4().hex}" + environment = os.environ | { + "HAYHOOKS_PIPELINES_DIR": str(pipelines_dir), + "HAYHOOKS_A2A_EXTERNAL_URL": f"http://127.0.0.1:{unused_tcp_port}", + "HAYHOOKS_A2A_TASK_STORE": "redis", + "HAYHOOKS_A2A_REDIS_URL": redis_url, + "HAYHOOKS_A2A_REDIS_KEY_PREFIX": f"{prefix}:a2a", + "HAYHOOKS_DURABLE_STORE": "redis", + "HAYHOOKS_DURABLE_REDIS_URL": redis_url, + "HAYHOOKS_DURABLE_REDIS_KEY_PREFIX": f"{prefix}:durable", + "HAYHOOKS_DURABLE_POLL_INTERVAL": "0.05", + "HAYHOOKS_DURABLE_LEASE_DURATION_MS": "250", + "HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS": "25", + } + base_url = f"http://127.0.0.1:{unused_tcp_port}" + factory = "tests.test_durable_process_recovery:create_a2a_recovery_app" + servers: list[subprocess.Popen[str]] = [] + + def submit(text: str) -> str: + message = {"messageId": f"message-{text}", "role": "ROLE_USER", "parts": [{"text": text}]} + result = _a2a_rpc( + base_url, + "SendMessage", + { + "message": message, + "configuration": {"returnImmediately": True}, + }, + f"send-{text}", + ) + return _a2a_task(result)["id"] + + try: + server = _start_server( + unused_tcp_port, + environment | {_CRASH_AFTER_A2A_SUBMIT_ENV: "1"}, + factory, + ) + servers.append(server) + _wait_for_server(server, base_url) + with pytest.raises(requests.RequestException): + submit("resume") + server.wait(timeout=3) + redis = Redis.from_url(redis_url) + try: + task_store = RedisTaskStore(redis, "durable_agent", key_prefix=f"{prefix}:a2a") + task_ids = redis.zrange(task_store._key("active"), 0, -1) + finally: + redis.close() + assert len(task_ids) == 1 + resumable_id = task_ids[0].decode() if isinstance(task_ids[0], bytes) else task_ids[0] + + server = _start_server(unused_tcp_port, environment, factory) + servers.append(server) + _wait_for_server(server, base_url) + + assert _wait_for_task_state(base_url, resumable_id, "TASK_STATE_INPUT_REQUIRED") + cancelable_id = submit("cancel") + _wait_for_task_state(base_url, cancelable_id, "TASK_STATE_INPUT_REQUIRED") + listed = _a2a_rpc(base_url, "ListTasks", {}, "list") + assert {task["id"] for task in listed["tasks"]} == {resumable_id, cancelable_id} + + resumed = _a2a_rpc( + base_url, + "SendMessage", + { + "message": { + "messageId": "message-approved", + "taskId": resumable_id, + "role": "ROLE_USER", + "parts": [{"text": "approved"}], + } + }, + "resume", + ) + if _a2a_task(resumed)["status"]["state"] != "TASK_STATE_COMPLETED": + _stop_server(server) + pytest.fail(f"{json.dumps(resumed)}\n{_server_error(server)}") + + canceled = _a2a_rpc(base_url, "CancelTask", {"id": cancelable_id}, "cancel") + if _a2a_task(canceled)["status"]["state"] != "TASK_STATE_CANCELED": + _wait_for_task_state(base_url, cancelable_id, "TASK_STATE_CANCELED") + finally: + for server in servers: + _stop_server(server) + _cleanup_redis(redis_url, prefix) diff --git a/tests/test_files/durable_process_recovery/pipeline_wrapper.py b/tests/test_files/durable_process_recovery/pipeline_wrapper.py new file mode 100644 index 00000000..003dc600 --- /dev/null +++ b/tests/test_files/durable_process_recovery/pipeline_wrapper.py @@ -0,0 +1,74 @@ +"""Deterministic durable wrapper used by the process-recovery smoke test.""" + +import sqlite3 +import time +from pathlib import Path + +from haystack import Pipeline, component +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper, DurableContext, current_execution_id + + +class RecoveryRequest(BaseModel): + database_path: str + ready_file: str + + +class RecoveryResult(BaseModel): + attempt: int + effect_applied: bool + + +@component +class Checkpoint: + @component.output_types(value=str) + def run(self, database_path: str) -> dict[str, str]: + execution_id = current_execution_id() + if execution_id is None: + msg = "checkpoint component requires a durable execution" + raise RuntimeError(msg) + with sqlite3.connect(database_path) as connection: + connection.execute("CREATE TABLE IF NOT EXISTS checkpoint_runs (execution_id TEXT PRIMARY KEY)") + connection.execute("INSERT INTO checkpoint_runs VALUES (?)", (execution_id,)) + return {"value": "checkpointed"} + + +@component +class Effect: + @component.output_types(effect_applied=bool) + def run(self, value: str, database_path: str, ready_file: str) -> dict[str, bool]: + del value + execution_id = current_execution_id() + if execution_id is None: + msg = "effect component requires a durable execution" + raise RuntimeError(msg) + with sqlite3.connect(database_path) as connection: + connection.execute("CREATE TABLE IF NOT EXISTS effects (execution_id TEXT PRIMARY KEY)") + effect_applied = ( + connection.execute("INSERT OR IGNORE INTO effects VALUES (?)", (execution_id,)).rowcount == 1 + ) + if effect_applied: + Path(ready_file).write_text("ready", encoding="utf-8") + time.sleep(60) + return {"effect_applied": effect_applied} + + +class PipelineWrapper(BasePipelineWrapper): + durable_revision = "process-recovery" + + def setup(self) -> None: + self.pipeline = Pipeline() + self.pipeline.add_component("checkpoint", Checkpoint()) + self.pipeline.add_component("effect", Effect()) + self.pipeline.connect("checkpoint.value", "effect.value") + + async def run_durable_async(self, context: DurableContext, request: RecoveryRequest) -> RecoveryResult: + outputs = await context.run_pipeline_async( + { + "checkpoint": {"database_path": request.database_path}, + "effect": {"database_path": request.database_path, "ready_file": request.ready_file}, + }, + checkpoint_at=["checkpoint", "effect"], + ) + return RecoveryResult(attempt=context.attempt, effect_applied=outputs["effect"]["effect_applied"]) diff --git a/tests/test_redis_execution_integration.py b/tests/test_redis_execution_integration.py new file mode 100644 index 00000000..d6bf027e --- /dev/null +++ b/tests/test_redis_execution_integration.py @@ -0,0 +1,235 @@ +"""Real-Redis contract tests for the isolated durable namespace.""" + +from __future__ import annotations + +import asyncio +from dataclasses import replace + +import pytest + +from hayhooks.durable.backend import ExecutionStoreConfig +from hayhooks.durable.engine import ( + Checkpoint, + Claim, + Complete, + ExecutionLeaseLostError, + ExecutionPayloadSizeError, + RecoverExpiredLease, + RequestCancellation, + ScheduleRetry, +) +from hayhooks.durable.manager import DurableExecutionManager +from hayhooks.durable.models import ( + ExecutionAdmissionError, + ExecutionCheckpoint, + ExecutionKind, + ExecutionRecord, + ExecutionStatus, +) +from hayhooks.durable.redis import RedisExecutionStore +from hayhooks.durable.store import ExecutionStore +from tests.durable_contract import assert_store_contract, control + +pytestmark = pytest.mark.integration + + +@pytest.fixture +async def store(isolated_redis): + redis, prefix = isolated_redis + config = ExecutionStoreConfig( + key_prefix=f"{prefix}:durable", + max_input_bytes=64, + max_checkpoint_bytes=64, + max_result_bytes=64, + max_error_bytes=64, + max_wait_bytes=64, + max_progress_events=2, + max_progress_event_bytes=32, + terminal_ttl_seconds=60, + ) + durable = RedisExecutionStore(redis, deployment="integration", config=config) + await durable.initialize() + yield redis, durable + + +async def _claim(durable: RedisExecutionStore, *, worker: str = "worker", lease_ms: int = 1_000): + run_id = await durable.read_candidate() + assert run_id is not None + return run_id, await durable.transition(run_id, Claim(worker, 0, lease_ms, 3, "rev-1"), candidate=True) + + +async def test_redis_store_matches_shared_contract(store) -> None: + _, durable = store + await assert_store_contract(durable) + + +async def test_reading_candidate_is_non_destructive(store) -> None: + _, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + assert await durable.read_candidate() == "run_1" + assert await durable.read_candidate() == "run_1" + + +async def test_three_concurrent_claimers_create_one_live_owner(store) -> None: + redis, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id = await durable.read_candidate() + assert run_id is not None + await asyncio.gather( + *( + durable.transition(run_id, Claim(f"worker-{index}", 0, 10_000, 3, "rev-1"), candidate=True) + for index in range(3) + ) + ) + current = await durable.get(run_id) + assert current is not None and current.status is ExecutionStatus.RUNNING + assert current.fence == current.run_attempt == 1 + assert await redis.zcard(durable.keys.lease_expiry) == 1 + assert await redis.zcard(durable.keys.runnable) == 0 + + +async def test_delayed_work_is_invisible_until_the_redis_deadline(store) -> None: + redis, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id, claimed = await _claim(durable) + await durable.transition(run_id, ScheduleRetry(claimed.next_control.fence, "worker", 0, 250, 2, b"retry")) + seconds, micros = await redis.time() + now_ms = int(seconds) * 1_000 + int(micros) // 1_000 + score = await redis.zscore(durable.keys.runnable, run_id) + assert score is not None and score > now_ms + assert await durable.read_candidate() is None + await asyncio.sleep(0.3) + assert await durable.read_candidate() == run_id + + +async def test_hundred_concurrent_submissions_succeed_when_admission_is_disabled(store) -> None: + _, durable = store + + async def submit(index: int): + return await durable.submit( + control(f"run_{index}", idempotency_digest=f"{index:064x}", binding_digest="b" * 64), + b"{}", + binding_digest="b" * 64, + ) + + results = await asyncio.gather(*(submit(index) for index in range(100))) + assert all(result.created for result in results) + assert (await durable.operational_counts())["nonterminal"] == 100 + + +async def test_global_admission_allows_replay_and_releases_capacity(store) -> None: + redis, durable = store + limited = RedisExecutionStore( + redis, + deployment="integration", + config=replace( + durable.config, + key_prefix=f"{durable.config.key_prefix}:limited", + max_nonterminal_executions=1, + ), + ) + await limited.initialize() + assert (await limited.submit(control(), b"{}", binding_digest="b" * 64)).created + replay = await limited.submit(control("replay"), b"{}", binding_digest="b" * 64) + assert not replay.created and replay.control.run_id == "run_1" + + second = control("run_2", idempotency_digest="c" * 64, binding_digest="d" * 64) + with pytest.raises(ExecutionAdmissionError): + await limited.submit(second, b"{}", binding_digest="d" * 64) + await limited.transition("run_1", RequestCancellation(0, "done")) + assert (await limited.submit(second, b"{}", binding_digest="d" * 64)).created + + +async def test_expired_lease_reenters_runnable(store) -> None: + _, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id, claimed = await _claim(durable, lease_ms=20) + await asyncio.sleep(0.03) + await durable.maintain(lambda fence, deadline: RecoverExpiredLease(0, fence, deadline, 3, "rev-1")) + recovered = await durable.get(run_id) + assert recovered is not None and recovered.status is ExecutionStatus.QUEUED + assert await durable.read_candidate() == run_id + assert claimed.next_control.fence == 1 + + +async def test_terminal_ttl_retains_data_and_rejects_stale_fences(store) -> None: + redis, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id, _ = await _claim(durable) + checkpointed = await durable.transition(run_id, Checkpoint(1, "worker", 0, 1_000, b"checkpoint")) + with pytest.raises(ExecutionLeaseLostError): + await durable.transition(run_id, Complete(0, "worker", 0, b"result")) + terminal = await durable.transition(run_id, Complete(checkpointed.next_control.fence, "worker", 0, b"result")) + assert terminal.next_control.terminal + async with redis.pipeline(transaction=False) as pipe: + for key in durable._execution_keys(run_id): + pipe.pttl(key) + ttls = await pipe.execute() + assert all(ttl == -2 or ttl > 0 for ttl in ttls) + assert await redis.pttl(durable.keys.idempotency("a" * 64)) > 0 + + +async def test_redis_rejects_oversized_payload_without_partial_write(store) -> None: + _, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id, claim = await _claim(durable) + with pytest.raises(ExecutionPayloadSizeError): + await durable.transition(run_id, Checkpoint(claim.next_control.fence, "worker", 0, 1_000, b"x" * 65)) + assert await durable.get(run_id) == claim.next_control + + +async def test_redis_adapter_runs_the_public_manager_contract(store) -> None: + redis, durable = store + config = replace( + durable.config, + key_prefix=f"{durable.config.key_prefix}:adapter", + max_input_bytes=512, + max_checkpoint_bytes=512, + max_result_bytes=512, + max_error_bytes=512, + max_wait_bytes=512, + max_progress_event_bytes=256, + ) + adapter_store = ExecutionStore( + RedisExecutionStore(redis, deployment="adapter", config=config), + definition_revision="rev-1", + lease_duration_ms=10_000, + max_run_attempts=3, + max_progress_events=2, + max_record_bytes=512, + ) + + async def runner(context): + context.record.application_state["phase"] = "running" + context.record.checkpoint = ExecutionCheckpoint(ExecutionKind.PIPELINE, {"component": "search"}) + await context.checkpoint() + await context.report_progress("started") + return {"answer": "done"} + + manager = DurableExecutionManager("adapter", adapter_store, runner, adapter=object(), poll_interval=0.001) + await manager.start() + try: + assert await adapter_store.submit( + ExecutionRecord( + execution_id="public-run", + execution_kind=ExecutionKind.PIPELINE, + deployment_name="adapter", + definition_revision="rev-1", + validated_input={"question": "hello"}, + operation_fingerprint="request-fingerprint", + max_progress_events=2, + max_record_bytes=512, + ) + ) + for _ in range(100): + record = await adapter_store.get("public-run") + if record is not None and record.terminal: + break + await asyncio.sleep(0.01) + else: + pytest.fail("Redis adapter did not complete the public manager execution") + finally: + await manager.close() + + assert record.status is ExecutionStatus.COMPLETED + assert record.result == {"answer": "done"} From 5b4884213270b6814553a81d36d65899d31619f6 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:09:05 +0200 Subject: [PATCH 04/28] feat(durable): expose portable managed runtime --- src/hayhooks/__init__.py | 16 +- src/hayhooks/durable/__init__.py | 93 +++ src/hayhooks/durable/runtime.py | 577 ++++++++++++++++++ .../server/utils/base_pipeline_wrapper.py | 25 + tests/test_pipeline_utils.py | 7 + 5 files changed, 716 insertions(+), 2 deletions(-) create mode 100644 src/hayhooks/durable/__init__.py create mode 100644 src/hayhooks/durable/runtime.py diff --git a/src/hayhooks/__init__.py b/src/hayhooks/__init__.py index fcc93a3a..0f876d8a 100644 --- a/src/hayhooks/__init__.py +++ b/src/hayhooks/__init__.py @@ -1,25 +1,35 @@ +"""Public Hayhooks authoring API.""" + +from hayhooks.a2a import A2APipelineWrapper from hayhooks.callbacks import default_on_pipeline_end, default_on_tool_call_end, default_on_tool_call_start +from hayhooks.durable import ExecutionProgress, ExecutionResult, current_durable_context, current_execution_id +from hayhooks.durable.context import DurableContext +from hayhooks.durable.models import ExecutionStatus from hayhooks.events import PipelineEvent from hayhooks.server.app import create_app, run_app from hayhooks.server.logger import log from hayhooks.server.pipelines.sse import SSEStream +from hayhooks.server.pipelines.streaming import async_streaming_generator, streaming_generator from hayhooks.server.pipelines.utils import ( - async_streaming_generator, chat_messages_from_openai_response, coerce_pipeline_inputs, get_input_files, get_last_user_input_text, get_last_user_message, is_user_message, - streaming_generator, ) from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper from hayhooks.server.utils.haystack_compat import AsyncPipeline, Pipeline from hayhooks.server.utils.yaml_pipeline_wrapper import YAMLPipelineWrapper __all__ = [ + "A2APipelineWrapper", "AsyncPipeline", "BasePipelineWrapper", + "DurableContext", + "ExecutionProgress", + "ExecutionResult", + "ExecutionStatus", "Pipeline", "PipelineEvent", "SSEStream", @@ -28,6 +38,8 @@ "chat_messages_from_openai_response", "coerce_pipeline_inputs", "create_app", + "current_durable_context", + "current_execution_id", "default_on_pipeline_end", "default_on_tool_call_end", "default_on_tool_call_start", diff --git a/src/hayhooks/durable/__init__.py b/src/hayhooks/durable/__init__.py new file mode 100644 index 00000000..4b2959e0 --- /dev/null +++ b/src/hayhooks/durable/__init__.py @@ -0,0 +1,93 @@ +"""Advanced durable execution contracts and safe public result models.""" + +from __future__ import annotations + +from datetime import datetime +from typing import TYPE_CHECKING, Any + +from pydantic import BaseModel, Field + +from hayhooks.durable.context import get_current_durable_context +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode +from hayhooks.durable.models import ExecutionStatus + +if TYPE_CHECKING: + from hayhooks.durable.runtime import DurableRuntime, ExecutionStoreProvider + from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider + + +class ExecutionProgress(BaseModel): + """Sanitized client-visible progress event.""" + + sequence: int + kind: str + message: str + timestamp: datetime + metadata: dict[str, Any] = Field(default_factory=dict) + + +class ExecutionResult(BaseModel): + """Safe durable REST/A2A execution projection.""" + + execution_id: str + status: ExecutionStatus + attempt: int + sequence: int + progress: list[ExecutionProgress] + result: Any | None = None + error: dict[str, Any] | None = None + waiting: dict[str, Any] | None = None + cancellation_requested_at: datetime | None = None + created_at: datetime + updated_at: datetime + links: dict[str, str] = Field(default_factory=dict) + + +def current_execution_id() -> str | None: + """Return the active durable execution ID for hooks and idempotent tools.""" + context = get_current_durable_context() + return context.execution_id if context is not None else None + + +def current_durable_context() -> Any | None: + """Return the active context for advanced hooks and tools.""" + return get_current_durable_context() + + +def __getattr__(name: str) -> Any: + """Lazily expose durable infrastructure without eager optional imports.""" + if name in {"DurableRuntime", "ExecutionStoreProvider", "durable_runtime"}: + from hayhooks.durable.runtime import DurableRuntime, ExecutionStoreProvider, durable_runtime + + return { + "DurableRuntime": DurableRuntime, + "ExecutionStoreProvider": ExecutionStoreProvider, + "durable_runtime": durable_runtime, + }[name] + if name in {"ExecutionStore", "InMemoryExecutionStoreProvider", "RedisExecutionStoreProvider"}: + from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider + + return { + "ExecutionStore": ExecutionStore, + "InMemoryExecutionStoreProvider": InMemoryExecutionStoreProvider, + "RedisExecutionStoreProvider": RedisExecutionStoreProvider, + }[name] + msg = f"module {__name__!r} has no attribute {name!r}" + raise AttributeError(msg) + + +__all__ = [ + "DurableAuthoringMode", + "DurableRuntime", + "ExecutionProgress", + "ExecutionResult", + "ExecutionStatus", + "ExecutionStore", + "ExecutionStoreProvider", + "InMemoryExecutionStoreProvider", + "RedisExecutionStoreProvider", + "current_durable_context", + "current_execution_id", + "durable_authoring_mode", + "durable_runtime", +] diff --git a/src/hayhooks/durable/runtime.py b/src/hayhooks/durable/runtime.py new file mode 100644 index 00000000..ecab856c --- /dev/null +++ b/src/hayhooks/durable/runtime.py @@ -0,0 +1,577 @@ +"""Runtime-owned durable deployment services for wrappers and A2A projections.""" + +from __future__ import annotations + +import asyncio +import hashlib +import inspect +import json +import uuid +from collections.abc import Awaitable, Callable, Mapping +from typing import Any, Protocol, cast, get_type_hints + +from pydantic import BaseModel, Field, TypeAdapter, ValidationError + +from hayhooks.durable.adapters import HaystackDurableAdapter, _run_fenced_thread, execution_kind +from hayhooks.durable.backend import ExecutionIdempotencyConflictError +from hayhooks.durable.context import DurableContext +from hayhooks.durable.manager import DurableExecutionManager +from hayhooks.durable.mode import DurableAuthoringMode, _durable_method_implementations, durable_authoring_mode +from hayhooks.durable.models import ExecutionKind, ExecutionRecord, JsonValue, json_safe, validate_json +from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider +from hayhooks.server.exceptions import PipelineWrapperError +from hayhooks.server.logger import log +from hayhooks.server.pipelines.registry import registry +from hayhooks.server.tracing import SPAN_DURABLE_ATTEMPT, SPAN_DURABLE_SUBMIT, build_trace_tags, trace_operation +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.settings import AppSettings, settings + + +class ExecutionStoreProvider(Protocol): + """Create application-owned stores for durable deployments.""" + + def create_execution_store(self, deployment_name: str) -> ExecutionStore: ... + + async def close(self) -> None: ... + + +def _runtime_settings(provider: ExecutionStoreProvider | None, app_settings: AppSettings | None) -> AppSettings | None: + provider_settings = getattr(provider, "app_settings", None) + if isinstance(provider_settings, AppSettings): + if app_settings is not None and app_settings != provider_settings: + msg = "Durable runtime and execution-store provider settings must match" + raise ValueError(msg) + app_settings = provider_settings + return app_settings.model_copy(deep=True) if app_settings is not None else None + + +class DurableDeployment: + """One deployment's records, manager, validated callable, and adapter.""" + + def __init__( + self, + name: str, + wrapper: BasePipelineWrapper, + provider: ExecutionStoreProvider, + *, + app_settings: AppSettings | None = None, + ) -> None: + self.name = name + self.wrapper = wrapper + self.app_settings = _runtime_settings(provider, app_settings) or settings.model_copy(deep=True) + pipeline = wrapper.pipeline + try: + kind = execution_kind(pipeline) + except TypeError as error: + raise PipelineWrapperError(str(error)) from error + self.authoring_mode = durable_authoring_mode(wrapper) + self.builtin_agent = kind is ExecutionKind.AGENT and self.authoring_mode is DurableAuthoringMode.MANAGED_AGENT + if self.builtin_agent: + self.method = None + self.is_async = True + self.request_type = _DurableAgentRequest + self.result_type = Any + else: + self.method, self.is_async, self.request_type, self.result_type = _durable_method_contract(wrapper) + self.result_adapter = TypeAdapter(self.result_type) if self.result_type is not Any else None + self.resume_type = getattr(wrapper, "durable_resume_model", None) + if self.resume_type is not None and ( + not inspect.isclass(self.resume_type) or not issubclass(self.resume_type, BaseModel) + ): + msg = "durable_resume_model must be a Pydantic model class or None" + raise PipelineWrapperError(msg) + revision = getattr(wrapper, "durable_revision", None) + if not isinstance(revision, str) or not revision.strip(): + msg = "Durable wrappers must declare a non-empty durable_revision" + raise PipelineWrapperError(msg) + self.kind = kind + self.revision = revision.strip() + self.adapter = HaystackDurableAdapter(pipeline, kind) + self.store = provider.create_execution_store(name) + self.store.set_definition_revision(self.revision) + self.manager = DurableExecutionManager( + name, + self.store, + self._run, + self.adapter, + concurrency=self.app_settings.durable_execution_concurrency, + poll_interval=self.app_settings.durable_poll_interval, + shutdown_grace_period=self.app_settings.durable_shutdown_grace_period, + max_attempts=self.app_settings.durable_max_attempts, + retry_base_delay=self.app_settings.durable_retry_base_delay, + retry_max_delay=self.app_settings.durable_retry_max_delay, + ) + + async def start(self) -> None: + """Prepare and activate this deployment's execution manager.""" + await self.manager.start() + + async def prepare(self) -> None: + """Initialize the execution store without allowing the candidate to claim work.""" + await self.manager.prepare() + + def activate(self) -> None: + """Start prepared workers at the deployment publication boundary.""" + self.manager.activate() + + def deactivate(self) -> None: + """Reject new submissions while an active deployment is being replaced.""" + self.manager.deactivate() + + async def quiesce(self) -> None: + """Close submission admission before this deployment is stopped.""" + await self.manager.quiesce() + + async def close(self) -> None: + """Stop this deployment's workers and submission admission.""" + await self.manager.close() + + async def submit( + self, + payload: Mapping[str, Any], + *, + execution_id: str | None = None, + owner_id: str | None = None, + ) -> tuple[bool, ExecutionRecord]: + """Validate and idempotently submit one durable execution.""" + if not self.manager.accepting: + msg = f"Durable deployment '{self.name}' is not accepting submissions" + raise RuntimeError(msg) + request = self.request_type.model_validate(dict(payload)) + if execution_id is not None and owner_id is not None: + execution_id = execution_id_for(owner_id, execution_id) + execution_id = execution_id or uuid.uuid4().hex + validated_input = cast( + dict[str, JsonValue], + validate_json( + request.model_dump(mode="json"), limit=self.app_settings.durable_max_record_bytes, label="request" + ), + ) + fingerprint_input = cast(dict[str, JsonValue], _canonical_json(request.model_dump(mode="python"))) + operation_fingerprint = _operation_fingerprint( + self.name, + self.revision, + fingerprint_input, + owner_id=owner_id, + ) + record = ExecutionRecord( + execution_id=execution_id, + execution_kind=self.kind, + deployment_name=self.name, + definition_revision=self.revision, + validated_input=validated_input, + operation_fingerprint=operation_fingerprint, + owner_id=owner_id, + max_progress_events=self.app_settings.durable_max_progress_events, + max_record_bytes=self.app_settings.durable_max_record_bytes, + ) + with trace_operation( + SPAN_DURABLE_SUBMIT, + tags=build_trace_tags( + { + "hayhooks.pipeline.name": self.name, + "hayhooks.durable.execution_id": execution_id, + "hayhooks.durable.definition_revision": self.revision, + } + ), + ) as span: + try: + created, persisted = await self.manager.submit_with_record(record) + except ExecutionIdempotencyConflictError as error: + msg = "Idempotency-Key was already used for a different durable operation" + raise IdempotencyConflictError(msg) from error + persisted = self._validated_record( + execution_id, + persisted, + owner_id=owner_id, + enforce_owner=owner_id is not None, + ) + if not created and persisted.operation_fingerprint != operation_fingerprint: + msg = "Idempotency-Key was already used for a different durable operation" + raise IdempotencyConflictError(msg) + span.set_tag("hayhooks.durable.idempotent_replay", not created) + log.bind( + deployment=self.name, + execution_id=persisted.execution_id, + revision=self.revision, + kind=self.kind.value, + created=created, + ).debug("Accepted durable execution submission") + return created, persisted + + async def get( + self, + execution_id: str, + *, + owner_id: str | None = None, + enforce_owner: bool = False, + allow_revision_mismatch: bool = False, + ) -> ExecutionRecord: + """Return one execution after owner and definition-revision validation.""" + record = await self.store.get(execution_id) + return self._validated_record( + execution_id, + record, + owner_id=owner_id, + enforce_owner=enforce_owner, + allow_revision_mismatch=allow_revision_mismatch, + ) + + async def request_cancel( + self, + execution_id: str, + *, + owner_id: str | None = None, + enforce_owner: bool = False, + reason: str | None = None, + ) -> bool: + """Request cooperative cancellation after validating record ownership.""" + await self.get( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + allow_revision_mismatch=True, + ) + accepted = await self.store.request_cancel(execution_id, reason) + log.bind(deployment=self.name, execution_id=execution_id, accepted=accepted).debug( + "Processed durable execution cancellation request" + ) + return accepted + + async def resume( + self, + execution_id: str, + update: JsonValue | None = None, + *, + owner_id: str | None = None, + enforce_owner: bool = False, + ) -> bool: + """Validate and enqueue a resume update for a waiting execution.""" + await self.get( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + allow_revision_mismatch=True, + ) + if self.resume_type is not None and update is None: + msg = f"Execution '{execution_id}' requires a resume request body" + raise ValueError(msg) + if self.resume_type is not None: + update = cast(JsonValue, self.resume_type.model_validate(update).model_dump(mode="json")) + resumed = await self.store.resume(execution_id, update) + log.bind(deployment=self.name, execution_id=execution_id, resumed=resumed).debug( + "Processed durable execution resume request" + ) + return resumed + + def _validated_record( + self, + execution_id: str, + record: ExecutionRecord | None, + *, + owner_id: str | None = None, + enforce_owner: bool = False, + allow_revision_mismatch: bool = False, + ) -> ExecutionRecord: + if record is None or record.deployment_name != self.name: + raise KeyError(execution_id) + if enforce_owner and record.owner_id != owner_id: + raise KeyError(execution_id) + if record.definition_revision != self.revision and not record.terminal and not allow_revision_mismatch: + msg = ( + f"Durable execution '{execution_id}' was created for a different definition revision and cannot resume." + ) + raise DefinitionRevisionConflictError(msg) + return record + + async def _run(self, context: DurableContext) -> JsonValue: + if context.record.definition_revision != self.revision: + msg = ( + f"Durable execution '{context.execution_id}' was created for a different " + "definition revision and cannot resume." + ) + raise DefinitionRevisionConflictError(msg) + with trace_operation( + SPAN_DURABLE_ATTEMPT, + tags=build_trace_tags( + { + "hayhooks.pipeline.name": self.name, + "hayhooks.durable.execution_id": context.execution_id, + "hayhooks.durable.attempt": context.attempt, + "hayhooks.durable.kind": self.kind.value, + "hayhooks.durable.queue_latency_ms": max( + 0, + int((context.record.updated_at - context.record.created_at).total_seconds() * 1_000), + ), + } + ), + ): + if self.builtin_agent: + request = _DurableAgentRequest.model_validate(context.record.validated_input) + from haystack.dataclasses import ChatMessage + + messages = [ChatMessage.from_dict(message) for message in request.messages] + resume_input = context.take_resume_input() if context.record.checkpoint is None else None + if isinstance(resume_input, dict): + resumed_messages = resume_input.get("messages") + if isinstance(resumed_messages, list): + messages.extend( + ChatMessage.from_dict(cast(dict[str, Any], message)) + for message in resumed_messages + if isinstance(message, dict) + ) + return json_safe(await context.run_agent_async(messages=messages)) + method = cast(Callable[[DurableContext, BaseModel], Any], self.method) + request = self.request_type.model_validate(context.record.validated_input) + if self.is_async: + result = await cast(Awaitable[Any], method(context, request)) + else: + result = await _run_fenced_thread(method, context, request) + if self.result_adapter is not None: + try: + result = self.result_adapter.validate_python(result) + except ValidationError as error: + msg = "Durable method result does not match its declared return annotation" + raise ValueError(msg) from error + serializer = getattr(result, "model_dump", None) + return json_safe(serializer(mode="json") if callable(serializer) else result) + + +class _DurableAgentRequest(BaseModel): + """Private A2A input mapping; REST wrappers always provide their own model.""" + + messages: list[dict[str, Any]] = Field(min_length=1) + + +class DurableRuntime: + """Application-owned provider lifecycle and deployed durable services.""" + + def __init__( + self, + provider: ExecutionStoreProvider | None = None, + *, + app_settings: AppSettings | None = None, + ) -> None: + self.provider = provider + self._app_settings = _runtime_settings(provider, app_settings) + self._deployments: dict[str, DurableDeployment] = {} + self._started = False + self._provider_close_task: asyncio.Task[None] | None = None + self._registry: Any | None = None + + def has_capability(self, wrapper: BasePipelineWrapper) -> bool: + return durable_authoring_mode(wrapper) is not DurableAuthoringMode.NONE + + @property + def started(self) -> bool: + return self._started + + @property + def app_settings(self) -> AppSettings: + """Return configured or provider settings, falling back to Hayhooks' global settings.""" + if self._app_settings is not None: + return self._app_settings + provider_settings = getattr(self.provider, "app_settings", None) + return provider_settings if isinstance(provider_settings, AppSettings) else settings + + def create_deployment(self, name: str, wrapper: BasePipelineWrapper) -> DurableDeployment | None: + """Build an uncached candidate so route closures cannot capture an old deployment.""" + if not self.has_capability(wrapper): + return None + return DurableDeployment(name, wrapper, self._provider(), app_settings=self.app_settings) + + def current_deployment(self, name: str) -> DurableDeployment | None: + """Return the currently published durable deployment, if any.""" + return self._deployments.get(name) + + def install_deployment(self, name: str, deployment: DurableDeployment | None) -> None: + """Publish a prepared deployment, or clear a removed durable capability.""" + if deployment is None: + self._deployments.pop(name, None) + else: + self._deployments[name] = deployment + + def deployment(self, name: str, wrapper: BasePipelineWrapper | None = None) -> DurableDeployment: + """Return the published deployment or create an inactive one for the wrapper.""" + existing = self._deployments.get(name) + if existing is not None and (wrapper is None or existing.wrapper is wrapper): + return existing + wrapper = wrapper or (self._registry.get(name) if self._registry is not None else None) + if wrapper is None or not self.has_capability(wrapper): + msg = f"Pipeline '{name}' does not expose durable execution" + raise KeyError(msg) + if existing is not None and existing.manager.started: + msg = ( + f"Pipeline '{name}' has an active durable deployment; " + "use the async deployment transaction to replace it" + ) + raise RuntimeError(msg) + deployment = DurableDeployment(name, wrapper, self._provider(), app_settings=self.app_settings) + self._deployments[name] = deployment + return deployment + + async def start(self) -> None: + """Start runtime-owned deployments after optional registry discovery.""" + if self._started: + return + started: list[DurableDeployment] = [] + self._started = True + try: + if self._registry is not None: + for name in self._registry.get_names(): + wrapper = self._registry.get(name) + if wrapper is not None and self.has_capability(wrapper): + self.deployment(name, wrapper) + for deployment in list(self._deployments.values()): + if deployment.manager.started: + continue + await deployment.start() + started.append(deployment) + if started: + log.bind(deployments=len(started), store=type(self.provider).__name__).info("Durable runtime ready") + except BaseException: + self._started = False + for deployment in reversed(started): + await deployment.close() + raise + + async def close(self) -> None: + """Stop deployments and close their shared provider after draining work.""" + self._started = False + deployments = list(self._deployments.values()) + log.bind(deployments=len(deployments)).debug("Closing durable runtime") + for deployment in reversed(deployments): + await deployment.close() + self._deployments.clear() + if self.provider is not None: + provider = self.provider + self.provider = None + draining = [deployment.manager for deployment in deployments if deployment.manager.draining] + if draining: + + async def close_after_drain() -> None: + await asyncio.gather(*(manager.wait_drained() for manager in draining)) + await provider.close() + + self._provider_close_task = asyncio.create_task( + close_after_drain(), + name="durable-provider-close", + ) + else: + await provider.close() + + async def health(self) -> dict[str, JsonValue]: + """Return aggregate health for all published durable deployments.""" + deployments = dict( + zip( + self._deployments, + await asyncio.gather( + *(deployment.manager.health_snapshot() for deployment in self._deployments.values()) + ), + strict=True, + ) + ) + return { + "healthy": all(bool(health["healthy"]) for health in deployments.values()), + "deployments": cast(JsonValue, deployments), + } + + def _provider(self) -> ExecutionStoreProvider: + if self.provider is None: + if self.app_settings.durable_store == "memory": + log.warning("Durable execution uses volatile in-memory storage; queued work is lost on process exit") + self.provider = InMemoryExecutionStoreProvider(app_settings=self.app_settings) + else: + self.provider = RedisExecutionStoreProvider(app_settings=self.app_settings) + return self.provider + + +class IdempotencyConflictError(RuntimeError): + """An idempotency key was reused with a different operation fingerprint.""" + + +class DefinitionRevisionConflictError(RuntimeError): + """A nonterminal record belongs to an incompatible deployment revision.""" + + code = "definition_revision_conflict" + + +def execution_id_for(owner_id: str, external_task_id: str) -> str: + """Return the internal fixed-size durable ID for one A2A task.""" + return hashlib.sha256(owner_id.encode("utf-8") + b"\0" + external_task_id.encode("utf-8")).hexdigest() + + +def _durable_method_contract(wrapper: BasePipelineWrapper) -> tuple[Any, bool, type[BaseModel], Any]: + sync, asynchronous = _durable_method_implementations(wrapper) + if sync == asynchronous: + msg = "Implement exactly one of run_durable and run_durable_async" + raise PipelineWrapperError(msg) + method = wrapper.run_durable_async if asynchronous else wrapper.run_durable + parameters = list(inspect.signature(method).parameters.values()) + expected_parameters = 2 + if len(parameters) != expected_parameters: + msg = "Durable methods must accept exactly (context: DurableContext, request: PydanticModel)" + raise PipelineWrapperError(msg) + try: + annotations = get_type_hints(method) + except (NameError, TypeError) as error: + msg = f"Invalid durable method annotation: {error}" + raise PipelineWrapperError(msg) from error + context_annotation = annotations.get(parameters[0].name) + request_type = annotations.get(parameters[1].name) + if context_annotation is not DurableContext: + msg = "The first durable method parameter must be annotated DurableContext" + raise PipelineWrapperError(msg) + if not inspect.isclass(request_type) or not issubclass(request_type, BaseModel): + msg = "The durable request parameter must be an annotated Pydantic model" + raise PipelineWrapperError(msg) + return method, asynchronous, cast(type[BaseModel], request_type), annotations.get("return", Any) + + +def _operation_fingerprint( + deployment_name: str, + definition_revision: str, + validated_input: Mapping[str, JsonValue], + *, + owner_id: str | None, +) -> str: + payload = json.dumps( + { + "deployment_name": deployment_name, + "definition_revision": definition_revision, + "validated_input": validated_input, + "owner_id": owner_id, + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + + +def _canonical_json(value: Any) -> JsonValue: + """Serialize Pydantic input deterministically without changing list semantics.""" + if isinstance(value, Mapping): + return {str(key): _canonical_json(item) for key, item in value.items()} + if isinstance(value, list | tuple): + return [_canonical_json(item) for item in value] + if isinstance(value, set | frozenset): + return sorted( + (_canonical_json(item) for item in value), + key=lambda item: json.dumps(item, ensure_ascii=False, sort_keys=True, separators=(",", ":")), + ) + return cast(JsonValue, TypeAdapter(Any).dump_python(value, mode="json")) + + +durable_runtime = DurableRuntime() +durable_runtime._registry = registry + + +__all__ = [ + "DefinitionRevisionConflictError", + "DurableDeployment", + "DurableRuntime", + "ExecutionStoreProvider", + "IdempotencyConflictError", + "durable_runtime", +] diff --git a/src/hayhooks/server/utils/base_pipeline_wrapper.py b/src/hayhooks/server/utils/base_pipeline_wrapper.py index 2e0ae372..1cd48c02 100644 --- a/src/hayhooks/server/utils/base_pipeline_wrapper.py +++ b/src/hayhooks/server/utils/base_pipeline_wrapper.py @@ -19,10 +19,19 @@ class BasePipelineWrapper(ABC): # (a list of dicts with "id", "name", "description", "tags", "examples") a2a_card: dict[str, Any] | None = None + # Optional Pydantic model used to type and validate durable resume input. + durable_resume_model: type[Any] | None = None + + # Required non-empty immutable build revision for durable wrappers. Use a + # Git SHA or image digest in production. + durable_revision: str | None = None + def __init__(self): self.pipeline = None self._is_run_api_implemented = False self._is_run_api_async_implemented = False + self._is_run_durable_implemented = False + self._is_run_durable_async_implemented = False self._is_run_chat_completion_implemented = False self._is_run_chat_completion_async_implemented = False self._is_run_response_implemented = False @@ -64,6 +73,22 @@ async def run_api_async(self): msg = "run_api_async not implemented" raise NotImplementedError(msg) + def run_durable(self): + """ + Run one durable execution in a worker thread. + + Override either this method or :meth:`run_durable_async`, never both. + The first argument after ``self`` must be ``DurableContext`` and the + second must be one annotated Pydantic request model. + """ + msg = "run_durable not implemented" + raise NotImplementedError(msg) + + async def run_durable_async(self): + """Run one durable execution without blocking the server event loop.""" + msg = "run_durable_async not implemented" + raise NotImplementedError(msg) + def run_chat_completion(self, model: str, messages: list[dict], body: dict) -> str | Generator: """ This method is called when a user sends an OpenAI-compatible chat completion request. diff --git a/tests/test_pipeline_utils.py b/tests/test_pipeline_utils.py index fbb6f4ce..5a6aab83 100644 --- a/tests/test_pipeline_utils.py +++ b/tests/test_pipeline_utils.py @@ -206,3 +206,10 @@ def test_preserves_extra_keys(self): ] result = get_input_files(items) assert result[0]["filename"] == "doc.pdf" + + +def test_legacy_pipeline_utilities_remain_publicly_importable(): + from hayhooks import chat_messages_from_openai_response, is_user_message + + assert callable(chat_messages_from_openai_response) + assert callable(is_user_message) From 6bca714fca444bc62ad0cfdfce318dab5d422121 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:09:20 +0200 Subject: [PATCH 05/28] feat(a2a): add durable task execution and recovery --- src/hayhooks/a2a.py | 95 ++++ src/hayhooks/server/a2a/__init__.py | 1 + src/hayhooks/server/a2a/app.py | 165 ++++++ src/hayhooks/server/a2a/cards.py | 94 ++++ src/hayhooks/server/a2a/durable_executor.py | 466 ++++++++++++++++ src/hayhooks/server/a2a/executor.py | 167 ++++++ src/hayhooks/server/a2a/imports.py | 43 ++ src/hayhooks/server/a2a/messages.py | 89 +++ src/hayhooks/server/a2a/redis_task_store.py | 506 ++++++++++++++++++ src/hayhooks/server/a2a/runtime.py | 259 +++++++++ src/hayhooks/server/utils/a2a_utils.py | 399 -------------- tests/test_a2a.py | 224 +++++--- tests/test_durable_a2a.py | 414 ++++++++++++++ .../pipeline_wrapper.py | 43 ++ tests/test_it_a2a_server.py | 231 +++++++- tests/test_redis_a2a_recovery_integration.py | 323 +++++++++++ tests/test_redis_task_store.py | 120 +++++ 17 files changed, 3160 insertions(+), 479 deletions(-) create mode 100644 src/hayhooks/a2a.py create mode 100644 src/hayhooks/server/a2a/__init__.py create mode 100644 src/hayhooks/server/a2a/app.py create mode 100644 src/hayhooks/server/a2a/cards.py create mode 100644 src/hayhooks/server/a2a/durable_executor.py create mode 100644 src/hayhooks/server/a2a/executor.py create mode 100644 src/hayhooks/server/a2a/imports.py create mode 100644 src/hayhooks/server/a2a/messages.py create mode 100644 src/hayhooks/server/a2a/redis_task_store.py create mode 100644 src/hayhooks/server/a2a/runtime.py delete mode 100644 src/hayhooks/server/utils/a2a_utils.py create mode 100644 tests/test_durable_a2a.py create mode 100644 tests/test_files/durable_a2a_process_recovery/pipeline_wrapper.py create mode 100644 tests/test_redis_a2a_recovery_integration.py create mode 100644 tests/test_redis_task_store.py diff --git a/src/hayhooks/a2a.py b/src/hayhooks/a2a.py new file mode 100644 index 00000000..1a21e856 --- /dev/null +++ b/src/hayhooks/a2a.py @@ -0,0 +1,95 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable + +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper + +if TYPE_CHECKING: + from a2a.server.context import ServerCallContext + from a2a.server.tasks import TaskStore + from a2a.types import Task + + from hayhooks.server.a2a.redis_task_store import RedisTaskStore, RedisTaskStoreProvider + + +class A2APipelineWrapper(BasePipelineWrapper): + """Base class for wrappers that expose a managed durable A2A Agent.""" + + durable: bool = True + + +def default_a2a_owner(context: ServerCallContext) -> str: + """Return the built-in stable owner for an A2A request.""" + user = context.user + if not user.is_authenticated: + return "anonymous" + if not user.user_name: + msg = "Authenticated A2A users must have a non-empty user name" + raise ValueError(msg) + return f"user:{user.user_name}" + + +def validate_a2a_owner(owner_id: str) -> str: + """Reject owner resolvers that cannot isolate persisted tasks.""" + if not isinstance(owner_id, str) or not owner_id: + msg = "A2A owner resolvers must return a non-empty string" + raise ValueError(msg) + return owner_id + + +class TaskStoreProvider(ABC): + """Create A2A SDK task stores for the agents mounted by the server.""" + + @abstractmethod + def create_task_store(self, agent_name: str) -> TaskStore: + """Return the task store for an exposed agent.""" + raise NotImplementedError + + async def initialize(self) -> None: + """Validate provider resources during A2A application startup.""" + return None + + async def health(self) -> dict[str, Any]: + """Return a payload-safe readiness projection for provider resources.""" + return {"healthy": True, "provider": type(self).__name__} + + async def close(self) -> None: + """Release resources owned by the provider when the A2A server stops.""" + return None + + +@runtime_checkable +class RecoverableTaskStore(Protocol): + """Optional Redis operations used to recover durable A2A projections.""" + + def owner_id_for_context(self, context: ServerCallContext) -> str: ... + + async def recoverable_task_batch( + self, cursor: int, limit: int + ) -> tuple[list[tuple[Task, str, int]], int | None]: ... + + async def save_projection(self, task: Task, owner: str, expected_version: int) -> bool: + """Persist a projected task through its optimistic version fence.""" + ... + + +def __getattr__(name: str) -> Any: + """Lazily expose optional Redis task-store types without importing A2A runtime.""" + if name in {"RedisTaskStore", "RedisTaskStoreProvider"}: + from hayhooks.server.a2a.redis_task_store import RedisTaskStore, RedisTaskStoreProvider + + return {"RedisTaskStore": RedisTaskStore, "RedisTaskStoreProvider": RedisTaskStoreProvider}[name] + msg = f"module {__name__!r} has no attribute {name!r}" + raise AttributeError(msg) + + +__all__ = [ + "A2APipelineWrapper", + "RecoverableTaskStore", + "RedisTaskStore", + "RedisTaskStoreProvider", + "TaskStoreProvider", + "default_a2a_owner", + "validate_a2a_owner", +] diff --git a/src/hayhooks/server/a2a/__init__.py b/src/hayhooks/server/a2a/__init__.py new file mode 100644 index 00000000..98799eaf --- /dev/null +++ b/src/hayhooks/server/a2a/__init__.py @@ -0,0 +1 @@ +"""A2A server implementation.""" diff --git a/src/hayhooks/server/a2a/app.py b/src/hayhooks/server/a2a/app.py new file mode 100644 index 00000000..49d7f373 --- /dev/null +++ b/src/hayhooks/server/a2a/app.py @@ -0,0 +1,165 @@ +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse +from starlette.routing import Mount, Route + +from hayhooks.a2a import TaskStoreProvider +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode +from hayhooks.durable.runtime import durable_runtime +from hayhooks.server.a2a.cards import create_agent_card, get_a2a_base_url, is_a2a_exposable +from hayhooks.server.a2a.executor import DurableAgentExecutor, create_agent_executor +from hayhooks.server.a2a.imports import DefaultRequestHandler, create_agent_card_routes, create_jsonrpc_routes +from hayhooks.server.a2a.runtime import A2ARuntime, TaskAwareRequestContextBuilder, create_task_store_provider +from hayhooks.server.logger import log +from hayhooks.server.pipelines.registry import registry +from hayhooks.server.tracing import configure_tracing, instrument_starlette_app +from hayhooks.settings import settings + +# Path reserved by the A2A app itself; a pipeline with this name cannot be mounted +_RESERVED_PATHS = frozenset({"status"}) + + +def _create_agent_mount(pipeline_name: str, base_url: str, runtime: A2ARuntime) -> Mount: + wrapper = registry.get(pipeline_name) + if wrapper is None: + msg = f"Pipeline '{pipeline_name}' not found" + raise ValueError(msg) + + card = create_agent_card(pipeline_name, base_url) + + task_store = runtime.create_task_store(pipeline_name) + agent_executor = create_agent_executor(wrapper, pipeline_name, task_store=task_store) + if isinstance(agent_executor, DurableAgentExecutor): + task_store = agent_executor.task_store + runtime.register_agent_executor(agent_executor) + request_handler = DefaultRequestHandler( + agent_executor=agent_executor, + task_store=task_store, + agent_card=card, + request_context_builder=TaskAwareRequestContextBuilder(task_store), + ) + routes = [ + *create_agent_card_routes(card), + *create_jsonrpc_routes(request_handler, rpc_url="/", enable_v0_3_compat=settings.a2a_v0_3_compat), + ] + log.debug( + "Created A2A mount for pipeline '{}' at '/{}' with v0.3_compat={}", + pipeline_name, + pipeline_name, + settings.a2a_v0_3_compat, + ) + return Mount(f"/{pipeline_name}", routes=routes) + + +def _create_app_task_store_provider(durable_agents_deployed: bool) -> TaskStoreProvider: + backend = settings.a2a_task_store + redis_url = settings.a2a_redis_url + redis_key_prefix = settings.a2a_redis_key_prefix + durable_redis = durable_agents_deployed and settings.durable_store == "redis" + if backend == "auto" and durable_redis: + backend = "redis" + redis_url = settings.durable_redis_url + redis_key_prefix = f"{settings.durable_redis_key_prefix.rstrip(':')}:a2a" + elif backend == "auto": + backend = "memory" + + return create_task_store_provider( + backend=backend, + redis_url=redis_url, + redis_key_prefix=redis_key_prefix, + ) + + +def _create_agent_mounts(base_url: str, runtime: A2ARuntime) -> tuple[list[str], list[Mount]]: + agent_names: list[str] = [] + mounts: list[Mount] = [] + for pipeline_name in registry.get_names(): + if not is_a2a_exposable(pipeline_name): + continue + if pipeline_name in _RESERVED_PATHS: + log.warning("Skipping pipeline '{}': the path is reserved by the A2A server", pipeline_name) + continue + try: + mounts.append(_create_agent_mount(pipeline_name, base_url, runtime)) + except Exception as error: + log.opt(exception=True).warning( + "Skipping pipeline '{}': failed to build A2A agent: {}", + pipeline_name, + error, + ) + continue + agent_names.append(pipeline_name) + log.info("Exposing pipeline '{}' as A2A agent at {}/{}/", pipeline_name, base_url, pipeline_name) + return agent_names, mounts + + +def create_a2a_app(*, base_url: str | None = None, debug: bool = False, runtime: A2ARuntime | None = None) -> Starlette: + """ + Create a Starlette app exposing deployed pipelines as A2A agents. + + Each exposable pipeline is mounted under ``/{pipeline_name}/`` with its + agent card at ``/{pipeline_name}/.well-known/agent-card.json`` and the + JSON-RPC binding at ``POST /{pipeline_name}/``. + """ + + durable_agents_deployed = any( + durable_authoring_mode(wrapper) is DurableAuthoringMode.MANAGED_AGENT + for name in registry.get_names() + if (wrapper := registry.get(name)) is not None + ) + if runtime is None: + runtime = A2ARuntime( + task_store_provider=_create_app_task_store_provider(durable_agents_deployed), + ) + log.info("Using A2A task store provider '{}'", type(runtime.task_store_provider).__name__) + base_url = (base_url or get_a2a_base_url()).rstrip("/") + if "//0.0.0.0" in base_url or "//[::]" in base_url: + log.warning( + "Agent cards will advertise the wildcard bind address ({}) which remote clients cannot connect to. " + "Set HAYHOOKS_A2A_EXTERNAL_URL (or --external-url) to the server's reachable base URL.", + base_url, + ) + + agent_names, mounts = _create_agent_mounts(base_url, runtime) + + if not agent_names: + log.warning( + "No pipelines exposable as A2A agents. " + "A pipeline must expose a managed durable Agent or implement run_chat_completion " + "or run_chat_completion_async." + ) + + async def handle_status(request: Request) -> JSONResponse: # noqa: ARG001 + health = await runtime.health() + healthy = bool(health["healthy"]) + return JSONResponse( + { + "status": "ok" if healthy else "unavailable", + "agents": agent_names, + "components": health["components"], + }, + status_code=200 if healthy else 503, + headers={"Cache-Control": "no-store"}, + ) + + @asynccontextmanager + async def lifespan(app: Starlette) -> AsyncIterator[None]: # noqa: ARG001 + try: + await durable_runtime.start() + await runtime.start() + yield + finally: + try: + await runtime.close() + finally: + await durable_runtime.close() + + app = Starlette(debug=debug, routes=[Route("/status", endpoint=handle_status), *mounts], lifespan=lifespan) + log.debug("Created A2A Starlette app with {} mounted agent(s): {}", len(agent_names), agent_names) + + configure_tracing() + instrument_starlette_app(app) + return app diff --git a/src/hayhooks/server/a2a/cards.py b/src/hayhooks/server/a2a/cards.py new file mode 100644 index 00000000..24114087 --- /dev/null +++ b/src/hayhooks/server/a2a/cards.py @@ -0,0 +1,94 @@ +from typing import Any + +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode +from hayhooks.server.a2a.imports import AgentCapabilities, AgentCard, AgentInterface, AgentSkill +from hayhooks.server.logger import log +from hayhooks.server.pipelines.registry import registry +from hayhooks.settings import settings + + +def get_a2a_base_url() -> str: + """Base URL advertised in agent cards, without trailing slash.""" + base_url = settings.a2a_external_url or f"http://{settings.a2a_host}:{settings.a2a_port}" + return base_url.rstrip("/") + + +def is_a2a_exposable(pipeline_name: str) -> bool: + """ + Whether a deployed pipeline can be exposed as an A2A agent. + + A pipeline is exposable when it provides a durable Agent or implements + ``run_chat_completion`` / ``run_chat_completion_async`` and does not set + ``skip_a2a = True``. + """ + pipeline_wrapper = registry.get(pipeline_name) + if pipeline_wrapper is None: + return False + + metadata = registry.get_metadata(name=pipeline_name) or {} + if metadata.get("skip_a2a"): + log.debug("Skipping pipeline '{}': skip_a2a is set", pipeline_name) + return False + + exposable = ( + durable_authoring_mode(pipeline_wrapper) is DurableAuthoringMode.MANAGED_AGENT + or pipeline_wrapper._is_run_chat_completion_implemented + or pipeline_wrapper._is_run_chat_completion_async_implemented + ) + if not exposable: + log.debug("Skipping pipeline '{}': no A2A or chat completion method implemented", pipeline_name) + return exposable + + +def create_agent_card(pipeline_name: str, base_url: str) -> "AgentCard": + """ + Build an A2A agent card for a deployed pipeline. + + Card fields are derived from the pipeline's registry metadata and can be + overridden via the wrapper's ``a2a_card`` class attribute. + """ + + metadata = registry.get_metadata(name=pipeline_name) or {} + overrides = metadata.get("a2a_card") or {} + + name = overrides.get("name") or pipeline_name + description = ( + overrides.get("description") + or metadata.get("description") + or f"Haystack pipeline '{pipeline_name}' deployed with Hayhooks" + ) + version = overrides.get("version") or "1.0.0" + agent_url = f"{base_url.rstrip('/')}/{pipeline_name}/" + + skills_spec: list[dict[str, Any]] = overrides.get("skills") or [ + {"id": pipeline_name, "name": name, "description": description, "tags": ["haystack", "hayhooks"]} + ] + skills = [ + AgentSkill( + id=skill.get("id", pipeline_name), + name=skill.get("name", name), + description=skill.get("description", description), + tags=list(skill.get("tags", [])), + examples=list(skill.get("examples", [])), + ) + for skill in skills_spec + ] + + log.debug( + "Built A2A agent card for pipeline '{}' with name='{}', url='{}', skills={}", + pipeline_name, + name, + agent_url, + [skill.id for skill in skills], + ) + + return AgentCard( + name=name, + description=description, + version=version, + default_input_modes=["text/plain"], + default_output_modes=["text/plain"], + capabilities=AgentCapabilities(streaming=True), + supported_interfaces=[AgentInterface(protocol_binding="JSONRPC", url=agent_url)], + skills=skills, + ) diff --git a/src/hayhooks/server/a2a/durable_executor.py b/src/hayhooks/server/a2a/durable_executor.py new file mode 100644 index 00000000..46168f1b --- /dev/null +++ b/src/hayhooks/server/a2a/durable_executor.py @@ -0,0 +1,466 @@ +"""A2A adapter for managed durable-Agent executions.""" + +from __future__ import annotations + +import asyncio +import builtins +import json +from typing import Any, cast + +from hayhooks.a2a import RecoverableTaskStore, default_a2a_owner +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStatus, ExecutionStoreError +from hayhooks.durable.runtime import DurableDeployment, execution_id_for +from hayhooks.server.a2a.imports import ( + AgentExecutor, + EventQueue, + InvalidParamsError, + RequestContext, + TaskStore, + TaskUpdater, + new_task_from_user_message, + new_text_part, +) +from hayhooks.server.a2a.messages import ( + build_haystack_messages, + build_haystack_resume_messages, + build_haystack_task_messages, + task_is_terminal, + task_matches_filters, +) +from hayhooks.server.logger import log +from hayhooks.settings import settings + +DURABLE_PROGRESS_ARTIFACT_NAME = "durable-progress" +DURABLE_RESULT_ARTIFACT_NAME = "durable-result" + + +class _TaskProjectionQueue: + """Apply A2A events to a transient task snapshot.""" + + def __init__(self, task: Any) -> None: + self.task = task + + async def enqueue_event(self, event: Any) -> None: + from a2a.server.tasks.task_manager import append_artifact_to_task + from a2a.types import TaskArtifactUpdateEvent, TaskStatusUpdateEvent + + if isinstance(event, TaskStatusUpdateEvent): + if self.task.status.HasField("message"): + self.task.history.append(self.task.status.message) + if event.metadata: + self.task.metadata.MergeFrom(event.metadata) + self.task.status.CopyFrom(event.status) + elif isinstance(event, TaskArtifactUpdateEvent): + append_artifact_to_task(self.task, event) + + +class DurableTaskStore(TaskStore): + """Read durable execution state through an ordinary A2A task store.""" + + def __init__(self, task_store: TaskStore, deployment: DurableDeployment) -> None: + self._task_store = task_store + self._deployment = deployment + self._read_through_task_ids: set[str] = set() + + async def save(self, task: Any, context: Any) -> None: + await self._task_store.save(task, context) + + async def get(self, task_id: str, context: Any) -> Any | None: + task = await self._task_store.get(task_id, context) + return await self._project(task, self.owner_id_for_context(context), context) + + async def list(self, params: Any, context: Any) -> Any: + from a2a.types import ListTasksResponse + from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE + from a2a.utils.task import decode_page_token, encode_page_token + + tasks = await self._all_tasks(params, context) + owner_id = self.owner_id_for_context(context) + projected: builtins.list[Any | None] = [] + # ponytail: fixed fan-out keeps Redis pools bounded; make adaptive only + # if task-list latency becomes a measured bottleneck. + for offset in range(0, len(tasks), 32): + projected.extend( + await asyncio.gather(*(self._project(task, owner_id, context) for task in tasks[offset : offset + 32])) + ) + timestamp_after = params.status_timestamp_after if params.HasField("status_timestamp_after") else None + filtered = [ + task for task in projected if task is not None and task_matches_filters(task, params, timestamp_after) + ] + page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE + start = 0 + if params.page_token: + task_id = decode_page_token(params.page_token) + try: + start = next(index for index, task in enumerate(filtered) if task.id == task_id) + 1 + except StopIteration as error: + msg = f"Invalid page token: {params.page_token}" + raise InvalidParamsError(msg) from error + page = filtered[start : start + page_size] + next_page_token = encode_page_token(page[-1].id) if start + len(page) < len(filtered) else None + return ListTasksResponse( + tasks=page, + next_page_token=next_page_token, + page_size=page_size, + total_size=len(filtered), + ) + + async def delete(self, task_id: str, context: Any) -> None: + await self._task_store.delete(task_id, context) + + def owner_id_for_context(self, context: Any) -> str: + resolver = getattr(self._task_store, "owner_id_for_context", None) + call_context = getattr(context, "call_context", context) + return resolver(call_context) if callable(resolver) else default_a2a_owner(call_context) + + async def _all_tasks(self, params: Any, context: Any) -> builtins.list[Any]: + """Read the owner task list before applying live durable-state filters.""" + # ponytail: scans one owner's A2A tasks for exact live-state filters; + # add a durable-state index only if this becomes a measured hot path. + request = type(params)() + request.CopyFrom(params) + request.ClearField("context_id") + request.ClearField("status") + request.ClearField("status_timestamp_after") + request.ClearField("page_token") + request.page_size = settings.a2a_list_scan_batch_size + tasks: list[Any] = [] + while True: + page = await self._task_store.list(request, context) + tasks.extend(page.tasks) + if not page.next_page_token: + return tasks + request.page_token = page.next_page_token + + async def _project(self, task: Any | None, owner_id: str, context: Any) -> Any | None: # noqa: C901 + if task is None: + return None + projected = type(task)() + projected.CopyFrom(task) + copy_version = getattr(self._task_store, "copy_task_version", None) + if callable(copy_version): + copy_version(task, projected) + settled = False + try: + record = await self._deployment.get( + execution_id_for(owner_id, task.id), + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) + except KeyError: + if task_is_terminal(task): + return projected + if task.HasField("status"): + from a2a.types import TaskState + + if task.status.state == TaskState.TASK_STATE_SUBMITTED: + return projected + updater = TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id) + await updater.failed( + message=updater.new_agent_message( + [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] + ) + ) + settled = True + else: + if _task_matches_record(task, record): + return projected + settled = await _project_record( + record, + TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id), + ) + if settled and task.id in self._read_through_task_ids: + try: + await self._task_store.save(projected, context) + except InvalidParamsError: + pass + else: + self._read_through_task_ids.discard(task.id) + return projected + + +class DurableAgentExecutor(AgentExecutor): + """Run and stream a managed durable Agent without a second lifecycle store.""" + + def __init__(self, pipeline_name: str, task_store: TaskStore, deployment: DurableDeployment) -> None: + self.pipeline_name = pipeline_name + self._task_store = task_store + self.task_store = DurableTaskStore(task_store, deployment) + self.deployment = deployment + self._closed = False + + def health(self) -> dict[str, Any]: + return {"healthy": not self._closed} + + async def start(self) -> None: + self._closed = False + if isinstance(self._task_store, RecoverableTaskStore): + await self._recover_tasks(cast(RecoverableTaskStore, self._task_store)) + + async def close(self) -> None: + self._closed = True + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + """Submit or resume a durable execution and project its state into A2A events.""" + task = _task(context) + self.task_store._read_through_task_ids.discard(task.id) + updater = TaskUpdater(event_queue, task.id, task.context_id) + owner_id = self.task_store.owner_id_for_context(context) + execution_id = execution_id_for(owner_id, task.id) + record = None + if context.current_task is not None: + try: + record = await self.deployment.get( + execution_id, + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) + except KeyError: + await updater.failed( + message=updater.new_agent_message( + [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] + ) + ) + return + if record is not None and record.status is ExecutionStatus.WAITING: + action = "resume" + resumed = await self.deployment.resume( + execution_id, + {"messages": [message.to_dict() for message in build_haystack_resume_messages(context)]}, + owner_id=owner_id, + enforce_owner=True, + ) + if not resumed: + msg = f"Task '{task.id}' is no longer accepting follow-up messages" + raise InvalidParamsError(msg) + elif record is None: + action = "submit" + await self._task_store.save(task, context.call_context) + try: + record = await self._submit(task.id, owner_id, build_haystack_messages(context)) + except ValueError as error: + log.bind( + pipeline_name=self.pipeline_name, + task_id=task.id, + error_type=type(error).__name__, + ).warning("Rejected durable A2A task submission") + await updater.failed( + message=updater.new_agent_message( + [new_text_part("The durable Agent submission was rejected (durable_submission_rejected).")] + ) + ) + return + execution_id = record.execution_id + else: + msg = f"Task '{task.id}' is already running and cannot accept another message" + raise InvalidParamsError(msg) + log.bind( + pipeline_name=self.pipeline_name, + task_id=task.id, + execution_id=execution_id, + action=action, + ).debug("Accepted durable A2A task action") + await updater.start_work() + await self._wait_for_update(execution_id, owner_id, updater) + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + """Request cancellation and continue projecting until the task settles.""" + task = context.current_task + if task is None: + return + self.task_store._read_through_task_ids.discard(task.id) + owner_id = self.task_store.owner_id_for_context(context) + execution_id = execution_id_for(owner_id, task.id) + try: + accepted = await self.deployment.request_cancel( + execution_id, + owner_id=owner_id, + enforce_owner=True, + ) + except KeyError: + return + if not accepted: + return + log.bind(pipeline_name=self.pipeline_name, task_id=task.id, execution_id=execution_id).debug( + "Accepted durable A2A task cancellation" + ) + updater = TaskUpdater(event_queue, task.id, task.context_id) + await updater.add_artifact( + [new_text_part("Cancellation requested")], + artifact_id=f"{task.id}-{DURABLE_PROGRESS_ARTIFACT_NAME}", + name=DURABLE_PROGRESS_ARTIFACT_NAME, + append=False, + ) + await self._wait_for_update(execution_id, owner_id, updater, terminal_only=True) + + async def _wait_for_update( + self, execution_id: str, owner_id: str, updater: TaskUpdater, *, terminal_only: bool = False + ) -> None: + last_sequence = -1 + while not self._closed: + try: + record = await self.deployment.get( + execution_id, + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) + except KeyError: + await updater.failed( + message=updater.new_agent_message( + [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] + ) + ) + return + except ExecutionStoreError: + await asyncio.sleep(max(0.1, settings.durable_poll_interval)) + continue + if record.sequence != last_sequence: + last_sequence = record.sequence + settled = await _project_record(record, updater) + if settled and (not terminal_only or record.status is not ExecutionStatus.WAITING): + log.bind( + pipeline_name=self.pipeline_name, + task_id=updater.task_id, + execution_id=execution_id, + status=record.status.value, + ).debug("Projected durable A2A task state") + return + await asyncio.sleep(max(0.1, settings.durable_poll_interval)) + + async def _submit(self, task_id: str, owner_id: str, messages: list[Any]) -> Any: + payload = {"messages": [message.to_dict() for message in messages]} + while not self._closed: + try: + return (await self.deployment.submit(payload, execution_id=task_id, owner_id=owner_id))[1] + except (ExecutionAdmissionError, ExecutionStoreError): + await asyncio.sleep(max(0.1, settings.durable_poll_interval)) + raise asyncio.CancelledError + + async def _recover_tasks(self, task_store: RecoverableTaskStore) -> None: # noqa: C901 + """Repair durable A2A tasks saved before this executor started.""" + cursor = 0 + recovered = 0 + while not self._closed: + tasks, next_cursor = await task_store.recoverable_task_batch(cursor, settings.a2a_list_scan_batch_size) + for task, owner_id, version in tasks: + execution_id = execution_id_for(owner_id, task.id) + try: + record = await self.deployment.get( + execution_id, + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) + except KeyError: + from a2a.types import TaskState + + if task.status.state != TaskState.TASK_STATE_SUBMITTED: + continue + try: + record = await self._submit(task.id, owner_id, build_haystack_task_messages(task)) + except ValueError as error: + record = None + log.bind( + pipeline_name=self.pipeline_name, + task_id=task.id, + error_type=type(error).__name__, + ).warning("Rejected recovered durable A2A task submission") + if record is not None and record.status in {ExecutionStatus.QUEUED, ExecutionStatus.RUNNING}: + self.task_store._read_through_task_ids.add(task.id) + if record is not None and _task_matches_record(task, record): + continue + projected = type(task)() + projected.CopyFrom(task) + updater = TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id) + if record is None: + await updater.failed( + message=updater.new_agent_message( + [new_text_part("The durable Agent submission was rejected (durable_submission_rejected).")] + ) + ) + else: + await _project_record(record, updater) + if not await task_store.save_projection(projected, owner_id, version): + continue + task.CopyFrom(projected) + recovered += 1 + if next_cursor is None: + if recovered: + log.bind(pipeline_name=self.pipeline_name, recovered=recovered).debug( + "Recovered durable A2A task projections" + ) + return + cursor = next_cursor + + +def _task(context: RequestContext) -> Any: + if context.current_task is not None: + return context.current_task + if context.message is None: + msg = "A2A request has neither a current task nor a message" + raise ValueError(msg) + return new_task_from_user_message(context.message) + + +async def _project_record(record: Any, updater: TaskUpdater) -> bool: + """Project one durable record and report whether the active request can stop waiting.""" + if record.progress: + await updater.add_artifact( + [new_text_part("\n".join(event.message for event in record.progress))], + artifact_id=f"{updater.task_id}-{DURABLE_PROGRESS_ARTIFACT_NAME}", + name=DURABLE_PROGRESS_ARTIFACT_NAME, + append=False, + ) + if record.status is ExecutionStatus.WAITING: + await updater.requires_input( + message=updater.new_agent_message([new_text_part("The durable Agent requires input to continue.")]) + ) + elif record.status is ExecutionStatus.COMPLETED: + await updater.add_artifact( + [new_text_part(_result_text(record.result))], + artifact_id=f"{updater.task_id}-{DURABLE_RESULT_ARTIFACT_NAME}", + name=DURABLE_RESULT_ARTIFACT_NAME, + append=False, + last_chunk=True, + ) + await updater.complete() + elif record.status is ExecutionStatus.FAILED: + text = record.error.message if record.error else "Durable Agent execution failed" + await updater.failed(message=updater.new_agent_message([new_text_part(text)])) + elif record.status is ExecutionStatus.CANCELED: + await updater.cancel( + message=updater.new_agent_message([new_text_part("The durable Agent execution was canceled.")]) + ) + else: + return False + return True + + +def _task_matches_record(task: Any, record: Any) -> bool: + from a2a.types import TaskState + + states = { + ExecutionStatus.WAITING: TaskState.TASK_STATE_INPUT_REQUIRED, + ExecutionStatus.COMPLETED: TaskState.TASK_STATE_COMPLETED, + ExecutionStatus.FAILED: TaskState.TASK_STATE_FAILED, + ExecutionStatus.CANCELED: TaskState.TASK_STATE_CANCELED, + } + return record.status in states and task.status.state == states[record.status] + + +def _result_text(result: Any) -> str: + if isinstance(result, dict): + last = result.get("last_message") + if isinstance(last, dict): + content = last.get("content") + if isinstance(content, str): + return content + if isinstance(content, list): + text_parts = [part.get("text", "") for part in content if isinstance(part, dict)] + if text := "".join(text_parts): + return text + return json.dumps(result, ensure_ascii=False, default=str) + return str(result or "") diff --git a/src/hayhooks/server/a2a/executor.py b/src/hayhooks/server/a2a/executor.py new file mode 100644 index 00000000..7c736d9e --- /dev/null +++ b/src/hayhooks/server/a2a/executor.py @@ -0,0 +1,167 @@ +"""A2A chat executor and authoring-mode selection.""" + +from __future__ import annotations + +import asyncio +import traceback +import uuid +from collections.abc import AsyncGenerator, AsyncIterator, Iterator +from typing import Any + +from fastapi.concurrency import iterate_in_threadpool, run_in_threadpool +from haystack.dataclasses import StreamingChunk + +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode +from hayhooks.durable.runtime import durable_runtime +from hayhooks.server.a2a.durable_executor import DurableAgentExecutor +from hayhooks.server.a2a.imports import ( + AgentExecutor, + EventQueue, + RequestContext, + TaskUpdater, + new_task_from_user_message, + new_text_part, +) +from hayhooks.server.a2a.messages import build_openai_messages +from hayhooks.server.logger import log +from hayhooks.server.tracing import SPAN_A2A_RUN_AGENT, build_trace_tags, trace_operation +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.settings import settings + +RESPONSE_ARTIFACT_NAME = "response" + + +def _stream_item_to_text(item: Any) -> str | None: + if isinstance(item, StreamingChunk): + return item.content or None + if isinstance(item, str): + return item or None + if isinstance(item, bytes): + return item.decode("utf-8", errors="replace") or None + return None + + +async def _iter_text_chunks(result: Any) -> AsyncGenerator[str, None]: + if isinstance(result, str): + yield result + elif isinstance(result, AsyncIterator): + async for item in result: + if text := _stream_item_to_text(item): + yield text + elif isinstance(result, Iterator): + async for item in iterate_in_threadpool(result): + if text := _stream_item_to_text(item): + yield text + else: + msg = f"run_chat_completion returned unsupported type '{type(result).__name__}'" + raise ValueError(msg) + + +async def _stream_result_as_artifact(result: Any, updater: TaskUpdater) -> None: + artifact_id = str(uuid.uuid4()) + first = True + + async def emit(text: str, *, last: bool) -> None: + nonlocal first + await updater.add_artifact( + [new_text_part(text)], + artifact_id=artifact_id, + name=RESPONSE_ARTIFACT_NAME, + append=not first, + last_chunk=last, + ) + first = False + + if isinstance(result, str): + await emit(result, last=True) + return + async for text in _iter_text_chunks(result): + await emit(text, last=False) + await emit("", last=True) + + +class ChatCompletionAgentExecutor(AgentExecutor): + """Run a deployed wrapper's existing chat-completion capability through A2A.""" + + def __init__(self, pipeline_name: str, pipeline_wrapper: BasePipelineWrapper) -> None: + self.pipeline_name = pipeline_name + self.pipeline_wrapper = pipeline_wrapper + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + """Run one A2A task and stream its result through task artifacts.""" + if context.current_task is not None: + task = context.current_task + elif context.message is not None: + task = new_task_from_user_message(context.message) + await event_queue.enqueue_event(task) + else: + msg = "A2A request has neither a current task nor a message" + raise ValueError(msg) + + updater = TaskUpdater(event_queue, task.id, task.context_id) + await updater.start_work() + task_log = log.bind(pipeline_name=self.pipeline_name, task_id=task.id) + task_log.debug("Started A2A task") + + with trace_operation( + SPAN_A2A_RUN_AGENT, + tags=build_trace_tags({"hayhooks.transport": "a2a", "hayhooks.pipeline.name": self.pipeline_name}), + ): + try: + messages = build_openai_messages(context) + if self.pipeline_wrapper._is_run_chat_completion_async_implemented: + result = await self.pipeline_wrapper.run_chat_completion_async( + model=self.pipeline_name, + messages=messages, + body={}, + ) + else: + result = await run_in_threadpool( + self.pipeline_wrapper.run_chat_completion, + model=self.pipeline_name, + messages=messages, + body={}, + ) + await _stream_result_as_artifact(result, updater) + except asyncio.CancelledError: + raise + except Exception as error: + message = f"Error running pipeline '{self.pipeline_name}' as A2A agent: {error}" + if settings.show_tracebacks: + message += f"\n{traceback.format_exc()}" + log.opt(exception=True).error(message) + await updater.failed(message=updater.new_agent_message([new_text_part(message)])) + return + await updater.complete() + task_log.debug("Completed A2A task") + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + """Cancel the current A2A task when the request references one.""" + if context.current_task is not None: + await TaskUpdater(event_queue, context.current_task.id, context.current_task.context_id).cancel() + log.bind(pipeline_name=self.pipeline_name, task_id=context.current_task.id).debug("Canceled A2A task") + + +def create_agent_executor( + wrapper: BasePipelineWrapper, + pipeline_name: str, + *, + task_store: Any | None = None, +) -> AgentExecutor: + """Select a managed durable Agent or chat-compatible executor.""" + if durable_authoring_mode(wrapper) is DurableAuthoringMode.MANAGED_AGENT: + if task_store is None: + msg = "A durable A2A Agent requires an A2A task store" + raise RuntimeError(msg) + # Constructing the deployment here validates the Haystack v3 Agent and the + # durable definition before the Agent Card is exposed. + deployment = durable_runtime.deployment(pipeline_name, wrapper) + return DurableAgentExecutor(pipeline_name, task_store, deployment) + return ChatCompletionAgentExecutor(pipeline_name, wrapper) + + +__all__ = [ + "ChatCompletionAgentExecutor", + "DurableAgentExecutor", + "create_agent_executor", +] diff --git a/src/hayhooks/server/a2a/imports.py b/src/hayhooks/server/a2a/imports.py new file mode 100644 index 00000000..d8ed167c --- /dev/null +++ b/src/hayhooks/server/a2a/imports.py @@ -0,0 +1,43 @@ +from haystack.lazy_imports import LazyImport + +INSTALL_A2A_MESSAGE = "Run 'pip install \"hayhooks[a2a]\"' to install A2A support." + +with LazyImport(INSTALL_A2A_MESSAGE) as a2a_import: + from a2a.helpers import get_message_text, new_task_from_user_message, new_text_part + from a2a.server.agent_execution import ( + AgentExecutor, + RequestContext, + RequestContextBuilder, + SimpleRequestContextBuilder, + ) + from a2a.server.events import EventQueue + from a2a.server.request_handlers import DefaultRequestHandler + from a2a.server.routes import create_agent_card_routes, create_jsonrpc_routes + from a2a.server.tasks import InMemoryTaskStore, TaskStore, TaskUpdater + from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill, Role + from a2a.utils.errors import InvalidParamsError + +a2a_import.check() + +__all__ = [ + "AgentCapabilities", + "AgentCard", + "AgentExecutor", + "AgentInterface", + "AgentSkill", + "DefaultRequestHandler", + "EventQueue", + "InMemoryTaskStore", + "InvalidParamsError", + "RequestContext", + "RequestContextBuilder", + "Role", + "SimpleRequestContextBuilder", + "TaskStore", + "TaskUpdater", + "create_agent_card_routes", + "create_jsonrpc_routes", + "get_message_text", + "new_task_from_user_message", + "new_text_part", +] diff --git a/src/hayhooks/server/a2a/messages.py b/src/hayhooks/server/a2a/messages.py new file mode 100644 index 00000000..e9e10483 --- /dev/null +++ b/src/hayhooks/server/a2a/messages.py @@ -0,0 +1,89 @@ +"""Conversions between A2A SDK messages and Haystack chat messages.""" + +from __future__ import annotations + +from typing import Any + +from hayhooks.server.a2a.imports import RequestContext, Role, get_message_text + + +def build_openai_messages(context: RequestContext) -> list[dict[str, str]]: + """Map A2A history plus the current message to OpenAI-compatible messages.""" + messages: list[dict[str, str]] = [] + history = list(context.current_task.history) if context.current_task else [] + history_ids = {message.message_id for message in history} + for message in history: + text = get_message_text(message) + if text: + messages.append({"role": "assistant" if message.role == Role.ROLE_AGENT else "user", "content": text}) + if context.message is not None and context.message.message_id not in history_ids: + text = get_message_text(context.message) + if text: + messages.append({"role": "user", "content": text}) + return messages + + +def _haystack_message(role: str, text: str) -> Any: + from haystack.dataclasses import ChatMessage + + return ChatMessage.from_assistant(text) if role == "assistant" else ChatMessage.from_user(text) + + +def build_haystack_messages(context: RequestContext) -> list[Any]: + return [_haystack_message(message["role"], message["content"]) for message in build_openai_messages(context)] + + +def build_haystack_task_messages(task: Any) -> list[Any]: + """Rebuild the original durable input from persisted A2A history.""" + return [ + _haystack_message("assistant" if message.role == Role.ROLE_AGENT else "user", text) + for message in task.history + if (text := get_message_text(message)) + ] + + +def build_haystack_resume_messages(context: RequestContext) -> list[Any]: + """Convert only the follow-up turn; recovered Agent state already contains history.""" + if context.message is None: + return [] + + text = get_message_text(context.message) + if not text: + return [] + + role = "assistant" if context.message.role == Role.ROLE_AGENT else "user" + return [_haystack_message(role, text)] + + +def task_is_terminal(task: Any) -> bool: + from a2a.types import TaskState + + return task.HasField("status") and task.status.state in { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_REJECTED, + } + + +def task_matches_filters(task: Any, params: Any, timestamp_after: Any | None) -> bool: + if params.context_id and task.context_id != params.context_id: + return False + if params.status and task.status.state != params.status: + return False + return timestamp_after is None or ( + task.HasField("status") + and task.status.HasField("timestamp") + and (task.status.timestamp.seconds, task.status.timestamp.nanos) + >= (timestamp_after.seconds, timestamp_after.nanos) + ) + + +__all__ = [ + "build_haystack_messages", + "build_haystack_resume_messages", + "build_haystack_task_messages", + "build_openai_messages", + "task_is_terminal", + "task_matches_filters", +] diff --git a/src/hayhooks/server/a2a/redis_task_store.py b/src/hayhooks/server/a2a/redis_task_store.py new file mode 100644 index 00000000..369ec026 --- /dev/null +++ b/src/hayhooks/server/a2a/redis_task_store.py @@ -0,0 +1,506 @@ +"""Optional Redis-backed persistence for A2A tasks.""" + +from __future__ import annotations + +import builtins +from collections import OrderedDict +from collections.abc import Callable, Mapping +from typing import TYPE_CHECKING, Any + +from google.protobuf.message import DecodeError + +from hayhooks.a2a import TaskStoreProvider, default_a2a_owner, validate_a2a_owner +from hayhooks.durable.backend import ( + DEFAULT_TRANSACTION_MAX_RETRIES, + ExecutionContentionError, + ExecutionStoreCorruptionError, +) +from hayhooks.durable.redis import digest, redis_time_ms, redis_transaction_backoff, redis_watch_error +from hayhooks.server.a2a.imports import InvalidParamsError, TaskStore +from hayhooks.server.a2a.messages import task_is_terminal, task_matches_filters +from hayhooks.settings import settings + +if TYPE_CHECKING: + from a2a.server.context import ServerCallContext + from a2a.types import ListTasksRequest, ListTasksResponse, Task + + +OwnerResolver = Callable[[Any], str] +_STALE_VERSION = -2 +_OWNER_MISMATCH = 0 + + +_default_owner_resolver = default_a2a_owner + + +class RedisTaskStore(TaskStore): + """Persist A2A tasks in Redis with agent- and owner-scoped keys.""" + + def __init__( + self, + redis: Any, + agent_name: str, + *, + key_prefix: str | None = None, + owner_resolver: OwnerResolver = _default_owner_resolver, + terminal_ttl_seconds: int | None = None, + ) -> None: + self.redis = redis + self.key_prefix = (key_prefix or settings.a2a_redis_key_prefix).rstrip(":") + self._base_key = f"{self.key_prefix}:v2:{{{digest('a2a-agent', agent_name)}}}" + self.owner_resolver = owner_resolver + self.terminal_ttl_seconds = terminal_ttl_seconds or settings.a2a_terminal_task_ttl_seconds + # Versions belong to the loaded protobuf snapshot, not merely its task + # ID. Multiple requests can legitimately hold different snapshots. A + # global LRU prevents one long-lived process from retaining every task + # object it has ever loaded. + self._loaded_task_versions: OrderedDict[tuple[str, int], tuple[Any, int]] = OrderedDict() + + def _remember_task_version(self, task: Task, version: int) -> None: + key = (task.id, id(task)) + self._loaded_task_versions[key] = (task, version) + self._loaded_task_versions.move_to_end(key) + while len(self._loaded_task_versions) > settings.a2a_task_snapshot_cache_size: + self._loaded_task_versions.popitem(last=False) + + def _loaded_task_version(self, task: Task) -> int: + key = (task.id, id(task)) + snapshot = self._loaded_task_versions.get(key) + if snapshot is None or snapshot[0] is not task: + return -1 + self._loaded_task_versions.move_to_end(key) + return snapshot[1] + + def copy_task_version(self, source: Task, target: Task) -> None: + """Preserve optimistic-write state when an adapter clones a loaded task.""" + version = self._loaded_task_version(source) + if version >= 0: + self._remember_task_version(target, version) + + def _forget_task_versions(self, task_id: str) -> None: + for key in [key for key in self._loaded_task_versions if key[0] == task_id]: + self._loaded_task_versions.pop(key, None) + + def _key(self, suffix: str) -> str: + return f"{self._base_key}:{suffix}" + + def _task_key(self, task_id: str) -> str: + return self._key(f"task:{digest('a2a-task', task_id)}") + + def _updates_key(self, owner: str) -> str: + return self._key(f"owner:{digest('a2a-owner', owner)}:updates") + + def owner_id_for_context(self, context: ServerCallContext) -> str: + """Resolve ownership once for every Redis-backed A2A operation.""" + return validate_a2a_owner(self.owner_resolver(context)) + + @staticmethod + def _deserialize(payload: bytes | str) -> Task: + from a2a.types import Task + + if isinstance(payload, str): + payload = payload.encode("utf-8") + task = Task() + task.ParseFromString(payload) + return task + + @staticmethod + def _serialize(task: Task) -> bytes: + return task.SerializeToString() + + @staticmethod + def _decode_value(value: bytes | str | int) -> str: + return value.decode("utf-8") if isinstance(value, bytes) else str(value) + + @classmethod + def _decode_snapshot(cls, task_id: str, values: Mapping[Any, Any]) -> tuple[Task, str, int, int]: + try: + snapshot = {cls._decode_value(key): value for key, value in values.items()} + owner = validate_a2a_owner(cls._decode_value(snapshot["owner"])) + version = int(snapshot["version"]) + terminal_expiry_ms = int(snapshot["terminal_expiry_ms"]) + task = cls._deserialize(snapshot["payload"]) + if task.id != task_id or version < 1 or terminal_expiry_ms < 0: + raise ValueError + except (AttributeError, DecodeError, KeyError, TypeError, UnicodeError, ValueError) as error: + msg = f"A2A task '{task_id}' has an invalid Redis snapshot" + raise ExecutionStoreCorruptionError(msg) from error + return task, owner, version, terminal_expiry_ms + + @staticmethod + def _task_score(task: Task) -> float: + if task.HasField("status") and task.status.HasField("timestamp"): + return task.status.timestamp.ToNanoseconds() / 1_000_000_000 + return -1.0 + + async def _save_payload( + self, + task: Task, + owner: str, + *, + expected_version: int | None = None, + ) -> int: + """Persist one task and its indexes through an optimistic transaction.""" + await self.cleanup_expired_tasks(limit=10) + expected_version = self._loaded_task_version(task) if expected_version is None else expected_version + task_key = self._task_key(task.id) + for attempt in range(DEFAULT_TRANSACTION_MAX_RETRIES): + async with self.redis.pipeline(transaction=True) as pipe: + try: + await pipe.watch(task_key) + values = await pipe.hgetall(task_key) + if values: + _, recorded_owner, current_version, _ = self._decode_snapshot(task.id, values) + if recorded_owner != owner: + return _OWNER_MISMATCH + if expected_version < 0 or current_version != expected_version: + return _STALE_VERSION + else: + if expected_version >= 0: + return _STALE_VERSION + current_version = 0 + + now_ms = await redis_time_ms(pipe) + version = current_version + 1 + terminal_expiry_ms = now_ms + self.terminal_ttl_seconds * 1_000 if task_is_terminal(task) else 0 + pipe.multi() + pipe.hset( + task_key, + mapping={ + "owner": owner, + "payload": self._serialize(task), + "version": version, + "terminal_expiry_ms": terminal_expiry_ms, + }, + ) + pipe.zadd(self._updates_key(owner), {task.id: self._task_score(task)}) + if terminal_expiry_ms: + pipe.zrem(self._key("active"), task.id) + pipe.zadd(self._key("terminal-expiry"), {task.id: terminal_expiry_ms}) + else: + pipe.zadd(self._key("active"), {task.id: self._task_score(task)}) + pipe.zrem(self._key("terminal-expiry"), task.id) + await pipe.execute() + self._remember_task_version(task, version) + return version + except redis_watch_error(): + await redis_transaction_backoff(attempt) + msg = "A2A task save transaction retry budget exhausted" + raise ExecutionContentionError(msg) + + async def _load_snapshots(self, task_ids: builtins.list[str]) -> builtins.list[tuple[Task, str, int, int] | None]: + if not task_ids: + return [] + async with self.redis.pipeline(transaction=False) as pipe: + for task_id in task_ids: + pipe.hgetall(self._task_key(task_id)) + values = await pipe.execute() + return [ + self._decode_snapshot(task_id, value) if value else None + for task_id, value in zip(task_ids, values, strict=True) + ] + + async def save(self, task: Task, context: ServerCallContext) -> None: + owner = self.owner_id_for_context(context) + saved = await self._save_payload(task, owner) + if saved == _STALE_VERSION: + msg = f"Task '{task.id}' has a stale projection version" + raise InvalidParamsError(msg) + if saved < 1: + msg = f"Task '{task.id}' belongs to another owner" + raise InvalidParamsError(msg) + + async def get(self, task_id: str, context: ServerCallContext) -> Task | None: + values = await self.redis.hgetall(self._task_key(task_id)) + if not values: + return None + task, owner, version, _ = self._decode_snapshot(task_id, values) + if owner != self.owner_id_for_context(context): + return None + self._remember_task_version(task, version) + return task + + async def save_projection(self, task: Task, owner: str, expected_version: int) -> bool: + return await self._save_payload(task, owner, expected_version=expected_version) > 0 + + async def recoverable_task_batch( + self, cursor: int, limit: int + ) -> tuple[builtins.list[tuple[Task, str, int]], int | None]: + """Return one active-task page for restart projection.""" + await self.cleanup_expired_tasks(limit=limit) + next_cursor, entries = await self.redis.zscan(self._key("active"), cursor=cursor, count=limit) + task_ids = [self._decode_value(raw_task_id) for raw_task_id, _score in entries] + tasks: builtins.list[tuple[Task, str, int]] = [] + for snapshot in await self._load_snapshots(task_ids): + if snapshot is None: + continue + task, owner, version, _ = snapshot + self._remember_task_version(task, version) + tasks.append((task, owner, version)) + return tasks, int(next_cursor) or None + + async def cleanup_expired_tasks(self, *, limit: int = 100) -> int: + now_ms = await redis_time_ms(self.redis) + expired = await self.redis.zrangebyscore( + self._key("terminal-expiry"), + "-inf", + now_ms, + start=0, + num=limit, + ) + removed = 0 + for raw_task_id in expired: + task_id = self._decode_value(raw_task_id) + deleted = await self._delete_payload(task_id, expired_before_ms=now_ms) + removed += deleted + if deleted: + self._forget_task_versions(task_id) + return removed + + async def _delete_payload( + self, + task_id: str, + *, + owner: str | None = None, + expired_before_ms: int | None = None, + ) -> int: + task_key = self._task_key(task_id) + for attempt in range(DEFAULT_TRANSACTION_MAX_RETRIES): + async with self.redis.pipeline(transaction=True) as pipe: + try: + await pipe.watch(task_key) + values = await pipe.hgetall(task_key) + if not values: + if expired_before_ms is None: + return 0 + pipe.multi() + pipe.zrem(self._key("active"), task_id) + pipe.zrem(self._key("terminal-expiry"), task_id) + await pipe.execute() + return 0 + + _, recorded_owner, _, terminal_expiry_ms = self._decode_snapshot(task_id, values) + if owner is not None and recorded_owner != owner: + return 0 + if expired_before_ms is not None and ( + not terminal_expiry_ms or terminal_expiry_ms > expired_before_ms + ): + pipe.multi() + if terminal_expiry_ms: + pipe.zadd(self._key("terminal-expiry"), {task_id: terminal_expiry_ms}) + else: + pipe.zrem(self._key("terminal-expiry"), task_id) + await pipe.execute() + return 0 + + pipe.multi() + pipe.delete(task_key) + pipe.zrem(self._updates_key(recorded_owner), task_id) + pipe.zrem(self._key("active"), task_id) + pipe.zrem(self._key("terminal-expiry"), task_id) + await pipe.execute() + return 1 + except redis_watch_error(): + await redis_transaction_backoff(attempt) + msg = "A2A task delete transaction retry budget exhausted" + raise ExecutionContentionError(msg) + + async def list(self, params: ListTasksRequest, context: ServerCallContext) -> ListTasksResponse: + await self.cleanup_expired_tasks(limit=100) + from a2a.types import ListTasksResponse + from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE + from a2a.utils.task import decode_page_token, encode_page_token + + page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE + if params.context_id or params.status or params.HasField("status_timestamp_after"): + page, total_size, next_page_token = await self._list_filtered( + params, + context, + page_size, + decode_page_token, + encode_page_token, + ) + else: + page, total_size, next_page_token = await self._list_by_recent_update( + params, + context, + page_size, + decode_page_token, + encode_page_token, + ) + return ListTasksResponse( + tasks=page, + next_page_token=next_page_token, + page_size=page_size, + total_size=total_size, + ) + + async def _list_by_recent_update( + self, + params: ListTasksRequest, + context: ServerCallContext, + page_size: int, + decode_page_token: Callable[[str], str], + encode_page_token: Callable[[str], str], + ) -> tuple[builtins.list[Task], int, str | None]: + owner = self.owner_id_for_context(context) + updates_key = self._updates_key(owner) + total_size = await self.redis.zcard(updates_key) + start_index = 0 + if params.page_token: + start_task_id = decode_page_token(params.page_token) + rank = await self.redis.zrevrank(updates_key, start_task_id) + if rank is None: + msg = f"Invalid page token: {params.page_token}" + raise InvalidParamsError(msg) + start_index = rank + 1 + + task_ids = await self.redis.zrevrange(updates_key, start_index, start_index + page_size - 1) + task_ids = [self._decode_value(task_id) for task_id in task_ids] + snapshots = await self._load_snapshots(task_ids) + page = [snapshot[0] for snapshot in snapshots if snapshot is not None and snapshot[1] == owner] + has_next_page = start_index + len(task_ids) < total_size + next_page_token = encode_page_token(task_ids[-1]) if task_ids and has_next_page else None + return page, total_size, next_page_token + + async def _list_filtered( + self, + params: ListTasksRequest, + context: ServerCallContext, + page_size: int, + decode_page_token: Callable[[str], str], + encode_page_token: Callable[[str], str], + ) -> tuple[builtins.list[Task], int, str | None]: + """Apply filters in bounded update-index batches with exact pagination.""" + owner = self.owner_id_for_context(context) + updates_key = self._updates_key(owner) + indexed_size = int(await self.redis.zcard(updates_key)) + start_rank = 0 + if params.page_token: + start_task_id = decode_page_token(params.page_token) + rank = await self.redis.zrevrank(updates_key, start_task_id) + if rank is None: + msg = f"Invalid page token: {params.page_token}" + raise InvalidParamsError(msg) + start_rank = int(rank) + 1 + + page: builtins.list[Task] = [] + matching_after_page = False + total_size = 0 + batch_size = settings.a2a_list_scan_batch_size + timestamp_after = params.status_timestamp_after if params.HasField("status_timestamp_after") else None + for offset in range(0, indexed_size, batch_size): + task_ids = await self.redis.zrevrange( + updates_key, + offset, + min(indexed_size - 1, offset + batch_size - 1), + ) + decoded_ids = [self._decode_value(task_id) for task_id in task_ids] + snapshots = await self._load_snapshots(decoded_ids) + for rank, snapshot in enumerate(snapshots, start=offset): + if snapshot is None or snapshot[1] != owner: + continue + task = snapshot[0] + if not task_matches_filters(task, params, timestamp_after): + continue + total_size += 1 + if rank < start_rank: + continue + if len(page) < page_size: + page.append(task) + else: + matching_after_page = True + + has_next_page = bool(page) and matching_after_page + next_page_token = encode_page_token(page[-1].id) if page and has_next_page else None + return page, total_size, next_page_token + + async def delete(self, task_id: str, context: ServerCallContext) -> None: + owner = self.owner_id_for_context(context) + if await self._delete_payload(task_id, owner=owner): + self._forget_task_versions(task_id) + + +class RedisTaskStoreProvider(TaskStoreProvider): + """ + Create one Redis task store per exposed agent. + + The provider intentionally does not import ``redis`` until it needs to + create a client, so importing Hayhooks remains possible without the A2A + optional dependencies installed. + """ + + def __init__( # noqa: PLR0913 - provider options mirror Redis connection settings + self, + redis_url: str | None = None, + *, + key_prefix: str | None = None, + redis: Any | None = None, + owner_resolver: OwnerResolver = _default_owner_resolver, + close_redis: bool = True, + terminal_ttl_seconds: int | None = None, + socket_timeout: float | None = None, + socket_connect_timeout: float | None = None, + health_check_interval: int | None = None, + ) -> None: + redis_url = redis_url or settings.a2a_redis_url + key_prefix = key_prefix or settings.a2a_redis_key_prefix + self.key_prefix = key_prefix.rstrip(":") + self.owner_resolver = owner_resolver + self.stores: dict[str, RedisTaskStore] = {} + self._close_redis = close_redis + self.terminal_ttl_seconds = terminal_ttl_seconds or settings.a2a_terminal_task_ttl_seconds + self.socket_timeout = socket_timeout if socket_timeout is not None else settings.a2a_redis_socket_timeout + self.socket_connect_timeout = ( + socket_connect_timeout if socket_connect_timeout is not None else settings.a2a_redis_socket_connect_timeout + ) + self.health_check_interval = ( + health_check_interval if health_check_interval is not None else settings.a2a_redis_health_check_interval + ) + if redis is None: + try: + from redis.asyncio import Redis + except ImportError as error: # pragma: no cover - depends on optional extras + msg = 'Redis task storage requires the A2A extra. Install with `pip install "hayhooks[a2a]"`.' + raise ImportError(msg) from error + redis = Redis.from_url( + redis_url, + decode_responses=False, + socket_timeout=self.socket_timeout, + socket_connect_timeout=self.socket_connect_timeout, + health_check_interval=self.health_check_interval, + ) + self.redis = redis + + def create_task_store(self, agent_name: str) -> RedisTaskStore: + if agent_name not in self.stores: + self.stores[agent_name] = RedisTaskStore( + self.redis, + agent_name, + key_prefix=self.key_prefix, + owner_resolver=self.owner_resolver, + terminal_ttl_seconds=self.terminal_ttl_seconds, + ) + return self.stores[agent_name] + + async def initialize(self) -> None: + """Fail A2A startup when its authoritative task store is unavailable.""" + await self.redis.ping() + + async def health(self) -> dict[str, Any]: + try: + await self.redis.ping() + except Exception as error: + return { + "healthy": False, + "provider": type(self).__name__, + "error": type(error).__name__, + } + return {"healthy": True, "provider": type(self).__name__} + + async def close(self) -> None: + if self._close_redis: + await self.redis.aclose() + + +__all__ = ["RedisTaskStore", "RedisTaskStoreProvider"] diff --git a/src/hayhooks/server/a2a/runtime.py b/src/hayhooks/server/a2a/runtime.py new file mode 100644 index 00000000..c678d9f1 --- /dev/null +++ b/src/hayhooks/server/a2a/runtime.py @@ -0,0 +1,259 @@ +import asyncio +from contextlib import suppress +from typing import Any + +from hayhooks.a2a import TaskStoreProvider +from hayhooks.durable.runtime import durable_runtime +from hayhooks.server.a2a.durable_executor import DurableAgentExecutor +from hayhooks.server.a2a.imports import ( + InMemoryTaskStore, + InvalidParamsError, + RequestContext, + RequestContextBuilder, + SimpleRequestContextBuilder, + TaskStore, +) +from hayhooks.server.logger import log +from hayhooks.settings import settings + + +class TaskAwareRequestContextBuilder(RequestContextBuilder): + """Infer and validate context identity for messages that continue a task.""" + + def __init__(self, task_store: "TaskStore") -> None: + self._task_store = task_store + self._delegate = SimpleRequestContextBuilder( + should_populate_referred_tasks=False, + task_store=task_store, + ) + + async def build( + self, + context: Any, + params: Any | None = None, + task_id: str | None = None, + context_id: str | None = None, + task: Any | None = None, + ) -> "RequestContext": + """ + Build a request context while preserving the context of an existing task. + + A2A permits a follow-up message to provide only ``task_id``. In that + case the server must infer ``context_id`` from the stored task. If the + client provides both identifiers, they must refer to the same task. + """ + existing_task = task + if task_id is not None and existing_task is None: + existing_task = await self._task_store.get(task_id, context) + + if existing_task is not None: + if context_id is not None and context_id != existing_task.context_id: + msg = ( + f"Message context_id '{context_id}' does not match context_id " + f"'{existing_task.context_id}' for task '{task_id}'" + ) + raise InvalidParamsError(message=msg) + context_id = existing_task.context_id + + # Preserve the SDK's concurrency behavior: ActiveTask refreshes + # current_task immediately before invoking the executor, so do not pass + # the independently loaded copy into the request context here. + return await self._delegate.build( + context=context, + params=params, + task_id=task_id, + context_id=context_id, + task=task, + ) + + +class InMemoryTaskStoreProvider(TaskStoreProvider): + """Provide an independent in-memory task store for each exposed agent.""" + + def create_task_store(self, agent_name: str) -> "TaskStore": # noqa: ARG002 + return InMemoryTaskStore() + + +def create_task_store_provider( + *, + backend: str = "auto", + redis_url: str | None = None, + redis_key_prefix: str | None = None, + redis: Any | None = None, + close_redis: bool = True, +) -> TaskStoreProvider: + """Create a built-in A2A task-store provider.""" + if backend == "redis": + from hayhooks.server.a2a.redis_task_store import RedisTaskStoreProvider + + return RedisTaskStoreProvider( + redis_url=redis_url, + key_prefix=redis_key_prefix, + redis=redis, + close_redis=close_redis, + ) + if backend in {"auto", "memory"}: + return InMemoryTaskStoreProvider() + msg = f"Unsupported A2A task-store backend '{backend}'; expected 'auto', 'memory', or 'redis'" + raise ValueError(msg) + + +class A2ARuntime: + """Owns A2A server resources shared by mounted agents.""" + + def __init__( + self, + task_store_provider: TaskStoreProvider | None = None, + ) -> None: + self.task_store_provider = task_store_provider or InMemoryTaskStoreProvider() + self._executors: list[DurableAgentExecutor] = [] + self._started_executors: list[DurableAgentExecutor] = [] + self._task_stores: list[TaskStore] = [] + self._maintenance_task: asyncio.Task[None] | None = None + self._started = False + + def register_agent_executor(self, executor: Any) -> None: + if isinstance(executor, DurableAgentExecutor): + self._executors.append(executor) + + async def start(self) -> None: + """Start lifecycle-aware executors after the application event loop is available.""" + try: + await self.task_store_provider.initialize() + for executor in self._executors: + self._started_executors.append(executor) + await executor.start() + if any(callable(getattr(store, "cleanup_expired_tasks", None)) for store in self._task_stores): + self._maintenance_task = asyncio.create_task( + self._maintain_task_stores(), + name="a2a-task-store-maintenance", + ) + self._started = True + except BaseException: + self._started = False + await self._close_executors() + raise + + def create_task_store(self, agent_name: str) -> "TaskStore": + task_store = self.task_store_provider.create_task_store(agent_name) + if not isinstance(task_store, TaskStore): + msg = ( + f"Task store provider {type(self.task_store_provider).__name__} returned " + f"{type(task_store).__name__} for agent '{agent_name}'; expected a2a.server.tasks.TaskStore" + ) + raise TypeError(msg) + if task_store not in self._task_stores: + self._task_stores.append(task_store) + return task_store + + async def close(self) -> None: + """Stop executor work before releasing shared task-store resources.""" + self._started = False + try: + if self._maintenance_task is not None: + self._maintenance_task.cancel() + with suppress(asyncio.CancelledError): + await self._maintenance_task + self._maintenance_task = None + await self._close_executors() + finally: + await self.task_store_provider.close() + + async def health(self) -> dict[str, Any]: + """Report operational readiness without changing the A2A protocol surface.""" + provider = await self._provider_health() + executor_health = { + f"{type(executor).__name__}:{index}": executor.health() for index, executor in enumerate(self._executors) + } + maintenance = self._maintenance_health() + components: dict[str, Any] = { + "task_store": provider, + "executors": executor_health, + "maintenance": maintenance, + } + if self._executors: + components["durable_execution"] = await self._durable_health() + + healthy = self._started and bool(provider.get("healthy", False)) and bool(maintenance["healthy"]) + healthy = healthy and all(bool(value.get("healthy", False)) for value in executor_health.values()) + durable = components.get("durable_execution") + if isinstance(durable, dict): + healthy = healthy and bool(durable.get("healthy", False)) + return { + "healthy": healthy, + "started": self._started, + "components": components, + } + + async def _provider_health(self) -> dict[str, Any]: + try: + value = await self.task_store_provider.health() + if isinstance(value, dict): + return value + return { + "healthy": False, + "provider": type(self.task_store_provider).__name__, + "error": "InvalidHealthPayload", + } + except asyncio.CancelledError: + raise + except Exception as error: + return { + "healthy": False, + "provider": type(self.task_store_provider).__name__, + "error": type(error).__name__, + } + + @staticmethod + async def _durable_health() -> dict[str, Any]: + try: + return await durable_runtime.health() + except asyncio.CancelledError: + raise + except Exception as error: + return {"healthy": False, "error": type(error).__name__} + + def _maintenance_health(self) -> dict[str, Any]: + task = self._maintenance_task + health: dict[str, Any] = { + "healthy": task is None or not task.done(), + "enabled": task is not None, + } + if task is None or not task.done(): + return health + if task.cancelled(): + health["error"] = "CancelledError" + elif error := task.exception(): + health["error"] = type(error).__name__ + return health + + async def _maintain_task_stores(self) -> None: + """Expire terminal tasks even when no later A2A request arrives.""" + intervals = [ + max(1.0, min(60.0, float(getattr(store, "terminal_ttl_seconds", 60)) / 10)) + for store in self._task_stores + if callable(getattr(store, "cleanup_expired_tasks", None)) + ] + interval = min(intervals, default=60.0) + while True: + await asyncio.sleep(interval) + for store in self._task_stores: + cleanup = getattr(store, "cleanup_expired_tasks", None) + if not callable(cleanup): + continue + try: + await cleanup(limit=settings.a2a_list_scan_batch_size) + except Exception as error: + log.opt(exception=error).warning("A2A terminal-task cleanup failed: {}", error) + + async def _close_executors(self) -> None: + for executor in reversed(self._started_executors): + try: + await executor.close() + except Exception as error: + log.opt(exception=True).warning( + "Error closing A2A executor lifecycle '{}': {}", + type(executor).__name__, + error, + ) + self._started_executors.clear() diff --git a/src/hayhooks/server/utils/a2a_utils.py b/src/hayhooks/server/utils/a2a_utils.py deleted file mode 100644 index 04884c89..00000000 --- a/src/hayhooks/server/utils/a2a_utils.py +++ /dev/null @@ -1,399 +0,0 @@ -import traceback -import uuid -from collections.abc import AsyncGenerator, AsyncIterator, Iterator -from typing import Any - -from fastapi.concurrency import iterate_in_threadpool, run_in_threadpool -from haystack.dataclasses import StreamingChunk -from haystack.lazy_imports import LazyImport -from starlette.applications import Starlette -from starlette.requests import Request -from starlette.responses import JSONResponse -from starlette.routing import Mount, Route - -from hayhooks.server.logger import log -from hayhooks.server.pipelines.registry import registry -from hayhooks.server.tracing import ( - SPAN_A2A_RUN_AGENT, - build_trace_tags, - configure_tracing, - instrument_starlette_app, - trace_operation, -) -from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper -from hayhooks.settings import settings - -# Lazily import A2A modules so the optional dependency is only required when used -with LazyImport("Run 'pip install \"hayhooks[a2a]\"' to install A2A support.") as a2a_import: - from a2a.helpers import get_message_text, new_task_from_user_message, new_text_part - from a2a.server.agent_execution import AgentExecutor, RequestContext - from a2a.server.events import EventQueue - from a2a.server.request_handlers import DefaultRequestHandler - from a2a.server.routes import create_agent_card_routes, create_jsonrpc_routes - from a2a.server.tasks import InMemoryTaskStore, TaskUpdater - from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill, Role - -RESPONSE_ARTIFACT_NAME = "response" - -# Path reserved by the A2A app itself; a pipeline with this name cannot be mounted -_RESERVED_PATHS = frozenset({"status"}) - - -def get_a2a_base_url() -> str: - """Base URL advertised in agent cards, without trailing slash.""" - base_url = settings.a2a_external_url or f"http://{settings.a2a_host}:{settings.a2a_port}" - return base_url.rstrip("/") - - -def is_a2a_exposable(pipeline_name: str) -> bool: - """ - Whether a deployed pipeline can be exposed as an A2A agent. - - A pipeline is exposable when it implements ``run_chat_completion`` or - ``run_chat_completion_async`` and does not set ``skip_a2a = True``. - """ - pipeline_wrapper = registry.get(pipeline_name) - if pipeline_wrapper is None: - return False - - metadata = registry.get_metadata(name=pipeline_name) or {} - if metadata.get("skip_a2a"): - log.debug("Skipping pipeline '{}': skip_a2a is set", pipeline_name) - return False - - exposable = ( - pipeline_wrapper._is_run_chat_completion_implemented - or pipeline_wrapper._is_run_chat_completion_async_implemented - ) - if not exposable: - log.debug("Skipping pipeline '{}': no chat completion method implemented", pipeline_name) - return exposable - - -def create_agent_card(pipeline_name: str, base_url: str) -> "AgentCard": - """ - Build an A2A agent card for a deployed pipeline. - - Card fields are derived from the pipeline's registry metadata and can be - overridden via the wrapper's ``a2a_card`` class attribute. - """ - a2a_import.check() - - metadata = registry.get_metadata(name=pipeline_name) or {} - overrides = metadata.get("a2a_card") or {} - - name = overrides.get("name") or pipeline_name - description = ( - overrides.get("description") - or metadata.get("description") - or f"Haystack pipeline '{pipeline_name}' deployed with Hayhooks" - ) - version = overrides.get("version") or "1.0.0" - agent_url = f"{base_url.rstrip('/')}/{pipeline_name}/" - - skills_spec: list[dict[str, Any]] = overrides.get("skills") or [ - {"id": pipeline_name, "name": name, "description": description, "tags": ["haystack", "hayhooks"]} - ] - skills = [ - AgentSkill( - id=skill.get("id", pipeline_name), - name=skill.get("name", name), - description=skill.get("description", description), - tags=list(skill.get("tags", [])), - examples=list(skill.get("examples", [])), - ) - for skill in skills_spec - ] - - log.debug( - "Built A2A agent card for pipeline '{}' with name='{}', url='{}', skills={}", - pipeline_name, - name, - agent_url, - [skill.id for skill in skills], - ) - - return AgentCard( - name=name, - description=description, - version=version, - default_input_modes=["text/plain"], - default_output_modes=["text/plain"], - capabilities=AgentCapabilities(streaming=True), - supported_interfaces=[AgentInterface(protocol_binding="JSONRPC", url=agent_url)], - skills=skills, - ) - - -def _stream_item_to_text(item: Any) -> str | None: - """ - Map a chat-completion stream item to response text. - - Returns None for items that carry no response text (UI events, empty chunks). - """ - if isinstance(item, StreamingChunk): - return item.content or None - if isinstance(item, str): - return item or None - if isinstance(item, bytes): - return item.decode("utf-8", errors="replace") or None - # PipelineEvent / dict items are UI-oriented events, not part of the text response - return None - - -def _build_openai_messages(context: "RequestContext") -> list[dict]: - """Map the A2A task history and current message to OpenAI-format messages.""" - messages: list[dict] = [] - - history = list(context.current_task.history) if context.current_task else [] - history_message_ids = {message.message_id for message in history} - for message in history: - text = get_message_text(message) - if text: - role = "assistant" if message.role == Role.ROLE_AGENT else "user" - messages.append({"role": role, "content": text}) - - # Append the incoming message unless it is already part of the task history - current = context.message - if current is not None and current.message_id not in history_message_ids: - current_text = get_message_text(current) - if current_text: - messages.append({"role": "user", "content": current_text}) - - log.debug( - "Mapped A2A request context to {} OpenAI message(s): history={}, current_message={}", - len(messages), - len(history), - current is not None, - ) - return messages - - -async def _run_chat_completion(pipeline_name: str, context: "RequestContext") -> Any: - """Run the pipeline's chat completion method (async preferred, sync via threadpool).""" - pipeline_wrapper: BasePipelineWrapper | None = registry.get(pipeline_name) - if pipeline_wrapper is None: - msg = f"Pipeline '{pipeline_name}' not found" - raise ValueError(msg) - - messages = _build_openai_messages(context) - - if pipeline_wrapper._is_run_chat_completion_async_implemented: - log.debug("Running pipeline '{}' as A2A agent via async chat completion", pipeline_name) - return await pipeline_wrapper.run_chat_completion_async(model=pipeline_name, messages=messages, body={}) - log.debug("Running pipeline '{}' as A2A agent via sync chat completion in threadpool", pipeline_name) - return await run_in_threadpool( - pipeline_wrapper.run_chat_completion, model=pipeline_name, messages=messages, body={} - ) - - -async def _iter_text_chunks(result: Any) -> AsyncGenerator[str, None]: - """ - Normalize a chat completion result (str or sync/async iterator) into text chunks. - - Raises ValueError for results outside the ``run_chat_completion`` contract - (e.g. None) so wrapper bugs surface as failed tasks instead of a completed - task whose response text is ``"None"``. - """ - if isinstance(result, str): - yield result - elif isinstance(result, AsyncIterator): - async for item in result: - text = _stream_item_to_text(item) - if text is not None: - yield text - elif isinstance(result, Iterator): - # Drain in a threadpool to keep the event loop free - async for item in iterate_in_threadpool(result): - text = _stream_item_to_text(item) - if text is not None: - yield text - else: - msg = f"run_chat_completion returned unsupported type '{type(result).__name__}'; expected str or generator" - raise ValueError(msg) - - -async def _stream_result_as_artifact(result: Any, updater: "TaskUpdater") -> None: - """ - Emit the chat completion result as a single ``response`` artifact. - - Generator results are streamed incrementally as artifact chunks - (``append=True``) so SSE clients receive text as it is produced; - the task manager aggregates chunks for non-streaming clients. - The last chunk is emitted with ``last_chunk=True``, so chunks are - held back one iteration until the end of the stream is known. - """ - artifact_id = str(uuid.uuid4()) - first = True - pending: str | None = None - - async def emit(text: str, *, last: bool) -> None: - nonlocal first - log.debug( - "Emitting A2A artifact chunk: artifact_id={}, append={}, last={}, chars={}", - artifact_id, - not first, - last, - len(text), - ) - await updater.add_artifact( - [new_text_part(text)], - artifact_id=artifact_id, - name=RESPONSE_ARTIFACT_NAME, - append=not first, - last_chunk=last, - ) - first = False - - async for text in _iter_text_chunks(result): - if pending is not None: - await emit(pending, last=False) - pending = text - await emit(pending if pending is not None else "", last=True) - - -async def _execute_agent_task(pipeline_name: str, context: "RequestContext", event_queue: "EventQueue") -> None: - """ - Run a pipeline's chat completion as an A2A task. - - Emits the event sequence required by the A2A spec: the Task first, - then a working status, artifact chunk(s), and a terminal state. - """ - if context.current_task is not None: - task = context.current_task - log.debug("Continuing A2A task '{}' for pipeline '{}'", task.id, pipeline_name) - elif context.message is not None: - task = new_task_from_user_message(context.message) - await event_queue.enqueue_event(task) - log.debug("Created A2A task '{}' for pipeline '{}'", task.id, pipeline_name) - else: - msg = "A2A request has neither a current task nor a message" - raise ValueError(msg) - - updater = TaskUpdater(event_queue, task.id, task.context_id) - await updater.start_work() - - with trace_operation( - SPAN_A2A_RUN_AGENT, - tags=build_trace_tags({"hayhooks.transport": "a2a", "hayhooks.pipeline.name": pipeline_name}), - ): - try: - result = await _run_chat_completion(pipeline_name, context) - await _stream_result_as_artifact(result, updater) - except Exception as exc: - msg = f"Error running pipeline '{pipeline_name}' as A2A agent: {exc}" - if settings.show_tracebacks: - msg += f"\n{traceback.format_exc()}" - log.opt(exception=True).error(msg) - await updater.failed(message=updater.new_agent_message([new_text_part(msg)])) - return - - await updater.complete() - log.debug("Completed A2A task '{}' for pipeline '{}'", task.id, pipeline_name) - - -def create_agent_executor(pipeline_name: str) -> "AgentExecutor": - """ - Create an ``AgentExecutor`` bridging A2A requests to the given pipeline. - - The class is defined inside this factory (instead of at module level) so - the module stays importable when the optional ``a2a-sdk`` dependency is - not installed. - """ - a2a_import.check() - - class HayhooksAgentExecutor(AgentExecutor): - """Runs a deployed pipeline's chat completion method as an A2A task.""" - - def __init__(self, name: str) -> None: - self.pipeline_name = name - - async def execute(self, context: "RequestContext", event_queue: "EventQueue") -> None: - await _execute_agent_task(self.pipeline_name, context, event_queue) - - async def cancel(self, context: "RequestContext", event_queue: "EventQueue") -> None: - # Best-effort: Hayhooks has no pipeline interruption primitive - task = context.current_task - if task is not None: - await TaskUpdater(event_queue, task.id, task.context_id).cancel() - - return HayhooksAgentExecutor(pipeline_name) - - -def _create_agent_mount(pipeline_name: str, base_url: str) -> Mount: - card = create_agent_card(pipeline_name, base_url) - request_handler = DefaultRequestHandler( - agent_executor=create_agent_executor(pipeline_name), - task_store=InMemoryTaskStore(), - agent_card=card, - ) - routes = [ - *create_agent_card_routes(card), - *create_jsonrpc_routes(request_handler, rpc_url="/", enable_v0_3_compat=settings.a2a_v0_3_compat), - ] - log.debug( - "Created A2A mount for pipeline '{}' at '/{}' with v0.3_compat={}", - pipeline_name, - pipeline_name, - settings.a2a_v0_3_compat, - ) - return Mount(f"/{pipeline_name}", routes=routes) - - -def create_a2a_app(*, base_url: str | None = None, debug: bool = False) -> Starlette: - """ - Create a Starlette app exposing deployed pipelines as A2A agents. - - Each exposable pipeline is mounted under ``/{pipeline_name}/`` with its - agent card at ``/{pipeline_name}/.well-known/agent-card.json`` and the - JSON-RPC binding at ``POST /{pipeline_name}/``. - - NOTE: mounts are built from the registry at startup; pipelines deployed or - undeployed at runtime require a restart to be reflected. - """ - a2a_import.check() - - base_url = (base_url or get_a2a_base_url()).rstrip("/") - if "//0.0.0.0" in base_url or "//[::]" in base_url: - log.warning( - "Agent cards will advertise the wildcard bind address ({}) which remote clients cannot connect to. " - "Set HAYHOOKS_A2A_EXTERNAL_URL (or --external-url) to the server's reachable base URL.", - base_url, - ) - - agent_names: list[str] = [] - mounts: list[Mount] = [] - for pipeline_name in registry.get_names(): - if not is_a2a_exposable(pipeline_name): - continue - - if pipeline_name in _RESERVED_PATHS: - log.warning("Skipping pipeline '{}': the path is reserved by the A2A server", pipeline_name) - continue - - # One failing agent card (e.g. a malformed a2a_card override) must not - # take down the other agents - try: - mounts.append(_create_agent_mount(pipeline_name, base_url)) - except Exception as e: - log.opt(exception=True).warning("Skipping pipeline '{}': failed to build A2A agent: {}", pipeline_name, e) - continue - - agent_names.append(pipeline_name) - log.info("Exposing pipeline '{}' as A2A agent at {}/{}/", pipeline_name, base_url, pipeline_name) - - if not agent_names: - log.warning( - "No pipelines exposable as A2A agents. " - "A pipeline must implement run_chat_completion or run_chat_completion_async." - ) - - async def handle_status(request: Request) -> JSONResponse: # noqa: ARG001 - return JSONResponse({"status": "ok", "agents": agent_names}) - - app = Starlette(debug=debug, routes=[Route("/status", endpoint=handle_status), *mounts]) - log.debug("Created A2A Starlette app with {} mounted agent(s): {}", len(agent_names), agent_names) - - configure_tracing() - instrument_starlette_app(app) - return app diff --git a/tests/test_a2a.py b/tests/test_a2a.py index 86011387..7a5bbdaa 100644 --- a/tests/test_a2a.py +++ b/tests/test_a2a.py @@ -6,18 +6,12 @@ from haystack.dataclasses import StreamingChunk from hayhooks.events import PipelineEvent +from hayhooks.server.a2a.cards import create_agent_card, get_a2a_base_url, is_a2a_exposable +from hayhooks.server.a2a.executor import RESPONSE_ARTIFACT_NAME, _stream_item_to_text, create_agent_executor +from hayhooks.server.a2a.messages import build_openai_messages +from hayhooks.server.logger import log from hayhooks.server.pipelines import registry from hayhooks.server.tracing import SPAN_A2A_RUN_AGENT -from hayhooks.server.utils.a2a_utils import ( - RESPONSE_ARTIFACT_NAME, - _build_openai_messages, - _execute_agent_task, - _stream_item_to_text, - create_agent_card, - create_agent_executor, - get_a2a_base_url, - is_a2a_exposable, -) from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper from hayhooks.server.utils.module_loader import _set_method_implementation_flags @@ -46,6 +40,12 @@ async def enqueue_event(self, event): self.events.append(event) +async def execute_agent_task(pipeline_name, context, event_queue): + wrapper = registry.get(pipeline_name) + assert wrapper is not None + await create_agent_executor(wrapper, pipeline_name).execute(context, event_queue) + + class AsyncChatWrapper(BasePipelineWrapper): def setup(self): self.pipeline = object() @@ -116,44 +116,33 @@ def get_artifact_events(events) -> list: # --- Exposure rules --- -def test_is_a2a_exposable_unknown_pipeline(): - assert not is_a2a_exposable("non_existent_pipeline") - - -def test_is_a2a_exposable_chat_async(): - register_wrapper("chat_agent", AsyncChatWrapper) - assert is_a2a_exposable("chat_agent") - - -def test_is_a2a_exposable_chat_sync(): - register_wrapper("sync_agent", SyncChatWrapper) - assert is_a2a_exposable("sync_agent") - - -def test_is_a2a_exposable_api_only(): - register_wrapper("api_only", ApiOnlyWrapper) - assert not is_a2a_exposable("api_only") - - -def test_is_a2a_exposable_skip_a2a(): - register_wrapper("chat_agent", AsyncChatWrapper, metadata={"skip_a2a": True}) - assert not is_a2a_exposable("chat_agent") +@pytest.mark.parametrize( + ("name", "wrapper", "metadata", "expected"), + [ + ("non_existent_pipeline", None, None, False), + ("chat_agent", AsyncChatWrapper, None, True), + ("sync_agent", SyncChatWrapper, None, True), + ("api_only", ApiOnlyWrapper, None, False), + ("skipped_agent", AsyncChatWrapper, {"skip_a2a": True}, False), + ], +) +def test_is_a2a_exposable(name, wrapper, metadata, expected): + if wrapper is not None: + register_wrapper(name, wrapper, metadata=metadata) + assert is_a2a_exposable(name) is expected # --- Base URL --- -def test_get_a2a_base_url_default(test_settings): - test_settings.a2a_external_url = "" - assert get_a2a_base_url() == f"http://{test_settings.a2a_host}:{test_settings.a2a_port}" - - -def test_get_a2a_base_url_external(test_settings): - test_settings.a2a_external_url = "https://agents.example.com/" - try: - assert get_a2a_base_url() == "https://agents.example.com" - finally: - test_settings.a2a_external_url = "" +@pytest.mark.parametrize( + ("external_url", "expected"), + [("", None), ("https://agents.example.com/", "https://agents.example.com")], +) +def test_get_a2a_base_url(test_settings, external_url, expected): + test_settings.a2a_external_url = external_url + expected = expected or f"http://{test_settings.a2a_host}:{test_settings.a2a_port}" + assert get_a2a_base_url() == expected # --- Agent card --- @@ -230,7 +219,7 @@ def test_stream_item_to_text(): def test_build_openai_messages_from_message_only(): context = make_context("what is the weather?") - assert _build_openai_messages(context) == [{"role": "user", "content": "what is the weather?"}] + assert build_openai_messages(context) == [{"role": "user", "content": "what is the weather?"}] def test_build_openai_messages_with_task_history(): @@ -242,7 +231,7 @@ def test_build_openai_messages_with_task_history(): task.history.append(new_text_message("first answer", role=Role.ROLE_AGENT)) context = make_context("second question", current_task=task) - assert _build_openai_messages(context) == [ + assert build_openai_messages(context) == [ {"role": "user", "content": "first question"}, {"role": "assistant", "content": "first answer"}, {"role": "user", "content": "second question"}, @@ -257,7 +246,7 @@ def test_build_openai_messages_deduplicates_current_message(): task = new_task_from_user_message(message) # history already contains the message context = SimpleNamespace(message=message, current_task=task) - assert _build_openai_messages(context) == [{"role": "user", "content": "hello"}] + assert build_openai_messages(context) == [{"role": "user", "content": "hello"}] def test_build_openai_messages_keeps_new_message_matching_history_text(): @@ -270,7 +259,7 @@ def test_build_openai_messages_keeps_new_message_matching_history_text(): task.history.append(new_text_message("yes", role=Role.ROLE_AGENT)) context = make_context("yes", current_task=task) - assert _build_openai_messages(context) == [ + assert build_openai_messages(context) == [ {"role": "user", "content": "continue?"}, {"role": "assistant", "content": "yes"}, {"role": "user", "content": "yes"}, @@ -287,7 +276,7 @@ async def test_execute_agent_task_string_result(): register_wrapper("sync_agent", SyncChatWrapper) queue = RecordingQueue() - await _execute_agent_task("sync_agent", make_context(), queue) + await execute_agent_task("sync_agent", make_context(), queue) assert isinstance(queue.events[0], Task) assert get_status_states(queue.events) == [TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_COMPLETED] @@ -306,18 +295,18 @@ async def test_execute_agent_task_streaming_result(): register_wrapper("chat_agent", AsyncChatWrapper) queue = RecordingQueue() - await _execute_agent_task("chat_agent", make_context("hi"), queue) + await execute_agent_task("chat_agent", make_context("hi"), queue) assert get_status_states(queue.events)[-1] == TaskState.TASK_STATE_COMPLETED artifact_events = get_artifact_events(queue.events) # PipelineEvent items are skipped, text chunks are streamed incrementally - assert len(artifact_events) == 3 + assert len(artifact_events) == 4 texts = [event.artifact.parts[0].text for event in artifact_events] - assert texts == ["Hello, ", "world", " (question: hi)"] - # All chunks belong to the same artifact; only the last one is marked last_chunk + assert texts == ["Hello, ", "world", " (question: hi)", ""] + # All chunks belong to the same artifact; an empty marker finalizes iterator output assert len({event.artifact.artifact_id for event in artifact_events}) == 1 - assert [event.last_chunk for event in artifact_events] == [False, False, True] + assert [event.last_chunk for event in artifact_events] == [False, False, False, True] assert artifact_events[0].append is False assert artifact_events[1].append is True @@ -329,7 +318,7 @@ async def test_execute_agent_task_error_sets_failed_state(): register_wrapper("failing_agent", FailingChatWrapper) queue = RecordingQueue() - await _execute_agent_task("failing_agent", make_context(), queue) + await execute_agent_task("failing_agent", make_context(), queue) states = get_status_states(queue.events) assert states[-1] == TaskState.TASK_STATE_FAILED @@ -351,33 +340,126 @@ def run_chat_completion(self, model: str, messages: list[dict], body: dict): register_wrapper("none_agent", NoneResultWrapper) queue = RecordingQueue() - await _execute_agent_task("none_agent", make_context(), queue) - - assert get_status_states(queue.events)[-1] == TaskState.TASK_STATE_FAILED - + await execute_agent_task("none_agent", make_context(), queue) -@pytest.mark.asyncio -async def test_execute_agent_task_unknown_pipeline_fails(): - from a2a.types import TaskState - - queue = RecordingQueue() - await _execute_agent_task("non_existent", make_context(), queue) assert get_status_states(queue.events)[-1] == TaskState.TASK_STATE_FAILED @pytest.mark.asyncio -async def test_execute_agent_task_emits_trace_span(recording_tracer): +async def test_execute_agent_task_emits_trace_and_safe_lifecycle_logs(recording_tracer): register_wrapper("sync_agent", SyncChatWrapper) - await _execute_agent_task("sync_agent", make_context(), RecordingQueue()) + records = [] + sink = log.add(lambda message: records.append(message.record), level="DEBUG") + try: + await execute_agent_task("sync_agent", make_context("private message"), RecordingQueue()) + finally: + log.remove(sink) spans = [span for span in recording_tracer.spans if span.operation_name == SPAN_A2A_RUN_AGENT] assert spans assert spans[-1].tags["hayhooks.pipeline.name"] == "sync_agent" assert spans[-1].tags["hayhooks.transport"] == "a2a" + lifecycle = [record for record in records if record["message"] in {"Started A2A task", "Completed A2A task"}] + assert [record["message"] for record in lifecycle] == ["Started A2A task", "Completed A2A task"] + assert all(record["extra"]["pipeline_name"] == "sync_agent" for record in lifecycle) + assert all(record["extra"]["task_id"] for record in lifecycle) + assert "private message" not in str(lifecycle) + + +def test_runtime_passes_agent_name_to_task_store_provider(): + from a2a.server.tasks import InMemoryTaskStore + + from hayhooks.a2a import TaskStoreProvider + from hayhooks.server.a2a.runtime import A2ARuntime + + class RecordingTaskStoreProvider(TaskStoreProvider): + def __init__(self): + self.agent_names = [] + + def create_task_store(self, agent_name): + self.agent_names.append(agent_name) + return InMemoryTaskStore() + + provider = RecordingTaskStoreProvider() + runtime = A2ARuntime(task_store_provider=provider) + + first_store = runtime.create_task_store("first_agent") + second_store = runtime.create_task_store("second_agent") + + assert isinstance(first_store, InMemoryTaskStore) + assert isinstance(second_store, InMemoryTaskStore) + assert first_store is not second_store + assert provider.agent_names == ["first_agent", "second_agent"] + + +def test_runtime_rejects_invalid_task_store_from_provider(): + from hayhooks.a2a import TaskStoreProvider + from hayhooks.server.a2a.runtime import A2ARuntime + + class InvalidTaskStoreProvider(TaskStoreProvider): + def create_task_store(self, _agent_name): + return object() + + runtime = A2ARuntime(task_store_provider=InvalidTaskStoreProvider()) + + with pytest.raises(TypeError, match=r"InvalidTaskStoreProvider.*invalid_agent"): + runtime.create_task_store("invalid_agent") + + +async def test_runtime_closes_task_store_provider(): + from a2a.server.tasks import InMemoryTaskStore + + from hayhooks.a2a import TaskStoreProvider + from hayhooks.server.a2a.runtime import A2ARuntime + + class CloseableTaskStoreProvider(TaskStoreProvider): + def __init__(self): + self.closed = False + + def create_task_store(self, _agent_name): + return InMemoryTaskStore() + + async def close(self): + self.closed = True + + provider = CloseableTaskStoreProvider() + + await A2ARuntime(task_store_provider=provider).close() + + assert provider.closed + + +@pytest.mark.parametrize( + ("health", "expected_error"), + [ + ({"healthy": False, "provider": "TestProvider", "error": "ConnectionError"}, "ConnectionError"), + (True, "InvalidHealthPayload"), + ], +) +def test_a2a_status_returns_503_for_unhealthy_task_store(health, expected_error): + from a2a.server.tasks import InMemoryTaskStore + from starlette.testclient import TestClient + + from hayhooks.a2a import TaskStoreProvider + from hayhooks.server.a2a.app import create_a2a_app + from hayhooks.server.a2a.runtime import A2ARuntime + + class TestProvider(TaskStoreProvider): + def create_task_store(self, _agent_name): + return InMemoryTaskStore() + + async def health(self): + return health + + register_wrapper("chat_agent", AsyncChatWrapper) + app = create_a2a_app( + base_url="http://test:1418", + runtime=A2ARuntime(task_store_provider=TestProvider()), + ) + with TestClient(app) as client: + response = client.get("/status") -def test_create_agent_executor(): - executor = create_agent_executor("some_pipeline") - assert executor.pipeline_name == "some_pipeline" - assert hasattr(executor, "execute") - assert hasattr(executor, "cancel") + assert response.status_code == 503 + assert response.json()["status"] == "unavailable" + assert response.json()["components"]["task_store"]["error"] == expected_error diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py new file mode 100644 index 00000000..3fdc8060 --- /dev/null +++ b/tests/test_durable_a2a.py @@ -0,0 +1,414 @@ +import asyncio +import importlib.metadata +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import TimeoutError as FutureTimeoutError +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi.testclient import TestClient + +from hayhooks.a2a import A2APipelineWrapper, TaskStoreProvider +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStatus, ExecutionStoreError +from hayhooks.durable.runtime import execution_id_for +from hayhooks.server.a2a.app import create_a2a_app +from hayhooks.server.a2a.durable_executor import DurableAgentExecutor, DurableTaskStore +from hayhooks.server.a2a.imports import TaskStore, new_task_from_user_message, new_text_part +from hayhooks.server.a2a.runtime import A2ARuntime +from hayhooks.server.pipelines.registry import registry +from hayhooks.settings import settings + +pytestmark = pytest.mark.skipif( + not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" +) + + +class _Deployment: + def __init__(self, status=ExecutionStatus.COMPLETED) -> None: + self.record = SimpleNamespace( + status=status, + progress=[], + result={"last_message": {"content": "recovered"}}, + error=None, + sequence=0, + ) + self.execution_id = None + self.submitted_payload = None + self.resume_update = None + self.cancel_requested = False + + async def start(self): + return None + + async def submit(self, payload, *, execution_id=None, owner_id=None): + self.submitted_payload = payload + self.execution_id = execution_id_for(owner_id, execution_id) if owner_id else execution_id + self.record.execution_id = self.execution_id + return True, self.record + + async def get(self, execution_id, **_kwargs): + if self.execution_id is not None and execution_id != self.execution_id: + raise KeyError(execution_id) + return self.record + + async def resume(self, execution_id, update, **_kwargs): + self.execution_id = execution_id + self.resume_update = update + self.record.status = ExecutionStatus.COMPLETED + self.record.result = {"last_message": {"content": "resumed"}} + self.record.sequence += 1 + return True + + async def request_cancel(self, _execution_id, **_kwargs): + self.cancel_requested = True + self.record.status = ExecutionStatus.CANCELED + self.record.sequence += 1 + return True + + +class _BlockingDeployment(_Deployment): + def __init__(self) -> None: + super().__init__() + self.submit_started = threading.Event() + self.allow_submit = threading.Event() + + async def submit(self, *args, **kwargs): + self.submit_started.set() + await asyncio.to_thread(self.allow_submit.wait) + return await super().submit(*args, **kwargs) + + +class _DurableHTTPWrapper(A2APipelineWrapper): + durable_revision = "durable-http-wrapper" + + def setup(self): + self.pipeline = object() + + +class _HTTPStore(TaskStore): + def __init__(self) -> None: + self.tasks = {} + + async def save(self, task, _context): + self.tasks[task.id] = task + + async def get(self, task_id, _context): + return self.tasks.get(task_id) + + async def list(self, _params, _context): + from a2a.types import ListTasksResponse + + return ListTasksResponse(tasks=list(self.tasks.values()), page_size=len(self.tasks), total_size=len(self.tasks)) + + async def delete(self, task_id, _context): + self.tasks.pop(task_id, None) + + +class _HTTPStoreProvider(TaskStoreProvider): + def __init__(self, store) -> None: + self.store = store + + def create_task_store(self, _agent_name): + return self.store + + +def _send_payload(text, *, task_id=None, return_immediately=False): + message = {"messageId": f"message-{text}", "role": "ROLE_USER", "parts": [{"text": text}]} + if task_id is not None: + message["taskId"] = task_id + params = {"message": message} + if return_immediately: + params["configuration"] = {"returnImmediately": True} + return {"jsonrpc": "2.0", "id": "send", "method": "SendMessage", "params": params} + + +def _get_payload(task_id): + return {"jsonrpc": "2.0", "id": "get", "method": "GetTask", "params": {"id": task_id}} + + +def _cancel_payload(task_id): + return {"jsonrpc": "2.0", "id": "cancel", "method": "CancelTask", "params": {"id": task_id}} + + +def _response_task(response): + result = response.json()["result"] + return result.get("task", result) + + +def _recoverable_task(): + from a2a.types import Message, Role + + return new_task_from_user_message( + Message( + message_id="message", + task_id="task", + context_id="context", + role=Role.ROLE_USER, + parts=[new_text_part("recover me")], + ) + ) + + +def _recovery_store(task, *, saved=True): + return SimpleNamespace( + recoverable_task_batch=AsyncMock(return_value=([(task, "owner", 1)], None)), + save_projection=AsyncMock(return_value=saved), + ) + + +def _http_app(store, deployment, monkeypatch): + wrapper = _DurableHTTPWrapper() + wrapper.setup() + registry.add("durable-agent", wrapper, metadata={"description": "durable agent"}) + monkeypatch.setattr("hayhooks.durable.runtime.durable_runtime.deployment", lambda *_args: deployment) + return create_a2a_app( + base_url="http://a2a-test:1418", + runtime=A2ARuntime(task_store_provider=_HTTPStoreProvider(store)), + ) + + +@pytest.fixture(autouse=True) +def _clean_registry(): + registry.clear() + yield + registry.clear() + + +@pytest.fixture +def http_store() -> _HTTPStore: + return _HTTPStore() + + +def test_a2a_http_reads_completion_from_durable_execution(monkeypatch, http_store) -> None: + app = _http_app(http_store, _Deployment(), monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + completed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) + + assert completed["status"]["state"] == "TASK_STATE_COMPLETED" + assert completed["artifacts"][-1]["name"] == "durable-result" + + +def test_a2a_http_waiting_task_resumes_with_only_the_follow_up(monkeypatch, http_store) -> None: + deployment = _Deployment(status=ExecutionStatus.WAITING) + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + waiting = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) + assert waiting["status"]["state"] == "TASK_STATE_INPUT_REQUIRED" + completed = _response_task( + client.post("/durable-agent/", json=_send_payload("follow up", task_id=waiting["id"])) + ) + + assert completed["status"]["state"] == "TASK_STATE_COMPLETED" + assert deployment.resume_update == { + "messages": [{"role": "user", "meta": {}, "name": None, "content": [{"text": "follow up"}]}] + } + + +def test_a2a_http_cancel_reaches_durable_execution(monkeypatch, http_store) -> None: + deployment = _Deployment(status=ExecutionStatus.RUNNING) + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) + canceled = _response_task(client.post("/durable-agent/", json=_cancel_payload(active["id"]))) + + assert deployment.cancel_requested + assert canceled["status"]["state"] == "TASK_STATE_CANCELED" + + +def test_return_immediately_waits_for_durable_submission(monkeypatch, http_store) -> None: + deployment = _BlockingDeployment() + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client, ThreadPoolExecutor() as pool: + response = pool.submit(client.post, "/durable-agent/", json=_send_payload("initial", return_immediately=True)) + try: + assert deployment.submit_started.wait(timeout=1) + assert http_store.tasks + with pytest.raises(FutureTimeoutError): + response.result(timeout=0.1) + finally: + deployment.allow_submit.set() + assert response.result(timeout=2).status_code == 200 + + +def test_returned_task_is_eventually_persisted_as_terminal(monkeypatch, http_store) -> None: + from a2a.types import TaskState + + deployment = _Deployment(status=ExecutionStatus.COMPLETED) + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) + deadline = time.monotonic() + 1 + while ( + http_store.tasks[active["id"]].status.state != TaskState.TASK_STATE_COMPLETED + and time.monotonic() < deadline + ): + time.sleep(0.01) + completed = _response_task(client.post("/durable-agent/", json=_get_payload(active["id"]))) + + assert completed["status"]["state"] == "TASK_STATE_COMPLETED" + assert http_store.tasks[active["id"]].status.state == TaskState.TASK_STATE_COMPLETED + + +async def test_expired_execution_preserves_retained_terminal_task(http_store) -> None: + from a2a.types import Task, TaskState + + task = Task(id="retained-terminal", context_id="context") + task.status.state = TaskState.TASK_STATE_COMPLETED + deployment = _Deployment() + deployment.get = AsyncMock(side_effect=KeyError("expired")) + + projected = await DurableTaskStore(http_store, deployment)._project(task, "owner", SimpleNamespace()) + + assert projected.status.state == TaskState.TASK_STATE_COMPLETED + + +async def test_missing_execution_preserves_a_task_awaiting_submission(http_store) -> None: + from a2a.types import TaskState + + task = _recoverable_task() + deployment = _Deployment() + deployment.get = AsyncMock(side_effect=KeyError("submission has not committed yet")) + + projected = await DurableTaskStore(http_store, deployment)._project(task, "owner", SimpleNamespace()) + + assert projected.status.state == TaskState.TASK_STATE_SUBMITTED + + +@pytest.mark.parametrize(("configured", "expected"), [(0.05, 0.1), (5.0, 5.0)]) +async def test_durable_a2a_polling_honors_its_configured_floor(monkeypatch, http_store, configured, expected) -> None: + executor = DurableAgentExecutor("agent", http_store, _Deployment(status=ExecutionStatus.RUNNING)) + delays = [] + + async def stop_after_one_poll(delay): + delays.append(delay) + executor._closed = True + + monkeypatch.setattr(settings, "durable_poll_interval", configured) + monkeypatch.setattr(asyncio, "sleep", stop_after_one_poll) + + await executor._wait_for_update("execution", "owner", object()) + + assert delays == [expected] + + +async def test_durable_a2a_polling_retries_store_outages(monkeypatch, http_store) -> None: + deployment = _Deployment() + deployment.get = AsyncMock( + side_effect=[ExecutionStoreError("offline"), ExecutionStoreError("still offline"), deployment.record] + ) + updater = SimpleNamespace(task_id="task", add_artifact=AsyncMock(), complete=AsyncMock()) + sleep = AsyncMock() + monkeypatch.setattr(asyncio, "sleep", sleep) + + await DurableAgentExecutor("agent", http_store, deployment)._wait_for_update("execution", "owner", updater) + + assert deployment.get.await_count == 3 + assert sleep.await_count == 2 + updater.complete.assert_awaited_once() + + +async def test_durable_a2a_cancellation_waits_past_input_required(monkeypatch, http_store) -> None: + waiting = _Deployment(status=ExecutionStatus.WAITING).record + canceled = _Deployment(status=ExecutionStatus.CANCELED).record + canceled.sequence = 1 + deployment = _Deployment() + deployment.get = AsyncMock(side_effect=[waiting, canceled]) + updater = SimpleNamespace( + task_id="task", + new_agent_message=lambda parts: parts, + requires_input=AsyncMock(), + cancel=AsyncMock(), + ) + sleep = AsyncMock() + monkeypatch.setattr(asyncio, "sleep", sleep) + + await DurableAgentExecutor("agent", http_store, deployment)._wait_for_update( + "execution", "owner", updater, terminal_only=True + ) + + updater.requires_input.assert_awaited_once() + updater.cancel.assert_awaited_once() + sleep.assert_awaited_once() + + +@pytest.mark.parametrize("failure", [ExecutionStoreError("offline"), ExecutionAdmissionError("test")]) +async def test_durable_a2a_submission_retries_transient_failures(monkeypatch, http_store, failure) -> None: + deployment = _Deployment() + deployment.submit = AsyncMock(side_effect=[failure, (True, deployment.record)]) + sleep = AsyncMock() + monkeypatch.setattr(asyncio, "sleep", sleep) + + record = await DurableAgentExecutor("agent", http_store, deployment)._submit("task", "owner", []) + + assert record is deployment.record + assert deployment.submit.await_count == 2 + sleep.assert_awaited_once() + + +def test_a2a_http_rejected_submission_is_persisted_as_failed(monkeypatch, http_store) -> None: + from a2a.types import TaskState + + deployment = _Deployment() + deployment.submit = AsyncMock(side_effect=ValueError("request is too large")) + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + failed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) + + assert failed["status"]["state"] == "TASK_STATE_FAILED" + assert http_store.tasks[failed["id"]].status.state == TaskState.TASK_STATE_FAILED + + +async def test_recovery_submits_a_persisted_task_without_an_execution(http_store) -> None: + from a2a.types import TaskState + + task = _recoverable_task() + recovery_store = _recovery_store(task) + deployment = _Deployment() + + async def missing_until_submitted(execution_id, **kwargs): + if deployment.execution_id is None: + raise KeyError(execution_id) + return await _Deployment.get(deployment, execution_id, **kwargs) + + deployment.get = missing_until_submitted + await DurableAgentExecutor("agent", http_store, deployment)._recover_tasks(recovery_store) + + assert deployment.execution_id is not None + assert deployment.submitted_payload["messages"][0]["content"] == [{"text": "recover me"}] + assert task.status.state == TaskState.TASK_STATE_COMPLETED + + +async def test_recovery_rejects_one_invalid_persisted_task_without_aborting(http_store) -> None: + from a2a.types import TaskState + + task = _recoverable_task() + recovery_store = _recovery_store(task) + deployment = _Deployment() + deployment.get = AsyncMock(side_effect=KeyError("missing")) + deployment.submit = AsyncMock(side_effect=ValueError("request is too large")) + + await DurableAgentExecutor("agent", http_store, deployment)._recover_tasks(recovery_store) + + assert task.status.state == TaskState.TASK_STATE_FAILED + recovery_store.save_projection.assert_awaited_once() + + +async def test_recovery_skips_a_projection_conflict(http_store) -> None: + from a2a.types import Task, TaskState + + task = Task(id="task", context_id="context") + task.status.state = TaskState.TASK_STATE_WORKING + recovery_store = _recovery_store(task, saved=False) + executor = DurableAgentExecutor("agent", http_store, _Deployment()) + + await executor._recover_tasks(recovery_store) + + recovery_store.save_projection.assert_awaited_once() diff --git a/tests/test_files/durable_a2a_process_recovery/pipeline_wrapper.py b/tests/test_files/durable_a2a_process_recovery/pipeline_wrapper.py new file mode 100644 index 00000000..b1e5e19d --- /dev/null +++ b/tests/test_files/durable_a2a_process_recovery/pipeline_wrapper.py @@ -0,0 +1,43 @@ +"""Deterministic durable A2A Agent used by the process-restart smoke test.""" + +from haystack import component +from haystack.components.agents import Agent +from haystack.components.agents.state import State +from haystack.dataclasses import ChatMessage +from haystack.hooks.from_function import FunctionHook + +from hayhooks import A2APipelineWrapper, current_durable_context + + +def require_approval(state: State) -> None: + context = current_durable_context() + if context is None or context.state.get("approval_requested"): + return + context.state["approval_requested"] = True + context.suspend_sync({"kind": "approval", "message": "Approve this task"}) + + +async def require_approval_async(state: State) -> None: + context = current_durable_context() + if context is None or context.state.get("approval_requested"): + return + context.state["approval_requested"] = True + await context.suspend({"kind": "approval", "message": "Approve this task"}) + + +@component +class FakeChatGenerator: + @component.output_types(replies=list[ChatMessage]) + def run(self, messages: list[ChatMessage], tools=None): + return {"replies": [ChatMessage.from_assistant("approved and complete")]} + + +class PipelineWrapper(A2APipelineWrapper): + durable_revision = "durable-a2a-process-recovery" + + def setup(self) -> None: + self.pipeline = Agent( + chat_generator=FakeChatGenerator(), + tools=[], + hooks={"before_llm": [FunctionHook(function=require_approval, async_function=require_approval_async)]}, + ) diff --git a/tests/test_it_a2a_server.py b/tests/test_it_a2a_server.py index 6e72780a..438f6291 100644 --- a/tests/test_it_a2a_server.py +++ b/tests/test_it_a2a_server.py @@ -1,3 +1,4 @@ +import asyncio import importlib.util import json @@ -5,9 +6,11 @@ import pytest from anyio import Path +from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.pipelines import registry -from hayhooks.server.utils.a2a_utils import create_a2a_app +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper from hayhooks.server.utils.deploy_utils import deploy_pipeline_files +from hayhooks.server.utils.module_loader import _set_method_implementation_flags A2A_AVAILABLE = importlib.util.find_spec("a2a") is not None @@ -42,7 +45,14 @@ async def a2a_client(): app = create_a2a_app(base_url=BASE_URL) transport = httpx.ASGITransport(app=app) - async with httpx.AsyncClient(transport=transport, base_url=BASE_URL, headers={"A2A-Version": "1.0"}) as client: + async with ( + app.router.lifespan_context(app), + httpx.AsyncClient( + transport=transport, + base_url=BASE_URL, + headers={"A2A-Version": "1.0"}, + ) as client, + ): yield client @@ -55,6 +65,23 @@ def send_message_payload(text: str, method: str = "SendMessage") -> dict: } +def get_task_payload(task_id: str, method: str = "GetTask") -> dict: + return {"jsonrpc": "2.0", "id": "get-task", "method": method, "params": {"id": task_id}} + + +def cancel_task_payload(task_id: str) -> dict: + return {"jsonrpc": "2.0", "id": "cancel-task", "method": "CancelTask", "params": {"id": task_id}} + + +def extract_task(response_payload: dict) -> dict: + result = response_payload["result"] + return result.get("task", result) + + +def artifact_text(task: dict) -> str: + return "".join(part["text"] for artifact in task.get("artifacts", []) for part in artifact["parts"]) + + def send_message_v0_3_payload(text: str, method: str = "message/send") -> dict: return { "jsonrpc": "2.0", @@ -70,12 +97,53 @@ def send_message_v0_3_payload(text: str, method: str = "message/send") -> dict: } +class ControlledLongRunningWrapper(BasePipelineWrapper): + emit_progress_chunk = True + + def setup(self): + self.pipeline = object() + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.cancelled = asyncio.Event() + + async def run_chat_completion_async(self, model: str, messages: list[dict], body: dict): + async def generator(): + self.entered.set() + if self.emit_progress_chunk: + yield "progress " + try: + await self.release.wait() + except asyncio.CancelledError: + self.cancelled.set() + raise + yield "done" + + return generator() + + +def register_test_wrapper(name: str, wrapper_cls: type[BasePipelineWrapper]) -> BasePipelineWrapper: + wrapper = wrapper_cls() + wrapper.setup() + _set_method_implementation_flags(wrapper) + registry.add( + name, + wrapper, + metadata={"description": f"{name} description", "skip_a2a": wrapper.skip_a2a, "a2a_card": wrapper.a2a_card}, + ) + return wrapper + + @pytest.mark.asyncio async def test_status_lists_exposed_agents(a2a_client): response = await a2a_client.get("/status") assert response.status_code == 200 # api_only has no chat completion method, so it must not be listed - assert response.json() == {"status": "ok", "agents": ["chat_agent", "async_chat_agent"]} + body = response.json() + assert body["status"] == "ok" + assert body["agents"] == ["chat_agent", "async_chat_agent"] + assert body["components"]["task_store"]["healthy"] + assert body["components"]["maintenance"]["healthy"] + assert response.headers["cache-control"] == "no-store" @pytest.mark.asyncio @@ -99,12 +167,6 @@ async def test_non_chat_pipeline_is_not_exposed(a2a_client): assert response.status_code == 404 -@pytest.mark.asyncio -async def test_unknown_agent_returns_404(a2a_client): - response = await a2a_client.get("/non_existent/.well-known/agent-card.json") - assert response.status_code == 404 - - @pytest.mark.asyncio async def test_send_message_returns_completed_task(a2a_client): response = await a2a_client.post("/chat_agent/", json=send_message_payload("What is Hayhooks?")) @@ -177,3 +239,154 @@ async def test_a2a_v1_method_with_v0_3_header_is_rejected(a2a_client): response = await a2a_client.post("/chat_agent/", json=send_message_payload("hi"), headers={"A2A-Version": "0.3"}) assert response.status_code == 200 assert response.json()["error"]["code"] == -32009 + + +@pytest.fixture +async def long_running_client(): + wrapper = register_test_wrapper("long_agent", ControlledLongRunningWrapper) + app = create_a2a_app(base_url=BASE_URL) + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url=BASE_URL, headers={"A2A-Version": "1.0"}) as client: + yield client, wrapper + + +async def poll_task_until_state(client: httpx.AsyncClient, task_id: str, expected_state: str) -> dict: + last_task = None + for _ in range(50): + response = await client.post("/long_agent/", json=get_task_payload(task_id)) + assert response.status_code == 200 + last_task = extract_task(response.json()) + if last_task["status"]["state"] == expected_state: + return last_task + await asyncio.sleep(0.01) + msg = f"Task {task_id} did not reach {expected_state}. Last task: {last_task}" + raise AssertionError(msg) + + +async def poll_task_until_artifact_text(client: httpx.AsyncClient, task_id: str, expected_text: str) -> dict: + last_task = None + for _ in range(50): + response = await client.post("/long_agent/", json=get_task_payload(task_id)) + assert response.status_code == 200 + last_task = extract_task(response.json()) + if artifact_text(last_task) == expected_text: + return last_task + await asyncio.sleep(0.01) + msg = f"Task {task_id} did not expose artifact text {expected_text!r}. Last task: {last_task}" + raise AssertionError(msg) + + +async def test_detached_send_returns_non_terminal_task(long_running_client): + client, wrapper = long_running_client + payload = send_message_payload("start") + payload["params"]["configuration"] = {"returnImmediately": True} + + response = await client.post("/long_agent/", json=payload) + + assert response.status_code == 200 + task = extract_task(response.json()) + assert task["id"] + assert task["contextId"] + assert task["status"]["state"] in {"TASK_STATE_SUBMITTED", "TASK_STATE_WORKING"} + assert wrapper.entered.is_set() + assert not wrapper.release.is_set() + + progress_task = await poll_task_until_artifact_text(client, task["id"], "progress ") + assert progress_task["status"]["state"] == "TASK_STATE_WORKING" + + wrapper.release.set() + last_task = None + for _ in range(50): + poll_response = await client.post( + "/long_agent/", + json=get_task_payload(task["id"], method="tasks/get"), + headers={"A2A-Version": "0.3"}, + ) + assert poll_response.status_code == 200 + last_task = extract_task(poll_response.json()) + if last_task["status"]["state"] == "completed": + break + await asyncio.sleep(0.01) + assert last_task is not None + assert last_task["status"]["state"] == "completed" + assert artifact_text(last_task) == "progress done" + + +async def test_default_send_remains_blocking(long_running_client): + client, wrapper = long_running_client + + request_task = asyncio.create_task(client.post("/long_agent/", json=send_message_payload("start"))) + await asyncio.wait_for(wrapper.entered.wait(), timeout=1) + await asyncio.sleep(0) + assert not request_task.done() + + wrapper.release.set() + response = await asyncio.wait_for(request_task, timeout=1) + task = extract_task(response.json()) + assert task["status"]["state"] == "TASK_STATE_COMPLETED" + assert artifact_text(task) == "progress done" + + +async def test_a2a_v0_3_blocking_false_returns_active_task(long_running_client): + client, wrapper = long_running_client + payload = send_message_v0_3_payload("start") + payload["params"]["configuration"] = {"blocking": False} + + response = await client.post("/long_agent/", json=payload, headers={"A2A-Version": "0.3"}) + assert response.status_code == 200 + task = extract_task(response.json()) + assert task["status"]["state"] in {"submitted", "working"} + + wrapper.release.set() + completed_task = await poll_task_until_state(client, task["id"], "TASK_STATE_COMPLETED") + assert artifact_text(completed_task) == "progress done" + + +async def test_subscribe_to_active_task(long_running_client): + client, wrapper = long_running_client + payload = send_message_payload("start") + payload["params"]["configuration"] = {"returnImmediately": True} + send_response = await client.post("/long_agent/", json=payload) + task = extract_task(send_response.json()) + + subscribe_payload = get_task_payload(task["id"], method="SubscribeToTask") + + async def read_subscription_events() -> list[dict]: + events = [] + async with client.stream("POST", "/long_agent/", json=subscribe_payload) as response: + assert response.status_code == 200 + async for line in response.aiter_lines(): + if line.startswith("data:"): + events.append(json.loads(line[len("data:") :])) + if ( + events[-1]["result"].get("statusUpdate", {}).get("status", {}).get("state") + == "TASK_STATE_COMPLETED" + ): + break + return events + + subscription_task = asyncio.create_task(read_subscription_events()) + await asyncio.sleep(0.01) + wrapper.release.set() + events = await asyncio.wait_for(subscription_task, timeout=1) + + assert next(iter(events[0]["result"].keys())) == "task" + assert "artifactUpdate" in [next(iter(event["result"].keys())) for event in events] + assert events[-1]["result"]["statusUpdate"]["status"]["state"] == "TASK_STATE_COMPLETED" + + +async def test_cooperative_async_cancellation(long_running_client): + client, wrapper = long_running_client + payload = send_message_payload("start") + payload["params"]["configuration"] = {"returnImmediately": True} + send_response = await client.post("/long_agent/", json=payload) + task = extract_task(send_response.json()) + + await asyncio.wait_for(wrapper.entered.wait(), timeout=1) + cancel_response = await client.post("/long_agent/", json=cancel_task_payload(task["id"])) + + assert cancel_response.status_code == 200 + canceled_task = extract_task(cancel_response.json()) + assert canceled_task["status"]["state"] == "TASK_STATE_CANCELED" + assert artifact_text(canceled_task) == "progress " + assert wrapper.cancelled.is_set() diff --git a/tests/test_redis_a2a_recovery_integration.py b/tests/test_redis_a2a_recovery_integration.py new file mode 100644 index 00000000..00add4af --- /dev/null +++ b/tests/test_redis_a2a_recovery_integration.py @@ -0,0 +1,323 @@ +"""Real-Redis regression coverage for durable A2A task recovery.""" + +from __future__ import annotations + +import asyncio +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from hayhooks.durable.models import ExecutionStatus +from hayhooks.durable.runtime import execution_id_for +from hayhooks.server.a2a.durable_executor import DurableAgentExecutor +from hayhooks.server.a2a.imports import InvalidParamsError +from hayhooks.server.a2a.redis_task_store import RedisTaskStore + +pytestmark = pytest.mark.integration + + +class _CompletedDeployment: + def __init__(self, execution_id: str) -> None: + self.execution_id = execution_id + self.record = SimpleNamespace( + status=ExecutionStatus.COMPLETED, + progress=[], + result={"last_message": {"content": "recovered"}}, + error=None, + sequence=1, + ) + + async def get(self, execution_id: str, **_kwargs): + if execution_id != self.execution_id: + raise KeyError(execution_id) + return self.record + + +@pytest.fixture +async def redis_task_store(isolated_redis): + redis, prefix = isolated_redis + store = RedisTaskStore(redis, "agent", key_prefix=f"{prefix}:a2a") + yield redis, store + + +def _task(task_id: str, seconds: int = 0): + from a2a.types import Task, TaskState + + task = Task(id=task_id, context_id=f"context-{task_id}") + task.status.state = TaskState.TASK_STATE_WORKING + task.status.timestamp.FromDatetime(datetime(2026, 1, 1, tzinfo=timezone.utc) + timedelta(seconds=seconds)) + return task + + +def _context(owner: str): + return SimpleNamespace(user=SimpleNamespace(is_authenticated=True, user_name=owner)) + + +async def test_redis_a2a_recovery_uses_atomic_projection_and_owner(redis_task_store) -> None: + from a2a.server.context import ServerCallContext + from a2a.types import TaskState + + redis, store = redis_task_store + context = ServerCallContext() + task = _task("client.task/" + "x" * 160) + execution_id = execution_id_for("anonymous", task.id) + await store.save(task, context) + + # A replacement executor sees the task through Redis, not a live request queue. + first = DurableAgentExecutor("agent", store, _CompletedDeployment(execution_id)) + await first.close() + second = DurableAgentExecutor("agent", store, _CompletedDeployment(execution_id)) + await second.start() + + recovered = await store.get(task.id, context) + assert recovered is not None and recovered.status.state == TaskState.TASK_STATE_COMPLETED + version = store._loaded_task_version(recovered) + assert version == 2 + + stale = type(recovered)() + stale.CopyFrom(recovered) + assert not await store.save_projection(stale, "anonymous", version - 1) + + await store.delete(task.id, context) + assert await store.get(task.id, context) is None + assert await redis.zscore(store._key("active"), task.id) is None + + +async def test_redis_a2a_read_through_persists_late_completion(redis_task_store) -> None: + from a2a.server.context import ServerCallContext + from a2a.types import TaskState + + redis, store = redis_task_store + context = ServerCallContext() + task = _task("late-completion") + execution_id = execution_id_for("anonymous", task.id) + deployment = _CompletedDeployment(execution_id) + deployment.record.status = ExecutionStatus.RUNNING + await store.save(task, context) + + executor = DurableAgentExecutor("agent", store, deployment) + await executor.start() + deployment.record.status = ExecutionStatus.COMPLETED + completed = await executor.task_store.get(task.id, context) + persisted = await store.get(task.id, context) + + assert completed.status.state == TaskState.TASK_STATE_COMPLETED + assert persisted.status.state == TaskState.TASK_STATE_COMPLETED + assert await redis.zscore(store._key("active"), task.id) is None + assert await redis.zscore(store._key("terminal-expiry"), task.id) is not None + + +async def test_redis_task_store_isolates_agents_and_owners(redis_task_store) -> None: + from a2a.types import TaskState + + redis, base_store = redis_task_store + store = RedisTaskStore(redis, "agent/one", key_prefix=base_store.key_prefix) + other_agent = RedisTaskStore(redis, "agent/two", key_prefix=base_store.key_prefix) + owner = _context("alice@example.com") + other_owner = _context("bob@example.com") + task = _task("task-1", 1) + await store.save(task, owner) + + assert (await store.get(task.id, owner)).id == task.id + recoverable, _ = await store.recoverable_task_batch(0, 1) + assert recoverable[0][1:] == ("user:alice@example.com", 1) + assert await store.get(task.id, other_owner) is None + assert await other_agent.get(task.id, owner) is None + task.status.state = TaskState.TASK_STATE_COMPLETED + await store.save(task, owner) + assert (await store.get(task.id, owner)).status.state == TaskState.TASK_STATE_COMPLETED + + with pytest.raises(InvalidParamsError, match="belongs to another owner"): + await store.save(_task(task.id, 2), other_owner) + assert (await store.get(task.id, owner)).status.state == TaskState.TASK_STATE_COMPLETED + + await store.delete(task.id, owner) + assert await store.get(task.id, owner) is None + + +async def test_redis_task_store_uses_one_owner_resolver_for_all_paths(redis_task_store) -> None: + redis, base_store = redis_task_store + store = RedisTaskStore( + redis, + "custom-owner-agent", + key_prefix=base_store.key_prefix, + owner_resolver=lambda _context: "tenant:alice", + ) + context = _context("ignored") + task = _task("client.task/with unicode-€", 1) + + await store.save(task, context) + + assert await store.get(task.id, context) is not None + recoverable, _ = await store.recoverable_task_batch(0, 1) + assert recoverable[0][1] == "tenant:alice" + + +async def test_all_task_store_writes_reject_a_stale_loaded_version(redis_task_store) -> None: + redis, base_store = redis_task_store + first = RedisTaskStore(redis, "stale-writes", key_prefix=base_store.key_prefix) + second = RedisTaskStore(redis, "stale-writes", key_prefix=base_store.key_prefix) + context = _context("alice@example.com") + task = _task("task", 1) + await first.save(task, context) + stale = await second.get(task.id, context) + assert stale is not None + + task.status.timestamp.FromDatetime(datetime(2026, 1, 2, tzinfo=timezone.utc)) + await first.save(task, context) + with pytest.raises(InvalidParamsError, match="stale projection version"): + await second.save(stale, context) + + +async def test_same_store_tracks_versions_per_loaded_task_snapshot(redis_task_store) -> None: + from a2a.types import TaskState + + _, store = redis_task_store + context = _context("alice@example.com") + await store.save(_task("task", 1), context) + first = await store.get("task", context) + stale = await store.get("task", context) + assert first is not None + assert stale is not None + + first.status.state = TaskState.TASK_STATE_COMPLETED + await store.save(first, context) + stale.status.state = TaskState.TASK_STATE_FAILED + with pytest.raises(InvalidParamsError, match="stale projection version"): + await store.save(stale, context) + + persisted = await store.get("task", context) + assert persisted is not None + assert persisted.status.state == TaskState.TASK_STATE_COMPLETED + + +async def test_concurrent_projection_writers_use_one_version_fence(redis_task_store) -> None: + redis, base_store = redis_task_store + first = RedisTaskStore(redis, "projection-race", key_prefix=base_store.key_prefix) + second = RedisTaskStore(redis, "projection-race", key_prefix=base_store.key_prefix) + context = _context("alice@example.com") + owner = "user:alice@example.com" + await first.save(_task("task", 1), context) + left = await first.get("task", context) + right = await second.get("task", context) + assert left is not None and right is not None + left.metadata["winner"] = "left" + right.metadata["winner"] = "right" + + results = await asyncio.gather( + first.save_projection(left, owner, 1), + second.save_projection(right, owner, 1), + ) + + assert sorted(results) == [False, True] + persisted = await first.get("task", context) + assert persisted is not None and persisted.metadata["winner"] in {"left", "right"} + + +async def test_cleanup_preserves_a_task_whose_terminal_ttl_was_extended(monkeypatch, redis_task_store) -> None: + from a2a.types import TaskState + + redis, base_store = redis_task_store + cleanup = RedisTaskStore(redis, "cleanup-race", key_prefix=base_store.key_prefix, terminal_ttl_seconds=60) + writer = RedisTaskStore(redis, "cleanup-race", key_prefix=base_store.key_prefix, terminal_ttl_seconds=60) + context = _context("alice@example.com") + task = _task("task", 1) + task.status.state = TaskState.TASK_STATE_COMPLETED + await cleanup.save(task, context) + current = await writer.get(task.id, context) + assert current is not None + await redis.zadd(cleanup._key("terminal-expiry"), {task.id: 0}) + + candidate_selected = asyncio.Event() + allow_delete = asyncio.Event() + delete_payload = cleanup._delete_payload + + async def delayed_delete(*args, **kwargs): + candidate_selected.set() + await allow_delete.wait() + return await delete_payload(*args, **kwargs) + + monkeypatch.setattr(cleanup, "_delete_payload", delayed_delete) + cleanup_call = asyncio.create_task(cleanup.cleanup_expired_tasks()) + await candidate_selected.wait() + current.metadata["updated"] = "yes" + await writer.save(current, context) + allow_delete.set() + + assert await cleanup_call == 0 + persisted = await cleanup.get(task.id, context) + assert persisted is not None and persisted.metadata["updated"] == "yes" + assert await redis.zscore(cleanup._key("terminal-expiry"), task.id) is not None + await redis.hset(cleanup._task_key(task.id), "terminal_expiry_ms", 1) + await redis.zadd(cleanup._key("terminal-expiry"), {task.id: 1}) + assert await cleanup.cleanup_expired_tasks() == 1 + assert await cleanup.get(task.id, context) is None + assert await redis.zscore(cleanup._updates_key("user:alice@example.com"), task.id) is None + + +async def test_redis_task_store_lists_with_filters_and_page_tokens(redis_task_store) -> None: + from a2a.types import ListTasksRequest + + _, store = redis_task_store + context = _context("alice@example.com") + for index in range(3): + await store.save(_task(f"task-{index}", index), context) + + first_page = await store.list(ListTasksRequest(page_size=2), context) + assert [task.id for task in first_page.tasks] == ["task-2", "task-1"] + assert first_page.total_size == 3 + assert first_page.next_page_token + + second_page = await store.list( + ListTasksRequest(page_size=2, page_token=first_page.next_page_token), + context, + ) + assert [task.id for task in second_page.tasks] == ["task-0"] + assert second_page.next_page_token == "" + + filtered_page = await store.list(ListTasksRequest(context_id="context-task-2"), context) + assert [task.id for task in filtered_page.tasks] == ["task-2"] + + with pytest.raises(InvalidParamsError, match="base64-encoded cursor"): + await store.list(ListTasksRequest(page_token="invalid"), context) # noqa: S106 + + +async def test_redis_task_store_compares_status_timestamps_as_timestamps(redis_task_store) -> None: + from a2a.types import ListTasksRequest + + _, store = redis_task_store + context = _context("alice@example.com") + task = _task("task") + task.status.timestamp.FromDatetime(datetime(2026, 1, 1, 0, 0, 0, 100_000, tzinfo=timezone.utc)) + await store.save(task, context) + + request = ListTasksRequest() + request.status_timestamp_after.FromDatetime(datetime(2026, 1, 1, tzinfo=timezone.utc)) + page = await store.list(request, context) + + assert [item.id for item in page.tasks] == ["task"] + + +async def test_filtered_task_listing_and_snapshot_cache_are_globally_bounded(monkeypatch, redis_task_store) -> None: + from a2a.types import ListTasksRequest + + from hayhooks.settings import settings + + monkeypatch.setattr(settings, "a2a_list_scan_batch_size", 2) + monkeypatch.setattr(settings, "a2a_task_snapshot_cache_size", 3) + redis, store = redis_task_store + load_snapshots = AsyncMock(side_effect=store._load_snapshots) + monkeypatch.setattr(store, "_load_snapshots", load_snapshots) + monkeypatch.setattr(redis, "hvals", AsyncMock(side_effect=AssertionError("unbounded task scan"))) + context = _context("alice@example.com") + for index in range(6): + await store.save(_task(f"task-{index}", index), context) + await store.get(f"task-{index}", context) + + page = await store.list(ListTasksRequest(context_id="context-task-5"), context) + + assert [task.id for task in page.tasks] == ["task-5"] + assert all(len(call.args[0]) <= 2 for call in load_snapshots.await_args_list) + assert len(store._loaded_task_versions) <= 3 diff --git a/tests/test_redis_task_store.py b/tests/test_redis_task_store.py new file mode 100644 index 00000000..7df73f75 --- /dev/null +++ b/tests/test_redis_task_store.py @@ -0,0 +1,120 @@ +import inspect +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from redis.cluster import key_slot + +from hayhooks.a2a import RedisTaskStore, RedisTaskStoreProvider, TaskStoreProvider +from hayhooks.durable.backend import ExecutionStoreCorruptionError +from hayhooks.durable.runtime import execution_id_for +from hayhooks.server.a2a.durable_executor import DurableTaskStore +from hayhooks.server.a2a.runtime import create_task_store_provider + + +def _context(owner: str): + return SimpleNamespace(user=SimpleNamespace(is_authenticated=True, user_name=owner)) + + +def test_default_owner_distinguishes_unauthenticated_and_anonymous_user(): + store = RedisTaskStore(AsyncMock(), "agent", key_prefix="test:a2a") + unauthenticated = SimpleNamespace(user=SimpleNamespace(is_authenticated=False, user_name="")) + + assert store.owner_id_for_context(unauthenticated) == "anonymous" + assert store.owner_id_for_context(_context("anonymous")) == "user:anonymous" + + +def test_durable_task_store_delegates_custom_owner_resolution(): + store = RedisTaskStore(AsyncMock(), "agent", key_prefix="test:a2a", owner_resolver=lambda _context: "tenant:alice") + + assert DurableTaskStore(store, object()).owner_id_for_context(_context("ignored")) == "tenant:alice" + + +def test_execution_id_for_accepts_arbitrary_a2a_task_ids(): + owner = "tenant:alice" + task_ids = ["", "client.task", "folder/task", "λ" * 200] + execution_ids = [execution_id_for(owner, task_id) for task_id in task_ids] + + assert all(len(execution_id) == 64 and execution_id.isalnum() for execution_id in execution_ids) + assert len(set(execution_ids)) == len(task_ids) + assert execution_id_for("tenant:bob", "client.task") != execution_id_for(owner, "client.task") + + +def test_redis_task_keys_share_one_cluster_slot_and_store_uses_no_eval(): + store = RedisTaskStore(AsyncMock(), "agent", key_prefix="test:a2a") + keys = ( + store._task_key("task"), + store._updates_key("tenant:alice"), + store._key("active"), + store._key("terminal-expiry"), + ) + + assert len({key_slot(key.encode()) for key in keys}) == 1 + assert ".eval(" not in inspect.getsource(RedisTaskStore) + + +def test_redis_task_store_normalizes_corrupt_protobuf_snapshots(): + with pytest.raises(ExecutionStoreCorruptionError, match="invalid Redis snapshot"): + RedisTaskStore._decode_snapshot( + "task", + { + b"owner": b"owner", + b"payload": b"not-a-protobuf", + b"version": b"1", + b"terminal_expiry_ms": b"0", + }, + ) + + +async def test_redis_task_store_provider_health_pings_redis(): + redis = AsyncMock() + provider = RedisTaskStoreProvider(redis=redis, close_redis=False) + + await provider.initialize() + assert (await provider.health())["healthy"] + + redis.ping.side_effect = ConnectionError("offline") + health = await provider.health() + assert not health["healthy"] + assert health["error"] == "ConnectionError" + + +async def test_redis_task_store_provider_is_cached_and_closes_redis(): + redis = AsyncMock() + provider = RedisTaskStoreProvider(redis=redis, key_prefix="test:a2a") + + first = provider.create_task_store("agent") + assert provider.create_task_store("agent") is first + assert isinstance(provider, TaskStoreProvider) + + await provider.close() + redis.aclose.assert_awaited_once() + + +def test_redis_task_store_defaults_use_app_settings(monkeypatch): + from hayhooks.settings import settings + + monkeypatch.setattr(settings, "a2a_redis_key_prefix", "configured:a2a:") + monkeypatch.setattr(settings, "a2a_redis_socket_timeout", 3.5) + monkeypatch.setattr(settings, "a2a_redis_socket_connect_timeout", 2.5) + monkeypatch.setattr(settings, "a2a_redis_health_check_interval", 20) + redis = AsyncMock() + + direct_store = RedisTaskStore(redis, "direct") + provider_store = RedisTaskStoreProvider(redis=redis, close_redis=False).create_task_store("provided") + + assert direct_store.key_prefix == "configured:a2a" + assert provider_store.key_prefix == "configured:a2a" + + provider = RedisTaskStoreProvider(redis=redis, close_redis=False) + assert provider.socket_timeout == 3.5 + assert provider.socket_connect_timeout == 2.5 + assert provider.health_check_interval == 20 + + +def test_create_task_store_provider_selects_builtin_backends(): + memory = create_task_store_provider() + assert type(memory).__name__ == "InMemoryTaskStoreProvider" + + redis = create_task_store_provider(backend="redis", redis_url="redis://localhost:6379/2") + assert isinstance(redis, RedisTaskStoreProvider) From e10e61d55f091b24f9b8d618102701e722b1490e Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:09:32 +0200 Subject: [PATCH 06/28] feat(server): integrate durable deployment lifecycle --- src/hayhooks/cli/a2a.py | 65 +++- src/hayhooks/server/app.py | 14 +- src/hayhooks/server/durable/__init__.py | 1 + src/hayhooks/server/durable/routes.py | 266 ++++++++++++++++ src/hayhooks/server/pipelines/models.py | 41 ++- src/hayhooks/server/routers/status.py | 10 +- src/hayhooks/server/tracing.py | 41 ++- src/hayhooks/server/utils/deploy_utils.py | 324 ++++++++++++++++--- src/hayhooks/server/utils/mcp_utils.py | 9 +- src/hayhooks/server/utils/module_loader.py | 11 +- tests/test_cli.py | 96 +++++- tests/test_deploy_performance.py | 35 +-- tests/test_deploy_utils.py | 25 +- tests/test_durable_deployment_lifecycle.py | 348 +++++++++++++++++++++ tests/test_it_deploy_files.py | 4 +- 15 files changed, 1188 insertions(+), 102 deletions(-) create mode 100644 src/hayhooks/server/durable/__init__.py create mode 100644 src/hayhooks/server/durable/routes.py create mode 100644 tests/test_durable_deployment_lifecycle.py diff --git a/src/hayhooks/cli/a2a.py b/src/hayhooks/cli/a2a.py index 2c0e0cbd..04764435 100644 --- a/src/hayhooks/cli/a2a.py +++ b/src/hayhooks/cli/a2a.py @@ -1,5 +1,5 @@ import sys -from typing import Annotated +from typing import Annotated, Literal import typer @@ -20,6 +20,38 @@ def run( # noqa: PLR0913 str | None, typer.Option("--external-url", help="Base URL advertised in agent cards (e.g. behind a reverse proxy)"), ] = None, + task_store: Annotated[ + Literal["auto", "memory", "redis"] | None, + typer.Option("--task-store", help="Built-in A2A task-store backend"), + ] = None, + a2a_redis_url: Annotated[ + str | None, + typer.Option("--a2a-redis-url", help="Redis URL for the built-in A2A task store"), + ] = None, + a2a_redis_key_prefix: Annotated[ + str | None, + typer.Option("--a2a-redis-key-prefix", help="Redis key prefix for the built-in A2A task store"), + ] = None, + execution_store: Annotated[ + Literal["memory", "redis"] | None, + typer.Option("--execution-store", help="Built-in durable execution-store backend"), + ] = None, + execution_redis_url: Annotated[ + str | None, + typer.Option("--execution-redis-url", help="Redis URL for durable execution storage"), + ] = None, + execution_redis_key_prefix: Annotated[ + str | None, + typer.Option("--execution-redis-key-prefix", help="Redis key prefix for durable execution storage"), + ] = None, + durable_execution_concurrency: Annotated[ + int | None, + typer.Option( + "--durable-execution-concurrency", + min=1, + help="Maximum concurrent durable Agent executions per deployed agent", + ), + ] = None, debug: Annotated[bool, typer.Option("--debug", help="If true, tracebacks should be returned on errors")] = False, ) -> None: """ @@ -28,13 +60,11 @@ def run( # noqa: PLR0913 # Lazy imports of settings, logger and uvicorn import uvicorn + from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.logger import intercept_stdlib_logging, log - from hayhooks.server.utils.a2a_utils import a2a_import, create_a2a_app from hayhooks.server.utils.deploy_utils import deploy_pipelines from hayhooks.settings import settings - a2a_import.check() - # Fill defaults from settings when command executes host = host or settings.a2a_host port = port or settings.a2a_port @@ -48,6 +78,27 @@ def run( # noqa: PLR0913 if external_url: settings.a2a_external_url = external_url + if task_store is not None: + settings.a2a_task_store = task_store + + if a2a_redis_url is not None: + settings.a2a_redis_url = a2a_redis_url + + if a2a_redis_key_prefix is not None: + settings.a2a_redis_key_prefix = a2a_redis_key_prefix + + if execution_store is not None: + settings.durable_store = execution_store + + if execution_redis_url is not None: + settings.durable_redis_url = execution_redis_url + + if execution_redis_key_prefix is not None: + settings.durable_redis_key_prefix = execution_redis_key_prefix + + if durable_execution_concurrency is not None: + settings.durable_execution_concurrency = durable_execution_concurrency + if additional_python_path: settings.additional_python_path = additional_python_path sys.path.append(additional_python_path) @@ -58,12 +109,16 @@ def run( # noqa: PLR0913 # Setup the Starlette app exposing pipelines as A2A agents log.debug( - "Starting A2A server with host={}, port={}, pipelines_dir={}, external_url={}, v0.3_compat={}", + "Starting A2A server with host={}, port={}, pipelines_dir={}, external_url={}, " + "v0.3_compat={}, task_store={}, durable_store={}, durable_execution_concurrency={}", host, port, pipelines_dir, settings.a2a_external_url or "", settings.a2a_v0_3_compat, + settings.a2a_task_store, + settings.durable_store, + settings.durable_execution_concurrency, ) app = create_a2a_app(debug=debug) diff --git a/src/hayhooks/server/app.py b/src/hayhooks/server/app.py index 80b24f4a..9a29a54b 100644 --- a/src/hayhooks/server/app.py +++ b/src/hayhooks/server/app.py @@ -20,6 +20,7 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles +from hayhooks.durable.runtime import durable_runtime from hayhooks.server.logger import RequestIdMiddleware, intercept_stdlib_logging, log, log_elapsed from hayhooks.server.routers import ( dashboard_router, @@ -241,17 +242,20 @@ def deploy_pipelines(app: FastAPI, pipelines_dir: PathLike | str) -> None: @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncIterator[None]: - if settings.pipelines_dir: - deploy_pipelines(app, settings.pipelines_dir) - # Capture the running loop so synchronous span recording can wake SSE # subscribers via call_soon_threadsafe. broadcaster = get_trace_stream_broadcaster() - broadcaster.set_loop(asyncio.get_running_loop()) try: + if settings.pipelines_dir: + deploy_pipelines(app, settings.pipelines_dir) + await durable_runtime.start() + broadcaster.set_loop(asyncio.get_running_loop()) yield finally: - broadcaster.clear_loop() + try: + await durable_runtime.close() + finally: + broadcaster.clear_loop() @lru_cache(maxsize=1) diff --git a/src/hayhooks/server/durable/__init__.py b/src/hayhooks/server/durable/__init__.py new file mode 100644 index 00000000..5005858d --- /dev/null +++ b/src/hayhooks/server/durable/__init__.py @@ -0,0 +1 @@ +"""HTTP-facing durable execution adapters.""" diff --git a/src/hayhooks/server/durable/routes.py b/src/hayhooks/server/durable/routes.py new file mode 100644 index 00000000..15fd18fc --- /dev/null +++ b/src/hayhooks/server/durable/routes.py @@ -0,0 +1,266 @@ +"""Typed REST resources that project the durable execution record.""" + +from __future__ import annotations + +import inspect +import re +from typing import Annotated, Any, cast +from urllib.parse import quote + +from fastapi import Body, FastAPI, Header, HTTPException, Path, Request, status +from fastapi.responses import Response +from fastapi.routing import APIRoute +from pydantic import ValidationError, create_model + +from hayhooks.durable import ExecutionResult +from hayhooks.durable.engine import RUN_ID_PATTERN +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError +from hayhooks.durable.runtime import ( + DefinitionRevisionConflictError, + DurableDeployment, + IdempotencyConflictError, + durable_runtime, +) +from hayhooks.server.pipelines.registry import registry +from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.settings import settings + +DURABLE_ROUTE_SUFFIXES = ( + "/run-durable", + "/executions/{execution_id}", + "/executions/{execution_id}/cancel", + "/executions/{execution_id}/resume", +) +_IDEMPOTENCY_KEY_PATTERN = re.compile(rf"^{RUN_ID_PATTERN}$") +_MAX_DURABLE_OWNER_LENGTH = 512 +_MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH = 63 +ExecutionId = Annotated[str, Path(pattern=rf"^{RUN_ID_PATTERN}$", min_length=1, max_length=128)] + + +def _execution_links(pipeline_name: str, execution_id: str) -> dict[str, str]: + root = f"/{pipeline_name}/executions/{quote(execution_id, safe='-._~')}" + return {"self": root, "cancel": f"{root}/cancel", "resume": f"{root}/resume"} + + +def _execution_result( + deployment: DurableDeployment, + record: Any, + *, + response_model: type[ExecutionResult] = ExecutionResult, +) -> ExecutionResult: + return response_model.model_validate(record.safe_view(links=_execution_links(deployment.name, record.execution_id))) + + +def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResult]: + if deployment.result_type is Any: + return ExecutionResult + return create_model( + f"{deployment.name.title().replace('-', '').replace('_', '')}ExecutionResult", + __base__=ExecutionResult, + result=(deployment.result_type | None, None), + ) + + +def _durable_owner(request: Request) -> tuple[str | None, bool]: + header = settings.durable_trusted_owner_header.strip() + if not header: + return None, False + owner = request.headers.get(header) + if not owner: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=f"Authenticated owner header '{header}' is required", + ) + if len(owner) > _MAX_DURABLE_OWNER_LENGTH: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Authenticated owner header '{header}' exceeds 512 characters", + ) + return owner, True + + +def _remove_pipeline_route(app: FastAPI, path: str, method: str) -> None: + for route in list(app.routes): + if isinstance(route, APIRoute) and route.path == path and route.methods is not None and method in route.methods: + app.routes.remove(route) + + +def _remove_durable_api_routes(app: FastAPI, pipeline_name: str) -> None: + root = f"/{pipeline_name}" + durable_paths = {f"{root}{suffix}" for suffix in DURABLE_ROUTE_SUFFIXES} + app.routes[:] = [route for route in app.routes if not (isinstance(route, APIRoute) and route.path in durable_paths)] + + +def add_durable_api_routes( # noqa: C901, PLR0915 - route-local handlers share generated models + app: FastAPI, + pipeline_name: str, + pipeline_wrapper: BasePipelineWrapper, + *, + deployment: DurableDeployment | None = None, + _defer_openapi_rebuild: bool, +) -> None: + """Register typed durable submission and control resources when opted in.""" + _remove_durable_api_routes(app, pipeline_name) + if not durable_runtime.has_capability(pipeline_wrapper): + if not _defer_openapi_rebuild: + app.openapi_schema = None + app.setup() + return + deployment = deployment or durable_runtime.deployment(pipeline_name, pipeline_wrapper) + request_model = deployment.request_type + response_model = _durable_response_model(deployment) + root = f"/{pipeline_name}" + + async def get_execution(execution_id: str, owner_id: str | None, enforce_owner: bool) -> Any: + return await deployment.get( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + allow_revision_mismatch=True, + ) + + async def submit( + run_req: Any, + response: Response, + request: Request, + idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), + ) -> ExecutionResult: + owner_id, owner_scoped = _durable_owner(request) + if idempotency_key is not None and _IDEMPOTENCY_KEY_PATTERN.fullmatch(idempotency_key) is None: + raise HTTPException( + status_code=422, + detail="Idempotency-Key must contain 1-128 letters, digits, underscores, or hyphens", + ) + if ( + owner_scoped + and idempotency_key is not None + and len(idempotency_key) > _MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH + ): + raise HTTPException( + status_code=422, + detail="Idempotency-Key must be at most 63 characters when owner scoping is enabled", + ) + try: + created, record = await deployment.submit( + run_req.model_dump(mode="json"), + execution_id=idempotency_key, + owner_id=owner_id, + ) + except IdempotencyConflictError as error: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error + except DefinitionRevisionConflictError as error: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error + except (ValidationError, ValueError) as error: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error + except ExecutionAdmissionError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=str(error), + headers={"Retry-After": str(error.retry_after_seconds)}, + ) from error + except (ExecutionStoreError, RuntimeError) as error: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error + response.status_code = status.HTTP_200_OK if not created and record.terminal else status.HTTP_202_ACCEPTED + response.headers["Location"] = _execution_links(deployment.name, record.execution_id)["self"] + if not created: + response.headers["Idempotent-Replay"] = "true" + return _execution_result(deployment, record, response_model=response_model) + + # FastAPI consumes the runtime annotation; keep the Python type checker on + # the stable BaseModel boundary while preserving the generated request schema. + submit.__annotations__["run_req"] = request_model + + async def inspect_execution(execution_id: ExecutionId, request: Request) -> ExecutionResult: + try: + owner_id, enforce_owner = _durable_owner(request) + record = await get_execution(execution_id, owner_id, enforce_owner) + return _execution_result(deployment, record) + except KeyError as error: + raise HTTPException(status_code=404, detail="Execution not found") from error + except ExecutionStoreError as error: + raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error + + async def cancel_execution(execution_id: ExecutionId, response: Response, request: Request) -> ExecutionResult: + try: + owner_id, enforce_owner = _durable_owner(request) + accepted = await deployment.request_cancel( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + ) + record = await get_execution(execution_id, owner_id, enforce_owner) + except KeyError as error: + raise HTTPException(status_code=404, detail="Execution not found") from error + except ExecutionStoreError as error: + raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error + response.status_code = status.HTTP_202_ACCEPTED if accepted else status.HTTP_200_OK + return _execution_result(deployment, record) + + async def resume_execution( + execution_id: ExecutionId, + response: Response, + request: Request, + update: Any = Body(default=None), # noqa: B008 + ) -> ExecutionResult: + try: + owner_id, enforce_owner = _durable_owner(request) + resumed = await deployment.resume( + execution_id, + update, + owner_id=owner_id, + enforce_owner=enforce_owner, + ) + if not resumed: + raise HTTPException(status_code=409, detail="Execution is not waiting") + record = await get_execution(execution_id, owner_id, enforce_owner) + except KeyError as error: + raise HTTPException(status_code=404, detail="Execution not found") from error + except (ValidationError, ValueError) as error: + raise HTTPException(status_code=422, detail=str(error)) from error + except ExecutionStoreError as error: + raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error + response.status_code = status.HTTP_202_ACCEPTED + return _execution_result(deployment, record) + + if deployment.resume_type is not None: + resume_execution.__annotations__["update"] = deployment.resume_type + signature = inspect.signature(resume_execution) + update_parameter = signature.parameters["update"].replace( + annotation=deployment.resume_type, + default=Body(), + ) + cast(Any, resume_execution).__signature__ = signature.replace( + parameters=[ + update_parameter if parameter.name == "update" else parameter + for parameter in signature.parameters.values() + ] + ) + + routes = [ + (f"{root}/run-durable", submit, ["POST"], f"{pipeline_name}_run_durable"), + (f"{root}/executions/{{execution_id}}", inspect_execution, ["GET"], f"{pipeline_name}_execution"), + (f"{root}/executions/{{execution_id}}/cancel", cancel_execution, ["POST"], f"{pipeline_name}_cancel"), + (f"{root}/executions/{{execution_id}}/resume", resume_execution, ["POST"], f"{pipeline_name}_resume"), + ] + for path, endpoint, methods, name in routes: + _remove_pipeline_route(app, path, methods[0]) + app.add_api_route( + path, + endpoint, + methods=methods, + name=name, + response_model=response_model if endpoint is submit else ExecutionResult, + tags=["durable executions"], + status_code=status.HTTP_202_ACCEPTED if methods == ["POST"] else status.HTTP_200_OK, + ) + + registry.update_metadata( + pipeline_name, + {"durable_request_model": request_model, "durable_response_model": response_model}, + ) + if not _defer_openapi_rebuild: + app.openapi_schema = None + app.setup() + + +__all__ = ["DURABLE_ROUTE_SUFFIXES", "add_durable_api_routes"] diff --git a/src/hayhooks/server/pipelines/models.py b/src/hayhooks/server/pipelines/models.py index aec68137..c973ae81 100644 --- a/src/hayhooks/server/pipelines/models.py +++ b/src/hayhooks/server/pipelines/models.py @@ -1,6 +1,6 @@ import inspect from collections.abc import AsyncGenerator, Callable, Generator -from typing import Any, get_origin +from typing import Any, get_origin, get_type_hints from docstring_parser.common import Docstring from fastapi.responses import Response, StreamingResponse @@ -10,6 +10,28 @@ from hayhooks.server.utils.yaml_utils import InputResolution, OutputResolution +def _create_schema_model(model_name: str, **fields: Any) -> type[BaseModel]: + """Create a dynamic API model that Pydantic can resolve during OpenAPI generation.""" + model = create_model(model_name, __module__=__name__, **fields) + + # FastAPI may rebuild a route's TypeAdapter long after the route is + # registered. Pydantic resolves the dynamically-created model by name in + # this module's namespace at that point, so keep it available there. + globals()[model_name] = model + model.model_rebuild() + return model + + +def _resolved_annotations(func: Callable) -> dict[str, Any]: + """Resolve postponed annotations in the wrapper module that declared them.""" + try: + return get_type_hints(func, include_extras=True) + except (NameError, TypeError) as error: + name = getattr(func, "__name__", type(func).__name__) + msg = f"Pipeline wrapper has an invalid type annotation on '{name}': {error}" + raise PipelineWrapperError(msg) from error + + def get_request_model_from_resolved_io( pipeline_name: str, declared_inputs: dict[str, InputResolution] ) -> type[BaseModel]: @@ -30,7 +52,7 @@ def get_request_model_from_resolved_io( default_value = ... if resolution.required else None fields[input_name] = (input_type, default_value) - return create_model(f"{pipeline_name.capitalize()}RunRequest", **fields) + return _create_schema_model(f"{pipeline_name.capitalize()}RunRequest", **fields) def get_response_model_from_resolved_io( @@ -52,7 +74,7 @@ def get_response_model_from_resolved_io( output_type = resolution.type fields[output_name] = (output_type, ...) - return create_model( + return _create_schema_model( f"{pipeline_name.capitalize()}RunResponse", result=(dict, Field(..., description="Pipeline result")) ) @@ -70,6 +92,7 @@ def create_request_model_from_callable(func: Callable, model_name: str, docstrin """ params = inspect.signature(func).parameters + annotations = _resolved_annotations(func) param_docs = {p.arg_name: p.description for p in docstring.params} fields: dict[str, Any] = {} @@ -77,9 +100,9 @@ def create_request_model_from_callable(func: Callable, model_name: str, docstrin default_value = ... if param.default == param.empty else param.default description = param_docs.get(name) or f"Parameter '{name}'" field_info = Field(default=default_value, description=description) - fields[name] = (param.annotation, field_info) + fields[name] = (annotations.get(name, param.annotation), field_info) - return create_model(f"{model_name}Request", **fields) + return _create_schema_model(f"{model_name}Request", **fields) def _is_streaming_type(return_type: type) -> bool: @@ -113,7 +136,7 @@ def create_response_model_from_callable( Pydantic model class for response, or None for streaming/file responses. """ - return_type = inspect.signature(func).return_annotation + return_type = _resolved_annotations(func).get("return", inspect.signature(func).return_annotation) if return_type is inspect.Signature.empty: msg = f"Pipeline wrapper is missing a return type for '{func.__name__}' method" # ty: ignore[unresolved-attribute] @@ -134,7 +157,9 @@ def create_response_model_from_callable( return_description = docstring.returns.description if docstring.returns else None - return create_model(f"{model_name}Response", result=(return_type, Field(..., description=return_description))) + return _create_schema_model( + f"{model_name}Response", result=(return_type, Field(..., description=return_description)) + ) def get_response_class_from_callable(func: Callable) -> type[Response] | None: @@ -154,7 +179,7 @@ def get_response_class_from_callable(func: Callable) -> type[Response] | None: * ``None`` for normal JSON endpoints (the caller should omit the ``response_class`` kwarg so FastAPI uses its default ``JSONResponse``). """ - return_type = inspect.signature(func).return_annotation + return_type = _resolved_annotations(func).get("return", inspect.signature(func).return_annotation) if return_type is inspect.Signature.empty: return None diff --git a/src/hayhooks/server/routers/status.py b/src/hayhooks/server/routers/status.py index 4b25e88a..9a2a67e1 100644 --- a/src/hayhooks/server/routers/status.py +++ b/src/hayhooks/server/routers/status.py @@ -1,6 +1,7 @@ from fastapi import APIRouter, HTTPException from pydantic import BaseModel, Field +from hayhooks.durable.runtime import durable_runtime from hayhooks.server.pipelines.registry import registry router = APIRouter() @@ -9,6 +10,7 @@ class StatusResponse(BaseModel): status: str = Field(description="The current status of the system, 'Up!' when operational") pipelines: list[str] = Field(description="List of all available pipeline names") + durable: dict = Field(default_factory=dict, description="Durable worker readiness and health") model_config = { "json_schema_extra": {"description": "Response model for the system status and available pipelines"} @@ -32,7 +34,10 @@ class PipelineStatusResponse(BaseModel): ) async def status_all() -> StatusResponse: pipelines = registry.get_names() - return StatusResponse(status="Up!", pipelines=pipelines) + durable_health = await durable_runtime.health() + if not durable_health["healthy"]: + raise HTTPException(status_code=503, detail={"status": "Degraded", "durable": durable_health}) + return StatusResponse(status="Up!", pipelines=pipelines, durable=durable_health) @router.get( @@ -46,4 +51,7 @@ async def status_all() -> StatusResponse: async def status(pipeline_name: str) -> PipelineStatusResponse: if pipeline_name not in registry.get_names(): raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") + deployment = durable_runtime.current_deployment(pipeline_name) + if deployment is not None and not deployment.manager.health["healthy"]: + raise HTTPException(status_code=503, detail=f"Pipeline '{pipeline_name}' has no live durable worker slots") return PipelineStatusResponse(status="Up!", pipeline=pipeline_name) diff --git a/src/hayhooks/server/tracing.py b/src/hayhooks/server/tracing.py index 9a883ab5..73e1e6e9 100644 --- a/src/hayhooks/server/tracing.py +++ b/src/hayhooks/server/tracing.py @@ -12,9 +12,10 @@ from __future__ import annotations +import importlib import os import traceback -from collections.abc import AsyncGenerator, Generator, Iterator, Mapping +from collections.abc import AsyncGenerator, AsyncIterator, Generator, Iterator, Mapping from contextlib import contextmanager, nullcontext from contextvars import ContextVar, Token, copy_context from time import monotonic, time @@ -47,6 +48,9 @@ SPAN_MCP_CALL_TOOL = "hayhooks.mcp.call_tool" SPAN_MCP_RUN_PIPELINE_TOOL = "hayhooks.mcp.run_pipeline_tool" SPAN_A2A_RUN_AGENT = "hayhooks.a2a.run_agent" +SPAN_A2A_DURABLE_PROJECT = "hayhooks.a2a.durable.project" +SPAN_DURABLE_SUBMIT = "hayhooks.durable.submit" +SPAN_DURABLE_ATTEMPT = "hayhooks.durable.attempt" _TAG_SUCCESS = "hayhooks.success" _TAG_ERROR_TYPE = "hayhooks.error.type" @@ -309,6 +313,35 @@ def _load_fastapi_instrumentor() -> type[Any] | None: return FastAPIInstrumentor +def _patch_fastapi_route_details() -> None: + """ + Make older OpenTelemetry FastAPI instrumentation safe with FastAPI routers. + + Recent FastAPI versions represent included routers as private route objects + without a ``path`` attribute. OpenTelemetry's route discovery dereferences + that attribute for partial matches, causing every request to fail before it + reaches the application. Preserve tracing while falling back to the + concrete request path until the dependency handles this router shape. + """ + try: + otel_fastapi: Any = importlib.import_module("opentelemetry.instrumentation.fastapi") + except ImportError: # pragma: no cover - guarded by the lazy import above + return + + original = otel_fastapi._get_route_details + if getattr(original, "_hayhooks_safe", False): + return + + def safe_get_route_details(scope: Any) -> str | None: + try: + return original(scope) + except AttributeError: + return scope.get("path") + + safe_get_route_details.__dict__["_hayhooks_safe"] = True + otel_fastapi._get_route_details = safe_get_route_details + + def _load_starlette_instrumentor() -> type[Any] | None: """Return Starlette OTel instrumentor when tracing extras are installed.""" try: @@ -460,6 +493,8 @@ def _instrument_app( return False try: + if framework_name == "FastAPI": + _patch_fastapi_route_details() instrument_kwargs: dict[str, Any] = {} if settings.tracing_excluded_spans: instrument_kwargs["exclude_spans"] = settings.tracing_excluded_spans @@ -695,7 +730,7 @@ def start_trace_operation( def trace_sync_stream( - stream: Generator[Any, None, None], + stream: Iterator[Any], operation_name: str, *, tags: Mapping[str, Any] | None = None, @@ -729,7 +764,7 @@ def traced_stream() -> Generator[Any, None, None]: def trace_async_stream( - stream: AsyncGenerator[Any, None], + stream: AsyncIterator[Any], operation_name: str, *, tags: Mapping[str, Any] | None = None, diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index fc90918e..dce0e889 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -2,11 +2,12 @@ import inspect import json import shutil +import sys import tempfile -import threading import time import traceback -from collections.abc import AsyncGenerator, Callable, Generator +from collections.abc import AsyncGenerator, Awaitable, Callable, Generator +from contextlib import nullcontext from functools import wraps from pathlib import Path from typing import Any, cast @@ -18,6 +19,9 @@ from fastapi.routing import APIRoute from pydantic import BaseModel +from hayhooks.durable.runtime import DurableDeployment, durable_runtime +from hayhooks.server.durable.routes import DURABLE_ROUTE_SUFFIXES as _DURABLE_ROUTE_SUFFIXES +from hayhooks.server.durable.routes import add_durable_api_routes as _add_durable_api_routes from hayhooks.server.exceptions import PipelineAlreadyExistsError, PipelineFilesError from hayhooks.server.logger import log, log_elapsed from hayhooks.server.pipelines.models import ( @@ -50,29 +54,189 @@ from hayhooks.server.utils.yaml_pipeline_wrapper import YAMLPipelineWrapper from hayhooks.settings import DeployConcurrencyPolicy, settings -# threading.Lock (not asyncio.Lock) because it's only acquired inside worker -# threads spawned by asyncio.to_thread, so never on the event loop itself. -_deploy_lock = threading.Lock() - - -def _with_deploy_lock(func: Callable) -> Callable: - """Wrap *func* so it acquires ``_deploy_lock`` before executing.""" - - @wraps(func) - def wrapper(*args, **kwargs): - with _deploy_lock: - return func(*args, **kwargs) - - return wrapper +_deployment_serial_lock = asyncio.Lock() +_deployment_publication_lock = asyncio.Lock() +_deployments_in_progress: set[str] = set() async def _offload(func: Callable, **kwargs: Any) -> Any: - """Run *func* in a thread, applying the deploy lock if policy is SERIALIZED.""" - if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED: - func = _with_deploy_lock(func) + """Run blocking pipeline preparation outside the event loop.""" return await asyncio.to_thread(func, **kwargs) +class _DeploymentSnapshot: + """Rollback state captured before preparation mutates files or loaded modules.""" + + def __init__(self, pipeline_name: str, app: FastAPI | None) -> None: + self.pipeline_name = pipeline_name + self.app = app + self.wrapper = registry.get(pipeline_name) + metadata = registry.get_metadata(pipeline_name) + self.metadata = dict(metadata) if metadata is not None else None + self.deployment = durable_runtime.current_deployment(pipeline_name) + self.routes = list(app.routes) if app is not None else None + self.openapi_schema = app.openapi_schema if app is not None else None + self.modules = { + name: module + for name, module in sys.modules.items() + if name == pipeline_name or name.startswith(f"{pipeline_name}.") + } + pipelines_dir = Path(settings.pipelines_dir) + source_dir = pipelines_dir / pipeline_name + sources = [source_dir] if source_dir.is_dir() else [] + sources.extend( + source + for extension in (".yml", ".yaml") + if (source := pipelines_dir / f"{pipeline_name}{extension}").is_file() + ) + self.backup_dir = Path(tempfile.mkdtemp(prefix="hayhooks-deploy-rollback-")) if sources else None + if source_dir.is_dir() and self.backup_dir is not None: + shutil.copytree(source_dir, self.backup_dir / "pipeline") + for extension in (".yml", ".yaml"): + source = pipelines_dir / f"{pipeline_name}{extension}" + if source.is_file() and self.backup_dir is not None: + shutil.copy2(source, self.backup_dir / f"pipeline{extension}") + + @classmethod + def capture(cls, pipeline_name: str, app: FastAPI | None) -> "_DeploymentSnapshot": + return cls(pipeline_name, app) + + def restore_publication(self) -> None: + registry.remove(self.pipeline_name) + if self.wrapper is not None: + registry.add(self.pipeline_name, self.wrapper, metadata=dict(self.metadata or {})) + durable_runtime.install_deployment(self.pipeline_name, self.deployment) + if self.app is not None and self.routes is not None: + self.app.routes[:] = self.routes + self.app.openapi_schema = self.openapi_schema + + def refresh_publication(self) -> None: + """Capture unrelated route changes made while this pipeline prepared.""" + if self.app is not None: + self.routes = list(self.app.routes) + self.openapi_schema = self.app.openapi_schema + + def restore_files_and_modules(self) -> None: + remove_pipeline_files(self.pipeline_name, settings.pipelines_dir) + pipelines_dir = Path(settings.pipelines_dir) + pipelines_dir.mkdir(parents=True, exist_ok=True) + if self.backup_dir is not None: + backup_pipeline = self.backup_dir / "pipeline" + if backup_pipeline.is_dir(): + shutil.copytree(backup_pipeline, pipelines_dir / self.pipeline_name) + for extension in (".yml", ".yaml"): + backup = self.backup_dir / f"pipeline{extension}" + if backup.is_file(): + shutil.copy2(backup, pipelines_dir / f"{self.pipeline_name}{extension}") + + unload_pipeline_modules(self.pipeline_name) + sys.modules.update(self.modules) + + def cleanup(self) -> None: + if self.backup_dir is not None: + shutil.rmtree(self.backup_dir, ignore_errors=True) + + +async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0915 + pipeline_name: str, + prepare: Callable[[], Awaitable[PreparedPipeline]], + *, + app: FastAPI | None, + overwrite: bool, + remove_files_before_prepare: bool, + cleanup_files_on_overwrite: bool, +) -> dict[str, str]: + """Prepare independently, then atomically publish or restore one pipeline.""" + dlog = log.bind( + pipeline_name=pipeline_name, + overwrite=overwrite, + deploy_concurrency=settings.deploy_concurrency.value, + ) + dlog.debug("Starting pipeline deployment transaction") + policy_lock = ( + _deployment_serial_lock if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED else nullcontext() + ) + snapshot: _DeploymentSnapshot | None = None + candidate: DurableDeployment | None = None + old_quiesced = False + registered = False + publication_started = False + async with policy_lock: + try: + async with _deployment_publication_lock: + if pipeline_name in _deployments_in_progress: + msg = f"Pipeline '{pipeline_name}' is already being deployed" + raise PipelineAlreadyExistsError(msg) + snapshot = _DeploymentSnapshot.capture(pipeline_name, app) + if snapshot.wrapper is not None and not overwrite: + msg = f"Pipeline '{pipeline_name}' already exists" + raise PipelineAlreadyExistsError(msg) + if snapshot.deployment is not None: + await snapshot.deployment.quiesce() + old_quiesced = True + await snapshot.deployment.close() + nonterminal = await _durable_nonterminal_count(snapshot.deployment) + if nonterminal: + # ponytail: reject all live-work replacements; isolated candidate imports can restore + # same-revision hot swaps if that capability becomes necessary. + dlog.bind(nonterminal=nonterminal).debug( + "Rejected pipeline replacement while durable work remains" + ) + msg = ( + f"Pipeline '{pipeline_name}' has {nonterminal} nonterminal durable execution(s); " + "complete or cancel them before replacing it" + ) + raise PipelineAlreadyExistsError(msg) + _deployments_in_progress.add(pipeline_name) + registered = True + + if remove_files_before_prepare: + remove_pipeline_files(pipeline_name, settings.pipelines_dir) + prepared = await prepare() + candidate = durable_runtime.create_deployment(prepared.name, prepared.wrapper) + if candidate is not None and durable_runtime.started: + await candidate.prepare() + dlog.bind(durable=candidate is not None).debug("Prepared pipeline deployment candidate") + + async with _deployment_publication_lock: + snapshot.refresh_publication() + publication_started = True + result = commit_prepared_pipeline( + prepared, + app=app, + overwrite=overwrite, + cleanup_files_on_overwrite=cleanup_files_on_overwrite, + _durable_deployment=candidate, + ) + durable_runtime.install_deployment(prepared.name, candidate) + if candidate is not None and durable_runtime.started: + candidate.activate() + dlog.debug("Published pipeline deployment") + return result + except BaseException as error: + dlog.bind(error_type=type(error).__name__).debug("Pipeline deployment transaction failed") + if snapshot is not None and (registered or old_quiesced): + async with _deployment_publication_lock: + try: + if registered: + if candidate is not None: + candidate.deactivate() + await candidate.close() + snapshot.restore_files_and_modules() + if publication_started: + snapshot.restore_publication() + finally: + if old_quiesced and snapshot.deployment is not None and durable_runtime.started: + await snapshot.deployment.start() + raise + finally: + if registered: + async with _deployment_publication_lock: + _deployments_in_progress.discard(pipeline_name) + if snapshot is not None: + snapshot.cleanup() + + async def deploy_pipeline_yaml_async( pipeline_name: str, source_code: str, @@ -83,17 +247,22 @@ async def deploy_pipeline_yaml_async( """ Async wrapper that offloads ``deploy_pipeline_yaml`` off the event loop. - Respects the ``deploy_concurrency`` setting: when *serialized* (default), a - global lock ensures only one deploy/undeploy runs at a time; when *parallel*, - the call runs in a thread without serialization. + Preparation respects ``deploy_concurrency``. Publication and lifecycle are + always serialized so registry, route, file, and runtime state change together. """ - return await _offload( - deploy_pipeline_yaml, - pipeline_name=pipeline_name, - source_code=source_code, + save_file = True if options is None else bool(options.get("save_file", True)) + return await _deploy_prepared_pipeline_async( + pipeline_name, + lambda: _offload( + prepare_pipeline_yaml, + pipeline_name=pipeline_name, + source_code=source_code, + options=options, + ), app=app, overwrite=overwrite, - options=options, + remove_files_before_prepare=overwrite and save_file, + cleanup_files_on_overwrite=overwrite and not save_file, ) @@ -105,19 +274,60 @@ async def deploy_pipeline_files_async( overwrite: bool = False, ) -> dict[str, str]: """Async wrapper that offloads ``deploy_pipeline_files`` off the event loop.""" - return await _offload( - deploy_pipeline_files, - pipeline_name=pipeline_name, - files=files, + return await _deploy_prepared_pipeline_async( + pipeline_name, + lambda: _offload( + prepare_pipeline_files, + pipeline_name=pipeline_name, + files=files, + save_files=save_files, + ), app=app, - save_files=save_files, overwrite=overwrite, + remove_files_before_prepare=overwrite and save_files, + cleanup_files_on_overwrite=overwrite and not save_files, + ) + + +async def undeploy_pipeline_async( + pipeline_name: str, + app: FastAPI | None = None, +) -> None: + """Atomically unpublish a pipeline before stopping its owned resources.""" + policy_lock = ( + _deployment_serial_lock if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED else nullcontext() ) + async with policy_lock, _deployment_publication_lock: + if pipeline_name in _deployments_in_progress: + raise HTTPException(status_code=409, detail=f"Pipeline '{pipeline_name}' is being deployed") + if registry.get(pipeline_name) is None: + raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") + deployment = durable_runtime.current_deployment(pipeline_name) + if deployment is not None: + await deployment.quiesce() + try: + await deployment.close() + if await _durable_nonterminal_count(deployment): + raise HTTPException( + status_code=409, + detail=( + f"Pipeline '{pipeline_name}' has durable executions that must be completed or canceled " + "before undeployment" + ), + ) + except BaseException: + if durable_runtime.started: + await deployment.start() + raise + + undeploy_pipeline(pipeline_name=pipeline_name, app=app) + durable_runtime.install_deployment(pipeline_name, None) -async def undeploy_pipeline_async(pipeline_name: str, app: FastAPI | None = None) -> None: - """Async wrapper that offloads ``undeploy_pipeline`` off the event loop.""" - return await _offload(undeploy_pipeline, pipeline_name=pipeline_name, app=app) +async def _durable_nonterminal_count(deployment: DurableDeployment) -> int: + """Return the authoritative number of executions that would be stranded.""" + counts = await deployment.store.operational_counts() + return int(counts.get("nonterminal", 0)) def _is_single_yaml_file(files: dict[str, str]) -> bool: @@ -442,10 +652,7 @@ async def _handle_request(run_req: BaseModel) -> Response | BaseModel: if response_model is None: return cast(Response | BaseModel, traced_result) - # response_model is built dynamically via create_model(..., result=(...)); cast to Any so ty does not - # treat the real `result` field as an extra, discarded argument. - response_instance = cast("Any", response_model)(result=traced_result) - return cast(Response | BaseModel, response_instance) + return cast(Response | BaseModel, cast(Any, response_model)(result=traced_result)) @handle_pipeline_exceptions() async def run_endpoint_with_files( @@ -466,6 +673,7 @@ def add_pipeline_api_route( pipeline_wrapper: BasePipelineWrapper, *, _defer_openapi_rebuild: bool = False, + _durable_deployment: DurableDeployment | None = None, ) -> None: """ Create or replace the wrapper-based pipeline run endpoint at /{pipeline_name}/run. @@ -499,6 +707,13 @@ def add_pipeline_api_route( f"Pipeline '{pipeline_name}' does not implement `run_api` or `run_api_async`. " f"Skipping /{pipeline_name}/run API route creation." ) + _add_durable_api_routes( + app, + pipeline_name, + pipeline_wrapper, + deployment=_durable_deployment, + _defer_openapi_rebuild=_defer_openapi_rebuild, + ) return docstring_content = inspect.getdoc(run_method_to_inspect) or "" @@ -542,6 +757,14 @@ def add_pipeline_api_route( app.add_api_route(**route_kwargs) + _add_durable_api_routes( + app, + pipeline_name, + pipeline_wrapper, + deployment=_durable_deployment, + _defer_openapi_rebuild=True, + ) + registry.update_metadata( pipeline_name, { @@ -570,6 +793,7 @@ def _register_prepared_pipeline( extra_metadata: dict[str, Any] | None = None, *, _defer_openapi_rebuild: bool = False, + _durable_deployment: DurableDeployment | None = None, ) -> dict[str, str]: """ Register a prepared pipeline wrapper and optionally add its API route. @@ -636,7 +860,13 @@ def _register_prepared_pipeline( # Create API route if app is provided if app: - add_pipeline_api_route(app, pipeline_name, pipeline_wrapper, _defer_openapi_rebuild=_defer_openapi_rebuild) + add_pipeline_api_route( + app, + pipeline_name, + pipeline_wrapper, + _defer_openapi_rebuild=_defer_openapi_rebuild, + _durable_deployment=_durable_deployment, + ) return {"name": pipeline_name} @@ -737,6 +967,7 @@ def commit_prepared_pipeline( *, _defer_openapi_rebuild: bool = False, cleanup_files_on_overwrite: bool = True, + _durable_deployment: DurableDeployment | None = None, ) -> dict[str, str]: """ Commit a prepared pipeline to the registry and (optionally) add its route. @@ -779,6 +1010,7 @@ def commit_prepared_pipeline( app=app, extra_metadata=prepared.extra_metadata, _defer_openapi_rebuild=_defer_openapi_rebuild, + _durable_deployment=_durable_deployment, ) @@ -994,13 +1226,13 @@ def undeploy_pipeline(pipeline_name: str, app: FastAPI | None = None) -> None: unload_pipeline_modules(pipeline_name) if app: - # Remove API routes for the pipeline - # All pipelines have a run endpoint at //run - routes_to_remove = [ - route for route in app.routes if isinstance(route, APIRoute) and route.path == f"/{pipeline_name}/run" + route_paths = { + f"/{pipeline_name}/run", + *(f"/{pipeline_name}{suffix}" for suffix in _DURABLE_ROUTE_SUFFIXES), + } + app.routes[:] = [ + route for route in app.routes if not (isinstance(route, APIRoute) and route.path in route_paths) ] - for route in routes_to_remove: - app.routes.remove(route) # Invalidate OpenAPI cache app.openapi_schema = None diff --git a/src/hayhooks/server/utils/mcp_utils.py b/src/hayhooks/server/utils/mcp_utils.py index 382e8441..04fd41ed 100644 --- a/src/hayhooks/server/utils/mcp_utils.py +++ b/src/hayhooks/server/utils/mcp_utils.py @@ -12,6 +12,7 @@ from starlette.routing import Mount, Route from starlette.types import Receive, Scope, Send +from hayhooks.durable.runtime import durable_runtime from hayhooks.server.logger import log from hayhooks.server.pipelines.registry import registry from hayhooks.server.routers.deploy import PipelineFilesRequest @@ -299,11 +300,15 @@ async def handle_streamable_http(scope: Scope, receive: Receive, send: Send) -> @asynccontextmanager async def lifespan(app: Starlette) -> AsyncIterator[None]: # noqa: ARG001 async with session_manager.run(): - log.info("Hayhooks MCP server started") try: + await durable_runtime.start() + log.info("Hayhooks MCP server started") yield finally: - log.info("Hayhooks MCP server shutting down...") + try: + await durable_runtime.close() + finally: + log.info("Hayhooks MCP server shutting down...") async def handle_sse(request): async with sse.connect_sse(request.scope, request.receive, request._send) as streams: diff --git a/src/hayhooks/server/utils/module_loader.py b/src/hayhooks/server/utils/module_loader.py index 9883ca18..543f56ff 100644 --- a/src/hayhooks/server/utils/module_loader.py +++ b/src/hayhooks/server/utils/module_loader.py @@ -12,6 +12,7 @@ from types import ModuleType from typing import NoReturn +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode from hayhooks.server.exceptions import PipelineModuleLoadError, PipelineWrapperError from hayhooks.server.logger import log from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper @@ -249,6 +250,8 @@ def _set_method_implementation_flags(pipeline_wrapper: BasePipelineWrapper) -> N methods_to_check = [ ("_is_run_api_implemented", "run_api"), ("_is_run_api_async_implemented", "run_api_async"), + ("_is_run_durable_implemented", "run_durable"), + ("_is_run_durable_async_implemented", "run_durable_async"), ("_is_run_chat_completion_implemented", "run_chat_completion"), ("_is_run_chat_completion_async_implemented", "run_chat_completion_async"), ("_is_run_response_implemented", "run_response"), @@ -296,10 +299,16 @@ def _validate_run_methods(pipeline_wrapper: BasePipelineWrapper) -> None: Raises: PipelineWrapperError: If no run methods are implemented. """ - has_run_method = any( + if pipeline_wrapper._is_run_durable_implemented and pipeline_wrapper._is_run_durable_async_implemented: + msg = "Implement at most one of run_durable and run_durable_async" + raise PipelineWrapperError(msg) + + has_run_method = durable_authoring_mode(pipeline_wrapper) is DurableAuthoringMode.MANAGED_AGENT or any( [ pipeline_wrapper._is_run_api_implemented, pipeline_wrapper._is_run_api_async_implemented, + pipeline_wrapper._is_run_durable_implemented, + pipeline_wrapper._is_run_durable_async_implemented, pipeline_wrapper._is_run_chat_completion_implemented, pipeline_wrapper._is_run_chat_completion_async_implemented, pipeline_wrapper._is_run_response_implemented, diff --git a/tests/test_cli.py b/tests/test_cli.py index 1fa8b8ce..861bfeee 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -22,6 +22,10 @@ def set_default_settings(): settings.dashboard_enabled = False settings.dashboard_path = "/dashboard" settings.dashboard_dist_dir = "dashboard/dist" + settings.durable_execution_concurrency = 1 + settings.a2a_task_store = "memory" + settings.a2a_redis_url = "redis://localhost:6379/0" + settings.a2a_redis_key_prefix = "hayhooks:a2a" def test_run_command_with_reload(monkeypatch): @@ -139,7 +143,8 @@ def fake_prepare_tracing_dashboard_assets(dashboard_dist_dir: str) -> str: def test_a2a_run_debug_enables_tracebacks(monkeypatch): import uvicorn - from hayhooks.server.utils import a2a_utils, deploy_utils + from hayhooks.server.a2a import app as a2a_app + from hayhooks.server.utils import deploy_utils from hayhooks.settings import settings calls = [] @@ -155,11 +160,7 @@ def fake_deploy_pipelines() -> None: monkeypatch.setattr(uvicorn, "run", fake_uvicorn_run) monkeypatch.setattr(deploy_utils, "deploy_pipelines", fake_deploy_pipelines) - monkeypatch.setattr(a2a_utils, "create_a2a_app", fake_create_a2a_app) - # Neutralize the optional-dependency guard so this CLI-plumbing test runs - # even when the optional 'a2a' package isn't installed. - monkeypatch.setattr(a2a_utils.a2a_import, "check", lambda: None) - + monkeypatch.setattr(a2a_app, "create_a2a_app", fake_create_a2a_app) settings.show_tracebacks = False result = runner.invoke(hayhooks_cli, ["a2a", "run", "--debug", "--pipelines-dir", "dummy_pipelines"]) @@ -168,6 +169,89 @@ def fake_deploy_pipelines() -> None: assert calls, "uvicorn.run was not called" +def test_a2a_run_sets_builtin_redis_task_store(monkeypatch): + import uvicorn + + from hayhooks.server.a2a import app as a2a_app + from hayhooks.server.utils import deploy_utils + + monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) + + result = runner.invoke( + hayhooks_cli, + [ + "a2a", + "run", + "--task-store", + "redis", + "--a2a-redis-url", + "redis://localhost:6379/4", + "--a2a-redis-key-prefix", + "demo:a2a", + "--pipelines-dir", + "dummy_pipelines", + ], + ) + + assert result.exit_code == 0, result.output + assert settings.a2a_task_store == "redis" + assert settings.a2a_redis_url == "redis://localhost:6379/4" + assert settings.a2a_redis_key_prefix == "demo:a2a" + + +def test_a2a_run_sets_durable_execution_concurrency(monkeypatch): + import uvicorn + + from hayhooks.server.a2a import app as a2a_app + from hayhooks.server.utils import deploy_utils + + monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) + + result = runner.invoke( + hayhooks_cli, + ["a2a", "run", "--durable-execution-concurrency", "3", "--pipelines-dir", "dummy_pipelines"], + ) + + assert result.exit_code == 0, result.output + assert settings.durable_execution_concurrency == 3 + + +def test_a2a_run_sets_durable_execution_store_configuration(monkeypatch): + import uvicorn + + from hayhooks.server.a2a import app as a2a_app + from hayhooks.server.utils import deploy_utils + + monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) + + result = runner.invoke( + hayhooks_cli, + [ + "a2a", + "run", + "--execution-store", + "redis", + "--execution-redis-url", + "redis://localhost:6379/5", + "--execution-redis-key-prefix", + "demo:durable", + "--pipelines-dir", + "dummy_pipelines", + ], + ) + + assert result.exit_code == 0, result.output + assert settings.durable_store == "redis" + assert settings.durable_redis_url == "redis://localhost:6379/5" + assert settings.durable_redis_key_prefix == "demo:durable" + + def test_status_command(monkeypatch): """ Test the status command. We patch make_request (used in the status command) diff --git a/tests/test_deploy_performance.py b/tests/test_deploy_performance.py index 4e02688e..722c1272 100644 --- a/tests/test_deploy_performance.py +++ b/tests/test_deploy_performance.py @@ -1,3 +1,5 @@ +import asyncio +import threading from pathlib import Path from unittest.mock import MagicMock @@ -104,31 +106,22 @@ async def test_undeploy_pipeline_async(monkeypatch): @pytest.mark.asyncio -async def test_serialized_policy_wraps_with_lock(monkeypatch): - monkeypatch.setattr(settings, "deploy_concurrency", DeployConcurrencyPolicy.SERIALIZED) +async def test_parallel_policy_prepares_different_pipelines_concurrently(monkeypatch): + monkeypatch.setattr(settings, "deploy_concurrency", DeployConcurrencyPolicy.PARALLEL) + original_prepare = prepare_pipeline_yaml + both_preparing = threading.Barrier(2) - lock_calls = [] - original_with_lock = __import__( - "hayhooks.server.utils.deploy_utils", fromlist=["_with_deploy_lock"] - )._with_deploy_lock + def synchronized_prepare(*args, **kwargs): + both_preparing.wait(timeout=2) + return original_prepare(*args, **kwargs) - def tracking_with_lock(func): - lock_calls.append(func.__name__) - return original_with_lock(func) + monkeypatch.setattr("hayhooks.server.utils.deploy_utils.prepare_pipeline_yaml", synchronized_prepare) - monkeypatch.setattr( - "hayhooks.server.utils.deploy_utils._with_deploy_lock", - tracking_with_lock, + await asyncio.gather( + deploy_pipeline_yaml_async("parallel_a", SAMPLE_YAML, options={"save_file": False}), + deploy_pipeline_yaml_async("parallel_b", SAMPLE_YAML, options={"save_file": False}), ) - await deploy_pipeline_yaml_async( - pipeline_name="lock_test", - source_code=SAMPLE_YAML, - options={"save_file": False}, - ) - - assert lock_calls == ["deploy_pipeline_yaml"] - def test_defer_openapi_rebuild_skips_setup(): mock_app = MagicMock(spec=FastAPI) @@ -156,7 +149,7 @@ def test_no_defer_calls_setup(): options={"save_file": False}, ) - mock_app.setup.assert_called() + mock_app.setup.assert_called_once() def test_rebuild_openapi(): diff --git a/tests/test_deploy_utils.py b/tests/test_deploy_utils.py index 1869b946..c5401337 100644 --- a/tests/test_deploy_utils.py +++ b/tests/test_deploy_utils.py @@ -4,7 +4,7 @@ import sys from collections.abc import AsyncGenerator, Callable, Generator from pathlib import Path -from typing import Any +from typing import Any, Literal import docstring_parser import pytest @@ -448,6 +448,26 @@ def sample_func_no_doc() -> int: assert "result" in schema["required"] +def test_callable_models_resolve_postponed_annotations(): + def func(action: Literal["start", "status"], execution_id: str | None = None) -> dict[str, Any]: + return {} + + func.__annotations__ = { + "action": 'Literal["start", "status"]', + "execution_id": "str | None", + "return": "dict[str, Any]", + } + docstring = docstring_parser.parse("") + + request_model = create_request_model_from_callable(func, "Postponed", docstring) + response_model = create_response_model_from_callable(func, "Postponed", docstring) + + request_schema = request_model.model_json_schema() + assert request_schema["properties"]["action"]["enum"] == ["start", "status"] + assert request_schema["properties"]["execution_id"]["anyOf"] == [{"type": "string"}, {"type": "null"}] + assert response_model.model_json_schema()["properties"]["result"]["type"] == "object" + + @pytest.mark.parametrize( "return_type", [ @@ -586,7 +606,8 @@ def setup(self): with pytest.raises( PipelineWrapperError, match=re.escape( - "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, run_response, or run_response_async must be implemented" + "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, run_response, " + "or run_response_async must be implemented" ), ): create_pipeline_wrapper_instance(module) diff --git a/tests/test_durable_deployment_lifecycle.py b/tests/test_durable_deployment_lifecycle.py new file mode 100644 index 00000000..4be2faa6 --- /dev/null +++ b/tests/test_durable_deployment_lifecycle.py @@ -0,0 +1,348 @@ +import importlib.metadata +import time + +import pytest +from fastapi.testclient import TestClient + +from hayhooks.durable.runtime import durable_runtime +from hayhooks.server.app import create_app +from hayhooks.server.pipelines.registry import registry +from hayhooks.settings import settings + +pytestmark = pytest.mark.skipif( + not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" +) + + +def _durable_source(*, field: str, increment: int, revision: str, result_field: str = "value") -> str: + return f""" +from haystack import Pipeline +from pydantic import BaseModel +from hayhooks import BasePipelineWrapper, DurableContext + +class Request(BaseModel): + {field}: int + +class Result(BaseModel): + {result_field}: int + +class PipelineWrapper(BasePipelineWrapper): + durable_revision = "{revision}" + + def setup(self): + self.pipeline = Pipeline() + + async def run_durable_async(self, context: DurableContext, request: Request) -> Result: + return Result({result_field}=request.{field} + {increment}) +""" + + +def _api_source(*, increment: int = 1) -> str: + return f""" +from haystack import Pipeline +from hayhooks import BasePipelineWrapper + +class PipelineWrapper(BasePipelineWrapper): + def setup(self): + self.pipeline = Pipeline() + + def run_api(self, value: int) -> int: + return value + {increment} +""" + + +def _waiting_source(*, revision: str) -> str: + return f""" +from haystack import Pipeline +from pydantic import BaseModel +from hayhooks import BasePipelineWrapper, DurableContext + +class Request(BaseModel): + value: int + +class PipelineWrapper(BasePipelineWrapper): + durable_revision = "{revision}" + + def setup(self): + self.pipeline = Pipeline() + + async def run_durable_async(self, context: DurableContext, request: Request) -> dict: + if context.resume_input is None: + await context.suspend({{"kind": "input", "message": "waiting"}}) + return {{"value": request.value}} +""" + + +def _blocking_source() -> str: + return """ +import threading +from haystack import Pipeline +from pydantic import BaseModel +from hayhooks import BasePipelineWrapper, DurableContext + +class Request(BaseModel): + value: int + +class PipelineWrapper(BasePipelineWrapper): + durable_revision = "blocking" + + def setup(self): + self.pipeline = Pipeline() + self.started = threading.Event() + self.release = threading.Event() + + def run_durable(self, context: DurableContext, request: Request) -> dict: + self.started.set() + assert self.release.wait(timeout=5) + return {"value": request.value} +""" + + +def _deploy(client: TestClient, source: str, *, overwrite: bool = False): + return client.post( + "/deploy_files", + json={ + "name": "job", + "files": {"pipeline_wrapper.py": source}, + "save_files": False, + "overwrite": overwrite, + }, + ) + + +def _wait_for_completion(client: TestClient, response) -> dict: + body = response.json() + for _ in range(200): + result = client.get(body["links"]["self"]) + if result.json()["status"] == "completed": + return result.json() + time.sleep(0.01) + pytest.fail("durable execution did not complete") + + +@pytest.fixture(autouse=True) +def _isolated_runtime(monkeypatch, tmp_path): + registry.clear() + monkeypatch.setattr(settings, "pipelines_dir", str(tmp_path)) + monkeypatch.setattr(settings, "durable_store", "memory") + monkeypatch.setattr(settings, "durable_poll_interval", 0.05) + yield + registry.clear() + + +def test_undeploy_removes_entire_durable_route_family() -> None: + app = create_app() + with TestClient(app) as client: + assert _deploy(client, _durable_source(field="value", increment=1, revision="first")).status_code == 200 + assert client.post("/undeploy/job").status_code == 200 + assert client.post("/job/run-durable", json={"value": 2}).status_code == 404 + assert client.get("/job/executions/missing").status_code == 404 + assert client.post("/job/executions/missing/cancel").status_code == 404 + assert client.post("/job/executions/missing/resume").status_code == 404 + assert durable_runtime.current_deployment("job") is None + + +def test_undeploy_refuses_to_strand_waiting_execution() -> None: + app = create_app() + with TestClient(app) as client: + assert _deploy(client, _waiting_source(revision="first")).status_code == 200 + submitted = client.post("/job/run-durable", json={"value": 1}) + for _ in range(100): + if client.get(submitted.json()["links"]["self"]).json()["status"] == "waiting": + break + time.sleep(0.01) + else: + pytest.fail("execution did not enter waiting before undeploy") + + blocked = client.post("/undeploy/job") + assert blocked.status_code == 409 + assert "completed or canceled" in blocked.json()["detail"] + assert client.get(submitted.json()["links"]["self"]).json()["status"] == "waiting" + assert client.post(f"{submitted.json()['links']['self']}/cancel").status_code == 202 + assert client.post("/undeploy/job").status_code == 200 + assert client.get(submitted.json()["links"]["self"]).status_code == 404 + + +def test_durable_overwrite_routes_bind_new_model_runner_and_revision() -> None: + app = create_app() + with TestClient(app) as client: + assert ( + _deploy( + client, + _durable_source( + field="old_value", + increment=1, + revision="first", + result_field="old_result", + ), + ).status_code + == 200 + ) + first = client.post("/job/run-durable", json={"old_value": 2}) + assert _wait_for_completion(client, first)["result"] == {"old_result": 3} + + assert ( + _deploy( + client, + _durable_source( + field="new_value", + increment=20, + revision="second", + result_field="new_result", + ), + overwrite=True, + ).status_code + == 200 + ) + assert client.get(first.json()["links"]["self"]).json()["result"] == {"old_result": 3} + assert client.post("/job/run-durable", json={"old_value": 2}).status_code == 422 + second = client.post("/job/run-durable", json={"new_value": 2}) + assert second.status_code == 202 + assert _wait_for_completion(client, second)["result"] == {"new_result": 22} + + +@pytest.mark.parametrize("operation", ["overwrite", "undeploy"]) +def test_failed_durable_preflight_restarts_existing_deployment(monkeypatch, operation: str) -> None: + app = create_app() + with TestClient(app) as client: + source = _durable_source(field="value", increment=1, revision="first") + assert _deploy(client, source).status_code == 200 + deployment = durable_runtime.current_deployment("job") + assert deployment is not None + + async def fail_counts(): + msg = "redis unavailable" + raise ConnectionError(msg) + + monkeypatch.setattr(deployment.store, "operational_counts", fail_counts) + if operation == "overwrite": + assert _deploy(client, source, overwrite=True).status_code == 500 + else: + with pytest.raises(ConnectionError, match="redis unavailable"): + client.post("/undeploy/job") + assert deployment.manager.accepting + assert client.post("/job/run-durable", json={"value": 1}).status_code == 202 + + +def test_overwrite_refuses_to_prepare_while_old_execution_is_waiting(tmp_path) -> None: + app = create_app() + with TestClient(app) as client: + assert _deploy(client, _waiting_source(revision="first")).status_code == 200 + submitted = client.post("/job/run-durable", json={"value": 2}) + url = submitted.json()["links"]["self"] + for _ in range(100): + waiting = client.get(url) + if waiting.json()["status"] == "waiting": + break + time.sleep(0.01) + else: + pytest.fail("old-revision execution did not enter waiting") + + preparation_marker = tmp_path / "replacement-prepared" + replacement_source = _durable_source(field="value", increment=20, revision="second").replace( + "from haystack import Pipeline", + f"from haystack import Pipeline\nfrom pathlib import Path\nPath({str(preparation_marker)!r}).touch()", + ) + replacement = _deploy( + client, + replacement_source, + overwrite=True, + ) + assert replacement.status_code == 409 + assert "complete or cancel" in replacement.json()["detail"] + assert not preparation_marker.exists() + assert client.get(url).json()["status"] == "waiting" + assert client.post(f"{url}/cancel").status_code == 202 + assert ( + _deploy( + client, + _durable_source(field="value", increment=20, revision="second"), + overwrite=True, + ).status_code + == 200 + ) + + +def test_undeploy_refuses_to_strand_thread_backed_work(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) + app = create_app() + with TestClient(app) as client: + assert _deploy(client, _blocking_source()).status_code == 200 + old_wrapper = registry.get("job") + assert old_wrapper is not None + submitted = client.post("/job/run-durable", json={"value": 2}) + assert old_wrapper.started.wait(timeout=1) + + assert client.post("/undeploy/job").status_code == 409 + old_wrapper.release.set() + for _ in range(100): + if client.get(submitted.json()["links"]["self"]).json()["status"] == "completed": + break + time.sleep(0.01) + else: + pytest.fail("durable work did not complete after rejected undeploy") + assert client.post("/undeploy/job").status_code == 200 + + +def test_durable_to_non_durable_overwrite_removes_control_routes() -> None: + app = create_app() + with TestClient(app) as client: + assert _deploy(client, _durable_source(field="value", increment=1, revision="first")).status_code == 200 + submitted = client.post("/job/run-durable", json={"value": 1}) + execution_id = submitted.json()["execution_id"] + _wait_for_completion(client, submitted) + + assert _deploy(client, _api_source(increment=5), overwrite=True).status_code == 200 + assert client.post("/job/run-durable", json={"value": 2}).status_code == 404 + assert client.get(f"/job/executions/{execution_id}").status_code == 404 + assert client.post("/job/run", json={"value": 2}).json() == {"result": 7} + + +def test_failed_commit_does_not_irreversibly_retire_old_durable_work(monkeypatch) -> None: + app = create_app() + with TestClient(app) as client: + source = _durable_source(field="value", increment=1, revision="first") + assert _deploy(client, source).status_code == 200 + submitted = client.post("/job/run-durable", json={"value": 1}) + url = submitted.json()["links"]["self"] + _wait_for_completion(client, submitted) + + def fail_commit(*_args, **_kwargs): + msg = "commit fault" + raise RuntimeError(msg) + + monkeypatch.setattr("hayhooks.server.utils.deploy_utils.commit_prepared_pipeline", fail_commit) + failed = _deploy(client, source, overwrite=True) + + assert failed.status_code == 500 + restored = client.get(url) + assert restored.status_code == 200 + assert restored.json()["status"] == "completed" + + +class _FailingStore: + async def initialize(self): + msg = "redis unavailable" + raise ConnectionError(msg) + + async def close(self): + return None + + +class _FailingProvider: + def create_execution_store(self, _deployment_name): + return _FailingStore() + + async def close(self): + return None + + +def test_store_initialization_failure_never_publishes_candidate(monkeypatch) -> None: + monkeypatch.setattr(durable_runtime, "provider", _FailingProvider()) + app = create_app() + with TestClient(app) as client: + failed = _deploy(client, _durable_source(field="value", increment=1, revision="first")) + + assert failed.status_code == 500 + assert registry.get("job") is None + assert client.post("/job/run-durable", json={"value": 1}).status_code == 404 diff --git a/tests/test_it_deploy_files.py b/tests/test_it_deploy_files.py index 7ae8a233..9c467cc0 100644 --- a/tests/test_it_deploy_files.py +++ b/tests/test_it_deploy_files.py @@ -154,8 +154,8 @@ def test_deploy_files_missing_required_methods(client, deploy_files) -> None: err_body: dict[str, Any] = response.json() assert ( - "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, run_response, or run_response_async must be implemented" - in err_body["detail"] + "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, run_response, " + "or run_response_async must be implemented" in err_body["detail"] ) From 9b2e80a1123990de86a728678baf1cb82a48a17b Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:09:51 +0200 Subject: [PATCH 07/28] feat(examples): add portable durable pipeline demo --- examples/durable-compose.yaml | 18 +++ examples/durable_execution/README.md | 70 ++++++++++ examples/durable_execution/demo.py | 129 +++++++++++++++++ .../pipelines/durable_job/pipeline_wrapper.py | 132 ++++++++++++++++++ 4 files changed, 349 insertions(+) create mode 100644 examples/durable-compose.yaml create mode 100644 examples/durable_execution/README.md create mode 100644 examples/durable_execution/demo.py create mode 100644 examples/durable_execution/pipelines/durable_job/pipeline_wrapper.py diff --git a/examples/durable-compose.yaml b/examples/durable-compose.yaml new file mode 100644 index 00000000..67cc5304 --- /dev/null +++ b/examples/durable-compose.yaml @@ -0,0 +1,18 @@ +services: + redis: + image: redis:7.4-alpine + command: + - redis-server + - --appendonly + - "yes" + - --appendfsync + - everysec + - --maxmemory-policy + - noeviction + ports: + - "6379:6379" + volumes: + - durable-redis:/data + +volumes: + durable-redis: diff --git a/examples/durable_execution/README.md b/examples/durable_execution/README.md new file mode 100644 index 00000000..83d808a3 --- /dev/null +++ b/examples/durable_execution/README.md @@ -0,0 +1,70 @@ +# Durable document-preparation Pipeline + +This is the canonical REST durable-execution example. The wrapper owns typed +application behavior; Hayhooks owns records, the Redis runnable queue, fenced workers, +checkpoints, retry delay, waiting/resume, cancellation, and retention. +The wrapper declares the stable `durable_revision` required by the engine; use +an image digest or Git SHA instead for production releases. + +The real Haystack Pipeline cleans and splits a document. Hayhooks persists a +`PipelineSnapshot` after `clean`, so a crash during the later delay resumes +from that checkpoint without repeating the completed cleaning step. + +The included `httpx` client pauses five seconds between requests, prints every +URL and response with Rich, automatically handles the retry and approval, and +stays alive while you restart Hayhooks. + +Run each command from the repository root. This is a local reliability +demonstration, not a production Redis configuration. + +1. Start Redis and install the example dependencies. + +```bash +docker compose -f examples/durable-compose.yaml up -d && python -m pip install -e ".[durable]" httpx +``` + +2. Set the durable settings. The five-second lease keeps recovery short. + +```bash +export HAYHOOKS_DURABLE_REDIS_URL=redis://localhost:6379/0 HAYHOOKS_DURABLE_LEASE_DURATION_MS=5000 HAYHOOKS_DURABLE_MAX_ATTEMPTS=4 +``` + +3. Open a first terminal and start Hayhooks. It prints the PID for the forced + crash. + +```bash +sh -c 'echo "Hayhooks PID: $$"; exec hayhooks run --pipelines-dir examples/durable_execution/pipelines' +``` + +4. Open a second terminal and run the client. It submits the document, shows + the intentional retry, waits for approval, approves it, and polls every five + seconds. + +```bash +python examples/durable_execution/demo.py +``` + +5. When the client prints “The clean checkpoint is persisted,” return to the + first terminal and press `Ctrl-C` to stop Hayhooks. The next client request + reports the expected connection failure and waits five seconds before trying + again. + +6. In the first terminal, start Hayhooks again with the same command. The + client detects recovery and prints the completed response. + +```bash +sh -c 'echo "Hayhooks PID: $$"; exec hayhooks run --pipelines-dir examples/durable_execution/pipelines' +``` + +The recovered Pipeline skips `clean`, repeats the interrupted `demo_delay`, +then runs `split`. Durable execution is at least once: if a Pipeline has an +external side effect, use the execution ID and logical step as its idempotency +key. + +Press Ctrl-C in the first terminal to stop Hayhooks, then stop Redis. Add `-v` +to remove the retained Redis volume before a clean rehearsal: + +```bash +docker compose -f examples/durable-compose.yaml down +# docker compose -f examples/durable-compose.yaml down -v +``` diff --git a/examples/durable_execution/demo.py b/examples/durable_execution/demo.py new file mode 100644 index 00000000..603cb820 --- /dev/null +++ b/examples/durable_execution/demo.py @@ -0,0 +1,129 @@ +"""Submit and follow the durable Pipeline recovery demonstration.""" + +from __future__ import annotations + +import argparse +import json +import time +import uuid +from collections.abc import Callable +from typing import Any +from urllib.parse import urljoin + +import httpx +from rich.console import Console +from rich.json import JSON +from rich.panel import Panel + +_PAUSE_SECONDS = 5 +_TERMINAL_STATUSES = {"completed", "failed", "canceled"} + + +def request(console: Console, client: httpx.Client, method: str, url: str, **kwargs: Any) -> httpx.Response | None: + console.print(Panel.fit(f"[bold cyan]{method}[/] {url}", title="Request", border_style="cyan")) + if payload := kwargs.get("json"): + console.print(JSON.from_data(payload)) + try: + response = client.request(method, url, **kwargs) + except httpx.HTTPError as error: + console.print(Panel(str(error), title="Connection failed", border_style="red")) + return None + try: + body = JSON.from_data(response.json()) + except json.JSONDecodeError: + body = response.text + console.print(Panel(body, title=f"{response.status_code} {response.url}", border_style="green")) + return response + + +def pause(console: Console) -> None: + console.print(f"[dim]Waiting {_PAUSE_SECONDS} seconds before the next request...[/]") + time.sleep(_PAUSE_SECONDS) + + +def poll_until( + console: Console, + client: httpx.Client, + execution_url: str, + matches: Callable[[dict[str, Any]], bool], +) -> dict[str, Any] | None: + while True: + response = request(console, client, "GET", execution_url) + if response is not None: + body = response.json() + if matches(body): + return body + if body["status"] in _TERMINAL_STATUSES: + return None + pause(console) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base-url", default="http://localhost:1416") + args = parser.parse_args() + base_url = args.base_url.rstrip("/") + execution_id = f"prepare-hayhooks-guide-{uuid.uuid4().hex[:12]}" + console = Console() + + with httpx.Client(timeout=2) as client: + while ( + submitted := request( + console, + client, + "POST", + f"{base_url}/durable_job/run-durable", + headers={"Idempotency-Key": execution_id}, + json={ + "documents": [ + { + "document_id": "hayhooks-guide", + "content": "Hayhooks durable Pipelines survive restarts.", + } + ], + "fail_first_attempt": True, + "require_approval": True, + "demo_delay_seconds": 30, + }, + ) + ) is None: + pause(console) + if submitted.is_error: + return 1 + links = submitted.json()["links"] + execution_url = urljoin(f"{base_url}/", links["self"]) + resume_url = urljoin(f"{base_url}/", links["resume"]) + + pause(console) + if poll_until(console, client, execution_url, lambda body: body["status"] == "waiting") is None: + return 1 + + pause(console) + resumed = request(console, client, "POST", resume_url, json={"approved": True}) + if resumed is None or resumed.is_error: + return 1 + + pause(console) + checkpoint = poll_until( + console, + client, + execution_url, + lambda body: any(event["kind"] == "demo_delay" for event in body["progress"]), + ) + if checkpoint is None: + return 1 + console.print( + Panel( + "Kill Hayhooks now, then restart it. This client will keep polling.", + title="The clean checkpoint is persisted", + border_style="yellow", + ) + ) + + pause(console) + completed = poll_until(console, client, execution_url, lambda body: body["status"] in _TERMINAL_STATUSES) + return 0 if completed is not None and completed["status"] == "completed" else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/examples/durable_execution/pipelines/durable_job/pipeline_wrapper.py b/examples/durable_execution/pipelines/durable_job/pipeline_wrapper.py new file mode 100644 index 00000000..998cee12 --- /dev/null +++ b/examples/durable_execution/pipelines/durable_job/pipeline_wrapper.py @@ -0,0 +1,132 @@ +"""A durable document-preparation Pipeline using real Haystack components.""" + +import time + +from haystack import Document, Pipeline, component +from haystack.components.preprocessors import DocumentCleaner, DocumentSplitter +from pydantic import BaseModel, Field + +from hayhooks import BasePipelineWrapper, DurableContext, current_durable_context + + +class SourceDocument(BaseModel): + """One raw document accepted by the durable indexing-preparation job.""" + + document_id: str = Field(min_length=1, max_length=128) + content: str = Field(min_length=1, max_length=10_000) + + +class DocumentPreparationRequest(BaseModel): + """Documents to clean and split into embedding-ready chunks.""" + + documents: list[SourceDocument] = Field(min_length=1, max_length=25) + fail_first_attempt: bool = False + require_approval: bool = False + demo_delay_seconds: float = Field(default=0, ge=0, le=300) + + +class ApprovalInput(BaseModel): + """Typed input accepted by the generated resume endpoint.""" + + approved: bool + + +class PreparedChunk(BaseModel): + """A compact, client-safe projection of a Haystack Document chunk.""" + + document_id: str + chunk_id: str + content: str + + +class DocumentPreparationResult(BaseModel): + """The chunks produced by the real Haystack preprocessing Pipeline.""" + + document_count: int + chunk_count: int + chunks: list[PreparedChunk] + + +@component +class DemoDelay: + """Optional pause after a persisted checkpoint, used only for restart demonstrations.""" + + @component.output_types(documents=list[Document]) + def run(self, documents: list[Document], seconds: float) -> dict[str, list[Document]]: + if seconds: + context = current_durable_context() + if context is not None: + context.report_progress_sync( + f"Checkpointed demo delay started for {seconds:g} seconds", + kind="demo_delay", + ) + time.sleep(seconds) + return {"documents": documents} + + +class PipelineWrapper(BasePipelineWrapper): + """Clean and chunk documents before a later embedding/indexing stage.""" + + durable_revision = "durable-document-preparation" + durable_resume_model = ApprovalInput + + def setup(self) -> None: + self.pipeline = Pipeline() + self.pipeline.add_component("clean", DocumentCleaner(remove_empty_lines=True)) + self.pipeline.add_component("demo_delay", DemoDelay()) + self.pipeline.add_component( + "split", + DocumentSplitter(split_by="word", split_length=80, split_overlap=10), + ) + self.pipeline.connect("clean.documents", "demo_delay.documents") + self.pipeline.connect("demo_delay.documents", "split.documents") + + async def run_durable_async( + self, context: DurableContext, request: DocumentPreparationRequest + ) -> DocumentPreparationResult: + await context.report_progress("Document preparation accepted", kind="accepted") + if request.fail_first_attempt and context.attempt == 1: + await context.report_progress("Demonstrating one bounded retry", kind="retry_demo") + await context.retry("Intentional first-attempt failure", delay=1) + + if request.require_approval and context.resume_input is None: + await context.suspend( + { + "kind": "approval", + "message": "Approve document preparation", + "expected_input_schema": ApprovalInput.model_json_schema(), + } + ) + if request.require_approval: + approval = ApprovalInput.model_validate(context.take_resume_input()) + if not approval.approved: + msg = "Document preparation was not approved" + raise ValueError(msg) + + documents = [ + Document(id=source.document_id, content=source.content, meta={"document_id": source.document_id}) + for source in request.documents + ] + outputs = await context.run_pipeline_async( + { + "clean": {"documents": documents}, + "demo_delay": {"seconds": request.demo_delay_seconds}, + }, + # The delay begins after clean has completed and the snapshot before + # demo_delay has been persisted, creating a reliable crash window. + checkpoint_at=["clean", "demo_delay", "split"], + ) + chunks = outputs["split"]["documents"] + await context.report_progress("Document preparation completed", kind="completed") + return DocumentPreparationResult( + document_count=len(documents), + chunk_count=len(chunks), + chunks=[ + PreparedChunk( + document_id=str(chunk.meta["document_id"]), + chunk_id=str(chunk.id), + content=chunk.content or "", + ) + for chunk in chunks + ], + ) From 80b23d9834795ac56f236c62eecb3a2d66debfd3 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:09:58 +0200 Subject: [PATCH 08/28] feat(examples): add long-running A2A agent demo --- .gitignore | 5 + examples/a2a_long_running/README.md | 204 ++++++++++++++++++ examples/a2a_long_running/demo.py | 149 +++++++++++++ .../long_running_agent/pipeline_wrapper.py | 175 +++++++++++++++ 4 files changed, 533 insertions(+) create mode 100644 examples/a2a_long_running/README.md create mode 100644 examples/a2a_long_running/demo.py create mode 100644 examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py diff --git a/.gitignore b/.gitignore index 048a3b8d..1345bd25 100644 --- a/.gitignore +++ b/.gitignore @@ -192,6 +192,11 @@ chainlit.md .env.local .env.*.local +# Local A2A long-running demo artifacts +.a2a-long-running-demo-session.json +/a2a-req +/a2a-res + # Temporary files *.tmp *.temp diff --git a/examples/a2a_long_running/README.md b/examples/a2a_long_running/README.md new file mode 100644 index 00000000..bf68b56d --- /dev/null +++ b/examples/a2a_long_running/README.md @@ -0,0 +1,204 @@ +# Durable long-running A2A Agent + +This example runs a Haystack Agent as a durable A2A task. Redis persists both +the execution checkpoints and the A2A task projection, so accepted work can +survive a Hayhooks restart. + +This is a local reliability demonstration, not a production Redis or ingress +configuration. Before a production evaluation, apply the +[controlled beta deployment profile](../../docs/advanced/durable-execution-operations.md#controlled-beta-deployment-profile). + +Run the setup commands from the repository root. The example requires `curl` +and `jq`. All durable examples use the same Compose service and Redis volume. + +```bash +docker compose -f examples/durable-compose.yaml up -d +python -m pip install -e ".[durable,a2a]" + +export OPENAI_API_KEY=... +export HAYHOOKS_DURABLE_REDIS_URL=redis://localhost:6379/0 +export HAYHOOKS_A2A_TASK_STORE=auto +export HAYHOOKS_EXAMPLE_TOOL_DELAY_SECONDS=15 +export HAYHOOKS_EXAMPLE_RECEIPT_DELAY_SECONDS=0 +export HAYHOOKS_DURABLE_LEASE_DURATION_MS=5000 +``` + +The five-second execution lease keeps the restart demonstration short. Active +workers heartbeat their claims; production deployments should tune this value +for their own workload and failure environment. + +Open a first terminal and start Hayhooks in the foreground. It prints the PID +used for the forced-crash demonstration: + +```bash +sh -c 'echo "Hayhooks PID: $$"; exec hayhooks a2a run --pipelines-dir examples/a2a_long_running/pipelines' +``` + +For a quick happy-path rehearsal, run the Rich demo client in a second +terminal. It creates unique message IDs, submits work, waits for approval, +approves it, and follows the task to completion: + +```bash +python examples/a2a_long_running/demo.py +``` + +Use the manual requests below when presenting the protocol or demonstrating +crash recovery and cancellation. + +In a second terminal, define a helper that sends the current request file and +prints the response: + +```bash +send_a2a_request() { + local request_file="$1" + local response_file="${2:-a2a-res}" + + curl -fsS http://localhost:1418/long_running_agent/ \ + -H 'content-type: application/json' \ + -H 'A2A-Version: 1.0' \ + --data-binary @"$request_file" \ + --output "$response_file" + + jq . "$response_file" +} +``` + +Submit detached work and save the returned task ID: + +```bash +jq -n '{ + "jsonrpc":"2.0", + "id":"submit", + "method":"SendMessage", + "params":{ + "message":{ + "messageId":"prepare-demo", + "role":"ROLE_USER", + "parts":[{ + "text":"Prepare this document for indexing. document_id: hayhooks-guide. content: Hayhooks durable A2A work survives restarts." + }] + }, + "configuration":{"returnImmediately":true} + } +}' > a2a-req + +send_a2a_request a2a-req a2a-res + +TASK_ID=$(jq -er '.result.task.id // .result.id' a2a-res) +``` + +Inspect the task with `GetTask`. Repeat these commands until its state is +`TASK_STATE_INPUT_REQUIRED`: + +```bash +jq -n --arg id "$TASK_ID" \ + '{"jsonrpc":"2.0","id":"poll","method":"GetTask","params":{"id":$id}}' \ + > a2a-req + +send_a2a_request a2a-req a2a-res +``` + +Approve the task with a follow-up A2A message. The persisted Agent checkpoint +already contains the original request: + +```bash +jq -n --arg task_id "$TASK_ID" '{ + "jsonrpc":"2.0", + "id":"resume", + "method":"SendMessage", + "params":{ + "message":{ + "messageId":"approval-demo", + "taskId":$task_id, + "role":"ROLE_USER", + "parts":[{"text":"Approved; proceed."}] + }, + "configuration":{"returnImmediately":true} + } +}' > a2a-req + +send_a2a_request a2a-req a2a-res +``` + +## Crash and replay + +Run the `GetTask` commands again until progress contains “Indexing effect +committed.” The tool has written its SQLite row and remains open for 15 seconds. + +In the second terminal, replace `` with the PID printed in the first +terminal to kill Hayhooks without graceful shutdown: + +```bash +kill -9 +``` + +Return to the first terminal and run the same foreground command again against +the same Redis: + +```bash +sh -c 'echo "Hayhooks PID: $$"; exec hayhooks a2a run --pipelines-dir examples/a2a_long_running/pipelines' +``` + +After about five seconds, run the same `GetTask` commands again. Redis reclaims +the interrupted execution and the Agent replays the tool from its previous +checkpoint. Continue inspecting until the task is terminal. + +The indexing tool uses the execution ID and document ID as a SQLite primary +key. The first attempt reports `side_effect_applied: true`; replay reports +`false` instead of inserting the same effect twice. This demonstrates the +at-least-once contract: Hayhooks protects execution state, while external +effects still require application-level idempotency. + +## Checkpoint efficiency: resume after indexing + +The Agent follows a realistic ingestion workflow: it first cleans and splits +the document, then makes a small catalog receipt update. The durable +`after_tool` hook checkpoints the completed indexing result before the Agent +asks the model to call the receipt tool. That means a crash while the receipt +is running resumes from the indexed document instead of repeating the expensive +cleaning and splitting work. + +For this demonstration, start a fresh task, approve it, and use no indexing +delay but a long receipt delay: + +```bash +export HAYHOOKS_EXAMPLE_TOOL_DELAY_SECONDS=0 +export HAYHOOKS_EXAMPLE_RECEIPT_DELAY_SECONDS=15 +``` + +Wait until task progress contains “Indexing is checkpointed; holding the +lightweight receipt update,” then kill and restart Hayhooks as in the crash +replay section. The recovered Agent receives the saved indexing result and +continues with `publish_indexing_receipt`; it must not call +`prepare_document_for_indexing` again. The indexing table retains the example's +idempotency proof; the receipt tool has no irreversible effect, so replaying it +is safe. + +## Cancellation + +For a cancellation demonstration, submit and approve a fresh task, then send: + +```bash +jq -n --arg id "$TASK_ID" \ + '{"jsonrpc":"2.0","id":"cancel","method":"CancelTask","params":{"id":$id}}' \ + > a2a-req + +send_a2a_request a2a-req a2a-res +``` + +Cancellation is cooperative. A synchronous tool may finish its current work +before the Agent reaches the next cancellation checkpoint. + +The SQLite database is a local effect-store demonstration. It defaults to the +operating system's temporary directory and survives a process restart on the +same machine. Set `HAYHOOKS_EXAMPLE_INDEX_DB` to a mounted path if the effect +must survive container replacement. + +Press Ctrl-C in the first terminal to stop Hayhooks, then stop Redis. The static message IDs are intended +for a clean rehearsal; remove the volume before starting again: + +```bash +docker compose -f examples/durable-compose.yaml down +# docker compose -f examples/durable-compose.yaml down -v +rm -f a2a-req a2a-res +``` diff --git a/examples/a2a_long_running/demo.py b/examples/a2a_long_running/demo.py new file mode 100644 index 00000000..f9413a47 --- /dev/null +++ b/examples/a2a_long_running/demo.py @@ -0,0 +1,149 @@ +"""Run the durable A2A example's submit, approval, and completion flow.""" + +from __future__ import annotations + +import argparse +import asyncio +import sys +import time + +import httpx +from a2a.client import A2AClientError, Client, ClientConfig, create_client +from a2a.helpers import new_text_message +from a2a.types import GetTaskRequest, Role, SendMessageConfiguration, SendMessageRequest, Task, TaskState +from rich.console import Console +from rich.panel import Panel +from rich.table import Table + +TERMINAL_STATES = { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_REJECTED, +} + + +class A2ADemoError(RuntimeError): + """A readable failure raised by the demo client.""" + + +async def send_message(client: Client, text: str, *, task_id: str = "") -> Task: + request = SendMessageRequest( + message=new_text_message(text, task_id=task_id or None, role=Role.ROLE_USER), + configuration=SendMessageConfiguration(return_immediately=True), + ) + async for response in client.send_message(request): + if response.HasField("task"): + return response.task + msg = "A2A response did not contain a task" + raise A2ADemoError(msg) + + +async def wait_for_state( # noqa: PLR0913 - polling controls are explicit CLI inputs + client: Client, + task_id: str, + expected: set[int], + *, + deadline: float, + poll_interval: float, + console: Console, +) -> Task: + previous_state = None + while time.monotonic() < deadline: + task = await client.get_task(GetTaskRequest(id=task_id)) + state = task.status.state + if state != previous_state: + console.print(f" [cyan]A2A state[/cyan] → [bold]{TaskState.Name(state)}[/bold]") + previous_state = state + if state in expected: + return task + if state in TERMINAL_STATES: + names = ", ".join(TaskState.Name(item) for item in sorted(expected)) + msg = f"Task became {TaskState.Name(state)} before reaching {names}" + raise A2ADemoError(msg) + await asyncio.sleep(poll_interval) + names = ", ".join(TaskState.Name(item) for item in sorted(expected)) + msg = f"Timed out waiting for {names}" + raise A2ADemoError(msg) + + +def print_summary(task: Task, console: Console) -> None: + table = Table(title="Durable A2A result", show_header=False) + table.add_column("Field", style="cyan") + table.add_column("Value") + table.add_row("Task", task.id) + table.add_row("State", TaskState.Name(task.status.state)) + table.add_row("History messages", str(len(task.history))) + table.add_row("Artifacts", ", ".join(artifact.name or "unnamed" for artifact in task.artifacts)) + console.print(table) + + result = next((artifact for artifact in task.artifacts if artifact.name == "durable-result"), None) + if result: + text = "\n".join(part.text for part in result.parts if part.WhichOneof("content") == "text") + if text: + console.print(Panel(text, title="Agent result", border_style="green")) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--url", default="http://localhost:1418/long_running_agent/", help="A2A agent base URL") + parser.add_argument("--document-id", default="hayhooks-guide") + parser.add_argument("--content", default="Hayhooks durable A2A work survives restarts.") + parser.add_argument("--timeout", type=float, default=120, help="Overall timeout in seconds") + parser.add_argument("--poll-interval", type=float, default=0.5) + return parser.parse_args() + + +async def run(args: argparse.Namespace, console: Console) -> None: + if args.timeout <= 0 or args.poll_interval <= 0: + msg = "timeout and poll interval must be positive" + raise A2ADemoError(msg) + + deadline = time.monotonic() + args.timeout + async with httpx.AsyncClient(timeout=15) as http: + client = await create_client( + args.url, + ClientConfig(streaming=False, polling=True, httpx_client=http), + ) + console.print(Panel.fit("Submit → input required → approve → complete", title="Durable A2A demo")) + task = await send_message( + client, + f"Prepare this document for indexing. document_id: {args.document_id}. content: {args.content}", + ) + console.print(f"[green]✓[/green] Submitted task [bold]{task.id}[/bold]") + await wait_for_state( + client, + task.id, + {TaskState.TASK_STATE_INPUT_REQUIRED}, + deadline=deadline, + poll_interval=args.poll_interval, + console=console, + ) + console.print("[green]✓[/green] Approval requested") + await send_message(client, "Approved; proceed.", task_id=task.id) + console.print("[green]✓[/green] Approval sent") + print_summary( + await wait_for_state( + client, + task.id, + {TaskState.TASK_STATE_COMPLETED}, + deadline=deadline, + poll_interval=args.poll_interval, + console=console, + ), + console, + ) + + +def main() -> int: + console = Console() + try: + asyncio.run(run(parse_args(), console)) + except (A2ADemoError, A2AClientError, httpx.HTTPError, ValueError) as error: + console.print(f"[bold red]Demo failed:[/bold red] {error}") + return 1 + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py b/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py new file mode 100644 index 00000000..cd952d6e --- /dev/null +++ b/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py @@ -0,0 +1,175 @@ +"""A durable A2A Agent that calls a real Haystack document-preparation Pipeline.""" + +import json +import os +import sqlite3 +import tempfile +import time +from pathlib import Path +from typing import Annotated + +from haystack import Document, Pipeline +from haystack.components.agents import Agent +from haystack.components.agents.state import State +from haystack.components.generators.chat import OpenAIChatGenerator +from haystack.components.preprocessors import DocumentCleaner, DocumentSplitter +from haystack.hooks.from_function import FunctionHook +from haystack.tools import tool + +from hayhooks import A2APipelineWrapper, current_durable_context, current_execution_id + +_MAX_DEMO_TOOL_DELAY_SECONDS = 300.0 + + +def _demo_delay_seconds(name: str, default: str = "0") -> float: + raw_delay = os.getenv(name, default) + try: + delay = float(raw_delay) + except ValueError as error: + msg = f"{name} must be a number" + raise ValueError(msg) from error + if not 0 <= delay <= _MAX_DEMO_TOOL_DELAY_SECONDS: + msg = f"{name} must be between 0 and 300" + raise ValueError(msg) + return delay + + +def require_approval(state: State) -> None: # noqa: ARG001 + """Suspend before the first model call so A2A exposes input-required.""" + context = current_durable_context() + if context is None or context.state.get("approval_requested"): + return + context.state["approval_requested"] = True + context.suspend_sync( + { + "kind": "approval", + "message": "Approve the indexing side effect", + "expected_input_schema": { + "type": "object", + "properties": {"message": {"type": "string"}}, + "required": ["message"], + }, + } + ) + + +async def require_approval_async(state: State) -> None: # noqa: ARG001 + context = current_durable_context() + if context is None or context.state.get("approval_requested"): + return + context.state["approval_requested"] = True + await context.suspend( + { + "kind": "approval", + "message": "Approve the indexing side effect", + "expected_input_schema": { + "type": "object", + "properties": {"message": {"type": "string"}}, + "required": ["message"], + }, + } + ) + + +@tool +def prepare_document_for_indexing( + document_id: Annotated[str, "A stable identifier for the source document"], + content: Annotated[str, "Raw document text to clean and split into chunks"], +) -> str: + """Clean, chunk, and idempotently record an indexing side effect.""" + context = current_durable_context() + execution_id = current_execution_id() + if context is None or execution_id is None: + msg = "This example tool must run inside a durable execution" + raise RuntimeError(msg) + effect_key = f"{execution_id}:index:{document_id}" + preparation_pipeline = Pipeline() + preparation_pipeline.add_component("clean", DocumentCleaner(remove_empty_lines=True)) + preparation_pipeline.add_component( + "split", + DocumentSplitter(split_by="word", split_length=80, split_overlap=10), + ) + preparation_pipeline.connect("clean.documents", "split.documents") + outputs = preparation_pipeline.run( + {"clean": {"documents": [Document(id=document_id, content=content, meta={"document_id": document_id})]}} + ) + chunks = outputs["split"]["documents"] + default_database = Path(tempfile.gettempdir()) / "hayhooks-durable-a2a.sqlite3" + database = os.getenv("HAYHOOKS_EXAMPLE_INDEX_DB", str(default_database)) + with sqlite3.connect(database) as connection: + connection.execute( + "CREATE TABLE IF NOT EXISTS indexing_effects " + "(idempotency_key TEXT PRIMARY KEY, document_id TEXT NOT NULL, chunk_count INTEGER NOT NULL)" + ) + cursor = connection.execute( + "INSERT OR IGNORE INTO indexing_effects (idempotency_key, document_id, chunk_count) VALUES (?, ?, ?)", + (effect_key, document_id, len(chunks)), + ) + applied = cursor.rowcount == 1 + + # Hold the tool open *after* its external effect. Killing the server in + # this window replays the tool from the previous Agent checkpoint, while + # the SQLite primary key proves that the effect is still applied once. + delay = _demo_delay_seconds("HAYHOOKS_EXAMPLE_TOOL_DELAY_SECONDS", "3") + context.report_progress_sync( + f"Indexing effect committed; holding the tool open for {delay:g} seconds", + kind="side_effect_committed", + metadata={ + "idempotency_key": effect_key, + "side_effect_applied": applied, + }, + ) + if delay: + time.sleep(delay) + + return json.dumps( + { + "document_id": document_id, + "chunk_count": len(chunks), + "idempotency_key": effect_key, + "side_effect_applied": applied, + "chunks": [{"chunk_id": str(chunk.id), "preview": (chunk.content or "")[:160]} for chunk in chunks], + } + ) + + +@tool +def publish_indexing_receipt( + document_id: Annotated[str, "The stable identifier returned by document preparation"], + chunk_count: Annotated[int, "The prepared chunk count returned by document preparation"], +) -> str: + """Perform the inexpensive follow-up step after document preparation succeeds.""" + context = current_durable_context() + if context is None: + msg = "This example tool must run inside a durable execution" + raise RuntimeError(msg) + + delay = _demo_delay_seconds("HAYHOOKS_EXAMPLE_RECEIPT_DELAY_SECONDS") + context.report_progress_sync( + f"Indexing is checkpointed; holding the lightweight receipt update for {delay:g} seconds", + kind="receipt_started", + ) + if delay: + time.sleep(delay) + + return json.dumps({"document_id": document_id, "chunk_count": chunk_count, "receipt": "published"}) + + +class PipelineWrapper(A2APipelineWrapper): + durable_revision = "a2a-long-running-agent" + """Let Hayhooks map this real tool-using Agent to durable A2A executions.""" + + def setup(self) -> None: + self.pipeline = Agent( + chat_generator=OpenAIChatGenerator(model="gpt-4o-mini"), + tools=[prepare_document_for_indexing, publish_indexing_receipt], + system_prompt=( + "You prepare documents for retrieval and publish their catalog status. When a user supplies a " + "document identifier and content, first call only prepare_document_for_indexing. After its result, " + "in a later tool turn call only publish_indexing_receipt with its document_id and chunk_count. " + "Do not prepare a document again once its successful result is present. Then report the number " + "of chunks and a concise readiness summary. Treat a follow-up approval message as authorization " + "to proceed." + ), + hooks={"before_llm": [FunctionHook(function=require_approval, async_function=require_approval_async)]}, + ) From 152b7abb5ff79c7fe9c47c945eab3789b1e12e77 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:10:09 +0200 Subject: [PATCH 09/28] docs(durable): explain engine and authoring model --- README.md | 20 +- docs/advanced/durable-engine-vs-temporal.md | 67 +++++++ docs/advanced/durable-engine.md | 202 ++++++++++++++++++++ docs/advanced/running-pipelines.md | 28 ++- docs/concepts/pipeline-wrapper.md | 154 +++++++++------ docs/features/a2a-support.md | 173 ++++++++++++++++- docs/guides/production-best-practices.md | 19 +- mkdocs.yml | 6 + 8 files changed, 579 insertions(+), 90 deletions(-) create mode 100644 docs/advanced/durable-engine-vs-temporal.md create mode 100644 docs/advanced/durable-engine.md diff --git a/README.md b/README.md index ae14176f..c5de60da 100644 --- a/README.md +++ b/README.md @@ -7,6 +7,7 @@ With Hayhooks, you can: - 📦 **Deploy your Haystack pipelines and agents as REST APIs** with maximum flexibility and minimal boilerplate code. - 🛠️ **Expose your Haystack pipelines and agents over the MCP protocol**, making them available as tools in AI dev environments like [Cursor](https://cursor.com) or [Claude Desktop](https://claude.ai/download). Under the hood, Hayhooks runs as an [MCP Server](https://modelcontextprotocol.io/docs/concepts/architecture), exposing each pipeline and agent as an [MCP Tool](https://modelcontextprotocol.io/docs/concepts/tools). - 🤝 **Expose your Haystack pipelines and agents over the [A2A protocol](https://a2a-protocol.org)** (`pip install "hayhooks[a2a]"`), so other agents can discover them through auto-generated agent cards and delegate tasks to them via `hayhooks a2a run`. +- ♻️ **Run Haystack 3 pipelines and agents as durable background work** (`pip install "hayhooks[durable]"`) with Redis-backed checkpoints, retries, cancellation, wait/resume, progress, and process recovery. - 💬 **Integrate your Haystack pipelines and agents with [Open WebUI](https://openwebui.com)** as OpenAI-compatible chat completion backends with streaming support. - 🖥️ **Embed a [Chainlit](https://chainlit.io/) chat UI** directly in Hayhooks with `pip install "hayhooks[chainlit]"` and `hayhooks run --with-chainlit` -- zero-configuration frontend with streaming, pipeline selection, and custom UI widgets. - 🕹️ **Control Hayhooks core API endpoints through chat** - deploy, undeploy, list, or run Haystack pipelines and agents by chatting with [Claude Desktop](https://claude.ai/download), [Cursor](https://cursor.com), or any other MCP client. @@ -53,6 +54,7 @@ from hayhooks import BasePipelineWrapper, async_streaming_generator def weather_function(location): return f"The weather in {location} is sunny." + weather_tool = Tool( name="weather_tool", description="Provides weather information for a given location.", @@ -64,6 +66,7 @@ weather_tool = Tool( function=weather_function, ) + class PipelineWrapper(BasePipelineWrapper): def setup(self) -> None: self.agent = Agent( @@ -73,7 +76,7 @@ class PipelineWrapper(BasePipelineWrapper): ) # This will create a POST /my_agent/run endpoint - # `question` will be the input argument and will be auto-validated by a Pydantic model + # `question` will be the input argument and will be auto-validated by a Pydantic model async def run_api_async(self, question: str) -> str: result = await self.agent.run_async(messages=[ChatMessage.from_user(question)]) return result["last_message"].text @@ -82,9 +85,7 @@ class PipelineWrapper(BasePipelineWrapper): async def run_chat_completion_async( self, model: str, messages: list[dict], body: dict ) -> AsyncGenerator[str, None]: - chat_messages = [ - ChatMessage.from_openai_dict_format(message) for message in messages - ] + chat_messages = [ChatMessage.from_openai_dict_format(message) for message in messages] return async_streaming_generator( pipeline=self.agent, @@ -154,11 +155,22 @@ Or chat with it in the [embedded Chainlit UI](docs/features/chainlit-integration - Built-in support for handling file uploads in pipelines - Perfect for RAG systems and document processing +### ♻️ Durable Execution + +- Run typed Pipeline and Agent work outside the request and recover it after a process restart +- Preserve checkpoints, bounded retries, progress, cancellation, wait/resume, idempotency, and owner isolation +- Project managed durable Agents over A2A and retain terminal results with Redis TTL + +Durable execution is intentionally an at-least-once engine for low-to-moderate workloads with one to three replicas. +It uses fixed polling and cooperative cancellation; it is not a DAG orchestrator, high-scale fair queue, live migration +system, or exactly-once boundary for external side effects. See the [supported scope and tradeoffs](docs/advanced/durable-engine.md#supported-scope-and-tradeoffs). + ## Next Steps - [Quick Start Guide](docs/getting-started/quick-start.md) - Get started with Hayhooks - [Installation](docs/getting-started/installation.md) - Install Hayhooks and dependencies - [Configuration](docs/getting-started/configuration.md) - Configure Hayhooks for your needs +- [Durable Engine](docs/advanced/durable-engine.md) - Understand restart-safe execution and its boundaries - [Tracing Dashboard Frontend](dashboard/README.md) - Local dashboard setup and frontend development commands - [Examples](docs/examples/overview.md) - Explore example implementations diff --git a/docs/advanced/durable-engine-vs-temporal.md b/docs/advanced/durable-engine-vs-temporal.md new file mode 100644 index 00000000..706a7266 --- /dev/null +++ b/docs/advanced/durable-engine-vs-temporal.md @@ -0,0 +1,67 @@ +# Hayhooks durable engine and Temporal + +Hayhooks provides focused durable execution for Haystack 3 Pipelines and Agents. Temporal is a general-purpose +durable workflow platform. + +## Core comparison + +| Requirement | Hayhooks durable engine | Temporal | +|---|---|---| +| Persistence and recovery | Redis records plus explicit Pipeline or Agent checkpoints | Event History plus deterministic [Workflow replay](https://docs.temporal.io/workflows) | +| Delivery safety | At-least-once with fenced leases and one active owner | Workflow logic is effectively once; [Activities](https://docs.temporal.io/activity-execution) may retry | +| Retries | Explicit, bounded retries from the latest checkpoint | Declarative, independently configurable [Activity retry policies](https://docs.temporal.io/encyclopedia/retry-policies) | +| Interaction | Typed inspect, wait/resume, progress, and result APIs | [Queries, Signals, and Updates](https://docs.temporal.io/encyclopedia/workflow-message-passing) plus durable Workflow state | +| Cancellation | Cooperative checks at safe boundaries | Cooperative cancellation with propagation policies | +| Orchestration | One Pipeline or Agent execution with delayed retries | Durable timers, Activities, [schedules](https://docs.temporal.io/schedule), and [Child Workflows](https://docs.temporal.io/child-workflows) | +| Versioning | Exact revision gate prevents incompatible recovery | Replay-safe patching and [Worker Versioning](https://docs.temporal.io/production-deployment/worker-deployments/worker-versioning) | +| Haystack example | Run a RAG Pipeline with `checkpoint_at=["generator"]`; after a crash, restore its `PipelineSnapshot` before generation | Invoke retrieval and generation as separate Activities wrapping Haystack components; completed Activity results are not repeated during Workflow replay | +| Best fit | Focused, moderate-scale durable Haystack workloads | Large-scale or cross-service orchestration around Haystack | + +## Hayhooks code map + +| Concern | Relevant implementation | +|---|---| +| Lifecycle and fencing | [`engine.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/engine.py) | +| Store contract and Redis persistence | [`store.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/store.py), [`redis.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/redis.py) | +| Pipeline and Agent checkpoints | [`adapters.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/adapters.py) | +| Retries and worker recovery | [`context.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/context.py), [`manager.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/manager.py) | +| Progress, wait/resume, and inspection | [`context.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/context.py), [`routes.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/server/durable/routes.py) | +| Cooperative cancellation | [`context.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/context.py), [`engine.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/engine.py) | +| Revision and deployment safety | [`runtime.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/durable/runtime.py), [`deploy_utils.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/server/utils/deploy_utils.py) | +| Durable A2A projection and recovery | [`durable_executor.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/server/a2a/durable_executor.py), [`redis_task_store.py`](https://github.com/deepset-ai/hayhooks/blob/main/src/hayhooks/server/a2a/redis_task_store.py) | + +Hayhooks makes an existing Haystack Pipeline or Agent durable with minimal restructuring. Temporal offers finer-grained +orchestration, but obtaining that granularity usually means deciding which Haystack operations should become separate +Activities; the whole Pipeline can still run as one Activity when independent step recovery is unnecessary. + +## What Temporal adds + +| Gain | Why it matters | +|---|---| +| Independent step policies | Retrieval, generation, payments, or notifications can have separate retries, timeouts, workers, and resource limits | +| Cross-service orchestration | A Workflow can coordinate Haystack with databases, external APIs, approval systems, and other services | +| Durable time and interaction | Native timers, schedules, callbacks, Signals, Queries, and Updates | +| Horizontal routing | Activities can run on different task queues, worker pools, languages, or infrastructure | +| Production operations | Execution history, search, UI, batch operations, metrics, and mature failure investigation | +| Long-running deployments | Worker versioning, pinning, gradual rollout, rollback, and Workflows that span releases | +| Reduced engine ownership | Temporal owns task delivery, persistence, recovery, and orchestration semantics instead of Hayhooks maintaining them in Redis | + +For restart-safe Haystack Pipelines and Agents with checkpoints, retries, cancellation, and resume, Temporal provides +little immediate functional gain and adds infrastructure plus integration work. It becomes valuable when an execution +grows into a long-lived, multi-step workflow spanning Haystack and other systems. + +Temporal does not remove the need for idempotent external side effects: an Activity may execute more than once when its +result is lost and the task is retried. + +## References + +- [Hayhooks durable engine](durable-engine.md) +- [Temporal Workflows](https://docs.temporal.io/workflows) +- [Temporal Activities](https://docs.temporal.io/activities) +- [Temporal Activity execution](https://docs.temporal.io/activity-execution) +- [Temporal retry policies](https://docs.temporal.io/encyclopedia/retry-policies) +- [Temporal Workflow message passing](https://docs.temporal.io/encyclopedia/workflow-message-passing) +- [Temporal Child Workflows](https://docs.temporal.io/child-workflows) +- [Temporal Schedules](https://docs.temporal.io/schedule) +- [Temporal Visibility](https://docs.temporal.io/visibility) +- [Temporal Worker Versioning](https://docs.temporal.io/production-deployment/worker-deployments/worker-versioning) diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md new file mode 100644 index 00000000..2fada2ce --- /dev/null +++ b/docs/advanced/durable-engine.md @@ -0,0 +1,202 @@ +# Durable engine + +Hayhooks durable execution provides detached, checkpointed work for **Haystack +3** Pipelines and Agents. + +It accepts validated work, runs it outside the request, resumes from Haystack +checkpoints, and recovers after a worker or process disappears. External +writes use application idempotency keys derived from the execution ID and +logical step, keeping replay safe when a process exits between an effect and +its next checkpoint. + +## Supported capabilities + +- Detached, typed Pipeline and Agent execution through durable REST endpoints. +- Pipeline snapshots and Agent state checkpoints with bounded retries, + progress, cancellation, and typed wait/resume. +- Redis-backed fenced claims and lease recovery across process restarts. +- Idempotent submission, optional owner-isolated REST access, and managed A2A + task projection. +- Native Redis TTL for terminal records, plus an equivalent volatile + in-memory backend for local development and tests. +- Explicit revision checks and safe deploy/undeploy handling. + +## Supported scope and tradeoffs + +- durable input before submission succeeds; +- at-least-once execution with one fenced worker owner at a time; +- checkpoints, progress, retries, cancellation, wait/resume, and terminal + results; and +- a public result that excludes private input, checkpoint, state, owner, and + fence details. + +The pure reducer in `hayhooks.durable.engine` is the only lifecycle +decision-maker. Storage atomically persists its plan and derives indexes from +the old and new control records. + +```text +queued ── claim ──> running ── complete/fail/cancel ──> terminal + │ │ + │ ├── checkpoint / heartbeat + │ ├── retry ──> queued (due later) + │ └── suspend ──> waiting ── resume ──> queued + └── cancel ──> terminal +``` + +## Embedding the runtime + +Applications can import `DurableRuntime`, `ExecutionStore`, +`ExecutionStoreProvider`, `InMemoryExecutionStoreProvider`, and +`RedisExecutionStoreProvider` directly from `hayhooks.durable`. A standalone +runtime starts only deployments attached to that runtime; it does not inspect +Hayhooks' process-global pipeline registry. + +This complete `app.py` embeds an in-memory durable worker in FastAPI. Its tool +simulates an eight-second upstream call so detached execution is easy to see: + +```python +import asyncio +from contextlib import asynccontextmanager +from typing import Annotated + +from fastapi import FastAPI, HTTPException, status +from haystack.components.agents import Agent +from haystack.components.generators.chat import OpenAIChatGenerator +from haystack.dataclasses import ChatMessage +from haystack.tools import tool +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper, DurableContext +from hayhooks.durable import DurableRuntime, ExecutionResult, InMemoryExecutionStoreProvider +from hayhooks.settings import AppSettings + + +class AgentRequest(BaseModel): + question: str + + +@tool +async def check_order(order_id: Annotated[str, "The customer's order ID"]) -> str: + """Return the current shipping status for an order.""" + # Intentional demo delay: replace it with a real upstream API call. + await asyncio.sleep(8) + # Read-only tools are replay-safe; make mutating tools idempotent. + return f"Order {order_id} shipped and arrives Friday." + + +class SupportAgentWrapper(BasePipelineWrapper): + # Bump this when checkpoint-relevant code or prompts change. + durable_revision = "support-agent-v1" + + def setup(self) -> None: + self.pipeline = Agent( + chat_generator=OpenAIChatGenerator(), + system_prompt="Help customers with their orders. Use the order tool when needed.", + tools=[check_order], + ) + + async def run_durable_async(self, context: DurableContext, request: AgentRequest) -> dict: + return await context.run_agent_async(messages=[ChatMessage.from_user(request.question)]) + + +durable_settings = AppSettings(durable_store="memory", durable_poll_interval=0.05) +provider = InMemoryExecutionStoreProvider(app_settings=durable_settings) +runtime = DurableRuntime(provider) + +wrapper = SupportAgentWrapper() +wrapper.setup() +deployment = runtime.deployment("support-agent", wrapper) + + +@asynccontextmanager +async def lifespan(_app: FastAPI): + try: + await runtime.start() + yield + finally: + await runtime.close() + + +app = FastAPI(lifespan=lifespan) + + +@app.post("/agent-runs", response_model=ExecutionResult, status_code=status.HTTP_202_ACCEPTED) +async def submit_agent_run(request: AgentRequest) -> ExecutionResult: + _, record = await deployment.submit(request.model_dump(mode="json")) + return ExecutionResult.model_validate(record.safe_view(links={"self": f"/agent-runs/{record.execution_id}"})) + + +@app.get("/agent-runs/{execution_id}", response_model=ExecutionResult) +async def get_agent_run(execution_id: str) -> ExecutionResult: + try: + record = await deployment.get(execution_id) + except KeyError as error: + raise HTTPException(status_code=404, detail="Agent run not found") from error + return ExecutionResult.model_validate(record.safe_view(links={"self": f"/agent-runs/{record.execution_id}"})) +``` + +Run it and submit work: + +```bash +pip install "hayhooks[durable]" +export OPENAI_API_KEY="your-api-key" +uvicorn app:app + +execution_id="$( + curl --fail --silent -X POST http://127.0.0.1:8000/agent-runs \ + -H 'content-type: application/json' \ + -d '{"question":"Where is order A-123?"}' | jq -r '.execution_id' +)" + +# The first poll should show `queued` or `running`; repeat until `completed`. +curl --fail --silent "http://127.0.0.1:8000/agent-runs/${execution_id}" | jq +``` + +The runtime owns provider shutdown. Built-in providers snapshot their settings, +and the runtime adopts that snapshot when a provider is supplied. Pass custom +settings once—either to a built-in provider as above, or to `DurableRuntime` +when it selects the default provider. Conflicting runtime and provider settings +are rejected before a deployment is created. + +## Redis layout + +Each deployment has an isolated namespace with controls and opaque input, +checkpoint, result, error, wait, and progress payload keys. It has exactly two +sorted-set indexes: + +| Key | Purpose | +|---|---| +| `runnable` | All queued work, scored by its retry deadline or immediate transition time. | +| `lease-expiry` | Running fences, scored by their Redis-server lease deadline. | + +The namespace also contains a `capacity` hash with only `nonterminal` and one +idempotency binding per execution. Terminal execution and idempotency keys use +native Redis TTL. The in-memory backend schedules equivalent cleanup. + +Workers poll `runnable` every configured poll interval, one second by default. +They use Redis `TIME` and read one due member without removing it. Multiple +replicas can observe it; the watched control and fence let exactly one +transition to `running`. Lease maintenance uses the same interval and recovers +at most 100 expired entries per pass. The default averages about 500 ms claim +latency and can add up to one second to claim or lease recovery. + +The controlled beta deployment profile uses one logical deployment with one to +three replicas and low-to-moderate load. Its two indexes, native TTL, and +single reducer keep the worker model observable during normal operation and +recovery. + +## Revisions and rollout + +Every durable wrapper, including a managed A2A Agent, declares a non-empty +`durable_revision`. Use an image digest or Git SHA in production and update it +with checkpoint-relevant code, prompts, configuration, or dependencies. Claims +and resumes verify that persisted work matches the active revision. + +## Operations + +`DurableExecutionManager.health_snapshot()` reports `nonterminal`, `runnable`, +and `lease_expiry`. Alert on sustained runnable growth, repeated lease recovery, +worker/store health failures, and runs that exceed their expected duration. + +See [Durable execution operations](durable-execution-operations.md) for +deployment, retention, and incident guidance. diff --git a/docs/advanced/running-pipelines.md b/docs/advanced/running-pipelines.md index 6566ca1a..86af7f5a 100644 --- a/docs/advanced/running-pipelines.md +++ b/docs/advanced/running-pipelines.md @@ -23,10 +23,7 @@ Execute deployed pipelines via CLI, HTTP API, or programmatically. ```python import requests - resp = requests.post( - "http://localhost:1416/my_pipeline/run", - json={"query": "What is Haystack?"} - ) + resp = requests.post("http://localhost:1416/my_pipeline/run", json={"query": "What is Haystack?"}) print(resp.json()) ``` @@ -36,14 +33,13 @@ Execute deployed pipelines via CLI, HTTP API, or programmatically. import httpx import asyncio + async def main(): async with httpx.AsyncClient() as client: - r = await client.post( - "http://localhost:1416/my_pipeline/run", - json={"query": "What is Haystack?"} - ) + r = await client.post("http://localhost:1416/my_pipeline/run", json={"query": "What is Haystack?"}) print(r.json()) + asyncio.run(main()) ``` @@ -95,10 +91,7 @@ See [File Upload Support](../features/file-upload-support.md) for implementation ```python import requests -resp = requests.post( - "http://localhost:1416/my_pipeline/run", - json={"query": "What is Haystack?"} -) +resp = requests.post("http://localhost:1416/my_pipeline/run", json={"query": "What is Haystack?"}) print(resp.json()) ``` @@ -110,14 +103,13 @@ print(resp.json()) import httpx import asyncio + async def main(): async with httpx.AsyncClient() as client: - r = await client.post( - "http://localhost:1416/my_pipeline/run", - json={"query": "What is Haystack?"} - ) + r = await client.post("http://localhost:1416/my_pipeline/run", json={"query": "What is Haystack?"}) print(r.json()) + asyncio.run(main()) ``` @@ -132,6 +124,7 @@ import requests from requests.exceptions import RequestException import time + def run_with_retry(pipeline_name, params, max_retries=3): url = f"http://localhost:1416/{pipeline_name}/run" @@ -143,7 +136,7 @@ def run_with_retry(pipeline_name, params, max_retries=3): except RequestException as e: if attempt == max_retries - 1: raise - time.sleep(2 ** attempt) # Exponential backoff + time.sleep(2**attempt) # Exponential backoff ``` ## Logging @@ -153,6 +146,7 @@ Add logging to your pipeline wrappers: ```python from hayhooks import log + class PipelineWrapper(BasePipelineWrapper): def run_api(self, query: str) -> str: log.info("Processing query: {}", query) diff --git a/docs/concepts/pipeline-wrapper.md b/docs/concepts/pipeline-wrapper.md index c21cd4ce..e6cb519f 100644 --- a/docs/concepts/pipeline-wrapper.md +++ b/docs/concepts/pipeline-wrapper.md @@ -20,6 +20,7 @@ from haystack import Pipeline from hayhooks import BasePipelineWrapper, get_last_user_message, async_streaming_generator, streaming_generator + class PipelineWrapper(BasePipelineWrapper): def setup(self) -> None: pipeline_yaml = (Path(__file__).parent / "pipeline.yml").read_text() @@ -67,7 +68,7 @@ def setup(self) -> None: {% endfor %} Answer the given question: {{query}} {% endmessage %}""", - required_variables="*" + required_variables="*", ) llm = OpenAIChatGenerator(model="gpt-4o-mini") @@ -165,6 +166,7 @@ Hayhooks can stream results from `run_api()` or `run_api_async()` when you retur from collections.abc import Generator from hayhooks import streaming_generator + def run_api(self, query: str) -> Generator: return streaming_generator( pipeline=self.pipeline, @@ -178,6 +180,7 @@ For async pipelines: from collections.abc import AsyncGenerator from hayhooks import async_streaming_generator + async def run_api_async(self, query: str) -> AsyncGenerator: return async_streaming_generator( pipeline=self.pipeline, @@ -198,6 +201,7 @@ If you need SSE (for browsers, [EventSource](https://developer.mozilla.org/en-US ```python from hayhooks import SSEStream, streaming_generator + def run_api(self, query: str): return SSEStream( streaming_generator( @@ -212,6 +216,7 @@ For async pipelines: ```python from hayhooks import SSEStream, async_streaming_generator + async def run_api_async(self, query: str): return SSEStream( async_streaming_generator( @@ -229,6 +234,7 @@ Hayhooks can return binary files (images, PDFs, audio, etc.) directly from `run_ import tempfile from fastapi.responses import FileResponse + def run_api(self, prompt: str) -> FileResponse: image = self.generate_image(prompt) @@ -250,6 +256,56 @@ For a full working example, see the [Image Generation example](https://github.co ## Optional Methods +### Durable execution + +Implement `run_durable()` or `run_durable_async()` on an ordinary `BasePipelineWrapper` to submit restart-safe work. +The method receives a `DurableContext` and one Pydantic request model; it returns a Pydantic result model. Hayhooks +owns records, worker lifecycle, Redis, cancellation, and resume. + +Ordinary exceptions from a durable wrapper are terminal failures. For a +transient dependency failure (for example an LLM timeout or rate limit), call +`await context.retry(...)` so Hayhooks persists the checkpoint and schedules a +bounded retry; do not simply re-raise the transient error. + +```python +from haystack import Pipeline +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper, DurableContext + + +class JobRequest(BaseModel): + source: str + + +class JobResult(BaseModel): + processed: int + + +class PipelineWrapper(BasePipelineWrapper): + durable_revision = "my-image-digest-or-git-sha" + + def setup(self) -> None: + self.pipeline = Pipeline() + # Add and connect components. + + async def run_durable_async(self, context: DurableContext, request: JobRequest) -> JobResult: + outputs = await context.run_pipeline_async({"source": {"value": request.source}}, checkpoint_at=["process"]) + return JobResult(processed=outputs["process"]["count"]) +``` + +Hayhooks exposes `POST /{pipeline}/run-durable`, `GET /{pipeline}/executions/{execution_id}`, and cancel/resume +endpoints. The submit response and status endpoint expose only a safe result view; validated inputs and checkpoint +state remain server-side. The built-in Redis store is the default. Set `HAYHOOKS_DURABLE_STORE=memory` only for +volatile local development. `run_pipeline_async()` uses a worker thread around +Haystack's synchronous snapshot API; it does not call `Pipeline.run_async()`. +See the [durable engine](../advanced/durable-engine.md) for checkpoint, +recovery, and side-effect boundaries. +Every durable wrapper must set a non-empty `durable_revision` class attribute +from the immutable build identifier; Hayhooks does not fingerprint source or +configuration automatically. See also the +[complete durable example](https://github.com/deepset-ai/hayhooks/tree/main/examples/durable_execution). + ### run_api_async() The asynchronous version of `run_api()` for better performance under high load. @@ -348,7 +404,7 @@ async def run_chat_completion_async(self, model: str, messages: list[dict], body return async_streaming_generator( pipeline=self.pipeline, pipeline_run_args={"prompt": {"query": question}}, - allow_sync_streaming_callbacks=True # ✅ Auto-detect and enable hybrid mode + allow_sync_streaming_callbacks=True, # ✅ Auto-detect and enable hybrid mode ) ``` @@ -371,12 +427,12 @@ When you set `allow_sync_streaming_callbacks=True`, the system enables **intelli ```python # Option 1: Strict mode (Default - Recommended) -allow_sync_streaming_callbacks=False +allow_sync_streaming_callbacks = False # → Raises error if sync-only components found # → Best for: New code, ensuring proper async components, best performance # Option 2: Auto-detection (Compatibility mode) -allow_sync_streaming_callbacks=True +allow_sync_streaming_callbacks = True # → Automatically detects and enables hybrid mode only when needed # → Best for: Legacy pipelines, components without async support, gradual migration ``` @@ -412,16 +468,14 @@ class SyncOnlyWrapper(BasePipelineWrapper): self.pipeline = Pipeline() self.pipeline.add_component("llm", SyncOnlyGenerator()) - async def run_chat_completion_async( - self, model: str, messages: list[dict], body: dict - ) -> AsyncGenerator: + async def run_chat_completion_async(self, model: str, messages: list[dict], body: dict) -> AsyncGenerator: question = get_last_user_message(messages) # Enable hybrid mode so the sync-only component can stream in an async pipeline return async_streaming_generator( pipeline=self.pipeline, pipeline_run_args={"llm": {"prompt": question}}, - allow_sync_streaming_callbacks=True # ✅ Handles sync component + allow_sync_streaming_callbacks=True, # ✅ Handles sync component ) ``` @@ -478,11 +532,8 @@ class MultiLLMWrapper(BasePipelineWrapper): self.pipeline.add_component( "prompt_1", ChatPromptBuilder( - template=[ - ChatMessage.from_system("You are a helpful assistant."), - ChatMessage.from_user("{{query}}") - ] - ) + template=[ChatMessage.from_system("You are a helpful assistant."), ChatMessage.from_user("{{query}}")] + ), ) self.pipeline.add_component("llm_1", OpenAIChatGenerator(model="gpt-4o-mini")) @@ -492,11 +543,9 @@ class MultiLLMWrapper(BasePipelineWrapper): ChatPromptBuilder( template=[ ChatMessage.from_system("You are a helpful assistant that refines responses."), - ChatMessage.from_user( - "Previous response: {{previous_response[0].text}}\n\nRefine this." - ) + ChatMessage.from_user("Previous response: {{previous_response[0].text}}\n\nRefine this."), ] - ) + ), ) self.pipeline.add_component("llm_2", OpenAIChatGenerator(model="gpt-4o-mini")) @@ -509,10 +558,7 @@ class MultiLLMWrapper(BasePipelineWrapper): question = get_last_user_message(messages) # By default, only llm_2 (the last streaming component) will stream - return streaming_generator( - pipeline=self.pipeline, - pipeline_run_args={"prompt_1": {"query": question}} - ) + return streaming_generator(pipeline=self.pipeline, pipeline_run_args={"prompt_1": {"query": question}}) ``` **What happens:** Only `llm_2` (the last streaming-capable component) streams its responses token by token. The first LLM (`llm_1`) executes normally without streaming, and only the final refined output streams to the user. @@ -529,7 +575,7 @@ def run_chat_completion(self, model: str, messages: list[dict], body: dict) -> G return streaming_generator( pipeline=self.pipeline, pipeline_run_args={"prompt_1": {"query": question}}, - streaming_components=["llm_1", "llm_2"] # Stream both components + streaming_components=["llm_1", "llm_2"], # Stream both components ) ``` @@ -539,16 +585,16 @@ You can also selectively enable streaming for specific components: ```python # Stream only the first LLM -streaming_components=["llm_1"] +streaming_components = ["llm_1"] # Stream only the second LLM (same as default) -streaming_components=["llm_2"] +streaming_components = ["llm_2"] # Stream ALL capable components (shorthand) -streaming_components="all" +streaming_components = "all" -# Stream ALL capable components (specific list) -streaming_components=["llm_1", "llm_2"] +# Stream ALL capable components (specific list) +streaming_components = ["llm_1", "llm_2"] ``` ### Using the "all" Keyword @@ -559,7 +605,7 @@ The `"all"` keyword is a convenient shorthand to enable streaming for all capabl return streaming_generator( pipeline=self.pipeline, pipeline_run_args={...}, - streaming_components="all" # Enable all streaming components + streaming_components="all", # Enable all streaming components ) ``` @@ -667,27 +713,24 @@ See the [Multi-LLM Streaming Example](https://github.com/deepset-ai/hayhooks/tre For streaming responses, pass `include_outputs_from` to `streaming_generator()` or `async_streaming_generator()`, and use the `on_pipeline_end` callback to access intermediate outputs. For example: ```python - def run_chat_completion(self, model: str, messages: List[dict], body: dict) -> Generator: - question = get_last_user_message(messages) +def run_chat_completion(self, model: str, messages: List[dict], body: dict) -> Generator: + question = get_last_user_message(messages) - # Store retrieved documents for citations - self.retrieved_docs = [] + # Store retrieved documents for citations + self.retrieved_docs = [] - def on_pipeline_end(result: dict[str, Any]) -> None: - # Access intermediate outputs here - if "retriever" in result: - self.retrieved_docs = result["retriever"]["documents"] - # Use for citations, logging, analytics, etc. + def on_pipeline_end(result: dict[str, Any]) -> None: + # Access intermediate outputs here + if "retriever" in result: + self.retrieved_docs = result["retriever"]["documents"] + # Use for citations, logging, analytics, etc. - return streaming_generator( - pipeline=self.pipeline, - pipeline_run_args={ - "retriever": {"query": question}, - "prompt_builder": {"query": question} - }, - include_outputs_from={"retriever"}, # Make retriever outputs available - on_pipeline_end=on_pipeline_end - ) + return streaming_generator( + pipeline=self.pipeline, + pipeline_run_args={"retriever": {"query": question}, "prompt_builder": {"query": question}}, + include_outputs_from={"retriever"}, # Make retriever outputs available + on_pipeline_end=on_pipeline_end, + ) ``` **What happens:** The `on_pipeline_end` callback receives both `llm` and `retriever` outputs in the `result` dict, allowing you to access retrieved documents alongside the generated response. @@ -704,12 +747,9 @@ async def run_chat_completion_async(self, model: str, messages: List[dict], body return async_streaming_generator( pipeline=self.async_pipeline, - pipeline_run_args={ - "retriever": {"query": question}, - "prompt_builder": {"query": question} - }, + pipeline_run_args={"retriever": {"query": question}, "prompt_builder": {"query": question}}, include_outputs_from={"retriever"}, - on_pipeline_end=on_pipeline_end + on_pipeline_end=on_pipeline_end, ) ``` @@ -719,10 +759,7 @@ For non-streaming `run_api` or `run_api_async` endpoints, pass `include_outputs_ ```python def run_api(self, query: str) -> dict: - result = self.pipeline.run( - data={"retriever": {"query": query}}, - include_outputs_from={"retriever"} - ) + result = self.pipeline.run(data={"retriever": {"query": query}}, include_outputs_from={"retriever"}) # Build custom response with both answer and sources return {"answer": result["llm"]["replies"][0], "sources": result["retriever"]["documents"]} ``` @@ -732,8 +769,7 @@ Same pattern for async: ```python async def run_api_async(self, query: str) -> dict: result = await self.async_pipeline.run_async( - data={"retriever": {"query": query}}, - include_outputs_from={"retriever"} + data={"retriever": {"query": query}}, include_outputs_from={"retriever"} ) return {"answer": result["llm"]["replies"][0], "sources": result["retriever"]["documents"]} ``` @@ -768,6 +804,7 @@ def on_reasoning( """ return text + def run_chat_completion(self, model: str, messages: list[dict], body: dict) -> Generator: return streaming_generator( pipeline=self.pipeline, @@ -793,6 +830,7 @@ Hayhooks can handle file uploads by adding a `files` parameter: ```python from fastapi import UploadFile + def run_api(self, files: list[UploadFile] | None = None, query: str = "") -> str: if files: # Process uploaded files @@ -841,6 +879,7 @@ Your pipeline wrapper may require additional dependencies: # pipeline_wrapper.py import trafilatura # Additional dependency + def run_api(self, urls: list[str], question: str) -> str: # Use additional library content = trafilatura.fetch(urls[0]) @@ -867,6 +906,7 @@ Implement proper error handling in production: from hayhooks import log from fastapi import HTTPException + class PipelineWrapper(BasePipelineWrapper): def setup(self) -> None: try: diff --git a/docs/features/a2a-support.md b/docs/features/a2a-support.md index e0da2f17..3dd08106 100644 --- a/docs/features/a2a-support.md +++ b/docs/features/a2a-support.md @@ -8,14 +8,16 @@ A2A complements [MCP support](mcp-support.md): MCP exposes pipelines as **tools* The Hayhooks A2A Server: -- Exposes every deployed pipeline that implements `run_chat_completion` or `run_chat_completion_async` as an A2A agent +- Exposes deployed chat and durable Agent wrappers as A2A agents - Serves a per-agent [Agent Card](#agent-cards) for discovery, auto-generated from the pipeline and customizable from the wrapper - Implements the JSON-RPC protocol binding of the [A2A specification](https://a2a-protocol.org/latest/specification/) (v1.0), including SSE streaming - Streams pipeline output incrementally as task artifact updates +- Supports detached long-running task execution with polling, subscription, and cooperative async cancellation ## Requirements -- Install with `pip install hayhooks[a2a]` (uses the official [a2a-sdk](https://github.com/a2aproject/a2a-python)) +- Install with `pip install hayhooks[a2a]` (uses the official [a2a-sdk](https://github.com/a2aproject/a2a-python) + and [redis-py](https://redis.io/docs/latest/develop/clients/redis-py/) clients) ## Getting Started @@ -44,11 +46,22 @@ HAYHOOKS_A2A_EXTERNAL_URL= # Base URL advertised in agent cards # (set when behind a reverse proxy) HAYHOOKS_A2A_V0_3_COMPAT=true # Also accept A2A spec 0.3 requests # (used by older clients and tools) +HAYHOOKS_A2A_TASK_STORE=auto # auto, memory, or redis +HAYHOOKS_A2A_REDIS_URL=redis://localhost:6379/0 +HAYHOOKS_A2A_REDIS_KEY_PREFIX=hayhooks:a2a +HAYHOOKS_DURABLE_STORE=redis # Redis by default; memory is volatile +HAYHOOKS_DURABLE_REDIS_URL=redis://localhost:6379/0 +HAYHOOKS_DURABLE_REDIS_KEY_PREFIX=hayhooks:durable +HAYHOOKS_DURABLE_EXECUTION_CONCURRENCY=1 + # Workers per deployed durable Agent ``` ## Which pipelines are exposed -A deployed pipeline is exposed as an A2A agent when it implements `run_chat_completion` or `run_chat_completion_async` — the same methods used by the [OpenAI-compatible chat endpoints](openai-compatibility.md). No extra method is needed. +A deployed pipeline is exposed as an A2A agent when it uses either authoring mode: + +- **Chat compatibility**: implement `run_chat_completion` or `run_chat_completion_async`, the same methods used by the [OpenAI-compatible chat endpoints](openai-compatibility.md). +- **Durable Agent**: inherit from `hayhooks.a2a.A2APipelineWrapper` and assign a Haystack 3 `Agent` to `self.pipeline`. Hayhooks supplies the detached executor, checkpoints, progress projection, and durable store. To exclude a chat-capable pipeline from A2A, set `skip_a2a` on the wrapper: @@ -65,7 +78,7 @@ Each exposed pipeline is mounted under its own path prefix: |----------|-------------| | `GET /{pipeline_name}/.well-known/agent-card.json` | The pipeline's agent card | | `POST /{pipeline_name}/` | JSON-RPC binding (`SendMessage`, `SendStreamingMessage`, `GetTask`, ...) | -| `GET /status` | Server status and the list of exposed agents | +| `GET /status` | Operational readiness and the list of exposed agents; not an A2A protocol method | For example, with a deployed `weather_agent` pipeline: @@ -73,6 +86,13 @@ For example, with a deployed `weather_agent` pipeline: curl http://localhost:1418/weather_agent/.well-known/agent-card.json ``` +The operational status endpoint returns `200` only when the configured task +store, executor lifecycle, maintenance loop, and (for managed +durable Agents) durable execution runtime are healthy. It returns `503` with +`status: unavailable` otherwise. The A2A specification does not define a +health method; Agent Cards and their advertised interfaces remain the +standards-compliant discovery and protocol surface. + ## Agent Cards Agent cards are generated automatically: the card name is the pipeline name, the description comes from the pipeline's registry metadata, and a single default skill is created. Override any of it with the `a2a_card` class attribute: @@ -136,14 +156,146 @@ curl -s http://localhost:1418/weather_agent/ \ "parts": [{"text": "Weather in Berlin?"}]}}}' ``` -## Task lifecycle and streaming +## Chat-compatibility task lifecycle and streaming -Each request is handled as an A2A task: +In chat-compatibility mode, each request is handled as an A2A task: 1. A `Task` is created from the incoming message. 2. The task transitions to `working` and the pipeline's chat completion method runs. 3. Pipeline output is emitted as a single `response` artifact. Streaming results (generators returned by `streaming_generator` / `async_streaming_generator`) are emitted incrementally as artifact chunk updates, so `SendStreamingMessage` clients receive text as it is produced. -4. The task ends in `completed` (or `failed`, with the error in the status message — enable `HAYHOOKS_SHOW_TRACEBACKS` to include tracebacks). +4. The task ends in `completed`, `failed`, or `canceled`. Enable `HAYHOOKS_SHOW_TRACEBACKS` to include tracebacks in failure messages. + +Native executors own their protocol lifecycle. Durable Agents instead project +their authoritative durable execution into an A2A task as described under +[Durable Haystack Agents](#durable-haystack-agents). + +By default, non-streaming `SendMessage` remains blocking for backward compatibility: the response is returned after the task reaches a terminal or interrupted state. + +For detached execution, set `configuration.returnImmediately`: + +```bash +curl -s http://localhost:1418/weather_agent/ \ + -H "Content-Type: application/json" -H "A2A-Version: 1.0" \ + -d '{"jsonrpc": "2.0", "id": "1", "method": "SendMessage", + "params": {"configuration": {"returnImmediately": true}, + "message": {"messageId": "m1", "role": "ROLE_USER", + "parts": [{"text": "Start the long task"}]}}}' +``` + +The response contains a non-terminal task. Poll it with `GetTask`: + +```bash +curl -s http://localhost:1418/weather_agent/ \ + -H "Content-Type: application/json" -H "A2A-Version: 1.0" \ + -d '{"jsonrpc": "2.0", "id": "2", "method": "GetTask", + "params": {"id": ""}}' +``` + +Or subscribe to an active task with `SubscribeToTask` to receive the latest task snapshot followed by task updates over SSE: + +```bash +curl -N http://localhost:1418/weather_agent/ \ + -H "Content-Type: application/json" -H "A2A-Version: 1.0" \ + -d '{"jsonrpc": "2.0", "id": "3", "method": "SubscribeToTask", + "params": {"id": ""}}' +``` + +When `HAYHOOKS_A2A_V0_3_COMPAT=true`, A2A 0.3 clients can request the same detached behavior with `configuration.blocking=false`. + +## Task storage + +With `HAYHOOKS_A2A_TASK_STORE=auto`, Hayhooks uses Redis for Redis-backed durable Agents and otherwise gives each exposed agent its own A2A SDK `InMemoryTaskStore`. An explicit `memory` choice is never overridden. + +Task storage is server infrastructure rather than pipeline configuration. Hayhooks includes independent in-memory and Redis-backed providers. Select the built-in Redis provider with `HAYHOOKS_A2A_TASK_STORE=redis` or `hayhooks a2a run --task-store redis`; configure its URL and key prefix with `HAYHOOKS_A2A_REDIS_URL` and `HAYHOOKS_A2A_REDIS_KEY_PREFIX`. + +The SDK in-memory store is for one Hayhooks process only. It loses tasks on +restart, and a request routed to another replica cannot see tasks created by +the first. Multi-replica deployments must set the task store to `redis` and +give every replica the same Redis URL and key prefix. `auto` does this only for +Redis-backed durable Agents, so scaled chat-compatible wrappers must select `redis` explicitly. + +The A2A extra includes the official Redis client. Redis task records are protobuf payloads scoped by agent and resolved owner. The configured owner resolver is the sole owner decision point; the default distinguishes unauthenticated requests from an authenticated user named `anonymous`. Durable execution IDs are internal fixed-size hashes of that owner and the opaque A2A task ID, so clients must keep using the original A2A task ID. + +Redis stores each task and its version in one hash, with derived owner, active, and expiry indexes. Writes use bounded `WATCH`/`MULTI` transactions and Redis server time; they do not require Lua or `EVAL`. All keys for one agent share a Redis Cluster hash slot. Restart recovery relies on the durable engine's idempotent submission and persists each complete task projection with one compare-and-set write, so a stale replica cannot overwrite a newer task state. Persistent task records do not replay historical live event queues. + +A durable Agent accepts another user message only while its execution is +waiting for input. A follow-up sent while it is running is rejected rather than +being treated as an idempotent redelivery; clients should wait for +`INPUT_REQUIRED` before continuing a task. + +Terminal tasks use `HAYHOOKS_A2A_TERMINAL_TASK_TTL_SECONDS`. Runtime maintenance performs cleanup even when no later A2A request arrives and removes the protobuf payload and task indexes. Execution-record retention remains independent. + +Applications constructing the server directly can use a configured provider: + +```python +from hayhooks.a2a import RedisTaskStoreProvider +from hayhooks.server.a2a.app import create_a2a_app +from hayhooks.server.a2a.runtime import A2ARuntime + +runtime = A2ARuntime( + task_store_provider=RedisTaskStoreProvider( + redis_url="redis://localhost:6379/0", + key_prefix="my-app:a2a", + ) +) +app = create_a2a_app(runtime=runtime) +``` + +Persisting task records alone does not make execution recoverable. For durable Agents, use Redis durable execution +(the default) and configure the A2A task store for the task history retention your clients need. Chat-compatible +execution remains process-local. + +### Durable Haystack Agents + +For the managed mode, use an `A2APipelineWrapper` with a Haystack 3 `Agent`. Hayhooks creates the +execution record using the A2A task ID, captures public Agent state at model/tool boundaries, and projects safe +progress, waiting, completion, failure, and cancellation states back to the A2A task. + +The execution record is authoritative; the A2A Task is its persisted client-facing projection. This is why the two +records have separate retention settings, and why a persistent A2A task store alone cannot recover interrupted work. + +```mermaid +flowchart LR + Client["A2A client"] --> Server["A2A server\nmanaged durable executor"] + Server --> Task["A2A Task store\nclient-facing task"] + Server --> Execution["Durable execution record\nsource of truth"] + Execution --> Manager["DurableExecutionManager\nclaim + fenced lease"] + Manager --> Agent["Haystack Agent\ntools + checkpoints"] + Agent --> Execution + Execution --> Projection["A2A task projection\nprogress / waiting / terminal state"] + Projection --> Task + Task --> Client +``` + +The projection updates the task; it does not run the Agent. After a restart, +the durable worker recovers execution. Startup then refreshes any persisted +active task snapshot once; later `GetTask` and list requests project the latest +durable state directly. + +```python +from haystack.components.agents import Agent +from haystack.components.generators.chat import OpenAIChatGenerator + +from hayhooks import A2APipelineWrapper + + +class PipelineWrapper(A2APipelineWrapper): + def setup(self) -> None: + self.pipeline = Agent(chat_generator=OpenAIChatGenerator(), tools=[]) +``` + +The wrapper does not create an executor, worker, record, queue, or Redis client. Durable execution uses Redis by +default; set `HAYHOOKS_DURABLE_STORE=memory` only for non-recoverable local development. Concurrent +execution is controlled by `HAYHOOKS_DURABLE_EXECUTION_CONCURRENCY` and requires the Agent, tools, and their shared +dependencies to be concurrency-safe. + +Snapshots, validated messages, and internal tool state remain server-side. A restarted Hayhooks process reclaims +incomplete Redis work from its last safe checkpoint. Tool effects before a checkpoint may be replayed, so tools should +be idempotent. Before exposing a durable Agent, apply the +[controlled beta deployment profile](../advanced/durable-execution-operations.md#controlled-beta-deployment-profile); +in particular, authenticate the A2A endpoint and enforce request and admission +limits at the gateway. See the +[durable A2A example](https://github.com/deepset-ai/hayhooks/tree/main/examples/a2a_long_running). ## Inspecting agents with a2a-inspector @@ -155,8 +307,9 @@ See [examples/a2a_multi_agent](https://github.com/deepset-ai/hayhooks/tree/main/ ## Current limitations -- **Request-bound task execution**: Hayhooks currently treats A2A as a chat-shaped bridge. Each task runs inside the request handler by calling `run_chat_completion` / `run_chat_completion_async`, so non-streaming `SendMessage` returns after the task has completed or failed. This means detached task execution via [`returnImmediately`](https://a2a-protocol.org/latest/specification/#322-sendmessageconfiguration), [`input-required`](https://a2a-protocol.org/latest/specification/#63-multi-turn-interaction) pauses, and [push notification delivery](https://a2a-protocol.org/latest/specification/#353-push-notification-delivery) are not supported yet. +- **Process-owned execution**: chat-compatible execution pauses while the A2A server is offline; durable Agents persist checkpoints and reclaim incomplete work after Hayhooks starts again. +- **Automatic task-store selection**: `auto` selects Redis only for Redis-backed durable Agents. Explicit `memory` stays process-local. A persistent task store preserves the protocol projection, but does not recover interrupted execution by itself. +- **Push notifications**: push notification delivery is not enabled yet, and agent cards do not advertise it. - **Static agents list**: A2A routes are built from the registry at startup. Pipelines deployed or undeployed at runtime require restarting `hayhooks a2a run`. -- **In-memory task store**: task state is kept in memory and lost on restart. - **Path-prefixed agent cards**: one server hosts many agents, so cards live under `/{pipeline_name}/.well-known/agent-card.json` instead of the domain root. If a consumer requires strict root-level discovery, run one A2A server instance per agent (separate `--pipelines-dir` and `--port`). -- **Cancellation is best-effort**: cancelling a task marks it canceled but does not interrupt a running pipeline. +- **Cancellation is cooperative**: async chat wrappers and durable Agents can observe cancellation. A durable A2A task reports cancellation requested first and becomes A2A canceled only after the execution record is terminal canceled. Synchronous work cannot be forcibly interrupted and retains its fenced claim until it returns. diff --git a/docs/guides/production-best-practices.md b/docs/guides/production-best-practices.md index 289fd0f7..4d521d76 100644 --- a/docs/guides/production-best-practices.md +++ b/docs/guides/production-best-practices.md @@ -100,6 +100,7 @@ Pipelines that spend most of their time waiting on external services -- LLM API ```python from hayhooks import BasePipelineWrapper, Pipeline + class PipelineWrapper(BasePipelineWrapper): def setup(self) -> None: self.pipeline = Pipeline() @@ -157,10 +158,10 @@ healthcheck: start_period: 40s ``` -For Kubernetes, use a liveness probe on the same endpoint: +For Kubernetes, use a readiness probe on the same endpoint: ```yaml -livenessProbe: +readinessProbe: httpGet: path: /status port: 1416 @@ -168,6 +169,20 @@ livenessProbe: periodSeconds: 30 ``` +## Treat Durable Execution as a Controlled Beta + +Durable Pipelines and managed durable A2A Agents have stricter requirements +than ordinary request/response pipelines. Before using them with production +traffic, follow the authoritative +[controlled beta deployment profile](../advanced/durable-execution-operations.md#controlled-beta-deployment-profile). +It covers the supported Redis topology, immutable revisions, drained upgrades, +concurrency, ingress limits, operation timeouts, idempotent effects, +observability, and incident recovery. + +Do not use rolling mixed-version upgrades for the durable engine, and do not +assume Pipeline or Agent checkpoints remain compatible across Haystack, +Hayhooks, or application upgrades. + ## Docker and Container Tips Follow these practices when running Hayhooks in containers: diff --git a/mkdocs.yml b/mkdocs.yml index 7e217187..7c450d51 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -30,6 +30,9 @@ nav: - Advanced Usage: - Running Pipelines: advanced/running-pipelines.md - Advanced Configuration: advanced/advanced-configuration.md + - Durable Engine: advanced/durable-engine.md + - Durable Engine vs Temporal: advanced/durable-engine-vs-temporal.md + - Durable Execution Operations: advanced/durable-execution-operations.md - Code Sharing: advanced/code-sharing.md - Guides: - Development Best Practices: guides/development-best-practices.md @@ -118,6 +121,9 @@ plugins: Advanced Usage: - advanced/running-pipelines.md: Execute Haystack Pipelines and manage runs programmatically. - advanced/advanced-configuration.md: Fine-tune advanced configuration settings. + - advanced/durable-engine.md: Current contract, state model, and architectural boundaries for durable Pipeline and Agent execution. + - advanced/durable-engine-vs-temporal.md: Compare Hayhooks durable execution with Temporal and understand when each fits. + - advanced/durable-execution-operations.md: Operate durable execution with explicit Redis, recovery, security, and scaling contracts. - advanced/code-sharing.md: Share Python code securely across deployments. Guides: - guides/development-best-practices.md: Development workflow tips for Hayhooks pipelines. From c93c86cf2605dff9eb4159b10bb83a068948b26a Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:10:19 +0200 Subject: [PATCH 10/28] docs(durable): add operations and configuration reference --- docs/advanced/durable-execution-operations.md | 53 +++++++ docs/examples/overview.md | 2 + docs/features/cli-commands.md | 8 + docs/getting-started/installation.md | 12 +- docs/reference/api-reference.md | 39 ++--- docs/reference/environment-variables.md | 140 ++++++++++++++++++ examples/README.md | 25 +++- 7 files changed, 257 insertions(+), 22 deletions(-) create mode 100644 docs/advanced/durable-execution-operations.md diff --git a/docs/advanced/durable-execution-operations.md b/docs/advanced/durable-execution-operations.md new file mode 100644 index 00000000..1e4cf652 --- /dev/null +++ b/docs/advanced/durable-execution-operations.md @@ -0,0 +1,53 @@ +# Durable execution operations + +Hayhooks durable execution provides fenced, at-least-once recovery. Use an +idempotency key derived from the execution ID and logical step for every +external write so recovered work remains safe to replay. + +## Controlled beta deployment profile + +- Use authenticated Redis 6.2+ with TLS, persistence, backups, and + `maxmemory-policy noeviction`. +- Run one logical Hayhooks deployment, normally one replica and at most two or + three, against one isolated namespace. +- Set a non-empty `durable_revision` on every durable Pipeline wrapper and + managed A2A Agent. An image digest or Git SHA is the recommended value. +- Start with `HAYHOOKS_DURABLE_EXECUTION_CONCURRENCY=1`; increase it only after + every Pipeline/Agent component and tool is proven concurrency-safe. +- Put REST and A2A behind authentication, request-size, rate, and tenant + controls. `HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS` is an optional + deployment-wide secondary admission cap. + +## Execution and recovery + +The namespace holds a control record, opaque payloads, one `runnable` +ZSET, one `lease-expiry` ZSET, a `nonterminal` capacity field, and idempotency +bindings. The control is authoritative; the indexes are derived atomically +with it. + +Workers poll due runnable work at the configured interval using Redis `TIME`. +Candidate reads are non-destructive. A watched control hash and monotonically +increasing fence make concurrent replica claims safe. Lease maintenance uses +the same interval and processes up to 100 expired fences. Delayed retries +remain in `runnable` with their Redis-server due timestamp and are invisible +until due. + +## Retention and rollout + +Terminal control/payload keys and their idempotency binding receive the +configured Redis TTL when a run first becomes terminal. Memory uses equivalent +internal cleanup. Do not delete records manually while they are nonterminal. + +Begin a new controlled-beta deployment with an empty durable namespace, then +retain its terminal records through the configured Redis TTL. + +## Health and incidents + +Health exposes `nonterminal`, `runnable`, and `lease_expiry`. Investigate a +growing runnable count, repeated lease recovery, store failures, or executions +that remain running/waiting longer than expected. Pause submissions, preserve +the Redis namespace, and inspect controls and fences before changing code or +restarting workers. + +Use Redis 6.2 or later. Monitor the durable counts alongside Redis availability +and latency to keep execution recovery healthy. diff --git a/docs/examples/overview.md b/docs/examples/overview.md index c8456f68..d7d2ba41 100644 --- a/docs/examples/overview.md +++ b/docs/examples/overview.md @@ -23,6 +23,8 @@ This page lists all maintained Hayhooks examples with detailed descriptions and | Example | Docs | Code | Description | |---|---|---|---| +| Durable Pipeline | [Pipeline wrapper](../concepts/pipeline-wrapper.md#durable-execution) | [GitHub](https://github.com/deepset-ai/hayhooks/tree/main/examples/durable_execution) | Typed REST wrapper with Redis-backed restart recovery, inspection, cancellation, and resume | +| Durable A2A Agent | [A2A Support](../features/a2a-support.md) | [GitHub](https://github.com/deepset-ai/hayhooks/tree/main/examples/a2a_long_running) | Haystack 3 Agent checkpoints, detached A2A tasks, progress, and restart recovery | | RAG: Indexing and Query with Elasticsearch | [rag-system.md](rag-system.md) | [GitHub](https://github.com/deepset-ai/hayhooks/tree/main/examples/rag_indexing_query) | Full indexing/query pipelines with Elasticsearch | | API Key Authentication | [advanced-configuration.md](../advanced/advanced-configuration.md) | [GitHub](https://github.com/deepset-ai/hayhooks/tree/main/examples/programmatic/api_key_auth) | Middleware-based API key auth with multi-key support and Swagger Authorize | diff --git a/docs/features/cli-commands.md b/docs/features/cli-commands.md index 3e9b0c9b..d2f6f427 100644 --- a/docs/features/cli-commands.md +++ b/docs/features/cli-commands.md @@ -143,6 +143,14 @@ hayhooks a2a run --host 0.0.0.0 --port 1418 | `--pipelines-dir` | | Directory for pipeline definitions | `./pipelines` | | `--additional-python-path` | | Additional Python path | `None` | | `--external-url` | | Base URL advertised in agent cards | `None` | +| `--task-store` | | Built-in A2A task-store backend: `auto`, `memory`, or `redis` | `auto` | +| `--a2a-redis-url` | | Redis URL for the built-in A2A task store | `redis://localhost:6379/0` | +| `--a2a-redis-key-prefix` | | Redis key prefix for the built-in A2A task store | `hayhooks:a2a` | +| `--execution-store` | | Built-in durable execution backend: `memory` or `redis` | `redis` | +| `--execution-redis-url` | | Redis URL for durable execution storage | `redis://localhost:6379/0` | +| `--execution-redis-key-prefix` | | Redis key prefix for durable execution storage | `hayhooks:durable` | +| `--durable-execution-concurrency` | | Maximum concurrent durable Agent executions per deployed agent | `1` | +| `--debug` | | Include tracebacks in errors | `false` | ## Pipeline Management Commands diff --git a/docs/getting-started/installation.md b/docs/getting-started/installation.md index d68539c5..fac962e3 100644 --- a/docs/getting-started/installation.md +++ b/docs/getting-started/installation.md @@ -38,7 +38,17 @@ This guide covers how to install Hayhooks and its dependencies. Includes all standard features plus [A2A Server](../features/a2a-support.md) support, exposing deployed pipelines and agents over the [A2A protocol](https://a2a-protocol.org) so other agents can discover - and delegate tasks to them. + and delegate tasks to them. This extra also includes Redis support for durable execution. + +=== "With Durable Execution" + + ```bash + pip install "hayhooks[durable]" + ``` + + Includes Redis support and Haystack 3, required for restart-safe Pipeline and Agent execution. See the + [durable Pipeline example](https://github.com/deepset-ai/hayhooks/tree/main/examples/durable_execution). + Durable execution requires Redis server 6.2 or newer. === "With Tracing Support" diff --git a/docs/reference/api-reference.md b/docs/reference/api-reference.md index eef7674b..dcfa949c 100644 --- a/docs/reference/api-reference.md +++ b/docs/reference/api-reference.md @@ -129,6 +129,24 @@ Execute a deployed pipeline. } ``` +#### Durable execution + +Wrappers that implement exactly one of `run_durable()` or +`run_durable_async()` expose these typed resources: + +| Endpoint | Description | +|---|---| +| `POST /{pipeline_name}/run-durable` | Validate and persist a detached execution; accepts an optional `Idempotency-Key` header | +| `GET /{pipeline_name}/executions/{execution_id}` | Inspect safe status, progress, waiting state, error, or result | +| `POST /{pipeline_name}/executions/{execution_id}/cancel` | Request cooperative cancellation | +| `POST /{pipeline_name}/executions/{execution_id}/resume` | Resume an execution waiting for input | + +Submission normally returns `202 Accepted` and a `Location` header. An +idempotent replay of a retained terminal execution returns `200 OK`. Validated +input, checkpoints, application state, ownership, and fence details remain +server-side. See [Pipeline wrapper durable execution](../concepts/pipeline-wrapper.md#durable-execution) +and the [durable engine contract](../advanced/durable-engine.md). + ### OpenAI Compatibility #### Chat Completion @@ -252,10 +270,7 @@ Currently, Hayhooks does not include built-in rate limiting. Consider implementi ```python import requests - response = requests.post( - "http://localhost:1416/chat_pipeline/run", - json={"query": "Hello!"} - ) + response = requests.post("http://localhost:1416/chat_pipeline/run", json={"query": "Hello!"}) print(response.json()) ``` @@ -287,12 +302,7 @@ Currently, Hayhooks does not include built-in rate limiting. Consider implementi response = requests.post( "http://localhost:1416/v1/chat/completions", - json={ - "model": "chat_pipeline", - "messages": [ - {"role": "user", "content": "Hello!"} - ] - } + json={"model": "chat_pipeline", "messages": [{"role": "user", "content": "Hello!"}]}, ) print(response.json()) ``` @@ -304,15 +314,10 @@ Currently, Hayhooks does not include built-in rate limiting. Consider implementi client = OpenAI( base_url="http://localhost:1416/v1", - api_key="not-needed" # Hayhooks doesn't require auth by default + api_key="not-needed", # Hayhooks doesn't require auth by default ) - response = client.chat.completions.create( - model="chat_pipeline", - messages=[ - {"role": "user", "content": "Hello!"} - ] - ) + response = client.chat.completions.create(model="chat_pipeline", messages=[{"role": "user", "content": "Hello!"}]) print(response.choices[0].message.content) ``` diff --git a/docs/reference/environment-variables.md b/docs/reference/environment-variables.md index 5e0cfd0b..fadb1cf8 100644 --- a/docs/reference/environment-variables.md +++ b/docs/reference/environment-variables.md @@ -141,6 +141,146 @@ export HAYHOOKS_DEPLOY_CONCURRENCY=parallel - Default: `true` - Description: Accept A2A spec 0.3 requests on the same endpoints. Many clients and tools (e.g. the a2a-inspector) still speak 0.3 during the 1.0 transition +### HAYHOOKS_A2A_TASK_STORE + +- Default: `auto` +- Description: Built-in A2A task-store backend +- Options: + - `auto`: Select Redis for Redis-backed durable A2A Agents, otherwise memory + - `memory`: Process-local task records; use only one Hayhooks process, not load-balanced replicas or restart-safe tasks + - `redis`: Persistent task records using the configured Redis URL and key prefix + +### HAYHOOKS_A2A_REDIS_URL + +- Default: `redis://localhost:6379/0` +- Description: Redis URL used by the built-in A2A task store + +### HAYHOOKS_A2A_REDIS_KEY_PREFIX + +- Default: `hayhooks:a2a` +- Description: Prefix applied to built-in Redis A2A task-store keys. Use a distinct prefix when multiple environments share Redis. + +### HAYHOOKS_A2A_REDIS_SOCKET_TIMEOUT / HAYHOOKS_A2A_REDIS_SOCKET_CONNECT_TIMEOUT + +- Default: `5.0` seconds each +- Description: Bound established-socket operations and new Redis connections for the built-in A2A task store. + +### HAYHOOKS_A2A_REDIS_HEALTH_CHECK_INTERVAL + +- Default: `30` seconds +- Description: Redis-py connection health-check interval for the built-in A2A task store. Set `0` to disable proactive checks. + +### HAYHOOKS_A2A_TERMINAL_TASK_TTL_SECONDS + +- Default: `604800` (seven days) +- Description: Retention window for terminal A2A tasks. Cleanup also removes their owner-update and recovery indexes. + +### HAYHOOKS_A2A_TASK_SNAPSHOT_CACHE_SIZE + +- Default: `1024` +- Description: Maximum loaded protobuf task snapshots retained per A2A Redis task-store instance for optimistic version checks. The cache is a global LRU across task IDs. + +### HAYHOOKS_A2A_LIST_SCAN_BATCH_SIZE + +- Default: `500` +- Description: Maximum task IDs and payloads loaded in one Redis batch while applying filtered A2A task-list queries. Exact filtered counts still require scanning the owner's update index. + +## Durable execution + +### HAYHOOKS_DURABLE_STORE + +- Default: `redis` +- Description: Built-in durable execution-store backend. Both choices use the same state reducer; Redis is required for restart recovery. +- Options: + - `memory`: Request-detached execution whose records are lost on process exit + - `redis`: Redis control records, runnable/lease indexes, checkpoints, cancellation, and restart recovery + +### HAYHOOKS_DURABLE_REDIS_URL + +- Default: `redis://localhost:6379/0` +- Description: Redis URL used by durable REST and A2A execution. + +### HAYHOOKS_DURABLE_REDIS_KEY_PREFIX + +- Default: `hayhooks:durable` +- Description: Key prefix used by durable execution records and queues. + +### HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT / HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT + +- Default: `5.0` seconds each +- Description: Bound established-socket operations and new Redis connections for durable execution. A timeout is reported as a store failure, causing worker backoff and readiness to return `503` rather than waiting indefinitely. + +### HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL + +- Default: `30` seconds +- Description: Redis-py connection health-check interval for durable execution. Set `0` to disable proactive checks. + +### HAYHOOKS_DURABLE_LEASE_DURATION_MS + +- Default: `30000` +- Description: Milliseconds for the renewable execution lease. Workers renew it at one-third of this duration. + +### HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS + +- Default: `1500` +- Description: Server-clock margin before a lease deadline within which an owned transition is rejected. This prevents a transition that is near expiry from committing after another worker may recover the lease. + +### HAYHOOKS_DURABLE_TERMINAL_TTL_SECONDS + +- Default: `604800` (seven days) +- Description: Retention period for terminal execution records. + +### HAYHOOKS_DURABLE_MAX_PROGRESS_EVENTS + +- Default: `100` +- Description: Maximum retained client-visible progress events per execution. + +### HAYHOOKS_DURABLE_MAX_RECORD_BYTES + +- Default: `1000000` +- Description: Maximum JSON payload size for validated input, a checkpoint, wait data, a result, or an error. The engine reserves space for each payload independently, so retained state remains bounded even while a run has input, a checkpoint, and progress history. + +### HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS + +- Default: `0` (disabled) +- Description: Maximum queued, running, or waiting executions for one logical deployment. New work is rejected atomically with `503` and `Retry-After` when the bound is full; idempotent replay remains available. + +### HAYHOOKS_DURABLE_SHUTDOWN_GRACE_PERIOD + +- Default: `5.0` +- Description: Seconds to wait for workers during shutdown. Synchronous work that exceeds the window retains its claim and heartbeat until the underlying thread exits. + +### HAYHOOKS_DURABLE_MAX_ATTEMPTS + +- Default: `3` +- Description: Maximum application attempts, including the first attempt and recovered abandoned claims, before retry exhaustion becomes terminal failure. + +### HAYHOOKS_DURABLE_RETRY_BASE_DELAY + +- Default: `1.0` +- Description: Base seconds for bounded exponential retry delay when application code does not provide a delay. + +### HAYHOOKS_DURABLE_RETRY_MAX_DELAY + +- Default: `60.0` +- Description: Maximum seconds for the next retry delay, including explicit application overrides. + +### HAYHOOKS_DURABLE_TRUSTED_OWNER_HEADER + +- Default: `""` (bearer execution-ID mode) +- Description: Trusted reverse-proxy header containing the authenticated owner. When set, submit, inspect, cancel, and resume enforce owner equality. +- Security: The proxy must remove client-supplied copies and inject this header over a trusted hop. + +### HAYHOOKS_DURABLE_EXECUTION_CONCURRENCY + +- Default: `1` +- Description: Maximum concurrent durable executions per deployment. Increase it only when the Pipeline or Agent and its shared dependencies are concurrency-safe. + +### HAYHOOKS_DURABLE_POLL_INTERVAL + +- Default: `1.0` +- Description: Shared worker and lease-maintenance polling interval in seconds. At concurrency 1, `0.25` uses about 16 idle Redis commands per second with 125 ms average pickup latency, `0.5` uses about 8 with 250 ms latency, and `1.0` uses about 4 with 500 ms latency. Counts are per durable deployment and replica. + ## Chainlit UI ### HAYHOOKS_CHAINLIT_ENABLED diff --git a/examples/README.md b/examples/README.md index 9889595e..39efef35 100644 --- a/examples/README.md +++ b/examples/README.md @@ -22,6 +22,8 @@ This directory contains various examples demonstrating different use cases and f | [responses_with_file_upload](./pipeline_wrappers/responses_with_file_upload/) | Agent-based Responses API with file reading | • Haystack Agent with `read_file` tool
• `run_response_async` with streaming
• `run_file_upload` with in-memory store
• `_strip_tool_calls` for agentic clients
• Codex CLI compatible | Building an agent that reads local files and uploaded files via the Responses API, compatible with Codex CLI and the OpenAI Python client | | [chat_completion_with_file_upload](./pipeline_wrappers/chat_completion_with_file_upload/) | Chat Completions API with `/v1/files` upload | • `run_chat_completion_async` with streaming
• `run_file_upload` with in-memory store
• Resolves `{"type": "file"}` content parts
• OpenAI file input format | Using the Chat Completions API with files uploaded via `/v1/files` and referenced using OpenAI's multi-part content format | | [a2a_multi_agent](./a2a_multi_agent/) | Two agents with their own MCP tools, communicating over A2A | • `hayhooks a2a run` hosting two agents
• Per-agent A2A agent cards
• Agent-to-agent delegation via A2A client tool
• One MCP tool server per agent (FastMCP)
• Streaming A2A client | Building multi-agent systems where Haystack Agents expose themselves over A2A and delegate tasks to each other while using MCP for their own tools | +| [a2a_long_running](./a2a_long_running/) | Recoverable OpenAI agent over A2A | • `OpenAIChatGenerator`-based Haystack Agent
• Redis-backed fenced execution checkpoints
• Hayhooks and client restart recovery
• `input-required` continuation
• Polling and durable cancellation | Building tool-using A2A agents whose accepted work resumes after process restarts | +| [durable_execution](./durable_execution/) | First-class durable Pipeline | • Typed `/run-durable` REST endpoint
• Built-in Redis store and fenced claims
• Restart recovery, inspection, cancellation, and resume | Running recoverable background jobs through an ordinary wrapper without A2A | | [rag_indexing_query](./rag_indexing_query/) | Complete RAG system with Elasticsearch | • Document indexing pipeline
• Query pipeline
• Elasticsearch integration
• Multiple file format support (PDF, Markdown, Text)
• Sentence transformers embeddings | Implementing production-ready RAG systems for document search and knowledge retrieval | | [shared_code_between_wrappers](./shared_code_between_wrappers/) | Code sharing between pipeline wrappers | • Shared library imports
• HAYHOOKS_ADDITIONAL_PYTHON_PATH
• Multiple deployment strategies
• Code reusability | Organizing complex projects with multiple pipelines that share common functionality | @@ -33,12 +35,27 @@ This directory contains various examples demonstrating different use cases and f ## Getting Started -Each example includes: +Examples intentionally stay lightweight. Every runnable example includes its +source files; examples that need non-default setup or dependencies also include +a dedicated README and/or `requirements.txt`. - **Pipeline wrapper implementation** (`pipeline_wrapper.py`) - **Pipeline configuration** (`.yml` files where applicable) -- **Dependencies** (`requirements.txt` where applicable) -- **Documentation** (individual README files with setup instructions) +- **Dependencies** (`requirements.txt` where a demo needs them) +- **Documentation** (individual README files where a demo needs dedicated setup instructions) + +For the durable examples, use this presentation order: + +1. [`durable_execution`](./durable_execution/) — the deterministic reference + for typed submission, retry, approval, checkpoints, crash recovery, and + cancellation. +2. [`a2a_long_running`](./a2a_long_running/) — durable Agent execution exposed + through standard A2A task lifecycle and continuation messages. + +Each durable example's Compose file publishes Redis on `localhost:6379`. +Run these examples one at a time, or change the host port and corresponding +`HAYHOOKS_DURABLE_REDIS_URL`. Each stack has its own named volume; `compose +down` retains it and `compose down -v` resets it. ## Common Prerequisites @@ -54,7 +71,7 @@ Most examples require: 1. Navigate to the `/examples` directory 2. Create and activate a virtual environment (recommended) 3. Install dependencies: `pip install -r requirements.txt` (if present) -4. Follow the specific example's README for deployment and testing +4. Follow the example-specific README when present; otherwise deploy its wrapper using the standard Hayhooks command ## Support From 683423f9d4cddee3c019f0b8906ccb27787be75a Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:10:29 +0200 Subject: [PATCH 11/28] ci: validate durable execution support --- .github/workflows/docs.yml | 3 +++ .github/workflows/pypi.yml | 5 ++++- .github/workflows/tests.yml | 33 +++++++++++++++++++++++++++------ 3 files changed, 34 insertions(+), 7 deletions(-) diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index fc03fdc8..0f9775f6 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -32,5 +32,8 @@ jobs: git config --global user.name "github-actions[bot]" git config --global user.email "github-actions[bot]@users.noreply.github.com" + - name: Validate docs in strict mode + run: hatch run docs:build --strict + - name: Deploy docs to GitHub Pages run: hatch run docs:deploy diff --git a/.github/workflows/pypi.yml b/.github/workflows/pypi.yml index ad055038..4e496e11 100644 --- a/.github/workflows/pypi.yml +++ b/.github/workflows/pypi.yml @@ -5,6 +5,9 @@ on: tags: - "v[0-9].[0-9]+.[0-9]+*" +env: + HATCH_VERSION: "1.16.5" + jobs: release-on-pypi: runs-on: ubuntu-latest @@ -17,7 +20,7 @@ jobs: uses: actions/checkout@f43a0e5ff2bd294095638e18286ca9a3d1956744 # v3.6.0 - name: Install Hatch - run: pip install hatch + run: pip install hatch==${{ env.HATCH_VERSION }} - name: Build run: hatch build diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 9a1e6c86..b00d5e59 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -40,6 +40,18 @@ jobs: tests-haystack-v3: runs-on: ubuntu-latest + services: + redis: + image: redis:6.2-alpine + ports: + - 6379:6379 + options: >- + --health-cmd "redis-cli ping" + --health-interval 2s + --health-timeout 2s + --health-retries 20 + env: + HAYHOOKS_TEST_REDIS_URL: redis://127.0.0.1:6379/15 strategy: matrix: python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"] @@ -55,15 +67,21 @@ jobs: - name: Install Haystack v3 into the test env run: | - hatch env create test - hatch -e test env run -- uv pip install --upgrade "haystack-ai>=3" - hatch -e test env run -- python -c "import haystack, sys; v = haystack.__version__; print('haystack', v); sys.exit(0 if int(v.split('.')[0]) >= 3 else 'expected Haystack v3+, got ' + v)" + hatch env create test-v3 + hatch -e test-v3 env run -- uv pip install --upgrade "haystack-ai>=3" + hatch -e test-v3 env run -- python -c "import haystack, sys; v = haystack.__version__; print('haystack', v); sys.exit(0 if int(v.split('.')[0]) >= 3 else 'expected Haystack v3+, got ' + v)" - - name: Run unit tests - run: hatch run test:unit + - name: Run tests + run: hatch run test-v3:all + + - name: Run durable process-recovery smoke test + if: matrix.python-version == '3.12' + env: + HAYHOOKS_TEST_PROCESS_RECOVERY: "1" + run: hatch run test-v3:all tests/test_durable_process_recovery.py - name: Ty - check types (Haystack v3) - run: hatch run test:types + run: hatch run test-v3:types linting: runs-on: ubuntu-slim @@ -83,6 +101,9 @@ jobs: - name: Ty - check types run: hatch run test:types + - name: Build documentation in strict mode + run: hatch run docs:build --strict + dashboard-tests: runs-on: ubuntu-latest steps: From dfd35a078e314554fee2fc6d45f018a18a18d970 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 12 Aug 2026 11:56:58 +0200 Subject: [PATCH 12/28] fix(durable): harden lifecycle edge cases --- docs/advanced/durable-engine.md | 14 +++-- docs/advanced/durable-execution-operations.md | 10 ++-- docs/features/a2a-support.md | 2 +- src/hayhooks/durable/manager.py | 29 ++++++--- src/hayhooks/durable/runtime.py | 20 +++++-- src/hayhooks/durable/store.py | 32 ++++++++-- src/hayhooks/server/a2a/durable_executor.py | 55 +++++++++++++---- tests/test_durable_a2a.py | 22 ++++++- tests/test_durable_deployment_lifecycle.py | 2 +- tests/test_durable_store.py | 60 +++++++++++++++++++ 10 files changed, 202 insertions(+), 44 deletions(-) diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 2fada2ce..5d72981a 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -156,7 +156,8 @@ The runtime owns provider shutdown. Built-in providers snapshot their settings, and the runtime adopts that snapshot when a provider is supplied. Pass custom settings once—either to a built-in provider as above, or to `DurableRuntime` when it selects the default provider. Conflicting runtime and provider settings -are rejected before a deployment is created. +are rejected before a deployment is created. The selected provider is fixed for +the runtime's lifetime; create a new runtime to change storage backends. ## Redis layout @@ -171,7 +172,9 @@ sorted-set indexes: The namespace also contains a `capacity` hash with only `nonterminal` and one idempotency binding per execution. Terminal execution and idempotency keys use -native Redis TTL. The in-memory backend schedules equivalent cleanup. +native Redis TTL. The in-memory backend schedules equivalent cleanup. The store +assigns progress sequence numbers when it commits each transition, preserving +checkpoint and cancellation progress when they race. Workers poll `runnable` every configured poll interval, one second by default. They use Redis `TIME` and read one due member without removing it. Multiple @@ -195,8 +198,11 @@ and resumes verify that persisted work matches the active revision. ## Operations `DurableExecutionManager.health_snapshot()` reports `nonterminal`, `runnable`, -and `lease_expiry`. Alert on sustained runnable growth, repeated lease recovery, -worker/store health failures, and runs that exceed their expected duration. +`lease_expiry`, and the current worker store-error streak. Repeated claim or +transition failures make readiness unhealthy until a worker completes a store +operation successfully. Alert on sustained runnable growth, repeated lease +recovery, worker/store health failures, and runs that exceed their expected +duration. See [Durable execution operations](durable-execution-operations.md) for deployment, retention, and incident guidance. diff --git a/docs/advanced/durable-execution-operations.md b/docs/advanced/durable-execution-operations.md index 1e4cf652..4bf32968 100644 --- a/docs/advanced/durable-execution-operations.md +++ b/docs/advanced/durable-execution-operations.md @@ -43,10 +43,12 @@ retain its terminal records through the configured Redis TTL. ## Health and incidents -Health exposes `nonterminal`, `runnable`, and `lease_expiry`. Investigate a -growing runnable count, repeated lease recovery, store failures, or executions -that remain running/waiting longer than expected. Pause submissions, preserve -the Redis namespace, and inspect controls and fences before changing code or +Health exposes `nonterminal`, `runnable`, `lease_expiry`, and +`worker_store_error_streak`. A claim or transition store failure makes readiness +unhealthy until that worker completes a store operation successfully. Investigate +a growing runnable count, repeated lease recovery, store failures, or executions +that remain running/waiting longer than expected. Pause submissions, preserve the +Redis namespace, and inspect controls and fences before changing code or restarting workers. Use Redis 6.2 or later. Monitor the durable counts alongside Redis availability diff --git a/docs/features/a2a-support.md b/docs/features/a2a-support.md index 3dd08106..53dc9d9f 100644 --- a/docs/features/a2a-support.md +++ b/docs/features/a2a-support.md @@ -223,7 +223,7 @@ waiting for input. A follow-up sent while it is running is rejected rather than being treated as an idempotent redelivery; clients should wait for `INPUT_REQUIRED` before continuing a task. -Terminal tasks use `HAYHOOKS_A2A_TERMINAL_TASK_TTL_SECONDS`. Runtime maintenance performs cleanup even when no later A2A request arrives and removes the protobuf payload and task indexes. Execution-record retention remains independent. +Terminal tasks use `HAYHOOKS_A2A_TERMINAL_TASK_TTL_SECONDS`. Runtime maintenance performs cleanup even when no later A2A request arrives and removes the protobuf payload and task indexes. Execution-record retention remains independent. If a task projection expires first, `GetTask` with the original task ID reconstructs its current state from the retained execution; the expired history and list entry remain gone. Applications constructing the server directly can use a configured provider: diff --git a/src/hayhooks/durable/manager.py b/src/hayhooks/durable/manager.py index bae7de29..69e1b8ce 100644 --- a/src/hayhooks/durable/manager.py +++ b/src/hayhooks/durable/manager.py @@ -102,6 +102,7 @@ def __init__( # noqa: PLR0913 self._draining_runs: set[asyncio.Future[JsonValue]] = set() self._maintenance_task: asyncio.Task[None] | None = None self._maintenance_error_streak = 0 + self._worker_store_error_streaks: dict[str, int] = {} self._prepared = False self._started = False self._accepting_claims = False @@ -136,6 +137,7 @@ def activate(self) -> None: self._accepting_claims = True self._submission_gate.activate() self._worker_generation += 1 + self._worker_store_error_streaks.clear() generation = self._worker_generation identity = f"{socket.gethostname()}-{uuid.uuid4().hex[:8]}" self._workers = [self._start_worker(identity, slot, generation) for slot in range(self.concurrency)] @@ -168,14 +170,22 @@ def health(self) -> dict[str, JsonValue]: maintenance_healthy = self._maintenance_task is None or ( maintenance_running and self._maintenance_error_streak == 0 ) + worker_store_error_streak = max(self._worker_store_error_streaks.values(), default=0) return { "healthy": not self._prepared - or (self._started and self._accepting_claims and running == self.concurrency and maintenance_healthy), + or ( + self._started + and self._accepting_claims + and running == self.concurrency + and maintenance_healthy + and worker_store_error_streak == 0 + ), "configured_slots": self.concurrency, "running_slots": running, "draining_slots": sum(not worker.done() for worker in self._draining_workers), "draining_runs": sum(not runner.done() for runner in self._draining_runs), "maintenance_running": maintenance_running, + "worker_store_error_streak": worker_store_error_streak, "accepting": self.accepting, } @@ -308,16 +318,17 @@ def _log_worker_failures(self, workers: set[asyncio.Task[None]]) -> None: ) async def _worker(self, worker_name: str, generation: int) -> None: - consecutive_store_errors = 0 + self._worker_store_error_streaks[worker_name] = 0 while self._accepting_claims and generation == self._worker_generation: try: claim = await self.store.claim_next(worker_name) except asyncio.CancelledError: raise except Exception as error: - consecutive_store_errors += 1 - await self._backoff_store_error(error, consecutive_store_errors, operation="claim") + self._worker_store_error_streaks[worker_name] += 1 + await self._backoff_store_error(error, self._worker_store_error_streaks[worker_name], operation="claim") continue + self._worker_store_error_streaks[worker_name] = 0 if claim is None: await asyncio.sleep(self.poll_interval) continue @@ -331,16 +342,18 @@ async def _worker(self, worker_name: str, generation: int) -> None: attempt_log.debug("Claimed durable execution") try: await self._process_claim(claim) - consecutive_store_errors = 0 + self._worker_store_error_streaks[worker_name] = 0 attempt_log.bind(status=claim.record.status.value).debug("Finished durable execution attempt") except ExecutionLeaseLostError: - consecutive_store_errors = 0 + self._worker_store_error_streaks[worker_name] = 0 log.warning("{} | lost durable execution claim {}", self.name, claim.record.execution_id) except asyncio.CancelledError: raise except Exception as error: - consecutive_store_errors += 1 - await self._backoff_store_error(error, consecutive_store_errors, operation="transition") + self._worker_store_error_streaks[worker_name] += 1 + await self._backoff_store_error( + error, self._worker_store_error_streaks[worker_name], operation="transition" + ) async def _process_claim(self, claim: ExecutionClaim) -> None: # noqa: C901, PLR0912 - explicit attempt outcomes """Run one fenced claim and leave every terminal decision to the store.""" diff --git a/src/hayhooks/durable/runtime.py b/src/hayhooks/durable/runtime.py index ecab856c..7b6eec12 100644 --- a/src/hayhooks/durable/runtime.py +++ b/src/hayhooks/durable/runtime.py @@ -341,6 +341,7 @@ class _DurableAgentRequest(BaseModel): """Private A2A input mapping; REST wrappers always provide their own model.""" messages: list[dict[str, Any]] = Field(min_length=1) + a2a_context_id: str | None = Field(default=None, min_length=1) class DurableRuntime: @@ -352,7 +353,7 @@ def __init__( *, app_settings: AppSettings | None = None, ) -> None: - self.provider = provider + self._store_provider = provider self._app_settings = _runtime_settings(provider, app_settings) self._deployments: dict[str, DurableDeployment] = {} self._started = False @@ -366,6 +367,11 @@ def has_capability(self, wrapper: BasePipelineWrapper) -> bool: def started(self) -> bool: return self._started + @property + def provider(self) -> ExecutionStoreProvider | None: + """Return the runtime-owned provider selected at construction or first use.""" + return self._store_provider + @property def app_settings(self) -> AppSettings: """Return configured or provider settings, falling back to Hayhooks' global settings.""" @@ -445,7 +451,7 @@ async def close(self) -> None: self._deployments.clear() if self.provider is not None: provider = self.provider - self.provider = None + self._store_provider = None draining = [deployment.manager for deployment in deployments if deployment.manager.draining] if draining: @@ -477,13 +483,15 @@ async def health(self) -> dict[str, JsonValue]: } def _provider(self) -> ExecutionStoreProvider: - if self.provider is None: + provider = self._store_provider + if provider is None: if self.app_settings.durable_store == "memory": log.warning("Durable execution uses volatile in-memory storage; queued work is lost on process exit") - self.provider = InMemoryExecutionStoreProvider(app_settings=self.app_settings) + provider = InMemoryExecutionStoreProvider(app_settings=self.app_settings) else: - self.provider = RedisExecutionStoreProvider(app_settings=self.app_settings) - return self.provider + provider = RedisExecutionStoreProvider(app_settings=self.app_settings) + self._store_provider = provider + return provider class IdempotencyConflictError(RuntimeError): diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py index 1e0f1307..e5f65489 100644 --- a/src/hayhooks/durable/store.py +++ b/src/hayhooks/durable/store.py @@ -32,10 +32,12 @@ normalize_cancellation_reason, ) from hayhooks.durable.engine import ExecutionStatus as EngineStatus +from hayhooks.durable.engine import ProgressEvent as EngineProgressEvent from hayhooks.durable.models import ( DEFAULT_MAX_PROGRESS_BYTES, ExecutionError, ExecutionKind, + ExecutionProgressEvent, ExecutionRecord, ExecutionRecordSizeError, ExecutionStatus, @@ -120,7 +122,7 @@ def __init__( self._lost = False self._lost_event = asyncio.Event() self._confirmed_until = confirmed_at + self.store.lease_safe_duration - self._persisted_progress_sequence = control.progress_sequence + self._last_persisted_progress = record.progress[-1] if record.progress else None @property def record(self) -> ExecutionRecord: @@ -138,7 +140,7 @@ async def __aenter__(self) -> ExecutionClaim: ) return self - async def __aexit__(self, exc_type: Any, exc: Any, traceback: Any) -> None: + async def __aexit__(self, _exc_type: Any, _exc: Any, _traceback: Any) -> None: if self._heartbeat is not None: self._heartbeat.cancel() with suppress(asyncio.CancelledError): @@ -244,7 +246,7 @@ async def _transition(self, command: Any) -> Any: except ExecutionLeaseLostError: self._mark_lost() raise - self._sync(plan.next_control, confirmed_at=confirmed_at) + self._sync(plan.next_control, confirmed_at=confirmed_at, progress_events=plan.progress_events) return plan async def _heartbeat_loop(self) -> None: @@ -261,18 +263,36 @@ async def _heartbeat_loop(self) -> None: continue def _new_progress(self) -> tuple[bytes, ...]: - events = tuple(event for event in self.record.progress if event.sequence > self._persisted_progress_sequence) + events = self._new_progress_events() return tuple( _encode(event.to_dict(), limit=self.store.config.max_progress_event_bytes, label="progress") for event in events ) - def _sync(self, control: Any, *, confirmed_at: float | None = None) -> None: + def _new_progress_events(self) -> tuple[ExecutionProgressEvent, ...]: + if self._last_persisted_progress is None: + return tuple(self.record.progress) + for index, event in enumerate(self.record.progress): + if event is self._last_persisted_progress: + return tuple(self.record.progress[index + 1 :]) + return tuple(self.record.progress) + + def _sync( + self, + control: Any, + *, + confirmed_at: float | None = None, + progress_events: tuple[EngineProgressEvent, ...] = (), + ) -> None: + if progress_events: + pending = self._new_progress_events() + for event, persisted in zip(pending, progress_events, strict=True): + event.sequence = persisted.sequence + self._last_persisted_progress = self.record.progress[-1] self.control = control self.record.attempt = control.run_attempt self.record.sequence = control.version self.record.status = control.status - self._persisted_progress_sequence = control.progress_sequence if confirmed_at is not None: self._confirmed_until = confirmed_at + self.store.lease_safe_duration diff --git a/src/hayhooks/server/a2a/durable_executor.py b/src/hayhooks/server/a2a/durable_executor.py index 46168f1b..747ea238 100644 --- a/src/hayhooks/server/a2a/durable_executor.py +++ b/src/hayhooks/server/a2a/durable_executor.py @@ -66,8 +66,31 @@ async def save(self, task: Any, context: Any) -> None: await self._task_store.save(task, context) async def get(self, task_id: str, context: Any) -> Any | None: + owner_id = self.owner_id_for_context(context) task = await self._task_store.get(task_id, context) - return await self._project(task, self.owner_id_for_context(context), context) + record = None + if task is None: + try: + record = await self._deployment.get( + execution_id_for(owner_id, task_id), + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) + except KeyError: + return None + context_id = record.validated_input.get("a2a_context_id") + if not isinstance(context_id, str) or not context_id: + return None + from a2a.types import Task, TaskState + + task = Task(id=task_id, context_id=context_id) + task.status.state = ( + TaskState.TASK_STATE_WORKING + if record.status is ExecutionStatus.RUNNING + else TaskState.TASK_STATE_SUBMITTED + ) + return await self._project(task, owner_id, context, record=record) async def list(self, params: Any, context: Any) -> Any: from a2a.types import ListTasksResponse @@ -132,7 +155,9 @@ async def _all_tasks(self, params: Any, context: Any) -> builtins.list[Any]: return tasks request.page_token = page.next_page_token - async def _project(self, task: Any | None, owner_id: str, context: Any) -> Any | None: # noqa: C901 + async def _project( # noqa: C901 + self, task: Any | None, owner_id: str, context: Any, *, record: Any | None = None + ) -> Any | None: if task is None: return None projected = type(task)() @@ -142,12 +167,13 @@ async def _project(self, task: Any | None, owner_id: str, context: Any) -> Any | copy_version(task, projected) settled = False try: - record = await self._deployment.get( - execution_id_for(owner_id, task.id), - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) + if record is None: + record = await self._deployment.get( + execution_id_for(owner_id, task.id), + owner_id=owner_id, + enforce_owner=True, + allow_revision_mismatch=True, + ) except KeyError: if task_is_terminal(task): return projected @@ -239,7 +265,7 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non action = "submit" await self._task_store.save(task, context.call_context) try: - record = await self._submit(task.id, owner_id, build_haystack_messages(context)) + record = await self._submit(task, owner_id, build_haystack_messages(context)) except ValueError as error: log.bind( pipeline_name=self.pipeline_name, @@ -330,11 +356,14 @@ async def _wait_for_update( return await asyncio.sleep(max(0.1, settings.durable_poll_interval)) - async def _submit(self, task_id: str, owner_id: str, messages: list[Any]) -> Any: - payload = {"messages": [message.to_dict() for message in messages]} + async def _submit(self, task: Any, owner_id: str, messages: list[Any]) -> Any: + payload = { + "messages": [message.to_dict() for message in messages], + "a2a_context_id": task.context_id, + } while not self._closed: try: - return (await self.deployment.submit(payload, execution_id=task_id, owner_id=owner_id))[1] + return (await self.deployment.submit(payload, execution_id=task.id, owner_id=owner_id))[1] except (ExecutionAdmissionError, ExecutionStoreError): await asyncio.sleep(max(0.1, settings.durable_poll_interval)) raise asyncio.CancelledError @@ -360,7 +389,7 @@ async def _recover_tasks(self, task_store: RecoverableTaskStore) -> None: # noq if task.status.state != TaskState.TASK_STATE_SUBMITTED: continue try: - record = await self._submit(task.id, owner_id, build_haystack_task_messages(task)) + record = await self._submit(task, owner_id, build_haystack_task_messages(task)) except ValueError as error: record = None log.bind( diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py index 3fdc8060..f3cbfcf1 100644 --- a/tests/test_durable_a2a.py +++ b/tests/test_durable_a2a.py @@ -33,6 +33,7 @@ def __init__(self, status=ExecutionStatus.COMPLETED) -> None: result={"last_message": {"content": "recovered"}}, error=None, sequence=0, + validated_input={}, ) self.execution_id = None self.submitted_payload = None @@ -44,6 +45,7 @@ async def start(self): async def submit(self, payload, *, execution_id=None, owner_id=None): self.submitted_payload = payload + self.record.validated_input = payload self.execution_id = execution_id_for(owner_id, execution_id) if owner_id else execution_id self.record.execution_id = self.execution_id return True, self.record @@ -191,6 +193,23 @@ def test_a2a_http_reads_completion_from_durable_execution(monkeypatch, http_stor assert completed["artifacts"][-1]["name"] == "durable-result" +def test_expired_task_projection_uses_the_retained_execution(monkeypatch, http_store) -> None: + deployment = _Deployment() + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + completed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) + http_store.tasks.clear() + deployment.submit = AsyncMock(side_effect=AssertionError("retained execution must not be resubmitted")) + replayed = _response_task(client.post("/durable-agent/", json=_get_payload(completed["id"]))) + + assert replayed["status"]["state"] == "TASK_STATE_COMPLETED" + assert replayed["contextId"] == completed["contextId"] + assert replayed["artifacts"][-1]["name"] == "durable-result" + deployment.submit.assert_not_awaited() + assert not http_store.tasks + + def test_a2a_http_waiting_task_resumes_with_only_the_follow_up(monkeypatch, http_store) -> None: deployment = _Deployment(status=ExecutionStatus.WAITING) app = _http_app(http_store, deployment, monkeypatch) @@ -345,7 +364,8 @@ async def test_durable_a2a_submission_retries_transient_failures(monkeypatch, ht sleep = AsyncMock() monkeypatch.setattr(asyncio, "sleep", sleep) - record = await DurableAgentExecutor("agent", http_store, deployment)._submit("task", "owner", []) + task = SimpleNamespace(id="task", context_id="context") + record = await DurableAgentExecutor("agent", http_store, deployment)._submit(task, "owner", []) assert record is deployment.record assert deployment.submit.await_count == 2 diff --git a/tests/test_durable_deployment_lifecycle.py b/tests/test_durable_deployment_lifecycle.py index 4be2faa6..25097636 100644 --- a/tests/test_durable_deployment_lifecycle.py +++ b/tests/test_durable_deployment_lifecycle.py @@ -338,7 +338,7 @@ async def close(self): def test_store_initialization_failure_never_publishes_candidate(monkeypatch) -> None: - monkeypatch.setattr(durable_runtime, "provider", _FailingProvider()) + monkeypatch.setattr(durable_runtime, "_store_provider", _FailingProvider()) app = create_app() with TestClient(app) as client: failed = _deploy(client, _durable_source(field="value", increment=1, revision="first")) diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py index d1ae29e2..278fda65 100644 --- a/tests/test_durable_store.py +++ b/tests/test_durable_store.py @@ -5,6 +5,7 @@ import asyncio import time from dataclasses import replace +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest @@ -140,6 +141,20 @@ def test_runtime_rejects_conflicting_builtin_provider_settings() -> None: DurableRuntime(provider, app_settings=AppSettings(durable_lease_duration_ms=60_000)) +async def test_runtime_provider_cannot_be_replaced(monkeypatch: pytest.MonkeyPatch) -> None: + provider = InMemoryExecutionStoreProvider() + close = AsyncMock() + monkeypatch.setattr(provider, "close", close) + runtime = DurableRuntime(provider) + + with pytest.raises(AttributeError): + setattr(runtime, "provider", InMemoryExecutionStoreProvider()) + + assert runtime.provider is provider + await runtime.close() + close.assert_awaited_once() + + async def test_store_preserves_public_checkpoint_progress_wait_resume_and_result_contract() -> None: store = _store( lease_duration_ms=10_000, @@ -221,6 +236,27 @@ async def gated_transition(run_id, command, *, candidate=False): assert [event.sequence for event in record.progress] == [1, 2] +async def test_checkpoint_keeps_progress_added_after_a_concurrent_cancellation() -> None: + store = _store() + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + + async with claim: + claim.record.append_progress("before cancellation") + assert await store.request_cancel("run_1") + await claim.checkpoint() + claim.record.append_progress("after cancellation") + await claim.checkpoint() + + record = await store.get("run_1") + assert record is not None + assert [(event.sequence, event.message) for event in record.progress] == [ + (2, "before cancellation"), + (3, "after cancellation"), + ] + + async def test_losing_resume_race_returns_false(monkeypatch: pytest.MonkeyPatch) -> None: core = InMemoryExecutionStore(deployment="deployment", config=_config()) store = ExecutionStore(core, definition_revision="rev-1") @@ -311,6 +347,30 @@ async def runner(context): await manager.close() +async def test_manager_health_reports_worker_store_failures() -> None: + store = SimpleNamespace( + initialize=AsyncMock(), + claim_next=AsyncMock(side_effect=ExecutionStoreError("claim failed")), + maintain=AsyncMock(), + operational_counts=AsyncMock(return_value={"nonterminal": 1, "runnable": 1, "lease_expiry": 0}), + ) + manager = DurableExecutionManager( + "deployment", store, AsyncMock(), adapter=object(), poll_interval=0.001, shutdown_grace_period=0.01 + ) + await manager.start() + try: + for _ in range(100): + if store.claim_next.await_count: + break + await asyncio.sleep(0.001) + health = await manager.health_snapshot() + finally: + await manager.close() + + assert not health["healthy"] + assert health["worker_store_error_streak"] >= 1 + + async def test_canceled_runner_restarts_its_worker_slot() -> None: store = _store( config=replace(_config(), lease_commit_safety_ms=1), From a3285c2a2e85bb2a9259aa657504a7e122afc569 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Thu, 13 Aug 2026 10:20:37 +0200 Subject: [PATCH 13/28] docs(durable): fix fragment and docstring placement in example --- docs/advanced/durable-engine.md | 2 +- .../pipelines/long_running_agent/pipeline_wrapper.py | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 5d72981a..7bfd4804 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -23,7 +23,7 @@ its next checkpoint. ## Supported scope and tradeoffs -- durable input before submission succeeds; +- validated durable input before submission succeeds; - at-least-once execution with one fenced worker owner at a time; - checkpoints, progress, retries, cancellation, wait/resume, and terminal results; and diff --git a/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py b/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py index cd952d6e..dc516a33 100644 --- a/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py +++ b/examples/a2a_long_running/pipelines/long_running_agent/pipeline_wrapper.py @@ -156,9 +156,10 @@ def publish_indexing_receipt( class PipelineWrapper(A2APipelineWrapper): - durable_revision = "a2a-long-running-agent" """Let Hayhooks map this real tool-using Agent to durable A2A executions.""" + durable_revision = "a2a-long-running-agent" + def setup(self) -> None: self.pipeline = Agent( chat_generator=OpenAIChatGenerator(model="gpt-4o-mini"), From fc4cf771b904de5ac5c4aa404545fe4bd96c4fc8 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Thu, 13 Aug 2026 15:21:13 +0200 Subject: [PATCH 14/28] Fix confirmed durable execution review findings --- src/hayhooks/durable/backend.py | 3 +- src/hayhooks/durable/engine.py | 31 +++- src/hayhooks/durable/manager.py | 2 - src/hayhooks/durable/redis.py | 17 +- src/hayhooks/durable/runtime.py | 1 - src/hayhooks/durable/store.py | 184 ++++++++++++-------- src/hayhooks/server/a2a/durable_executor.py | 34 ++-- src/hayhooks/server/a2a/executor.py | 1 - src/hayhooks/server/durable/routes.py | 15 +- src/hayhooks/server/pipelines/models.py | 10 +- src/hayhooks/server/routers/status.py | 2 - src/hayhooks/server/utils/deploy_utils.py | 37 ++-- tests/test_a2a.py | 8 +- tests/test_deploy_performance.py | 20 +++ tests/test_deploy_utils.py | 11 ++ tests/test_durable_a2a.py | 27 +++ tests/test_durable_execution.py | 33 +++- tests/test_durable_redis_codec.py | 24 ++- tests/test_durable_store.py | 145 ++++++++++++++- tests/test_it_status.py | 13 ++ tests/test_redis_execution_integration.py | 18 ++ 21 files changed, 502 insertions(+), 134 deletions(-) diff --git a/src/hayhooks/durable/backend.py b/src/hayhooks/durable/backend.py index 26d74bae..53675b46 100644 --- a/src/hayhooks/durable/backend.py +++ b/src/hayhooks/durable/backend.py @@ -17,6 +17,7 @@ Fail, Heartbeat, PayloadKind, + ReleaseClaim, RequestCancellation, Resume, ScheduleRetry, @@ -135,7 +136,7 @@ def parse_lease_member(value: str) -> tuple[str, int]: def bind_command(command: ExecutionCommand, *, now_ms: int, lease_commit_safety_ms: int) -> ExecutionCommand: """Apply the backend clock and lease safety policy before reduction.""" - if isinstance(command, (Heartbeat, Checkpoint, ScheduleRetry, Suspend, Complete, Fail)): + if isinstance(command, (ReleaseClaim, Heartbeat, Checkpoint, ScheduleRetry, Suspend, Complete, Fail)): return replace(command, now_ms=now_ms, lease_commit_safety_ms=lease_commit_safety_ms) return replace(command, now_ms=now_ms) diff --git a/src/hayhooks/durable/engine.py b/src/hayhooks/durable/engine.py index 34b9d5d0..3d1cff34 100644 --- a/src/hayhooks/durable/engine.py +++ b/src/hayhooks/durable/engine.py @@ -174,6 +174,14 @@ class Claim: worker_revision: str +@dataclass(frozen=True, slots=True) +class ReleaseClaim: + fence: int + worker_id: str + now_ms: int = 0 + lease_commit_safety_ms: int = 0 + + @dataclass(frozen=True, slots=True) class Heartbeat: fence: int @@ -229,6 +237,7 @@ class Resume: worker_revision: str checkpoint: bytes | None = None progress_events: tuple[bytes, ...] = () + expected_version: int | None = None @dataclass(frozen=True, slots=True) @@ -262,6 +271,7 @@ class RecoverExpiredLease: ExecutionCommand = ( Claim + | ReleaseClaim | Heartbeat | Checkpoint | RequestCancellation @@ -316,6 +326,8 @@ def decide(control: ExecutionControl, command: ExecutionCommand) -> TransitionPl """ if isinstance(command, Claim): return _claim(control, command) + if isinstance(command, ReleaseClaim): + return _release_claim(control, command) if isinstance(command, Heartbeat): _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) return TransitionPlan( @@ -391,6 +403,19 @@ def _claim(control: ExecutionControl, command: Claim) -> TransitionPlan: return TransitionPlan(next_control, lease_index_update=LeaseIndexUpdate(deadline, next_control.fence)) +def _release_claim(control: ExecutionControl, command: ReleaseClaim) -> TransitionPlan: + _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + next_control = _business( + control, + command.now_ms, + status=ExecutionStatus.QUEUED, + run_attempt=control.run_attempt - 1, + lease_owner=None, + lease_expires_at_ms=None, + ) + return TransitionPlan(next_control, lease_index_update=LeaseIndexUpdate(None, control.fence)) + + def _cancel(control: ExecutionControl, command: RequestCancellation) -> TransitionPlan: if control.terminal or control.cancel_requested_at_ms is not None: return TransitionPlan(control) @@ -476,12 +501,12 @@ def _suspend(control: ExecutionControl, command: Suspend) -> TransitionPlan: def _resume(control: ExecutionControl, command: Resume) -> TransitionPlan: + if command.expected_version is not None and control.version != command.expected_version: + raise InvalidExecutionTransitionError("execution changed before it could resume") if control.status is not ExecutionStatus.WAITING: raise InvalidExecutionTransitionError("only waiting executions can resume") if control.definition_revision != command.worker_revision: - return _terminal( - control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, b"definition revision is incompatible" - ) + raise InvalidExecutionTransitionError("definition revision is incompatible") if control.cancel_requested_at_ms is not None: return _terminal(control, command.now_ms, ExecutionStatus.CANCELED, None, None) writes = (PayloadWrite(PayloadKind.CHECKPOINT, command.checkpoint),) if command.checkpoint is not None else () diff --git a/src/hayhooks/durable/manager.py b/src/hayhooks/durable/manager.py index 69e1b8ce..a6b2ea2e 100644 --- a/src/hayhooks/durable/manager.py +++ b/src/hayhooks/durable/manager.py @@ -332,8 +332,6 @@ async def _worker(self, worker_name: str, generation: int) -> None: if claim is None: await asyncio.sleep(self.poll_interval) continue - if not self._accepting_claims or generation != self._worker_generation: - return attempt_log = log.bind( deployment=self.name, execution_id=claim.record.execution_id, diff --git a/src/hayhooks/durable/redis.py b/src/hayhooks/durable/redis.py index 41d90ff3..a22b303c 100644 --- a/src/hayhooks/durable/redis.py +++ b/src/hayhooks/durable/redis.py @@ -34,6 +34,7 @@ ExecutionNotFoundError, ExecutionPayloadSizeError, ExecutionStatus, + Heartbeat, InvalidExecutionTransitionError, PayloadKind, TransitionPlan, @@ -271,7 +272,9 @@ async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> di async def read_progress(self, run_id: str) -> list[bytes]: return [bytes(value) for value in await self.redis.lrange(self.keys.progress(run_id), 0, -1)] - async def transition(self, run_id: str, command: ExecutionCommand, *, candidate: bool = False) -> TransitionPlan: + async def transition( # noqa: C901 + self, run_id: str, command: ExecutionCommand, *, candidate: bool = False + ) -> TransitionPlan: validate_command_payloads(command, self.config) control_key = self.keys.control(run_id) for attempt in range(self.config.transaction_max_retries): @@ -311,7 +314,17 @@ async def transition(self, run_id: str, command: ExecutionCommand, *, candidate: await pipe.execute() return TransitionPlan(current) pipe.multi() - self._apply_plan(pipe, current, plan) + if isinstance(command, Heartbeat): + lease = plan.lease_index_update + if lease is None or lease.deadline_ms is None: + raise AssertionError("heartbeat must extend its lease") + pipe.hset(control_key, "lease_expires_at_ms", lease.deadline_ms) + pipe.zadd( + self.keys.lease_expiry, + {RedisKeys.lease_member(run_id, lease.fence): lease.deadline_ms}, + ) + else: + self._apply_plan(pipe, current, plan) await pipe.execute() return plan except redis_watch_error(): diff --git a/src/hayhooks/durable/runtime.py b/src/hayhooks/durable/runtime.py index 7b6eec12..aae15dd4 100644 --- a/src/hayhooks/durable/runtime.py +++ b/src/hayhooks/durable/runtime.py @@ -251,7 +251,6 @@ async def resume( execution_id, owner_id=owner_id, enforce_owner=enforce_owner, - allow_revision_mismatch=True, ) if self.resume_type is not None and update is None: msg = f"Execution '{execution_id}' requires a resume request body" diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py index e5f65489..c76c17b1 100644 --- a/src/hayhooks/durable/store.py +++ b/src/hayhooks/durable/store.py @@ -24,6 +24,7 @@ InvalidExecutionTransitionError, PayloadKind, RecoverExpiredLease, + ReleaseClaim, RequestCancellation, Resume, ScheduleRetry, @@ -118,6 +119,7 @@ def __init__( self._record = record self.worker_id = worker_id self._heartbeat: asyncio.Task[None] | None = None + self._transition_lock = asyncio.Lock() self._finished = False self._lost = False self._lost_event = asyncio.Event() @@ -133,7 +135,8 @@ def lost_event(self) -> asyncio.Event: return self._lost_event async def __aenter__(self) -> ExecutionClaim: - await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) + async with self._transition_lock: + await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) self._heartbeat = asyncio.create_task( self._heartbeat_loop(), name=f"durable-heartbeat:{self.record.execution_id}", @@ -147,17 +150,20 @@ async def __aexit__(self, _exc_type: Any, _exc: Any, _traceback: Any) -> None: await self._heartbeat async def checkpoint(self) -> None: - self._ensure_owned() - await self._transition( - Checkpoint( - self.control.fence, - self.worker_id, - 0, - self.store.lease_duration_ms, - self.store._snapshot(self.record), - self._new_progress(), + async with self._transition_lock: + self._ensure_owned() + progress, progress_payloads = self._new_progress() + await self._transition( + Checkpoint( + self.control.fence, + self.worker_id, + 0, + self.store.lease_duration_ms, + self.store._snapshot(self.record), + progress_payloads, + ), + record_progress=progress, ) - ) async def cancellation_requested(self) -> bool: control = await self.store._core_call( @@ -175,68 +181,76 @@ async def cancellation_requested(self) -> bool: return False async def complete(self) -> None: - self._ensure_owned() - if self.record.status is ExecutionStatus.CANCELED and self.control.cancel_requested_at_ms is None: - cancellation = await self.store._core_call( - "request cancellation", - self.store.core.transition( - self.record.execution_id, - RequestCancellation(0, self.record.cancel_reason), - ), - ) - self._sync(cancellation.next_control) - if self.record.status is ExecutionStatus.FAILED: - error = self.record.error or ExecutionError(type="ExecutionError", message="Execution failed") + async with self._transition_lock: + self._ensure_owned() + if self.record.status is ExecutionStatus.CANCELED and self.control.cancel_requested_at_ms is None: + cancellation = await self.store._core_call( + "request cancellation", + self.store.core.transition( + self.record.execution_id, + RequestCancellation(0, self.record.cancel_reason), + ), + ) + self._sync(cancellation.next_control) + progress, progress_payloads = self._new_progress() + if self.record.status is ExecutionStatus.FAILED: + error = self.record.error or ExecutionError(type="ExecutionError", message="Execution failed") + await self._transition( + Fail( + self.control.fence, + self.worker_id, + 0, + _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="error"), + progress_payloads, + ), + record_progress=progress, + ) + else: + await self._transition( + Complete( + self.control.fence, + self.worker_id, + 0, + _encode(self.record.result, limit=self.store.config.max_result_bytes, label="result"), + progress_payloads, + ), + record_progress=progress, + ) + self._finished = True + + async def suspend(self) -> None: + async with self._transition_lock: + self._ensure_owned() + progress, progress_payloads = self._new_progress() await self._transition( - Fail( + Suspend( self.control.fence, self.worker_id, 0, - _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="error"), - self._new_progress(), - ) + self.store._snapshot(self.record), + _encode(self.record.wait, limit=self.store.config.max_wait_bytes, label="wait"), + progress_payloads, + ), + record_progress=progress, ) - else: + self._finished = True + + async def retry(self, error: ExecutionError, *, delay: float) -> None: + async with self._transition_lock: + self._ensure_owned() await self._transition( - Complete( + ScheduleRetry( self.control.fence, self.worker_id, 0, - _encode(self.record.result, limit=self.store.config.max_result_bytes, label="result"), - self._new_progress(), + max(0, round(delay * 1_000)), + self.store.max_application_retries, + _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="retry error"), ) ) - self._finished = True - - async def suspend(self) -> None: - self._ensure_owned() - await self._transition( - Suspend( - self.control.fence, - self.worker_id, - 0, - self.store._snapshot(self.record), - _encode(self.record.wait, limit=self.store.config.max_wait_bytes, label="wait"), - self._new_progress(), - ) - ) - self._finished = True - - async def retry(self, error: ExecutionError, *, delay: float) -> None: - self._ensure_owned() - await self._transition( - ScheduleRetry( - self.control.fence, - self.worker_id, - 0, - max(0, round(delay * 1_000)), - self.store.max_application_retries, - _encode(error.to_dict(), limit=self.store.config.max_error_bytes, label="retry error"), - ) - ) - self._finished = True + self._finished = True - async def _transition(self, command: Any) -> Any: + async def _transition(self, command: Any, *, record_progress: tuple[ExecutionProgressEvent, ...] = ()) -> Any: confirmed_at = time.monotonic() try: plan = await self.store._core_call( @@ -246,14 +260,24 @@ async def _transition(self, command: Any) -> Any: except ExecutionLeaseLostError: self._mark_lost() raise - self._sync(plan.next_control, confirmed_at=confirmed_at, progress_events=plan.progress_events) + self._sync( + plan.next_control, + confirmed_at=confirmed_at, + record_progress=record_progress, + progress_events=plan.progress_events, + ) return plan async def _heartbeat_loop(self) -> None: while not self._finished and not self._lost: await asyncio.sleep(self.store.heartbeat_interval) try: - await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) + async with self._transition_lock: + if self._finished or self._lost: + return + await self._transition( + Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms) + ) except ExecutionLeaseLostError: return except Exception: @@ -262,11 +286,14 @@ async def _heartbeat_loop(self) -> None: return continue - def _new_progress(self) -> tuple[bytes, ...]: + def _new_progress(self) -> tuple[tuple[ExecutionProgressEvent, ...], tuple[bytes, ...]]: events = self._new_progress_events() - return tuple( - _encode(event.to_dict(), limit=self.store.config.max_progress_event_bytes, label="progress") - for event in events + return ( + events, + tuple( + _encode(event.to_dict(), limit=self.store.config.max_progress_event_bytes, label="progress") + for event in events + ), ) def _new_progress_events(self) -> tuple[ExecutionProgressEvent, ...]: @@ -282,13 +309,13 @@ def _sync( control: Any, *, confirmed_at: float | None = None, + record_progress: tuple[ExecutionProgressEvent, ...] = (), progress_events: tuple[EngineProgressEvent, ...] = (), ) -> None: if progress_events: - pending = self._new_progress_events() - for event, persisted in zip(pending, progress_events, strict=True): + for event, persisted in zip(record_progress, progress_events, strict=True): event.sequence = persisted.sequence - self._last_persisted_progress = self.record.progress[-1] + self._last_persisted_progress = record_progress[-1] self.control = control self.record.attempt = control.run_attempt self.record.sequence = control.version @@ -397,8 +424,14 @@ async def claim_next(self, worker_name: str) -> ExecutionClaim | None: return None if plan.next_control.status is not EngineStatus.RUNNING: return None - view = await self._read_view(run_id) + try: + view = await self._read_view(run_id) + except BaseException: + with suppress(Exception): + await self._release_claim(plan.next_control, worker_name) + raise if view is None: + await self._release_claim(plan.next_control, worker_name) return None current, record = view if ( @@ -406,9 +439,17 @@ async def claim_next(self, worker_name: str) -> ExecutionClaim | None: or current.fence != plan.next_control.fence or current.lease_owner != worker_name ): + await self._release_claim(plan.next_control, worker_name) return None return ExecutionClaim(self, current, record, worker_name, confirmed_at) + async def _release_claim(self, control: ExecutionControl, worker_name: str) -> None: + with suppress(ExecutionLeaseLostError, ExecutionNotFoundError, InvalidExecutionTransitionError): + await self._core_call( + "release undelivered execution claim", + self.core.transition(control.run_id, ReleaseClaim(control.fence, worker_name)), + ) + async def request_cancel(self, execution_id: str, reason: str | None = None) -> bool: """Persist a cancellation request, returning whether it was accepted.""" control = await self._core_call("read execution for cancellation", self.core.get(execution_id)) @@ -458,6 +499,7 @@ async def resume(self, execution_id: str, update: JsonValue | None = None) -> bo self.definition_revision or control.definition_revision, self._snapshot(record), (_encode(event.to_dict(), limit=self.config.max_progress_event_bytes, label="progress"),), + expected_version=control.version, ), ), ) diff --git a/src/hayhooks/server/a2a/durable_executor.py b/src/hayhooks/server/a2a/durable_executor.py index 747ea238..843d8a55 100644 --- a/src/hayhooks/server/a2a/durable_executor.py +++ b/src/hayhooks/server/a2a/durable_executor.py @@ -9,7 +9,7 @@ from hayhooks.a2a import RecoverableTaskStore, default_a2a_owner from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStatus, ExecutionStoreError -from hayhooks.durable.runtime import DurableDeployment, execution_id_for +from hayhooks.durable.runtime import DefinitionRevisionConflictError, DurableDeployment, execution_id_for from hayhooks.server.a2a.imports import ( AgentExecutor, EventQueue, @@ -244,20 +244,30 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non allow_revision_mismatch=True, ) except KeyError: - await updater.failed( - message=updater.new_agent_message( - [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] + from a2a.types import TaskState + + if task.status.state != TaskState.TASK_STATE_SUBMITTED: + await updater.failed( + message=updater.new_agent_message( + [ + new_text_part( + "The durable Agent execution record is missing (durable_execution_missing)." + ) + ] + ) ) - ) - return + return if record is not None and record.status is ExecutionStatus.WAITING: action = "resume" - resumed = await self.deployment.resume( - execution_id, - {"messages": [message.to_dict() for message in build_haystack_resume_messages(context)]}, - owner_id=owner_id, - enforce_owner=True, - ) + try: + resumed = await self.deployment.resume( + execution_id, + {"messages": [message.to_dict() for message in build_haystack_resume_messages(context)]}, + owner_id=owner_id, + enforce_owner=True, + ) + except DefinitionRevisionConflictError as error: + raise InvalidParamsError(str(error)) from error if not resumed: msg = f"Task '{task.id}' is no longer accepting follow-up messages" raise InvalidParamsError(msg) diff --git a/src/hayhooks/server/a2a/executor.py b/src/hayhooks/server/a2a/executor.py index 7c736d9e..996dc523 100644 --- a/src/hayhooks/server/a2a/executor.py +++ b/src/hayhooks/server/a2a/executor.py @@ -77,7 +77,6 @@ async def emit(text: str, *, last: bool) -> None: return async for text in _iter_text_chunks(result): await emit(text, last=False) - await emit("", last=True) class ChatCompletionAgentExecutor(AgentExecutor): diff --git a/src/hayhooks/server/durable/routes.py b/src/hayhooks/server/durable/routes.py index 15fd18fc..4823a18b 100644 --- a/src/hayhooks/server/durable/routes.py +++ b/src/hayhooks/server/durable/routes.py @@ -23,7 +23,6 @@ ) from hayhooks.server.pipelines.registry import registry from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper -from hayhooks.settings import settings DURABLE_ROUTE_SUFFIXES = ( "/run-durable", @@ -61,8 +60,8 @@ def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResu ) -def _durable_owner(request: Request) -> tuple[str | None, bool]: - header = settings.durable_trusted_owner_header.strip() +def _durable_owner(request: Request, deployment: DurableDeployment) -> tuple[str | None, bool]: + header = deployment.app_settings.durable_trusted_owner_header.strip() if not header: return None, False owner = request.headers.get(header) @@ -125,7 +124,7 @@ async def submit( request: Request, idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), ) -> ExecutionResult: - owner_id, owner_scoped = _durable_owner(request) + owner_id, owner_scoped = _durable_owner(request, deployment) if idempotency_key is not None and _IDEMPOTENCY_KEY_PATTERN.fullmatch(idempotency_key) is None: raise HTTPException( status_code=422, @@ -172,7 +171,7 @@ async def submit( async def inspect_execution(execution_id: ExecutionId, request: Request) -> ExecutionResult: try: - owner_id, enforce_owner = _durable_owner(request) + owner_id, enforce_owner = _durable_owner(request, deployment) record = await get_execution(execution_id, owner_id, enforce_owner) return _execution_result(deployment, record) except KeyError as error: @@ -182,7 +181,7 @@ async def inspect_execution(execution_id: ExecutionId, request: Request) -> Exec async def cancel_execution(execution_id: ExecutionId, response: Response, request: Request) -> ExecutionResult: try: - owner_id, enforce_owner = _durable_owner(request) + owner_id, enforce_owner = _durable_owner(request, deployment) accepted = await deployment.request_cancel( execution_id, owner_id=owner_id, @@ -203,7 +202,7 @@ async def resume_execution( update: Any = Body(default=None), # noqa: B008 ) -> ExecutionResult: try: - owner_id, enforce_owner = _durable_owner(request) + owner_id, enforce_owner = _durable_owner(request, deployment) resumed = await deployment.resume( execution_id, update, @@ -215,6 +214,8 @@ async def resume_execution( record = await get_execution(execution_id, owner_id, enforce_owner) except KeyError as error: raise HTTPException(status_code=404, detail="Execution not found") from error + except DefinitionRevisionConflictError as error: + raise HTTPException(status_code=409, detail=str(error)) from error except (ValidationError, ValueError) as error: raise HTTPException(status_code=422, detail=str(error)) from error except ExecutionStoreError as error: diff --git a/src/hayhooks/server/pipelines/models.py b/src/hayhooks/server/pipelines/models.py index c973ae81..355215ae 100644 --- a/src/hayhooks/server/pipelines/models.py +++ b/src/hayhooks/server/pipelines/models.py @@ -11,15 +11,7 @@ def _create_schema_model(model_name: str, **fields: Any) -> type[BaseModel]: - """Create a dynamic API model that Pydantic can resolve during OpenAPI generation.""" - model = create_model(model_name, __module__=__name__, **fields) - - # FastAPI may rebuild a route's TypeAdapter long after the route is - # registered. Pydantic resolves the dynamically-created model by name in - # this module's namespace at that point, so keep it available there. - globals()[model_name] = model - model.model_rebuild() - return model + return create_model(model_name, **fields) def _resolved_annotations(func: Callable) -> dict[str, Any]: diff --git a/src/hayhooks/server/routers/status.py b/src/hayhooks/server/routers/status.py index 9a2a67e1..15196e25 100644 --- a/src/hayhooks/server/routers/status.py +++ b/src/hayhooks/server/routers/status.py @@ -35,8 +35,6 @@ class PipelineStatusResponse(BaseModel): async def status_all() -> StatusResponse: pipelines = registry.get_names() durable_health = await durable_runtime.health() - if not durable_health["healthy"]: - raise HTTPException(status_code=503, detail={"status": "Degraded", "durable": durable_health}) return StatusResponse(status="Up!", pipelines=pipelines, durable=durable_health) diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index dce0e889..6a1a1d03 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -4,10 +4,11 @@ import shutil import sys import tempfile +import threading import time import traceback from collections.abc import AsyncGenerator, Awaitable, Callable, Generator -from contextlib import nullcontext +from contextlib import asynccontextmanager, nullcontext from functools import wraps from pathlib import Path from typing import Any, cast @@ -54,8 +55,9 @@ from hayhooks.server.utils.yaml_pipeline_wrapper import YAMLPipelineWrapper from hayhooks.settings import DeployConcurrencyPolicy, settings -_deployment_serial_lock = asyncio.Lock() -_deployment_publication_lock = asyncio.Lock() +# These APIs can be called from multiple event loops, so the locks must be process-wide. +_deployment_serial_lock = threading.Lock() +_deployment_publication_lock = threading.Lock() _deployments_in_progress: set[str] = set() @@ -64,6 +66,15 @@ async def _offload(func: Callable, **kwargs: Any) -> Any: return await asyncio.to_thread(func, **kwargs) +@asynccontextmanager +async def _deployment_lock(lock: threading.Lock): + await run_in_threadpool(lock.acquire) + try: + yield + finally: + lock.release() + + class _DeploymentSnapshot: """Rollback state captured before preparation mutates files or loaded modules.""" @@ -154,7 +165,9 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 ) dlog.debug("Starting pipeline deployment transaction") policy_lock = ( - _deployment_serial_lock if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED else nullcontext() + _deployment_lock(_deployment_serial_lock) + if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED + else nullcontext() ) snapshot: _DeploymentSnapshot | None = None candidate: DurableDeployment | None = None @@ -163,7 +176,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 publication_started = False async with policy_lock: try: - async with _deployment_publication_lock: + async with _deployment_lock(_deployment_publication_lock): if pipeline_name in _deployments_in_progress: msg = f"Pipeline '{pipeline_name}' is already being deployed" raise PipelineAlreadyExistsError(msg) @@ -198,7 +211,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 await candidate.prepare() dlog.bind(durable=candidate is not None).debug("Prepared pipeline deployment candidate") - async with _deployment_publication_lock: + async with _deployment_lock(_deployment_publication_lock): snapshot.refresh_publication() publication_started = True result = commit_prepared_pipeline( @@ -216,7 +229,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 except BaseException as error: dlog.bind(error_type=type(error).__name__).debug("Pipeline deployment transaction failed") if snapshot is not None and (registered or old_quiesced): - async with _deployment_publication_lock: + async with _deployment_lock(_deployment_publication_lock): try: if registered: if candidate is not None: @@ -231,7 +244,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 raise finally: if registered: - async with _deployment_publication_lock: + async with _deployment_lock(_deployment_publication_lock): _deployments_in_progress.discard(pipeline_name) if snapshot is not None: snapshot.cleanup() @@ -295,9 +308,11 @@ async def undeploy_pipeline_async( ) -> None: """Atomically unpublish a pipeline before stopping its owned resources.""" policy_lock = ( - _deployment_serial_lock if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED else nullcontext() + _deployment_lock(_deployment_serial_lock) + if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED + else nullcontext() ) - async with policy_lock, _deployment_publication_lock: + async with policy_lock, _deployment_lock(_deployment_publication_lock): if pipeline_name in _deployments_in_progress: raise HTTPException(status_code=409, detail=f"Pipeline '{pipeline_name}' is being deployed") if registry.get(pipeline_name) is None: @@ -320,7 +335,7 @@ async def undeploy_pipeline_async( await deployment.start() raise - undeploy_pipeline(pipeline_name=pipeline_name, app=app) + await _offload(undeploy_pipeline, pipeline_name=pipeline_name, app=app) durable_runtime.install_deployment(pipeline_name, None) diff --git a/tests/test_a2a.py b/tests/test_a2a.py index 7a5bbdaa..90002026 100644 --- a/tests/test_a2a.py +++ b/tests/test_a2a.py @@ -301,12 +301,12 @@ async def test_execute_agent_task_streaming_result(): artifact_events = get_artifact_events(queue.events) # PipelineEvent items are skipped, text chunks are streamed incrementally - assert len(artifact_events) == 4 + assert len(artifact_events) == 3 texts = [event.artifact.parts[0].text for event in artifact_events] - assert texts == ["Hello, ", "world", " (question: hi)", ""] - # All chunks belong to the same artifact; an empty marker finalizes iterator output + assert texts == ["Hello, ", "world", " (question: hi)"] + # All chunks belong to the same artifact; terminal task status closes iterator output. assert len({event.artifact.artifact_id for event in artifact_events}) == 1 - assert [event.last_chunk for event in artifact_events] == [False, False, False, True] + assert [event.last_chunk for event in artifact_events] == [False, False, False] assert artifact_events[0].append is False assert artifact_events[1].append is True diff --git a/tests/test_deploy_performance.py b/tests/test_deploy_performance.py index 722c1272..f71fad40 100644 --- a/tests/test_deploy_performance.py +++ b/tests/test_deploy_performance.py @@ -10,6 +10,8 @@ from hayhooks.server.app import create_app from hayhooks.server.pipelines.registry import registry from hayhooks.server.utils.deploy_utils import ( + _deployment_publication_lock, + _deployment_serial_lock, commit_prepared_pipeline, deploy_pipeline_files_async, deploy_pipeline_yaml, @@ -101,8 +103,26 @@ async def test_undeploy_pipeline_async(monkeypatch): ) assert registry.get("undeploy_async_test") is not None + caller_thread = threading.get_ident() + undeploy_threads = [] + from hayhooks.server.utils import deploy_utils + + original_undeploy = deploy_utils.undeploy_pipeline + + def tracked_undeploy(*args, **kwargs): + undeploy_threads.append(threading.get_ident()) + return original_undeploy(*args, **kwargs) + + monkeypatch.setattr(deploy_utils, "undeploy_pipeline", tracked_undeploy) await undeploy_pipeline_async(pipeline_name="undeploy_async_test") assert registry.get("undeploy_async_test") is None + assert undeploy_threads and undeploy_threads[0] != caller_thread + + +def test_deploy_locks_are_safe_across_event_loops(): + lock_type = type(threading.Lock()) + assert isinstance(_deployment_serial_lock, lock_type) + assert isinstance(_deployment_publication_lock, lock_type) @pytest.mark.asyncio diff --git a/tests/test_deploy_utils.py b/tests/test_deploy_utils.py index c5401337..132a1754 100644 --- a/tests/test_deploy_utils.py +++ b/tests/test_deploy_utils.py @@ -13,6 +13,7 @@ from haystack import Pipeline from hayhooks.server.exceptions import PipelineFilesError, PipelineModuleLoadError, PipelineWrapperError +from hayhooks.server.pipelines import models as pipeline_models from hayhooks.server.pipelines import registry from hayhooks.server.pipelines.sse import SSEStream from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper @@ -367,6 +368,16 @@ def sample_func(name: str, age: int = 25, optional: str = ""): assert "optional" not in schema.get("required", []) +def test_dynamic_schema_models_are_not_retained_in_module_globals(): + def sample_func(value: str): + pass + + model = create_request_model_from_callable(sample_func, "Ephemeral", docstring_parser.parse("")) + + assert model.__name__ == "EphemeralRequest" + assert "EphemeralRequest" not in vars(pipeline_models) + + def test_create_request_model_no_docstring(): def sample_func_no_doc(name: str, age: int = 30): pass diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py index f3cbfcf1..cbea8d51 100644 --- a/tests/test_durable_a2a.py +++ b/tests/test_durable_a2a.py @@ -300,6 +300,33 @@ async def test_missing_execution_preserves_a_task_awaiting_submission(http_store assert projected.status.state == TaskState.TASK_STATE_SUBMITTED +async def test_submitted_task_with_a_missing_execution_is_resubmitted(http_store) -> None: + task = _recoverable_task() + deployment = _Deployment() + + async def missing_until_submitted(execution_id, **kwargs): + if deployment.execution_id is None: + raise KeyError(execution_id) + return await _Deployment.get(deployment, execution_id, **kwargs) + + class Queue: + def __init__(self): + self.events = [] + + async def enqueue_event(self, event): + self.events.append(event) + + deployment.get = missing_until_submitted + executor = DurableAgentExecutor("agent", http_store, deployment) + executor.task_store.owner_id_for_context = lambda _context: "owner" + context = SimpleNamespace(current_task=task, message=None, call_context=SimpleNamespace()) + + await executor.execute(context, Queue()) + + assert deployment.execution_id == execution_id_for("owner", task.id) + assert deployment.submitted_payload["messages"][0]["content"] == [{"text": "recover me"}] + + @pytest.mark.parametrize(("configured", "expected"), [(0.05, 0.1), (5.0, 5.0)]) async def test_durable_a2a_polling_honors_its_configured_floor(monkeypatch, http_store, configured, expected) -> None: executor = DurableAgentExecutor("agent", http_store, _Deployment(status=ExecutionStatus.RUNNING)) diff --git a/tests/test_durable_execution.py b/tests/test_durable_execution.py index d42112b7..110f765f 100644 --- a/tests/test_durable_execution.py +++ b/tests/test_durable_execution.py @@ -46,7 +46,7 @@ load_pipeline_module, unload_pipeline_modules, ) -from hayhooks.settings import settings +from hayhooks.settings import AppSettings, settings pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -613,15 +613,18 @@ def test_durable_waiting_resume_is_typed_private_and_revision_safe(monkeypatch) } deployment = durable_runtime.current_deployment("approval") assert deployment is not None - revision = deployment.revision - deployment.revision = "replacement" missing = client.post(f"{url}/resume") assert missing.status_code == 422 invalid = client.post(f"{url}/resume", json={"approved": "not-a-bool"}) assert invalid.status_code == 422 + revision = deployment.revision + deployment.revision = "replacement" + conflict = client.post(f"{url}/resume", json={"approved": True}) + assert conflict.status_code == 409 + assert client.get(url).json()["status"] == "waiting" + deployment.revision = revision resumed = client.post(f"{url}/resume", json={"approved": True}) assert resumed.status_code == 202 - deployment.revision = revision completed = _wait_for_status(client, url, "completed", "resumed execution did not complete") assert completed.json()["result"] == {"value": 7} @@ -656,6 +659,28 @@ def test_durable_rest_enforces_configured_trusted_owner_header(monkeypatch) -> N assert "exceeds 512 characters" in oversized.json()["detail"] +def test_durable_rest_uses_the_deployments_owner_header_setting(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_trusted_owner_header", "") + app_settings = AppSettings(durable_store="memory", durable_trusted_owner_header="X-Embedded-Owner") + provider = InMemoryExecutionStoreProvider(app_settings=app_settings) + wrapper = Wrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + deployment = DurableDeployment("embedded", wrapper, provider, app_settings=app_settings) + registry.add("embedded", wrapper) + app = create_app() + add_pipeline_api_route(app, "embedded", wrapper, _durable_deployment=deployment) + + with TestClient(app) as client: + assert client.get("/embedded/executions/missing").status_code == 401 + authenticated = client.get( + "/embedded/executions/missing", + headers={"X-Embedded-Owner": "alice"}, + ) + + assert authenticated.status_code == 404 + + def test_durable_deployment_requires_an_explicit_revision() -> None: class MissingRevisionWrapper(BasePipelineWrapper): def setup(self) -> None: diff --git a/tests/test_durable_redis_codec.py b/tests/test_durable_redis_codec.py index 6aebb382..2143aa5f 100644 --- a/tests/test_durable_redis_codec.py +++ b/tests/test_durable_redis_codec.py @@ -2,12 +2,12 @@ from __future__ import annotations -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock import pytest from hayhooks.durable.backend import MAINTENANCE_BATCH_SIZE, ExecutionStoreCorruptionError -from hayhooks.durable.engine import MAX_CONTROL_SCALAR_BYTES, initial_control +from hayhooks.durable.engine import MAX_CONTROL_SCALAR_BYTES, Claim, Heartbeat, decide, initial_control from hayhooks.durable.redis import RedisExecutionStore, RedisKeys, decode_control, digest, encode_control @@ -111,3 +111,23 @@ async def test_maintenance_reads_a_fixed_batch_of_due_leases() -> None: num=MAINTENANCE_BATCH_SIZE, withscores=True, ) + + +async def test_heartbeat_writes_only_lease_fields() -> None: + current = decide(control(), Claim("worker", 100, 10_000, 3, "rev-1")).next_control + pipe = MagicMock() + pipe.__aenter__ = AsyncMock(return_value=pipe) + pipe.__aexit__ = AsyncMock(return_value=None) + pipe.watch = AsyncMock() + pipe.hgetall = AsyncMock(return_value=encode_control(current)) + pipe.time = AsyncMock(return_value=(1, 0)) + pipe.execute = AsyncMock(return_value=[]) + redis = Mock() + redis.pipeline.return_value = pipe + store = RedisExecutionStore(redis, deployment="deployment") + + await store.transition("run-1", Heartbeat(1, "worker", 0, 2_000)) + + pipe.hset.assert_called_once_with(store.keys.control("run-1"), "lease_expires_at_ms", 3_000) + pipe.zadd.assert_called_once_with(store.keys.lease_expiry, {RedisKeys.lease_member("run-1", 1): 3_000}) + pipe.zrem.assert_not_called() diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py index 278fda65..0ebd76c3 100644 --- a/tests/test_durable_store.py +++ b/tests/test_durable_store.py @@ -12,7 +12,7 @@ from hayhooks.durable.backend import ExecutionStoreConfig from hayhooks.durable.context import RESUME_INPUT_KEY -from hayhooks.durable.engine import Claim, Heartbeat, RequestCancellation, Resume +from hayhooks.durable.engine import Checkpoint, Claim, Heartbeat, RequestCancellation, Resume from hayhooks.durable.manager import DurableExecutionManager from hayhooks.durable.models import ( ExecutionCheckpoint, @@ -148,7 +148,7 @@ async def test_runtime_provider_cannot_be_replaced(monkeypatch: pytest.MonkeyPat runtime = DurableRuntime(provider) with pytest.raises(AttributeError): - setattr(runtime, "provider", InMemoryExecutionStoreProvider()) + setattr(runtime, "provider", InMemoryExecutionStoreProvider()) # noqa: B010 - immutable property check assert runtime.provider is provider await runtime.close() @@ -257,6 +257,42 @@ async def test_checkpoint_keeps_progress_added_after_a_concurrent_cancellation() ] +async def test_concurrent_checkpoints_persist_each_progress_event_once(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + persisted = asyncio.Event() + release = asyncio.Event() + original_transition = core.transition + first = True + + async def gated_transition(run_id, command, *, candidate=False): + nonlocal first + plan = await original_transition(run_id, command, candidate=candidate) + if isinstance(command, Checkpoint) and first: + first = False + persisted.set() + await release.wait() + return plan + + monkeypatch.setattr(core, "transition", gated_transition) + claim.record.append_progress("first") + first_checkpoint = asyncio.create_task(claim.checkpoint()) + await persisted.wait() + claim.record.append_progress("second") + claim.record.append_progress("third") + second_checkpoint = asyncio.create_task(claim.checkpoint()) + await asyncio.sleep(0) + release.set() + await asyncio.gather(first_checkpoint, second_checkpoint) + + record = await store.get("run_1") + assert record is not None + assert [(event.sequence, event.message) for event in record.progress] == [(2, "second"), (3, "third")] + + async def test_losing_resume_race_returns_false(monkeypatch: pytest.MonkeyPatch) -> None: core = InMemoryExecutionStore(deployment="deployment", config=_config()) store = ExecutionStore(core, definition_revision="rev-1") @@ -289,6 +325,111 @@ async def gated_transition(run_id, command, *, candidate=False): assert record is not None and record.status is ExecutionStatus.CANCELED +async def test_delayed_resume_cannot_overwrite_a_newer_waiting_checkpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + claim.record.status = ExecutionStatus.WAITING + claim.record.wait = {"kind": "approval"} + await claim.suspend() + + entered = asyncio.Event() + release = asyncio.Event() + original_transition = core.transition + first = True + + async def gated_transition(run_id, command, *, candidate=False): + nonlocal first + if isinstance(command, Resume) and first: + first = False + entered.set() + await release.wait() + return await original_transition(run_id, command, candidate=candidate) + + monkeypatch.setattr(core, "transition", gated_transition) + stale_resume = asyncio.create_task(store.resume("run_1", {"generation": 1})) + await entered.wait() + assert await store.resume("run_1", {"generation": 2}) + newer = await store.claim_next("worker") + assert newer is not None + newer.record.application_state["generation"] = 2 + newer.record.checkpoint = ExecutionCheckpoint(ExecutionKind.PIPELINE, {"generation": 2}) + newer.record.status = ExecutionStatus.WAITING + newer.record.wait = {"kind": "approval"} + await newer.suspend() + release.set() + + assert not await stale_resume + record = await store.get("run_1") + assert record is not None + assert record.status is ExecutionStatus.WAITING + assert record.checkpoint == ExecutionCheckpoint(ExecutionKind.PIPELINE, {"generation": 2}) + assert record.application_state[RESUME_INPUT_KEY] == {"generation": 2} + + +async def test_revision_mismatched_resume_leaves_the_execution_waiting() -> None: + store = _store() + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + claim.record.status = ExecutionStatus.WAITING + claim.record.wait = {"kind": "approval"} + await claim.suspend() + + store.set_definition_revision("rev-2") + assert not await store.resume("run_1", {"approved": True}) + + record = await store.get("run_1") + assert record is not None + assert record.status is ExecutionStatus.WAITING + assert record.error is None + + +async def test_failed_post_claim_read_releases_the_attempt(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + + async def failed_read(_execution_id): + msg = "read failed" + raise RuntimeError(msg) + + monkeypatch.setattr(store, "_read_view", failed_read) + with pytest.raises(RuntimeError, match="read failed"): + await store.claim_next("worker") + + control = await core.get("run_1") + assert control is not None + assert control.status.value == "queued" + assert control.run_attempt == 0 + assert control.lease_owner is None + assert (await core.operational_counts())["lease_expiry"] == 0 + + +async def test_worker_finishes_a_claim_acquired_during_deactivation(monkeypatch: pytest.MonkeyPatch) -> None: + store = _store() + manager = DurableExecutionManager("deployment", store, AsyncMock(), adapter=object()) + claim = SimpleNamespace(record=SimpleNamespace(execution_id="run_1", attempt=1, status=ExecutionStatus.COMPLETED)) + + async def claim_while_deactivating(_worker_name): + manager._accepting_claims = False + return claim + + process_claim = AsyncMock() + monkeypatch.setattr(store, "claim_next", claim_while_deactivating) + monkeypatch.setattr(manager, "_process_claim", process_claim) + manager._accepting_claims = True + manager._worker_generation = 1 + + await manager._worker("worker", 1) + + process_claim.assert_awaited_once_with(claim) + + async def test_retry_exhaustion_persists_its_progress_event() -> None: store = _store(max_run_attempts=1) diff --git a/tests/test_it_status.py b/tests/test_it_status.py index 53c9f622..89abf8b1 100644 --- a/tests/test_it_status.py +++ b/tests/test_it_status.py @@ -1,8 +1,11 @@ from pathlib import Path +from unittest.mock import AsyncMock import pytest +from hayhooks.durable.runtime import durable_runtime from hayhooks.server.pipelines import registry +from hayhooks.server.routers.status import status_all @pytest.fixture(autouse=True) @@ -39,3 +42,13 @@ def test_status_no_pipelines(client, status_pipeline): assert status_response.status_code == 200 assert "pipelines" in status_response.json() assert len(status_response.json()["pipelines"]) == 0 + + +async def test_global_status_remains_a_liveness_probe_when_durable_is_unhealthy(monkeypatch): + health = {"healthy": False, "deployments": {"job": {"healthy": False}}} + monkeypatch.setattr(durable_runtime, "health", AsyncMock(return_value=health)) + + response = await status_all() + + assert response.status == "Up!" + assert response.durable == health diff --git a/tests/test_redis_execution_integration.py b/tests/test_redis_execution_integration.py index d6bf027e..9975e6a8 100644 --- a/tests/test_redis_execution_integration.py +++ b/tests/test_redis_execution_integration.py @@ -14,6 +14,7 @@ Complete, ExecutionLeaseLostError, ExecutionPayloadSizeError, + Heartbeat, RecoverExpiredLease, RequestCancellation, ScheduleRetry, @@ -102,6 +103,23 @@ async def test_delayed_work_is_invisible_until_the_redis_deadline(store) -> None assert await durable.read_candidate() == run_id +async def test_heartbeat_updates_only_the_lease_path(store, monkeypatch) -> None: + _, durable = store + await durable.submit(control(), b"{}", binding_digest="b" * 64) + run_id, claimed = await _claim(durable) + before = claimed.next_control.lease_expires_at_ms + monkeypatch.setattr( + durable, + "_apply_plan", + lambda *_args, **_kwargs: pytest.fail("heartbeat rewrote the full execution control"), + ) + + heartbeat = await durable.transition(run_id, Heartbeat(1, "worker", 0, 2_000)) + + assert heartbeat.next_control.lease_expires_at_ms is not None + assert heartbeat.next_control.lease_expires_at_ms > before + + async def test_hundred_concurrent_submissions_succeed_when_admission_is_disabled(store) -> None: _, durable = store From 6b95d31b208bd1084659b85742c455e0c532ceb9 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Mon, 17 Aug 2026 13:19:56 +0200 Subject: [PATCH 15/28] feat(durable): add portable FastAPI integration --- docs/advanced/durable-engine.md | 144 +-- .../durable-fastapi-integration-plan.md | 876 ++++++++++++++++++ docs/reference/api-reference.md | 5 + src/hayhooks/cli/a2a.py | 7 +- src/hayhooks/cli/mcp.py | 14 +- src/hayhooks/durable/__init__.py | 89 +- src/hayhooks/durable/context.py | 5 + src/hayhooks/durable/fastapi.py | 246 +++++ src/hayhooks/durable/runtime.py | 89 +- src/hayhooks/durable/settings.py | 51 + src/hayhooks/durable/store.py | 72 +- src/hayhooks/server/a2a/app.py | 44 +- src/hayhooks/server/a2a/executor.py | 6 +- src/hayhooks/server/a2a/runtime.py | 12 +- src/hayhooks/server/app.py | 9 +- src/hayhooks/server/durable/routes.py | 277 +----- src/hayhooks/server/routers/status.py | 11 +- src/hayhooks/server/utils/deploy_utils.py | 123 ++- src/hayhooks/server/utils/mcp_utils.py | 60 +- tests/test_cli.py | 15 +- tests/test_durable_a2a.py | 7 +- tests/test_durable_deployment_lifecycle.py | 9 +- tests/test_durable_execution.py | 128 ++- tests/test_durable_fastapi.py | 193 ++++ tests/test_durable_store.py | 32 +- tests/test_it_status.py | 13 +- 26 files changed, 2033 insertions(+), 504 deletions(-) create mode 100644 docs/advanced/durable-fastapi-integration-plan.md create mode 100644 src/hayhooks/durable/fastapi.py create mode 100644 src/hayhooks/durable/settings.py create mode 100644 tests/test_durable_fastapi.py diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 7bfd4804..2ee8123b 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -45,67 +45,60 @@ queued ── claim ──> running ── complete/fail/cancel ──> terminal ## Embedding the runtime -Applications can import `DurableRuntime`, `ExecutionStore`, -`ExecutionStoreProvider`, `InMemoryExecutionStoreProvider`, and -`RedisExecutionStoreProvider` directly from `hayhooks.durable`. A standalone -runtime starts only deployments attached to that runtime; it does not inspect -Hayhooks' process-global pipeline registry. +Applications can import the runtime, providers, deployment contracts, public +exceptions, and FastAPI adapter directly from `hayhooks.durable`. A standalone +runtime starts only deployments attached to that runtime; it never inspects the +Hayhooks pipeline registry. -This complete `app.py` embeds an in-memory durable worker in FastAPI. Its tool -simulates an eight-second upstream call so detached execution is easy to see: +This complete `app.py` adds an authenticated durable API to an existing FastAPI +application. Authentication middleware is expected to set a stable principal +on `request.state` before the owner dependency runs: ```python -import asyncio from contextlib import asynccontextmanager -from typing import Annotated -from fastapi import FastAPI, HTTPException, status -from haystack.components.agents import Agent -from haystack.components.generators.chat import OpenAIChatGenerator -from haystack.dataclasses import ChatMessage -from haystack.tools import tool +from fastapi import FastAPI, Request +from haystack import Pipeline from pydantic import BaseModel from hayhooks import BasePipelineWrapper, DurableContext -from hayhooks.durable import DurableRuntime, ExecutionResult, InMemoryExecutionStoreProvider -from hayhooks.settings import AppSettings +from hayhooks.durable import DurableRuntime, DurableSettings, create_durable_router -class AgentRequest(BaseModel): - question: str +class JobRequest(BaseModel): + document_id: str -@tool -async def check_order(order_id: Annotated[str, "The customer's order ID"]) -> str: - """Return the current shipping status for an order.""" - # Intentional demo delay: replace it with a real upstream API call. - await asyncio.sleep(8) - # Read-only tools are replay-safe; make mutating tools idempotent. - return f"Order {order_id} shipped and arrives Friday." +class JobResult(BaseModel): + indexed: bool -class SupportAgentWrapper(BasePipelineWrapper): - # Bump this when checkpoint-relevant code or prompts change. - durable_revision = "support-agent-v1" +class JobWrapper(BasePipelineWrapper): + durable_revision = "job-v1" def setup(self) -> None: - self.pipeline = Agent( - chat_generator=OpenAIChatGenerator(), - system_prompt="Help customers with their orders. Use the order tool when needed.", - tools=[check_order], - ) + self.pipeline = Pipeline() - async def run_durable_async(self, context: DurableContext, request: AgentRequest) -> dict: - return await context.run_agent_async(messages=[ChatMessage.from_user(request.question)]) + async def run_durable_async(self, context: DurableContext, request: JobRequest) -> JobResult: + # Replace this example body with checkpointed Pipeline work. + return JobResult(indexed=bool(request.document_id)) -durable_settings = AppSettings(durable_store="memory", durable_poll_interval=0.05) -provider = InMemoryExecutionStoreProvider(app_settings=durable_settings) -runtime = DurableRuntime(provider) +runtime = DurableRuntime( + durable_settings=DurableSettings( + durable_store="memory", # Development only; use Redis in production. + durable_poll_interval=0.05, + ) +) -wrapper = SupportAgentWrapper() +wrapper = JobWrapper() wrapper.setup() -deployment = runtime.deployment("support-agent", wrapper) +deployment = runtime.deployment("jobs", wrapper) + + +def current_owner_id(request: Request) -> str: + principal = request.state.principal + return f"{principal.tenant_id}:{principal.subject_id}" @asynccontextmanager @@ -118,46 +111,57 @@ async def lifespan(_app: FastAPI): app = FastAPI(lifespan=lifespan) +app.include_router( + create_durable_router(deployment, owner_id_dependency=current_owner_id), + prefix="/jobs", +) +``` +The adapter exposes typed submit, inspect, cancel, and resume routes. It does +not start workers or own the runtime. Host middleware and dependencies retain +control of authentication and authorization; the durable layer persists only +the stable owner ID returned by the dependency. Durable wrapper code can read +that value through `context.owner_id` after process recovery. -@app.post("/agent-runs", response_model=ExecutionResult, status_code=status.HTTP_202_ACCEPTED) -async def submit_agent_run(request: AgentRequest) -> ExecutionResult: - _, record = await deployment.submit(request.model_dump(mode="json")) - return ExecutionResult.model_validate(record.safe_view(links={"self": f"/agent-runs/{record.execution_id}"})) - +Passing `owner_id_dependency=None` is an explicit unscoped security choice: -@app.get("/agent-runs/{execution_id}", response_model=ExecutionResult) -async def get_agent_run(execution_id: str) -> ExecutionResult: - try: - record = await deployment.get(execution_id) - except KeyError as error: - raise HTTPException(status_code=404, detail="Agent run not found") from error - return ExecutionResult.model_validate(record.safe_view(links={"self": f"/agent-runs/{record.execution_id}"})) +```python +app.include_router( + create_durable_router(deployment, owner_id_dependency=None), + prefix="/internal-jobs", +) ``` -Run it and submit work: +In this mode, possession of the unguessable execution ID grants access. Use it +only behind one application-wide authorization boundary or for local +development. -```bash -pip install "hayhooks[durable]" -export OPENAI_API_KEY="your-api-key" -uvicorn app:app +For production Redis, give each application/environment an isolated prefix: -execution_id="$( - curl --fail --silent -X POST http://127.0.0.1:8000/agent-runs \ - -H 'content-type: application/json' \ - -d '{"question":"Where is order A-123?"}' | jq -r '.execution_id' -)" +```python +from hayhooks.durable import DurableRuntime, RedisExecutionStoreProvider -# The first poll should show `queued` or `running`; repeat until `completed`. -curl --fail --silent "http://127.0.0.1:8000/agent-runs/${execution_id}" | jq +provider = RedisExecutionStoreProvider( + redis_url="redis://localhost:6379/0", + key_prefix="myapp:production:durable", +) +runtime = DurableRuntime(provider) ``` -The runtime owns provider shutdown. Built-in providers snapshot their settings, -and the runtime adopts that snapshot when a provider is supplied. Pass custom -settings once—either to a built-in provider as above, or to `DurableRuntime` -when it selects the default provider. Conflicting runtime and provider settings -are rejected before a deployment is created. The selected provider is fixed for -the runtime's lifetime; create a new runtime to change storage backends. +When the host owns an existing binary Redis client, pass `redis=client` and +`close_redis=False`, then close the durable runtime before closing the client. +Do not use `decode_responses=True`. + +Every Uvicorn worker owns a runtime, Redis pool, and worker tasks. Redis leases +and fences coordinate them, so effective concurrency is `processes × +durable_execution_concurrency` per deployment. Keep wrapper revisions identical +across replicas and start with one to three processes and conservative +concurrency. + +Built-in providers snapshot `DurableSettings`, and the runtime adopts that +snapshot when a provider is supplied. Conflicting runtime and provider settings +are rejected before deployment creation. The selected provider is fixed for the +runtime's lifetime; create a new runtime to change storage backends. ## Redis layout diff --git a/docs/advanced/durable-fastapi-integration-plan.md b/docs/advanced/durable-fastapi-integration-plan.md new file mode 100644 index 00000000..c7bb5c7e --- /dev/null +++ b/docs/advanced/durable-fastapi-integration-plan.md @@ -0,0 +1,876 @@ +# Portable FastAPI durable integration plan + +**Status:** Proposed + +**Scope:** Hayhooks durable REST integration, authentication composition, and runtime ownership + +**Primary objective:** Make Hayhooks use the same public durable integration that an independent FastAPI application uses + +## Summary + +Hayhooks already has a durable execution engine with Redis-backed recovery, +fenced leases, idempotent submission, retries, progress, cancellation, and +wait/resume. The engine is usable from a standalone `DurableRuntime`, but the +current REST integration is implemented inside the Hayhooks server and reaches +into process-global runtime and registry state. + +This plan introduces one public FastAPI adapter and makes Hayhooks consume it: + +1. Applications own a `DurableRuntime` and start and close it in their lifespan. +2. `create_durable_router()` returns a standard FastAPI `APIRouter` for one + `DurableDeployment`. +3. An ordinary FastAPI dependency supplies a stable owner ID when authenticated + ownership isolation is required. +4. Hayhooks keeps only a small internal mount/unmount shim for its dynamic + pipeline deployment lifecycle. +5. Hayhooks REST, A2A, and MCP applications move from the process-global runtime + to application-owned runtime instances. + +The durable engine remains unaware of JWTs, cookies, sessions, middleware, +FastAPI application state, and the Hayhooks registry. Redis remains the source +of truth for execution state. + +## Goals + +- Provide a simple, documented public integration for an existing FastAPI + application. +- Compose with authentication middleware and existing FastAPI dependencies. +- Preserve typed request, resume, and result schemas in OpenAPI. +- Preserve all current durable REST paths and response semantics. +- Make ownership enforcement consistent for submit, inspect, cancel, and + resume. +- Make Hayhooks dogfood the public router rather than maintaining separate REST + handlers. +- Make each application instance own its runtime, workers, and provider + lifecycle. +- Preserve dynamic deploy, overwrite, rollback, and undeploy behavior in the + Hayhooks server. +- Keep Redis keys and persisted execution records unchanged. +- Maintain the current controlled-beta multi-process execution model. + +## Non-goals + +- A generic non-Haystack job framework. +- Authentication or authorization implemented by the durable engine. +- Persisting access tokens, sessions, or complete principal objects. +- A custom FastAPI middleware supplied by Hayhooks. +- A mounted durable sub-application. +- A `DurableRouter` subclass or a stateful integration/service container. +- Per-operation authorization hooks in the first public API. +- Separating API-only and worker-only processes in this change. +- Changing the Redis schema, execution state machine, or delivery semantics. +- Replacing Hayhooks' dynamic pipeline deployment transaction. + +## Current state + +### Durable runtime + +`DurableRuntime` owns deployment managers and a shared execution-store provider. +`DurableDeployment` owns the typed wrapper contract, Haystack adapter, store, +and worker manager. Redis or the in-memory reference backend owns persisted +execution data. + +The current standalone embedding path is valid, but it requires applications to +write their own HTTP handlers. The public facade also omits several types that +an embedding application naturally needs. + +### REST transport + +The current durable REST handlers live in +`hayhooks.server.durable.routes`. That module currently combines four concerns: + +- reusable durable HTTP behavior; +- Hayhooks trusted-header owner resolution; +- process-global runtime discovery; +- Hayhooks registry metadata and dynamic route mutation. + +Only the first concern belongs in a portable FastAPI adapter. + +### Runtime ownership + +Hayhooks REST, A2A, MCP, status, and deployment code currently import the same +module-level `durable_runtime`. This makes multiple application instances in a +single process share deployments, provider ownership, and shutdown behavior. +An independent FastAPI application would instead create and own its runtime. + +## Architectural decision + +### Chosen design + +Use a public `APIRouter` factory and a FastAPI owner-ID dependency: + +```python +from collections.abc import Awaitable, Callable + + +def create_durable_router( + deployment: DurableDeployment, + *, + owner_id_dependency: Callable[..., str | Awaitable[str]] | None, +) -> APIRouter: + ... +``` + +The owner dependency is a required keyword. Passing `None` is an explicit +choice to use unscoped bearer-by-execution-ID access. Supplying a dependency +enables owner isolation. + +The function returns routes only. It does not start workers, own the runtime, +mutate a registry, or alter an application. The caller uses ordinary FastAPI +composition: + +```python +app.include_router( + create_durable_router( + deployment, + owner_id_dependency=current_owner_id, + ), + prefix="/jobs", +) +``` + +### Why `APIRouter` + +- It is FastAPI's native unit of route composition. +- Host middleware automatically wraps included routes. +- Dependencies compose with existing authentication and authorization. +- Typed request and response models remain visible in OpenAPI. +- Applications control prefixes, tags, and router-level dependencies. +- Hayhooks can include the same router dynamically and retain its existing + route replacement logic. +- The adapter has no lifecycle or global state of its own. + +### Why a dependency rather than a callback + +A callback invoked manually by the route handler would need to reproduce part +of FastAPI's dependency system. A dependency already supports: + +- middleware-populated `request.state`; +- nested `Depends(...)` authentication dependencies; +- OAuth/OpenAPI security dependencies; +- async and sync implementations; +- application-specific `HTTPException` responses; +- dependency overrides in tests. + +The durable adapter needs only the resulting stable owner ID. It does not need +to know how authentication was performed. + +### Request and execution flow + +```mermaid +flowchart LR + request["HTTP request"] --> middleware["Host auth middleware"] + middleware --> permission["Host authorization dependencies"] + permission --> owner["Owner-ID dependency"] + owner --> router["Hayhooks durable router"] + router --> deployment["DurableDeployment"] + deployment <--> redis["Redis execution store"] + deployment --> worker["Process-local durable worker"] + worker --> wrapper["PipelineWrapper with DurableContext"] +``` + +Middleware runs before FastAPI dependency resolution. The host application +therefore retains control of authentication, request context, and broad API +authorization. The owner dependency reduces that context to the stable string +needed for record isolation. The router validates and translates HTTP; the +deployment and Redis store perform durable execution and recovery; the wrapper +remains independent of HTTP. + +### Rejected alternatives + +| Alternative | Reason for rejection | +|---|---| +| Durable authentication middleware | Couples the engine to authentication and duplicates host middleware | +| Mounted FastAPI/Starlette sub-app | Makes OpenAPI, prefixes, host middleware, and dynamic replacement harder | +| Runtime callback hooks | Bypass FastAPI dependency injection and make error/security behavior bespoke | +| Stateful integration class | Adds lifecycle and registration state that already exists in FastAPI and `DurableRuntime` | +| Manually documented endpoint examples only | Leaves each application to duplicate validation, errors, ownership, and links | +| `DurableRuntime.create_router()` | Couples the runtime layer directly to FastAPI | + +## State ownership + +Storing the runtime on `app.state` is intentional. It stores a process-local +service handle, not durable execution data. + +| Location | Owned data | +|---|---| +| `app.state` | Runtime object, deployment definitions, worker tasks, provider/client handles | +| Python process | Wrapper instances, active call stacks, event-loop tasks | +| Redis | Controls, inputs, checkpoints, progress, waits, results, errors, leases, indexes, idempotency bindings | + +After a process restart, the application reconstructs the runtime and wrapper +definitions. Redis retains nonterminal work. An expired running lease is +recovered and made claimable according to the existing state machine. + +No execution record, checkpoint, or result should be copied into `app.state`. +This centralized-state guarantee applies when using the Redis provider. The +in-memory provider is deliberately volatile, process-local, and suitable only +for development and tests. + +The public router does not need to read `app.state`; it closes over one +`DurableDeployment`. Keeping the runtime on `app.state` is recommended when +status routes, deployment utilities, or other application components need the +same process-local handle. Hayhooks itself needs that access. A small external +application may instead keep the runtime only in its application-factory +lifespan closure. + +## Target external-application experience + +### Wrapper authoring + +Pipeline wrappers remain independent of HTTP and authentication: + +```python +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper +from hayhooks.durable import DurableContext + + +class JobRequest(BaseModel): + document_id: str + + +class JobResult(BaseModel): + indexed: bool + + +class JobWrapper(BasePipelineWrapper): + durable_revision = "job-v1" + + def setup(self) -> None: + self.pipeline = build_pipeline() + + async def run_durable_async( + self, + context: DurableContext, + request: JobRequest, + ) -> JobResult: + result = await context.run_pipeline_async( + {"loader": {"document_id": request.document_id}}, + checkpoint_at=["loader"], + ) + return JobResult(indexed=bool(result["loader"]["indexed"])) +``` + +The deployment continues to derive its Pydantic request and result contracts +from the wrapper method annotations. + +### Authentication middleware + +When middleware has already authenticated the request and populated +`request.state`: + +```python +from fastapi import Request + + +def current_owner_id(request: Request) -> str: + principal = request.state.principal + return f"{principal.tenant_id}:{principal.subject_id}" +``` + +When the application already uses dependencies: + +```python +from typing import Annotated + +from fastapi import Depends + + +async def current_owner_id( + principal: Annotated[Principal, Depends(require_principal)], +) -> str: + return f"{principal.tenant_id}:{principal.subject_id}" +``` + +Both forms have the same durable behavior. + +### Runtime and router + +```python +from contextlib import asynccontextmanager + +from fastapi import Depends, FastAPI + +from hayhooks.durable import ( + DurableRuntime, + RedisExecutionStoreProvider, + create_durable_router, +) + + +@asynccontextmanager +async def lifespan(app: FastAPI): + runtime = app.state.durable_runtime + try: + await runtime.start() + yield + finally: + await runtime.close() + + +def create_app() -> FastAPI: + provider = RedisExecutionStoreProvider( + redis_url="redis://localhost:6379/0", + key_prefix="myapp:durable", + ) + runtime = DurableRuntime(provider) + + wrapper = JobWrapper() + wrapper.setup() + deployment = runtime.deployment("jobs", wrapper) + + app = FastAPI(lifespan=lifespan) + app.state.durable_runtime = runtime + app.include_router( + create_durable_router( + deployment, + owner_id_dependency=current_owner_id, + ), + prefix="/jobs", + dependencies=[Depends(require_jobs_permission)], + ) + return app + + +app = create_app() +``` + +The router-level authorization dependency is optional. Ownership and general +authorization remain distinct: + +- `require_jobs_permission` decides whether the caller may use the job API. +- `current_owner_id` provides the stable identity used to isolate records. + +## Public router behavior + +### Routes + +The returned router contains relative paths so the host application controls +the prefix: + +| Method | Relative path | Behavior | +|---|---|---| +| `POST` | `/run-durable` | Validate and submit detached work | +| `GET` | `/executions/{execution_id}` | Inspect safe execution state | +| `POST` | `/executions/{execution_id}/cancel` | Request cooperative cancellation | +| `POST` | `/executions/{execution_id}/resume` | Resume waiting work with optional typed input | + +Hayhooks includes the router at `/{pipeline_name}`, preserving all existing +paths. + +### Dependency binding inside the factory + +The factory fixes scoped versus unscoped behavior once, when the router is +created. When an owner dependency is supplied, each route binds it with +`Depends(...)`, validates its resolved return value, and sets +`enforce_owner=True`. When `None` is supplied, each route binds a private +constant dependency that returns `None`, and sets `enforce_owner=False`. + +Do not let a configured dependency return `None` to select unscoped behavior +at request time. That would turn an authentication bug into an authorization +bypass. The dependency's sync or async execution remains FastAPI's +responsibility. + +### Response behavior + +The extraction must preserve: + +| Situation | Response | +|---|---| +| New submission | `202 Accepted` | +| Nonterminal idempotent replay | `202 Accepted` plus `Idempotent-Replay: true` | +| Retained terminal replay | `200 OK` plus `Idempotent-Replay: true` | +| Accepted cancellation | `202 Accepted` | +| Already terminal cancellation | `200 OK` | +| Successful resume | `202 Accepted` | +| Missing or foreign-owned execution | `404 Not Found` | +| Idempotency or revision conflict | `409 Conflict` | +| Execution is not waiting | `409 Conflict` | +| Invalid ID, request, or resume body | `422 Unprocessable Entity` | +| Admission limit | `503 Service Unavailable` plus `Retry-After` | +| Execution-store outage | `503 Service Unavailable` | + +The `Location` and result links must be generated with named-route resolution +through `request.url_for(...)`, rather than by concatenating the deployment +name. Each factory result must give its routes deployment-unique names, such as +`hayhooks.durable.{deployment_name}.inspect`, so multiple durable deployments +cannot resolve one another's links. Hayhooks deployment names are unique within +an application; including the same deployment router more than once is outside +the initial contract. + +Named resolution keeps links correct when the router is included below +additional application prefixes or root paths. Preserve the current relative +link contract by using the resolved URL's path component for response links and +the `Location` header. + +### Typed OpenAPI models + +The adapter keeps the current dynamic request and result model behavior: + +- submission uses `deployment.request_type`; +- result fields use `deployment.result_type` when declared; +- resume uses `deployment.resume_type` when declared; +- inspect, cancel, and resume return the safe execution projection; +- private input, state, checkpoints, owner, and fencing details remain absent. + +The implementation may continue setting endpoint annotations/signatures after +handler construction because FastAPI consumes those annotations during route +registration. + +## Ownership and authentication contract + +### Owner ID rules + +When `owner_id_dependency` is supplied, the adapter must require: + +- a string; +- at least one character; +- no more than 512 characters; +- a stable value across token refreshes and process restarts. + +Recommended values are immutable application IDs, for example +`tenant_uuid:user_uuid`. Do not use access tokens, session IDs, emails, or +display names. + +The host application chooses the ownership granularity. Use a tenant ID for +tenant-owned jobs, a user ID for user-owned jobs, or a stable compound ID when +both boundaries matter. + +The dependency should perform authentication and may raise the host +application's normal `401` or `403`. The durable adapter must not replace those +responses. + +When the dependency is configured, a missing or invalid owner must fail closed. +It must never switch the request to unscoped access. + +### Enforcement + +The adapter passes the owner to every deployment operation: + +```python +await deployment.submit(..., owner_id=owner_id) +await deployment.get(..., owner_id=owner_id, enforce_owner=True) +await deployment.request_cancel(..., owner_id=owner_id, enforce_owner=True) +await deployment.resume(..., owner_id=owner_id, enforce_owner=True) +``` + +The existing deployment behavior returns `KeyError` for both missing records +and owner mismatches. The router maps both to `404`, avoiding an execution-ID +existence oracle. + +Owner-scoped submission continues deriving the internal execution ID from the +owner and caller-provided idempotency key. The same external key can therefore +be used independently by different owners. + +### Unscoped mode + +Passing `owner_id_dependency=None` explicitly retains the current behavior: + +- records have no owner; +- possession of a sufficiently unguessable execution ID grants access; +- the router does not enforce owner matching. + +This is useful for local development and services protected by a single +application-wide authorization boundary. Documentation must label it as an +explicit security choice, not an authentication default. + +### Wrapper access to identity + +Background work cannot depend on an HTTP request, middleware state, cookies, +or a current token. Those values do not exist after process recovery. + +Add a minimal public property to `DurableContext`: + +```python +@property +def owner_id(self) -> str | None: + return self.record.owner_id +``` + +This lets durable application code use the persisted stable identity without +exposing the full private execution record as its normal API. + +Roles and permissions should be checked before submission. Full principal +objects and tokens must not be persisted automatically. If a job needs +additional trusted identifiers, they must be deliberately represented as +non-secret validated input or durable application state. + +## Hayhooks dogfooding design + +### Public adapter boundary + +Create `src/hayhooks/durable/fastapi.py`. It owns all reusable HTTP behavior and +imports only durable public/infrastructure types plus FastAPI/Pydantic. + +It must not import: + +- `hayhooks.server.pipelines.registry`; +- the module-level `durable_runtime`; +- `hayhooks.settings`; +- `BasePipelineWrapper`; +- deployment or route mutation utilities from `hayhooks.server`. + +### Hayhooks server shim + +Reduce `hayhooks.server.durable.routes` to Hayhooks-specific composition: + +1. Determine whether the deployment has durable capability. +2. Build the trusted-header owner dependency from that deployment's settings. +3. Remove the previous durable route family for the pipeline. +4. Include `create_durable_router(deployment, ...)` at + `/{pipeline_name}`. +5. Invalidate/rebuild OpenAPI according to the existing deferred-rebuild flag. + +The trusted-header dependency remains a Hayhooks server concern because a +third-party application should normally use its authenticated principal rather +than trust a configurable raw header. + +The current `durable_request_model` and `durable_response_model` registry +metadata is not read elsewhere in the codebase. Remove those writes rather than +adding a public result object solely to preserve unused internal metadata. + +### Dynamic deployment lifecycle + +Hayhooks must preserve the existing publication transaction: + +1. Capture the current wrapper, routes, OpenAPI schema, modules, files, and + durable deployment. +2. Quiesce and close the previous deployment when replacing it. +3. Reject replacement while nonterminal durable work would be stranded. +4. Prepare the new wrapper and durable deployment before publication. +5. Build/include routes whose closures reference the new deployment. +6. Publish registry and runtime state and activate the candidate without an + intervening await point. +7. Restore the previous routes, runtime deployment, modules, files, and worker + state if publication fails. + +`APIRouter` inclusion produces ordinary `APIRoute` instances on the application, +so the current path-based removal and route-list snapshot rollback remain +usable. Do not introduce a route-mount handle or registration class unless the +existing mechanism proves insufficient in tests. + +## Application-owned runtime design + +### REST factory + +Update the application factory without breaking existing callers: + +```python +def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: + runtime = durable_runtime or DurableRuntime() + app = FastAPI(...) + app.state.durable_runtime = runtime + ... +``` + +The lifespan reads the runtime from the application instance, starts it after +startup pipeline preparation, and closes it before application-owned dependent +resources are closed. + +Status endpoints should read the runtime from `request.app.state`, not import a +singleton. + +Deployment helpers should receive the runtime explicitly when no application is +available, or use the runtime attached to the supplied application. Avoid a +fallback that silently selects the process-global runtime. + +### A2A and MCP factories + +Apply the same ownership rule to other server factories: + +- `create_a2a_app(..., durable_runtime=runtime)`; +- `create_agent_executor(..., durable_runtime=runtime)`; +- A2A health reads its supplied runtime; +- MCP server/application factories retain and close their supplied runtime. + +CLI commands construct one runtime for the server they launch and pass that +same instance through construction and lifespan. + +### Compatibility singleton + +Keep `hayhooks.durable.durable_runtime` temporarily as a compatibility export, +but stop using it inside Hayhooks. Document application-owned `DurableRuntime` +as the supported integration. Removal of the singleton, if desired, should be a +separate deprecation decision. + +Removing internal use of the singleton also removes the need for the private +runtime-to-registry attachment. Registry discovery remains a Hayhooks server +responsibility; standalone runtimes continue to know only deployments that the +application explicitly attaches. + +## Multi-process behavior + +Every Uvicorn worker process creates its own application state, runtime, worker +tasks, and Redis client/pool. All processes use the same Redis namespace and +definition revision for a logical deployment. + +The existing Redis fences and leases coordinate claims across processes. The +effective execution concurrency is: + +```text +application processes × durable_execution_concurrency per deployment +``` + +The first rollout should remain within the controlled-beta profile of one to +three replicas with conservative per-process concurrency. Application state is +not shared and must never be treated as a cross-process coordination mechanism. + +## Implementation phases + +### Phase 0: Lock the existing HTTP contract + +Before moving handlers, add or identify focused tests for every existing route, +status code, response header, owner behavior, typed model, and error mapping. +These tests define extraction compatibility. + +No implementation behavior changes in this phase. + +### Phase 1: Add the public FastAPI adapter + +Add: + +- `src/hayhooks/durable/fastapi.py`; +- `tests/test_durable_fastapi.py` containing a standalone FastAPI app. + +Update: + +- `src/hayhooks/durable/__init__.py` to lazily export + `create_durable_router`, `DurableContext`, `DurableDeployment`, and the public + durable exceptions; +- `src/hayhooks/durable/context.py` with `owner_id`; +- durable embedding documentation with authenticated and unscoped examples. + +Move, rather than rewrite, the existing validation, response projection, and +error mapping. Keep the diff behavior-preserving. + +### Phase 2: Make Hayhooks REST consume the adapter + +Update: + +- `src/hayhooks/server/durable/routes.py` to contain only trusted-header + dependency construction and dynamic route composition; +- `src/hayhooks/server/utils/deploy_utils.py` to pass the prepared + `DurableDeployment` into the composition shim; +- durable REST and deployment lifecycle tests. + +Delete the duplicated endpoint handlers from the server module. Preserve route +paths and operation behavior. + +### Phase 3: Replace internal global-runtime use + +Update REST, status, deploy/undeploy, A2A, MCP, and their app factories to pass +an application-owned runtime explicitly. Store the REST runtime on `app.state`. + +Keep each protocol migration reviewable. A safe order is: + +1. REST app and status endpoints; +2. REST dynamic deployment transaction; +3. A2A app, executors, and health; +4. MCP server and lifespan; +5. remove all internal imports of the singleton; +6. remove the private registry attachment from `runtime.py`. + +The phase is complete when searching `src/hayhooks/server` finds no import of +the module-level `durable_runtime`. + +### Phase 4: Focus durable configuration + +Introduce a durable-only configuration model containing Redis connection, +retention, record limits, lease, retry, polling, shutdown, admission, and +concurrency settings. + +Allow built-in providers and `DurableRuntime` to consume it without requiring +the full Hayhooks `AppSettings`. Retain an internal conversion from Hayhooks +settings for server compatibility. + +Keep this separate from route extraction so configuration changes cannot hide +HTTP regressions. + +### Phase 5: Release and adoption + +- Run the external FastAPI example against a built wheel rather than the source + checkout. +- Update API and durable-engine documentation. +- Note application-owned runtime and router integration in release notes. +- Pin the target application's trial to the released version. +- Start with one authenticated deployment and a dedicated Redis namespace. + +## Test plan + +### Public adapter tests + +The standalone test application must not import `hayhooks.server.app`, the +Hayhooks registry, or the global runtime. + +Cover: + +- typed submission and result validation; +- typed resume input; +- all four routes; +- generated OpenAPI schemas; +- new submission and idempotent replay headers/statuses; +- cancellation before and after terminal state; +- wait/resume and revision conflict; +- invalid execution IDs and idempotency keys; +- record-size and admission failures; +- normalized execution-store failures; +- safe views excluding private fields; +- link resolution beneath an additional `include_router(prefix=...)` prefix. + +### Authentication and ownership matrix + +| Scenario | Expected result | +|---|---| +| Middleware rejects unauthenticated caller | Host application's `401` | +| Owner dependency raises `403` | Host application's `403` | +| Configured dependency returns empty/invalid owner | Fail closed; never unscoped | +| Alice submits and inspects | Success | +| Bob inspects Alice's execution | `404` | +| Bob cancels Alice's execution | `404` | +| Bob resumes Alice's execution | `404` | +| Alice reuses same key and input | Idempotent replay | +| Alice reuses same key with different input | `409` | +| Bob uses Alice's external key | Independent execution | +| Explicit unscoped router | Existing bearer-ID behavior | + +Provide one test where middleware writes `request.state.principal`, and another +where the owner dependency depends on an existing FastAPI authentication +dependency. + +### Hayhooks dogfooding regression tests + +Preserve and extend tests for: + +- startup deployment route creation; +- dynamic deployment after runtime start; +- durable-to-durable overwrite; +- durable-to-nondurable overwrite; +- undeploy route removal; +- refusal to strand queued, running, or waiting work; +- candidate preparation failure; +- publication failure and route rollback; +- OpenAPI replacement with new request/result/resume types; +- trusted owner header compatibility; +- deferred OpenAPI rebuild during batch startup. + +### Runtime ownership tests + +- Two FastAPI applications in one process receive distinct runtimes. +- A deployment installed in app A is absent from app B. +- Closing app A does not close app B's provider or workers. +- A supplied runtime is the one used by routes, status, and deployment helpers. +- REST, A2A, and MCP lifespans close only the runtime they own. +- Startup failure still closes the corresponding runtime exactly once. + +### Existing engine tests + +The following remain mandatory: + +- reducer and reference-store contract tests; +- Redis transaction and concurrent-claim tests; +- manager retry, cancellation, and lease tests; +- process-kill/restart recovery; +- A2A recovery tests; +- type checking and linting. + +## Compatibility requirements + +### HTTP compatibility + +- Existing Hayhooks durable paths remain unchanged. +- Existing request and response payloads remain unchanged. +- Status codes and headers remain unchanged. +- Owner mismatch remains indistinguishable from a missing execution. +- Current trusted owner header behavior remains available in the Hayhooks + server shim. + +### Python compatibility + +- Existing direct imports continue to work. +- New public imports are additive. +- The global `durable_runtime` remains importable during the compatibility + period. +- `create_app()` with no arguments continues to work, but owns a new runtime. +- Supplying a runtime is keyword-only. + +### Persistence compatibility + +This work must not modify: + +- Redis key construction; +- control serialization; +- payload formats; +- idempotency digests; +- execution state transitions; +- lease or fence semantics. + +No namespace migration is required for this integration refactor. + +## Operational guidance + +- Use an isolated Redis key prefix per application/environment. +- Use Redis 6.2 or later with persistence, backups, TLS/authentication as + appropriate, and `maxmemory-policy noeviction`. +- If the host application supplies a Redis client, durable storage requires + binary responses (`decode_responses=False`). +- Set `close_redis=False` when the host owns the client lifecycle. +- Close the durable runtime before closing a shared Redis client. +- Keep wrapper definition revisions identical across replicas. +- Account for Uvicorn workers when setting durable concurrency. +- Continue making every external side effect idempotent because execution is + at least once. + +## Risks and mitigations + +| Risk | Mitigation | +|---|---| +| Router extraction changes an HTTP edge case | Lock the current contract before moving handlers | +| Included router links ignore a host prefix | Generate links from named routes and the current request | +| Owner dependency accidentally returns `None` | Fail closed whenever a dependency is configured | +| Authentication data is expected during background recovery | Persist only owner ID and deliberate validated identifiers | +| Dynamic overwrite leaves old handlers active | Remove the known route family before including the candidate router | +| Publication failure loses previous routes | Retain the existing route-list snapshot rollback | +| Multiple app instances share shutdown state | Construct and attach one runtime per app | +| Multi-worker deployments create unexpected concurrency | Document and test the process × slot calculation | +| Scope expands into a workflow framework | Keep generic runners, schedules, and transport-independent auth out of this work | + +## Acceptance criteria + +The work is complete when all of the following are true: + +1. An independent FastAPI application can integrate durable execution using + only public Hayhooks imports. +2. That application can supply identity from middleware or an existing FastAPI + authentication dependency. +3. Hayhooks' own durable REST endpoints are created by the same public router + factory. +4. The public router has no dependency on the Hayhooks registry, global runtime, + or server settings. +5. Hayhooks server modules no longer import the global durable runtime. +6. Every app/server factory owns and closes its runtime. +7. Durable wrappers remain HTTP-independent and can read their stable owner via + `DurableContext.owner_id`. +8. Existing REST paths, schemas, status codes, headers, and owner behavior are + preserved. +9. Dynamic deployment, rollback, and undeploy tests remain green. +10. Live Redis and process-recovery tests remain green. +11. Redis data written before this refactor remains readable without migration. +12. Documentation includes authenticated, unscoped, shared-Redis, and + multi-worker examples. + +## Deliberately deferred extensions + +Add these only when a real integration requires them: + +- per-operation authorization dependencies; +- API-only versus worker-only runtime modes; +- generic non-Haystack runners; +- richer persisted principal metadata; +- customizable durable route shapes; +- a stable Redis schema migration framework. + +The first portable release should contain the smallest complete boundary: +application-owned runtime, one public router factory, and one owner-ID +dependency. diff --git a/docs/reference/api-reference.md b/docs/reference/api-reference.md index dcfa949c..3a07802e 100644 --- a/docs/reference/api-reference.md +++ b/docs/reference/api-reference.md @@ -147,6 +147,11 @@ input, checkpoints, application state, ownership, and fence details remain server-side. See [Pipeline wrapper durable execution](../concepts/pipeline-wrapper.md#durable-execution) and the [durable engine contract](../advanced/durable-engine.md). +Existing FastAPI applications can expose the same routes with the public +`hayhooks.durable.create_durable_router()` factory. The application owns a +`DurableRuntime` in its lifespan and may provide a normal FastAPI dependency +that returns a stable owner ID. See [Embedding the runtime](../advanced/durable-engine.md#embedding-the-runtime). + ### OpenAI Compatibility #### Chat Completion diff --git a/src/hayhooks/cli/a2a.py b/src/hayhooks/cli/a2a.py index 04764435..c34fcc1a 100644 --- a/src/hayhooks/cli/a2a.py +++ b/src/hayhooks/cli/a2a.py @@ -60,6 +60,7 @@ def run( # noqa: PLR0913 # Lazy imports of settings, logger and uvicorn import uvicorn + from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.logger import intercept_stdlib_logging, log from hayhooks.server.utils.deploy_utils import deploy_pipelines @@ -104,8 +105,10 @@ def run( # noqa: PLR0913 sys.path.append(additional_python_path) log.trace("Added '{}' to sys.path", additional_python_path) + durable_runtime = DurableRuntime(app_settings=settings) + # Deploy the pipelines - deploy_pipelines() + deploy_pipelines(durable_runtime=durable_runtime) # Setup the Starlette app exposing pipelines as A2A agents log.debug( @@ -120,7 +123,7 @@ def run( # noqa: PLR0913 settings.durable_store, settings.durable_execution_concurrency, ) - app = create_a2a_app(debug=debug) + app = create_a2a_app(debug=debug, durable_runtime=durable_runtime) # Run the A2A server # NOTE: reload and workers options are not supported in this context diff --git a/src/hayhooks/cli/mcp.py b/src/hayhooks/cli/mcp.py index 87008365..8003c188 100644 --- a/src/hayhooks/cli/mcp.py +++ b/src/hayhooks/cli/mcp.py @@ -33,6 +33,7 @@ def run( # noqa: PLR0913 # Lazy imports of settings, logger and uvicorn import uvicorn + from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.logger import intercept_stdlib_logging, log from hayhooks.server.utils.deploy_utils import deploy_pipelines from hayhooks.server.utils.mcp_utils import create_mcp_server, create_starlette_app @@ -54,14 +55,21 @@ def run( # noqa: PLR0913 sys.path.append(additional_python_path) log.trace("Added '{}' to sys.path", additional_python_path) + durable_runtime = DurableRuntime(app_settings=settings) + # Deploy the pipelines - deploy_pipelines() + deploy_pipelines(durable_runtime=durable_runtime) # Setup the MCP server - server: Server = create_mcp_server() + server: Server = create_mcp_server(durable_runtime=durable_runtime) # Setup the Starlette app - app = create_starlette_app(server, debug=debug, json_response=json_response) + app = create_starlette_app( + server, + debug=debug, + json_response=json_response, + durable_runtime=durable_runtime, + ) # Run the MCP server # NOTE: reload and workers options are not supported in this context diff --git a/src/hayhooks/durable/__init__.py b/src/hayhooks/durable/__init__.py index 4b2959e0..b349bbd2 100644 --- a/src/hayhooks/durable/__init__.py +++ b/src/hayhooks/durable/__init__.py @@ -12,7 +12,24 @@ from hayhooks.durable.models import ExecutionStatus if TYPE_CHECKING: - from hayhooks.durable.runtime import DurableRuntime, ExecutionStoreProvider + from hayhooks.durable.context import DurableContext + from hayhooks.durable.fastapi import create_durable_router + from hayhooks.durable.models import ( + ExecutionAdmissionError, + ExecutionCanceledError, + ExecutionRecordSizeError, + ExecutionStoreError, + ExecutionSuspendedError, + RetryableExecutionError, + ) + from hayhooks.durable.runtime import ( + DefinitionRevisionConflictError, + DurableDeployment, + DurableRuntime, + ExecutionStoreProvider, + IdempotencyConflictError, + ) + from hayhooks.durable.settings import DurableSettings from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider @@ -56,14 +73,68 @@ def current_durable_context() -> Any | None: def __getattr__(name: str) -> Any: """Lazily expose durable infrastructure without eager optional imports.""" - if name in {"DurableRuntime", "ExecutionStoreProvider", "durable_runtime"}: - from hayhooks.durable.runtime import DurableRuntime, ExecutionStoreProvider, durable_runtime + if name == "DurableContext": + from hayhooks.durable.context import DurableContext + + return DurableContext + if name == "create_durable_router": + from hayhooks.durable.fastapi import create_durable_router + + return create_durable_router + if name == "DurableSettings": + from hayhooks.durable.settings import DurableSettings + + return DurableSettings + if name in { + "DefinitionRevisionConflictError", + "DurableDeployment", + "DurableRuntime", + "ExecutionStoreProvider", + "IdempotencyConflictError", + "durable_runtime", + }: + from hayhooks.durable.runtime import ( + DefinitionRevisionConflictError, + DurableDeployment, + DurableRuntime, + ExecutionStoreProvider, + IdempotencyConflictError, + durable_runtime, + ) return { + "DefinitionRevisionConflictError": DefinitionRevisionConflictError, + "DurableDeployment": DurableDeployment, "DurableRuntime": DurableRuntime, "ExecutionStoreProvider": ExecutionStoreProvider, + "IdempotencyConflictError": IdempotencyConflictError, "durable_runtime": durable_runtime, }[name] + if name in { + "ExecutionAdmissionError", + "ExecutionCanceledError", + "ExecutionRecordSizeError", + "ExecutionStoreError", + "ExecutionSuspendedError", + "RetryableExecutionError", + }: + from hayhooks.durable.models import ( + ExecutionAdmissionError, + ExecutionCanceledError, + ExecutionRecordSizeError, + ExecutionStoreError, + ExecutionSuspendedError, + RetryableExecutionError, + ) + + return { + "ExecutionAdmissionError": ExecutionAdmissionError, + "ExecutionCanceledError": ExecutionCanceledError, + "ExecutionRecordSizeError": ExecutionRecordSizeError, + "ExecutionStoreError": ExecutionStoreError, + "ExecutionSuspendedError": ExecutionSuspendedError, + "RetryableExecutionError": RetryableExecutionError, + }[name] if name in {"ExecutionStore", "InMemoryExecutionStoreProvider", "RedisExecutionStoreProvider"}: from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider @@ -77,15 +148,27 @@ def __getattr__(name: str) -> Any: __all__ = [ + "DefinitionRevisionConflictError", "DurableAuthoringMode", + "DurableContext", + "DurableDeployment", "DurableRuntime", + "DurableSettings", + "ExecutionAdmissionError", + "ExecutionCanceledError", "ExecutionProgress", + "ExecutionRecordSizeError", "ExecutionResult", "ExecutionStatus", "ExecutionStore", + "ExecutionStoreError", "ExecutionStoreProvider", + "ExecutionSuspendedError", + "IdempotencyConflictError", "InMemoryExecutionStoreProvider", "RedisExecutionStoreProvider", + "RetryableExecutionError", + "create_durable_router", "current_durable_context", "current_execution_id", "durable_authoring_mode", diff --git a/src/hayhooks/durable/context.py b/src/hayhooks/durable/context.py index 50d07df4..c9e012ff 100644 --- a/src/hayhooks/durable/context.py +++ b/src/hayhooks/durable/context.py @@ -66,6 +66,11 @@ def execution_id(self) -> str: def attempt(self) -> int: return self.record.attempt + @property + def owner_id(self) -> str | None: + """Return the stable owner identity persisted with this execution.""" + return self.record.owner_id + @property def state(self) -> dict[str, JsonValue]: return self.record.application_state diff --git a/src/hayhooks/durable/fastapi.py b/src/hayhooks/durable/fastapi.py new file mode 100644 index 00000000..568a6911 --- /dev/null +++ b/src/hayhooks/durable/fastapi.py @@ -0,0 +1,246 @@ +"""FastAPI routes for one application-owned durable deployment.""" + +from __future__ import annotations + +import inspect +import re +from collections.abc import Awaitable, Callable +from typing import Annotated, Any, cast + +from fastapi import APIRouter, Body, Depends, Header, HTTPException, Path, Request, Response, status +from pydantic import ValidationError, create_model + +from hayhooks.durable import ExecutionResult +from hayhooks.durable.engine import RUN_ID_PATTERN +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError +from hayhooks.durable.runtime import DefinitionRevisionConflictError, DurableDeployment, IdempotencyConflictError + +_IDEMPOTENCY_KEY_PATTERN = re.compile(rf"^{RUN_ID_PATTERN}$") +_MAX_OWNER_LENGTH = 512 +_MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH = 63 +ExecutionId = Annotated[str, Path(pattern=rf"^{RUN_ID_PATTERN}$", min_length=1, max_length=128)] +OwnerIdDependency = Callable[..., str | Awaitable[str]] + + +def _unscoped_owner() -> None: + return None + + +def _validated_owner(owner_id: Any, *, enforce_owner: bool) -> str | None: + if not enforce_owner: + return None + if not isinstance(owner_id, str) or not owner_id or len(owner_id) > _MAX_OWNER_LENGTH: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="The configured owner dependency must return a non-empty string of at most 512 characters", + ) + return owner_id + + +def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResult]: + if deployment.result_type is Any: + return ExecutionResult + return create_model( + f"{deployment.name.title().replace('-', '').replace('_', '')}ExecutionResult", + __base__=ExecutionResult, + result=(deployment.result_type | None, None), + ) + + +def create_durable_router( # noqa: C901, PLR0915 - handlers share one deployment and generated models + deployment: DurableDeployment, + *, + owner_id_dependency: OwnerIdDependency | None, +) -> APIRouter: + """ + Return the typed HTTP API for one durable deployment. + + Passing ``None`` explicitly enables unscoped bearer-by-execution-ID access. + The caller owns the deployment and its runtime lifecycle. + """ + router = APIRouter() + owner_dependency = owner_id_dependency or _unscoped_owner + enforce_owner = owner_id_dependency is not None + response_model = _durable_response_model(deployment) + route_names = { + "submit": f"hayhooks.durable.{deployment.name}.submit", + "inspect": f"hayhooks.durable.{deployment.name}.inspect", + "cancel": f"hayhooks.durable.{deployment.name}.cancel", + "resume": f"hayhooks.durable.{deployment.name}.resume", + } + + def execution_links(request: Request, execution_id: str) -> dict[str, str]: + root = request.url_for(route_names["inspect"], execution_id=execution_id).path + return { + "self": root, + "cancel": request.url_for(route_names["cancel"], execution_id=execution_id).path, + "resume": request.url_for(route_names["resume"], execution_id=execution_id).path, + } + + def execution_result( + request: Request, + record: Any, + *, + model: type[ExecutionResult] = ExecutionResult, + ) -> ExecutionResult: + return model.model_validate(record.safe_view(links=execution_links(request, record.execution_id))) + + async def get_execution(execution_id: str, owner_id: str | None) -> Any: + return await deployment.get( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + allow_revision_mismatch=True, + ) + + async def submit( + run_req: Any, + response: Response, + request: Request, + owner_id: Any = Depends(owner_dependency), # noqa: B008 + idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), + ) -> ExecutionResult: + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + if idempotency_key is not None and _IDEMPOTENCY_KEY_PATTERN.fullmatch(idempotency_key) is None: + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail="Idempotency-Key must contain 1-128 letters, digits, underscores, or hyphens", + ) + if ( + enforce_owner + and idempotency_key is not None + and len(idempotency_key) > _MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH + ): + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail="Idempotency-Key must be at most 63 characters when owner scoping is enabled", + ) + try: + created, record = await deployment.submit( + run_req.model_dump(mode="json"), + execution_id=idempotency_key, + owner_id=owner_id, + ) + except (IdempotencyConflictError, DefinitionRevisionConflictError) as error: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error + except (ValidationError, ValueError) as error: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error + except ExecutionAdmissionError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=str(error), + headers={"Retry-After": str(error.retry_after_seconds)}, + ) from error + except (ExecutionStoreError, RuntimeError) as error: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error + response.status_code = status.HTTP_200_OK if not created and record.terminal else status.HTTP_202_ACCEPTED + result = execution_result(request, record, model=response_model) + response.headers["Location"] = result.links["self"] + if not created: + response.headers["Idempotent-Replay"] = "true" + return result + + submit.__annotations__["run_req"] = deployment.request_type + + async def inspect_execution( + execution_id: ExecutionId, + request: Request, + owner_id: Any = Depends(owner_dependency), # noqa: B008 + ) -> ExecutionResult: + try: + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + return execution_result(request, await get_execution(execution_id, owner_id)) + except KeyError as error: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error + except ExecutionStoreError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Durable execution store is unavailable", + ) from error + + async def cancel_execution( + execution_id: ExecutionId, + response: Response, + request: Request, + owner_id: Any = Depends(owner_dependency), # noqa: B008 + ) -> ExecutionResult: + try: + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + accepted = await deployment.request_cancel( + execution_id, + owner_id=owner_id, + enforce_owner=enforce_owner, + ) + record = await get_execution(execution_id, owner_id) + except KeyError as error: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error + except ExecutionStoreError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Durable execution store is unavailable", + ) from error + response.status_code = status.HTTP_202_ACCEPTED if accepted else status.HTTP_200_OK + return execution_result(request, record) + + async def resume_execution( + execution_id: ExecutionId, + response: Response, + request: Request, + owner_id: Any = Depends(owner_dependency), # noqa: B008 + update: Any = Body(default=None), # noqa: B008 + ) -> ExecutionResult: + try: + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + resumed = await deployment.resume( + execution_id, + update, + owner_id=owner_id, + enforce_owner=enforce_owner, + ) + if not resumed: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Execution is not waiting") + record = await get_execution(execution_id, owner_id) + except KeyError as error: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error + except DefinitionRevisionConflictError as error: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error + except (ValidationError, ValueError) as error: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error + except ExecutionStoreError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Durable execution store is unavailable", + ) from error + response.status_code = status.HTTP_202_ACCEPTED + return execution_result(request, record) + + if deployment.resume_type is not None: + resume_execution.__annotations__["update"] = deployment.resume_type + signature = inspect.signature(resume_execution) + update_parameter = signature.parameters["update"].replace(annotation=deployment.resume_type, default=Body()) + cast(Any, resume_execution).__signature__ = signature.replace( + parameters=[ + update_parameter if parameter.name == "update" else parameter + for parameter in signature.parameters.values() + ] + ) + + for path, endpoint, methods, name in ( + ("/run-durable", submit, ["POST"], route_names["submit"]), + ("/executions/{execution_id}", inspect_execution, ["GET"], route_names["inspect"]), + ("/executions/{execution_id}/cancel", cancel_execution, ["POST"], route_names["cancel"]), + ("/executions/{execution_id}/resume", resume_execution, ["POST"], route_names["resume"]), + ): + router.add_api_route( + path, + endpoint, + methods=methods, + name=name, + response_model=response_model if endpoint is submit else ExecutionResult, + tags=["durable executions"], + status_code=status.HTTP_202_ACCEPTED if methods == ["POST"] else status.HTTP_200_OK, + ) + return router + + +__all__ = ["OwnerIdDependency", "create_durable_router"] diff --git a/src/hayhooks/durable/runtime.py b/src/hayhooks/durable/runtime.py index aae15dd4..266f1c78 100644 --- a/src/hayhooks/durable/runtime.py +++ b/src/hayhooks/durable/runtime.py @@ -18,13 +18,12 @@ from hayhooks.durable.manager import DurableExecutionManager from hayhooks.durable.mode import DurableAuthoringMode, _durable_method_implementations, durable_authoring_mode from hayhooks.durable.models import ExecutionKind, ExecutionRecord, JsonValue, json_safe, validate_json +from hayhooks.durable.settings import DurableSettings from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider from hayhooks.server.exceptions import PipelineWrapperError from hayhooks.server.logger import log -from hayhooks.server.pipelines.registry import registry from hayhooks.server.tracing import SPAN_DURABLE_ATTEMPT, SPAN_DURABLE_SUBMIT, build_trace_tags, trace_operation from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper -from hayhooks.settings import AppSettings, settings class ExecutionStoreProvider(Protocol): @@ -35,14 +34,24 @@ def create_execution_store(self, deployment_name: str) -> ExecutionStore: ... async def close(self) -> None: ... -def _runtime_settings(provider: ExecutionStoreProvider | None, app_settings: AppSettings | None) -> AppSettings | None: - provider_settings = getattr(provider, "app_settings", None) - if isinstance(provider_settings, AppSettings): - if app_settings is not None and app_settings != provider_settings: +def _runtime_settings( + provider: ExecutionStoreProvider | None, + durable_settings: DurableSettings | None, + app_settings: Any | None, +) -> DurableSettings | None: + if durable_settings is not None and app_settings is not None: + msg = "Pass durable_settings or app_settings, not both" + raise ValueError(msg) + requested = durable_settings or ( + DurableSettings.from_app_settings(app_settings) if app_settings is not None else None + ) + provider_settings = getattr(provider, "settings", None) + if isinstance(provider_settings, DurableSettings): + if requested is not None and requested != provider_settings: msg = "Durable runtime and execution-store provider settings must match" raise ValueError(msg) - app_settings = provider_settings - return app_settings.model_copy(deep=True) if app_settings is not None else None + requested = provider_settings + return requested.model_copy(deep=True) if requested is not None else None class DurableDeployment: @@ -54,11 +63,12 @@ def __init__( wrapper: BasePipelineWrapper, provider: ExecutionStoreProvider, *, - app_settings: AppSettings | None = None, + durable_settings: DurableSettings | None = None, + app_settings: Any | None = None, ) -> None: self.name = name self.wrapper = wrapper - self.app_settings = _runtime_settings(provider, app_settings) or settings.model_copy(deep=True) + self.settings = _runtime_settings(provider, durable_settings, app_settings) or DurableSettings() pipeline = wrapper.pipeline try: kind = execution_kind(pipeline) @@ -94,12 +104,12 @@ def __init__( self.store, self._run, self.adapter, - concurrency=self.app_settings.durable_execution_concurrency, - poll_interval=self.app_settings.durable_poll_interval, - shutdown_grace_period=self.app_settings.durable_shutdown_grace_period, - max_attempts=self.app_settings.durable_max_attempts, - retry_base_delay=self.app_settings.durable_retry_base_delay, - retry_max_delay=self.app_settings.durable_retry_max_delay, + concurrency=self.settings.durable_execution_concurrency, + poll_interval=self.settings.durable_poll_interval, + shutdown_grace_period=self.settings.durable_shutdown_grace_period, + max_attempts=self.settings.durable_max_attempts, + retry_base_delay=self.settings.durable_retry_base_delay, + retry_max_delay=self.settings.durable_retry_max_delay, ) async def start(self) -> None: @@ -144,7 +154,7 @@ async def submit( validated_input = cast( dict[str, JsonValue], validate_json( - request.model_dump(mode="json"), limit=self.app_settings.durable_max_record_bytes, label="request" + request.model_dump(mode="json"), limit=self.settings.durable_max_record_bytes, label="request" ), ) fingerprint_input = cast(dict[str, JsonValue], _canonical_json(request.model_dump(mode="python"))) @@ -162,8 +172,8 @@ async def submit( validated_input=validated_input, operation_fingerprint=operation_fingerprint, owner_id=owner_id, - max_progress_events=self.app_settings.durable_max_progress_events, - max_record_bytes=self.app_settings.durable_max_record_bytes, + max_progress_events=self.settings.durable_max_progress_events, + max_record_bytes=self.settings.durable_max_record_bytes, ) with trace_operation( SPAN_DURABLE_SUBMIT, @@ -350,14 +360,14 @@ def __init__( self, provider: ExecutionStoreProvider | None = None, *, - app_settings: AppSettings | None = None, + durable_settings: DurableSettings | None = None, + app_settings: Any | None = None, ) -> None: self._store_provider = provider - self._app_settings = _runtime_settings(provider, app_settings) + self._durable_settings = _runtime_settings(provider, durable_settings, app_settings) self._deployments: dict[str, DurableDeployment] = {} self._started = False self._provider_close_task: asyncio.Task[None] | None = None - self._registry: Any | None = None def has_capability(self, wrapper: BasePipelineWrapper) -> bool: return durable_authoring_mode(wrapper) is not DurableAuthoringMode.NONE @@ -372,18 +382,22 @@ def provider(self) -> ExecutionStoreProvider | None: return self._store_provider @property - def app_settings(self) -> AppSettings: - """Return configured or provider settings, falling back to Hayhooks' global settings.""" - if self._app_settings is not None: - return self._app_settings - provider_settings = getattr(self.provider, "app_settings", None) - return provider_settings if isinstance(provider_settings, AppSettings) else settings + def settings(self) -> DurableSettings: + """Return the runtime's durable-only configuration.""" + if self._durable_settings is None: + self._durable_settings = DurableSettings() + return self._durable_settings + + @property + def app_settings(self) -> DurableSettings: + """Compatibility alias for :attr:`settings`.""" + return self.settings def create_deployment(self, name: str, wrapper: BasePipelineWrapper) -> DurableDeployment | None: """Build an uncached candidate so route closures cannot capture an old deployment.""" if not self.has_capability(wrapper): return None - return DurableDeployment(name, wrapper, self._provider(), app_settings=self.app_settings) + return DurableDeployment(name, wrapper, self._provider(), durable_settings=self.settings) def current_deployment(self, name: str) -> DurableDeployment | None: """Return the currently published durable deployment, if any.""" @@ -401,7 +415,6 @@ def deployment(self, name: str, wrapper: BasePipelineWrapper | None = None) -> D existing = self._deployments.get(name) if existing is not None and (wrapper is None or existing.wrapper is wrapper): return existing - wrapper = wrapper or (self._registry.get(name) if self._registry is not None else None) if wrapper is None or not self.has_capability(wrapper): msg = f"Pipeline '{name}' does not expose durable execution" raise KeyError(msg) @@ -411,22 +424,17 @@ def deployment(self, name: str, wrapper: BasePipelineWrapper | None = None) -> D "use the async deployment transaction to replace it" ) raise RuntimeError(msg) - deployment = DurableDeployment(name, wrapper, self._provider(), app_settings=self.app_settings) + deployment = DurableDeployment(name, wrapper, self._provider(), durable_settings=self.settings) self._deployments[name] = deployment return deployment async def start(self) -> None: - """Start runtime-owned deployments after optional registry discovery.""" + """Start runtime-owned deployments.""" if self._started: return started: list[DurableDeployment] = [] self._started = True try: - if self._registry is not None: - for name in self._registry.get_names(): - wrapper = self._registry.get(name) - if wrapper is not None and self.has_capability(wrapper): - self.deployment(name, wrapper) for deployment in list(self._deployments.values()): if deployment.manager.started: continue @@ -484,11 +492,11 @@ async def health(self) -> dict[str, JsonValue]: def _provider(self) -> ExecutionStoreProvider: provider = self._store_provider if provider is None: - if self.app_settings.durable_store == "memory": + if self.settings.durable_store == "memory": log.warning("Durable execution uses volatile in-memory storage; queued work is lost on process exit") - provider = InMemoryExecutionStoreProvider(app_settings=self.app_settings) + provider = InMemoryExecutionStoreProvider(durable_settings=self.settings) else: - provider = RedisExecutionStoreProvider(app_settings=self.app_settings) + provider = RedisExecutionStoreProvider(durable_settings=self.settings) self._store_provider = provider return provider @@ -571,7 +579,6 @@ def _canonical_json(value: Any) -> JsonValue: durable_runtime = DurableRuntime() -durable_runtime._registry = registry __all__ = [ diff --git a/src/hayhooks/durable/settings.py b/src/hayhooks/durable/settings.py new file mode 100644 index 00000000..04efe59e --- /dev/null +++ b/src/hayhooks/durable/settings.py @@ -0,0 +1,51 @@ +"""Configuration owned by the portable durable runtime.""" + +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, Field, model_validator +from typing_extensions import Self + + +class DurableSettings(BaseModel): + """Durable storage, retention, retry, lease, and worker settings.""" + + durable_store: Literal["memory", "redis"] = "redis" + durable_redis_url: str = "redis://localhost:6379/0" + durable_redis_key_prefix: str = "hayhooks:durable" + durable_redis_socket_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + durable_redis_socket_connect_timeout: float = Field(default=5.0, gt=0.0, le=300.0) + durable_redis_health_check_interval: int = Field(default=30, ge=0, le=3_600) + durable_terminal_ttl_seconds: int = Field(default=604_800, ge=1) + durable_max_progress_events: int = Field(default=100, ge=1, le=10_000) + durable_max_record_bytes: int = Field(default=1_000_000, ge=1_024) + durable_max_nonterminal_executions: int = Field(default=0, ge=0) + durable_shutdown_grace_period: float = Field(default=5.0, ge=0.0) + durable_max_attempts: int = Field(default=3, ge=1, le=1_000) + durable_retry_base_delay: float = Field(default=1.0, ge=0.0, le=86_400.0) + durable_retry_max_delay: float = Field(default=60.0, ge=0.0, le=604_800.0) + durable_poll_interval: float = Field(default=1.0, ge=0.05, le=60.0) + durable_lease_duration_ms: int = Field(default=30_000, ge=1, le=86_400_000) + durable_lease_commit_safety_ms: int = Field(default=1_500, ge=0, le=86_400_000) + durable_execution_concurrency: int = Field(default=1, ge=1, le=128) + + @model_validator(mode="after") + def _validate_lease_margin(self) -> Self: + if self.durable_lease_commit_safety_ms >= self.durable_lease_duration_ms: + msg = "durable_lease_commit_safety_ms must be smaller than durable_lease_duration_ms" + raise ValueError(msg) + if self.durable_lease_duration_ms - self.durable_lease_commit_safety_ms <= max( + 10, self.durable_lease_duration_ms / 3 + ): + msg = "durable lease duration minus commit safety must exceed the heartbeat interval" + raise ValueError(msg) + return self + + @classmethod + def from_app_settings(cls, app_settings: Any) -> DurableSettings: + """Copy durable fields from Hayhooks settings without retaining that dependency.""" + return cls(**{name: getattr(app_settings, name) for name in cls.model_fields}) + + +__all__ = ["DurableSettings"] diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py index c76c17b1..23f11147 100644 --- a/src/hayhooks/durable/store.py +++ b/src/hayhooks/durable/store.py @@ -47,8 +47,8 @@ ) from hayhooks.durable.redis import RedisExecutionStore, digest from hayhooks.durable.reference import InMemoryExecutionStore +from hayhooks.durable.settings import DurableSettings from hayhooks.server.logger import log -from hayhooks.settings import AppSettings, settings _RECORD_PAYLOADS = ( PayloadKind.INPUT, @@ -655,27 +655,34 @@ def __init__( # noqa: PLR0913 - mirrors the configurable Redis task-store provi redis: Any | None = None, key_prefix: str | None = None, close_redis: bool = True, - app_settings: AppSettings | None = None, + durable_settings: DurableSettings | None = None, + app_settings: Any | None = None, socket_timeout: float | None = None, socket_connect_timeout: float | None = None, health_check_interval: int | None = None, ) -> None: - source_settings = app_settings if app_settings is not None else settings - self.app_settings = source_settings.model_copy(deep=True) - self.config = _config(app_settings=self.app_settings, key_prefix=key_prefix) + if durable_settings is not None and app_settings is not None: + msg = "Pass durable_settings or app_settings, not both" + raise ValueError(msg) + self.settings = ( + durable_settings + or (DurableSettings.from_app_settings(app_settings) if app_settings is not None else DurableSettings()) + ).model_copy(deep=True) + self.app_settings = self.settings + self.config = _config(durable_settings=self.settings, key_prefix=key_prefix) self.close_redis = close_redis self.socket_timeout = ( - socket_timeout if socket_timeout is not None else self.app_settings.durable_redis_socket_timeout + socket_timeout if socket_timeout is not None else self.settings.durable_redis_socket_timeout ) self.socket_connect_timeout = ( socket_connect_timeout if socket_connect_timeout is not None - else self.app_settings.durable_redis_socket_connect_timeout + else self.settings.durable_redis_socket_connect_timeout ) self.health_check_interval = ( health_check_interval if health_check_interval is not None - else self.app_settings.durable_redis_health_check_interval + else self.settings.durable_redis_health_check_interval ) if redis is None: try: @@ -684,7 +691,7 @@ def __init__( # noqa: PLR0913 - mirrors the configurable Redis task-store provi msg = 'Durable Redis storage requires `pip install "hayhooks[durable]`.' raise ImportError(msg) from error redis = Redis.from_url( - redis_url or self.app_settings.durable_redis_url, + redis_url or self.settings.durable_redis_url, decode_responses=False, socket_timeout=self.socket_timeout, socket_connect_timeout=self.socket_connect_timeout, @@ -700,7 +707,7 @@ def create_execution_store(self, deployment_name: str) -> ExecutionStore: self.cores[deployment_name] = core # A candidate deployment must not mutate the active deployment's accepted # definition revision while it is still preparing or rolling back. - return _execution_store(core, app_settings=self.app_settings) + return _execution_store(core, durable_settings=self.settings) async def close(self) -> None: if self.close_redis: @@ -710,10 +717,21 @@ async def close(self) -> None: class InMemoryExecutionStoreProvider: """Volatile reference backend for local development and tests.""" - def __init__(self, *, app_settings: AppSettings | None = None) -> None: - source_settings = app_settings if app_settings is not None else settings - self.app_settings = source_settings.model_copy(deep=True) - self.config = _config(app_settings=self.app_settings) + def __init__( + self, + *, + durable_settings: DurableSettings | None = None, + app_settings: Any | None = None, + ) -> None: + if durable_settings is not None and app_settings is not None: + msg = "Pass durable_settings or app_settings, not both" + raise ValueError(msg) + self.settings = ( + durable_settings + or (DurableSettings.from_app_settings(app_settings) if app_settings is not None else DurableSettings()) + ).model_copy(deep=True) + self.app_settings = self.settings + self.config = _config(durable_settings=self.settings) self.cores: dict[str, InMemoryExecutionStore] = {} def create_execution_store(self, deployment_name: str) -> ExecutionStore: @@ -721,37 +739,37 @@ def create_execution_store(self, deployment_name: str) -> ExecutionStore: if core is None: core = InMemoryExecutionStore(deployment=deployment_name, config=self.config) self.cores[deployment_name] = core - return _execution_store(core, app_settings=self.app_settings) + return _execution_store(core, durable_settings=self.settings) async def close(self) -> None: return None -def _execution_store(core: ExecutionBackend, *, app_settings: AppSettings) -> ExecutionStore: +def _execution_store(core: ExecutionBackend, *, durable_settings: DurableSettings) -> ExecutionStore: """Build the same public adapter for both built-in backend implementations.""" return ExecutionStore( core, - lease_duration_ms=app_settings.durable_lease_duration_ms, - max_run_attempts=app_settings.durable_max_attempts, - max_progress_events=app_settings.durable_max_progress_events, - max_record_bytes=app_settings.durable_max_record_bytes, + lease_duration_ms=durable_settings.durable_lease_duration_ms, + max_run_attempts=durable_settings.durable_max_attempts, + max_progress_events=durable_settings.durable_max_progress_events, + max_record_bytes=durable_settings.durable_max_record_bytes, ) -def _config(*, app_settings: AppSettings, key_prefix: str | None = None) -> ExecutionStoreConfig: - max_record = app_settings.durable_max_record_bytes +def _config(*, durable_settings: DurableSettings, key_prefix: str | None = None) -> ExecutionStoreConfig: + max_record = durable_settings.durable_max_record_bytes progress_bytes = DEFAULT_MAX_PROGRESS_BYTES return ExecutionStoreConfig( - key_prefix=key_prefix or app_settings.durable_redis_key_prefix, - lease_commit_safety_ms=app_settings.durable_lease_commit_safety_ms, - terminal_ttl_seconds=app_settings.durable_terminal_ttl_seconds, - max_nonterminal_executions=app_settings.durable_max_nonterminal_executions, + key_prefix=key_prefix or durable_settings.durable_redis_key_prefix, + lease_commit_safety_ms=durable_settings.durable_lease_commit_safety_ms, + terminal_ttl_seconds=durable_settings.durable_terminal_ttl_seconds, + max_nonterminal_executions=durable_settings.durable_max_nonterminal_executions, max_input_bytes=max_record, max_checkpoint_bytes=max_record, max_result_bytes=max_record, max_error_bytes=max_record, max_wait_bytes=max_record, - max_progress_events=app_settings.durable_max_progress_events, + max_progress_events=durable_settings.durable_max_progress_events, max_progress_event_bytes=progress_bytes, ) diff --git a/src/hayhooks/server/a2a/app.py b/src/hayhooks/server/a2a/app.py index 49d7f373..43ab46fa 100644 --- a/src/hayhooks/server/a2a/app.py +++ b/src/hayhooks/server/a2a/app.py @@ -8,7 +8,7 @@ from hayhooks.a2a import TaskStoreProvider from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.a2a.cards import create_agent_card, get_a2a_base_url, is_a2a_exposable from hayhooks.server.a2a.executor import DurableAgentExecutor, create_agent_executor from hayhooks.server.a2a.imports import DefaultRequestHandler, create_agent_card_routes, create_jsonrpc_routes @@ -22,7 +22,12 @@ _RESERVED_PATHS = frozenset({"status"}) -def _create_agent_mount(pipeline_name: str, base_url: str, runtime: A2ARuntime) -> Mount: +def _create_agent_mount( + pipeline_name: str, + base_url: str, + runtime: A2ARuntime, + durable_runtime: DurableRuntime, +) -> Mount: wrapper = registry.get(pipeline_name) if wrapper is None: msg = f"Pipeline '{pipeline_name}' not found" @@ -31,7 +36,12 @@ def _create_agent_mount(pipeline_name: str, base_url: str, runtime: A2ARuntime) card = create_agent_card(pipeline_name, base_url) task_store = runtime.create_task_store(pipeline_name) - agent_executor = create_agent_executor(wrapper, pipeline_name, task_store=task_store) + agent_executor = create_agent_executor( + wrapper, + pipeline_name, + task_store=task_store, + durable_runtime=durable_runtime, + ) if isinstance(agent_executor, DurableAgentExecutor): task_store = agent_executor.task_store runtime.register_agent_executor(agent_executor) @@ -73,7 +83,11 @@ def _create_app_task_store_provider(durable_agents_deployed: bool) -> TaskStoreP ) -def _create_agent_mounts(base_url: str, runtime: A2ARuntime) -> tuple[list[str], list[Mount]]: +def _create_agent_mounts( + base_url: str, + runtime: A2ARuntime, + durable_runtime: DurableRuntime, +) -> tuple[list[str], list[Mount]]: agent_names: list[str] = [] mounts: list[Mount] = [] for pipeline_name in registry.get_names(): @@ -83,7 +97,7 @@ def _create_agent_mounts(base_url: str, runtime: A2ARuntime) -> tuple[list[str], log.warning("Skipping pipeline '{}': the path is reserved by the A2A server", pipeline_name) continue try: - mounts.append(_create_agent_mount(pipeline_name, base_url, runtime)) + mounts.append(_create_agent_mount(pipeline_name, base_url, runtime, durable_runtime)) except Exception as error: log.opt(exception=True).warning( "Skipping pipeline '{}': failed to build A2A agent: {}", @@ -96,7 +110,13 @@ def _create_agent_mounts(base_url: str, runtime: A2ARuntime) -> tuple[list[str], return agent_names, mounts -def create_a2a_app(*, base_url: str | None = None, debug: bool = False, runtime: A2ARuntime | None = None) -> Starlette: +def create_a2a_app( + *, + base_url: str | None = None, + debug: bool = False, + runtime: A2ARuntime | None = None, + durable_runtime: DurableRuntime | None = None, +) -> Starlette: """ Create a Starlette app exposing deployed pipelines as A2A agents. @@ -110,10 +130,18 @@ def create_a2a_app(*, base_url: str | None = None, debug: bool = False, runtime: for name in registry.get_names() if (wrapper := registry.get(name)) is not None ) + durable_runtime = durable_runtime or (runtime.durable_runtime if runtime is not None else None) + durable_runtime = durable_runtime or DurableRuntime(app_settings=settings) if runtime is None: runtime = A2ARuntime( task_store_provider=_create_app_task_store_provider(durable_agents_deployed), + durable_runtime=durable_runtime, ) + elif runtime.durable_runtime is None: + runtime.durable_runtime = durable_runtime + elif runtime.durable_runtime is not durable_runtime: + msg = "A2A and durable runtimes must reference the same DurableRuntime" + raise ValueError(msg) log.info("Using A2A task store provider '{}'", type(runtime.task_store_provider).__name__) base_url = (base_url or get_a2a_base_url()).rstrip("/") if "//0.0.0.0" in base_url or "//[::]" in base_url: @@ -123,7 +151,7 @@ def create_a2a_app(*, base_url: str | None = None, debug: bool = False, runtime: base_url, ) - agent_names, mounts = _create_agent_mounts(base_url, runtime) + agent_names, mounts = _create_agent_mounts(base_url, runtime, durable_runtime) if not agent_names: log.warning( @@ -158,6 +186,8 @@ async def lifespan(app: Starlette) -> AsyncIterator[None]: # noqa: ARG001 await durable_runtime.close() app = Starlette(debug=debug, routes=[Route("/status", endpoint=handle_status), *mounts], lifespan=lifespan) + app.state.durable_runtime = durable_runtime + app.state.a2a_runtime = runtime log.debug("Created A2A Starlette app with {} mounted agent(s): {}", len(agent_names), agent_names) configure_tracing() diff --git a/src/hayhooks/server/a2a/executor.py b/src/hayhooks/server/a2a/executor.py index 996dc523..a298ff53 100644 --- a/src/hayhooks/server/a2a/executor.py +++ b/src/hayhooks/server/a2a/executor.py @@ -12,7 +12,7 @@ from haystack.dataclasses import StreamingChunk from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.a2a.durable_executor import DurableAgentExecutor from hayhooks.server.a2a.imports import ( AgentExecutor, @@ -146,12 +146,16 @@ def create_agent_executor( pipeline_name: str, *, task_store: Any | None = None, + durable_runtime: DurableRuntime | None = None, ) -> AgentExecutor: """Select a managed durable Agent or chat-compatible executor.""" if durable_authoring_mode(wrapper) is DurableAuthoringMode.MANAGED_AGENT: if task_store is None: msg = "A durable A2A Agent requires an A2A task store" raise RuntimeError(msg) + if durable_runtime is None: + msg = "A durable A2A Agent requires an application-owned DurableRuntime" + raise RuntimeError(msg) # Constructing the deployment here validates the Haystack v3 Agent and the # durable definition before the Agent Card is exposed. deployment = durable_runtime.deployment(pipeline_name, wrapper) diff --git a/src/hayhooks/server/a2a/runtime.py b/src/hayhooks/server/a2a/runtime.py index c678d9f1..77de927b 100644 --- a/src/hayhooks/server/a2a/runtime.py +++ b/src/hayhooks/server/a2a/runtime.py @@ -3,7 +3,7 @@ from typing import Any from hayhooks.a2a import TaskStoreProvider -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.a2a.durable_executor import DurableAgentExecutor from hayhooks.server.a2a.imports import ( InMemoryTaskStore, @@ -104,6 +104,8 @@ class A2ARuntime: def __init__( self, task_store_provider: TaskStoreProvider | None = None, + *, + durable_runtime: DurableRuntime | None = None, ) -> None: self.task_store_provider = task_store_provider or InMemoryTaskStoreProvider() self._executors: list[DurableAgentExecutor] = [] @@ -111,6 +113,7 @@ def __init__( self._task_stores: list[TaskStore] = [] self._maintenance_task: asyncio.Task[None] | None = None self._started = False + self.durable_runtime = durable_runtime def register_agent_executor(self, executor: Any) -> None: if isinstance(executor, DurableAgentExecutor): @@ -204,10 +207,11 @@ async def _provider_health(self) -> dict[str, Any]: "error": type(error).__name__, } - @staticmethod - async def _durable_health() -> dict[str, Any]: + async def _durable_health(self) -> dict[str, Any]: try: - return await durable_runtime.health() + if self.durable_runtime is None: + return {"healthy": False, "error": "DurableRuntimeUnavailable"} + return await self.durable_runtime.health() except asyncio.CancelledError: raise except Exception as error: diff --git a/src/hayhooks/server/app.py b/src/hayhooks/server/app.py index 9a29a54b..387442cb 100644 --- a/src/hayhooks/server/app.py +++ b/src/hayhooks/server/app.py @@ -20,7 +20,7 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.logger import RequestIdMiddleware, intercept_stdlib_logging, log, log_elapsed from hayhooks.server.routers import ( dashboard_router, @@ -248,12 +248,12 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: try: if settings.pipelines_dir: deploy_pipelines(app, settings.pipelines_dir) - await durable_runtime.start() + await app.state.durable_runtime.start() broadcaster.set_loop(asyncio.get_running_loop()) yield finally: try: - await durable_runtime.close() + await app.state.durable_runtime.close() finally: broadcaster.clear_loop() @@ -280,7 +280,7 @@ def get_package_version() -> str: return "0.0.0" -def create_app() -> FastAPI: +def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: """ Create and configure a FastAPI application. @@ -311,6 +311,7 @@ def create_app() -> FastAPI: app_params["root_path"] = root_path app = FastAPI(**app_params) + app.state.durable_runtime = durable_runtime or DurableRuntime(app_settings=settings) configure_tracing() app.add_middleware(RequestIdMiddleware) diff --git a/src/hayhooks/server/durable/routes.py b/src/hayhooks/server/durable/routes.py index 4823a18b..3f247c0c 100644 --- a/src/hayhooks/server/durable/routes.py +++ b/src/hayhooks/server/durable/routes.py @@ -1,28 +1,13 @@ -"""Typed REST resources that project the durable execution record.""" +"""Hayhooks-specific composition for the public durable FastAPI adapter.""" from __future__ import annotations -import inspect -import re -from typing import Annotated, Any, cast -from urllib.parse import quote - -from fastapi import Body, FastAPI, Header, HTTPException, Path, Request, status -from fastapi.responses import Response +from fastapi import FastAPI, HTTPException, Request, status from fastapi.routing import APIRoute -from pydantic import ValidationError, create_model -from hayhooks.durable import ExecutionResult -from hayhooks.durable.engine import RUN_ID_PATTERN -from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError -from hayhooks.durable.runtime import ( - DefinitionRevisionConflictError, - DurableDeployment, - IdempotencyConflictError, - durable_runtime, -) -from hayhooks.server.pipelines.registry import registry -from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.durable.fastapi import OwnerIdDependency, create_durable_router +from hayhooks.durable.runtime import DurableDeployment +from hayhooks.settings import settings DURABLE_ROUTE_SUFFIXES = ( "/run-durable", @@ -30,238 +15,66 @@ "/executions/{execution_id}/cancel", "/executions/{execution_id}/resume", ) -_IDEMPOTENCY_KEY_PATTERN = re.compile(rf"^{RUN_ID_PATTERN}$") -_MAX_DURABLE_OWNER_LENGTH = 512 -_MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH = 63 -ExecutionId = Annotated[str, Path(pattern=rf"^{RUN_ID_PATTERN}$", min_length=1, max_length=128)] - - -def _execution_links(pipeline_name: str, execution_id: str) -> dict[str, str]: - root = f"/{pipeline_name}/executions/{quote(execution_id, safe='-._~')}" - return {"self": root, "cancel": f"{root}/cancel", "resume": f"{root}/resume"} - +_MAX_OWNER_LENGTH = 512 -def _execution_result( - deployment: DurableDeployment, - record: Any, - *, - response_model: type[ExecutionResult] = ExecutionResult, -) -> ExecutionResult: - return response_model.model_validate(record.safe_view(links=_execution_links(deployment.name, record.execution_id))) - - -def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResult]: - if deployment.result_type is Any: - return ExecutionResult - return create_model( - f"{deployment.name.title().replace('-', '').replace('_', '')}ExecutionResult", - __base__=ExecutionResult, - result=(deployment.result_type | None, None), - ) - -def _durable_owner(request: Request, deployment: DurableDeployment) -> tuple[str | None, bool]: - header = deployment.app_settings.durable_trusted_owner_header.strip() +def _trusted_owner_dependency() -> OwnerIdDependency | None: + header = settings.durable_trusted_owner_header.strip() if not header: - return None, False - owner = request.headers.get(header) - if not owner: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=f"Authenticated owner header '{header}' is required", - ) - if len(owner) > _MAX_DURABLE_OWNER_LENGTH: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Authenticated owner header '{header}' exceeds 512 characters", - ) - return owner, True + return None + def trusted_owner(request: Request) -> str: + owner = request.headers.get(header) + if not owner: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=f"Authenticated owner header '{header}' is required", + ) + if len(owner) > _MAX_OWNER_LENGTH: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Authenticated owner header '{header}' exceeds 512 characters", + ) + return owner -def _remove_pipeline_route(app: FastAPI, path: str, method: str) -> None: - for route in list(app.routes): - if isinstance(route, APIRoute) and route.path == path and route.methods is not None and method in route.methods: - app.routes.remove(route) + return trusted_owner -def _remove_durable_api_routes(app: FastAPI, pipeline_name: str) -> None: +def remove_durable_api_routes(app: FastAPI, pipeline_name: str) -> None: root = f"/{pipeline_name}" durable_paths = {f"{root}{suffix}" for suffix in DURABLE_ROUTE_SUFFIXES} - app.routes[:] = [route for route in app.routes if not (isinstance(route, APIRoute) and route.path in durable_paths)] + app.routes[:] = [ + route + for route in app.routes + if not ( + (isinstance(route, APIRoute) and route.path in durable_paths) + or getattr(getattr(route, "original_router", None), "_hayhooks_durable_pipeline", None) == pipeline_name + ) + ] -def add_durable_api_routes( # noqa: C901, PLR0915 - route-local handlers share generated models +def add_durable_api_routes( app: FastAPI, pipeline_name: str, - pipeline_wrapper: BasePipelineWrapper, + deployment: DurableDeployment | None, *, - deployment: DurableDeployment | None = None, _defer_openapi_rebuild: bool, ) -> None: - """Register typed durable submission and control resources when opted in.""" - _remove_durable_api_routes(app, pipeline_name) - if not durable_runtime.has_capability(pipeline_wrapper): - if not _defer_openapi_rebuild: - app.openapi_schema = None - app.setup() - return - deployment = deployment or durable_runtime.deployment(pipeline_name, pipeline_wrapper) - request_model = deployment.request_type - response_model = _durable_response_model(deployment) - root = f"/{pipeline_name}" - - async def get_execution(execution_id: str, owner_id: str | None, enforce_owner: bool) -> Any: - return await deployment.get( - execution_id, - owner_id=owner_id, - enforce_owner=enforce_owner, - allow_revision_mismatch=True, + """Replace one pipeline's durable route family with the public adapter.""" + remove_durable_api_routes(app, pipeline_name) + if deployment is not None: + router = create_durable_router( + deployment, + owner_id_dependency=_trusted_owner_dependency(), ) - - async def submit( - run_req: Any, - response: Response, - request: Request, - idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), - ) -> ExecutionResult: - owner_id, owner_scoped = _durable_owner(request, deployment) - if idempotency_key is not None and _IDEMPOTENCY_KEY_PATTERN.fullmatch(idempotency_key) is None: - raise HTTPException( - status_code=422, - detail="Idempotency-Key must contain 1-128 letters, digits, underscores, or hyphens", - ) - if ( - owner_scoped - and idempotency_key is not None - and len(idempotency_key) > _MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH - ): - raise HTTPException( - status_code=422, - detail="Idempotency-Key must be at most 63 characters when owner scoping is enabled", - ) - try: - created, record = await deployment.submit( - run_req.model_dump(mode="json"), - execution_id=idempotency_key, - owner_id=owner_id, - ) - except IdempotencyConflictError as error: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error - except DefinitionRevisionConflictError as error: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error - except (ValidationError, ValueError) as error: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error - except ExecutionAdmissionError as error: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail=str(error), - headers={"Retry-After": str(error.retry_after_seconds)}, - ) from error - except (ExecutionStoreError, RuntimeError) as error: - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error - response.status_code = status.HTTP_200_OK if not created and record.terminal else status.HTTP_202_ACCEPTED - response.headers["Location"] = _execution_links(deployment.name, record.execution_id)["self"] - if not created: - response.headers["Idempotent-Replay"] = "true" - return _execution_result(deployment, record, response_model=response_model) - - # FastAPI consumes the runtime annotation; keep the Python type checker on - # the stable BaseModel boundary while preserving the generated request schema. - submit.__annotations__["run_req"] = request_model - - async def inspect_execution(execution_id: ExecutionId, request: Request) -> ExecutionResult: - try: - owner_id, enforce_owner = _durable_owner(request, deployment) - record = await get_execution(execution_id, owner_id, enforce_owner) - return _execution_result(deployment, record) - except KeyError as error: - raise HTTPException(status_code=404, detail="Execution not found") from error - except ExecutionStoreError as error: - raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error - - async def cancel_execution(execution_id: ExecutionId, response: Response, request: Request) -> ExecutionResult: - try: - owner_id, enforce_owner = _durable_owner(request, deployment) - accepted = await deployment.request_cancel( - execution_id, - owner_id=owner_id, - enforce_owner=enforce_owner, - ) - record = await get_execution(execution_id, owner_id, enforce_owner) - except KeyError as error: - raise HTTPException(status_code=404, detail="Execution not found") from error - except ExecutionStoreError as error: - raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error - response.status_code = status.HTTP_202_ACCEPTED if accepted else status.HTTP_200_OK - return _execution_result(deployment, record) - - async def resume_execution( - execution_id: ExecutionId, - response: Response, - request: Request, - update: Any = Body(default=None), # noqa: B008 - ) -> ExecutionResult: - try: - owner_id, enforce_owner = _durable_owner(request, deployment) - resumed = await deployment.resume( - execution_id, - update, - owner_id=owner_id, - enforce_owner=enforce_owner, - ) - if not resumed: - raise HTTPException(status_code=409, detail="Execution is not waiting") - record = await get_execution(execution_id, owner_id, enforce_owner) - except KeyError as error: - raise HTTPException(status_code=404, detail="Execution not found") from error - except DefinitionRevisionConflictError as error: - raise HTTPException(status_code=409, detail=str(error)) from error - except (ValidationError, ValueError) as error: - raise HTTPException(status_code=422, detail=str(error)) from error - except ExecutionStoreError as error: - raise HTTPException(status_code=503, detail="Durable execution store is unavailable") from error - response.status_code = status.HTTP_202_ACCEPTED - return _execution_result(deployment, record) - - if deployment.resume_type is not None: - resume_execution.__annotations__["update"] = deployment.resume_type - signature = inspect.signature(resume_execution) - update_parameter = signature.parameters["update"].replace( - annotation=deployment.resume_type, - default=Body(), - ) - cast(Any, resume_execution).__signature__ = signature.replace( - parameters=[ - update_parameter if parameter.name == "update" else parameter - for parameter in signature.parameters.values() - ] + router._hayhooks_durable_pipeline = pipeline_name # ty: ignore[unresolved-attribute] + app.include_router( + router, + prefix=f"/{pipeline_name}", ) - - routes = [ - (f"{root}/run-durable", submit, ["POST"], f"{pipeline_name}_run_durable"), - (f"{root}/executions/{{execution_id}}", inspect_execution, ["GET"], f"{pipeline_name}_execution"), - (f"{root}/executions/{{execution_id}}/cancel", cancel_execution, ["POST"], f"{pipeline_name}_cancel"), - (f"{root}/executions/{{execution_id}}/resume", resume_execution, ["POST"], f"{pipeline_name}_resume"), - ] - for path, endpoint, methods, name in routes: - _remove_pipeline_route(app, path, methods[0]) - app.add_api_route( - path, - endpoint, - methods=methods, - name=name, - response_model=response_model if endpoint is submit else ExecutionResult, - tags=["durable executions"], - status_code=status.HTTP_202_ACCEPTED if methods == ["POST"] else status.HTTP_200_OK, - ) - - registry.update_metadata( - pipeline_name, - {"durable_request_model": request_model, "durable_response_model": response_model}, - ) if not _defer_openapi_rebuild: app.openapi_schema = None app.setup() -__all__ = ["DURABLE_ROUTE_SUFFIXES", "add_durable_api_routes"] +__all__ = ["DURABLE_ROUTE_SUFFIXES", "add_durable_api_routes", "remove_durable_api_routes"] diff --git a/src/hayhooks/server/routers/status.py b/src/hayhooks/server/routers/status.py index 15196e25..b8b8f351 100644 --- a/src/hayhooks/server/routers/status.py +++ b/src/hayhooks/server/routers/status.py @@ -1,7 +1,6 @@ -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Request from pydantic import BaseModel, Field -from hayhooks.durable.runtime import durable_runtime from hayhooks.server.pipelines.registry import registry router = APIRouter() @@ -32,9 +31,9 @@ class PipelineStatusResponse(BaseModel): summary="Get status of all pipelines", description="Returns the system status and a list of all available pipelines.", ) -async def status_all() -> StatusResponse: +async def status_all(request: Request) -> StatusResponse: pipelines = registry.get_names() - durable_health = await durable_runtime.health() + durable_health = await request.app.state.durable_runtime.health() return StatusResponse(status="Up!", pipelines=pipelines, durable=durable_health) @@ -46,10 +45,10 @@ async def status_all() -> StatusResponse: summary="Get status of a specific pipeline", description="Returns the status of a specific pipeline. Returns 404 if the pipeline doesn't exist.", ) -async def status(pipeline_name: str) -> PipelineStatusResponse: +async def status(request: Request, pipeline_name: str) -> PipelineStatusResponse: if pipeline_name not in registry.get_names(): raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") - deployment = durable_runtime.current_deployment(pipeline_name) + deployment = request.app.state.durable_runtime.current_deployment(pipeline_name) if deployment is not None and not deployment.manager.health["healthy"]: raise HTTPException(status_code=503, detail=f"Pipeline '{pipeline_name}' has no live durable worker slots") return PipelineStatusResponse(status="Up!", pipeline=pipeline_name) diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index 6a1a1d03..8cf6a5a1 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -20,9 +20,10 @@ from fastapi.routing import APIRoute from pydantic import BaseModel -from hayhooks.durable.runtime import DurableDeployment, durable_runtime -from hayhooks.server.durable.routes import DURABLE_ROUTE_SUFFIXES as _DURABLE_ROUTE_SUFFIXES +from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode +from hayhooks.durable.runtime import DurableDeployment, DurableRuntime from hayhooks.server.durable.routes import add_durable_api_routes as _add_durable_api_routes +from hayhooks.server.durable.routes import remove_durable_api_routes as _remove_durable_api_routes from hayhooks.server.exceptions import PipelineAlreadyExistsError, PipelineFilesError from hayhooks.server.logger import log, log_elapsed from hayhooks.server.pipelines.models import ( @@ -61,6 +62,27 @@ _deployments_in_progress: set[str] = set() +def _app_runtime(app: FastAPI | None, runtime: DurableRuntime | None = None) -> DurableRuntime | None: + if runtime is not None: + return runtime + state = getattr(app, "state", None) if app is not None else None + candidate = getattr(state, "durable_runtime", None) + return candidate if isinstance(candidate, DurableRuntime) else None + + +def _deployment_candidate( + runtime: DurableRuntime | None, + name: str, + wrapper: BasePipelineWrapper, +) -> DurableDeployment | None: + if runtime is not None: + return runtime.create_deployment(name, wrapper) + if durable_authoring_mode(wrapper) is not DurableAuthoringMode.NONE: + msg = "Durable pipeline deployment requires an application-owned DurableRuntime" + raise RuntimeError(msg) + return None + + async def _offload(func: Callable, **kwargs: Any) -> Any: """Run blocking pipeline preparation outside the event loop.""" return await asyncio.to_thread(func, **kwargs) @@ -78,13 +100,14 @@ async def _deployment_lock(lock: threading.Lock): class _DeploymentSnapshot: """Rollback state captured before preparation mutates files or loaded modules.""" - def __init__(self, pipeline_name: str, app: FastAPI | None) -> None: + def __init__(self, pipeline_name: str, app: FastAPI | None, runtime: DurableRuntime | None) -> None: self.pipeline_name = pipeline_name self.app = app self.wrapper = registry.get(pipeline_name) metadata = registry.get_metadata(pipeline_name) self.metadata = dict(metadata) if metadata is not None else None - self.deployment = durable_runtime.current_deployment(pipeline_name) + self.runtime = runtime + self.deployment = runtime.current_deployment(pipeline_name) if runtime is not None else None self.routes = list(app.routes) if app is not None else None self.openapi_schema = app.openapi_schema if app is not None else None self.modules = { @@ -109,14 +132,20 @@ def __init__(self, pipeline_name: str, app: FastAPI | None) -> None: shutil.copy2(source, self.backup_dir / f"pipeline{extension}") @classmethod - def capture(cls, pipeline_name: str, app: FastAPI | None) -> "_DeploymentSnapshot": - return cls(pipeline_name, app) + def capture( + cls, + pipeline_name: str, + app: FastAPI | None, + runtime: DurableRuntime | None, + ) -> "_DeploymentSnapshot": + return cls(pipeline_name, app, runtime) def restore_publication(self) -> None: registry.remove(self.pipeline_name) if self.wrapper is not None: registry.add(self.pipeline_name, self.wrapper, metadata=dict(self.metadata or {})) - durable_runtime.install_deployment(self.pipeline_name, self.deployment) + if self.runtime is not None: + self.runtime.install_deployment(self.pipeline_name, self.deployment) if self.app is not None and self.routes is not None: self.app.routes[:] = self.routes self.app.openapi_schema = self.openapi_schema @@ -153,11 +182,13 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 prepare: Callable[[], Awaitable[PreparedPipeline]], *, app: FastAPI | None, + durable_runtime: DurableRuntime | None, overwrite: bool, remove_files_before_prepare: bool, cleanup_files_on_overwrite: bool, ) -> dict[str, str]: """Prepare independently, then atomically publish or restore one pipeline.""" + durable_runtime = _app_runtime(app, durable_runtime) dlog = log.bind( pipeline_name=pipeline_name, overwrite=overwrite, @@ -180,7 +211,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 if pipeline_name in _deployments_in_progress: msg = f"Pipeline '{pipeline_name}' is already being deployed" raise PipelineAlreadyExistsError(msg) - snapshot = _DeploymentSnapshot.capture(pipeline_name, app) + snapshot = _DeploymentSnapshot.capture(pipeline_name, app, durable_runtime) if snapshot.wrapper is not None and not overwrite: msg = f"Pipeline '{pipeline_name}' already exists" raise PipelineAlreadyExistsError(msg) @@ -206,8 +237,8 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 if remove_files_before_prepare: remove_pipeline_files(pipeline_name, settings.pipelines_dir) prepared = await prepare() - candidate = durable_runtime.create_deployment(prepared.name, prepared.wrapper) - if candidate is not None and durable_runtime.started: + candidate = _deployment_candidate(durable_runtime, prepared.name, prepared.wrapper) + if candidate is not None and durable_runtime is not None and durable_runtime.started: await candidate.prepare() dlog.bind(durable=candidate is not None).debug("Prepared pipeline deployment candidate") @@ -220,9 +251,9 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 overwrite=overwrite, cleanup_files_on_overwrite=cleanup_files_on_overwrite, _durable_deployment=candidate, + durable_runtime=durable_runtime, ) - durable_runtime.install_deployment(prepared.name, candidate) - if candidate is not None and durable_runtime.started: + if candidate is not None and durable_runtime is not None and durable_runtime.started: candidate.activate() dlog.debug("Published pipeline deployment") return result @@ -239,7 +270,12 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 if publication_started: snapshot.restore_publication() finally: - if old_quiesced and snapshot.deployment is not None and durable_runtime.started: + if ( + old_quiesced + and snapshot.deployment is not None + and durable_runtime is not None + and durable_runtime.started + ): await snapshot.deployment.start() raise finally: @@ -250,12 +286,14 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 snapshot.cleanup() -async def deploy_pipeline_yaml_async( +async def deploy_pipeline_yaml_async( # noqa: PLR0913 - stable public deployment API pipeline_name: str, source_code: str, app: FastAPI | None = None, overwrite: bool = False, options: dict[str, Any] | None = None, + *, + durable_runtime: DurableRuntime | None = None, ) -> dict[str, str]: """ Async wrapper that offloads ``deploy_pipeline_yaml`` off the event loop. @@ -273,18 +311,21 @@ async def deploy_pipeline_yaml_async( options=options, ), app=app, + durable_runtime=durable_runtime, overwrite=overwrite, remove_files_before_prepare=overwrite and save_file, cleanup_files_on_overwrite=overwrite and not save_file, ) -async def deploy_pipeline_files_async( +async def deploy_pipeline_files_async( # noqa: PLR0913 - stable public deployment API pipeline_name: str, files: dict[str, str], app: FastAPI | None = None, save_files: bool = True, overwrite: bool = False, + *, + durable_runtime: DurableRuntime | None = None, ) -> dict[str, str]: """Async wrapper that offloads ``deploy_pipeline_files`` off the event loop.""" return await _deploy_prepared_pipeline_async( @@ -296,6 +337,7 @@ async def deploy_pipeline_files_async( save_files=save_files, ), app=app, + durable_runtime=durable_runtime, overwrite=overwrite, remove_files_before_prepare=overwrite and save_files, cleanup_files_on_overwrite=overwrite and not save_files, @@ -305,8 +347,11 @@ async def deploy_pipeline_files_async( async def undeploy_pipeline_async( pipeline_name: str, app: FastAPI | None = None, + *, + durable_runtime: DurableRuntime | None = None, ) -> None: """Atomically unpublish a pipeline before stopping its owned resources.""" + durable_runtime = _app_runtime(app, durable_runtime) policy_lock = ( _deployment_lock(_deployment_serial_lock) if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED @@ -317,7 +362,7 @@ async def undeploy_pipeline_async( raise HTTPException(status_code=409, detail=f"Pipeline '{pipeline_name}' is being deployed") if registry.get(pipeline_name) is None: raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") - deployment = durable_runtime.current_deployment(pipeline_name) + deployment = durable_runtime.current_deployment(pipeline_name) if durable_runtime is not None else None if deployment is not None: await deployment.quiesce() try: @@ -331,12 +376,13 @@ async def undeploy_pipeline_async( ), ) except BaseException: - if durable_runtime.started: + if durable_runtime is not None and durable_runtime.started: await deployment.start() raise await _offload(undeploy_pipeline, pipeline_name=pipeline_name, app=app) - durable_runtime.install_deployment(pipeline_name, None) + if durable_runtime is not None: + durable_runtime.install_deployment(pipeline_name, None) async def _durable_nonterminal_count(deployment: DurableDeployment) -> int: @@ -706,6 +752,12 @@ def add_pipeline_api_route( - Updates registry metadata with request/response models and file requirement flag """ clog = log.bind(pipeline_name=pipeline_name) + deployment = _durable_deployment + if deployment is None: + runtime = _app_runtime(app) + deployment = _deployment_candidate(runtime, pipeline_name, pipeline_wrapper) + if runtime is not None: + runtime.install_deployment(pipeline_name, deployment) # Determine which run_api method to use (prefer async if available) if pipeline_wrapper._is_run_api_async_implemented: @@ -725,8 +777,7 @@ def add_pipeline_api_route( _add_durable_api_routes( app, pipeline_name, - pipeline_wrapper, - deployment=_durable_deployment, + deployment, _defer_openapi_rebuild=_defer_openapi_rebuild, ) return @@ -775,8 +826,7 @@ def add_pipeline_api_route( _add_durable_api_routes( app, pipeline_name, - pipeline_wrapper, - deployment=_durable_deployment, + deployment, _defer_openapi_rebuild=True, ) @@ -983,6 +1033,7 @@ def commit_prepared_pipeline( _defer_openapi_rebuild: bool = False, cleanup_files_on_overwrite: bool = True, _durable_deployment: DurableDeployment | None = None, + durable_runtime: DurableRuntime | None = None, ) -> dict[str, str]: """ Commit a prepared pipeline to the registry and (optionally) add its route. @@ -997,6 +1048,10 @@ def commit_prepared_pipeline( _defer_openapi_rebuild: Forwarded to route registration. cleanup_files_on_overwrite: If ``True``, remove persisted files when replacing an existing pipeline. """ + durable_runtime = _app_runtime(app, durable_runtime) + candidate = _durable_deployment + if candidate is None: + candidate = _deployment_candidate(durable_runtime, prepared.name, prepared.wrapper) with trace_operation( SPAN_PIPELINE_DEPLOY_COMMIT, tags=build_trace_tags( @@ -1019,17 +1074,20 @@ def commit_prepared_pipeline( if cleanup_files_on_overwrite: remove_pipeline_files(prepared.name, settings.pipelines_dir) - return _register_prepared_pipeline( + result = _register_prepared_pipeline( pipeline_name=prepared.name, pipeline_wrapper=prepared.wrapper, app=app, extra_metadata=prepared.extra_metadata, _defer_openapi_rebuild=_defer_openapi_rebuild, - _durable_deployment=_durable_deployment, + _durable_deployment=candidate, ) + if durable_runtime is not None: + durable_runtime.install_deployment(prepared.name, candidate) + return result -def deploy_pipeline_files( +def deploy_pipeline_files( # noqa: PLR0913 - stable public deployment API pipeline_name: str, files: dict[str, str], app: FastAPI | None = None, @@ -1037,6 +1095,7 @@ def deploy_pipeline_files( overwrite: bool = False, *, _defer_openapi_rebuild: bool = False, + durable_runtime: DurableRuntime | None = None, ) -> dict[str, str]: """ Deploy a pipeline from Python files (pipeline_wrapper.py and other files). @@ -1085,10 +1144,11 @@ def deploy_pipeline_files( overwrite=overwrite, _defer_openapi_rebuild=_defer_openapi_rebuild, cleanup_files_on_overwrite=cleanup_files_on_overwrite, + durable_runtime=durable_runtime, ) -def deploy_pipeline_yaml( +def deploy_pipeline_yaml( # noqa: PLR0913 - stable public deployment API pipeline_name: str, source_code: str, app: FastAPI | None = None, @@ -1096,6 +1156,7 @@ def deploy_pipeline_yaml( options: dict[str, Any] | None = None, *, _defer_openapi_rebuild: bool = False, + durable_runtime: DurableRuntime | None = None, ) -> dict[str, str]: """ Deploy a YAML pipeline to the FastAPI application with IO declared in the YAML. @@ -1147,6 +1208,7 @@ def deploy_pipeline_yaml( overwrite=overwrite, _defer_openapi_rebuild=_defer_openapi_rebuild, cleanup_files_on_overwrite=cleanup_files_on_overwrite, + durable_runtime=durable_runtime, ) @@ -1180,7 +1242,7 @@ def read_pipeline_files_from_dir(dir_path: Path) -> dict[str, str]: return files -def deploy_pipelines() -> None: +def deploy_pipelines(*, durable_runtime: DurableRuntime | None = None) -> None: """Deploy pipelines from the configured directory""" # Imported here to avoid a circular import (hayhooks.server.app imports this module) from hayhooks.server.app import init_pipeline_dir @@ -1201,6 +1263,7 @@ def deploy_pipelines() -> None: pipeline_name=pipeline_dir.name, files=read_pipeline_files_from_dir(pipeline_dir), save_files=False, # Files already exist on disk + durable_runtime=durable_runtime, ) except Exception as e: log.warning("Skipping pipeline directory '{}': {}", pipeline_dir, e) @@ -1241,13 +1304,11 @@ def undeploy_pipeline(pipeline_name: str, app: FastAPI | None = None) -> None: unload_pipeline_modules(pipeline_name) if app: - route_paths = { - f"/{pipeline_name}/run", - *(f"/{pipeline_name}{suffix}" for suffix in _DURABLE_ROUTE_SUFFIXES), - } + route_paths = {f"/{pipeline_name}/run"} app.routes[:] = [ route for route in app.routes if not (isinstance(route, APIRoute) and route.path in route_paths) ] + _remove_durable_api_routes(app, pipeline_name) # Invalidate OpenAPI cache app.openapi_schema = None diff --git a/src/hayhooks/server/utils/mcp_utils.py b/src/hayhooks/server/utils/mcp_utils.py index 04fd41ed..009058f1 100644 --- a/src/hayhooks/server/utils/mcp_utils.py +++ b/src/hayhooks/server/utils/mcp_utils.py @@ -12,7 +12,7 @@ from starlette.routing import Mount, Route from starlette.types import Receive, Scope, Send -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.logger import log from hayhooks.server.pipelines.registry import registry from hayhooks.server.routers.deploy import PipelineFilesRequest @@ -158,7 +158,9 @@ async def notify_client(server: "Server") -> None: await server.request_context.session.send_tool_list_changed() -async def _handle_deploy_pipeline(arguments: dict[str, Any], span: Any) -> list["TextContent"]: +async def _handle_deploy_pipeline( + arguments: dict[str, Any], span: Any, durable_runtime: DurableRuntime +) -> list["TextContent"]: span.set_tag("hayhooks.pipeline.name", arguments.get("name")) result = await deploy_pipeline_files_async( pipeline_name=arguments["name"], @@ -166,34 +168,41 @@ async def _handle_deploy_pipeline(arguments: dict[str, Any], span: Any) -> list[ app=None, save_files=arguments["save_files"], overwrite=arguments["overwrite"], + durable_runtime=durable_runtime, ) return [TextContent(type="text", text=f"Pipeline '{result['name']}' deployed successfully")] -async def _handle_get_all_pipeline_statuses(_arguments: dict[str, Any], _span: Any) -> list["TextContent"]: +async def _handle_get_all_pipeline_statuses( + _arguments: dict[str, Any], _span: Any, _durable_runtime: DurableRuntime +) -> list["TextContent"]: pipelines_str = "\n".join(registry.get_names()) return [TextContent(type="text", text=f"Available pipelines:\n{pipelines_str}")] -async def _handle_get_pipeline_status(arguments: dict[str, Any], span: Any) -> list["TextContent"]: +async def _handle_get_pipeline_status( + arguments: dict[str, Any], span: Any, _durable_runtime: DurableRuntime +) -> list["TextContent"]: pipeline_name = arguments["pipeline_name"] span.set_tag("hayhooks.pipeline.name", pipeline_name) is_deployed = pipeline_name in registry.get_names() return [TextContent(type="text", text=f"Pipeline '{pipeline_name}' is deployed: {is_deployed}")] -async def _handle_undeploy_pipeline(arguments: dict[str, Any], span: Any) -> list["TextContent"]: +async def _handle_undeploy_pipeline( + arguments: dict[str, Any], span: Any, durable_runtime: DurableRuntime +) -> list["TextContent"]: pipeline_name = arguments["pipeline_name"] span.set_tag("hayhooks.pipeline.name", pipeline_name) # app=None: the MCP server doesn't own FastAPI routes - await undeploy_pipeline_async(pipeline_name=pipeline_name) + await undeploy_pipeline_async(pipeline_name=pipeline_name, durable_runtime=durable_runtime) return [TextContent(type="text", text=f"Pipeline '{pipeline_name}' undeployed")] # Core tools that trigger a ``tools/list_changed`` notification after execution. _MUTATING_CORE_TOOLS: frozenset[str] = frozenset({CoreTools.DEPLOY_PIPELINE.value, CoreTools.UNDEPLOY_PIPELINE.value}) -_CoreToolHandler = Callable[[dict[str, Any], Any], Awaitable[list["TextContent"]]] +_CoreToolHandler = Callable[[dict[str, Any], Any, DurableRuntime], Awaitable[list["TextContent"]]] _CORE_TOOL_HANDLERS: dict[str, _CoreToolHandler] = { CoreTools.DEPLOY_PIPELINE.value: _handle_deploy_pipeline, @@ -215,10 +224,16 @@ async def _run_pipeline_tool(name: str, arguments: dict[str, Any], span: Any) -> raise Exception(msg) from exc -def create_mcp_server(name: str = "hayhooks-mcp-server") -> "Server": +def create_mcp_server( + name: str = "hayhooks-mcp-server", + *, + durable_runtime: DurableRuntime | None = None, +) -> "Server": mcp_import.check() + durable_runtime = durable_runtime or DurableRuntime(app_settings=settings) server: Server = Server(name) + server._hayhooks_durable_runtime = durable_runtime # ty: ignore[unresolved-attribute] @server.list_tools() async def list_tools() -> list[Tool]: @@ -261,7 +276,7 @@ async def call_tool(name: str, arguments: dict[str, Any]) -> list["TextContent"] ) as span: try: if handler is not None: - return await handler(arguments, span) + return await handler(arguments, span, durable_runtime) return await _run_pipeline_tool(name, arguments, span) except Exception as exc: msg = f"General unhandled error in call_tool for tool '{name}': {exc}" @@ -277,11 +292,23 @@ async def call_tool(name: str, arguments: dict[str, Any]) -> list["TextContent"] return server -def create_starlette_app(server: "Server", *, debug: bool = False, json_response: bool = False) -> "Starlette": +def create_starlette_app( + server: "Server", + *, + debug: bool = False, + json_response: bool = False, + durable_runtime: DurableRuntime | None = None, +) -> "Starlette": """ Create a Starlette app for the MCP server. """ mcp_import.check() + server_runtime = getattr(server, "_hayhooks_durable_runtime", None) + if durable_runtime is None: + durable_runtime = server_runtime or DurableRuntime(app_settings=settings) + elif server_runtime is not None and server_runtime is not durable_runtime: + msg = "MCP server and application must reference the same DurableRuntime" + raise ValueError(msg) # Setup the Streamable HTTP session manager session_manager = StreamableHTTPSessionManager( @@ -299,16 +326,16 @@ async def handle_streamable_http(scope: Scope, receive: Receive, send: Send) -> @asynccontextmanager async def lifespan(app: Starlette) -> AsyncIterator[None]: # noqa: ARG001 - async with session_manager.run(): - try: + try: + async with session_manager.run(): await durable_runtime.start() log.info("Hayhooks MCP server started") yield + finally: + try: + await durable_runtime.close() finally: - try: - await durable_runtime.close() - finally: - log.info("Hayhooks MCP server shutting down...") + log.info("Hayhooks MCP server shutting down...") async def handle_sse(request): async with sse.connect_sse(request.scope, request.receive, request._send) as streams: @@ -327,6 +354,7 @@ async def handle_status(request): # noqa: ARG001 ], lifespan=lifespan, ) + app.state.durable_runtime = durable_runtime configure_tracing() instrument_starlette_app(app) diff --git a/tests/test_cli.py b/tests/test_cli.py index 861bfeee..9fb325ed 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -148,15 +148,17 @@ def test_a2a_run_debug_enables_tracebacks(monkeypatch): from hayhooks.settings import settings calls = [] + runtimes = [] def fake_uvicorn_run(*args, **kwargs): calls.append((args, kwargs)) - def fake_create_a2a_app(*, debug: bool = False): + def fake_create_a2a_app(*, debug: bool = False, durable_runtime=None): + runtimes.append(durable_runtime) return object() - def fake_deploy_pipelines() -> None: - return None + def fake_deploy_pipelines(*, durable_runtime=None) -> None: + runtimes.append(durable_runtime) monkeypatch.setattr(uvicorn, "run", fake_uvicorn_run) monkeypatch.setattr(deploy_utils, "deploy_pipelines", fake_deploy_pipelines) @@ -167,6 +169,7 @@ def fake_deploy_pipelines() -> None: assert result.exit_code == 0, result.output assert settings.show_tracebacks is True assert calls, "uvicorn.run was not called" + assert runtimes[0] is runtimes[1] def test_a2a_run_sets_builtin_redis_task_store(monkeypatch): @@ -176,7 +179,7 @@ def test_a2a_run_sets_builtin_redis_task_store(monkeypatch): from hayhooks.server.utils import deploy_utils monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) result = runner.invoke( @@ -208,7 +211,7 @@ def test_a2a_run_sets_durable_execution_concurrency(monkeypatch): from hayhooks.server.utils import deploy_utils monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) result = runner.invoke( @@ -227,7 +230,7 @@ def test_a2a_run_sets_durable_execution_store_configuration(monkeypatch): from hayhooks.server.utils import deploy_utils monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda: None) + monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) result = runner.invoke( diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py index cbea8d51..493c3e4f 100644 --- a/tests/test_durable_a2a.py +++ b/tests/test_durable_a2a.py @@ -12,7 +12,8 @@ from hayhooks.a2a import A2APipelineWrapper, TaskStoreProvider from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStatus, ExecutionStoreError -from hayhooks.durable.runtime import execution_id_for +from hayhooks.durable.runtime import DurableRuntime, execution_id_for +from hayhooks.durable.store import InMemoryExecutionStoreProvider from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.a2a.durable_executor import DurableAgentExecutor, DurableTaskStore from hayhooks.server.a2a.imports import TaskStore, new_task_from_user_message, new_text_part @@ -164,10 +165,12 @@ def _http_app(store, deployment, monkeypatch): wrapper = _DurableHTTPWrapper() wrapper.setup() registry.add("durable-agent", wrapper, metadata={"description": "durable agent"}) - monkeypatch.setattr("hayhooks.durable.runtime.durable_runtime.deployment", lambda *_args: deployment) + durable_runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + monkeypatch.setattr(durable_runtime, "deployment", lambda *_args: deployment) return create_a2a_app( base_url="http://a2a-test:1418", runtime=A2ARuntime(task_store_provider=_HTTPStoreProvider(store)), + durable_runtime=durable_runtime, ) diff --git a/tests/test_durable_deployment_lifecycle.py b/tests/test_durable_deployment_lifecycle.py index 25097636..633c2d09 100644 --- a/tests/test_durable_deployment_lifecycle.py +++ b/tests/test_durable_deployment_lifecycle.py @@ -4,7 +4,7 @@ import pytest from fastapi.testclient import TestClient -from hayhooks.durable.runtime import durable_runtime +from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.app import create_app from hayhooks.server.pipelines.registry import registry from hayhooks.settings import settings @@ -139,7 +139,7 @@ def test_undeploy_removes_entire_durable_route_family() -> None: assert client.get("/job/executions/missing").status_code == 404 assert client.post("/job/executions/missing/cancel").status_code == 404 assert client.post("/job/executions/missing/resume").status_code == 404 - assert durable_runtime.current_deployment("job") is None + assert app.state.durable_runtime.current_deployment("job") is None def test_undeploy_refuses_to_strand_waiting_execution() -> None: @@ -207,7 +207,7 @@ def test_failed_durable_preflight_restarts_existing_deployment(monkeypatch, oper with TestClient(app) as client: source = _durable_source(field="value", increment=1, revision="first") assert _deploy(client, source).status_code == 200 - deployment = durable_runtime.current_deployment("job") + deployment = app.state.durable_runtime.current_deployment("job") assert deployment is not None async def fail_counts(): @@ -338,8 +338,7 @@ async def close(self): def test_store_initialization_failure_never_publishes_candidate(monkeypatch) -> None: - monkeypatch.setattr(durable_runtime, "_store_provider", _FailingProvider()) - app = create_app() + app = create_app(durable_runtime=DurableRuntime(_FailingProvider(), app_settings=settings)) with TestClient(app) as client: failed = _deploy(client, _durable_source(field="value", increment=1, revision="first")) diff --git a/tests/test_durable_execution.py b/tests/test_durable_execution.py index 110f765f..e3c92baf 100644 --- a/tests/test_durable_execution.py +++ b/tests/test_durable_execution.py @@ -4,6 +4,7 @@ import threading import time from pathlib import Path +from unittest.mock import AsyncMock, MagicMock import pytest from fastapi.testclient import TestClient @@ -26,13 +27,8 @@ ) from hayhooks.durable.context import execution_context_scope from hayhooks.durable.models import ExecutionCheckpoint, ExecutionKind, ExecutionRecord, ExecutionStatus -from hayhooks.durable.runtime import ( - DurableDeployment, - DurableRuntime, - _canonical_json, - _operation_fingerprint, - durable_runtime, -) +from hayhooks.durable.runtime import DurableDeployment, DurableRuntime, _canonical_json, _operation_fingerprint +from hayhooks.durable.settings import DurableSettings from hayhooks.durable.store import InMemoryExecutionStoreProvider from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.app import create_app @@ -46,7 +42,7 @@ load_pipeline_module, unload_pipeline_modules, ) -from hayhooks.settings import AppSettings, settings +from hayhooks.settings import settings pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -97,13 +93,75 @@ async def test_portable_runtime_starts_its_own_deployments() -> None: await runtime.close() -def _create_mcp_app(): - return create_starlette_app(create_mcp_server()) +async def test_rest_app_runtimes_are_isolated(monkeypatch) -> None: + monkeypatch.setattr(settings, "pipelines_dir", "") + provider_a = InMemoryExecutionStoreProvider() + provider_b = InMemoryExecutionStoreProvider() + close_a = AsyncMock(wraps=provider_a.close) + close_b = AsyncMock(wraps=provider_b.close) + monkeypatch.setattr(provider_a, "close", close_a) + monkeypatch.setattr(provider_b, "close", close_b) + runtime_a = DurableRuntime(provider_a) + runtime_b = DurableRuntime(provider_b) + wrapper = Wrapper() + wrapper.setup() + runtime_a.deployment("only-a", wrapper) + app_a = create_app(durable_runtime=runtime_a) + app_b = create_app(durable_runtime=runtime_b) + + assert app_a.state.durable_runtime is runtime_a + assert app_b.state.durable_runtime is runtime_b + assert runtime_b.current_deployment("only-a") is None + + lifespan_a = app_a.router.lifespan_context(app_a) + lifespan_b = app_b.router.lifespan_context(app_b) + await lifespan_a.__aenter__() + await lifespan_b.__aenter__() + try: + await lifespan_a.__aexit__(None, None, None) + assert not runtime_a.started + assert runtime_b.started + close_a.assert_awaited_once() + close_b.assert_not_awaited() + finally: + await lifespan_b.__aexit__(None, None, None) + close_b.assert_awaited_once() + + +def test_app_factories_retain_the_supplied_runtime() -> None: + rest_runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + a2a_runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + mcp_runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + + rest = create_app(durable_runtime=rest_runtime) + a2a = create_a2a_app(durable_runtime=a2a_runtime) + mcp_server = create_mcp_server(durable_runtime=mcp_runtime) + mcp = create_starlette_app(mcp_server, durable_runtime=mcp_runtime) + + assert rest.state.durable_runtime is rest_runtime + assert a2a.state.durable_runtime is a2a_runtime + assert mcp.state.durable_runtime is mcp_runtime + assert create_app().state.durable_runtime is not create_app().state.durable_runtime + + +def _create_rest_app(runtime: DurableRuntime): + return create_app(durable_runtime=runtime) + + +def _create_a2a_test_app(runtime: DurableRuntime): + return create_a2a_app(durable_runtime=runtime) + + +def _create_mcp_app(runtime: DurableRuntime): + return create_starlette_app( + create_mcp_server(durable_runtime=runtime), + durable_runtime=runtime, + ) @pytest.mark.parametrize( "app_factory", - [create_app, create_a2a_app, _create_mcp_app], + [_create_rest_app, _create_a2a_test_app, _create_mcp_app], ids=["rest", "a2a", "mcp"], ) async def test_app_lifespans_close_durable_runtime_when_start_fails(monkeypatch, app_factory) -> None: @@ -118,9 +176,10 @@ async def close() -> None: closed += 1 monkeypatch.setattr(settings, "pipelines_dir", "") - monkeypatch.setattr(durable_runtime, "start", fail_start) - monkeypatch.setattr(durable_runtime, "close", close) - app = app_factory() + runtime = DurableRuntime(InMemoryExecutionStoreProvider()) + monkeypatch.setattr(runtime, "start", fail_start) + monkeypatch.setattr(runtime, "close", close) + app = app_factory(runtime) with pytest.raises(Exception) as exc_info: async with app.router.lifespan_context(app): @@ -130,6 +189,25 @@ async def close() -> None: assert closed == 1 +async def test_mcp_lifespan_closes_durable_runtime_when_session_manager_start_fails(monkeypatch) -> None: + session_manager = MagicMock() + session_manager.run.return_value.__aenter__.side_effect = RuntimeError("MCP session manager startup failed") + provider = InMemoryExecutionStoreProvider() + close = AsyncMock(wraps=provider.close) + monkeypatch.setattr(provider, "close", close) + monkeypatch.setattr( + "hayhooks.server.utils.mcp_utils.StreamableHTTPSessionManager", lambda **_kwargs: session_manager + ) + runtime = DurableRuntime(provider) + app = _create_mcp_app(runtime) + + with pytest.raises(RuntimeError, match="MCP session manager startup failed"): + async with app.router.lifespan_context(app): + pass + + close.assert_awaited_once() + + def _checkpoint_test_tool(value: str) -> str: return value @@ -506,7 +584,7 @@ def test_durable_rest_can_inspect_and_cancel_an_execution_from_an_old_revision(m with TestClient(app) as client: submitted = client.post("/rolling/run-durable", json={"value": 1}) assert wrapper.started.wait(timeout=5) - deployment = durable_runtime.current_deployment("rolling") + deployment = app.state.durable_runtime.current_deployment("rolling") assert deployment is not None deployment.revision = "replacement" @@ -589,11 +667,11 @@ def test_durable_rest_bounds_owner_scoped_idempotency_keys(monkeypatch) -> None: def test_durable_rest_maps_oversized_validated_request_to_422(monkeypatch) -> None: - monkeypatch.setattr(settings, "durable_max_record_bytes", 5) + monkeypatch.setattr(settings, "durable_max_record_bytes", 1_024) app = _durable_app(monkeypatch, "job", Wrapper()) with TestClient(app) as client: - response = client.post("/job/run-durable", json={"value": 123}) + response = client.post("/job/run-durable", json={"value": int("1" * 1_100)}) assert response.status_code == 422 assert "durable execution limit" in response.json()["detail"] @@ -611,7 +689,7 @@ def test_durable_waiting_resume_is_typed_private_and_revision_safe(monkeypatch) "message": "Approve this job", "expected_input_schema": ResumeInput.model_json_schema(), } - deployment = durable_runtime.current_deployment("approval") + deployment = app.state.durable_runtime.current_deployment("approval") assert deployment is not None missing = client.post(f"{url}/resume") assert missing.status_code == 422 @@ -659,14 +737,14 @@ def test_durable_rest_enforces_configured_trusted_owner_header(monkeypatch) -> N assert "exceeds 512 characters" in oversized.json()["detail"] -def test_durable_rest_uses_the_deployments_owner_header_setting(monkeypatch) -> None: - monkeypatch.setattr(settings, "durable_trusted_owner_header", "") - app_settings = AppSettings(durable_store="memory", durable_trusted_owner_header="X-Embedded-Owner") - provider = InMemoryExecutionStoreProvider(app_settings=app_settings) +def test_durable_rest_uses_the_server_owner_header_setting(monkeypatch) -> None: + monkeypatch.setattr(settings, "durable_trusted_owner_header", "X-Embedded-Owner") + durable_settings = DurableSettings(durable_store="memory") + provider = InMemoryExecutionStoreProvider(durable_settings=durable_settings) wrapper = Wrapper() wrapper.setup() _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("embedded", wrapper, provider, app_settings=app_settings) + deployment = DurableDeployment("embedded", wrapper, provider, durable_settings=durable_settings) registry.add("embedded", wrapper) app = create_app() add_pipeline_api_route(app, "embedded", wrapper, _durable_deployment=deployment) @@ -723,8 +801,8 @@ def run_durable(self, context: DurableContext, request: Request) -> Result: assert release.wait(timeout=5) return Result(value=request.value) - monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) - provider = InMemoryExecutionStoreProvider() + durable_settings = DurableSettings(durable_store="memory", durable_shutdown_grace_period=0.001) + provider = InMemoryExecutionStoreProvider(durable_settings=durable_settings) wrapper = BlockingWrapper() wrapper.setup() _set_method_implementation_flags(wrapper) diff --git a/tests/test_durable_fastapi.py b/tests/test_durable_fastapi.py new file mode 100644 index 00000000..274ae20a --- /dev/null +++ b/tests/test_durable_fastapi.py @@ -0,0 +1,193 @@ +"""Standalone contract tests for the public FastAPI durable adapter.""" + +from __future__ import annotations + +import time +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Annotated + +import pytest +from fastapi import APIRouter, Depends, FastAPI, HTTPException, Request +from fastapi.responses import JSONResponse +from fastapi.testclient import TestClient +from haystack import Pipeline +from pydantic import BaseModel + +from hayhooks import BasePipelineWrapper +from hayhooks.durable import ( + DurableContext, + DurableRuntime, + DurableSettings, + InMemoryExecutionStoreProvider, + create_durable_router, +) + + +class JobRequest(BaseModel): + value: int + + +class JobResult(BaseModel): + value: int + owner_id: str | None + + +class ResumeInput(BaseModel): + approved: bool + + +class JobWrapper(BasePipelineWrapper): + durable_revision = "portable-job-v1" + durable_resume_model = ResumeInput + + def setup(self) -> None: + self.pipeline = Pipeline() + + async def run_durable_async(self, context: DurableContext, request: JobRequest) -> JobResult: + if request.value == 0 and context.resume_input is None: + await context.suspend({"kind": "approval", "private": "hidden"}) + approved = context.take_resume_input() + value = request.value if approved is None or ResumeInput.model_validate(approved).approved else -1 + return JobResult(value=value + 1, owner_id=context.owner_id) + + +def _require_principal(request: Request) -> str: + if request.headers.get("X-Deny"): + raise HTTPException(status_code=403, detail="Forbidden") + return request.headers.get("X-Owner", "") + + +async def _dependency_owner_id(principal: Annotated[str, Depends(_require_principal)]) -> str: + return principal + + +def _app( + owner_dependency=None, + *, + middleware_auth: bool = False, +) -> tuple[FastAPI, DurableRuntime]: + durable_settings = DurableSettings(durable_store="memory", durable_poll_interval=0.05) + runtime = DurableRuntime(InMemoryExecutionStoreProvider(durable_settings=durable_settings)) + wrapper = JobWrapper() + wrapper.setup() + deployment = runtime.deployment("jobs", wrapper) + + @asynccontextmanager + async def lifespan(_app: FastAPI): + try: + await runtime.start() + yield + finally: + await runtime.close() + + app = FastAPI(lifespan=lifespan) + if middleware_auth: + + @app.middleware("http") + async def authenticate(request: Request, call_next): + owner = request.headers.get("X-Owner") + if owner is None: + return JSONResponse({"detail": "Unauthorized"}, status_code=401) + request.state.principal = SimpleNamespace(owner_id=owner) + return await call_next(request) + + outer = APIRouter(prefix="/api") + outer.include_router( + create_durable_router(deployment, owner_id_dependency=owner_dependency), + prefix="/jobs", + ) + app.include_router(outer) + return app, runtime + + +def _wait(client: TestClient, url: str, expected: str) -> dict: + for _ in range(200): + response = client.get(url) + if response.json()["status"] == expected: + return response.json() + time.sleep(0.01) + pytest.fail(f"execution did not become {expected}") + + +def test_public_router_is_typed_prefix_safe_and_supports_all_routes() -> None: + app, _ = _app(owner_dependency=None) + with TestClient(app) as client: + submitted = client.post("/api/jobs/run-durable", json={"value": 0}) + assert submitted.status_code == 202 + assert submitted.headers["Location"].startswith("/api/jobs/executions/") + links = submitted.json()["links"] + assert set(links) == {"self", "cancel", "resume"} + waiting = _wait(client, links["self"], "waiting") + assert waiting["waiting"] == {"kind": "approval"} + assert client.post(links["resume"], json={"approved": "invalid"}).status_code == 422 + resumed = client.post(links["resume"], json={"approved": True}) + assert resumed.status_code == 202 + completed = _wait(client, links["self"], "completed") + assert completed["result"] == {"value": 1, "owner_id": None} + assert client.post(links["cancel"]).status_code == 200 + + openapi = app.openapi() + paths = openapi["paths"] + assert "/api/jobs/run-durable" in paths + assert "JobRequest" in str(paths["/api/jobs/run-durable"]["post"]["requestBody"]) + submit_response = paths["/api/jobs/run-durable"]["post"]["responses"] + assert "JobsExecutionResult" in str(submit_response) + assert "JobResult" in str(openapi["components"]["schemas"]["JobsExecutionResult"]) + assert "ResumeInput" in str(paths["/api/jobs/executions/{execution_id}/resume"]["post"]["requestBody"]) + + +def test_middleware_owner_dependency_isolates_every_operation_and_idempotency() -> None: + def current_owner_id(request: Request) -> str: + return request.state.principal.owner_id + + app, _ = _app(current_owner_id, middleware_auth=True) + alice = {"X-Owner": "alice", "Idempotency-Key": "same"} + bob = {"X-Owner": "bob", "Idempotency-Key": "same"} + with TestClient(app) as client: + assert client.post("/api/jobs/run-durable", json={"value": 2}).status_code == 401 + submitted = client.post("/api/jobs/run-durable", json={"value": 2}, headers=alice) + replay = client.post("/api/jobs/run-durable", json={"value": 2}, headers=alice) + assert replay.headers["Idempotent-Replay"] == "true" + url = submitted.json()["links"]["self"] + assert client.get(url, headers={"X-Owner": "alice"}).status_code == 200 + assert client.get(url, headers={"X-Owner": "bob"}).status_code == 404 + assert client.post(f"{url}/cancel", headers={"X-Owner": "bob"}).status_code == 404 + assert client.post(f"{url}/resume", headers={"X-Owner": "bob"}, json={"approved": True}).status_code == 404 + conflict = client.post("/api/jobs/run-durable", json={"value": 3}, headers=alice) + independent = client.post("/api/jobs/run-durable", json={"value": 2}, headers=bob) + assert conflict.status_code == 409 + assert independent.status_code == 202 + assert independent.json()["execution_id"] != submitted.json()["execution_id"] + + +def test_owner_dependency_composes_with_dependencies_and_fails_closed() -> None: + app, _ = _app(_dependency_owner_id) + with TestClient(app) as client: + assert client.get("/api/jobs/executions/missing", headers={"X-Deny": "yes"}).status_code == 403 + invalid = client.post("/api/jobs/run-durable", json={"value": 1}) + oversized = client.post( + "/api/jobs/run-durable", + json={"value": 1}, + headers={"X-Owner": "x" * 513}, + ) + assert invalid.status_code == 500 + assert oversized.status_code == 500 + + +def test_unscoped_router_uses_bearer_execution_ids() -> None: + app, _ = _app(owner_dependency=None) + with TestClient(app) as client: + submitted = client.post( + "/api/jobs/run-durable", + json={"value": 3}, + headers={"Idempotency-Key": "shared"}, + ) + replay = client.post( + "/api/jobs/run-durable", + json={"value": 3}, + headers={"Idempotency-Key": "shared"}, + ) + assert submitted.status_code == 202 + assert replay.headers["Idempotent-Replay"] == "true" + assert client.get(submitted.json()["links"]["self"]).status_code == 200 diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py index 0ebd76c3..d6645e4b 100644 --- a/tests/test_durable_store.py +++ b/tests/test_durable_store.py @@ -26,6 +26,7 @@ ) from hayhooks.durable.reference import InMemoryExecutionStore from hayhooks.durable.runtime import DurableRuntime +from hayhooks.durable.settings import DurableSettings from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider from hayhooks.settings import AppSettings, settings @@ -65,7 +66,7 @@ def _record() -> ExecutionRecord: def test_builtin_providers_snapshot_explicit_durable_settings() -> None: - app_settings = AppSettings( + durable_settings = DurableSettings( durable_redis_key_prefix="portable:durable", durable_lease_duration_ms=45_000, durable_lease_commit_safety_ms=2_000, @@ -75,10 +76,10 @@ def test_builtin_providers_snapshot_explicit_durable_settings() -> None: durable_max_progress_events=17, durable_max_record_bytes=32_768, ) - memory_store = InMemoryExecutionStoreProvider(app_settings=app_settings).create_execution_store("portable") + memory_store = InMemoryExecutionStoreProvider(durable_settings=durable_settings).create_execution_store("portable") redis_provider = RedisExecutionStoreProvider( redis=AsyncMock(), - app_settings=app_settings, + durable_settings=durable_settings, socket_timeout=1.5, socket_connect_timeout=2.5, health_check_interval=0, @@ -101,8 +102,8 @@ def test_builtin_providers_snapshot_explicit_durable_settings() -> None: def test_runtime_uses_its_explicit_settings_for_the_default_provider() -> None: - app_settings = AppSettings(durable_store="memory", durable_lease_duration_ms=45_000) - runtime = DurableRuntime(app_settings=app_settings) + durable_settings = DurableSettings(durable_store="memory", durable_lease_duration_ms=45_000) + runtime = DurableRuntime(durable_settings=durable_settings) provider = runtime._provider() @@ -110,22 +111,22 @@ def test_runtime_uses_its_explicit_settings_for_the_default_provider() -> None: assert provider.app_settings.durable_lease_duration_ms == 45_000 -async def test_runtime_uses_implicit_provider_settings_until_close(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_runtime_defaults_are_independent_of_hayhooks_settings(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(settings, "durable_store", "memory") original_attempts = settings.durable_max_attempts - runtime = DurableRuntime() + runtime = DurableRuntime(durable_settings=DurableSettings(durable_store="memory")) provider = runtime._provider() monkeypatch.setattr(settings, "durable_max_attempts", original_attempts + 1) assert runtime.app_settings.durable_max_attempts == provider.app_settings.durable_max_attempts == original_attempts await runtime.close() - assert runtime.app_settings.durable_max_attempts == original_attempts + 1 + assert runtime.settings.durable_max_attempts == original_attempts def test_runtime_uses_supplied_builtin_provider_as_its_settings_source() -> None: - app_settings = AppSettings(durable_store="memory", durable_lease_duration_ms=45_000, durable_max_attempts=7) - provider = InMemoryExecutionStoreProvider(app_settings=app_settings) + durable_settings = DurableSettings(durable_store="memory", durable_lease_duration_ms=45_000, durable_max_attempts=7) + provider = InMemoryExecutionStoreProvider(durable_settings=durable_settings) runtime = DurableRuntime(provider) store = provider.create_execution_store("portable") @@ -135,10 +136,17 @@ def test_runtime_uses_supplied_builtin_provider_as_its_settings_source() -> None def test_runtime_rejects_conflicting_builtin_provider_settings() -> None: - provider = InMemoryExecutionStoreProvider(app_settings=AppSettings(durable_lease_duration_ms=45_000)) + provider = InMemoryExecutionStoreProvider(durable_settings=DurableSettings(durable_lease_duration_ms=45_000)) with pytest.raises(ValueError, match="settings must match"): - DurableRuntime(provider, app_settings=AppSettings(durable_lease_duration_ms=60_000)) + DurableRuntime(provider, durable_settings=DurableSettings(durable_lease_duration_ms=60_000)) + + +def test_runtime_converts_hayhooks_settings_for_server_compatibility() -> None: + runtime = DurableRuntime(app_settings=AppSettings(durable_store="memory", durable_max_attempts=7)) + + assert runtime.settings.durable_store == "memory" + assert runtime.settings.durable_max_attempts == 7 async def test_runtime_provider_cannot_be_replaced(monkeypatch: pytest.MonkeyPatch) -> None: diff --git a/tests/test_it_status.py b/tests/test_it_status.py index 89abf8b1..c799c0e2 100644 --- a/tests/test_it_status.py +++ b/tests/test_it_status.py @@ -3,9 +3,7 @@ import pytest -from hayhooks.durable.runtime import durable_runtime from hayhooks.server.pipelines import registry -from hayhooks.server.routers.status import status_all @pytest.fixture(autouse=True) @@ -44,11 +42,12 @@ def test_status_no_pipelines(client, status_pipeline): assert len(status_response.json()["pipelines"]) == 0 -async def test_global_status_remains_a_liveness_probe_when_durable_is_unhealthy(monkeypatch): +def test_global_status_remains_a_liveness_probe_when_durable_is_unhealthy(monkeypatch, client): health = {"healthy": False, "deployments": {"job": {"healthy": False}}} - monkeypatch.setattr(durable_runtime, "health", AsyncMock(return_value=health)) + monkeypatch.setattr(client.app.state.durable_runtime, "health", AsyncMock(return_value=health)) - response = await status_all() + response = client.get("/status") - assert response.status == "Up!" - assert response.durable == health + assert response.status_code == 200 + assert response.json()["status"] == "Up!" + assert response.json()["durable"] == health From 11d286b4541fd12a682aa4f8b1ce776e2fa4c1fa Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Mon, 17 Aug 2026 13:28:09 +0200 Subject: [PATCH 16/28] fix(tests): align durable runtime ownership --- tests/test_durable_fastapi.py | 5 +++++ tests/test_durable_process_recovery.py | 8 +++++--- 2 files changed, 10 insertions(+), 3 deletions(-) diff --git a/tests/test_durable_fastapi.py b/tests/test_durable_fastapi.py index 274ae20a..a20918f9 100644 --- a/tests/test_durable_fastapi.py +++ b/tests/test_durable_fastapi.py @@ -2,6 +2,7 @@ from __future__ import annotations +import importlib.metadata import time from contextlib import asynccontextmanager from types import SimpleNamespace @@ -23,6 +24,10 @@ create_durable_router, ) +pytestmark = pytest.mark.skipif( + not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" +) + class JobRequest(BaseModel): value: int diff --git a/tests/test_durable_process_recovery.py b/tests/test_durable_process_recovery.py index 9115e03b..3bddc23c 100644 --- a/tests/test_durable_process_recovery.py +++ b/tests/test_durable_process_recovery.py @@ -126,11 +126,13 @@ def _cleanup_redis(redis_url: str, prefix: str) -> None: def create_a2a_recovery_app(): """Build the process-test A2A app after loading its pipeline fixture.""" - from hayhooks.durable.runtime import DurableDeployment + from hayhooks.durable.runtime import DurableDeployment, DurableRuntime from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.utils.deploy_utils import deploy_pipelines + from hayhooks.settings import settings - deploy_pipelines() + durable_runtime = DurableRuntime(app_settings=settings) + deploy_pipelines(durable_runtime=durable_runtime) if os.getenv(_CRASH_AFTER_A2A_SUBMIT_ENV) == "1": submit = DurableDeployment.submit @@ -140,7 +142,7 @@ async def submit_then_crash(self, *args, **kwargs): return result DurableDeployment.submit = submit_then_crash - return create_a2a_app() + return create_a2a_app(durable_runtime=durable_runtime) def _a2a_rpc(base_url: str, method: str, params: dict, request_id: str) -> dict: From d2ab7c4255c6d5a2876b42bc40c6e0d97b9c7b2d Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Mon, 17 Aug 2026 14:42:13 +0200 Subject: [PATCH 17/28] fix(durable): resolve final integration contracts --- docs/advanced/durable-engine.md | 5 +-- docs/advanced/durable-execution-operations.md | 9 ++--- docs/guides/production-best-practices.md | 7 ++-- docs/reference/api-reference.md | 6 +++- docs/reference/environment-variables.md | 2 +- src/hayhooks/durable/store.py | 6 ++-- src/hayhooks/server/a2a/durable_executor.py | 21 ++++++----- src/hayhooks/server/a2a/executor.py | 1 + tests/test_a2a.py | 8 ++--- tests/test_durable_a2a.py | 15 +++++++- tests/test_durable_execution.py | 2 ++ tests/test_durable_store.py | 35 +++++++++++++++++++ 12 files changed, 88 insertions(+), 29 deletions(-) diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 2ee8123b..22544207 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -203,8 +203,9 @@ and resumes verify that persisted work matches the active revision. `DurableExecutionManager.health_snapshot()` reports `nonterminal`, `runnable`, `lease_expiry`, and the current worker store-error streak. Repeated claim or -transition failures make readiness unhealthy until a worker completes a store -operation successfully. Alert on sustained runnable growth, repeated lease +transition failures make the deployment health snapshot unhealthy until a +worker completes a store operation successfully; `/status/{pipeline_name}` +then returns `503`. Alert on sustained runnable growth, repeated lease recovery, worker/store health failures, and runs that exceed their expected duration. diff --git a/docs/advanced/durable-execution-operations.md b/docs/advanced/durable-execution-operations.md index 4bf32968..74f6ba16 100644 --- a/docs/advanced/durable-execution-operations.md +++ b/docs/advanced/durable-execution-operations.md @@ -44,11 +44,12 @@ retain its terminal records through the configured Redis TTL. ## Health and incidents Health exposes `nonterminal`, `runnable`, `lease_expiry`, and -`worker_store_error_streak`. A claim or transition store failure makes readiness -unhealthy until that worker completes a store operation successfully. Investigate +`worker_store_error_streak`. A claim or transition store failure marks the +deployment health snapshot unhealthy until that worker completes a store +operation successfully; `/status/{pipeline_name}` then returns `503`. Investigate a growing runnable count, repeated lease recovery, store failures, or executions -that remain running/waiting longer than expected. Pause submissions, preserve the -Redis namespace, and inspect controls and fences before changing code or +that remain running/waiting longer than expected. Pause submissions, preserve +the Redis namespace, and inspect controls and fences before changing code or restarting workers. Use Redis 6.2 or later. Monitor the durable counts alongside Redis availability diff --git a/docs/guides/production-best-practices.md b/docs/guides/production-best-practices.md index 4d521d76..a30e7a4a 100644 --- a/docs/guides/production-best-practices.md +++ b/docs/guides/production-best-practices.md @@ -158,10 +158,10 @@ healthcheck: start_period: 40s ``` -For Kubernetes, use a readiness probe on the same endpoint: +For Kubernetes, use the same endpoint as a liveness probe: ```yaml -readinessProbe: +livenessProbe: httpGet: path: /status port: 1416 @@ -169,6 +169,9 @@ readinessProbe: periodSeconds: 30 ``` +For a durable pipeline, use `/status/{pipeline_name}` as its readiness probe. +It returns `503` when that deployment has no healthy durable worker slots. + ## Treat Durable Execution as a Controlled Beta Durable Pipelines and managed durable A2A Agents have stricter requirements diff --git a/docs/reference/api-reference.md b/docs/reference/api-reference.md index 3a07802e..e7cab852 100644 --- a/docs/reference/api-reference.md +++ b/docs/reference/api-reference.md @@ -99,7 +99,11 @@ Get status of all deployed pipelines. "pipeline1", "pipeline2" ], - "status": "Up!" + "status": "Up!", + "durable": { + "healthy": true, + "deployments": {} + } } ``` diff --git a/docs/reference/environment-variables.md b/docs/reference/environment-variables.md index fadb1cf8..4448a162 100644 --- a/docs/reference/environment-variables.md +++ b/docs/reference/environment-variables.md @@ -208,7 +208,7 @@ export HAYHOOKS_DEPLOY_CONCURRENCY=parallel ### HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT / HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT - Default: `5.0` seconds each -- Description: Bound established-socket operations and new Redis connections for durable execution. A timeout is reported as a store failure, causing worker backoff and readiness to return `503` rather than waiting indefinitely. +- Description: Bound established-socket operations and new Redis connections for durable execution. A timeout is reported as a store failure, causing worker backoff and marking the affected durable deployment unhealthy rather than waiting indefinitely. ### HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py index 23f11147..b32d9404 100644 --- a/src/hayhooks/durable/store.py +++ b/src/hayhooks/durable/store.py @@ -456,7 +456,7 @@ async def request_cancel(self, execution_id: str, reason: str | None = None) -> if control is None: return False if control.terminal: - return control.status is EngineStatus.CANCELED + return False event = _encode( { "sequence": control.progress_sequence + 1, @@ -468,14 +468,14 @@ async def request_cancel(self, execution_id: str, reason: str | None = None) -> limit=self.config.max_progress_event_bytes, label="progress", ) - await self._core_call( + plan = await self._core_call( "request cancellation", self.core.transition( execution_id, RequestCancellation(0, normalize_cancellation_reason(reason), (event,)), ), ) - return True + return bool(plan.progress_events) async def resume(self, execution_id: str, update: JsonValue | None = None) -> bool: """Resume a waiting execution with an optional JSON-safe application update.""" diff --git a/src/hayhooks/server/a2a/durable_executor.py b/src/hayhooks/server/a2a/durable_executor.py index 843d8a55..ded4ed61 100644 --- a/src/hayhooks/server/a2a/durable_executor.py +++ b/src/hayhooks/server/a2a/durable_executor.py @@ -317,18 +317,17 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None ) except KeyError: return - if not accepted: - return - log.bind(pipeline_name=self.pipeline_name, task_id=task.id, execution_id=execution_id).debug( - "Accepted durable A2A task cancellation" - ) updater = TaskUpdater(event_queue, task.id, task.context_id) - await updater.add_artifact( - [new_text_part("Cancellation requested")], - artifact_id=f"{task.id}-{DURABLE_PROGRESS_ARTIFACT_NAME}", - name=DURABLE_PROGRESS_ARTIFACT_NAME, - append=False, - ) + if accepted: + log.bind(pipeline_name=self.pipeline_name, task_id=task.id, execution_id=execution_id).debug( + "Accepted durable A2A task cancellation" + ) + await updater.add_artifact( + [new_text_part("Cancellation requested")], + artifact_id=f"{task.id}-{DURABLE_PROGRESS_ARTIFACT_NAME}", + name=DURABLE_PROGRESS_ARTIFACT_NAME, + append=False, + ) await self._wait_for_update(execution_id, owner_id, updater, terminal_only=True) async def _wait_for_update( diff --git a/src/hayhooks/server/a2a/executor.py b/src/hayhooks/server/a2a/executor.py index a298ff53..8b5dde65 100644 --- a/src/hayhooks/server/a2a/executor.py +++ b/src/hayhooks/server/a2a/executor.py @@ -77,6 +77,7 @@ async def emit(text: str, *, last: bool) -> None: return async for text in _iter_text_chunks(result): await emit(text, last=False) + await emit("", last=True) class ChatCompletionAgentExecutor(AgentExecutor): diff --git a/tests/test_a2a.py b/tests/test_a2a.py index 90002026..187e3ffe 100644 --- a/tests/test_a2a.py +++ b/tests/test_a2a.py @@ -301,12 +301,12 @@ async def test_execute_agent_task_streaming_result(): artifact_events = get_artifact_events(queue.events) # PipelineEvent items are skipped, text chunks are streamed incrementally - assert len(artifact_events) == 3 + assert len(artifact_events) == 4 texts = [event.artifact.parts[0].text for event in artifact_events] - assert texts == ["Hello, ", "world", " (question: hi)"] - # All chunks belong to the same artifact; terminal task status closes iterator output. + assert texts == ["Hello, ", "world", " (question: hi)", ""] + # All chunks belong to the same artifact and the final chunk closes it. assert len({event.artifact.artifact_id for event in artifact_events}) == 1 - assert [event.last_chunk for event in artifact_events] == [False, False, False] + assert [event.last_chunk for event in artifact_events] == [False, False, False, True] assert artifact_events[0].append is False assert artifact_events[1].append is True diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py index 493c3e4f..59788692 100644 --- a/tests/test_durable_a2a.py +++ b/tests/test_durable_a2a.py @@ -40,6 +40,7 @@ def __init__(self, status=ExecutionStatus.COMPLETED) -> None: self.submitted_payload = None self.resume_update = None self.cancel_requested = False + self.cancel_accepted = True async def start(self): return None @@ -68,7 +69,7 @@ async def request_cancel(self, _execution_id, **_kwargs): self.cancel_requested = True self.record.status = ExecutionStatus.CANCELED self.record.sequence += 1 - return True + return self.cancel_accepted class _BlockingDeployment(_Deployment): @@ -242,6 +243,18 @@ def test_a2a_http_cancel_reaches_durable_execution(monkeypatch, http_store) -> N assert canceled["status"]["state"] == "TASK_STATE_CANCELED" +def test_a2a_http_cancel_projects_a_terminal_race(monkeypatch, http_store) -> None: + deployment = _Deployment(status=ExecutionStatus.RUNNING) + deployment.cancel_accepted = False + app = _http_app(http_store, deployment, monkeypatch) + + with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) + canceled = _response_task(client.post("/durable-agent/", json=_cancel_payload(active["id"]))) + + assert canceled["status"]["state"] == "TASK_STATE_CANCELED" + + def test_return_immediately_waits_for_durable_submission(monkeypatch, http_store) -> None: deployment = _BlockingDeployment() app = _http_app(http_store, deployment, monkeypatch) diff --git a/tests/test_durable_execution.py b/tests/test_durable_execution.py index e3c92baf..805b45fd 100644 --- a/tests/test_durable_execution.py +++ b/tests/test_durable_execution.py @@ -594,6 +594,8 @@ def test_durable_rest_can_inspect_and_cancel_an_execution_from_an_old_revision(m assert canceled.status_code == 202 assert canceled.json()["cancellation_requested_at"] is not None wrapper.release.set() + _wait_for_status(client, links["self"], "canceled", "durable execution did not cancel") + assert client.post(links["cancel"]).status_code == 200 finally: wrapper.release.set() diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py index d6645e4b..863eed42 100644 --- a/tests/test_durable_store.py +++ b/tests/test_durable_store.py @@ -244,6 +244,41 @@ async def gated_transition(run_id, command, *, candidate=False): assert [event.sequence for event in record.progress] == [1, 2] +async def test_cancel_reports_only_a_newly_accepted_request() -> None: + store = _store() + await store.submit(_record()) + + assert await store.request_cancel("run_1") + assert not await store.request_cancel("run_1") + + +async def test_cancel_losing_a_terminal_transition_race_returns_false(monkeypatch: pytest.MonkeyPatch) -> None: + core = InMemoryExecutionStore(deployment="deployment", config=_config()) + store = ExecutionStore(core, definition_revision="rev-1") + await store.submit(_record()) + claim = await store.claim_next("worker") + assert claim is not None + entered = asyncio.Event() + release = asyncio.Event() + original_transition = core.transition + + async def gated_transition(run_id, command, *, candidate=False): + if isinstance(command, RequestCancellation): + entered.set() + await release.wait() + return await original_transition(run_id, command, candidate=candidate) + + monkeypatch.setattr(core, "transition", gated_transition) + async with claim: + cancellation = asyncio.create_task(store.request_cancel("run_1")) + await entered.wait() + claim.record.status = ExecutionStatus.COMPLETED + claim.record.result = {"answer": "done"} + await claim.complete() + release.set() + assert not await cancellation + + async def test_checkpoint_keeps_progress_added_after_a_concurrent_cancellation() -> None: store = _store() await store.submit(_record()) From eba32efc316dd5a160aaeb7d2a2f7ecbd8dcf256 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Tue, 18 Aug 2026 08:48:25 +0200 Subject: [PATCH 18/28] minor portability fixes --- docs/advanced/durable-engine.md | 7 +- .../durable-fastapi-integration-plan.md | 876 ------------------ src/hayhooks/durable/fastapi.py | 13 +- src/hayhooks/server/app.py | 20 +- src/hayhooks/server/pipelines/registry.py | 18 +- src/hayhooks/server/routers/dashboard.py | 6 +- src/hayhooks/server/routers/draw.py | 7 +- src/hayhooks/server/routers/openai.py | 121 ++- src/hayhooks/server/routers/status.py | 6 +- src/hayhooks/server/utils/deploy_utils.py | 68 +- tests/conftest.py | 2 +- tests/test_deploy_performance.py | 3 + tests/test_deploy_utils.py | 2 + tests/test_durable_deployment_lifecycle.py | 39 +- tests/test_registry.py | 4 +- 15 files changed, 208 insertions(+), 984 deletions(-) delete mode 100644 docs/advanced/durable-fastapi-integration-plan.md diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 22544207..0be7c6b2 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -50,9 +50,10 @@ exceptions, and FastAPI adapter directly from `hayhooks.durable`. A standalone runtime starts only deployments attached to that runtime; it never inspects the Hayhooks pipeline registry. -This complete `app.py` adds an authenticated durable API to an existing FastAPI -application. Authentication middleware is expected to set a stable principal -on `request.state` before the owner dependency runs: +This durable integration fragment assumes the host FastAPI application already +has authentication middleware that sets a stable principal on `request.state` +before the owner dependency runs. The authentication middleware itself is +application-specific and intentionally omitted: ```python from contextlib import asynccontextmanager diff --git a/docs/advanced/durable-fastapi-integration-plan.md b/docs/advanced/durable-fastapi-integration-plan.md deleted file mode 100644 index c7bb5c7e..00000000 --- a/docs/advanced/durable-fastapi-integration-plan.md +++ /dev/null @@ -1,876 +0,0 @@ -# Portable FastAPI durable integration plan - -**Status:** Proposed - -**Scope:** Hayhooks durable REST integration, authentication composition, and runtime ownership - -**Primary objective:** Make Hayhooks use the same public durable integration that an independent FastAPI application uses - -## Summary - -Hayhooks already has a durable execution engine with Redis-backed recovery, -fenced leases, idempotent submission, retries, progress, cancellation, and -wait/resume. The engine is usable from a standalone `DurableRuntime`, but the -current REST integration is implemented inside the Hayhooks server and reaches -into process-global runtime and registry state. - -This plan introduces one public FastAPI adapter and makes Hayhooks consume it: - -1. Applications own a `DurableRuntime` and start and close it in their lifespan. -2. `create_durable_router()` returns a standard FastAPI `APIRouter` for one - `DurableDeployment`. -3. An ordinary FastAPI dependency supplies a stable owner ID when authenticated - ownership isolation is required. -4. Hayhooks keeps only a small internal mount/unmount shim for its dynamic - pipeline deployment lifecycle. -5. Hayhooks REST, A2A, and MCP applications move from the process-global runtime - to application-owned runtime instances. - -The durable engine remains unaware of JWTs, cookies, sessions, middleware, -FastAPI application state, and the Hayhooks registry. Redis remains the source -of truth for execution state. - -## Goals - -- Provide a simple, documented public integration for an existing FastAPI - application. -- Compose with authentication middleware and existing FastAPI dependencies. -- Preserve typed request, resume, and result schemas in OpenAPI. -- Preserve all current durable REST paths and response semantics. -- Make ownership enforcement consistent for submit, inspect, cancel, and - resume. -- Make Hayhooks dogfood the public router rather than maintaining separate REST - handlers. -- Make each application instance own its runtime, workers, and provider - lifecycle. -- Preserve dynamic deploy, overwrite, rollback, and undeploy behavior in the - Hayhooks server. -- Keep Redis keys and persisted execution records unchanged. -- Maintain the current controlled-beta multi-process execution model. - -## Non-goals - -- A generic non-Haystack job framework. -- Authentication or authorization implemented by the durable engine. -- Persisting access tokens, sessions, or complete principal objects. -- A custom FastAPI middleware supplied by Hayhooks. -- A mounted durable sub-application. -- A `DurableRouter` subclass or a stateful integration/service container. -- Per-operation authorization hooks in the first public API. -- Separating API-only and worker-only processes in this change. -- Changing the Redis schema, execution state machine, or delivery semantics. -- Replacing Hayhooks' dynamic pipeline deployment transaction. - -## Current state - -### Durable runtime - -`DurableRuntime` owns deployment managers and a shared execution-store provider. -`DurableDeployment` owns the typed wrapper contract, Haystack adapter, store, -and worker manager. Redis or the in-memory reference backend owns persisted -execution data. - -The current standalone embedding path is valid, but it requires applications to -write their own HTTP handlers. The public facade also omits several types that -an embedding application naturally needs. - -### REST transport - -The current durable REST handlers live in -`hayhooks.server.durable.routes`. That module currently combines four concerns: - -- reusable durable HTTP behavior; -- Hayhooks trusted-header owner resolution; -- process-global runtime discovery; -- Hayhooks registry metadata and dynamic route mutation. - -Only the first concern belongs in a portable FastAPI adapter. - -### Runtime ownership - -Hayhooks REST, A2A, MCP, status, and deployment code currently import the same -module-level `durable_runtime`. This makes multiple application instances in a -single process share deployments, provider ownership, and shutdown behavior. -An independent FastAPI application would instead create and own its runtime. - -## Architectural decision - -### Chosen design - -Use a public `APIRouter` factory and a FastAPI owner-ID dependency: - -```python -from collections.abc import Awaitable, Callable - - -def create_durable_router( - deployment: DurableDeployment, - *, - owner_id_dependency: Callable[..., str | Awaitable[str]] | None, -) -> APIRouter: - ... -``` - -The owner dependency is a required keyword. Passing `None` is an explicit -choice to use unscoped bearer-by-execution-ID access. Supplying a dependency -enables owner isolation. - -The function returns routes only. It does not start workers, own the runtime, -mutate a registry, or alter an application. The caller uses ordinary FastAPI -composition: - -```python -app.include_router( - create_durable_router( - deployment, - owner_id_dependency=current_owner_id, - ), - prefix="/jobs", -) -``` - -### Why `APIRouter` - -- It is FastAPI's native unit of route composition. -- Host middleware automatically wraps included routes. -- Dependencies compose with existing authentication and authorization. -- Typed request and response models remain visible in OpenAPI. -- Applications control prefixes, tags, and router-level dependencies. -- Hayhooks can include the same router dynamically and retain its existing - route replacement logic. -- The adapter has no lifecycle or global state of its own. - -### Why a dependency rather than a callback - -A callback invoked manually by the route handler would need to reproduce part -of FastAPI's dependency system. A dependency already supports: - -- middleware-populated `request.state`; -- nested `Depends(...)` authentication dependencies; -- OAuth/OpenAPI security dependencies; -- async and sync implementations; -- application-specific `HTTPException` responses; -- dependency overrides in tests. - -The durable adapter needs only the resulting stable owner ID. It does not need -to know how authentication was performed. - -### Request and execution flow - -```mermaid -flowchart LR - request["HTTP request"] --> middleware["Host auth middleware"] - middleware --> permission["Host authorization dependencies"] - permission --> owner["Owner-ID dependency"] - owner --> router["Hayhooks durable router"] - router --> deployment["DurableDeployment"] - deployment <--> redis["Redis execution store"] - deployment --> worker["Process-local durable worker"] - worker --> wrapper["PipelineWrapper with DurableContext"] -``` - -Middleware runs before FastAPI dependency resolution. The host application -therefore retains control of authentication, request context, and broad API -authorization. The owner dependency reduces that context to the stable string -needed for record isolation. The router validates and translates HTTP; the -deployment and Redis store perform durable execution and recovery; the wrapper -remains independent of HTTP. - -### Rejected alternatives - -| Alternative | Reason for rejection | -|---|---| -| Durable authentication middleware | Couples the engine to authentication and duplicates host middleware | -| Mounted FastAPI/Starlette sub-app | Makes OpenAPI, prefixes, host middleware, and dynamic replacement harder | -| Runtime callback hooks | Bypass FastAPI dependency injection and make error/security behavior bespoke | -| Stateful integration class | Adds lifecycle and registration state that already exists in FastAPI and `DurableRuntime` | -| Manually documented endpoint examples only | Leaves each application to duplicate validation, errors, ownership, and links | -| `DurableRuntime.create_router()` | Couples the runtime layer directly to FastAPI | - -## State ownership - -Storing the runtime on `app.state` is intentional. It stores a process-local -service handle, not durable execution data. - -| Location | Owned data | -|---|---| -| `app.state` | Runtime object, deployment definitions, worker tasks, provider/client handles | -| Python process | Wrapper instances, active call stacks, event-loop tasks | -| Redis | Controls, inputs, checkpoints, progress, waits, results, errors, leases, indexes, idempotency bindings | - -After a process restart, the application reconstructs the runtime and wrapper -definitions. Redis retains nonterminal work. An expired running lease is -recovered and made claimable according to the existing state machine. - -No execution record, checkpoint, or result should be copied into `app.state`. -This centralized-state guarantee applies when using the Redis provider. The -in-memory provider is deliberately volatile, process-local, and suitable only -for development and tests. - -The public router does not need to read `app.state`; it closes over one -`DurableDeployment`. Keeping the runtime on `app.state` is recommended when -status routes, deployment utilities, or other application components need the -same process-local handle. Hayhooks itself needs that access. A small external -application may instead keep the runtime only in its application-factory -lifespan closure. - -## Target external-application experience - -### Wrapper authoring - -Pipeline wrappers remain independent of HTTP and authentication: - -```python -from pydantic import BaseModel - -from hayhooks import BasePipelineWrapper -from hayhooks.durable import DurableContext - - -class JobRequest(BaseModel): - document_id: str - - -class JobResult(BaseModel): - indexed: bool - - -class JobWrapper(BasePipelineWrapper): - durable_revision = "job-v1" - - def setup(self) -> None: - self.pipeline = build_pipeline() - - async def run_durable_async( - self, - context: DurableContext, - request: JobRequest, - ) -> JobResult: - result = await context.run_pipeline_async( - {"loader": {"document_id": request.document_id}}, - checkpoint_at=["loader"], - ) - return JobResult(indexed=bool(result["loader"]["indexed"])) -``` - -The deployment continues to derive its Pydantic request and result contracts -from the wrapper method annotations. - -### Authentication middleware - -When middleware has already authenticated the request and populated -`request.state`: - -```python -from fastapi import Request - - -def current_owner_id(request: Request) -> str: - principal = request.state.principal - return f"{principal.tenant_id}:{principal.subject_id}" -``` - -When the application already uses dependencies: - -```python -from typing import Annotated - -from fastapi import Depends - - -async def current_owner_id( - principal: Annotated[Principal, Depends(require_principal)], -) -> str: - return f"{principal.tenant_id}:{principal.subject_id}" -``` - -Both forms have the same durable behavior. - -### Runtime and router - -```python -from contextlib import asynccontextmanager - -from fastapi import Depends, FastAPI - -from hayhooks.durable import ( - DurableRuntime, - RedisExecutionStoreProvider, - create_durable_router, -) - - -@asynccontextmanager -async def lifespan(app: FastAPI): - runtime = app.state.durable_runtime - try: - await runtime.start() - yield - finally: - await runtime.close() - - -def create_app() -> FastAPI: - provider = RedisExecutionStoreProvider( - redis_url="redis://localhost:6379/0", - key_prefix="myapp:durable", - ) - runtime = DurableRuntime(provider) - - wrapper = JobWrapper() - wrapper.setup() - deployment = runtime.deployment("jobs", wrapper) - - app = FastAPI(lifespan=lifespan) - app.state.durable_runtime = runtime - app.include_router( - create_durable_router( - deployment, - owner_id_dependency=current_owner_id, - ), - prefix="/jobs", - dependencies=[Depends(require_jobs_permission)], - ) - return app - - -app = create_app() -``` - -The router-level authorization dependency is optional. Ownership and general -authorization remain distinct: - -- `require_jobs_permission` decides whether the caller may use the job API. -- `current_owner_id` provides the stable identity used to isolate records. - -## Public router behavior - -### Routes - -The returned router contains relative paths so the host application controls -the prefix: - -| Method | Relative path | Behavior | -|---|---|---| -| `POST` | `/run-durable` | Validate and submit detached work | -| `GET` | `/executions/{execution_id}` | Inspect safe execution state | -| `POST` | `/executions/{execution_id}/cancel` | Request cooperative cancellation | -| `POST` | `/executions/{execution_id}/resume` | Resume waiting work with optional typed input | - -Hayhooks includes the router at `/{pipeline_name}`, preserving all existing -paths. - -### Dependency binding inside the factory - -The factory fixes scoped versus unscoped behavior once, when the router is -created. When an owner dependency is supplied, each route binds it with -`Depends(...)`, validates its resolved return value, and sets -`enforce_owner=True`. When `None` is supplied, each route binds a private -constant dependency that returns `None`, and sets `enforce_owner=False`. - -Do not let a configured dependency return `None` to select unscoped behavior -at request time. That would turn an authentication bug into an authorization -bypass. The dependency's sync or async execution remains FastAPI's -responsibility. - -### Response behavior - -The extraction must preserve: - -| Situation | Response | -|---|---| -| New submission | `202 Accepted` | -| Nonterminal idempotent replay | `202 Accepted` plus `Idempotent-Replay: true` | -| Retained terminal replay | `200 OK` plus `Idempotent-Replay: true` | -| Accepted cancellation | `202 Accepted` | -| Already terminal cancellation | `200 OK` | -| Successful resume | `202 Accepted` | -| Missing or foreign-owned execution | `404 Not Found` | -| Idempotency or revision conflict | `409 Conflict` | -| Execution is not waiting | `409 Conflict` | -| Invalid ID, request, or resume body | `422 Unprocessable Entity` | -| Admission limit | `503 Service Unavailable` plus `Retry-After` | -| Execution-store outage | `503 Service Unavailable` | - -The `Location` and result links must be generated with named-route resolution -through `request.url_for(...)`, rather than by concatenating the deployment -name. Each factory result must give its routes deployment-unique names, such as -`hayhooks.durable.{deployment_name}.inspect`, so multiple durable deployments -cannot resolve one another's links. Hayhooks deployment names are unique within -an application; including the same deployment router more than once is outside -the initial contract. - -Named resolution keeps links correct when the router is included below -additional application prefixes or root paths. Preserve the current relative -link contract by using the resolved URL's path component for response links and -the `Location` header. - -### Typed OpenAPI models - -The adapter keeps the current dynamic request and result model behavior: - -- submission uses `deployment.request_type`; -- result fields use `deployment.result_type` when declared; -- resume uses `deployment.resume_type` when declared; -- inspect, cancel, and resume return the safe execution projection; -- private input, state, checkpoints, owner, and fencing details remain absent. - -The implementation may continue setting endpoint annotations/signatures after -handler construction because FastAPI consumes those annotations during route -registration. - -## Ownership and authentication contract - -### Owner ID rules - -When `owner_id_dependency` is supplied, the adapter must require: - -- a string; -- at least one character; -- no more than 512 characters; -- a stable value across token refreshes and process restarts. - -Recommended values are immutable application IDs, for example -`tenant_uuid:user_uuid`. Do not use access tokens, session IDs, emails, or -display names. - -The host application chooses the ownership granularity. Use a tenant ID for -tenant-owned jobs, a user ID for user-owned jobs, or a stable compound ID when -both boundaries matter. - -The dependency should perform authentication and may raise the host -application's normal `401` or `403`. The durable adapter must not replace those -responses. - -When the dependency is configured, a missing or invalid owner must fail closed. -It must never switch the request to unscoped access. - -### Enforcement - -The adapter passes the owner to every deployment operation: - -```python -await deployment.submit(..., owner_id=owner_id) -await deployment.get(..., owner_id=owner_id, enforce_owner=True) -await deployment.request_cancel(..., owner_id=owner_id, enforce_owner=True) -await deployment.resume(..., owner_id=owner_id, enforce_owner=True) -``` - -The existing deployment behavior returns `KeyError` for both missing records -and owner mismatches. The router maps both to `404`, avoiding an execution-ID -existence oracle. - -Owner-scoped submission continues deriving the internal execution ID from the -owner and caller-provided idempotency key. The same external key can therefore -be used independently by different owners. - -### Unscoped mode - -Passing `owner_id_dependency=None` explicitly retains the current behavior: - -- records have no owner; -- possession of a sufficiently unguessable execution ID grants access; -- the router does not enforce owner matching. - -This is useful for local development and services protected by a single -application-wide authorization boundary. Documentation must label it as an -explicit security choice, not an authentication default. - -### Wrapper access to identity - -Background work cannot depend on an HTTP request, middleware state, cookies, -or a current token. Those values do not exist after process recovery. - -Add a minimal public property to `DurableContext`: - -```python -@property -def owner_id(self) -> str | None: - return self.record.owner_id -``` - -This lets durable application code use the persisted stable identity without -exposing the full private execution record as its normal API. - -Roles and permissions should be checked before submission. Full principal -objects and tokens must not be persisted automatically. If a job needs -additional trusted identifiers, they must be deliberately represented as -non-secret validated input or durable application state. - -## Hayhooks dogfooding design - -### Public adapter boundary - -Create `src/hayhooks/durable/fastapi.py`. It owns all reusable HTTP behavior and -imports only durable public/infrastructure types plus FastAPI/Pydantic. - -It must not import: - -- `hayhooks.server.pipelines.registry`; -- the module-level `durable_runtime`; -- `hayhooks.settings`; -- `BasePipelineWrapper`; -- deployment or route mutation utilities from `hayhooks.server`. - -### Hayhooks server shim - -Reduce `hayhooks.server.durable.routes` to Hayhooks-specific composition: - -1. Determine whether the deployment has durable capability. -2. Build the trusted-header owner dependency from that deployment's settings. -3. Remove the previous durable route family for the pipeline. -4. Include `create_durable_router(deployment, ...)` at - `/{pipeline_name}`. -5. Invalidate/rebuild OpenAPI according to the existing deferred-rebuild flag. - -The trusted-header dependency remains a Hayhooks server concern because a -third-party application should normally use its authenticated principal rather -than trust a configurable raw header. - -The current `durable_request_model` and `durable_response_model` registry -metadata is not read elsewhere in the codebase. Remove those writes rather than -adding a public result object solely to preserve unused internal metadata. - -### Dynamic deployment lifecycle - -Hayhooks must preserve the existing publication transaction: - -1. Capture the current wrapper, routes, OpenAPI schema, modules, files, and - durable deployment. -2. Quiesce and close the previous deployment when replacing it. -3. Reject replacement while nonterminal durable work would be stranded. -4. Prepare the new wrapper and durable deployment before publication. -5. Build/include routes whose closures reference the new deployment. -6. Publish registry and runtime state and activate the candidate without an - intervening await point. -7. Restore the previous routes, runtime deployment, modules, files, and worker - state if publication fails. - -`APIRouter` inclusion produces ordinary `APIRoute` instances on the application, -so the current path-based removal and route-list snapshot rollback remain -usable. Do not introduce a route-mount handle or registration class unless the -existing mechanism proves insufficient in tests. - -## Application-owned runtime design - -### REST factory - -Update the application factory without breaking existing callers: - -```python -def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: - runtime = durable_runtime or DurableRuntime() - app = FastAPI(...) - app.state.durable_runtime = runtime - ... -``` - -The lifespan reads the runtime from the application instance, starts it after -startup pipeline preparation, and closes it before application-owned dependent -resources are closed. - -Status endpoints should read the runtime from `request.app.state`, not import a -singleton. - -Deployment helpers should receive the runtime explicitly when no application is -available, or use the runtime attached to the supplied application. Avoid a -fallback that silently selects the process-global runtime. - -### A2A and MCP factories - -Apply the same ownership rule to other server factories: - -- `create_a2a_app(..., durable_runtime=runtime)`; -- `create_agent_executor(..., durable_runtime=runtime)`; -- A2A health reads its supplied runtime; -- MCP server/application factories retain and close their supplied runtime. - -CLI commands construct one runtime for the server they launch and pass that -same instance through construction and lifespan. - -### Compatibility singleton - -Keep `hayhooks.durable.durable_runtime` temporarily as a compatibility export, -but stop using it inside Hayhooks. Document application-owned `DurableRuntime` -as the supported integration. Removal of the singleton, if desired, should be a -separate deprecation decision. - -Removing internal use of the singleton also removes the need for the private -runtime-to-registry attachment. Registry discovery remains a Hayhooks server -responsibility; standalone runtimes continue to know only deployments that the -application explicitly attaches. - -## Multi-process behavior - -Every Uvicorn worker process creates its own application state, runtime, worker -tasks, and Redis client/pool. All processes use the same Redis namespace and -definition revision for a logical deployment. - -The existing Redis fences and leases coordinate claims across processes. The -effective execution concurrency is: - -```text -application processes × durable_execution_concurrency per deployment -``` - -The first rollout should remain within the controlled-beta profile of one to -three replicas with conservative per-process concurrency. Application state is -not shared and must never be treated as a cross-process coordination mechanism. - -## Implementation phases - -### Phase 0: Lock the existing HTTP contract - -Before moving handlers, add or identify focused tests for every existing route, -status code, response header, owner behavior, typed model, and error mapping. -These tests define extraction compatibility. - -No implementation behavior changes in this phase. - -### Phase 1: Add the public FastAPI adapter - -Add: - -- `src/hayhooks/durable/fastapi.py`; -- `tests/test_durable_fastapi.py` containing a standalone FastAPI app. - -Update: - -- `src/hayhooks/durable/__init__.py` to lazily export - `create_durable_router`, `DurableContext`, `DurableDeployment`, and the public - durable exceptions; -- `src/hayhooks/durable/context.py` with `owner_id`; -- durable embedding documentation with authenticated and unscoped examples. - -Move, rather than rewrite, the existing validation, response projection, and -error mapping. Keep the diff behavior-preserving. - -### Phase 2: Make Hayhooks REST consume the adapter - -Update: - -- `src/hayhooks/server/durable/routes.py` to contain only trusted-header - dependency construction and dynamic route composition; -- `src/hayhooks/server/utils/deploy_utils.py` to pass the prepared - `DurableDeployment` into the composition shim; -- durable REST and deployment lifecycle tests. - -Delete the duplicated endpoint handlers from the server module. Preserve route -paths and operation behavior. - -### Phase 3: Replace internal global-runtime use - -Update REST, status, deploy/undeploy, A2A, MCP, and their app factories to pass -an application-owned runtime explicitly. Store the REST runtime on `app.state`. - -Keep each protocol migration reviewable. A safe order is: - -1. REST app and status endpoints; -2. REST dynamic deployment transaction; -3. A2A app, executors, and health; -4. MCP server and lifespan; -5. remove all internal imports of the singleton; -6. remove the private registry attachment from `runtime.py`. - -The phase is complete when searching `src/hayhooks/server` finds no import of -the module-level `durable_runtime`. - -### Phase 4: Focus durable configuration - -Introduce a durable-only configuration model containing Redis connection, -retention, record limits, lease, retry, polling, shutdown, admission, and -concurrency settings. - -Allow built-in providers and `DurableRuntime` to consume it without requiring -the full Hayhooks `AppSettings`. Retain an internal conversion from Hayhooks -settings for server compatibility. - -Keep this separate from route extraction so configuration changes cannot hide -HTTP regressions. - -### Phase 5: Release and adoption - -- Run the external FastAPI example against a built wheel rather than the source - checkout. -- Update API and durable-engine documentation. -- Note application-owned runtime and router integration in release notes. -- Pin the target application's trial to the released version. -- Start with one authenticated deployment and a dedicated Redis namespace. - -## Test plan - -### Public adapter tests - -The standalone test application must not import `hayhooks.server.app`, the -Hayhooks registry, or the global runtime. - -Cover: - -- typed submission and result validation; -- typed resume input; -- all four routes; -- generated OpenAPI schemas; -- new submission and idempotent replay headers/statuses; -- cancellation before and after terminal state; -- wait/resume and revision conflict; -- invalid execution IDs and idempotency keys; -- record-size and admission failures; -- normalized execution-store failures; -- safe views excluding private fields; -- link resolution beneath an additional `include_router(prefix=...)` prefix. - -### Authentication and ownership matrix - -| Scenario | Expected result | -|---|---| -| Middleware rejects unauthenticated caller | Host application's `401` | -| Owner dependency raises `403` | Host application's `403` | -| Configured dependency returns empty/invalid owner | Fail closed; never unscoped | -| Alice submits and inspects | Success | -| Bob inspects Alice's execution | `404` | -| Bob cancels Alice's execution | `404` | -| Bob resumes Alice's execution | `404` | -| Alice reuses same key and input | Idempotent replay | -| Alice reuses same key with different input | `409` | -| Bob uses Alice's external key | Independent execution | -| Explicit unscoped router | Existing bearer-ID behavior | - -Provide one test where middleware writes `request.state.principal`, and another -where the owner dependency depends on an existing FastAPI authentication -dependency. - -### Hayhooks dogfooding regression tests - -Preserve and extend tests for: - -- startup deployment route creation; -- dynamic deployment after runtime start; -- durable-to-durable overwrite; -- durable-to-nondurable overwrite; -- undeploy route removal; -- refusal to strand queued, running, or waiting work; -- candidate preparation failure; -- publication failure and route rollback; -- OpenAPI replacement with new request/result/resume types; -- trusted owner header compatibility; -- deferred OpenAPI rebuild during batch startup. - -### Runtime ownership tests - -- Two FastAPI applications in one process receive distinct runtimes. -- A deployment installed in app A is absent from app B. -- Closing app A does not close app B's provider or workers. -- A supplied runtime is the one used by routes, status, and deployment helpers. -- REST, A2A, and MCP lifespans close only the runtime they own. -- Startup failure still closes the corresponding runtime exactly once. - -### Existing engine tests - -The following remain mandatory: - -- reducer and reference-store contract tests; -- Redis transaction and concurrent-claim tests; -- manager retry, cancellation, and lease tests; -- process-kill/restart recovery; -- A2A recovery tests; -- type checking and linting. - -## Compatibility requirements - -### HTTP compatibility - -- Existing Hayhooks durable paths remain unchanged. -- Existing request and response payloads remain unchanged. -- Status codes and headers remain unchanged. -- Owner mismatch remains indistinguishable from a missing execution. -- Current trusted owner header behavior remains available in the Hayhooks - server shim. - -### Python compatibility - -- Existing direct imports continue to work. -- New public imports are additive. -- The global `durable_runtime` remains importable during the compatibility - period. -- `create_app()` with no arguments continues to work, but owns a new runtime. -- Supplying a runtime is keyword-only. - -### Persistence compatibility - -This work must not modify: - -- Redis key construction; -- control serialization; -- payload formats; -- idempotency digests; -- execution state transitions; -- lease or fence semantics. - -No namespace migration is required for this integration refactor. - -## Operational guidance - -- Use an isolated Redis key prefix per application/environment. -- Use Redis 6.2 or later with persistence, backups, TLS/authentication as - appropriate, and `maxmemory-policy noeviction`. -- If the host application supplies a Redis client, durable storage requires - binary responses (`decode_responses=False`). -- Set `close_redis=False` when the host owns the client lifecycle. -- Close the durable runtime before closing a shared Redis client. -- Keep wrapper definition revisions identical across replicas. -- Account for Uvicorn workers when setting durable concurrency. -- Continue making every external side effect idempotent because execution is - at least once. - -## Risks and mitigations - -| Risk | Mitigation | -|---|---| -| Router extraction changes an HTTP edge case | Lock the current contract before moving handlers | -| Included router links ignore a host prefix | Generate links from named routes and the current request | -| Owner dependency accidentally returns `None` | Fail closed whenever a dependency is configured | -| Authentication data is expected during background recovery | Persist only owner ID and deliberate validated identifiers | -| Dynamic overwrite leaves old handlers active | Remove the known route family before including the candidate router | -| Publication failure loses previous routes | Retain the existing route-list snapshot rollback | -| Multiple app instances share shutdown state | Construct and attach one runtime per app | -| Multi-worker deployments create unexpected concurrency | Document and test the process × slot calculation | -| Scope expands into a workflow framework | Keep generic runners, schedules, and transport-independent auth out of this work | - -## Acceptance criteria - -The work is complete when all of the following are true: - -1. An independent FastAPI application can integrate durable execution using - only public Hayhooks imports. -2. That application can supply identity from middleware or an existing FastAPI - authentication dependency. -3. Hayhooks' own durable REST endpoints are created by the same public router - factory. -4. The public router has no dependency on the Hayhooks registry, global runtime, - or server settings. -5. Hayhooks server modules no longer import the global durable runtime. -6. Every app/server factory owns and closes its runtime. -7. Durable wrappers remain HTTP-independent and can read their stable owner via - `DurableContext.owner_id`. -8. Existing REST paths, schemas, status codes, headers, and owner behavior are - preserved. -9. Dynamic deployment, rollback, and undeploy tests remain green. -10. Live Redis and process-recovery tests remain green. -11. Redis data written before this refactor remains readable without migration. -12. Documentation includes authenticated, unscoped, shared-Redis, and - multi-worker examples. - -## Deliberately deferred extensions - -Add these only when a real integration requires them: - -- per-operation authorization dependencies; -- API-only versus worker-only runtime modes; -- generic non-Haystack runners; -- richer persisted principal metadata; -- customizable durable route shapes; -- a stable Redis schema migration framework. - -The first portable release should contain the smallest complete boundary: -application-owned runtime, one public router factory, and one owner-ID -dependency. diff --git a/src/hayhooks/durable/fastapi.py b/src/hayhooks/durable/fastapi.py index 568a6911..84d867ee 100644 --- a/src/hayhooks/durable/fastapi.py +++ b/src/hayhooks/durable/fastapi.py @@ -3,7 +3,6 @@ from __future__ import annotations import inspect -import re from collections.abc import Awaitable, Callable from typing import Annotated, Any, cast @@ -15,7 +14,6 @@ from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError from hayhooks.durable.runtime import DefinitionRevisionConflictError, DurableDeployment, IdempotencyConflictError -_IDEMPOTENCY_KEY_PATTERN = re.compile(rf"^{RUN_ID_PATTERN}$") _MAX_OWNER_LENGTH = 512 _MAX_OWNER_SCOPED_IDEMPOTENCY_KEY_LENGTH = 63 ExecutionId = Annotated[str, Path(pattern=rf"^{RUN_ID_PATTERN}$", min_length=1, max_length=128)] @@ -98,14 +96,13 @@ async def submit( response: Response, request: Request, owner_id: Any = Depends(owner_dependency), # noqa: B008 - idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), + idempotency_key: str | None = Header( + default=None, + alias="Idempotency-Key", + pattern=rf"^{RUN_ID_PATTERN}$", + ), ) -> ExecutionResult: owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) - if idempotency_key is not None and _IDEMPOTENCY_KEY_PATTERN.fullmatch(idempotency_key) is None: - raise HTTPException( - status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, - detail="Idempotency-Key must contain 1-128 letters, digits, underscores, or hyphens", - ) if ( enforce_owner and idempotency_key is not None diff --git a/src/hayhooks/server/app.py b/src/hayhooks/server/app.py index 387442cb..1afb3337 100644 --- a/src/hayhooks/server/app.py +++ b/src/hayhooks/server/app.py @@ -22,14 +22,9 @@ from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.logger import RequestIdMiddleware, intercept_stdlib_logging, log, log_elapsed -from hayhooks.server.routers import ( - dashboard_router, - deploy_router, - draw_router, - openai_router, - status_router, - undeploy_router, -) +from hayhooks.server.pipelines.registry import PipelineRegistry +from hayhooks.server.routers import dashboard_router, deploy_router, draw_router, status_router, undeploy_router +from hayhooks.server.routers.openai import create_openai_router from hayhooks.server.tracing import ( SPAN_PIPELINE_STARTUP_DEPLOY, build_trace_tags, @@ -280,7 +275,11 @@ def get_package_version() -> str: return "0.0.0" -def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: +def create_app( + *, + durable_runtime: DurableRuntime | None = None, + pipeline_registry: PipelineRegistry | None = None, +) -> FastAPI: """ Create and configure a FastAPI application. @@ -312,6 +311,7 @@ def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: app = FastAPI(**app_params) app.state.durable_runtime = durable_runtime or DurableRuntime(app_settings=settings) + app.state.pipeline_registry = pipeline_registry or PipelineRegistry() configure_tracing() app.add_middleware(RequestIdMiddleware) @@ -335,7 +335,7 @@ def create_app(*, durable_runtime: DurableRuntime | None = None) -> FastAPI: app.include_router(draw_router) app.include_router(deploy_router) app.include_router(undeploy_router) - app.include_router(openai_router) + app.include_router(create_openai_router(app.state.pipeline_registry)) app.include_router(dashboard_router, prefix=settings.dashboard_path) _mount_dashboard_ui(app) diff --git a/src/hayhooks/server/pipelines/registry.py b/src/hayhooks/server/pipelines/registry.py index 8a4de4b6..066b59fa 100644 --- a/src/hayhooks/server/pipelines/registry.py +++ b/src/hayhooks/server/pipelines/registry.py @@ -4,7 +4,7 @@ from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper -class _PipelineRegistry: +class PipelineRegistry: """ Registry for pipeline wrappers. @@ -74,4 +74,18 @@ def clear(self) -> None: self._metadata.clear() -registry = _PipelineRegistry() +registry = PipelineRegistry() + + +def get_pipeline_registry(app: Any | None = None) -> PipelineRegistry: + """Return the registry owned by *app*, or the process registry for app-less callers.""" + if app is None: + return registry + + app_registry = getattr(app.state, "pipeline_registry", None) + if isinstance(app_registry, PipelineRegistry): + return app_registry + + app_registry = PipelineRegistry() + app.state.pipeline_registry = app_registry + return app_registry diff --git a/src/hayhooks/server/routers/dashboard.py b/src/hayhooks/server/routers/dashboard.py index 9af9c318..967b44d0 100644 --- a/src/hayhooks/server/routers/dashboard.py +++ b/src/hayhooks/server/routers/dashboard.py @@ -8,7 +8,7 @@ from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field -from hayhooks.server.pipelines.registry import registry +from hayhooks.server.pipelines.registry import get_pipeline_registry from hayhooks.server.utils.live_trace_buffer import clear_live_traces, get_recent_traces from hayhooks.server.utils.live_trace_stream import get_trace_stream_broadcaster from hayhooks.settings import settings @@ -79,8 +79,8 @@ class DashboardUiConfigResponse(BaseModel): summary="List dashboard entry points", description="Returns deployed Hayhooks pipelines used as dashboard entry points.", ) -async def entrypoints() -> EntrypointsResponse: - return EntrypointsResponse(entrypoints=sorted(registry.get_names())) +async def entrypoints(request: Request) -> EntrypointsResponse: + return EntrypointsResponse(entrypoints=sorted(get_pipeline_registry(request.app).get_names())) @router.get( diff --git a/src/hayhooks/server/routers/draw.py b/src/hayhooks/server/routers/draw.py index 5621aa9b..b07c029f 100644 --- a/src/hayhooks/server/routers/draw.py +++ b/src/hayhooks/server/routers/draw.py @@ -1,11 +1,11 @@ import tempfile from pathlib import Path -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Request from fastapi import Path as PathParam from fastapi.responses import FileResponse -from hayhooks.server.pipelines.registry import registry +from hayhooks.server.pipelines.registry import get_pipeline_registry from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper router = APIRouter() @@ -26,9 +26,10 @@ }, ) async def draw( + request: Request, pipeline_name: str = PathParam(description="Name of the pipeline to visualize", examples=["my_pipeline"]), ) -> FileResponse: - pipeline = registry.get(pipeline_name) + pipeline = get_pipeline_registry(request.app).get(pipeline_name) if isinstance(pipeline, BasePipelineWrapper): pipeline = pipeline.pipeline diff --git a/src/hayhooks/server/routers/openai.py b/src/hayhooks/server/routers/openai.py index 8577f3fd..7bcc4dcd 100644 --- a/src/hayhooks/server/routers/openai.py +++ b/src/hayhooks/server/routers/openai.py @@ -1,6 +1,7 @@ import time from collections.abc import AsyncGenerator, Generator from dataclasses import dataclass +from functools import partial from typing import Any from uuid import uuid4 @@ -16,7 +17,7 @@ from haystack.dataclasses import StreamingChunk from hayhooks.server.logger import log -from hayhooks.server.pipelines.registry import registry +from hayhooks.server.pipelines.registry import PipelineRegistry, registry from hayhooks.server.tracing import ( SPAN_OPENAI_FILE_UPLOAD, SPAN_OPENAI_RUN, @@ -60,8 +61,8 @@ class _OpenAIDispatch: ) -def _list_models() -> list[str]: - return registry.get_names() +def _list_models(pipeline_registry: PipelineRegistry = registry) -> list[str]: + return pipeline_registry.get_names() def _chunk_to_text(chunk: Any) -> str: @@ -82,9 +83,9 @@ async def _collect_async_generator(gen: AsyncGenerator) -> str: return "".join([_chunk_to_text(chunk) async for chunk in gen]) -def _resolve_pipeline_wrapper(model: str) -> BasePipelineWrapper: +def _resolve_pipeline_wrapper(model: str, pipeline_registry: PipelineRegistry = registry) -> BasePipelineWrapper: """Look up *model* in the registry, raising 404 if it isn't a pipeline wrapper.""" - pipeline_wrapper = registry.get(model) + pipeline_wrapper = pipeline_registry.get(model) if not isinstance(pipeline_wrapper, BasePipelineWrapper): raise HTTPException(status_code=404, detail=f"Pipeline '{model}' not found or not a pipeline wrapper") return pipeline_wrapper @@ -140,6 +141,7 @@ async def _run_pipeline_method( model: str, kwargs: dict[str, Any], body: dict[str, Any], + pipeline_registry: PipelineRegistry = registry, ) -> str | Generator | AsyncGenerator: """Shared dispatch logic for chat completions and responses endpoints.""" stream_requested = bool(body.get("stream", False)) @@ -153,7 +155,7 @@ async def _run_pipeline_method( ) if stream_requested: try: - wrapper = _resolve_pipeline_wrapper(model) + wrapper = _resolve_pipeline_wrapper(model, pipeline_registry) mode, method_name = _select_execution_mode(wrapper, dispatch) result = await _invoke_pipeline_method( wrapper, mode=mode, method_name=method_name, model=model, call_kwargs={**kwargs, "body": body} @@ -180,7 +182,7 @@ async def _run_pipeline_method( return normalized_result with trace_operation(SPAN_OPENAI_RUN, tags=trace_tags) as span: - wrapper = _resolve_pipeline_wrapper(model) + wrapper = _resolve_pipeline_wrapper(model, pipeline_registry) mode, method_name = _select_execution_mode(wrapper, dispatch) span.set_tag("hayhooks.openai.execution_mode", mode) result = await _invoke_pipeline_method( @@ -190,27 +192,54 @@ async def _run_pipeline_method( async def _run_completion( - model: str, messages: list[dict[str, Any]], body: dict[str, Any] + model: str, + messages: list[dict[str, Any]], + body: dict[str, Any], + *, + pipeline_registry: PipelineRegistry = registry, ) -> str | Generator | AsyncGenerator: - return await _run_pipeline_method(_CHAT_COMPLETION_DISPATCH, model=model, kwargs={"messages": messages}, body=body) + return await _run_pipeline_method( + _CHAT_COMPLETION_DISPATCH, + model=model, + kwargs={"messages": messages}, + body=body, + pipeline_registry=pipeline_registry, + ) async def _run_response( - model: str, input_items: list[dict[str, Any]], body: dict[str, Any] + model: str, + input_items: list[dict[str, Any]], + body: dict[str, Any], + *, + pipeline_registry: PipelineRegistry = registry, ) -> str | Generator | AsyncGenerator: - return await _run_pipeline_method(_RESPONSE_DISPATCH, model=model, kwargs={"input_items": input_items}, body=body) + return await _run_pipeline_method( + _RESPONSE_DISPATCH, + model=model, + kwargs={"input_items": input_items}, + body=body, + pipeline_registry=pipeline_registry, + ) -def _find_file_upload_wrapper() -> BasePipelineWrapper | None: +def _find_file_upload_wrapper(pipeline_registry: PipelineRegistry = registry) -> BasePipelineWrapper | None: """Find the first registered pipeline wrapper that implements ``run_file_upload``.""" - for name in registry.get_names(): - wrapper = registry.get(name) + for name in pipeline_registry.get_names(): + wrapper = pipeline_registry.get(name) if isinstance(wrapper, BasePipelineWrapper) and wrapper._is_run_file_upload_implemented: return wrapper return None -async def _run_file_upload(filename: str | None, content_type: str | None, content: bytes, purpose: str) -> FileObject: +async def _run_file_upload( + filename: str | None, + content_type: str | None, + content: bytes, + purpose: str, + *, + pipeline_registry: PipelineRegistry = registry, +) -> FileObject: with trace_operation( SPAN_OPENAI_FILE_UPLOAD, tags=build_trace_tags( @@ -223,7 +252,7 @@ async def _run_file_upload(filename: str | None, content_type: str | None, conte } ), ): - wrapper = _find_file_upload_wrapper() + wrapper = _find_file_upload_wrapper(pipeline_registry) if wrapper is not None: result = await run_in_threadpool(wrapper.run_file_upload, filename, content_type, content, purpose) if isinstance(result, FileObject): @@ -254,39 +283,35 @@ async def _run_file_upload(filename: str | None, content_type: str | None, conte ) -router = APIRouter() - -router.include_router( - create_models_router( - list_models=_list_models, - owned_by="hayhooks", - tags=["openai"], +def create_openai_router(pipeline_registry: PipelineRegistry = registry) -> APIRouter: + """Create OpenAI-compatible routes bound to one pipeline registry.""" + list_models = partial(_list_models, pipeline_registry) + run_completion = partial(_run_completion, pipeline_registry=pipeline_registry) + run_response = partial(_run_response, pipeline_registry=pipeline_registry) + run_file_upload = partial(_run_file_upload, pipeline_registry=pipeline_registry) + + router = APIRouter() + router.include_router(create_models_router(list_models=list_models, owned_by="hayhooks", tags=["openai"])) + router.include_router( + create_chat_completion_router( + list_models=list_models, + run_completion=run_completion, + owned_by="hayhooks", + tags=["openai"], + include_models_endpoints=False, + ) ) -) - -router.include_router( - create_chat_completion_router( - list_models=_list_models, - run_completion=_run_completion, - owned_by="hayhooks", - tags=["openai"], - include_models_endpoints=False, + router.include_router( + create_responses_router( + list_models=list_models, + run_response=run_response, + owned_by="hayhooks", + tags=["openai"], + include_models_endpoints=False, + ) ) -) + router.include_router(create_files_router(run_file_upload=run_file_upload, tags=["openai"])) + return router -router.include_router( - create_responses_router( - list_models=_list_models, - run_response=_run_response, - owned_by="hayhooks", - tags=["openai"], - include_models_endpoints=False, - ) -) -router.include_router( - create_files_router( - run_file_upload=_run_file_upload, - tags=["openai"], - ) -) +router = create_openai_router() diff --git a/src/hayhooks/server/routers/status.py b/src/hayhooks/server/routers/status.py index b8b8f351..ef6aa696 100644 --- a/src/hayhooks/server/routers/status.py +++ b/src/hayhooks/server/routers/status.py @@ -1,7 +1,7 @@ from fastapi import APIRouter, HTTPException, Request from pydantic import BaseModel, Field -from hayhooks.server.pipelines.registry import registry +from hayhooks.server.pipelines.registry import get_pipeline_registry router = APIRouter() @@ -32,7 +32,7 @@ class PipelineStatusResponse(BaseModel): description="Returns the system status and a list of all available pipelines.", ) async def status_all(request: Request) -> StatusResponse: - pipelines = registry.get_names() + pipelines = get_pipeline_registry(request.app).get_names() durable_health = await request.app.state.durable_runtime.health() return StatusResponse(status="Up!", pipelines=pipelines, durable=durable_health) @@ -46,7 +46,7 @@ async def status_all(request: Request) -> StatusResponse: description="Returns the status of a specific pipeline. Returns 404 if the pipeline doesn't exist.", ) async def status(request: Request, pipeline_name: str) -> PipelineStatusResponse: - if pipeline_name not in registry.get_names(): + if pipeline_name not in get_pipeline_registry(request.app).get_names(): raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") deployment = request.app.state.durable_runtime.current_deployment(pipeline_name) if deployment is not None and not deployment.manager.health["healthy"]: diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index 8cf6a5a1..48ea0679 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -31,7 +31,7 @@ create_response_model_from_callable, get_response_class_from_callable, ) -from hayhooks.server.pipelines.registry import registry +from hayhooks.server.pipelines.registry import PipelineRegistry, get_pipeline_registry from hayhooks.server.pipelines.sse import SSEStream from hayhooks.server.tracing import ( SPAN_PIPELINE_DEPLOY, @@ -59,15 +59,17 @@ # These APIs can be called from multiple event loops, so the locks must be process-wide. _deployment_serial_lock = threading.Lock() _deployment_publication_lock = threading.Lock() -_deployments_in_progress: set[str] = set() +_deployments_in_progress: set[tuple[PipelineRegistry, str]] = set() def _app_runtime(app: FastAPI | None, runtime: DurableRuntime | None = None) -> DurableRuntime | None: - if runtime is not None: - return runtime state = getattr(app, "state", None) if app is not None else None candidate = getattr(state, "durable_runtime", None) - return candidate if isinstance(candidate, DurableRuntime) else None + app_runtime = candidate if isinstance(candidate, DurableRuntime) else None + if runtime is not None and app_runtime is not None and runtime is not app_runtime: + msg = "The explicit DurableRuntime does not belong to the supplied FastAPI app" + raise ValueError(msg) + return runtime or app_runtime def _deployment_candidate( @@ -100,11 +102,18 @@ async def _deployment_lock(lock: threading.Lock): class _DeploymentSnapshot: """Rollback state captured before preparation mutates files or loaded modules.""" - def __init__(self, pipeline_name: str, app: FastAPI | None, runtime: DurableRuntime | None) -> None: + def __init__( + self, + pipeline_name: str, + app: FastAPI | None, + runtime: DurableRuntime | None, + pipeline_registry: PipelineRegistry, + ) -> None: self.pipeline_name = pipeline_name self.app = app - self.wrapper = registry.get(pipeline_name) - metadata = registry.get_metadata(pipeline_name) + self.registry = pipeline_registry + self.wrapper = pipeline_registry.get(pipeline_name) + metadata = pipeline_registry.get_metadata(pipeline_name) self.metadata = dict(metadata) if metadata is not None else None self.runtime = runtime self.deployment = runtime.current_deployment(pipeline_name) if runtime is not None else None @@ -137,13 +146,14 @@ def capture( pipeline_name: str, app: FastAPI | None, runtime: DurableRuntime | None, + pipeline_registry: PipelineRegistry, ) -> "_DeploymentSnapshot": - return cls(pipeline_name, app, runtime) + return cls(pipeline_name, app, runtime, pipeline_registry) def restore_publication(self) -> None: - registry.remove(self.pipeline_name) + self.registry.remove(self.pipeline_name) if self.wrapper is not None: - registry.add(self.pipeline_name, self.wrapper, metadata=dict(self.metadata or {})) + self.registry.add(self.pipeline_name, self.wrapper, metadata=dict(self.metadata or {})) if self.runtime is not None: self.runtime.install_deployment(self.pipeline_name, self.deployment) if self.app is not None and self.routes is not None: @@ -189,6 +199,8 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 ) -> dict[str, str]: """Prepare independently, then atomically publish or restore one pipeline.""" durable_runtime = _app_runtime(app, durable_runtime) + pipeline_registry = get_pipeline_registry(app) + deployment_key = (pipeline_registry, pipeline_name) dlog = log.bind( pipeline_name=pipeline_name, overwrite=overwrite, @@ -208,10 +220,10 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 async with policy_lock: try: async with _deployment_lock(_deployment_publication_lock): - if pipeline_name in _deployments_in_progress: + if deployment_key in _deployments_in_progress: msg = f"Pipeline '{pipeline_name}' is already being deployed" raise PipelineAlreadyExistsError(msg) - snapshot = _DeploymentSnapshot.capture(pipeline_name, app, durable_runtime) + snapshot = _DeploymentSnapshot.capture(pipeline_name, app, durable_runtime, pipeline_registry) if snapshot.wrapper is not None and not overwrite: msg = f"Pipeline '{pipeline_name}' already exists" raise PipelineAlreadyExistsError(msg) @@ -231,7 +243,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 "complete or cancel them before replacing it" ) raise PipelineAlreadyExistsError(msg) - _deployments_in_progress.add(pipeline_name) + _deployments_in_progress.add(deployment_key) registered = True if remove_files_before_prepare: @@ -281,7 +293,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 finally: if registered: async with _deployment_lock(_deployment_publication_lock): - _deployments_in_progress.discard(pipeline_name) + _deployments_in_progress.discard(deployment_key) if snapshot is not None: snapshot.cleanup() @@ -352,15 +364,17 @@ async def undeploy_pipeline_async( ) -> None: """Atomically unpublish a pipeline before stopping its owned resources.""" durable_runtime = _app_runtime(app, durable_runtime) + pipeline_registry = get_pipeline_registry(app) + deployment_key = (pipeline_registry, pipeline_name) policy_lock = ( _deployment_lock(_deployment_serial_lock) if settings.deploy_concurrency == DeployConcurrencyPolicy.SERIALIZED else nullcontext() ) async with policy_lock, _deployment_lock(_deployment_publication_lock): - if pipeline_name in _deployments_in_progress: + if deployment_key in _deployments_in_progress: raise HTTPException(status_code=409, detail=f"Pipeline '{pipeline_name}' is being deployed") - if registry.get(pipeline_name) is None: + if pipeline_registry.get(pipeline_name) is None: raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") deployment = durable_runtime.current_deployment(pipeline_name) if durable_runtime is not None else None if deployment is not None: @@ -752,6 +766,7 @@ def add_pipeline_api_route( - Updates registry metadata with request/response models and file requirement flag """ clog = log.bind(pipeline_name=pipeline_name) + pipeline_registry = get_pipeline_registry(app) deployment = _durable_deployment if deployment is None: runtime = _app_runtime(app) @@ -830,7 +845,7 @@ def add_pipeline_api_route( _defer_openapi_rebuild=True, ) - registry.update_metadata( + pipeline_registry.update_metadata( pipeline_name, { "request_model": RunRequest, @@ -880,9 +895,10 @@ def _register_prepared_pipeline( PipelineAlreadyExistsError: If the pipeline already exists at commit time. """ clog = log.bind(pipeline_name=pipeline_name) + pipeline_registry = get_pipeline_registry(app) # Commit resolves overwrite semantics before registration. - if registry.get(pipeline_name): + if pipeline_registry.get(pipeline_name): msg = f"Pipeline '{pipeline_name}' already exists" raise PipelineAlreadyExistsError(msg) @@ -920,7 +936,7 @@ def _register_prepared_pipeline( # Add wrapper to registry clog.debug("Adding pipeline to registry with metadata: {}", metadata) - registry.add(pipeline_name, pipeline_wrapper, metadata=metadata) + pipeline_registry.add(pipeline_name, pipeline_wrapper, metadata=metadata) clog.success("Pipeline '{}' successfully added to registry", pipeline_name) # Create API route if app is provided @@ -1049,6 +1065,7 @@ def commit_prepared_pipeline( cleanup_files_on_overwrite: If ``True``, remove persisted files when replacing an existing pipeline. """ durable_runtime = _app_runtime(app, durable_runtime) + pipeline_registry = get_pipeline_registry(app) candidate = _durable_deployment if candidate is None: candidate = _deployment_candidate(durable_runtime, prepared.name, prepared.wrapper) @@ -1064,13 +1081,13 @@ def commit_prepared_pipeline( } ), ): - if registry.get(prepared.name) is not None: + if pipeline_registry.get(prepared.name) is not None: if not overwrite: msg = f"Pipeline '{prepared.name}' already exists" raise PipelineAlreadyExistsError(msg) log.bind(pipeline_name=prepared.name).debug("Clearing existing pipeline '{}'", prepared.name) - registry.remove(prepared.name) + pipeline_registry.remove(prepared.name) if cleanup_files_on_overwrite: remove_pipeline_files(prepared.name, settings.pipelines_dir) @@ -1120,6 +1137,7 @@ def deploy_pipeline_files( # noqa: PLR0913 - stable public deployment API PipelineModuleLoadError: If loading the pipeline module fails. PipelineWrapperError: If wrapper creation or setup fails. """ + durable_runtime = _app_runtime(app, durable_runtime) with trace_operation( SPAN_PIPELINE_DEPLOY, tags=build_trace_tags( @@ -1183,6 +1201,7 @@ def deploy_pipeline_yaml( # noqa: PLR0913 - stable public deployment API ValueError: If the YAML cannot be parsed into a Pipeline. InvalidYamlIOError: If the YAML is missing inputs/outputs declarations. """ + durable_runtime = _app_runtime(app, durable_runtime) with trace_operation( SPAN_PIPELINE_DEPLOY, tags=build_trace_tags( @@ -1283,6 +1302,7 @@ def undeploy_pipeline(pipeline_name: str, app: FastAPI | None = None) -> None: Raises: HTTPException: If the pipeline is not found in the registry (404). """ + pipeline_registry = get_pipeline_registry(app) with trace_operation( SPAN_PIPELINE_UNDEPLOY, tags=build_trace_tags( @@ -1294,11 +1314,11 @@ def undeploy_pipeline(pipeline_name: str, app: FastAPI | None = None) -> None: ), ): # Check if pipeline exists in registry - if pipeline_name not in registry.get_names(): + if pipeline_name not in pipeline_registry.get_names(): raise HTTPException(status_code=404, detail=f"Pipeline '{pipeline_name}' not found") # Remove pipeline from registry - registry.remove(pipeline_name) + pipeline_registry.remove(pipeline_name) # Clean up sys.modules for wrapper-based pipelines unload_pipeline_modules(pipeline_name) diff --git a/tests/conftest.py b/tests/conftest.py index 1e719ece..636077eb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -139,7 +139,7 @@ def test_settings(): @pytest.fixture(scope="session", autouse=True) def test_app(): - return create_app() + return create_app(pipeline_registry=registry) @pytest.fixture diff --git a/tests/test_deploy_performance.py b/tests/test_deploy_performance.py index f71fad40..dfc01538 100644 --- a/tests/test_deploy_performance.py +++ b/tests/test_deploy_performance.py @@ -1,6 +1,7 @@ import asyncio import threading from pathlib import Path +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -146,6 +147,7 @@ def synchronized_prepare(*args, **kwargs): def test_defer_openapi_rebuild_skips_setup(): mock_app = MagicMock(spec=FastAPI) mock_app.routes = [] + mock_app.state = SimpleNamespace(pipeline_registry=registry) deploy_pipeline_yaml( pipeline_name="defer_test", @@ -161,6 +163,7 @@ def test_defer_openapi_rebuild_skips_setup(): def test_no_defer_calls_setup(): mock_app = MagicMock(spec=FastAPI) mock_app.routes = [] + mock_app.state = SimpleNamespace(pipeline_registry=registry) deploy_pipeline_yaml( pipeline_name="no_defer_test", diff --git a/tests/test_deploy_utils.py b/tests/test_deploy_utils.py index 132a1754..e43fdad6 100644 --- a/tests/test_deploy_utils.py +++ b/tests/test_deploy_utils.py @@ -4,6 +4,7 @@ import sys from collections.abc import AsyncGenerator, Callable, Generator from pathlib import Path +from types import SimpleNamespace from typing import Any, Literal import docstring_parser @@ -708,6 +709,7 @@ def test_deploy_pipeline_files_without_adding_api_route(test_settings, mocker): def test_deploy_pipeline_files_skip_mcp(mocker): mock_app = mocker.Mock() mock_app.routes = [] + mock_app.state = SimpleNamespace(pipeline_registry=registry) # This pipeline wrapper has skip_mcp class attribute set to True test_file_path = Path("tests/test_files/files/chat_with_website_mcp_skip/pipeline_wrapper.py") diff --git a/tests/test_durable_deployment_lifecycle.py b/tests/test_durable_deployment_lifecycle.py index 633c2d09..acd30d96 100644 --- a/tests/test_durable_deployment_lifecycle.py +++ b/tests/test_durable_deployment_lifecycle.py @@ -7,6 +7,7 @@ from hayhooks.durable.runtime import DurableRuntime from hayhooks.server.app import create_app from hayhooks.server.pipelines.registry import registry +from hayhooks.server.utils.deploy_utils import deploy_pipeline_files from hayhooks.settings import settings pytestmark = pytest.mark.skipif( @@ -268,7 +269,7 @@ def test_undeploy_refuses_to_strand_thread_backed_work(monkeypatch) -> None: app = create_app() with TestClient(app) as client: assert _deploy(client, _blocking_source()).status_code == 200 - old_wrapper = registry.get("job") + old_wrapper = app.state.pipeline_registry.get("job") assert old_wrapper is not None submitted = client.post("/job/run-durable", json={"value": 2}) assert old_wrapper.started.wait(timeout=1) @@ -284,6 +285,42 @@ def test_undeploy_refuses_to_strand_thread_backed_work(monkeypatch) -> None: assert client.post("/undeploy/job").status_code == 200 +def test_fastapi_apps_own_pipeline_publication_and_allow_the_same_name() -> None: + app_a = create_app() + app_b = create_app() + + with TestClient(app_a) as client_a, TestClient(app_b) as client_b: + assert _deploy(client_a, _durable_source(field="value", increment=1, revision="app-a")).status_code == 200 + + assert client_a.get("/status").json()["pipelines"] == ["job"] + assert client_b.get("/status").json()["pipelines"] == [] + assert client_b.get("/status/job").status_code == 404 + assert client_b.post("/job/run-durable", json={"value": 1}).status_code == 404 + + assert _deploy(client_b, _durable_source(field="value", increment=10, revision="app-b")).status_code == 200 + submitted_a = client_a.post("/job/run-durable", json={"value": 1}) + submitted_b = client_b.post("/job/run-durable", json={"value": 1}) + assert _wait_for_completion(client_a, submitted_a)["result"] == {"value": 2} + assert _wait_for_completion(client_b, submitted_b)["result"] == {"value": 11} + + +def test_deploy_rejects_a_runtime_owned_by_another_app() -> None: + app = create_app() + other_runtime = DurableRuntime(app_settings=settings) + + with pytest.raises(ValueError, match="does not belong"): + deploy_pipeline_files( + "job", + {"pipeline_wrapper.py": _api_source()}, + app=app, + save_files=False, + durable_runtime=other_runtime, + ) + + assert app.state.pipeline_registry.get("job") is None + assert all(getattr(route, "path", None) != "/job/run" for route in app.routes) + + def test_durable_to_non_durable_overwrite_removes_control_routes() -> None: app = create_app() with TestClient(app) as client: diff --git a/tests/test_registry.py b/tests/test_registry.py index 6f14c1bf..02a2850d 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -4,14 +4,14 @@ from haystack import Document, Pipeline from hayhooks.server.exceptions import PipelineNotFoundError -from hayhooks.server.pipelines.registry import _PipelineRegistry +from hayhooks.server.pipelines.registry import PipelineRegistry from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper from hayhooks.server.utils.yaml_pipeline_wrapper import YAMLPipelineWrapper @pytest.fixture def pipeline_registry(): - return _PipelineRegistry() + return PipelineRegistry() @pytest.fixture From e6360a065175d65643e5a6ac77ed5956d0bdce0c Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 19 Aug 2026 11:06:13 +0200 Subject: [PATCH 19/28] refactoring --- src/hayhooks/durable/__init__.py | 161 ++-------- src/hayhooks/durable/adapters.py | 15 +- src/hayhooks/durable/backend.py | 54 +--- src/hayhooks/durable/engine.py | 23 +- src/hayhooks/durable/fastapi.py | 112 +++---- src/hayhooks/durable/manager.py | 3 - src/hayhooks/durable/models.py | 58 ++-- src/hayhooks/durable/redis.py | 46 +-- src/hayhooks/durable/reference.py | 9 +- src/hayhooks/durable/runtime.py | 42 +-- src/hayhooks/durable/settings.py | 35 +- src/hayhooks/durable/store.py | 57 +--- src/hayhooks/server/a2a/durable_executor.py | 204 +++++------- src/hayhooks/server/a2a/imports.py | 27 +- src/hayhooks/server/a2a/messages.py | 8 +- src/hayhooks/server/a2a/redis_task_store.py | 79 ++--- src/hayhooks/server/pipelines/models.py | 14 +- src/hayhooks/server/pipelines/registry.py | 12 +- src/hayhooks/server/routers/__init__.py | 3 +- src/hayhooks/server/routers/openai.py | 21 +- src/hayhooks/server/tracing.py | 1 - src/hayhooks/server/utils/deploy_utils.py | 31 +- src/hayhooks/settings.py | 49 +-- tests/durable_contract.py | 15 + tests/durable_helpers.py | 64 ++++ tests/test_a2a.py | 83 ++--- tests/test_cli.py | 162 ++++------ tests/test_durable_a2a.py | 91 ++---- tests/test_durable_deployment_lifecycle.py | 94 ++---- tests/test_durable_execution.py | 318 +++++++------------ tests/test_durable_fastapi.py | 15 +- tests/test_durable_process_recovery.py | 46 ++- tests/test_durable_reference.py | 29 +- tests/test_durable_store.py | 114 +++---- tests/test_it_a2a_server.py | 75 ++--- tests/test_redis_a2a_recovery_integration.py | 22 +- tests/test_redis_execution_integration.py | 26 +- tests/test_settings.py | 94 +++--- 38 files changed, 868 insertions(+), 1444 deletions(-) create mode 100644 tests/durable_helpers.py diff --git a/src/hayhooks/durable/__init__.py b/src/hayhooks/durable/__init__.py index b349bbd2..cc206e7f 100644 --- a/src/hayhooks/durable/__init__.py +++ b/src/hayhooks/durable/__init__.py @@ -1,63 +1,31 @@ -"""Advanced durable execution contracts and safe public result models.""" +"""Public durable execution API.""" from __future__ import annotations -from datetime import datetime -from typing import TYPE_CHECKING, Any - -from pydantic import BaseModel, Field - -from hayhooks.durable.context import get_current_durable_context +from hayhooks.durable.context import DurableContext, get_current_durable_context +from hayhooks.durable.fastapi import create_durable_router from hayhooks.durable.mode import DurableAuthoringMode, durable_authoring_mode -from hayhooks.durable.models import ExecutionStatus - -if TYPE_CHECKING: - from hayhooks.durable.context import DurableContext - from hayhooks.durable.fastapi import create_durable_router - from hayhooks.durable.models import ( - ExecutionAdmissionError, - ExecutionCanceledError, - ExecutionRecordSizeError, - ExecutionStoreError, - ExecutionSuspendedError, - RetryableExecutionError, - ) - from hayhooks.durable.runtime import ( - DefinitionRevisionConflictError, - DurableDeployment, - DurableRuntime, - ExecutionStoreProvider, - IdempotencyConflictError, - ) - from hayhooks.durable.settings import DurableSettings - from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider - - -class ExecutionProgress(BaseModel): - """Sanitized client-visible progress event.""" - - sequence: int - kind: str - message: str - timestamp: datetime - metadata: dict[str, Any] = Field(default_factory=dict) - - -class ExecutionResult(BaseModel): - """Safe durable REST/A2A execution projection.""" - - execution_id: str - status: ExecutionStatus - attempt: int - sequence: int - progress: list[ExecutionProgress] - result: Any | None = None - error: dict[str, Any] | None = None - waiting: dict[str, Any] | None = None - cancellation_requested_at: datetime | None = None - created_at: datetime - updated_at: datetime - links: dict[str, str] = Field(default_factory=dict) +from hayhooks.durable.models import ( + ExecutionAdmissionError, + ExecutionCanceledError, + ExecutionProgress, + ExecutionRecordSizeError, + ExecutionResult, + ExecutionStatus, + ExecutionStoreError, + ExecutionSuspendedError, + RetryableExecutionError, +) +from hayhooks.durable.runtime import ( + DefinitionRevisionConflictError, + DurableDeployment, + DurableRuntime, + ExecutionStoreProvider, + IdempotencyConflictError, + durable_runtime, +) +from hayhooks.durable.settings import DurableSettings +from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider def current_execution_id() -> str | None: @@ -66,86 +34,7 @@ def current_execution_id() -> str | None: return context.execution_id if context is not None else None -def current_durable_context() -> Any | None: - """Return the active context for advanced hooks and tools.""" - return get_current_durable_context() - - -def __getattr__(name: str) -> Any: - """Lazily expose durable infrastructure without eager optional imports.""" - if name == "DurableContext": - from hayhooks.durable.context import DurableContext - - return DurableContext - if name == "create_durable_router": - from hayhooks.durable.fastapi import create_durable_router - - return create_durable_router - if name == "DurableSettings": - from hayhooks.durable.settings import DurableSettings - - return DurableSettings - if name in { - "DefinitionRevisionConflictError", - "DurableDeployment", - "DurableRuntime", - "ExecutionStoreProvider", - "IdempotencyConflictError", - "durable_runtime", - }: - from hayhooks.durable.runtime import ( - DefinitionRevisionConflictError, - DurableDeployment, - DurableRuntime, - ExecutionStoreProvider, - IdempotencyConflictError, - durable_runtime, - ) - - return { - "DefinitionRevisionConflictError": DefinitionRevisionConflictError, - "DurableDeployment": DurableDeployment, - "DurableRuntime": DurableRuntime, - "ExecutionStoreProvider": ExecutionStoreProvider, - "IdempotencyConflictError": IdempotencyConflictError, - "durable_runtime": durable_runtime, - }[name] - if name in { - "ExecutionAdmissionError", - "ExecutionCanceledError", - "ExecutionRecordSizeError", - "ExecutionStoreError", - "ExecutionSuspendedError", - "RetryableExecutionError", - }: - from hayhooks.durable.models import ( - ExecutionAdmissionError, - ExecutionCanceledError, - ExecutionRecordSizeError, - ExecutionStoreError, - ExecutionSuspendedError, - RetryableExecutionError, - ) - - return { - "ExecutionAdmissionError": ExecutionAdmissionError, - "ExecutionCanceledError": ExecutionCanceledError, - "ExecutionRecordSizeError": ExecutionRecordSizeError, - "ExecutionStoreError": ExecutionStoreError, - "ExecutionSuspendedError": ExecutionSuspendedError, - "RetryableExecutionError": RetryableExecutionError, - }[name] - if name in {"ExecutionStore", "InMemoryExecutionStoreProvider", "RedisExecutionStoreProvider"}: - from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider - - return { - "ExecutionStore": ExecutionStore, - "InMemoryExecutionStoreProvider": InMemoryExecutionStoreProvider, - "RedisExecutionStoreProvider": RedisExecutionStoreProvider, - }[name] - msg = f"module {__name__!r} has no attribute {name!r}" - raise AttributeError(msg) - +current_durable_context = get_current_durable_context __all__ = [ "DefinitionRevisionConflictError", diff --git a/src/hayhooks/durable/adapters.py b/src/hayhooks/durable/adapters.py index 41f820e3..e8fffc39 100644 --- a/src/hayhooks/durable/adapters.py +++ b/src/hayhooks/durable/adapters.py @@ -267,6 +267,13 @@ async def checkpoint_after_run_async(state: State) -> None: self.pipeline._hayhooks_durable_hooks_installed = True +def chat_messages(values: Any) -> list[Any]: + """Decode a persisted list of serialized chat messages.""" + if not isinstance(values, list): + return [] + return [ChatMessage.from_dict(value) for value in values if isinstance(value, dict)] + + def _current_durable_context() -> DurableContext | None: from hayhooks.durable.context import get_current_durable_context @@ -366,10 +373,8 @@ def _restore_agent_state(context: DurableContext, state: Any) -> None: if live_hook_context is not None: state.data["hook_context"] = live_hook_context resume = context.take_resume_input() - if isinstance(resume, dict) and isinstance(resume.get("messages"), list): - state.data.setdefault("messages", []).extend( - ChatMessage.from_dict(message) for message in resume["messages"] if isinstance(message, dict) - ) + if isinstance(resume, dict): + state.data.setdefault("messages", []).extend(chat_messages(resume.get("messages"))) def execution_kind(pipeline: Any) -> ExecutionKind: @@ -383,4 +388,4 @@ def execution_kind(pipeline: Any) -> ExecutionKind: raise TypeError(msg) -__all__ = ["HaystackDurableAdapter", "execution_kind", "require_haystack_v3"] +__all__ = ["HaystackDurableAdapter", "chat_messages", "execution_kind", "require_haystack_v3"] diff --git a/src/hayhooks/durable/backend.py b/src/hayhooks/durable/backend.py index 53675b46..f9937276 100644 --- a/src/hayhooks/durable/backend.py +++ b/src/hayhooks/durable/backend.py @@ -9,19 +9,10 @@ from typing import Any, Protocol from hayhooks.durable.engine import ( - Checkpoint, - Complete, ExecutionCommand, ExecutionControl, ExecutionPayloadSizeError, - Fail, - Heartbeat, PayloadKind, - ReleaseClaim, - RequestCancellation, - Resume, - ScheduleRetry, - Suspend, TransitionPlan, validate_run_id, ) @@ -113,6 +104,11 @@ async def maintain(self, command_factory: Callable[[int, int], ExecutionCommand] async def operational_counts(self) -> dict[str, int]: ... +def runnable_score(control: ExecutionControl) -> int: + """Score a queued control in the runnable index; both backends must agree.""" + return control.available_at_ms if control.available_at_ms is not None else control.updated_at_ms + + def parse_idempotency_binding(value: str) -> tuple[str, str]: """Decode the execution ID and request digest stored for an idempotency key.""" run_id, separator, binding = value.partition("|") @@ -136,7 +132,7 @@ def parse_lease_member(value: str) -> tuple[str, int]: def bind_command(command: ExecutionCommand, *, now_ms: int, lease_commit_safety_ms: int) -> ExecutionCommand: """Apply the backend clock and lease safety policy before reduction.""" - if isinstance(command, (ReleaseClaim, Heartbeat, Checkpoint, ScheduleRetry, Suspend, Complete, Fail)): + if "lease_commit_safety_ms" in command.__dataclass_fields__: return replace(command, now_ms=now_ms, lease_commit_safety_ms=lease_commit_safety_ms) return replace(command, now_ms=now_ms) @@ -153,37 +149,13 @@ def check_admission(raw: Mapping[Any, Any], config: ExecutionStoreConfig) -> Non def validate_command_payloads(command: ExecutionCommand, config: ExecutionStoreConfig) -> None: """Reject command payloads before either backend attempts a transition.""" - checks: tuple[tuple[str, bytes | None, int], ...] = () - if isinstance(command, Checkpoint): - checks = ( - ("checkpoint", command.payload, config.max_checkpoint_bytes), - *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), - ) - elif isinstance(command, Suspend): - checks = ( - ("checkpoint", command.checkpoint, config.max_checkpoint_bytes), - ("wait", command.wait, config.max_wait_bytes), - *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), - ) - elif isinstance(command, Resume): - checks = ( - ("checkpoint", command.checkpoint, config.max_checkpoint_bytes), - *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), - ) - elif isinstance(command, RequestCancellation): - checks = tuple(("progress", event, config.max_progress_event_bytes) for event in command.progress_events) - elif isinstance(command, Complete): - checks = ( - ("result", command.result, config.max_result_bytes), - *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), - ) - elif isinstance(command, Fail): - checks = ( - ("error", command.error, config.max_error_bytes), - *(("progress", event, config.max_progress_event_bytes) for event in command.progress_events), - ) - elif isinstance(command, ScheduleRetry): - checks = (("error", command.error, config.max_error_bytes),) + checks = [ + (name, getattr(command, name, None), getattr(config, f"max_{name}_bytes")) + for name in ("checkpoint", "wait", "result", "error") + ] + checks += [ + ("progress", event, config.max_progress_event_bytes) for event in getattr(command, "progress_events", ()) + ] for label, payload, limit in checks: if payload is not None and len(payload) > limit: raise ExecutionPayloadSizeError(f"{label} payload exceeds its configured byte limit") diff --git a/src/hayhooks/durable/engine.py b/src/hayhooks/durable/engine.py index 3d1cff34..3fbf2a9e 100644 --- a/src/hayhooks/durable/engine.py +++ b/src/hayhooks/durable/engine.py @@ -197,7 +197,7 @@ class Checkpoint: worker_id: str now_ms: int lease_duration_ms: int - payload: bytes + checkpoint: bytes progress_events: tuple[bytes, ...] = () lease_commit_safety_ms: int = 0 @@ -330,9 +330,10 @@ def decide(control: ExecutionControl, command: ExecutionCommand) -> TransitionPl return _release_claim(control, command) if isinstance(command, Heartbeat): _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + deadline = command.now_ms + command.lease_duration_ms return TransitionPlan( - replace(control, lease_expires_at_ms=command.now_ms + command.lease_duration_ms), - lease_index_update=LeaseIndexUpdate(command.now_ms + command.lease_duration_ms, control.fence), + replace(control, lease_expires_at_ms=deadline), + lease_index_update=LeaseIndexUpdate(deadline, control.fence), ) if isinstance(command, Checkpoint): _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) @@ -344,7 +345,7 @@ def decide(control: ExecutionControl, command: ExecutionCommand) -> TransitionPl ) return TransitionPlan( next_control, - payload_writes=(PayloadWrite(PayloadKind.CHECKPOINT, command.payload),), + payload_writes=(PayloadWrite(PayloadKind.CHECKPOINT, command.checkpoint),), progress_events=_progress_events(control.progress_sequence, command.progress_events), lease_index_update=LeaseIndexUpdate(next_control.lease_expires_at_ms, control.fence), ) @@ -356,21 +357,17 @@ def decide(control: ExecutionControl, command: ExecutionCommand) -> TransitionPl return _suspend(control, command) if isinstance(command, Resume): return _resume(control, command) - if isinstance(command, Complete): + if isinstance(command, (Complete, Fail)): _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) + completed = isinstance(command, Complete) return _terminal_or_canceled( control, command.now_ms, - ExecutionStatus.COMPLETED, - PayloadKind.RESULT, - command.result, + ExecutionStatus.COMPLETED if completed else ExecutionStatus.FAILED, + PayloadKind.RESULT if completed else PayloadKind.ERROR, + command.result if isinstance(command, Complete) else command.error, command.progress_events, ) - if isinstance(command, Fail): - _owned(control, command.fence, command.worker_id, command.now_ms, command.lease_commit_safety_ms) - return _terminal_or_canceled( - control, command.now_ms, ExecutionStatus.FAILED, PayloadKind.ERROR, command.error, command.progress_events - ) if isinstance(command, RecoverExpiredLease): return _recover(control, command) raise TypeError(f"unsupported execution command {type(command).__name__}") diff --git a/src/hayhooks/durable/fastapi.py b/src/hayhooks/durable/fastapi.py index 84d867ee..0f1bd97a 100644 --- a/src/hayhooks/durable/fastapi.py +++ b/src/hayhooks/durable/fastapi.py @@ -4,14 +4,14 @@ import inspect from collections.abc import Awaitable, Callable +from functools import wraps from typing import Annotated, Any, cast from fastapi import APIRouter, Body, Depends, Header, HTTPException, Path, Request, Response, status from pydantic import ValidationError, create_model -from hayhooks.durable import ExecutionResult from hayhooks.durable.engine import RUN_ID_PATTERN -from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError +from hayhooks.durable.models import ExecutionAdmissionError, ExecutionResult, ExecutionStoreError from hayhooks.durable.runtime import DefinitionRevisionConflictError, DurableDeployment, IdempotencyConflictError _MAX_OWNER_LENGTH = 512 @@ -35,6 +35,34 @@ def _validated_owner(owner_id: Any, *, enforce_owner: bool) -> str | None: return owner_id +def _translate_errors(handler: Callable[..., Awaitable[Any]]) -> Callable[..., Awaitable[Any]]: + """Map durable domain failures onto the HTTP contract shared by every handler.""" + + @wraps(handler) + async def wrapper(*args: Any, **kwargs: Any) -> Any: + try: + return await handler(*args, **kwargs) + except KeyError as error: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error + except (IdempotencyConflictError, DefinitionRevisionConflictError) as error: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error + except (ValidationError, ValueError) as error: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error + except ExecutionAdmissionError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=str(error), + headers={"Retry-After": str(error.retry_after_seconds)}, + ) from error + except ExecutionStoreError as error: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Durable execution store is unavailable", + ) from error + + return wrapper + + def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResult]: if deployment.result_type is Any: return ExecutionResult @@ -45,7 +73,7 @@ def _durable_response_model(deployment: DurableDeployment) -> type[ExecutionResu ) -def create_durable_router( # noqa: C901, PLR0915 - handlers share one deployment and generated models +def create_durable_router( # noqa: C901 - one factory owns every generated route for a deployment deployment: DurableDeployment, *, owner_id_dependency: OwnerIdDependency | None, @@ -118,17 +146,14 @@ async def submit( execution_id=idempotency_key, owner_id=owner_id, ) - except (IdempotencyConflictError, DefinitionRevisionConflictError) as error: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error - except (ValidationError, ValueError) as error: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error - except ExecutionAdmissionError as error: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail=str(error), - headers={"Retry-After": str(error.retry_after_seconds)}, - ) from error - except (ExecutionStoreError, RuntimeError) as error: + except ( + IdempotencyConflictError, + DefinitionRevisionConflictError, + ExecutionAdmissionError, + ExecutionStoreError, + ): + raise # the shared translator owns these status codes + except RuntimeError as error: raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(error)) from error response.status_code = status.HTTP_200_OK if not created and record.terminal else status.HTTP_202_ACCEPTED result = execution_result(request, record, model=response_model) @@ -144,16 +169,8 @@ async def inspect_execution( request: Request, owner_id: Any = Depends(owner_dependency), # noqa: B008 ) -> ExecutionResult: - try: - owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) - return execution_result(request, await get_execution(execution_id, owner_id)) - except KeyError as error: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error - except ExecutionStoreError as error: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Durable execution store is unavailable", - ) from error + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + return execution_result(request, await get_execution(execution_id, owner_id)) async def cancel_execution( execution_id: ExecutionId, @@ -161,21 +178,9 @@ async def cancel_execution( request: Request, owner_id: Any = Depends(owner_dependency), # noqa: B008 ) -> ExecutionResult: - try: - owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) - accepted = await deployment.request_cancel( - execution_id, - owner_id=owner_id, - enforce_owner=enforce_owner, - ) - record = await get_execution(execution_id, owner_id) - except KeyError as error: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error - except ExecutionStoreError as error: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Durable execution store is unavailable", - ) from error + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + accepted = await deployment.request_cancel(execution_id, owner_id=owner_id, enforce_owner=enforce_owner) + record = await get_execution(execution_id, owner_id) response.status_code = status.HTTP_202_ACCEPTED if accepted else status.HTTP_200_OK return execution_result(request, record) @@ -186,30 +191,11 @@ async def resume_execution( owner_id: Any = Depends(owner_dependency), # noqa: B008 update: Any = Body(default=None), # noqa: B008 ) -> ExecutionResult: - try: - owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) - resumed = await deployment.resume( - execution_id, - update, - owner_id=owner_id, - enforce_owner=enforce_owner, - ) - if not resumed: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Execution is not waiting") - record = await get_execution(execution_id, owner_id) - except KeyError as error: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Execution not found") from error - except DefinitionRevisionConflictError as error: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(error)) from error - except (ValidationError, ValueError) as error: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(error)) from error - except ExecutionStoreError as error: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Durable execution store is unavailable", - ) from error + owner_id = _validated_owner(owner_id, enforce_owner=enforce_owner) + if not await deployment.resume(execution_id, update, owner_id=owner_id, enforce_owner=enforce_owner): + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Execution is not waiting") response.status_code = status.HTTP_202_ACCEPTED - return execution_result(request, record) + return execution_result(request, await get_execution(execution_id, owner_id)) if deployment.resume_type is not None: resume_execution.__annotations__["update"] = deployment.resume_type @@ -230,7 +216,7 @@ async def resume_execution( ): router.add_api_route( path, - endpoint, + _translate_errors(endpoint), methods=methods, name=name, response_model=response_model if endpoint is submit else ExecutionResult, diff --git a/src/hayhooks/durable/manager.py b/src/hayhooks/durable/manager.py index a6b2ea2e..255a9bb6 100644 --- a/src/hayhooks/durable/manager.py +++ b/src/hayhooks/durable/manager.py @@ -405,7 +405,6 @@ async def _process_claim(self, claim: ExecutionClaim) -> None: # noqa: C901, PL claim.record.error = None claim.record.status = ExecutionStatus.COMPLETED claim.record.wait = None - claim.record.retry_at = None claim.record.touch() await claim.complete() @@ -492,8 +491,6 @@ async def _fail_oversized_record(self, claim: ExecutionClaim) -> None: record.progress = [] record.result = None record.error = None - record.last_retry_error = None - record.retry_at = None record.mark_failed( ExecutionError( type="ExecutionRecordTooLarge", diff --git a/src/hayhooks/durable/models.py b/src/hayhooks/durable/models.py index 80824f54..6dca34e3 100644 --- a/src/hayhooks/durable/models.py +++ b/src/hayhooks/durable/models.py @@ -16,6 +16,8 @@ from enum import Enum from typing import Any, TypeAlias, cast +from pydantic import BaseModel, Field + from hayhooks.durable.engine import ExecutionLeaseLostError as _ExecutionLeaseLostError from hayhooks.durable.engine import ExecutionStatus, normalize_cancellation_reason @@ -117,6 +119,33 @@ def __init__(self, message: str, *, delay: float = 0.0) -> None: self.delay = max(0.0, delay) +class ExecutionProgress(BaseModel): + """Sanitized client-visible progress event.""" + + sequence: int + kind: str + message: str + timestamp: datetime + metadata: dict[str, Any] = Field(default_factory=dict) + + +class ExecutionResult(BaseModel): + """Safe durable REST/A2A execution projection.""" + + execution_id: str + status: ExecutionStatus + attempt: int + sequence: int + progress: list[ExecutionProgress] + result: Any | None = None + error: dict[str, Any] | None = None + waiting: dict[str, Any] | None = None + cancellation_requested_at: datetime | None = None + created_at: datetime + updated_at: datetime + links: dict[str, str] = Field(default_factory=dict) + + class ExecutionKind(str, Enum): """The Haystack object adapter selected for an execution.""" @@ -239,8 +268,6 @@ class ExecutionRecord: progress: list[ExecutionProgressEvent] = field(default_factory=list) result: JsonValue | None = None error: ExecutionError | None = None - last_retry_error: ExecutionError | None = None - retry_at: datetime | None = None cancel_requested_at: datetime | None = None cancel_reason: str | None = None created_at: datetime = field(default_factory=utc_now) @@ -259,29 +286,18 @@ def __post_init__(self) -> None: if self.cancel_requested_at is not None: self.cancel_requested_at = _as_utc(self.cancel_requested_at) self.cancel_reason = normalize_cancellation_reason(self.cancel_reason) - self.validated_input = cast( - dict[str, JsonValue], - validate_json(self.validated_input, limit=self.max_record_bytes, label="validated input"), - ) - self.application_state = cast( - dict[str, JsonValue], - validate_json(self.application_state, limit=self.max_record_bytes, label="application state"), - ) + self.validated_input = self._bounded_dict(self.validated_input, "validated input") + self.application_state = self._bounded_dict(self.application_state, "application state") if self.wait is not None: - self.wait = cast(dict[str, JsonValue], validate_json(self.wait, limit=self.max_record_bytes, label="wait")) + self.wait = self._bounded_dict(self.wait, "wait") if self.result is not None: self.result = validate_json(self.result, limit=self.max_record_bytes, label="result") - if self.checkpoint is not None and not isinstance(self.checkpoint, ExecutionCheckpoint): - self.checkpoint = ExecutionCheckpoint.from_dict(cast(Mapping[str, Any], self.checkpoint)) if self.checkpoint is not None: - self.checkpoint.data = cast( - dict[str, JsonValue], - validate_json(self.checkpoint.data, limit=self.max_record_bytes, label="checkpoint"), - ) + if not isinstance(self.checkpoint, ExecutionCheckpoint): + self.checkpoint = ExecutionCheckpoint.from_dict(cast(Mapping[str, Any], self.checkpoint)) + self.checkpoint.data = self._bounded_dict(self.checkpoint.data, "checkpoint") if self.error is not None and not isinstance(self.error, ExecutionError): self.error = ExecutionError.from_dict(cast(Mapping[str, Any], self.error)) - if self.last_retry_error is not None and not isinstance(self.last_retry_error, ExecutionError): - self.last_retry_error = ExecutionError.from_dict(cast(Mapping[str, Any], self.last_retry_error)) self.progress = [ event if isinstance(event, ExecutionProgressEvent) @@ -294,6 +310,9 @@ def __post_init__(self) -> None: def terminal(self) -> bool: return self.status.terminal + def _bounded_dict(self, value: Any, label: str) -> dict[str, JsonValue]: + return cast(dict[str, JsonValue], validate_json(value, limit=self.max_record_bytes, label=label)) + def touch(self) -> None: self.sequence += 1 self.updated_at = utc_now() @@ -320,7 +339,6 @@ def mark_canceled(self) -> None: self.status = ExecutionStatus.CANCELED self.error = None self.result = None - self.retry_at = None self.wait = None self.touch() diff --git a/src/hayhooks/durable/redis.py b/src/hayhooks/durable/redis.py index a22b303c..697ac2a6 100644 --- a/src/hayhooks/durable/redis.py +++ b/src/hayhooks/durable/redis.py @@ -25,6 +25,7 @@ check_admission, parse_idempotency_binding, parse_lease_member, + runnable_score, validate_command_payloads, ) from hayhooks.durable.engine import ( @@ -71,23 +72,11 @@ def capacity(self) -> str: def control(self, run_id: str) -> str: return f"{self._execution_base(run_id)}:control" - def input(self, run_id: str) -> str: - return f"{self._execution_base(run_id)}:input" - - def checkpoint(self, run_id: str) -> str: - return f"{self._execution_base(run_id)}:checkpoint" - - def result(self, run_id: str) -> str: - return f"{self._execution_base(run_id)}:result" - - def error(self, run_id: str) -> str: - return f"{self._execution_base(run_id)}:error" - def progress(self, run_id: str) -> str: return f"{self._execution_base(run_id)}:progress" - def wait(self, run_id: str) -> str: - return f"{self._execution_base(run_id)}:wait" + def payload(self, run_id: str, kind: PayloadKind) -> str: + return f"{self._execution_base(run_id)}:{kind.value}" def idempotency(self, idempotency_digest: str) -> str: if not re.fullmatch(r"[a-f0-9]{64}", idempotency_digest): @@ -265,7 +254,7 @@ async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> di return {} async with self.redis.pipeline(transaction=False) as pipe: for kind in kinds: - pipe.get(self._payload_key(run_id, kind)) + pipe.get(self.keys.payload(run_id, kind)) values = await pipe.execute() return {kind: bytes(value) if value is not None else None for kind, value in zip(kinds, values, strict=True)} @@ -310,7 +299,7 @@ async def transition( # noqa: C901 pipe.multi() pipe.zrem(self.keys.runnable, run_id) if current.status is ExecutionStatus.QUEUED: - pipe.zadd(self.keys.runnable, {run_id: _runnable_score(current)}) + pipe.zadd(self.keys.runnable, {run_id: runnable_score(current)}) await pipe.execute() return TransitionPlan(current) pipe.multi() @@ -379,16 +368,16 @@ def _apply_plan( # noqa: C901 if removed_fields: pipe.hdel(self.keys.control(next_control.run_id), *removed_fields) for write in plan.payload_writes: - pipe.set(self._payload_key(next_control.run_id, write.kind), write.data) + pipe.set(self.keys.payload(next_control.run_id, write.kind), write.data) for kind in plan.payload_deletes: - pipe.delete(self._payload_key(next_control.run_id, kind)) + pipe.delete(self.keys.payload(next_control.run_id, kind)) for event in plan.progress_events: pipe.rpush(self.keys.progress(next_control.run_id), event.data) pipe.ltrim(self.keys.progress(next_control.run_id), -self.config.max_progress_events, -1) pipe.zrem(self.keys.runnable, next_control.run_id) if next_control.status is ExecutionStatus.QUEUED: - pipe.zadd(self.keys.runnable, {next_control.run_id: _runnable_score(next_control)}) + pipe.zadd(self.keys.runnable, {next_control.run_id: runnable_score(next_control)}) if plan.lease_index_update is not None: member = RedisKeys.lease_member(next_control.run_id, plan.lease_index_update.fence) @@ -418,27 +407,10 @@ async def _time_ms(self) -> int: def _execution_keys(self, run_id: str) -> tuple[str, ...]: return ( self.keys.control(run_id), - self.keys.input(run_id), - self.keys.checkpoint(run_id), - self.keys.result(run_id), - self.keys.error(run_id), self.keys.progress(run_id), - self.keys.wait(run_id), + *(self.keys.payload(run_id, kind) for kind in PayloadKind), ) - def _payload_key(self, run_id: str, kind: PayloadKind) -> str: - return { - PayloadKind.INPUT: self.keys.input, - PayloadKind.CHECKPOINT: self.keys.checkpoint, - PayloadKind.RESULT: self.keys.result, - PayloadKind.ERROR: self.keys.error, - PayloadKind.WAIT: self.keys.wait, - }[kind](run_id) - - -def _runnable_score(control: ExecutionControl) -> int: - return control.available_at_ms if control.available_at_ms is not None else control.updated_at_ms - def _text(value: str | bytes | int) -> str: return value.decode() if isinstance(value, bytes) else str(value) diff --git a/src/hayhooks/durable/reference.py b/src/hayhooks/durable/reference.py index e847a1dc..e50b898b 100644 --- a/src/hayhooks/durable/reference.py +++ b/src/hayhooks/durable/reference.py @@ -16,6 +16,7 @@ check_admission, parse_idempotency_binding, parse_lease_member, + runnable_score, validate_command_payloads, ) from hayhooks.durable.engine import ( @@ -96,7 +97,7 @@ async def transition(self, run_id: str, command: ExecutionCommand, *, candidate: raise self._runnable.pop(run_id, None) if current.status is ExecutionStatus.QUEUED: - self._runnable[run_id] = _runnable_score(current) + self._runnable[run_id] = runnable_score(current) return TransitionPlan(current) self._apply_plan(current, plan) return plan @@ -157,7 +158,7 @@ def _apply_plan(self, current: ExecutionControl, plan: TransitionPlan) -> None: self._runnable.pop(next_control.run_id, None) if next_control.status is ExecutionStatus.QUEUED: - self._runnable[next_control.run_id] = _runnable_score(next_control) + self._runnable[next_control.run_id] = runnable_score(next_control) if plan.lease_index_update is not None: member = f"{next_control.run_id}|{plan.lease_index_update.fence}" @@ -175,7 +176,3 @@ def _apply_plan(self, current: ExecutionControl, plan: TransitionPlan) -> None: next_control.idempotency_digest, f"{next_control.run_id}|{next_control.idempotency_binding_digest}", ) - - -def _runnable_score(control: ExecutionControl) -> int: - return control.available_at_ms if control.available_at_ms is not None else control.updated_at_ms diff --git a/src/hayhooks/durable/runtime.py b/src/hayhooks/durable/runtime.py index 266f1c78..86ecb14b 100644 --- a/src/hayhooks/durable/runtime.py +++ b/src/hayhooks/durable/runtime.py @@ -12,13 +12,13 @@ from pydantic import BaseModel, Field, TypeAdapter, ValidationError -from hayhooks.durable.adapters import HaystackDurableAdapter, _run_fenced_thread, execution_kind +from hayhooks.durable.adapters import HaystackDurableAdapter, _run_fenced_thread, chat_messages, execution_kind from hayhooks.durable.backend import ExecutionIdempotencyConflictError from hayhooks.durable.context import DurableContext from hayhooks.durable.manager import DurableExecutionManager from hayhooks.durable.mode import DurableAuthoringMode, _durable_method_implementations, durable_authoring_mode from hayhooks.durable.models import ExecutionKind, ExecutionRecord, JsonValue, json_safe, validate_json -from hayhooks.durable.settings import DurableSettings +from hayhooks.durable.settings import DurableSettings, resolve_durable_settings from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider from hayhooks.server.exceptions import PipelineWrapperError from hayhooks.server.logger import log @@ -39,19 +39,15 @@ def _runtime_settings( durable_settings: DurableSettings | None, app_settings: Any | None, ) -> DurableSettings | None: - if durable_settings is not None and app_settings is not None: - msg = "Pass durable_settings or app_settings, not both" - raise ValueError(msg) - requested = durable_settings or ( - DurableSettings.from_app_settings(app_settings) if app_settings is not None else None - ) + """Prefer a supplied provider's settings, rejecting a conflicting explicit request.""" + requested = resolve_durable_settings(durable_settings, app_settings) provider_settings = getattr(provider, "settings", None) - if isinstance(provider_settings, DurableSettings): - if requested is not None and requested != provider_settings: - msg = "Durable runtime and execution-store provider settings must match" - raise ValueError(msg) - requested = provider_settings - return requested.model_copy(deep=True) if requested is not None else None + if not isinstance(provider_settings, DurableSettings): + return requested + if requested is not None and requested != provider_settings: + msg = "Durable runtime and execution-store provider settings must match" + raise ValueError(msg) + return provider_settings.model_copy(deep=True) class DurableDeployment: @@ -317,18 +313,11 @@ async def _run(self, context: DurableContext) -> JsonValue: ): if self.builtin_agent: request = _DurableAgentRequest.model_validate(context.record.validated_input) - from haystack.dataclasses import ChatMessage - - messages = [ChatMessage.from_dict(message) for message in request.messages] + messages = chat_messages(request.messages) + # A fresh attempt replays the resume turn; a recovered checkpoint already holds it. resume_input = context.take_resume_input() if context.record.checkpoint is None else None if isinstance(resume_input, dict): - resumed_messages = resume_input.get("messages") - if isinstance(resumed_messages, list): - messages.extend( - ChatMessage.from_dict(cast(dict[str, Any], message)) - for message in resumed_messages - if isinstance(message, dict) - ) + messages.extend(chat_messages(resume_input.get("messages"))) return json_safe(await context.run_agent_async(messages=messages)) method = cast(Callable[[DurableContext, BaseModel], Any], self.method) request = self.request_type.model_validate(context.record.validated_input) @@ -388,11 +377,6 @@ def settings(self) -> DurableSettings: self._durable_settings = DurableSettings() return self._durable_settings - @property - def app_settings(self) -> DurableSettings: - """Compatibility alias for :attr:`settings`.""" - return self.settings - def create_deployment(self, name: str, wrapper: BasePipelineWrapper) -> DurableDeployment | None: """Build an uncached candidate so route closures cannot capture an old deployment.""" if not self.has_capability(wrapper): diff --git a/src/hayhooks/durable/settings.py b/src/hayhooks/durable/settings.py index 04efe59e..268ee301 100644 --- a/src/hayhooks/durable/settings.py +++ b/src/hayhooks/durable/settings.py @@ -11,6 +11,8 @@ class DurableSettings(BaseModel): """Durable storage, retention, retry, lease, and worker settings.""" + # Durable executions use Redis by default. Memory is an explicit volatile + # development/test choice and is never selected after a Redis failure. durable_store: Literal["memory", "redis"] = "redis" durable_redis_url: str = "redis://localhost:6379/0" durable_redis_key_prefix: str = "hayhooks:durable" @@ -20,14 +22,21 @@ class DurableSettings(BaseModel): durable_terminal_ttl_seconds: int = Field(default=604_800, ge=1) durable_max_progress_events: int = Field(default=100, ge=1, le=10_000) durable_max_record_bytes: int = Field(default=1_000_000, ge=1_024) + # Zero disables the deployment-wide queued/running/waiting admission cap. durable_max_nonterminal_executions: int = Field(default=0, ge=0) durable_shutdown_grace_period: float = Field(default=5.0, ge=0.0) durable_max_attempts: int = Field(default=3, ge=1, le=1_000) durable_retry_base_delay: float = Field(default=1.0, ge=0.0, le=86_400.0) durable_retry_max_delay: float = Field(default=60.0, ge=0.0, le=604_800.0) + # Shared worker and lease-maintenance polling per deployment/replica at concurrency 1: + # 0.25 s -> ~16 idle Redis commands/s and ~125 ms average pickup latency. + # 0.50 s -> ~8 idle Redis commands/s and ~250 ms average pickup latency. + # 1.00 s -> ~4 idle Redis commands/s and ~500 ms average pickup latency (default). durable_poll_interval: float = Field(default=1.0, ge=0.05, le=60.0) durable_lease_duration_ms: int = Field(default=30_000, ge=1, le=86_400_000) durable_lease_commit_safety_ms: int = Field(default=1_500, ge=0, le=86_400_000) + # Keep the default conservative until Agent tools and shared components are + # proven concurrency-safe. durable_execution_concurrency: int = Field(default=1, ge=1, le=128) @model_validator(mode="after") @@ -44,8 +53,28 @@ def _validate_lease_margin(self) -> Self: @classmethod def from_app_settings(cls, app_settings: Any) -> DurableSettings: - """Copy durable fields from Hayhooks settings without retaining that dependency.""" - return cls(**{name: getattr(app_settings, name) for name in cls.model_fields}) + """ + Snapshot the durable fields of Hayhooks settings, dropping the rest. + ``AppSettings`` inherits this model, so the snapshot is built explicitly + rather than through ``cls`` to keep the runtime free of server settings. + """ + return DurableSettings(**{name: getattr(app_settings, name) for name in DurableSettings.model_fields}) -__all__ = ["DurableSettings"] + +def resolve_durable_settings( + durable_settings: DurableSettings | None = None, + app_settings: Any | None = None, +) -> DurableSettings | None: + """Return an owned copy of whichever settings source the caller supplied.""" + if durable_settings is not None and app_settings is not None: + msg = "Pass durable_settings or app_settings, not both" + raise ValueError(msg) + if durable_settings is not None: + return durable_settings.model_copy(deep=True) + if app_settings is not None: + return DurableSettings.from_app_settings(app_settings) + return None + + +__all__ = ["DurableSettings", "resolve_durable_settings"] diff --git a/src/hayhooks/durable/store.py b/src/hayhooks/durable/store.py index b32d9404..d3f52007 100644 --- a/src/hayhooks/durable/store.py +++ b/src/hayhooks/durable/store.py @@ -47,7 +47,7 @@ ) from hayhooks.durable.redis import RedisExecutionStore, digest from hayhooks.durable.reference import InMemoryExecutionStore -from hayhooks.durable.settings import DurableSettings +from hayhooks.durable.settings import DurableSettings, resolve_durable_settings from hayhooks.server.logger import log _RECORD_PAYLOADS = ( @@ -116,24 +116,16 @@ def __init__( ) -> None: self.store = store self.control = control - self._record = record + self.record = record self.worker_id = worker_id + self.lost_event = asyncio.Event() self._heartbeat: asyncio.Task[None] | None = None self._transition_lock = asyncio.Lock() self._finished = False self._lost = False - self._lost_event = asyncio.Event() self._confirmed_until = confirmed_at + self.store.lease_safe_duration self._last_persisted_progress = record.progress[-1] if record.progress else None - @property - def record(self) -> ExecutionRecord: - return self._record - - @property - def lost_event(self) -> asyncio.Event: - return self._lost_event - async def __aenter__(self) -> ExecutionClaim: async with self._transition_lock: await self._transition(Heartbeat(self.control.fence, self.worker_id, 0, self.store.lease_duration_ms)) @@ -331,7 +323,7 @@ def _ensure_owned(self) -> None: def _mark_lost(self) -> None: if not self._lost: self._lost = True - self._lost_event.set() + self.lost_event.set() class ExecutionStore: @@ -612,8 +604,6 @@ def _record( progress=progress, result=result if control.status is EngineStatus.COMPLETED else None, error=error if control.status is EngineStatus.FAILED else None, - last_retry_error=error if not control.terminal else None, - retry_at=_datetime(control.available_at_ms) if control.available_at_ms is not None else None, cancel_requested_at=( _datetime(control.cancel_requested_at_ms) if control.cancel_requested_at_ms is not None else None ), @@ -648,7 +638,7 @@ def _snapshot(self, record: ExecutionRecord) -> bytes: class RedisExecutionStoreProvider: """Application-owned Redis client and deployment stores.""" - def __init__( # noqa: PLR0913 - mirrors the configurable Redis task-store provider + def __init__( # noqa: PLR0913 - Redis connection and settings sources self, redis_url: str | None = None, *, @@ -657,33 +647,13 @@ def __init__( # noqa: PLR0913 - mirrors the configurable Redis task-store provi close_redis: bool = True, durable_settings: DurableSettings | None = None, app_settings: Any | None = None, - socket_timeout: float | None = None, - socket_connect_timeout: float | None = None, - health_check_interval: int | None = None, ) -> None: - if durable_settings is not None and app_settings is not None: - msg = "Pass durable_settings or app_settings, not both" - raise ValueError(msg) - self.settings = ( - durable_settings - or (DurableSettings.from_app_settings(app_settings) if app_settings is not None else DurableSettings()) - ).model_copy(deep=True) - self.app_settings = self.settings + self.settings = resolve_durable_settings(durable_settings, app_settings) or DurableSettings() self.config = _config(durable_settings=self.settings, key_prefix=key_prefix) self.close_redis = close_redis - self.socket_timeout = ( - socket_timeout if socket_timeout is not None else self.settings.durable_redis_socket_timeout - ) - self.socket_connect_timeout = ( - socket_connect_timeout - if socket_connect_timeout is not None - else self.settings.durable_redis_socket_connect_timeout - ) - self.health_check_interval = ( - health_check_interval - if health_check_interval is not None - else self.settings.durable_redis_health_check_interval - ) + self.socket_timeout = self.settings.durable_redis_socket_timeout + self.socket_connect_timeout = self.settings.durable_redis_socket_connect_timeout + self.health_check_interval = self.settings.durable_redis_health_check_interval if redis is None: try: from redis.asyncio import Redis @@ -723,14 +693,7 @@ def __init__( durable_settings: DurableSettings | None = None, app_settings: Any | None = None, ) -> None: - if durable_settings is not None and app_settings is not None: - msg = "Pass durable_settings or app_settings, not both" - raise ValueError(msg) - self.settings = ( - durable_settings - or (DurableSettings.from_app_settings(app_settings) if app_settings is not None else DurableSettings()) - ).model_copy(deep=True) - self.app_settings = self.settings + self.settings = resolve_durable_settings(durable_settings, app_settings) or DurableSettings() self.config = _config(durable_settings=self.settings) self.cores: dict[str, InMemoryExecutionStore] = {} diff --git a/src/hayhooks/server/a2a/durable_executor.py b/src/hayhooks/server/a2a/durable_executor.py index ded4ed61..01569a49 100644 --- a/src/hayhooks/server/a2a/durable_executor.py +++ b/src/hayhooks/server/a2a/durable_executor.py @@ -11,12 +11,21 @@ from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStatus, ExecutionStoreError from hayhooks.durable.runtime import DefinitionRevisionConflictError, DurableDeployment, execution_id_for from hayhooks.server.a2a.imports import ( + DEFAULT_LIST_TASKS_PAGE_SIZE, AgentExecutor, EventQueue, InvalidParamsError, + ListTasksResponse, RequestContext, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatusUpdateEvent, TaskStore, TaskUpdater, + append_artifact_to_task, + decode_page_token, + encode_page_token, new_task_from_user_message, new_text_part, ) @@ -32,6 +41,28 @@ DURABLE_PROGRESS_ARTIFACT_NAME = "durable-progress" DURABLE_RESULT_ARTIFACT_NAME = "durable-result" +_MISSING_RECORD = "The durable Agent execution record is missing (durable_execution_missing)." +_REJECTED_SUBMISSION = "The durable Agent submission was rejected (durable_submission_rejected)." + + +def _clone(task: Any) -> Any: + copy = type(task)() + copy.CopyFrom(task) + return copy + + +def _projection_updater(projected: Any) -> TaskUpdater: + """Build an updater that mutates a detached task snapshot instead of a queue.""" + return TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), projected.id, projected.context_id) + + +async def _fail(updater: TaskUpdater, text: str) -> None: + await updater.failed(message=updater.new_agent_message([new_text_part(text)])) + + +async def _record(deployment: DurableDeployment, execution_id: str, owner_id: str) -> Any: + """Read one owner-scoped record, tolerating a superseded definition revision.""" + return await deployment.get(execution_id, owner_id=owner_id, enforce_owner=True, allow_revision_mismatch=True) class _TaskProjectionQueue: @@ -41,9 +72,6 @@ def __init__(self, task: Any) -> None: self.task = task async def enqueue_event(self, event: Any) -> None: - from a2a.server.tasks.task_manager import append_artifact_to_task - from a2a.types import TaskArtifactUpdateEvent, TaskStatusUpdateEvent - if isinstance(event, TaskStatusUpdateEvent): if self.task.status.HasField("message"): self.task.history.append(self.task.status.message) @@ -71,19 +99,12 @@ async def get(self, task_id: str, context: Any) -> Any | None: record = None if task is None: try: - record = await self._deployment.get( - execution_id_for(owner_id, task_id), - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) + record = await _record(self._deployment, execution_id_for(owner_id, task_id), owner_id) except KeyError: return None context_id = record.validated_input.get("a2a_context_id") if not isinstance(context_id, str) or not context_id: return None - from a2a.types import Task, TaskState - task = Task(id=task_id, context_id=context_id) task.status.state = ( TaskState.TASK_STATE_WORKING @@ -93,10 +114,6 @@ async def get(self, task_id: str, context: Any) -> Any | None: return await self._project(task, owner_id, context, record=record) async def list(self, params: Any, context: Any) -> Any: - from a2a.types import ListTasksResponse - from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE - from a2a.utils.task import decode_page_token, encode_page_token - tasks = await self._all_tasks(params, context) owner_id = self.owner_id_for_context(context) projected: builtins.list[Any | None] = [] @@ -155,47 +172,28 @@ async def _all_tasks(self, params: Any, context: Any) -> builtins.list[Any]: return tasks request.page_token = page.next_page_token - async def _project( # noqa: C901 - self, task: Any | None, owner_id: str, context: Any, *, record: Any | None = None - ) -> Any | None: + async def _project(self, task: Any | None, owner_id: str, context: Any, *, record: Any | None = None) -> Any | None: if task is None: return None - projected = type(task)() - projected.CopyFrom(task) + projected = _clone(task) copy_version = getattr(self._task_store, "copy_task_version", None) if callable(copy_version): copy_version(task, projected) - settled = False try: - if record is None: - record = await self._deployment.get( - execution_id_for(owner_id, task.id), - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) + record = record or await _record(self._deployment, execution_id_for(owner_id, task.id), owner_id) except KeyError: - if task_is_terminal(task): + if ( + task_is_terminal(task) + or not task.HasField("status") + or task.status.state == TaskState.TASK_STATE_SUBMITTED + ): return projected - if task.HasField("status"): - from a2a.types import TaskState - - if task.status.state == TaskState.TASK_STATE_SUBMITTED: - return projected - updater = TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id) - await updater.failed( - message=updater.new_agent_message( - [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] - ) - ) - settled = True + await _fail(_projection_updater(projected), _MISSING_RECORD) + settled = True else: if _task_matches_record(task, record): return projected - settled = await _project_record( - record, - TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id), - ) + settled = await _project_record(record, _projection_updater(projected)) if settled and task.id in self._read_through_task_ids: try: await self._task_store.save(projected, context) @@ -237,25 +235,10 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non record = None if context.current_task is not None: try: - record = await self.deployment.get( - execution_id, - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) + record = await _record(self.deployment, execution_id, owner_id) except KeyError: - from a2a.types import TaskState - if task.status.state != TaskState.TASK_STATE_SUBMITTED: - await updater.failed( - message=updater.new_agent_message( - [ - new_text_part( - "The durable Agent execution record is missing (durable_execution_missing)." - ) - ] - ) - ) + await _fail(updater, _MISSING_RECORD) return if record is not None and record.status is ExecutionStatus.WAITING: action = "resume" @@ -282,11 +265,7 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non task_id=task.id, error_type=type(error).__name__, ).warning("Rejected durable A2A task submission") - await updater.failed( - message=updater.new_agent_message( - [new_text_part("The durable Agent submission was rejected (durable_submission_rejected).")] - ) - ) + await _fail(updater, _REJECTED_SUBMISSION) return execution_id = record.execution_id else: @@ -336,18 +315,9 @@ async def _wait_for_update( last_sequence = -1 while not self._closed: try: - record = await self.deployment.get( - execution_id, - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) + record = await _record(self.deployment, execution_id, owner_id) except KeyError: - await updater.failed( - message=updater.new_agent_message( - [new_text_part("The durable Agent execution record is missing (durable_execution_missing).")] - ) - ) + await _fail(updater, _MISSING_RECORD) return except ExecutionStoreError: await asyncio.sleep(max(0.1, settings.durable_poll_interval)) @@ -377,54 +347,14 @@ async def _submit(self, task: Any, owner_id: str, messages: list[Any]) -> Any: await asyncio.sleep(max(0.1, settings.durable_poll_interval)) raise asyncio.CancelledError - async def _recover_tasks(self, task_store: RecoverableTaskStore) -> None: # noqa: C901 + async def _recover_tasks(self, task_store: RecoverableTaskStore) -> None: """Repair durable A2A tasks saved before this executor started.""" cursor = 0 recovered = 0 while not self._closed: tasks, next_cursor = await task_store.recoverable_task_batch(cursor, settings.a2a_list_scan_batch_size) for task, owner_id, version in tasks: - execution_id = execution_id_for(owner_id, task.id) - try: - record = await self.deployment.get( - execution_id, - owner_id=owner_id, - enforce_owner=True, - allow_revision_mismatch=True, - ) - except KeyError: - from a2a.types import TaskState - - if task.status.state != TaskState.TASK_STATE_SUBMITTED: - continue - try: - record = await self._submit(task, owner_id, build_haystack_task_messages(task)) - except ValueError as error: - record = None - log.bind( - pipeline_name=self.pipeline_name, - task_id=task.id, - error_type=type(error).__name__, - ).warning("Rejected recovered durable A2A task submission") - if record is not None and record.status in {ExecutionStatus.QUEUED, ExecutionStatus.RUNNING}: - self.task_store._read_through_task_ids.add(task.id) - if record is not None and _task_matches_record(task, record): - continue - projected = type(task)() - projected.CopyFrom(task) - updater = TaskUpdater(cast(EventQueue, _TaskProjectionQueue(projected)), task.id, task.context_id) - if record is None: - await updater.failed( - message=updater.new_agent_message( - [new_text_part("The durable Agent submission was rejected (durable_submission_rejected).")] - ) - ) - else: - await _project_record(record, updater) - if not await task_store.save_projection(projected, owner_id, version): - continue - task.CopyFrom(projected) - recovered += 1 + recovered += await self._recover_task(task_store, task, owner_id, version) if next_cursor is None: if recovered: log.bind(pipeline_name=self.pipeline_name, recovered=recovered).debug( @@ -433,6 +363,38 @@ async def _recover_tasks(self, task_store: RecoverableTaskStore) -> None: # noq return cursor = next_cursor + async def _recover_task(self, task_store: RecoverableTaskStore, task: Any, owner_id: str, version: int) -> bool: + """Reconcile one saved task with its durable record, resubmitting when it never started.""" + try: + record = await _record(self.deployment, execution_id_for(owner_id, task.id), owner_id) + except KeyError: + if task.status.state != TaskState.TASK_STATE_SUBMITTED: + return False + try: + record = await self._submit(task, owner_id, build_haystack_task_messages(task)) + except ValueError as error: + record = None + log.bind( + pipeline_name=self.pipeline_name, + task_id=task.id, + error_type=type(error).__name__, + ).warning("Rejected recovered durable A2A task submission") + if record is not None: + if record.status in {ExecutionStatus.QUEUED, ExecutionStatus.RUNNING}: + self.task_store._read_through_task_ids.add(task.id) + if _task_matches_record(task, record): + return False + projected = _clone(task) + updater = _projection_updater(projected) + if record is None: + await _fail(updater, _REJECTED_SUBMISSION) + else: + await _project_record(record, updater) + if not await task_store.save_projection(projected, owner_id, version): + return False + task.CopyFrom(projected) + return True + def _task(context: RequestContext) -> Any: if context.current_task is not None: @@ -478,8 +440,6 @@ async def _project_record(record: Any, updater: TaskUpdater) -> bool: def _task_matches_record(task: Any, record: Any) -> bool: - from a2a.types import TaskState - states = { ExecutionStatus.WAITING: TaskState.TASK_STATE_INPUT_REQUIRED, ExecutionStatus.COMPLETED: TaskState.TASK_STATE_COMPLETED, diff --git a/src/hayhooks/server/a2a/imports.py b/src/hayhooks/server/a2a/imports.py index d8ed167c..78092838 100644 --- a/src/hayhooks/server/a2a/imports.py +++ b/src/hayhooks/server/a2a/imports.py @@ -14,12 +14,28 @@ from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.routes import create_agent_card_routes, create_jsonrpc_routes from a2a.server.tasks import InMemoryTaskStore, TaskStore, TaskUpdater - from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill, Role + from a2a.server.tasks.task_manager import append_artifact_to_task + from a2a.types import ( + AgentCapabilities, + AgentCard, + AgentInterface, + AgentSkill, + ListTasksRequest, + ListTasksResponse, + Role, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatusUpdateEvent, + ) + from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from a2a.utils.errors import InvalidParamsError + from a2a.utils.task import decode_page_token, encode_page_token a2a_import.check() __all__ = [ + "DEFAULT_LIST_TASKS_PAGE_SIZE", "AgentCapabilities", "AgentCard", "AgentExecutor", @@ -29,14 +45,23 @@ "EventQueue", "InMemoryTaskStore", "InvalidParamsError", + "ListTasksRequest", + "ListTasksResponse", "RequestContext", "RequestContextBuilder", "Role", "SimpleRequestContextBuilder", + "Task", + "TaskArtifactUpdateEvent", + "TaskState", + "TaskStatusUpdateEvent", "TaskStore", "TaskUpdater", + "append_artifact_to_task", "create_agent_card_routes", "create_jsonrpc_routes", + "decode_page_token", + "encode_page_token", "get_message_text", "new_task_from_user_message", "new_text_part", diff --git a/src/hayhooks/server/a2a/messages.py b/src/hayhooks/server/a2a/messages.py index e9e10483..6329daa0 100644 --- a/src/hayhooks/server/a2a/messages.py +++ b/src/hayhooks/server/a2a/messages.py @@ -4,7 +4,9 @@ from typing import Any -from hayhooks.server.a2a.imports import RequestContext, Role, get_message_text +from haystack.dataclasses import ChatMessage + +from hayhooks.server.a2a.imports import RequestContext, Role, TaskState, get_message_text def build_openai_messages(context: RequestContext) -> list[dict[str, str]]: @@ -24,8 +26,6 @@ def build_openai_messages(context: RequestContext) -> list[dict[str, str]]: def _haystack_message(role: str, text: str) -> Any: - from haystack.dataclasses import ChatMessage - return ChatMessage.from_assistant(text) if role == "assistant" else ChatMessage.from_user(text) @@ -56,8 +56,6 @@ def build_haystack_resume_messages(context: RequestContext) -> list[Any]: def task_is_terminal(task: Any) -> bool: - from a2a.types import TaskState - return task.HasField("status") and task.status.state in { TaskState.TASK_STATE_COMPLETED, TaskState.TASK_STATE_CANCELED, diff --git a/src/hayhooks/server/a2a/redis_task_store.py b/src/hayhooks/server/a2a/redis_task_store.py index 369ec026..ecbe873a 100644 --- a/src/hayhooks/server/a2a/redis_task_store.py +++ b/src/hayhooks/server/a2a/redis_task_store.py @@ -16,13 +16,21 @@ ExecutionStoreCorruptionError, ) from hayhooks.durable.redis import digest, redis_time_ms, redis_transaction_backoff, redis_watch_error -from hayhooks.server.a2a.imports import InvalidParamsError, TaskStore +from hayhooks.server.a2a.imports import ( + DEFAULT_LIST_TASKS_PAGE_SIZE, + InvalidParamsError, + ListTasksRequest, + ListTasksResponse, + Task, + TaskStore, + decode_page_token, + encode_page_token, +) from hayhooks.server.a2a.messages import task_is_terminal, task_matches_filters from hayhooks.settings import settings if TYPE_CHECKING: from a2a.server.context import ServerCallContext - from a2a.types import ListTasksRequest, ListTasksResponse, Task OwnerResolver = Callable[[Any], str] @@ -96,8 +104,6 @@ def owner_id_for_context(self, context: ServerCallContext) -> str: @staticmethod def _deserialize(payload: bytes | str) -> Task: - from a2a.types import Task - if isinstance(payload, str): payload = payload.encode("utf-8") task = Task() @@ -307,27 +313,10 @@ async def _delete_payload( async def list(self, params: ListTasksRequest, context: ServerCallContext) -> ListTasksResponse: await self.cleanup_expired_tasks(limit=100) - from a2a.types import ListTasksResponse - from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE - from a2a.utils.task import decode_page_token, encode_page_token - page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE - if params.context_id or params.status or params.HasField("status_timestamp_after"): - page, total_size, next_page_token = await self._list_filtered( - params, - context, - page_size, - decode_page_token, - encode_page_token, - ) - else: - page, total_size, next_page_token = await self._list_by_recent_update( - params, - context, - page_size, - decode_page_token, - encode_page_token, - ) + filtered = bool(params.context_id or params.status or params.HasField("status_timestamp_after")) + lister = self._list_filtered if filtered else self._list_by_recent_update + page, total_size, next_page_token = await lister(params, context, page_size) return ListTasksResponse( tasks=page, next_page_token=next_page_token, @@ -335,25 +324,23 @@ async def list(self, params: ListTasksRequest, context: ServerCallContext) -> Li total_size=total_size, ) + async def _start_rank(self, page_token: str, updates_key: str) -> int: + """Resolve a page token to the next index in the owner update index.""" + if not page_token: + return 0 + rank = await self.redis.zrevrank(updates_key, decode_page_token(page_token)) + if rank is None: + msg = f"Invalid page token: {page_token}" + raise InvalidParamsError(msg) + return int(rank) + 1 + async def _list_by_recent_update( - self, - params: ListTasksRequest, - context: ServerCallContext, - page_size: int, - decode_page_token: Callable[[str], str], - encode_page_token: Callable[[str], str], + self, params: ListTasksRequest, context: ServerCallContext, page_size: int ) -> tuple[builtins.list[Task], int, str | None]: owner = self.owner_id_for_context(context) updates_key = self._updates_key(owner) total_size = await self.redis.zcard(updates_key) - start_index = 0 - if params.page_token: - start_task_id = decode_page_token(params.page_token) - rank = await self.redis.zrevrank(updates_key, start_task_id) - if rank is None: - msg = f"Invalid page token: {params.page_token}" - raise InvalidParamsError(msg) - start_index = rank + 1 + start_index = await self._start_rank(params.page_token, updates_key) task_ids = await self.redis.zrevrange(updates_key, start_index, start_index + page_size - 1) task_ids = [self._decode_value(task_id) for task_id in task_ids] @@ -364,25 +351,13 @@ async def _list_by_recent_update( return page, total_size, next_page_token async def _list_filtered( - self, - params: ListTasksRequest, - context: ServerCallContext, - page_size: int, - decode_page_token: Callable[[str], str], - encode_page_token: Callable[[str], str], + self, params: ListTasksRequest, context: ServerCallContext, page_size: int ) -> tuple[builtins.list[Task], int, str | None]: """Apply filters in bounded update-index batches with exact pagination.""" owner = self.owner_id_for_context(context) updates_key = self._updates_key(owner) indexed_size = int(await self.redis.zcard(updates_key)) - start_rank = 0 - if params.page_token: - start_task_id = decode_page_token(params.page_token) - rank = await self.redis.zrevrank(updates_key, start_task_id) - if rank is None: - msg = f"Invalid page token: {params.page_token}" - raise InvalidParamsError(msg) - start_rank = int(rank) + 1 + start_rank = await self._start_rank(params.page_token, updates_key) page: builtins.list[Task] = [] matching_after_page = False diff --git a/src/hayhooks/server/pipelines/models.py b/src/hayhooks/server/pipelines/models.py index 355215ae..fd66adee 100644 --- a/src/hayhooks/server/pipelines/models.py +++ b/src/hayhooks/server/pipelines/models.py @@ -10,10 +10,6 @@ from hayhooks.server.utils.yaml_utils import InputResolution, OutputResolution -def _create_schema_model(model_name: str, **fields: Any) -> type[BaseModel]: - return create_model(model_name, **fields) - - def _resolved_annotations(func: Callable) -> dict[str, Any]: """Resolve postponed annotations in the wrapper module that declared them.""" try: @@ -44,7 +40,7 @@ def get_request_model_from_resolved_io( default_value = ... if resolution.required else None fields[input_name] = (input_type, default_value) - return _create_schema_model(f"{pipeline_name.capitalize()}RunRequest", **fields) + return create_model(f"{pipeline_name.capitalize()}RunRequest", **fields) def get_response_model_from_resolved_io( @@ -66,7 +62,7 @@ def get_response_model_from_resolved_io( output_type = resolution.type fields[output_name] = (output_type, ...) - return _create_schema_model( + return create_model( f"{pipeline_name.capitalize()}RunResponse", result=(dict, Field(..., description="Pipeline result")) ) @@ -94,7 +90,7 @@ def create_request_model_from_callable(func: Callable, model_name: str, docstrin field_info = Field(default=default_value, description=description) fields[name] = (annotations.get(name, param.annotation), field_info) - return _create_schema_model(f"{model_name}Request", **fields) + return create_model(f"{model_name}Request", **fields) def _is_streaming_type(return_type: type) -> bool: @@ -149,9 +145,7 @@ def create_response_model_from_callable( return_description = docstring.returns.description if docstring.returns else None - return _create_schema_model( - f"{model_name}Response", result=(return_type, Field(..., description=return_description)) - ) + return create_model(f"{model_name}Response", result=(return_type, Field(..., description=return_description))) def get_response_class_from_callable(func: Callable) -> type[Response] | None: diff --git a/src/hayhooks/server/pipelines/registry.py b/src/hayhooks/server/pipelines/registry.py index 066b59fa..c8ee588d 100644 --- a/src/hayhooks/server/pipelines/registry.py +++ b/src/hayhooks/server/pipelines/registry.py @@ -79,13 +79,5 @@ def clear(self) -> None: def get_pipeline_registry(app: Any | None = None) -> PipelineRegistry: """Return the registry owned by *app*, or the process registry for app-less callers.""" - if app is None: - return registry - - app_registry = getattr(app.state, "pipeline_registry", None) - if isinstance(app_registry, PipelineRegistry): - return app_registry - - app_registry = PipelineRegistry() - app.state.pipeline_registry = app_registry - return app_registry + owned = getattr(getattr(app, "state", None), "pipeline_registry", None) + return owned if isinstance(owned, PipelineRegistry) else registry diff --git a/src/hayhooks/server/routers/__init__.py b/src/hayhooks/server/routers/__init__.py index 8aabea99..0899d70d 100644 --- a/src/hayhooks/server/routers/__init__.py +++ b/src/hayhooks/server/routers/__init__.py @@ -1,8 +1,7 @@ from hayhooks.server.routers.dashboard import router as dashboard_router from hayhooks.server.routers.deploy import router as deploy_router from hayhooks.server.routers.draw import router as draw_router -from hayhooks.server.routers.openai import router as openai_router from hayhooks.server.routers.status import router as status_router from hayhooks.server.routers.undeploy import router as undeploy_router -__all__ = ["dashboard_router", "deploy_router", "draw_router", "openai_router", "status_router", "undeploy_router"] +__all__ = ["dashboard_router", "deploy_router", "draw_router", "status_router", "undeploy_router"] diff --git a/src/hayhooks/server/routers/openai.py b/src/hayhooks/server/routers/openai.py index 7bcc4dcd..8e902d3a 100644 --- a/src/hayhooks/server/routers/openai.py +++ b/src/hayhooks/server/routers/openai.py @@ -17,7 +17,7 @@ from haystack.dataclasses import StreamingChunk from hayhooks.server.logger import log -from hayhooks.server.pipelines.registry import PipelineRegistry, registry +from hayhooks.server.pipelines.registry import PipelineRegistry from hayhooks.server.tracing import ( SPAN_OPENAI_FILE_UPLOAD, SPAN_OPENAI_RUN, @@ -61,7 +61,7 @@ class _OpenAIDispatch: ) -def _list_models(pipeline_registry: PipelineRegistry = registry) -> list[str]: +def _list_models(pipeline_registry: PipelineRegistry) -> list[str]: return pipeline_registry.get_names() @@ -83,7 +83,7 @@ async def _collect_async_generator(gen: AsyncGenerator) -> str: return "".join([_chunk_to_text(chunk) async for chunk in gen]) -def _resolve_pipeline_wrapper(model: str, pipeline_registry: PipelineRegistry = registry) -> BasePipelineWrapper: +def _resolve_pipeline_wrapper(model: str, pipeline_registry: PipelineRegistry) -> BasePipelineWrapper: """Look up *model* in the registry, raising 404 if it isn't a pipeline wrapper.""" pipeline_wrapper = pipeline_registry.get(model) if not isinstance(pipeline_wrapper, BasePipelineWrapper): @@ -141,7 +141,7 @@ async def _run_pipeline_method( model: str, kwargs: dict[str, Any], body: dict[str, Any], - pipeline_registry: PipelineRegistry = registry, + pipeline_registry: PipelineRegistry, ) -> str | Generator | AsyncGenerator: """Shared dispatch logic for chat completions and responses endpoints.""" stream_requested = bool(body.get("stream", False)) @@ -196,7 +196,7 @@ async def _run_completion( messages: list[dict[str, Any]], body: dict[str, Any], *, - pipeline_registry: PipelineRegistry = registry, + pipeline_registry: PipelineRegistry, ) -> str | Generator | AsyncGenerator: return await _run_pipeline_method( _CHAT_COMPLETION_DISPATCH, @@ -212,7 +212,7 @@ async def _run_response( input_items: list[dict[str, Any]], body: dict[str, Any], *, - pipeline_registry: PipelineRegistry = registry, + pipeline_registry: PipelineRegistry, ) -> str | Generator | AsyncGenerator: return await _run_pipeline_method( _RESPONSE_DISPATCH, @@ -223,7 +223,7 @@ async def _run_response( ) -def _find_file_upload_wrapper(pipeline_registry: PipelineRegistry = registry) -> BasePipelineWrapper | None: +def _find_file_upload_wrapper(pipeline_registry: PipelineRegistry) -> BasePipelineWrapper | None: """Find the first registered pipeline wrapper that implements ``run_file_upload``.""" for name in pipeline_registry.get_names(): wrapper = pipeline_registry.get(name) @@ -238,7 +238,7 @@ async def _run_file_upload( content: bytes, purpose: str, *, - pipeline_registry: PipelineRegistry = registry, + pipeline_registry: PipelineRegistry, ) -> FileObject: with trace_operation( SPAN_OPENAI_FILE_UPLOAD, @@ -283,7 +283,7 @@ async def _run_file_upload( ) -def create_openai_router(pipeline_registry: PipelineRegistry = registry) -> APIRouter: +def create_openai_router(pipeline_registry: PipelineRegistry) -> APIRouter: """Create OpenAI-compatible routes bound to one pipeline registry.""" list_models = partial(_list_models, pipeline_registry) run_completion = partial(_run_completion, pipeline_registry=pipeline_registry) @@ -312,6 +312,3 @@ def create_openai_router(pipeline_registry: PipelineRegistry = registry) -> APIR ) router.include_router(create_files_router(run_file_upload=run_file_upload, tags=["openai"])) return router - - -router = create_openai_router() diff --git a/src/hayhooks/server/tracing.py b/src/hayhooks/server/tracing.py index 73e1e6e9..1a920a0b 100644 --- a/src/hayhooks/server/tracing.py +++ b/src/hayhooks/server/tracing.py @@ -48,7 +48,6 @@ SPAN_MCP_CALL_TOOL = "hayhooks.mcp.call_tool" SPAN_MCP_RUN_PIPELINE_TOOL = "hayhooks.mcp.run_pipeline_tool" SPAN_A2A_RUN_AGENT = "hayhooks.a2a.run_agent" -SPAN_A2A_DURABLE_PROJECT = "hayhooks.a2a.durable.project" SPAN_DURABLE_SUBMIT = "hayhooks.durable.submit" SPAN_DURABLE_ATTEMPT = "hayhooks.durable.attempt" diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index 48ea0679..4a9118c3 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -126,29 +126,18 @@ def __init__( } pipelines_dir = Path(settings.pipelines_dir) source_dir = pipelines_dir / pipeline_name - sources = [source_dir] if source_dir.is_dir() else [] - sources.extend( + yaml_sources = [ source for extension in (".yml", ".yaml") if (source := pipelines_dir / f"{pipeline_name}{extension}").is_file() - ) - self.backup_dir = Path(tempfile.mkdtemp(prefix="hayhooks-deploy-rollback-")) if sources else None - if source_dir.is_dir() and self.backup_dir is not None: - shutil.copytree(source_dir, self.backup_dir / "pipeline") - for extension in (".yml", ".yaml"): - source = pipelines_dir / f"{pipeline_name}{extension}" - if source.is_file() and self.backup_dir is not None: - shutil.copy2(source, self.backup_dir / f"pipeline{extension}") - - @classmethod - def capture( - cls, - pipeline_name: str, - app: FastAPI | None, - runtime: DurableRuntime | None, - pipeline_registry: PipelineRegistry, - ) -> "_DeploymentSnapshot": - return cls(pipeline_name, app, runtime, pipeline_registry) + ] + self.backup_dir = None + if source_dir.is_dir() or yaml_sources: + self.backup_dir = Path(tempfile.mkdtemp(prefix="hayhooks-deploy-rollback-")) + if source_dir.is_dir(): + shutil.copytree(source_dir, self.backup_dir / "pipeline") + for source in yaml_sources: + shutil.copy2(source, self.backup_dir / f"pipeline{source.suffix}") def restore_publication(self) -> None: self.registry.remove(self.pipeline_name) @@ -223,7 +212,7 @@ async def _deploy_prepared_pipeline_async( # noqa: C901, PLR0912, PLR0913, PLR0 if deployment_key in _deployments_in_progress: msg = f"Pipeline '{pipeline_name}' is already being deployed" raise PipelineAlreadyExistsError(msg) - snapshot = _DeploymentSnapshot.capture(pipeline_name, app, durable_runtime, pipeline_registry) + snapshot = _DeploymentSnapshot(pipeline_name, app, durable_runtime, pipeline_registry) if snapshot.wrapper is not None and not overwrite: msg = f"Pipeline '{pipeline_name}' already exists" raise PipelineAlreadyExistsError(msg) diff --git a/src/hayhooks/settings.py b/src/hayhooks/settings.py index bb57f477..5b6c1173 100644 --- a/src/hayhooks/settings.py +++ b/src/hayhooks/settings.py @@ -3,10 +3,10 @@ from typing import Literal from dotenv import find_dotenv, load_dotenv -from pydantic import Field, model_validator +from pydantic import Field from pydantic_settings import BaseSettings, SettingsConfigDict -from typing_extensions import Self +from hayhooks.durable.settings import DurableSettings from hayhooks.server.logger import log load_dotenv(dotenv_path=find_dotenv(usecwd=True)) @@ -34,7 +34,7 @@ class StartupDeployStrategy(str, Enum): PARALLEL = "parallel" -class AppSettings(BaseSettings): +class AppSettings(DurableSettings, BaseSettings): # Root path for the FastAPI app root_path: str = "" @@ -96,53 +96,10 @@ class AppSettings(BaseSettings): a2a_task_snapshot_cache_size: int = Field(default=1_024, ge=1, le=1_000_000) a2a_list_scan_batch_size: int = Field(default=500, ge=1, le=10_000) - # Durable executions use Redis by default. Memory is an explicit volatile - # development/test choice and is never selected after a Redis failure. - durable_store: Literal["memory", "redis"] = "redis" - durable_redis_url: str = "redis://localhost:6379/0" - durable_redis_key_prefix: str = "hayhooks:durable" - durable_redis_socket_timeout: float = Field(default=5.0, gt=0.0, le=300.0) - durable_redis_socket_connect_timeout: float = Field(default=5.0, gt=0.0, le=300.0) - durable_redis_health_check_interval: int = Field(default=30, ge=0, le=3_600) - - # Retention and safety limits are application settings, not wrapper API. - durable_terminal_ttl_seconds: int = Field(default=604_800, ge=1) - durable_max_progress_events: int = Field(default=100, ge=1, le=10_000) - durable_max_record_bytes: int = Field(default=1_000_000, ge=1_024) - # Zero disables the deployment-wide queued/running/waiting admission cap. - durable_max_nonterminal_executions: int = Field(default=0, ge=0) - durable_shutdown_grace_period: float = Field(default=5.0, ge=0.0) - durable_max_attempts: int = Field(default=3, ge=1, le=1_000) - durable_retry_base_delay: float = Field(default=1.0, ge=0.0, le=86_400.0) - durable_retry_max_delay: float = Field(default=60.0, ge=0.0, le=604_800.0) - # Shared worker and lease-maintenance polling per deployment/replica at concurrency 1: - # 0.25 s -> ~16 idle Redis commands/s and ~125 ms average pickup latency. - # 0.50 s -> ~8 idle Redis commands/s and ~250 ms average pickup latency. - # 1.00 s -> ~4 idle Redis commands/s and ~500 ms average pickup latency (default). - durable_poll_interval: float = Field(default=1.0, ge=0.05, le=60.0) # When configured, a trusted reverse proxy must strip any client-supplied # value and inject the authenticated owner. Empty means bearer-ID mode. durable_trusted_owner_header: str = "" - durable_lease_duration_ms: int = Field(default=30_000, ge=1, le=86_400_000) - durable_lease_commit_safety_ms: int = Field(default=1_500, ge=0, le=86_400_000) - - # Keep the default conservative until Agent tools and shared components are - # proven concurrency-safe. - durable_execution_concurrency: int = Field(default=1, ge=1, le=128) - - @model_validator(mode="after") - def _validate_durable_lease_margin(self) -> Self: - if self.durable_lease_commit_safety_ms >= self.durable_lease_duration_ms: - msg = "durable_lease_commit_safety_ms must be smaller than durable_lease_duration_ms" - raise ValueError(msg) - if self.durable_lease_duration_ms - self.durable_lease_commit_safety_ms <= max( - 10, self.durable_lease_duration_ms / 3 - ): - msg = "durable lease duration minus commit safety must exceed the heartbeat interval" - raise ValueError(msg) - return self - # Disable SSL verification when making requests from the CLI disable_ssl: bool = False diff --git a/tests/durable_contract.py b/tests/durable_contract.py index 3b5c7d48..4e407da0 100644 --- a/tests/durable_contract.py +++ b/tests/durable_contract.py @@ -6,6 +6,7 @@ import pytest +from hayhooks.durable.backend import ExecutionStoreConfig from hayhooks.durable.engine import ( Checkpoint, Claim, @@ -19,6 +20,20 @@ from hayhooks.durable.redis import ExecutionIdempotencyConflictError +def contract_config(**changes: Any) -> ExecutionStoreConfig: + """Small-limit configuration both backends use for the shared contract.""" + limits: dict[str, Any] = { + "max_input_bytes": 64, + "max_checkpoint_bytes": 64, + "max_result_bytes": 64, + "max_error_bytes": 64, + "max_wait_bytes": 64, + "max_progress_events": 2, + "max_progress_event_bytes": 32, + } + return ExecutionStoreConfig(**{**limits, **changes}) + + def control(run_id: str = "run_1", *, idempotency_digest: str = "a" * 64, binding_digest: str = "b" * 64): return initial_control( run_id=run_id, diff --git a/tests/durable_helpers.py b/tests/durable_helpers.py new file mode 100644 index 00000000..f442a2d2 --- /dev/null +++ b/tests/durable_helpers.py @@ -0,0 +1,64 @@ +"""Polling helpers shared by the durable test modules.""" + +from __future__ import annotations + +import asyncio +import inspect +import time +from collections.abc import Callable +from typing import Any + +import pytest + +_ATTEMPTS = 200 +_DELAY = 0.01 + + +def wait_until(predicate: Callable[[], Any], message: str, *, attempts: int = _ATTEMPTS, delay: float = _DELAY) -> Any: + """Poll *predicate* until it returns something truthy, then return it.""" + for _ in range(attempts): + if value := predicate(): + return value + time.sleep(delay) + pytest.fail(message) + + +async def wait_until_async( + predicate: Callable[[], Any], message: str, *, attempts: int = _ATTEMPTS, delay: float = _DELAY +) -> Any: + """Async counterpart to :func:`wait_until`; the predicate may be awaitable.""" + for _ in range(attempts): + value = predicate() + if inspect.isawaitable(value): + value = await value + if value: + return value + await asyncio.sleep(delay) + pytest.fail(message) + + +def wait_for_status(client: Any, url: str, status: str, **kwargs: Any) -> dict: + """Poll a durable execution resource until it reports *status*.""" + + def ready() -> dict | None: + body = client.get(url).json() + return body if body.get("status") == status else None + + return wait_until(ready, f"execution at {url} did not become {status}", **kwargs) + + +async def wait_for_record( + source: Any, + execution_id: str, + predicate: Callable[[Any], bool] = lambda record: record.terminal, + *, + message: str = "durable execution did not reach its expected state", + **kwargs: Any, +) -> Any: + """Poll a store or deployment until its record satisfies *predicate*.""" + + async def ready() -> Any: + record = await source.get(execution_id) + return record if record is not None and predicate(record) else None + + return await wait_until_async(ready, message, **kwargs) diff --git a/tests/test_a2a.py b/tests/test_a2a.py index 187e3ffe..ba6e3eec 100644 --- a/tests/test_a2a.py +++ b/tests/test_a2a.py @@ -3,12 +3,17 @@ from types import SimpleNamespace import pytest +from a2a.server.tasks import InMemoryTaskStore from haystack.dataclasses import StreamingChunk +from starlette.testclient import TestClient +from hayhooks.a2a import TaskStoreProvider from hayhooks.events import PipelineEvent +from hayhooks.server.a2a.app import create_a2a_app from hayhooks.server.a2a.cards import create_agent_card, get_a2a_base_url, is_a2a_exposable from hayhooks.server.a2a.executor import RESPONSE_ARTIFACT_NAME, _stream_item_to_text, create_agent_executor from hayhooks.server.a2a.messages import build_openai_messages +from hayhooks.server.a2a.runtime import A2ARuntime from hayhooks.server.logger import log from hayhooks.server.pipelines import registry from hayhooks.server.tracing import SPAN_A2A_RUN_AGENT @@ -30,6 +35,26 @@ def cleanup_test_pipelines(): registry.clear() +class RecordingTaskStoreProvider(TaskStoreProvider): + """Task-store provider each runtime test bends to the behavior it needs.""" + + def __init__(self, *, store=InMemoryTaskStore, health=None): + self.agent_names = [] + self.closed = False + self._store = store + self._health = health + + def create_task_store(self, agent_name): + self.agent_names.append(agent_name) + return self._store() + + async def health(self): + return self._health if self._health is not None else await super().health() + + async def close(self): + self.closed = True + + class RecordingQueue: """Minimal EventQueue stand-in recording enqueued events.""" @@ -367,19 +392,6 @@ async def test_execute_agent_task_emits_trace_and_safe_lifecycle_logs(recording_ def test_runtime_passes_agent_name_to_task_store_provider(): - from a2a.server.tasks import InMemoryTaskStore - - from hayhooks.a2a import TaskStoreProvider - from hayhooks.server.a2a.runtime import A2ARuntime - - class RecordingTaskStoreProvider(TaskStoreProvider): - def __init__(self): - self.agent_names = [] - - def create_task_store(self, agent_name): - self.agent_names.append(agent_name) - return InMemoryTaskStore() - provider = RecordingTaskStoreProvider() runtime = A2ARuntime(task_store_provider=provider) @@ -387,42 +399,19 @@ def create_task_store(self, agent_name): second_store = runtime.create_task_store("second_agent") assert isinstance(first_store, InMemoryTaskStore) - assert isinstance(second_store, InMemoryTaskStore) assert first_store is not second_store assert provider.agent_names == ["first_agent", "second_agent"] def test_runtime_rejects_invalid_task_store_from_provider(): - from hayhooks.a2a import TaskStoreProvider - from hayhooks.server.a2a.runtime import A2ARuntime - - class InvalidTaskStoreProvider(TaskStoreProvider): - def create_task_store(self, _agent_name): - return object() + runtime = A2ARuntime(task_store_provider=RecordingTaskStoreProvider(store=object)) - runtime = A2ARuntime(task_store_provider=InvalidTaskStoreProvider()) - - with pytest.raises(TypeError, match=r"InvalidTaskStoreProvider.*invalid_agent"): + with pytest.raises(TypeError, match=r"RecordingTaskStoreProvider.*invalid_agent"): runtime.create_task_store("invalid_agent") async def test_runtime_closes_task_store_provider(): - from a2a.server.tasks import InMemoryTaskStore - - from hayhooks.a2a import TaskStoreProvider - from hayhooks.server.a2a.runtime import A2ARuntime - - class CloseableTaskStoreProvider(TaskStoreProvider): - def __init__(self): - self.closed = False - - def create_task_store(self, _agent_name): - return InMemoryTaskStore() - - async def close(self): - self.closed = True - - provider = CloseableTaskStoreProvider() + provider = RecordingTaskStoreProvider() await A2ARuntime(task_store_provider=provider).close() @@ -437,24 +426,10 @@ async def close(self): ], ) def test_a2a_status_returns_503_for_unhealthy_task_store(health, expected_error): - from a2a.server.tasks import InMemoryTaskStore - from starlette.testclient import TestClient - - from hayhooks.a2a import TaskStoreProvider - from hayhooks.server.a2a.app import create_a2a_app - from hayhooks.server.a2a.runtime import A2ARuntime - - class TestProvider(TaskStoreProvider): - def create_task_store(self, _agent_name): - return InMemoryTaskStore() - - async def health(self): - return health - register_wrapper("chat_agent", AsyncChatWrapper) app = create_a2a_app( base_url="http://test:1418", - runtime=A2ARuntime(task_store_provider=TestProvider()), + runtime=A2ARuntime(task_store_provider=RecordingTaskStoreProvider(health=health)), ) with TestClient(app) as client: diff --git a/tests/test_cli.py b/tests/test_cli.py index 9fb325ed..eb383672 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1,3 +1,5 @@ +from types import SimpleNamespace + import pytest from typer.testing import CliRunner @@ -140,119 +142,85 @@ def fake_prepare_tracing_dashboard_assets(dashboard_dist_dir: str) -> str: assert settings.dashboard_dist_dir == str(built_dist_dir) -def test_a2a_run_debug_enables_tracebacks(monkeypatch): +@pytest.fixture +def a2a_run(monkeypatch): + """Neutralize the A2A server plumbing and return a `hayhooks a2a run` invoker.""" import uvicorn from hayhooks.server.a2a import app as a2a_app from hayhooks.server.utils import deploy_utils - from hayhooks.settings import settings - - calls = [] - runtimes = [] - def fake_uvicorn_run(*args, **kwargs): - calls.append((args, kwargs)) + recorded = SimpleNamespace(runtimes=[], uvicorn_calls=[]) def fake_create_a2a_app(*, debug: bool = False, durable_runtime=None): - runtimes.append(durable_runtime) + recorded.runtimes.append(durable_runtime) return object() - def fake_deploy_pipelines(*, durable_runtime=None) -> None: - runtimes.append(durable_runtime) - - monkeypatch.setattr(uvicorn, "run", fake_uvicorn_run) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", fake_deploy_pipelines) - monkeypatch.setattr(a2a_app, "create_a2a_app", fake_create_a2a_app) - settings.show_tracebacks = False - result = runner.invoke(hayhooks_cli, ["a2a", "run", "--debug", "--pipelines-dir", "dummy_pipelines"]) - - assert result.exit_code == 0, result.output - assert settings.show_tracebacks is True - assert calls, "uvicorn.run was not called" - assert runtimes[0] is runtimes[1] - - -def test_a2a_run_sets_builtin_redis_task_store(monkeypatch): - import uvicorn - - from hayhooks.server.a2a import app as a2a_app - from hayhooks.server.utils import deploy_utils - - monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) - monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) - - result = runner.invoke( - hayhooks_cli, - [ - "a2a", - "run", - "--task-store", - "redis", - "--a2a-redis-url", - "redis://localhost:6379/4", - "--a2a-redis-key-prefix", - "demo:a2a", - "--pipelines-dir", - "dummy_pipelines", - ], - ) - - assert result.exit_code == 0, result.output - assert settings.a2a_task_store == "redis" - assert settings.a2a_redis_url == "redis://localhost:6379/4" - assert settings.a2a_redis_key_prefix == "demo:a2a" - - -def test_a2a_run_sets_durable_execution_concurrency(monkeypatch): - import uvicorn - - from hayhooks.server.a2a import app as a2a_app - from hayhooks.server.utils import deploy_utils - - monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) - monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) - - result = runner.invoke( - hayhooks_cli, - ["a2a", "run", "--durable-execution-concurrency", "3", "--pipelines-dir", "dummy_pipelines"], + monkeypatch.setattr(uvicorn, "run", lambda *args, **kwargs: recorded.uvicorn_calls.append((args, kwargs))) + monkeypatch.setattr( + deploy_utils, "deploy_pipelines", lambda *, durable_runtime=None: recorded.runtimes.append(durable_runtime) ) + monkeypatch.setattr(a2a_app, "create_a2a_app", fake_create_a2a_app) - assert result.exit_code == 0, result.output - assert settings.durable_execution_concurrency == 3 - + def run(*args): + result = runner.invoke(hayhooks_cli, ["a2a", "run", *args, "--pipelines-dir", "dummy_pipelines"]) + assert result.exit_code == 0, result.output + return recorded -def test_a2a_run_sets_durable_execution_store_configuration(monkeypatch): - import uvicorn + return run - from hayhooks.server.a2a import app as a2a_app - from hayhooks.server.utils import deploy_utils - monkeypatch.setattr(uvicorn, "run", lambda *_args, **_kwargs: None) - monkeypatch.setattr(deploy_utils, "deploy_pipelines", lambda **_kwargs: None) - monkeypatch.setattr(a2a_app, "create_a2a_app", lambda **_kwargs: object()) +def test_a2a_run_debug_enables_tracebacks(a2a_run): + settings.show_tracebacks = False - result = runner.invoke( - hayhooks_cli, - [ - "a2a", - "run", - "--execution-store", - "redis", - "--execution-redis-url", - "redis://localhost:6379/5", - "--execution-redis-key-prefix", - "demo:durable", - "--pipelines-dir", - "dummy_pipelines", - ], - ) + recorded = a2a_run("--debug") - assert result.exit_code == 0, result.output - assert settings.durable_store == "redis" - assert settings.durable_redis_url == "redis://localhost:6379/5" - assert settings.durable_redis_key_prefix == "demo:durable" + assert settings.show_tracebacks is True + assert recorded.uvicorn_calls, "uvicorn.run was not called" + # Deployment and the app must share one runtime, or workers would never see the pipelines. + assert recorded.runtimes[0] is recorded.runtimes[1] + + +@pytest.mark.parametrize( + ("args", "expected"), + [ + pytest.param( + ["--task-store", "redis", "--a2a-redis-url", "redis://host:6379/4", "--a2a-redis-key-prefix", "demo:a2a"], + { + "a2a_task_store": "redis", + "a2a_redis_url": "redis://host:6379/4", + "a2a_redis_key_prefix": "demo:a2a", + }, + id="a2a-task-store", + ), + pytest.param( + ["--durable-execution-concurrency", "3"], + {"durable_execution_concurrency": 3}, + id="durable-concurrency", + ), + pytest.param( + [ + "--execution-store", + "redis", + "--execution-redis-url", + "redis://host:6379/5", + "--execution-redis-key-prefix", + "demo:durable", + ], + { + "durable_store": "redis", + "durable_redis_url": "redis://host:6379/5", + "durable_redis_key_prefix": "demo:durable", + }, + id="durable-execution-store", + ), + ], +) +def test_a2a_run_applies_store_options(a2a_run, args, expected): + a2a_run(*args) + + for field, value in expected.items(): + assert getattr(settings, field) == value def test_status_command(monkeypatch): diff --git a/tests/test_durable_a2a.py b/tests/test_durable_a2a.py index 59788692..caace088 100644 --- a/tests/test_durable_a2a.py +++ b/tests/test_durable_a2a.py @@ -1,13 +1,14 @@ import asyncio import importlib.metadata import threading -import time from concurrent.futures import ThreadPoolExecutor from concurrent.futures import TimeoutError as FutureTimeoutError +from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import AsyncMock import pytest +from a2a.types import ListTasksResponse, Message, Role, Task, TaskState from fastapi.testclient import TestClient from hayhooks.a2a import A2APipelineWrapper, TaskStoreProvider @@ -20,6 +21,7 @@ from hayhooks.server.a2a.runtime import A2ARuntime from hayhooks.server.pipelines.registry import registry from hayhooks.settings import settings +from tests.durable_helpers import wait_until pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -102,8 +104,6 @@ async def get(self, task_id, _context): return self.tasks.get(task_id) async def list(self, _params, _context): - from a2a.types import ListTasksResponse - return ListTasksResponse(tasks=list(self.tasks.values()), page_size=len(self.tasks), total_size=len(self.tasks)) async def delete(self, task_id, _context): @@ -142,8 +142,6 @@ def _response_task(response): def _recoverable_task(): - from a2a.types import Message, Role - return new_task_from_user_message( Message( message_id="message", @@ -187,21 +185,29 @@ def http_store() -> _HTTPStore: return _HTTPStore() -def test_a2a_http_reads_completion_from_durable_execution(monkeypatch, http_store) -> None: - app = _http_app(http_store, _Deployment(), monkeypatch) +@pytest.fixture +def a2a_client(monkeypatch, http_store): + """Mount one durable deployment and yield a version-pinned A2A client.""" + + @contextmanager + def _client(deployment): + with TestClient(_http_app(http_store, deployment, monkeypatch), headers={"A2A-Version": "1.0"}) as client: + yield client + + return _client + - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: +def test_a2a_http_reads_completion_from_durable_execution(a2a_client) -> None: + with a2a_client(_Deployment()) as client: completed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) assert completed["status"]["state"] == "TASK_STATE_COMPLETED" assert completed["artifacts"][-1]["name"] == "durable-result" -def test_expired_task_projection_uses_the_retained_execution(monkeypatch, http_store) -> None: +def test_expired_task_projection_uses_the_retained_execution(a2a_client, http_store) -> None: deployment = _Deployment() - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: completed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) http_store.tasks.clear() deployment.submit = AsyncMock(side_effect=AssertionError("retained execution must not be resubmitted")) @@ -214,11 +220,9 @@ def test_expired_task_projection_uses_the_retained_execution(monkeypatch, http_s assert not http_store.tasks -def test_a2a_http_waiting_task_resumes_with_only_the_follow_up(monkeypatch, http_store) -> None: +def test_a2a_http_waiting_task_resumes_with_only_the_follow_up(a2a_client) -> None: deployment = _Deployment(status=ExecutionStatus.WAITING) - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: waiting = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) assert waiting["status"]["state"] == "TASK_STATE_INPUT_REQUIRED" completed = _response_task( @@ -231,11 +235,9 @@ def test_a2a_http_waiting_task_resumes_with_only_the_follow_up(monkeypatch, http } -def test_a2a_http_cancel_reaches_durable_execution(monkeypatch, http_store) -> None: +def test_a2a_http_cancel_reaches_durable_execution(a2a_client) -> None: deployment = _Deployment(status=ExecutionStatus.RUNNING) - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) canceled = _response_task(client.post("/durable-agent/", json=_cancel_payload(active["id"]))) @@ -243,23 +245,19 @@ def test_a2a_http_cancel_reaches_durable_execution(monkeypatch, http_store) -> N assert canceled["status"]["state"] == "TASK_STATE_CANCELED" -def test_a2a_http_cancel_projects_a_terminal_race(monkeypatch, http_store) -> None: +def test_a2a_http_cancel_projects_a_terminal_race(a2a_client) -> None: deployment = _Deployment(status=ExecutionStatus.RUNNING) deployment.cancel_accepted = False - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) canceled = _response_task(client.post("/durable-agent/", json=_cancel_payload(active["id"]))) assert canceled["status"]["state"] == "TASK_STATE_CANCELED" -def test_return_immediately_waits_for_durable_submission(monkeypatch, http_store) -> None: +def test_return_immediately_waits_for_durable_submission(a2a_client, http_store) -> None: deployment = _BlockingDeployment() - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client, ThreadPoolExecutor() as pool: + with a2a_client(deployment) as client, ThreadPoolExecutor() as pool: response = pool.submit(client.post, "/durable-agent/", json=_send_payload("initial", return_immediately=True)) try: assert deployment.submit_started.wait(timeout=1) @@ -271,20 +269,15 @@ def test_return_immediately_waits_for_durable_submission(monkeypatch, http_store assert response.result(timeout=2).status_code == 200 -def test_returned_task_is_eventually_persisted_as_terminal(monkeypatch, http_store) -> None: - from a2a.types import TaskState - +def test_returned_task_is_eventually_persisted_as_terminal(a2a_client, http_store) -> None: deployment = _Deployment(status=ExecutionStatus.COMPLETED) - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: active = _response_task(client.post("/durable-agent/", json=_send_payload("initial", return_immediately=True))) - deadline = time.monotonic() + 1 - while ( - http_store.tasks[active["id"]].status.state != TaskState.TASK_STATE_COMPLETED - and time.monotonic() < deadline - ): - time.sleep(0.01) + wait_until( + lambda: http_store.tasks[active["id"]].status.state == TaskState.TASK_STATE_COMPLETED, + "returned task was never persisted as completed", + attempts=100, + ) completed = _response_task(client.post("/durable-agent/", json=_get_payload(active["id"]))) assert completed["status"]["state"] == "TASK_STATE_COMPLETED" @@ -292,8 +285,6 @@ def test_returned_task_is_eventually_persisted_as_terminal(monkeypatch, http_sto async def test_expired_execution_preserves_retained_terminal_task(http_store) -> None: - from a2a.types import Task, TaskState - task = Task(id="retained-terminal", context_id="context") task.status.state = TaskState.TASK_STATE_COMPLETED deployment = _Deployment() @@ -305,8 +296,6 @@ async def test_expired_execution_preserves_retained_terminal_task(http_store) -> async def test_missing_execution_preserves_a_task_awaiting_submission(http_store) -> None: - from a2a.types import TaskState - task = _recoverable_task() deployment = _Deployment() deployment.get = AsyncMock(side_effect=KeyError("submission has not committed yet")) @@ -415,14 +404,10 @@ async def test_durable_a2a_submission_retries_transient_failures(monkeypatch, ht sleep.assert_awaited_once() -def test_a2a_http_rejected_submission_is_persisted_as_failed(monkeypatch, http_store) -> None: - from a2a.types import TaskState - +def test_a2a_http_rejected_submission_is_persisted_as_failed(a2a_client, http_store) -> None: deployment = _Deployment() deployment.submit = AsyncMock(side_effect=ValueError("request is too large")) - app = _http_app(http_store, deployment, monkeypatch) - - with TestClient(app, headers={"A2A-Version": "1.0"}) as client: + with a2a_client(deployment) as client: failed = _response_task(client.post("/durable-agent/", json=_send_payload("initial"))) assert failed["status"]["state"] == "TASK_STATE_FAILED" @@ -430,8 +415,6 @@ def test_a2a_http_rejected_submission_is_persisted_as_failed(monkeypatch, http_s async def test_recovery_submits_a_persisted_task_without_an_execution(http_store) -> None: - from a2a.types import TaskState - task = _recoverable_task() recovery_store = _recovery_store(task) deployment = _Deployment() @@ -450,8 +433,6 @@ async def missing_until_submitted(execution_id, **kwargs): async def test_recovery_rejects_one_invalid_persisted_task_without_aborting(http_store) -> None: - from a2a.types import TaskState - task = _recoverable_task() recovery_store = _recovery_store(task) deployment = _Deployment() @@ -465,8 +446,6 @@ async def test_recovery_rejects_one_invalid_persisted_task_without_aborting(http async def test_recovery_skips_a_projection_conflict(http_store) -> None: - from a2a.types import Task, TaskState - task = Task(id="task", context_id="context") task.status.state = TaskState.TASK_STATE_WORKING recovery_store = _recovery_store(task, saved=False) diff --git a/tests/test_durable_deployment_lifecycle.py b/tests/test_durable_deployment_lifecycle.py index acd30d96..4db76b66 100644 --- a/tests/test_durable_deployment_lifecycle.py +++ b/tests/test_durable_deployment_lifecycle.py @@ -1,5 +1,4 @@ import importlib.metadata -import time import pytest from fastapi.testclient import TestClient @@ -9,6 +8,7 @@ from hayhooks.server.pipelines.registry import registry from hayhooks.server.utils.deploy_utils import deploy_pipeline_files from hayhooks.settings import settings +from tests.durable_helpers import wait_for_status pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -111,14 +111,12 @@ def _deploy(client: TestClient, source: str, *, overwrite: bool = False): ) +def _deploy_ok(client: TestClient, source: str, *, overwrite: bool = False) -> None: + assert _deploy(client, source, overwrite=overwrite).status_code == 200 + + def _wait_for_completion(client: TestClient, response) -> dict: - body = response.json() - for _ in range(200): - result = client.get(body["links"]["self"]) - if result.json()["status"] == "completed": - return result.json() - time.sleep(0.01) - pytest.fail("durable execution did not complete") + return wait_for_status(client, response.json()["links"]["self"], "completed") @pytest.fixture(autouse=True) @@ -134,7 +132,7 @@ def _isolated_runtime(monkeypatch, tmp_path): def test_undeploy_removes_entire_durable_route_family() -> None: app = create_app() with TestClient(app) as client: - assert _deploy(client, _durable_source(field="value", increment=1, revision="first")).status_code == 200 + _deploy_ok(client, _durable_source(field="value", increment=1, revision="first")) assert client.post("/undeploy/job").status_code == 200 assert client.post("/job/run-durable", json={"value": 2}).status_code == 404 assert client.get("/job/executions/missing").status_code == 404 @@ -146,14 +144,9 @@ def test_undeploy_removes_entire_durable_route_family() -> None: def test_undeploy_refuses_to_strand_waiting_execution() -> None: app = create_app() with TestClient(app) as client: - assert _deploy(client, _waiting_source(revision="first")).status_code == 200 + _deploy_ok(client, _waiting_source(revision="first")) submitted = client.post("/job/run-durable", json={"value": 1}) - for _ in range(100): - if client.get(submitted.json()["links"]["self"]).json()["status"] == "waiting": - break - time.sleep(0.01) - else: - pytest.fail("execution did not enter waiting before undeploy") + wait_for_status(client, submitted.json()["links"]["self"], "waiting") blocked = client.post("/undeploy/job") assert blocked.status_code == 409 @@ -167,33 +160,14 @@ def test_undeploy_refuses_to_strand_waiting_execution() -> None: def test_durable_overwrite_routes_bind_new_model_runner_and_revision() -> None: app = create_app() with TestClient(app) as client: - assert ( - _deploy( - client, - _durable_source( - field="old_value", - increment=1, - revision="first", - result_field="old_result", - ), - ).status_code - == 200 - ) + _deploy_ok(client, _durable_source(field="old_value", increment=1, revision="first", result_field="old_result")) first = client.post("/job/run-durable", json={"old_value": 2}) assert _wait_for_completion(client, first)["result"] == {"old_result": 3} - assert ( - _deploy( - client, - _durable_source( - field="new_value", - increment=20, - revision="second", - result_field="new_result", - ), - overwrite=True, - ).status_code - == 200 + _deploy_ok( + client, + _durable_source(field="new_value", increment=20, revision="second", result_field="new_result"), + overwrite=True, ) assert client.get(first.json()["links"]["self"]).json()["result"] == {"old_result": 3} assert client.post("/job/run-durable", json={"old_value": 2}).status_code == 422 @@ -207,7 +181,7 @@ def test_failed_durable_preflight_restarts_existing_deployment(monkeypatch, oper app = create_app() with TestClient(app) as client: source = _durable_source(field="value", increment=1, revision="first") - assert _deploy(client, source).status_code == 200 + _deploy_ok(client, source) deployment = app.state.durable_runtime.current_deployment("job") assert deployment is not None @@ -228,16 +202,10 @@ async def fail_counts(): def test_overwrite_refuses_to_prepare_while_old_execution_is_waiting(tmp_path) -> None: app = create_app() with TestClient(app) as client: - assert _deploy(client, _waiting_source(revision="first")).status_code == 200 + _deploy_ok(client, _waiting_source(revision="first")) submitted = client.post("/job/run-durable", json={"value": 2}) url = submitted.json()["links"]["self"] - for _ in range(100): - waiting = client.get(url) - if waiting.json()["status"] == "waiting": - break - time.sleep(0.01) - else: - pytest.fail("old-revision execution did not enter waiting") + wait_for_status(client, url, "waiting") preparation_marker = tmp_path / "replacement-prepared" replacement_source = _durable_source(field="value", increment=20, revision="second").replace( @@ -254,21 +222,14 @@ def test_overwrite_refuses_to_prepare_while_old_execution_is_waiting(tmp_path) - assert not preparation_marker.exists() assert client.get(url).json()["status"] == "waiting" assert client.post(f"{url}/cancel").status_code == 202 - assert ( - _deploy( - client, - _durable_source(field="value", increment=20, revision="second"), - overwrite=True, - ).status_code - == 200 - ) + _deploy_ok(client, _durable_source(field="value", increment=20, revision="second"), overwrite=True) def test_undeploy_refuses_to_strand_thread_backed_work(monkeypatch) -> None: monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) app = create_app() with TestClient(app) as client: - assert _deploy(client, _blocking_source()).status_code == 200 + _deploy_ok(client, _blocking_source()) old_wrapper = app.state.pipeline_registry.get("job") assert old_wrapper is not None submitted = client.post("/job/run-durable", json={"value": 2}) @@ -276,12 +237,7 @@ def test_undeploy_refuses_to_strand_thread_backed_work(monkeypatch) -> None: assert client.post("/undeploy/job").status_code == 409 old_wrapper.release.set() - for _ in range(100): - if client.get(submitted.json()["links"]["self"]).json()["status"] == "completed": - break - time.sleep(0.01) - else: - pytest.fail("durable work did not complete after rejected undeploy") + _wait_for_completion(client, submitted) assert client.post("/undeploy/job").status_code == 200 @@ -290,14 +246,14 @@ def test_fastapi_apps_own_pipeline_publication_and_allow_the_same_name() -> None app_b = create_app() with TestClient(app_a) as client_a, TestClient(app_b) as client_b: - assert _deploy(client_a, _durable_source(field="value", increment=1, revision="app-a")).status_code == 200 + _deploy_ok(client_a, _durable_source(field="value", increment=1, revision="app-a")) assert client_a.get("/status").json()["pipelines"] == ["job"] assert client_b.get("/status").json()["pipelines"] == [] assert client_b.get("/status/job").status_code == 404 assert client_b.post("/job/run-durable", json={"value": 1}).status_code == 404 - assert _deploy(client_b, _durable_source(field="value", increment=10, revision="app-b")).status_code == 200 + _deploy_ok(client_b, _durable_source(field="value", increment=10, revision="app-b")) submitted_a = client_a.post("/job/run-durable", json={"value": 1}) submitted_b = client_b.post("/job/run-durable", json={"value": 1}) assert _wait_for_completion(client_a, submitted_a)["result"] == {"value": 2} @@ -324,12 +280,12 @@ def test_deploy_rejects_a_runtime_owned_by_another_app() -> None: def test_durable_to_non_durable_overwrite_removes_control_routes() -> None: app = create_app() with TestClient(app) as client: - assert _deploy(client, _durable_source(field="value", increment=1, revision="first")).status_code == 200 + _deploy_ok(client, _durable_source(field="value", increment=1, revision="first")) submitted = client.post("/job/run-durable", json={"value": 1}) execution_id = submitted.json()["execution_id"] _wait_for_completion(client, submitted) - assert _deploy(client, _api_source(increment=5), overwrite=True).status_code == 200 + _deploy_ok(client, _api_source(increment=5), overwrite=True) assert client.post("/job/run-durable", json={"value": 2}).status_code == 404 assert client.get(f"/job/executions/{execution_id}").status_code == 404 assert client.post("/job/run", json={"value": 2}).json() == {"result": 7} @@ -339,7 +295,7 @@ def test_failed_commit_does_not_irreversibly_retire_old_durable_work(monkeypatch app = create_app() with TestClient(app) as client: source = _durable_source(field="value", increment=1, revision="first") - assert _deploy(client, source).status_code == 200 + _deploy_ok(client, source) submitted = client.post("/job/run-durable", json={"value": 1}) url = submitted.json()["links"]["self"] _wait_for_completion(client, submitted) diff --git a/tests/test_durable_execution.py b/tests/test_durable_execution.py index 805b45fd..a66d5cd1 100644 --- a/tests/test_durable_execution.py +++ b/tests/test_durable_execution.py @@ -2,7 +2,7 @@ import importlib.metadata import json import threading -import time +from contextlib import asynccontextmanager, contextmanager from pathlib import Path from unittest.mock import AsyncMock, MagicMock @@ -43,6 +43,7 @@ unload_pipeline_modules, ) from hayhooks.settings import settings +from tests.durable_helpers import wait_for_record, wait_for_status, wait_until_async pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -224,14 +225,15 @@ async def cancellation_requested(self) -> bool: return False -def _agent_record(execution_id: str, deployment_name: str = "agent") -> ExecutionRecord: - return ExecutionRecord( - execution_id=execution_id, - execution_kind=ExecutionKind.AGENT, - deployment_name=deployment_name, - definition_revision="revision", - validated_input={"messages": []}, - ) +def _agent_record(execution_id: str, deployment_name: str = "agent", **changes) -> ExecutionRecord: + fields = { + "execution_id": execution_id, + "execution_kind": ExecutionKind.AGENT, + "deployment_name": deployment_name, + "definition_revision": "revision", + "validated_input": {"messages": []}, + } + return ExecutionRecord(**{**fields, **changes}) class Wrapper(BasePipelineWrapper): @@ -278,10 +280,7 @@ async def run_durable_async(self, context: DurableContext, request: Request) -> async def test_durable_lifecycle_logs_identifiers_without_payload(monkeypatch) -> None: monkeypatch.setattr(settings, "durable_poll_interval", 0.01) - wrapper = Wrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("logged-job", wrapper, InMemoryExecutionStoreProvider()) + deployment = _deployment("logged-job") records = [] sink = log.add(lambda message: records.append(message.record), level="DEBUG") try: @@ -292,12 +291,10 @@ async def test_durable_lifecycle_logs_identifiers_without_payload(monkeypatch) - "Claimed durable execution", "Finished durable execution attempt", } - for _ in range(200): - if expected <= {record["message"] for record in records}: - break - await asyncio.sleep(0.005) - else: - pytest.fail("durable lifecycle logs were not emitted") + await wait_until_async( + lambda: expected <= {record["message"] for record in records}, + "durable lifecycle logs were not emitted", + ) finally: await deployment.close() log.remove(sink) @@ -311,10 +308,7 @@ async def test_durable_lifecycle_logs_identifiers_without_payload(monkeypatch) - async def test_quiesce_waits_for_an_admitted_submission_and_rejects_later_ones(monkeypatch) -> None: monkeypatch.setattr(settings, "durable_poll_interval", 0.01) monkeypatch.setattr(settings, "durable_shutdown_grace_period", 0.001) - wrapper = Wrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("admission-gate", wrapper, InMemoryExecutionStoreProvider()) + deployment = _deployment("admission-gate") await deployment.start() entered = asyncio.Event() release = asyncio.Event() @@ -375,24 +369,13 @@ async def test_deployment_claims_reject_incompatible_work_without_a_revision_sca wrapper = Wrapper() wrapper.durable_revision = "current" - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("revision-safe", wrapper, provider) - await deployment.start() - try: - for _ in range(100): - queued_old = await store.get("queued-old") - if queued_old is not None and queued_old.terminal: - break - await asyncio.sleep(0.001) - else: - pytest.fail("incompatible queued work was not rejected by its first claim") + async with _started(_deployment("revision-safe", wrapper, provider)): + queued_old = await wait_for_record( + store, "queued-old", message="incompatible queued work was not rejected by its first claim" + ) waiting_old = await store.get("waiting-old") waiting_current = await store.get("waiting-current") - finally: - await deployment.close() - assert queued_old is not None assert queued_old.status is ExecutionStatus.FAILED assert queued_old.error is not None assert waiting_old is not None @@ -484,6 +467,34 @@ def setup(self) -> None: self.pipeline = Agent(chat_generator=FakeChatGenerator(), tools=[]) +def _deployment(name: str, wrapper: BasePipelineWrapper | None = None, provider=None, **options) -> DurableDeployment: + """Prepare a wrapper the way the module loader does, then deploy it.""" + wrapper = wrapper or Wrapper() + wrapper.setup() + _set_method_implementation_flags(wrapper) + return DurableDeployment(name, wrapper, provider or InMemoryExecutionStoreProvider(), **options) + + +@contextmanager +def _example_module(module_name: str, source: Path): + """Load one bundled example wrapper module and always unload it.""" + module = load_pipeline_module(module_name, source) + try: + yield module + finally: + unload_pipeline_modules(module_name) + + +@asynccontextmanager +async def _started(deployment: DurableDeployment): + """Run one deployment for the body of a test and always close it.""" + await deployment.start() + try: + yield deployment + finally: + await deployment.close() + + @pytest.fixture(autouse=True) def clean_registry(): registry.clear() @@ -493,6 +504,7 @@ def clean_registry(): def _durable_app(monkeypatch, name, wrapper): monkeypatch.setattr(settings, "durable_store", "memory") + monkeypatch.setattr(settings, "durable_poll_interval", 0.05) wrapper.setup() _set_method_implementation_flags(wrapper) registry.add(name, wrapper) @@ -501,15 +513,6 @@ def _durable_app(monkeypatch, name, wrapper): return app -def _wait_for_status(client, url, status, message): - for _ in range(100): - record = client.get(url) - if record.json()["status"] == status: - return record - time.sleep(0.01) - pytest.fail(message) - - def test_durable_rest_submission_is_direct_typed_and_idempotent(monkeypatch) -> None: app = _durable_app(monkeypatch, "job", Wrapper()) @@ -533,9 +536,9 @@ def test_durable_rest_submission_is_direct_typed_and_idempotent(monkeypatch) -> "updated_at", "links", } - inspected = _wait_for_status(client, body["links"]["self"], "completed", "durable execution did not complete") + inspected = wait_for_status(client, body["links"]["self"], "completed") - assert inspected.json()["result"] == {"value": 5} + assert inspected["result"] == {"value": 5} assert submitted.headers["Location"] == body["links"]["self"] @@ -594,7 +597,7 @@ def test_durable_rest_can_inspect_and_cancel_an_execution_from_an_old_revision(m assert canceled.status_code == 202 assert canceled.json()["cancellation_requested_at"] is not None wrapper.release.set() - _wait_for_status(client, links["self"], "canceled", "durable execution did not cancel") + wait_for_status(client, links["self"], "canceled") assert client.post(links["cancel"]).status_code == 200 finally: wrapper.release.set() @@ -606,10 +609,10 @@ def test_durable_result_annotation_is_validated_before_completion(monkeypatch) - with TestClient(app) as client: submitted = client.post("/invalid-result/run-durable", json={"value": 4}) url = submitted.json()["links"]["self"] - inspected = _wait_for_status(client, url, "failed", "invalid durable result did not become a terminal failure") + inspected = wait_for_status(client, url, "failed") - assert inspected.json()["result"] is None - assert inspected.json()["error"] == { + assert inspected["result"] is None + assert inspected["error"] == { "type": "ValueError", "message": "Durable method result does not match its declared return annotation", "retryable": False, @@ -685,8 +688,8 @@ def test_durable_waiting_resume_is_typed_private_and_revision_safe(monkeypatch) with TestClient(app) as client: submitted = client.post("/approval/run-durable", json={"value": 7}) url = submitted.json()["links"]["self"] - waiting = _wait_for_status(client, url, "waiting", "execution did not wait") - assert waiting.json()["waiting"] == { + waiting = wait_for_status(client, url, "waiting") + assert waiting["waiting"] == { "kind": "approval", "message": "Approve this job", "expected_input_schema": ResumeInput.model_json_schema(), @@ -705,9 +708,9 @@ def test_durable_waiting_resume_is_typed_private_and_revision_safe(monkeypatch) deployment.revision = revision resumed = client.post(f"{url}/resume", json={"approved": True}) assert resumed.status_code == 202 - completed = _wait_for_status(client, url, "completed", "resumed execution did not complete") + completed = wait_for_status(client, url, "completed") - assert completed.json()["result"] == {"value": 7} + assert completed["result"] == {"value": 7} openapi = app.openapi() resume_schema = openapi["paths"]["/approval/executions/{execution_id}/resume"]["post"]["requestBody"]["content"][ "application/json" @@ -739,28 +742,6 @@ def test_durable_rest_enforces_configured_trusted_owner_header(monkeypatch) -> N assert "exceeds 512 characters" in oversized.json()["detail"] -def test_durable_rest_uses_the_server_owner_header_setting(monkeypatch) -> None: - monkeypatch.setattr(settings, "durable_trusted_owner_header", "X-Embedded-Owner") - durable_settings = DurableSettings(durable_store="memory") - provider = InMemoryExecutionStoreProvider(durable_settings=durable_settings) - wrapper = Wrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("embedded", wrapper, provider, durable_settings=durable_settings) - registry.add("embedded", wrapper) - app = create_app() - add_pipeline_api_route(app, "embedded", wrapper, _durable_deployment=deployment) - - with TestClient(app) as client: - assert client.get("/embedded/executions/missing").status_code == 401 - authenticated = client.get( - "/embedded/executions/missing", - headers={"X-Embedded-Owner": "alice"}, - ) - - assert authenticated.status_code == 404 - - def test_durable_deployment_requires_an_explicit_revision() -> None: class MissingRevisionWrapper(BasePipelineWrapper): def setup(self) -> None: @@ -769,11 +750,8 @@ def setup(self) -> None: async def run_durable_async(self, context: DurableContext, request: Request) -> Result: return Result(value=request.value) - wrapper = MissingRevisionWrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) with pytest.raises(Exception, match="non-empty durable_revision"): - DurableDeployment("missing-revision", wrapper, InMemoryExecutionStoreProvider()) + _deployment("missing-revision", MissingRevisionWrapper()) def test_sync_durable_wrapper_uses_context_sync_controls(monkeypatch) -> None: @@ -782,10 +760,10 @@ def test_sync_durable_wrapper_uses_context_sync_controls(monkeypatch) -> None: with TestClient(app) as client: submitted = client.post("/sync-job/run-durable", json={"value": 4}) url = submitted.json()["links"]["self"] - inspected = _wait_for_status(client, url, "completed", "sync durable execution did not complete") + inspected = wait_for_status(client, url, "completed") - assert inspected.json()["result"] == {"value": 6} - assert inspected.json()["progress"][0]["message"] == "working in a worker thread" + assert inspected["result"] == {"value": 6} + assert inspected["progress"][0]["message"] == "working in a worker thread" async def test_sync_work_retains_claim_after_shutdown_grace_until_thread_exits(monkeypatch) -> None: @@ -805,10 +783,7 @@ def run_durable(self, context: DurableContext, request: Request) -> Result: durable_settings = DurableSettings(durable_store="memory", durable_shutdown_grace_period=0.001) provider = InMemoryExecutionStoreProvider(durable_settings=durable_settings) - wrapper = BlockingWrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("blocking", wrapper, provider) + deployment = _deployment("blocking", BlockingWrapper(), provider) await deployment.start() _, submitted = await deployment.submit({"value": 9}) assert await asyncio.to_thread(started.wait, 1) @@ -851,25 +826,13 @@ def blocking_run(_context, _data, *, checkpoint_at): async def test_pipeline_snapshot_round_trip_skips_completed_components_after_retry(monkeypatch) -> None: monkeypatch.setattr(settings, "durable_retry_base_delay", 0) monkeypatch.setattr(settings, "durable_retry_max_delay", 0) - provider = InMemoryExecutionStoreProvider() wrapper = CheckpointPipelineWrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("checkpoint-pipeline", wrapper, provider) - await deployment.start() - try: + async with _started(_deployment("checkpoint-pipeline", wrapper)) as deployment: _, submitted = await deployment.submit({"value": 3}) - for _ in range(200): - completed = await deployment.store.get(submitted.execution_id) - if completed is not None and completed.terminal: - break - await asyncio.sleep(0.005) - else: - pytest.fail("checkpointed Pipeline did not finish its retry") - finally: - await deployment.close() + completed = await wait_for_record( + deployment.store, submitted.execution_id, message="checkpointed Pipeline did not finish its retry" + ) - assert completed is not None assert completed.status.value == "completed" assert completed.result == {"value": 5} assert completed.attempt == 2 @@ -1110,42 +1073,19 @@ async def run_agent_async(self, context, *, messages, **_kwargs): return {"messages": [message.to_dict() for message in state.data["messages"]]} wrapper = BuiltinAgentWrapper() - wrapper.setup() - _set_method_implementation_flags(wrapper) - deployment = DurableDeployment("builtin-agent", wrapper, InMemoryExecutionStoreProvider()) + deployment = _deployment("builtin-agent", wrapper) checkpoint_state = State( schema={"messages": {"type": list[ChatMessage]}}, data={"messages": [ChatMessage.from_user("before restart")]}, ) - record = ExecutionRecord( - execution_id="agent-resume", - execution_kind=ExecutionKind.AGENT, - deployment_name="builtin-agent", + checkpoint_context = DurableContext(_Claim(_agent_record("checkpoint", "builtin-agent")), adapter=object()) + record = _agent_record( + "agent-resume", + "builtin-agent", definition_revision=deployment.revision, validated_input={"messages": [ChatMessage.from_user("initial").to_dict()]}, - checkpoint=ExecutionCheckpoint( - ExecutionKind.AGENT, - _checkpoint_data( - checkpoint_state, - DurableContext( - _Claim( - ExecutionRecord( - execution_id="checkpoint", - execution_kind=ExecutionKind.AGENT, - deployment_name="builtin-agent", - definition_revision=deployment.revision, - validated_input={"messages": []}, - ) - ), - adapter=object(), - ), - ), - ), - application_state={ - "__hayhooks_resume_input": { - "messages": [ChatMessage.from_user("after restart").to_dict()], - } - }, + checkpoint=ExecutionCheckpoint(ExecutionKind.AGENT, _checkpoint_data(checkpoint_state, checkpoint_context)), + application_state={"__hayhooks_resume_input": {"messages": [ChatMessage.from_user("after restart").to_dict()]}}, ) adapter = RestoringAdapter() context = DurableContext(_Claim(record), adapter=adapter) @@ -1162,10 +1102,10 @@ def test_durable_agent_uses_native_run_and_public_hooks(monkeypatch) -> None: with TestClient(app) as client: submitted = client.post("/agent/run-durable", json={"message": "hello"}) url = submitted.json()["links"]["self"] - inspected = _wait_for_status(client, url, "completed", "durable Agent did not complete") + inspected = wait_for_status(client, url, "completed") - assert inspected.json()["result"]["last_message"]["content"][0]["text"] == "done" - assert inspected.json()["progress"][0]["kind"] == "checkpoint" + assert inspected["result"]["last_message"]["content"][0]["text"] == "done" + assert inspected["progress"][0]["kind"] == "checkpoint" @pytest.mark.parametrize( @@ -1177,69 +1117,55 @@ def test_durable_agent_uses_native_run_and_public_hooks(monkeypatch) -> None: ) def test_durable_examples_load(monkeypatch, module_name, source, kind) -> None: monkeypatch.setenv("OPENAI_API_KEY", "test-key") - module = load_pipeline_module(module_name, source) - try: - wrapper = create_pipeline_wrapper_instance(module) - deployment = DurableDeployment(module_name, wrapper, InMemoryExecutionStoreProvider()) + with _example_module(module_name, source) as module: + deployment = DurableDeployment( + module_name, create_pipeline_wrapper_instance(module), InMemoryExecutionStoreProvider() + ) assert deployment.kind is kind assert deployment.revision - finally: - unload_pipeline_modules(module_name) async def test_durable_execution_example_completes_retry_approval_and_real_pipeline(monkeypatch) -> None: monkeypatch.setattr(settings, "durable_retry_base_delay", 0) monkeypatch.setattr(settings, "durable_retry_max_delay", 0) module_name = "durable_execution_end_to_end_example" - module = load_pipeline_module(module_name, _DURABLE_EXECUTION_EXAMPLE) - deployment = None - try: + with _example_module(module_name, _DURABLE_EXECUTION_EXAMPLE) as module: wrapper = create_pipeline_wrapper_instance(module) deployment = DurableDeployment("durable-execution-example", wrapper, InMemoryExecutionStoreProvider()) - await deployment.start() - _, submitted = await deployment.submit( - { - "documents": [{"document_id": "guide", "content": "durable document preparation"}], - "fail_first_attempt": True, - "require_approval": True, - "demo_delay_seconds": 0.01, - } - ) + async with _started(deployment): + _, submitted = await deployment.submit( + { + "documents": [{"document_id": "guide", "content": "durable document preparation"}], + "fail_first_attempt": True, + "require_approval": True, + "demo_delay_seconds": 0.01, + } + ) + await wait_for_record( + deployment, + submitted.execution_id, + lambda record: record.status is ExecutionStatus.WAITING, + message="durable execution example did not reach approval", + ) - for _ in range(200): - waiting = await deployment.get(submitted.execution_id) - if waiting.status is ExecutionStatus.WAITING: - break - await asyncio.sleep(0.005) - else: - pytest.fail("durable execution example did not reach approval") - - assert await deployment.resume(submitted.execution_id, {"approved": True}) - for _ in range(200): - completed = await deployment.get(submitted.execution_id) - if completed.terminal: - break - await asyncio.sleep(0.005) - else: - pytest.fail("durable execution example did not complete") - - assert completed.status is ExecutionStatus.COMPLETED - assert completed.attempt == 3 - assert completed.result["document_count"] == 1 - assert completed.result["chunk_count"] == 1 - assert completed.checkpoint is not None - assert {event.kind for event in completed.progress} >= { - "accepted", - "retry_demo", - "waiting", - "checkpoint", - "demo_delay", - "completed", - } - finally: - if deployment is not None: - await deployment.close() - unload_pipeline_modules(module_name) + assert await deployment.resume(submitted.execution_id, {"approved": True}) + completed = await wait_for_record( + deployment, submitted.execution_id, message="durable execution example did not complete" + ) + + assert completed.status is ExecutionStatus.COMPLETED + assert completed.attempt == 3 + assert completed.result["document_count"] == 1 + assert completed.result["chunk_count"] == 1 + assert completed.checkpoint is not None + assert {event.kind for event in completed.progress} >= { + "accepted", + "retry_demo", + "waiting", + "checkpoint", + "demo_delay", + "completed", + } async def test_durable_a2a_example_tool_replays_its_external_effect_idempotently(monkeypatch, tmp_path) -> None: @@ -1247,8 +1173,7 @@ async def test_durable_a2a_example_tool_replays_its_external_effect_idempotently monkeypatch.setenv("HAYHOOKS_EXAMPLE_INDEX_DB", str(tmp_path / "indexing-effects.sqlite3")) monkeypatch.setenv("HAYHOOKS_EXAMPLE_TOOL_DELAY_SECONDS", "0") module_name = "durable_a2a_tool_example" - module = load_pipeline_module(module_name, _DURABLE_A2A_EXAMPLE) - try: + with _example_module(module_name, _DURABLE_A2A_EXAMPLE) as module: record = _agent_record("a2a-tool-replay", "long-running-agent") claim = _Claim(record) context = DurableContext(claim, adapter=object()) @@ -1268,9 +1193,4 @@ async def test_durable_a2a_example_tool_replays_its_external_effect_idempotently assert json.loads(first)["side_effect_applied"] is True assert json.loads(replay)["side_effect_applied"] is False assert claim.checkpoints == 2 - assert [event.kind for event in record.progress] == [ - "side_effect_committed", - "side_effect_committed", - ] - finally: - unload_pipeline_modules(module_name) + assert [event.kind for event in record.progress] == ["side_effect_committed", "side_effect_committed"] diff --git a/tests/test_durable_fastapi.py b/tests/test_durable_fastapi.py index a20918f9..58f29eda 100644 --- a/tests/test_durable_fastapi.py +++ b/tests/test_durable_fastapi.py @@ -3,7 +3,6 @@ from __future__ import annotations import importlib.metadata -import time from contextlib import asynccontextmanager from types import SimpleNamespace from typing import Annotated @@ -23,6 +22,7 @@ InMemoryExecutionStoreProvider, create_durable_router, ) +from tests.durable_helpers import wait_for_status pytestmark = pytest.mark.skipif( not importlib.metadata.version("haystack-ai").startswith("3."), reason="durable execution requires Haystack 3" @@ -106,15 +106,6 @@ async def authenticate(request: Request, call_next): return app, runtime -def _wait(client: TestClient, url: str, expected: str) -> dict: - for _ in range(200): - response = client.get(url) - if response.json()["status"] == expected: - return response.json() - time.sleep(0.01) - pytest.fail(f"execution did not become {expected}") - - def test_public_router_is_typed_prefix_safe_and_supports_all_routes() -> None: app, _ = _app(owner_dependency=None) with TestClient(app) as client: @@ -123,12 +114,12 @@ def test_public_router_is_typed_prefix_safe_and_supports_all_routes() -> None: assert submitted.headers["Location"].startswith("/api/jobs/executions/") links = submitted.json()["links"] assert set(links) == {"self", "cancel", "resume"} - waiting = _wait(client, links["self"], "waiting") + waiting = wait_for_status(client, links["self"], "waiting") assert waiting["waiting"] == {"kind": "approval"} assert client.post(links["resume"], json={"approved": "invalid"}).status_code == 422 resumed = client.post(links["resume"], json={"approved": True}) assert resumed.status_code == 202 - completed = _wait(client, links["self"], "completed") + completed = wait_for_status(client, links["self"], "completed") assert completed["result"] == {"value": 1, "owner_id": None} assert client.post(links["cancel"]).status_code == 200 diff --git a/tests/test_durable_process_recovery.py b/tests/test_durable_process_recovery.py index 3bddc23c..740ca44b 100644 --- a/tests/test_durable_process_recovery.py +++ b/tests/test_durable_process_recovery.py @@ -10,7 +10,6 @@ import sqlite3 import subprocess import sys -import time import uuid from pathlib import Path @@ -19,6 +18,7 @@ from redis import Redis from hayhooks.server.a2a.redis_task_store import RedisTaskStore +from tests.durable_helpers import wait_until pytestmark = [ pytest.mark.integration, @@ -72,46 +72,42 @@ def _stop_server(server: subprocess.Popen[str]) -> None: server.wait(timeout=3) +def _wait(predicate, message: str): + """Poll for up to ten seconds; every wait in this module shares that budget.""" + return wait_until(predicate, message, attempts=200, delay=0.05) + + def _server_error(server: subprocess.Popen[str]) -> str: output = server.stdout.read() if server.stdout is not None else "" return f"durable test server exited with {server.returncode}:\n{output}" def _wait_for_server(server: subprocess.Popen[str], base_url: str) -> None: - deadline = time.monotonic() + 10 - while time.monotonic() < deadline: + def ready() -> bool: if server.poll() is not None: pytest.fail(_server_error(server)) try: - if requests.get(f"{base_url}/status", timeout=0.25).status_code == 200: - return + return requests.get(f"{base_url}/status", timeout=0.25).status_code == 200 except requests.RequestException: - pass - time.sleep(0.05) - pytest.fail("durable test server did not become ready") + return False + + _wait(ready, "durable test server did not become ready") def _wait_for_file(path: Path) -> None: - deadline = time.monotonic() + 10 - while time.monotonic() < deadline: - if path.exists(): - return - time.sleep(0.05) - pytest.fail("durable test wrapper did not reach its crash window") + _wait(path.exists, "durable test wrapper did not reach its crash window") def _wait_for_completion(execution_url: str) -> dict: - deadline = time.monotonic() + 10 - while time.monotonic() < deadline: + def completed() -> dict | None: response = requests.get(execution_url, timeout=0.5) response.raise_for_status() execution = response.json() - if execution["status"] == "completed": - return execution if execution["status"] in {"failed", "canceled"}: pytest.fail(f"durable execution ended as {execution['status']}: {execution}") - time.sleep(0.05) - pytest.fail("durable execution did not recover to completion") + return execution if execution["status"] == "completed" else None + + return _wait(completed, "durable execution did not recover to completion") def _cleanup_redis(redis_url: str, prefix: str) -> None: @@ -159,13 +155,11 @@ def _a2a_rpc(base_url: str, method: str, params: dict, request_id: str) -> dict: def _wait_for_task_state(base_url: str, task_id: str, state: str) -> dict: - deadline = time.monotonic() + 10 - while time.monotonic() < deadline: + def ready() -> dict | None: task = _a2a_rpc(base_url, "GetTask", {"id": task_id}, f"get-{task_id}") - if task["status"]["state"] == state: - return task - time.sleep(0.05) - pytest.fail(f"A2A task '{task_id}' did not reach {state}") + return task if task["status"]["state"] == state else None + + return _wait(ready, f"A2A task '{task_id}' did not reach {state}") def _a2a_task(result: dict) -> dict: diff --git a/tests/test_durable_reference.py b/tests/test_durable_reference.py index 28efe872..c8db7f4a 100644 --- a/tests/test_durable_reference.py +++ b/tests/test_durable_reference.py @@ -4,7 +4,6 @@ import pytest -from hayhooks.durable.backend import ExecutionStoreConfig from hayhooks.durable.engine import ( Checkpoint, Claim, @@ -14,39 +13,17 @@ Suspend, ) from hayhooks.durable.reference import InMemoryExecutionStore -from tests.durable_contract import assert_store_contract, control +from tests.durable_contract import assert_store_contract, contract_config, control async def test_reference_store_matches_contract() -> None: - store = InMemoryExecutionStore( - deployment="integration", - config=ExecutionStoreConfig( - max_input_bytes=64, - max_checkpoint_bytes=64, - max_result_bytes=64, - max_error_bytes=64, - max_wait_bytes=64, - max_progress_events=2, - max_progress_event_bytes=32, - ), - ) + store = InMemoryExecutionStore(deployment="integration", config=contract_config()) await store.initialize() await assert_store_contract(store) async def test_reference_rejects_oversized_payload_before_transition() -> None: - store = InMemoryExecutionStore( - deployment="integration", - config=ExecutionStoreConfig( - max_input_bytes=64, - max_checkpoint_bytes=8, - max_result_bytes=64, - max_error_bytes=64, - max_wait_bytes=64, - max_progress_events=2, - max_progress_event_bytes=32, - ), - ) + store = InMemoryExecutionStore(deployment="integration", config=contract_config(max_checkpoint_bytes=8)) await store.submit(control(), b"{}", binding_digest="b" * 64) run_id = await store.read_candidate() assert run_id is not None diff --git a/tests/test_durable_store.py b/tests/test_durable_store.py index 863eed42..9521a958 100644 --- a/tests/test_durable_store.py +++ b/tests/test_durable_store.py @@ -4,6 +4,7 @@ import asyncio import time +from contextlib import asynccontextmanager from dataclasses import replace from types import SimpleNamespace from unittest.mock import AsyncMock @@ -29,6 +30,7 @@ from hayhooks.durable.settings import DurableSettings from hayhooks.durable.store import ExecutionStore, InMemoryExecutionStoreProvider, RedisExecutionStoreProvider from hayhooks.settings import AppSettings, settings +from tests.durable_helpers import wait_for_record, wait_until_async def _config() -> ExecutionStoreConfig: @@ -51,6 +53,17 @@ def _store(*, config: ExecutionStoreConfig | None = None, **options) -> Executio ) +@asynccontextmanager +async def _running_manager(store, runner, **options): + """Run one manager for the body of a test and always shut it down.""" + manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001, **options) + await manager.start() + try: + yield manager + finally: + await manager.close() + + def _record() -> ExecutionRecord: return ExecutionRecord( execution_id="run_1", @@ -75,15 +88,12 @@ def test_builtin_providers_snapshot_explicit_durable_settings() -> None: durable_max_attempts=7, durable_max_progress_events=17, durable_max_record_bytes=32_768, + durable_redis_socket_timeout=1.5, + durable_redis_socket_connect_timeout=2.5, + durable_redis_health_check_interval=0, ) memory_store = InMemoryExecutionStoreProvider(durable_settings=durable_settings).create_execution_store("portable") - redis_provider = RedisExecutionStoreProvider( - redis=AsyncMock(), - durable_settings=durable_settings, - socket_timeout=1.5, - socket_connect_timeout=2.5, - health_check_interval=0, - ) + redis_provider = RedisExecutionStoreProvider(redis=AsyncMock(), durable_settings=durable_settings) redis_store = redis_provider.create_execution_store("portable") for store in (memory_store, redis_store): @@ -108,7 +118,7 @@ def test_runtime_uses_its_explicit_settings_for_the_default_provider() -> None: provider = runtime._provider() assert isinstance(provider, InMemoryExecutionStoreProvider) - assert provider.app_settings.durable_lease_duration_ms == 45_000 + assert provider.settings.durable_lease_duration_ms == 45_000 async def test_runtime_defaults_are_independent_of_hayhooks_settings(monkeypatch: pytest.MonkeyPatch) -> None: @@ -119,7 +129,7 @@ async def test_runtime_defaults_are_independent_of_hayhooks_settings(monkeypatch monkeypatch.setattr(settings, "durable_max_attempts", original_attempts + 1) - assert runtime.app_settings.durable_max_attempts == provider.app_settings.durable_max_attempts == original_attempts + assert runtime.settings.durable_max_attempts == provider.settings.durable_max_attempts == original_attempts await runtime.close() assert runtime.settings.durable_max_attempts == original_attempts @@ -131,8 +141,8 @@ def test_runtime_uses_supplied_builtin_provider_as_its_settings_source() -> None store = provider.create_execution_store("portable") - assert runtime.app_settings.durable_lease_duration_ms == store.lease_duration_ms == 45_000 - assert runtime.app_settings.durable_max_attempts == store.max_run_attempts == 7 + assert runtime.settings.durable_lease_duration_ms == store.lease_duration_ms == 45_000 + assert runtime.settings.durable_max_attempts == store.max_run_attempts == 7 def test_runtime_rejects_conflicting_builtin_provider_settings() -> None: @@ -479,23 +489,10 @@ async def test_retry_exhaustion_persists_its_progress_event() -> None: async def runner(context): await context.retry("again", delay=0) - manager = DurableExecutionManager( - "deployment", store, runner, adapter=object(), poll_interval=0.001, max_attempts=1 - ) - await manager.start() - try: + async with _running_manager(store, runner, max_attempts=1): await store.submit(_record()) - for _ in range(100): - record = await store.get("run_1") - if record is not None and record.terminal: - break - await asyncio.sleep(0.001) - else: - pytest.fail("retry exhaustion did not become terminal") - finally: - await manager.close() + record = await wait_for_record(store, "run_1", message="retry exhaustion did not become terminal") - assert record is not None assert [event.kind for event in record.progress] == ["retry_exhausted"] @@ -512,23 +509,14 @@ async def runner(context): await context.report_progress("started") return {"answer": "done"} - manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) - await manager.start() - try: + async with _running_manager(store, runner): await store.submit(_record()) - for _ in range(100): - record = await store.get("run_1") - if record and record.terminal: - break - await asyncio.sleep(0.001) - else: - pytest.fail("manager did not complete the submitted execution") - assert record.status is ExecutionStatus.COMPLETED - assert record.result == {"answer": "done"} - assert record.application_state == {"phase": "running"} - assert [event.message for event in record.progress] == ["started"] - finally: - await manager.close() + record = await wait_for_record(store, "run_1", message="manager did not complete the submitted execution") + + assert record.status is ExecutionStatus.COMPLETED + assert record.result == {"answer": "done"} + assert record.application_state == {"phase": "running"} + assert [event.message for event in record.progress] == ["started"] async def test_manager_health_reports_worker_store_failures() -> None: @@ -538,18 +526,9 @@ async def test_manager_health_reports_worker_store_failures() -> None: maintain=AsyncMock(), operational_counts=AsyncMock(return_value={"nonterminal": 1, "runnable": 1, "lease_expiry": 0}), ) - manager = DurableExecutionManager( - "deployment", store, AsyncMock(), adapter=object(), poll_interval=0.001, shutdown_grace_period=0.01 - ) - await manager.start() - try: - for _ in range(100): - if store.claim_next.await_count: - break - await asyncio.sleep(0.001) + async with _running_manager(store, AsyncMock(), shutdown_grace_period=0.01) as manager: + await wait_until_async(lambda: store.claim_next.await_count, "worker never attempted a claim") health = await manager.health_snapshot() - finally: - await manager.close() assert not health["healthy"] assert health["worker_store_error_streak"] >= 1 @@ -569,22 +548,11 @@ async def runner(_context): raise asyncio.CancelledError return {"answer": "done"} - manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) - await manager.start() - try: + async with _running_manager(store, runner): await store.submit(_record()) - for _ in range(200): - record = await store.get("run_1") - if record is not None and record.terminal: - break - await asyncio.sleep(0.005) - else: - pytest.fail("canceled runner did not recover") - finally: - await manager.close() + record = await wait_for_record(store, "run_1", message="canceled runner did not recover") assert calls == 2 - assert record is not None assert record.status is ExecutionStatus.COMPLETED @@ -636,19 +604,9 @@ async def runner(_context): raise RetryableExecutionError(msg, delay=0.01) return {"answer": "done"} - manager = DurableExecutionManager("deployment", store, runner, adapter=object(), poll_interval=0.001) - await manager.start() - try: + async with _running_manager(store, runner): await store.submit(_record()) - for _ in range(100): - record = await store.get("run_1") - if record is not None and record.terminal: - break - await asyncio.sleep(0.005) - else: - pytest.fail("in-memory durable retry did not become due") - finally: - await manager.close() + record = await wait_for_record(store, "run_1", message="in-memory durable retry did not become due") assert attempts == 2 assert record.status is ExecutionStatus.COMPLETED diff --git a/tests/test_it_a2a_server.py b/tests/test_it_a2a_server.py index 438f6291..c5b249f5 100644 --- a/tests/test_it_a2a_server.py +++ b/tests/test_it_a2a_server.py @@ -250,65 +250,56 @@ async def long_running_client(): yield client, wrapper -async def poll_task_until_state(client: httpx.AsyncClient, task_id: str, expected_state: str) -> dict: +async def poll_task(client: httpx.AsyncClient, task_id: str, ready, description: str, *, legacy: bool = False) -> dict: + """Poll one task until *ready* accepts its projection; ``legacy`` uses the v0.3 alias.""" + payload = get_task_payload(task_id, method="tasks/get" if legacy else "GetTask") + headers = {"A2A-Version": "0.3"} if legacy else None last_task = None for _ in range(50): - response = await client.post("/long_agent/", json=get_task_payload(task_id)) + response = await client.post("/long_agent/", json=payload, headers=headers) assert response.status_code == 200 last_task = extract_task(response.json()) - if last_task["status"]["state"] == expected_state: + if ready(last_task): return last_task await asyncio.sleep(0.01) - msg = f"Task {task_id} did not reach {expected_state}. Last task: {last_task}" + msg = f"Task {task_id} did not {description}. Last task: {last_task}" raise AssertionError(msg) -async def poll_task_until_artifact_text(client: httpx.AsyncClient, task_id: str, expected_text: str) -> dict: - last_task = None - for _ in range(50): - response = await client.post("/long_agent/", json=get_task_payload(task_id)) - assert response.status_code == 200 - last_task = extract_task(response.json()) - if artifact_text(last_task) == expected_text: - return last_task - await asyncio.sleep(0.01) - msg = f"Task {task_id} did not expose artifact text {expected_text!r}. Last task: {last_task}" - raise AssertionError(msg) +async def start_detached_task(client: httpx.AsyncClient) -> dict: + """Send one message that returns before the agent has finished.""" + payload = send_message_payload("start") + payload["params"]["configuration"] = {"returnImmediately": True} + response = await client.post("/long_agent/", json=payload) + assert response.status_code == 200 + return extract_task(response.json()) async def test_detached_send_returns_non_terminal_task(long_running_client): client, wrapper = long_running_client - payload = send_message_payload("start") - payload["params"]["configuration"] = {"returnImmediately": True} - response = await client.post("/long_agent/", json=payload) + task = await start_detached_task(client) - assert response.status_code == 200 - task = extract_task(response.json()) assert task["id"] assert task["contextId"] assert task["status"]["state"] in {"TASK_STATE_SUBMITTED", "TASK_STATE_WORKING"} assert wrapper.entered.is_set() assert not wrapper.release.is_set() - progress_task = await poll_task_until_artifact_text(client, task["id"], "progress ") + progress_task = await poll_task( + client, task["id"], lambda task: artifact_text(task) == "progress ", "expose its progress artifact" + ) assert progress_task["status"]["state"] == "TASK_STATE_WORKING" wrapper.release.set() - last_task = None - for _ in range(50): - poll_response = await client.post( - "/long_agent/", - json=get_task_payload(task["id"], method="tasks/get"), - headers={"A2A-Version": "0.3"}, - ) - assert poll_response.status_code == 200 - last_task = extract_task(poll_response.json()) - if last_task["status"]["state"] == "completed": - break - await asyncio.sleep(0.01) - assert last_task is not None - assert last_task["status"]["state"] == "completed" + # The v0.3 alias must project the same task with lowercase states. + last_task = await poll_task( + client, + task["id"], + lambda task: task["status"]["state"] == "completed", + "complete", + legacy=True, + ) assert artifact_text(last_task) == "progress done" @@ -338,16 +329,15 @@ async def test_a2a_v0_3_blocking_false_returns_active_task(long_running_client): assert task["status"]["state"] in {"submitted", "working"} wrapper.release.set() - completed_task = await poll_task_until_state(client, task["id"], "TASK_STATE_COMPLETED") + completed_task = await poll_task( + client, task["id"], lambda task: task["status"]["state"] == "TASK_STATE_COMPLETED", "complete" + ) assert artifact_text(completed_task) == "progress done" async def test_subscribe_to_active_task(long_running_client): client, wrapper = long_running_client - payload = send_message_payload("start") - payload["params"]["configuration"] = {"returnImmediately": True} - send_response = await client.post("/long_agent/", json=payload) - task = extract_task(send_response.json()) + task = await start_detached_task(client) subscribe_payload = get_task_payload(task["id"], method="SubscribeToTask") @@ -377,10 +367,7 @@ async def read_subscription_events() -> list[dict]: async def test_cooperative_async_cancellation(long_running_client): client, wrapper = long_running_client - payload = send_message_payload("start") - payload["params"]["configuration"] = {"returnImmediately": True} - send_response = await client.post("/long_agent/", json=payload) - task = extract_task(send_response.json()) + task = await start_detached_task(client) await asyncio.wait_for(wrapper.entered.wait(), timeout=1) cancel_response = await client.post("/long_agent/", json=cancel_task_payload(task["id"])) diff --git a/tests/test_redis_a2a_recovery_integration.py b/tests/test_redis_a2a_recovery_integration.py index 00add4af..026b36fb 100644 --- a/tests/test_redis_a2a_recovery_integration.py +++ b/tests/test_redis_a2a_recovery_integration.py @@ -8,6 +8,8 @@ from unittest.mock import AsyncMock import pytest +from a2a.server.context import ServerCallContext +from a2a.types import ListTasksRequest, Task, TaskState from hayhooks.durable.models import ExecutionStatus from hayhooks.durable.runtime import execution_id_for @@ -43,8 +45,6 @@ async def redis_task_store(isolated_redis): def _task(task_id: str, seconds: int = 0): - from a2a.types import Task, TaskState - task = Task(id=task_id, context_id=f"context-{task_id}") task.status.state = TaskState.TASK_STATE_WORKING task.status.timestamp.FromDatetime(datetime(2026, 1, 1, tzinfo=timezone.utc) + timedelta(seconds=seconds)) @@ -56,9 +56,6 @@ def _context(owner: str): async def test_redis_a2a_recovery_uses_atomic_projection_and_owner(redis_task_store) -> None: - from a2a.server.context import ServerCallContext - from a2a.types import TaskState - redis, store = redis_task_store context = ServerCallContext() task = _task("client.task/" + "x" * 160) @@ -86,9 +83,6 @@ async def test_redis_a2a_recovery_uses_atomic_projection_and_owner(redis_task_st async def test_redis_a2a_read_through_persists_late_completion(redis_task_store) -> None: - from a2a.server.context import ServerCallContext - from a2a.types import TaskState - redis, store = redis_task_store context = ServerCallContext() task = _task("late-completion") @@ -110,8 +104,6 @@ async def test_redis_a2a_read_through_persists_late_completion(redis_task_store) async def test_redis_task_store_isolates_agents_and_owners(redis_task_store) -> None: - from a2a.types import TaskState - redis, base_store = redis_task_store store = RedisTaskStore(redis, "agent/one", key_prefix=base_store.key_prefix) other_agent = RedisTaskStore(redis, "agent/two", key_prefix=base_store.key_prefix) @@ -172,8 +164,6 @@ async def test_all_task_store_writes_reject_a_stale_loaded_version(redis_task_st async def test_same_store_tracks_versions_per_loaded_task_snapshot(redis_task_store) -> None: - from a2a.types import TaskState - _, store = redis_task_store context = _context("alice@example.com") await store.save(_task("task", 1), context) @@ -217,8 +207,6 @@ async def test_concurrent_projection_writers_use_one_version_fence(redis_task_st async def test_cleanup_preserves_a_task_whose_terminal_ttl_was_extended(monkeypatch, redis_task_store) -> None: - from a2a.types import TaskState - redis, base_store = redis_task_store cleanup = RedisTaskStore(redis, "cleanup-race", key_prefix=base_store.key_prefix, terminal_ttl_seconds=60) writer = RedisTaskStore(redis, "cleanup-race", key_prefix=base_store.key_prefix, terminal_ttl_seconds=60) @@ -258,8 +246,6 @@ async def delayed_delete(*args, **kwargs): async def test_redis_task_store_lists_with_filters_and_page_tokens(redis_task_store) -> None: - from a2a.types import ListTasksRequest - _, store = redis_task_store context = _context("alice@example.com") for index in range(3): @@ -285,8 +271,6 @@ async def test_redis_task_store_lists_with_filters_and_page_tokens(redis_task_st async def test_redis_task_store_compares_status_timestamps_as_timestamps(redis_task_store) -> None: - from a2a.types import ListTasksRequest - _, store = redis_task_store context = _context("alice@example.com") task = _task("task") @@ -301,8 +285,6 @@ async def test_redis_task_store_compares_status_timestamps_as_timestamps(redis_t async def test_filtered_task_listing_and_snapshot_cache_are_globally_bounded(monkeypatch, redis_task_store) -> None: - from a2a.types import ListTasksRequest - from hayhooks.settings import settings monkeypatch.setattr(settings, "a2a_list_scan_batch_size", 2) diff --git a/tests/test_redis_execution_integration.py b/tests/test_redis_execution_integration.py index 9975e6a8..6abf63d5 100644 --- a/tests/test_redis_execution_integration.py +++ b/tests/test_redis_execution_integration.py @@ -7,7 +7,6 @@ import pytest -from hayhooks.durable.backend import ExecutionStoreConfig from hayhooks.durable.engine import ( Checkpoint, Claim, @@ -29,7 +28,8 @@ ) from hayhooks.durable.redis import RedisExecutionStore from hayhooks.durable.store import ExecutionStore -from tests.durable_contract import assert_store_contract, control +from tests.durable_contract import assert_store_contract, contract_config, control +from tests.durable_helpers import wait_for_record pytestmark = pytest.mark.integration @@ -37,17 +37,7 @@ @pytest.fixture async def store(isolated_redis): redis, prefix = isolated_redis - config = ExecutionStoreConfig( - key_prefix=f"{prefix}:durable", - max_input_bytes=64, - max_checkpoint_bytes=64, - max_result_bytes=64, - max_error_bytes=64, - max_wait_bytes=64, - max_progress_events=2, - max_progress_event_bytes=32, - terminal_ttl_seconds=60, - ) + config = contract_config(key_prefix=f"{prefix}:durable", terminal_ttl_seconds=60) durable = RedisExecutionStore(redis, deployment="integration", config=config) await durable.initialize() yield redis, durable @@ -239,13 +229,9 @@ async def runner(context): max_record_bytes=512, ) ) - for _ in range(100): - record = await adapter_store.get("public-run") - if record is not None and record.terminal: - break - await asyncio.sleep(0.01) - else: - pytest.fail("Redis adapter did not complete the public manager execution") + record = await wait_for_record( + adapter_store, "public-run", message="Redis adapter did not complete the public manager execution" + ) finally: await manager.close() diff --git a/tests/test_settings.py b/tests/test_settings.py index fe329090..43dda9a0 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -50,67 +50,45 @@ def test_env_var_prefix(monkeypatch): assert settings.port == 5678 -def test_durable_redis_settings_defaults(monkeypatch): - names = ( - "HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS", - "HAYHOOKS_DURABLE_POLL_INTERVAL", - "HAYHOOKS_DURABLE_LEASE_DURATION_MS", - "HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS", - "HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT", - "HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT", - "HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL", - ) - for name in names: - monkeypatch.delenv(name, raising=False) - - settings = AppSettings() - - assert settings.durable_lease_duration_ms == 30_000 - assert settings.durable_lease_commit_safety_ms == 1_500 - assert settings.durable_poll_interval == 1.0 - assert settings.durable_max_nonterminal_executions == 0 - assert settings.durable_redis_socket_timeout == 5.0 - assert settings.durable_redis_socket_connect_timeout == 5.0 - assert settings.durable_redis_health_check_interval == 30 - - -def test_durable_redis_settings_from_environment(monkeypatch): - monkeypatch.setenv("HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS", "250") - monkeypatch.setenv("HAYHOOKS_DURABLE_POLL_INTERVAL", "0.5") - monkeypatch.setenv("HAYHOOKS_DURABLE_LEASE_DURATION_MS", "45000") - monkeypatch.setenv("HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS", "2000") - monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT", "3.5") - monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT", "2.5") - monkeypatch.setenv("HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL", "20") - - settings = AppSettings() - - assert settings.durable_lease_duration_ms == 45_000 - assert settings.durable_lease_commit_safety_ms == 2_000 - assert settings.durable_poll_interval == 0.5 - assert settings.durable_max_nonterminal_executions == 250 - assert settings.durable_redis_socket_timeout == 3.5 - assert settings.durable_redis_socket_connect_timeout == 2.5 - assert settings.durable_redis_health_check_interval == 20 - - -def test_durable_lease_safety_margin_must_leave_time_for_a_commit() -> None: - with pytest.raises(ValueError, match="durable_lease_commit_safety_ms"): - AppSettings(durable_lease_duration_ms=1_000, durable_lease_commit_safety_ms=1_000) - with pytest.raises(ValueError, match="heartbeat interval"): - AppSettings(durable_lease_duration_ms=1, durable_lease_commit_safety_ms=0) - with pytest.raises(ValueError, match="heartbeat interval"): - AppSettings(durable_lease_duration_ms=1_000, durable_lease_commit_safety_ms=700) - - -def test_a2a_task_store_bounds_from_environment(monkeypatch): - monkeypatch.setenv("HAYHOOKS_A2A_TASK_SNAPSHOT_CACHE_SIZE", "64") - monkeypatch.setenv("HAYHOOKS_A2A_LIST_SCAN_BATCH_SIZE", "25") +# (env var, settings field, default, env value, parsed value) +_ENV_SETTINGS = [ + ("HAYHOOKS_DURABLE_LEASE_DURATION_MS", "durable_lease_duration_ms", 30_000, "45000", 45_000), + ("HAYHOOKS_DURABLE_LEASE_COMMIT_SAFETY_MS", "durable_lease_commit_safety_ms", 1_500, "2000", 2_000), + ("HAYHOOKS_DURABLE_POLL_INTERVAL", "durable_poll_interval", 1.0, "0.5", 0.5), + ("HAYHOOKS_DURABLE_MAX_NONTERMINAL_EXECUTIONS", "durable_max_nonterminal_executions", 0, "250", 250), + ("HAYHOOKS_DURABLE_REDIS_SOCKET_TIMEOUT", "durable_redis_socket_timeout", 5.0, "3.5", 3.5), + ("HAYHOOKS_DURABLE_REDIS_SOCKET_CONNECT_TIMEOUT", "durable_redis_socket_connect_timeout", 5.0, "2.5", 2.5), + ("HAYHOOKS_DURABLE_REDIS_HEALTH_CHECK_INTERVAL", "durable_redis_health_check_interval", 30, "20", 20), + ("HAYHOOKS_A2A_TASK_SNAPSHOT_CACHE_SIZE", "a2a_task_snapshot_cache_size", 1_024, "64", 64), + ("HAYHOOKS_A2A_LIST_SCAN_BATCH_SIZE", "a2a_list_scan_batch_size", 500, "25", 25), +] + + +@pytest.mark.parametrize("from_environment", [False, True], ids=["defaults", "environment"]) +def test_durable_and_a2a_settings_follow_the_environment(monkeypatch, from_environment): + for name, _field, _default, value, _expected in _ENV_SETTINGS: + if from_environment: + monkeypatch.setenv(name, value) + else: + monkeypatch.delenv(name, raising=False) configured = AppSettings() - assert configured.a2a_task_snapshot_cache_size == 64 - assert configured.a2a_list_scan_batch_size == 25 + for _name, field, default, _value, expected in _ENV_SETTINGS: + assert getattr(configured, field) == (expected if from_environment else default) + + +@pytest.mark.parametrize( + ("duration_ms", "safety_ms", "match"), + [ + (1_000, 1_000, "durable_lease_commit_safety_ms"), + (1, 0, "heartbeat interval"), + (1_000, 700, "heartbeat interval"), + ], +) +def test_durable_lease_safety_margin_must_leave_time_for_a_commit(duration_ms, safety_ms, match) -> None: + with pytest.raises(ValueError, match=match): + AppSettings(durable_lease_duration_ms=duration_ms, durable_lease_commit_safety_ms=safety_ms) def test_cors(): From 048a58f55c4f33a23727ab100025a019c7e080c6 Mon Sep 17 00:00:00 2001 From: Michele Pangrazzi Date: Wed, 19 Aug 2026 14:33:43 +0200 Subject: [PATCH 20/28] feat(durable): add SSE execution streaming with reattachable chunk log --- docs/advanced/durable-engine.md | 86 ++++++++++- docs/advanced/durable-execution-operations.md | 30 +++- examples/README.md | 6 +- examples/durable_chat_with_website/README.md | 110 +++++++++++++ examples/durable_chat_with_website/demo.py | 119 ++++++++++++++ .../chat_with_website/chat_with_website.yml | 48 ++++++ .../chat_with_website/pipeline_wrapper.py | 83 ++++++++++ src/hayhooks/durable/backend.py | 35 ++++- src/hayhooks/durable/context.py | 17 +- src/hayhooks/durable/fastapi.py | 82 +++++++++- src/hayhooks/durable/models.py | 1 + src/hayhooks/durable/redis.py | 37 +++++ src/hayhooks/durable/reference.py | 30 ++++ src/hayhooks/durable/settings.py | 4 + src/hayhooks/durable/store.py | 33 +++- src/hayhooks/server/durable/routes.py | 1 + src/hayhooks/server/tracing.py | 3 +- tests/durable_contract.py | 23 ++- tests/durable_helpers.py | 91 +++++++++++ tests/test_durable_execution.py | 146 +++++++++++++++++- tests/test_durable_fastapi.py | 101 +++++++++++- tests/test_durable_process_recovery.py | 96 ++---------- tests/test_durable_reference.py | 16 ++ tests/test_it_durable_chat_with_website.py | 119 ++++++++++++++ tests/test_redis_execution_integration.py | 28 ++++ 25 files changed, 1233 insertions(+), 112 deletions(-) create mode 100644 examples/durable_chat_with_website/README.md create mode 100644 examples/durable_chat_with_website/demo.py create mode 100644 examples/durable_chat_with_website/pipelines/chat_with_website/chat_with_website.yml create mode 100644 examples/durable_chat_with_website/pipelines/chat_with_website/pipeline_wrapper.py create mode 100644 tests/test_it_durable_chat_with_website.py diff --git a/docs/advanced/durable-engine.md b/docs/advanced/durable-engine.md index 0be7c6b2..ae5e792a 100644 --- a/docs/advanced/durable-engine.md +++ b/docs/advanced/durable-engine.md @@ -15,6 +15,7 @@ its next checkpoint. - Pipeline snapshots and Agent state checkpoints with bounded retries, progress, cancellation, and typed wait/resume. - Redis-backed fenced claims and lease recovery across process restarts. +- Live SSE chunk streaming that clients can detach from and reattach to. - Idempotent submission, optional owner-isolated REST access, and managed A2A task projection. - Native Redis TTL for terminal records, plus an equivalent volatile @@ -167,13 +168,14 @@ runtime's lifetime; create a new runtime to change storage backends. ## Redis layout Each deployment has an isolated namespace with controls and opaque input, -checkpoint, result, error, wait, and progress payload keys. It has exactly two -sorted-set indexes: +checkpoint, result, error, wait, and progress payload keys, plus one `chunks` +stream per execution. It has exactly two sorted-set indexes: | Key | Purpose | |---|---| | `runnable` | All queued work, scored by its retry deadline or immediate transition time. | | `lease-expiry` | Running fences, scored by their Redis-server lease deadline. | +| `exec::chunks` | Bounded append-only display chunks, read by the execution SSE stream. | The namespace also contains a `capacity` hash with only `nonterminal` and one idempotency binding per execution. Terminal execution and idempotency keys use @@ -193,6 +195,86 @@ three replicas and low-to-moderate load. Its two indexes, native TTL, and single reducer keep the worker model observable during normal operation and recovery. +## Streaming chunks + +`GET /{pipeline}/executions/{execution_id}/stream` is a Server-Sent Events +stream of one execution's display chunks, followed by a terminal `completed`, +`failed`, or `canceled` event carrying the same public projection the inspect +route returns. The submit response advertises it as the `stream` link. + +Chunks are best-effort display data, deliberately outside the durable fence: a +chunk append is a single `XADD` with no transaction, so token-rate streaming +cannot contend with the heartbeat. Nothing about a chunk can fail a run +either: an oversized payload (chunks are capped at 64 KB) or a backend blip +drops that chunk and logs it, because replaying a pipeline to recover a display +token is never the right trade. Progress events remain the coarse durable audit +trail. + +```python +class StreamingWrapper(BasePipelineWrapper): + durable_revision = "streaming-v1" + + async def run_durable_async(self, context: DurableContext, request: Question) -> Answer: + result = await context.run_agent_async( + messages=[ChatMessage.from_user(request.query)], + streaming_callback=context.stream_chunk, + ) + return Answer(reply=result["messages"][-1].text) +``` + +Bind the callback to the component when the work is a Pipeline rather than an +Agent. `async_streaming_generator` passes it per run in `pipeline_run_args`, and +that does not survive here: `run_pipeline_async` data is serialized into the +`PipelineSnapshot`, Haystack drops the callable it cannot serialize, and +`Pipeline.run` rebuilds its `data` from the snapshot when resuming, so a callback +passed as run data disappears at the first checkpoint. + +Binding one callback to a shared component is safe for the same reason +`async_streaming_generator` can hand the same module-level +`_async_streaming_callback` to every concurrent run: the callback carries no +per-run state and resolves its destination on each call from a `ContextVar`. +Hayhooks routes on `_ASYNC_STREAMING_QUEUE`; the durable path routes on the +execution context, which the engine sets per execution task and `asyncio.to_thread` +copies into the Pipeline's worker thread: + +```python +def stream_to_execution(chunk: StreamingChunk) -> None: + if context := current_durable_context(): + context.stream_chunk_sync(chunk) + + +class StreamingPipelineWrapper(BasePipelineWrapper): + durable_revision = "streaming-pipeline-v1" + + def setup(self) -> None: + self.pipeline = Pipeline.loads(...) + self.pipeline.get_component("llm").streaming_callback = stream_to_execution +``` + +`run_pipeline_async` drives the Pipeline on a worker thread, which is why that +path uses `context.stream_chunk_sync` rather than the awaitable form. What +`pipeline_run_args` injection actually buys `async_streaming_generator` is +control over *which* components stream without permanently mutating a shared +Pipeline; here the callback simply does nothing outside a durable execution. A +run-time `streaming_callback` still takes precedence over a bound one, so an +ordinary streaming endpoint on the same wrapper is unaffected. See +`examples/durable_chat_with_website` for the whole wrapper. Each SSE event +carries the entry ID as `id:` and the producing `attempt` in its payload; a +client resets its buffer when `attempt` increases, because a retried attempt +re-streams from its checkpoint. Reconnecting clients resend `Last-Event-ID` +automatically and resume from that cursor. + +`durable_max_stream_chunks` bounds the log per execution (10 000 by default); +`0` disables chunk production entirely while leaving the endpoint working. +`durable_max_stream_chunk_bytes` caps a single chunk at 64 KB by default; an +oversized chunk is dropped, never failed. +The log expires with its execution under `durable_terminal_ttl_seconds`. + +A stream that breaks after its headers were sent has no status code left to +report with, so it ends in an `error` event instead. Treat it the way a client +treats a dropped connection: reattach with `Last-Event-ID`, or read the +execution's terminal state from the inspect route. + ## Revisions and rollout Every durable wrapper, including a managed A2A Agent, declares a non-empty diff --git a/docs/advanced/durable-execution-operations.md b/docs/advanced/durable-execution-operations.md index 74f6ba16..9317f584 100644 --- a/docs/advanced/durable-execution-operations.md +++ b/docs/advanced/durable-execution-operations.md @@ -20,10 +20,10 @@ external write so recovered work remains safe to replay. ## Execution and recovery -The namespace holds a control record, opaque payloads, one `runnable` -ZSET, one `lease-expiry` ZSET, a `nonterminal` capacity field, and idempotency -bindings. The control is authoritative; the indexes are derived atomically -with it. +The namespace holds a control record, opaque payloads, one bounded `chunks` +stream per execution, one `runnable` ZSET, one `lease-expiry` ZSET, a +`nonterminal` capacity field, and idempotency bindings. The control is +authoritative; the indexes are derived atomically with it. Workers poll due runnable work at the configured interval using Redis `TIME`. Candidate reads are non-destructive. A watched control hash and monotonically @@ -34,13 +34,29 @@ until due. ## Retention and rollout -Terminal control/payload keys and their idempotency binding receive the -configured Redis TTL when a run first becomes terminal. Memory uses equivalent -internal cleanup. Do not delete records manually while they are nonterminal. +Terminal control, payload, and chunk keys and their idempotency binding +receive the configured Redis TTL when a run first becomes terminal. +`HAYHOOKS_DURABLE_MAX_STREAM_CHUNKS` bounds each execution's SSE chunk log and +`0` disables it, which is the kill switch if streaming misbehaves. +`HAYHOOKS_DURABLE_MAX_STREAM_CHUNK_BYTES` caps one chunk (64 KB by default); +oversized chunks are dropped, never failed. A stale append from a worker that +lost its lease after the run went terminal sets the TTL on the recreated +chunk key itself, so the log still expires with its execution. Memory uses +equivalent internal cleanup. Do not delete records manually while they are +nonterminal. Begin a new controlled-beta deployment with an empty durable namespace, then retain its terminal records through the configured Redis TTL. +## Streaming load + +Every attached stream holds one Redis connection blocked in `XREAD` for up to +500 ms per poll cycle, so the number of concurrent viewers counts directly +against Redis client capacity. A terminal event arrives within about 500 ms of +the last chunk once the stream goes quiet; a run that keeps generating after a +cancellation request reports terminal only after it stops producing chunks. +Chunks are display data and are never a reason to replay a run. + ## Health and incidents Health exposes `nonterminal`, `runnable`, `lease_expiry`, and diff --git a/examples/README.md b/examples/README.md index 39efef35..fa5efd70 100644 --- a/examples/README.md +++ b/examples/README.md @@ -24,6 +24,7 @@ This directory contains various examples demonstrating different use cases and f | [a2a_multi_agent](./a2a_multi_agent/) | Two agents with their own MCP tools, communicating over A2A | • `hayhooks a2a run` hosting two agents
• Per-agent A2A agent cards
• Agent-to-agent delegation via A2A client tool
• One MCP tool server per agent (FastMCP)
• Streaming A2A client | Building multi-agent systems where Haystack Agents expose themselves over A2A and delegate tasks to each other while using MCP for their own tools | | [a2a_long_running](./a2a_long_running/) | Recoverable OpenAI agent over A2A | • `OpenAIChatGenerator`-based Haystack Agent
• Redis-backed fenced execution checkpoints
• Hayhooks and client restart recovery
• `input-required` continuation
• Polling and durable cancellation | Building tool-using A2A agents whose accepted work resumes after process restarts | | [durable_execution](./durable_execution/) | First-class durable Pipeline | • Typed `/run-durable` REST endpoint
• Built-in Redis store and fenced claims
• Restart recovery, inspection, cancellation, and resume | Running recoverable background jobs through an ordinary wrapper without A2A | +| [durable_chat_with_website](./durable_chat_with_website/) | Durable Pipeline that streams its answer | • SSE token streaming over `/executions/{id}/stream`
• Detach and reattach with `Last-Event-ID`
• Checkpointed page fetch, streamed generation
• Bounded chunk log outside the durable fence | Giving a recoverable job a live UI without making display chunks part of durable state | | [rag_indexing_query](./rag_indexing_query/) | Complete RAG system with Elasticsearch | • Document indexing pipeline
• Query pipeline
• Elasticsearch integration
• Multiple file format support (PDF, Markdown, Text)
• Sentence transformers embeddings | Implementing production-ready RAG systems for document search and knowledge retrieval | | [shared_code_between_wrappers](./shared_code_between_wrappers/) | Code sharing between pipeline wrappers | • Shared library imports
• HAYHOOKS_ADDITIONAL_PYTHON_PATH
• Multiple deployment strategies
• Code reusability | Organizing complex projects with multiple pipelines that share common functionality | @@ -49,7 +50,10 @@ For the durable examples, use this presentation order: 1. [`durable_execution`](./durable_execution/) — the deterministic reference for typed submission, retry, approval, checkpoints, crash recovery, and cancellation. -2. [`a2a_long_running`](./a2a_long_running/) — durable Agent execution exposed +2. [`durable_chat_with_website`](./durable_chat_with_website/) — the same + engine with a live SSE token stream, showing where display chunks sit + relative to the durable fence. +3. [`a2a_long_running`](./a2a_long_running/) — durable Agent execution exposed through standard A2A task lifecycle and continuation messages. Each durable example's Compose file publishes Redis on `localhost:6379`. diff --git a/examples/durable_chat_with_website/README.md b/examples/durable_chat_with_website/README.md new file mode 100644 index 00000000..01f3739a --- /dev/null +++ b/examples/durable_chat_with_website/README.md @@ -0,0 +1,110 @@ +# Durable chat-with-website with token streaming + +A durable Pipeline that fetches live web pages, answers a question about them, +and streams the answer token by token over Server-Sent Events. It is the +streaming counterpart to `examples/durable_execution`: the same engine owns +records, the Redis runnable queue, fenced workers, checkpoints, and retention, +while the SSE stream carries display chunks alongside it. + +The two halves are deliberately separate: + +- **Durable state** is fenced and checkpointed. A `PipelineSnapshot` is + persisted once the pages are fetched and converted, so a restart resumes into + generation instead of hitting the network again. +- **Chunks are display data.** They live in a bounded append-only log outside + the fence, so a token cannot contend with the lease heartbeat, and a dropped + token can never fail or replay the execution. + +The streaming callback is bound to the `llm` component in `setup()` rather than +passed per run the way `async_streaming_generator` passes it. Per-run injection +cannot work under checkpointing: run data is serialized into the +`PipelineSnapshot`, Haystack drops the callable it cannot serialize, and +`Pipeline.run` rebuilds its `data` from the snapshot when resuming, so the +callback disappears at the first checkpoint. + +Sharing one bound callback across concurrent executions is safe, and for exactly +the reason `async_streaming_generator` is: its `_async_streaming_callback` is +also a single module-level function handed to every concurrent run, and the +per-run destination comes from a `ContextVar` resolved on each call. Hayhooks +routes on `_ASYNC_STREAMING_QUEUE`; this routes on the durable execution +context. Per-run injection buys that helper control over *which* components +stream, not isolation between runs. A run-time `streaming_callback` still takes +precedence, so ordinary streaming endpoints on the same wrapper behave normally. + +The included `httpx` client submits a question, streams the answer, then +deliberately drops the connection mid-answer and reattaches with +`Last-Event-ID` to show that nothing in between is lost. + +Run each command from the repository root. This is a local demonstration, not a +production Redis configuration. + +1. Start Redis and install the example dependencies. + +```bash +docker compose -f examples/durable-compose.yaml up -d && python -m pip install -e ".[durable]" httpx rich +``` + +2. Point Hayhooks at Redis and supply an OpenAI key. The five-second lease keeps + the restart demonstration in step 5 short; the 30-second default would leave a + killed execution waiting that long before another worker may reclaim it. + +```bash +export HAYHOOKS_DURABLE_REDIS_URL=redis://localhost:6379/0 OPENAI_API_KEY=sk-... HAYHOOKS_DURABLE_LEASE_DURATION_MS=5000 +``` + +3. In a first terminal, start Hayhooks. + +```bash +hayhooks run --pipelines-dir examples/durable_chat_with_website/pipelines +``` + +4. In a second terminal, run the client. + +```bash +python examples/durable_chat_with_website/demo.py +``` + +The answer prints token by token. After eight tokens the client drops the +connection on purpose, prints the `Last-Event-ID` it reached, reattaches, and +finishes the answer without a gap. The stream ends with a `completed` event +carrying the same projection the inspect route returns. + +5. To watch the durable half, restart Hayhooks while a request is in flight. + +The client reports the broken stream and keeps reattaching. Once the lease +expires, a worker in the restarted process reclaims the execution, and the +progress log records what the second attempt actually did: + +``` +fetch Fetching 2 page(s) +checkpoint Checkpoint saved before pipeline component 'prompt' +resume Resuming from the fetch checkpoint +completed Answer complete +``` + +The checkpoint lands about a second after submission, so kill the server after +that to see the `resume` line. Kill it sooner and the second attempt honestly +reports a second `fetch`, because there was no snapshot to resume from yet. +The retried attempt replays from its checkpoint, so the client prints an +attempt marker when tokens repeat. A server-side `error` event is treated the +same way as a dropped connection: reattach from `Last-Event-ID`. + +Expect recovery to take roughly the lease duration plus one poll interval, so +about six seconds with the settings above. The same bounded retry also covers a +transient generator failure, which is the everyday reason that checkpoint earns +its keep. + +Streaming is bounded and can be switched off without touching the wrapper: + +```bash +# 10 000 chunks per execution by default; 0 disables the log and leaves the endpoint working +export HAYHOOKS_DURABLE_MAX_STREAM_CHUNKS=0 +``` + +Stop Hayhooks with Ctrl-C, then stop Redis. Add `-v` to remove the retained +volume before a clean rehearsal: + +```bash +docker compose -f examples/durable-compose.yaml down +# docker compose -f examples/durable-compose.yaml down -v +``` diff --git a/examples/durable_chat_with_website/demo.py b/examples/durable_chat_with_website/demo.py new file mode 100644 index 00000000..1a60ef13 --- /dev/null +++ b/examples/durable_chat_with_website/demo.py @@ -0,0 +1,119 @@ +"""Submit a durable chat-with-website question and follow its token stream.""" + +from __future__ import annotations + +import argparse +import json +import time +import uuid +from collections.abc import Iterator +from typing import Any +from urllib.parse import urljoin + +import httpx +from rich.console import Console +from rich.json import JSON +from rich.panel import Panel + +_PAUSE_SECONDS = 2 +# Detaching mid-answer is the point of the demo: the execution keeps running and +# the chunk log keeps every token until this client reattaches. +_DETACH_AFTER_CHUNKS = 8 + + +def sse_events(client: httpx.Client, url: str, cursor: str | None) -> Iterator[dict[str, str]]: + """Yield one dict of SSE fields per event, resuming from *cursor* when given.""" + headers = {"Last-Event-ID": cursor} if cursor else {} + with client.stream("GET", url, headers=headers, timeout=None) as response: + response.raise_for_status() + fields: dict[str, str] = {} + for line in response.iter_lines(): + if line.startswith(":"): + continue # a heartbeat comment, sent while the execution is quiet + if line: + name, _, value = line.partition(": ") + fields[name] = value + elif fields: + yield fields + fields = {} + + +class StreamDroppedError(Exception): + """The server ended the stream with an error event instead of a terminal one.""" + + +def follow( + console: Console, client: httpx.Client, url: str, cursor: str | None, *, detach_after: int | None +) -> tuple[str | None, dict[str, Any] | None]: + """Print chunks from one connection; return the cursor and the terminal body.""" + attempt: int | None = None + for printed, event in enumerate(sse_events(client, url, cursor), start=1): + cursor = event.get("id", cursor) + if event["event"] == "error": + raise StreamDroppedError(json.loads(event["data"]).get("detail", "execution stream interrupted")) + if event["event"] != "chunk": + console.print() + return cursor, json.loads(event["data"]) + chunk = json.loads(event["data"]) + if attempt is None: + attempt = chunk["attempt"] + elif chunk["attempt"] != attempt: + # A retried attempt replays from its checkpoint, so tokens printed + # before the crash arrive again. Printing cannot retract them; a + # client with a rewritable buffer would reset it here instead. + attempt = chunk["attempt"] + console.print() + console.print(f"[dim]Attempt {attempt}: the execution resumed from its checkpoint, so tokens repeat.[/dim]") + console.print(chunk["payload"]["content"] or "", end="") + if printed == detach_after: + return cursor, None + return cursor, None + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base-url", default="http://localhost:1416") + parser.add_argument("--question", default="What does Haystack do, and how does Redis fit in?") + args = parser.parse_args() + base_url = args.base_url.rstrip("/") + console = Console() + + with httpx.Client(timeout=10) as client: + submitted = client.post( + f"{base_url}/chat_with_website/run-durable", + headers={"Idempotency-Key": f"ask-{uuid.uuid4().hex[:12]}"}, + json={"question": args.question}, + ) + console.print(Panel(JSON.from_data(submitted.json()), title=f"{submitted.status_code} submitted")) + if submitted.is_error: + return 1 + stream_url = urljoin(f"{base_url}/", submitted.json()["links"]["stream"]) + + console.print(Panel.fit(f"[bold cyan]GET[/] {stream_url}", title="Streaming", border_style="cyan")) + cursor, terminal, detached = None, None, False + while terminal is None: + try: + cursor, terminal = follow( + console, client, stream_url, cursor, detach_after=None if detached else _DETACH_AFTER_CHUNKS + ) + except (httpx.HTTPError, StreamDroppedError) as error: + console.print() + console.print(Panel(str(error), title="Stream interrupted", border_style="red")) + if terminal is None: + detached = True + console.print() + console.print( + Panel( + f"Reattaching from Last-Event-ID {cursor}. Nothing between here and there is lost.", + title="Detached", + border_style="yellow", + ) + ) + time.sleep(_PAUSE_SECONDS) + + console.print(Panel(JSON.from_data(terminal), title=f"Terminal event: {terminal['status']}")) + return 0 if terminal["status"] == "completed" else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/examples/durable_chat_with_website/pipelines/chat_with_website/chat_with_website.yml b/examples/durable_chat_with_website/pipelines/chat_with_website/chat_with_website.yml new file mode 100644 index 00000000..86c58ba6 --- /dev/null +++ b/examples/durable_chat_with_website/pipelines/chat_with_website/chat_with_website.yml @@ -0,0 +1,48 @@ +components: + converter: + type: haystack.components.converters.html.HTMLToDocument + init_parameters: + extraction_kwargs: null + + fetcher: + init_parameters: + raise_on_failure: true + retry_attempts: 2 + timeout: 3 + user_agents: + - haystack/LinkContentFetcher/2.0.0b8 + type: haystack.components.fetchers.link_content.LinkContentFetcher + + llm: + init_parameters: + api_key: + env_vars: + - OPENAI_API_KEY + strict: true + type: env_var + generation_kwargs: {} + model: gpt-4o-mini + type: haystack.components.generators.chat.openai.OpenAIChatGenerator + + prompt: + init_parameters: + template: | + {% message role="user" %} + According to the contents of this website: + {% for document in documents %} + {{document.content}} + {% endfor %} + Answer the given question: {{query}} + {% endmessage %} + required_variables: "*" + type: haystack.components.builders.chat_prompt_builder.ChatPromptBuilder + +connections: + - receiver: converter.sources + sender: fetcher.streams + - receiver: prompt.documents + sender: converter.documents + - receiver: llm.messages + sender: prompt.prompt + +metadata: {} \ No newline at end of file diff --git a/examples/durable_chat_with_website/pipelines/chat_with_website/pipeline_wrapper.py b/examples/durable_chat_with_website/pipelines/chat_with_website/pipeline_wrapper.py new file mode 100644 index 00000000..3a90655a --- /dev/null +++ b/examples/durable_chat_with_website/pipelines/chat_with_website/pipeline_wrapper.py @@ -0,0 +1,83 @@ +"""A durable chat-with-website Pipeline that streams its answer while it runs.""" + +from pathlib import Path + +from haystack import Pipeline +from haystack.core.errors import PipelineRuntimeError +from haystack.dataclasses import StreamingChunk +from pydantic import BaseModel, Field + +from hayhooks import BasePipelineWrapper, DurableContext, current_durable_context + +DEFAULT_URLS = ["https://haystack.deepset.ai", "https://www.redis.io"] + + +class ChatRequest(BaseModel): + """A question to answer from the live contents of a few web pages.""" + + question: str = Field(min_length=1, max_length=2_000) + urls: list[str] = Field(default=DEFAULT_URLS, min_length=1, max_length=5) + + +class ChatAnswer(BaseModel): + """The finished answer, already delivered token by token over the stream.""" + + reply: str + urls: list[str] + + +def stream_to_execution(chunk: StreamingChunk) -> None: + """ + Forward one generated token to the SSE stream of whichever execution is running. + + Resolving the execution per call, rather than closing over one, is what lets a + single shared Pipeline serve concurrent durable executions without crossing + their streams. + """ + if context := current_durable_context(): + context.stream_chunk_sync(chunk) + + +class PipelineWrapper(BasePipelineWrapper): + """Fetch pages durably, then stream the generated answer to the SSE endpoint.""" + + durable_revision = "durable-chat-with-website-v1" + + def setup(self) -> None: + self.pipeline = Pipeline.loads((Path(__file__).parent / "chat_with_website.yml").read_text()) + # `async_streaming_generator` passes its callback per run in `pipeline_run_args`. + # That cannot work under checkpointing: run data is serialized into the + # PipelineSnapshot, Haystack drops the callable it cannot serialize, and + # `Pipeline.run` rebuilds `data` from the snapshot on resume, so the callback + # is gone from the first checkpoint on. Binding it to the component survives. + # + # Sharing one bound callback across concurrent executions is safe for the same + # reason `_async_streaming_callback` is: it is a module-level function that + # resolves its destination per call from a ContextVar. Hayhooks routes on + # `_ASYNC_STREAMING_QUEUE`; this routes on the durable execution context. A + # run-time `streaming_callback` still wins, so ordinary streaming endpoints on + # this wrapper are unaffected. + self.pipeline.get_component("llm").streaming_callback = stream_to_execution + + async def run_durable_async(self, context: DurableContext, request: ChatRequest) -> ChatAnswer: + # This body re-runs from the top on every attempt, so the message has to say + # what the attempt will actually do rather than what the first one did. + resumed = context.record.checkpoint is not None + await context.report_progress( + "Resuming from the fetch checkpoint" if resumed else f"Fetching {len(request.urls)} page(s)", + kind="resume" if resumed else "fetch", + ) + try: + outputs = await context.run_pipeline_async( + {"fetcher": {"urls": request.urls}, "prompt": {"query": request.question}}, + # Fetching is the slow, flaky step. Checkpointing once it is done means a + # later attempt resumes into generation instead of hitting the network again. + checkpoint_at=["prompt"], + ) + except PipelineRuntimeError as error: + # Without this, a transient generator failure ends the execution and the + # checkpoint above never pays for itself. The retry is bounded by + # `durable_max_attempts`. + await context.retry(f"Pipeline attempt failed: {error}") + await context.report_progress("Answer complete", kind="completed") + return ChatAnswer(reply=outputs["llm"]["replies"][0].text, urls=request.urls) diff --git a/src/hayhooks/durable/backend.py b/src/hayhooks/durable/backend.py index f9937276..d3115d0e 100644 --- a/src/hayhooks/durable/backend.py +++ b/src/hayhooks/durable/backend.py @@ -4,6 +4,7 @@ from __future__ import annotations import json +import re from collections.abc import Callable, Mapping from dataclasses import dataclass, replace from typing import Any, Protocol @@ -19,6 +20,9 @@ from hayhooks.durable.models import ExecutionAdmissionError, ExecutionStoreError MAINTENANCE_BATCH_SIZE = 100 +CHUNK_CURSOR_START = "0-0" +_CHUNK_CURSOR = re.compile(r"^\d{1,20}-\d{1,20}$") +_MAX_STREAM_ID_PART = 2**64 - 1 DEFAULT_TRANSACTION_MAX_RETRIES = 8 DEFAULT_TRANSACTION_BACKOFF_MAX_MS = 25 @@ -56,6 +60,9 @@ class ExecutionStoreConfig: max_wait_bytes: int = 64_000 max_progress_events: int = 100 max_progress_event_bytes: int = 8_192 + # Zero disables the append-only display chunk log; see append_chunk. + max_stream_chunks: int = 10_000 + max_stream_chunk_bytes: int = 64_000 def __post_init__(self) -> None: for name in ( @@ -68,10 +75,19 @@ def __post_init__(self) -> None: "max_wait_bytes", "max_progress_events", "max_progress_event_bytes", + "max_stream_chunk_bytes", ): if getattr(self, name) < 1: raise ValueError(f"{name} must be positive") - if min(self.transaction_backoff_max_ms, self.lease_commit_safety_ms, self.max_nonterminal_executions) < 0: + if ( + min( + self.transaction_backoff_max_ms, + self.lease_commit_safety_ms, + self.max_nonterminal_executions, + self.max_stream_chunks, + ) + < 0 + ): raise ValueError("durable limits cannot be negative") @@ -93,6 +109,10 @@ async def read_payloads(self, run_id: str, kinds: tuple[PayloadKind, ...]) -> di async def read_progress(self, run_id: str) -> list[bytes]: ... + async def append_chunk(self, run_id: str, attempt: int, chunk: bytes) -> None: ... + + async def read_chunks(self, run_id: str, after: str, *, block_ms: int) -> list[tuple[str, int, bytes]]: ... + async def transition( self, run_id: str, command: ExecutionCommand, *, candidate: bool = False ) -> TransitionPlan: ... @@ -104,6 +124,19 @@ async def maintain(self, command_factory: Callable[[int, int], ExecutionCommand] async def operational_counts(self) -> dict[str, int]: ... +def parse_chunk_cursor(value: str) -> tuple[int, int]: + """Decode a chunk-log entry ID; the cursor arrives from an untrusted SSE client.""" + if not _CHUNK_CURSOR.match(value): + raise ValueError("chunk cursor must be a '