From 19b99b4ba999388955b6d6adf5a479b3fb30bc44 Mon Sep 17 00:00:00 2001 From: Brian O'Kelley Date: Fri, 25 Sep 2026 11:04:54 +0000 Subject: [PATCH 1/2] fix(reporting): settle service lifecycle before resource cleanup --- src/adcp/reporting/__init__.py | 28 + src/adcp/reporting/_settlement.py | 47 ++ src/adcp/reporting/inline_source.py | 15 +- src/adcp/reporting/ledger/producer.py | 12 +- src/adcp/reporting/service.py | 194 +++-- src/adcp/reporting/service_lifecycle.py | 346 ++++++++ ...st_reliable_reporting_service_lifecycle.py | 214 +++++ tests/test_reliable_reporting_lifecycle.py | 752 ++++++++++++++++++ tests/test_reliable_reporting_service.py | 22 +- .../reliable_reporting_lifecycle.py | 55 ++ 10 files changed, 1605 insertions(+), 80 deletions(-) create mode 100644 src/adcp/reporting/_settlement.py create mode 100644 src/adcp/reporting/service_lifecycle.py create mode 100644 tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py create mode 100644 tests/test_reliable_reporting_lifecycle.py create mode 100644 tests/type_checks/reliable_reporting_lifecycle.py diff --git a/src/adcp/reporting/__init__.py b/src/adcp/reporting/__init__.py index 0e642143e..ec249e22c 100644 --- a/src/adcp/reporting/__init__.py +++ b/src/adcp/reporting/__init__.py @@ -92,6 +92,21 @@ from adcp.reporting.service import ( ReportingAdapter as ReportingAdapter, ) + from adcp.reporting.service_lifecycle import ( + ReliableReportingServiceError as ReliableReportingServiceError, + ) + from adcp.reporting.service_lifecycle import ( + ReliableReportingShutdownTimeoutError as ReliableReportingShutdownTimeoutError, + ) + from adcp.reporting.service_lifecycle import ( + ReliableReportingState as ReliableReportingState, + ) + from adcp.reporting.service_lifecycle import ( + ReliableReportingUnavailableError as ReliableReportingUnavailableError, + ) + from adcp.reporting.service_lifecycle import ( + ReportingServiceResource as ReportingServiceResource, + ) from adcp.reporting.testing import ( DeterministicReportingClock as DeterministicReportingClock, ) @@ -121,6 +136,14 @@ _LAZY_EXPORTS = { "ReliableReportingConfigurationError": ("service", "ReliableReportingConfigurationError"), "ReliableReportingService": ("service", "ReliableReportingService"), + "ReliableReportingServiceError": ("service_lifecycle", "ReliableReportingServiceError"), + "ReliableReportingShutdownTimeoutError": ( + "service_lifecycle", + "ReliableReportingShutdownTimeoutError", + ), + "ReliableReportingState": ("service_lifecycle", "ReliableReportingState"), + "ReliableReportingUnavailableError": ("service_lifecycle", "ReliableReportingUnavailableError"), + "ReportingServiceResource": ("service_lifecycle", "ReportingServiceResource"), "ReportingAccountContext": ("service", "ReportingAccountContext"), "ReportingAdapter": ("service", "ReportingAdapter"), "DeterministicReportingClock": ("testing", "DeterministicReportingClock"), @@ -161,6 +184,11 @@ def __dir__() -> list[str]: "ReportingTier", "ReliableReportingConfigurationError", "ReliableReportingService", + "ReliableReportingServiceError", + "ReliableReportingShutdownTimeoutError", + "ReliableReportingState", + "ReliableReportingUnavailableError", + "ReportingServiceResource", "ReportingAccountContext", "ReportingAdapter", "DeterministicReportingClock", diff --git a/src/adcp/reporting/_settlement.py b/src/adcp/reporting/_settlement.py new file mode 100644 index 000000000..5d50ddc45 --- /dev/null +++ b/src/adcp/reporting/_settlement.py @@ -0,0 +1,47 @@ +"""Settlement of work Python cannot safely interrupt.""" + +from __future__ import annotations + +import asyncio +from typing import Any, TypeVar + +_T = TypeVar("_T") + + +async def settle_task(task: asyncio.Task[_T]) -> _T: + """Defer caller cancellation until an owned task has actually finished. + + Shield alone only keeps the child alive: the parent still returns early on + cancellation. Retain the child and consume its outcome, including under + repeated cancellation, before propagating the caller's cancellation. This + barrier deliberately has no timeout; a supervisor can time out its *wait* + while retaining ownership of the still-running operation. + """ + cancelled: asyncio.CancelledError | None = None + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError as error: + cancelled = error + except Exception: + break # result() below retrieves the child's exception exactly once + if cancelled is not None: + if not task.cancelled(): + task.exception() + raise cancelled + return task.result() + + +async def cancel_and_settle(task: asyncio.Task[Any]) -> None: + """Request cancellation, then join without losing repeated caller cancels. + + Joining through a task that consumes the child's cancellation lets us + distinguish its expected CancelledError from a new cancellation of the + caller. The latter is propagated only after the child has settled. + """ + task.cancel() + + async def join() -> None: + await asyncio.gather(task, return_exceptions=True) + + await settle_task(asyncio.create_task(join())) diff --git a/src/adcp/reporting/inline_source.py b/src/adcp/reporting/inline_source.py index 71bb3bad0..bc2aca2a0 100644 --- a/src/adcp/reporting/inline_source.py +++ b/src/adcp/reporting/inline_source.py @@ -73,6 +73,7 @@ from pydantic import TypeAdapter, ValidationError +from adcp.reporting._settlement import settle_task from adcp.reporting.currency import ( ReportingCurrencyError, validate_currency, @@ -470,7 +471,7 @@ async def stage( digest = hashlib.sha256(payload).hexdigest() object_ref = f"{source_execution_key}.{ordinal}" target = self._path(account_id, object_ref, digest) - await asyncio.to_thread(self._write, target, payload) + await settle_task(asyncio.create_task(asyncio.to_thread(self._write, target, payload))) return object_ref, digest @staticmethod @@ -499,7 +500,9 @@ async def read( cancel: asyncio.Event, ) -> bytes: target = self._path(account_id, object_ref, object_generation) - payload: bytes = await asyncio.to_thread(target.read_bytes) + payload: bytes = await settle_task( + asyncio.create_task(asyncio.to_thread(target.read_bytes)) + ) if hashlib.sha256(payload).hexdigest() != object_generation: raise OSError("staged object bytes no longer match their pinned generation") return payload @@ -724,7 +727,13 @@ async def _invoke( if _is_async_callable(self._fetch): answer = await self._fetch(request) else: - answer = await asyncio.shield(asyncio.to_thread(self._fetch, request)) + # Keep both the thread and a dynamically returned awaitable owned + # until they settle. Shield alone abandons the wait on cancellation. + async def run_sync() -> Any: + result = await asyncio.to_thread(self._fetch, request) + return await result if isinstance(result, Awaitable) else result + + answer = await settle_task(asyncio.create_task(run_sync())) if isinstance(answer, Awaitable): # A plain callable that hands back a coroutine; await it here # rather than letting it reach the coercion step unfinished. diff --git a/src/adcp/reporting/ledger/producer.py b/src/adcp/reporting/ledger/producer.py index fc54ea2b3..ce552e7f4 100644 --- a/src/adcp/reporting/ledger/producer.py +++ b/src/adcp/reporting/ledger/producer.py @@ -37,6 +37,7 @@ from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, TypeAlias +from adcp.reporting._settlement import cancel_and_settle from adcp.reporting.canonical_json import canonical_json_utf8_v1 from adcp.reporting.currency import ( ReportingCurrencyError, @@ -832,13 +833,22 @@ async def acquire_obligation( constituents=constituents, ) cancel = asyncio.Event() + execution = asyncio.create_task(self._source.execute(request, cancel=cancel)) try: result = await asyncio.wait_for( - self._source.execute(request, cancel=cancel), + asyncio.shield(execution), timeout=self._offerings.slice_timeout.total_seconds(), ) + except asyncio.CancelledError: + cancel.set() + await cancel_and_settle(execution) + raise except asyncio.TimeoutError: cancel.set() + # wait_for's own cancellation join can be interrupted by a second + # cancellation of this producer (notably on Python 3.10). Retain + # the execution explicitly until even synchronous work has settled. + await cancel_and_settle(execution) turn.slices_failed.append(obligation.reporting_obligation_id) self._note_escalation(obligation, turn, now=now) return None diff --git a/src/adcp/reporting/service.py b/src/adcp/reporting/service.py index 6d97c4e5f..e7226d596 100644 --- a/src/adcp/reporting/service.py +++ b/src/adcp/reporting/service.py @@ -44,6 +44,15 @@ WorkerTurn, ) from adcp.reporting.ledger.store import LedgerConflictError, decode_cursor +from adcp.reporting.service_lifecycle import ( + FailureComponent, + ReliableReportingServiceError, + ReliableReportingShutdownTimeoutError, + ReliableReportingState, + ReliableReportingUnavailableError, + ReportingServiceResource, + _ServiceLifecycle, +) from adcp.reporting.source import ( AuthoritativeOfferingV1, ProvisionalSnapshotOfferingV1, @@ -56,7 +65,11 @@ "AdapterRegistration", "ReliableReportingConfigurationError", "ReliableReportingService", + "ReliableReportingServiceError", + "ReliableReportingShutdownTimeoutError", + "ReliableReportingState", "ReliableReportingTurn", + "ReliableReportingUnavailableError", "ReportingAccountContext", "ReportingAdapter", "ReportingAdapterRegistry", @@ -64,6 +77,7 @@ "ReportingCallerResolver", "ReportingContextResolver", "ReportingReceiptHandler", + "ReportingServiceResource", "ReportingWorkerErrorHandler", ] @@ -320,7 +334,13 @@ def did_work(self) -> bool: class ReliableReportingService: - """High-level owner of adapters, ledger handlers, workers, and lifecycle.""" + """High-level owner of adapters, ledger handlers, workers, and lifecycle. + + Admitted calls own a task so transport cancellation cannot interrupt their + cleanup. Invoke them outside caller-owned ledger transactions; use the + low-level store API when composing an ambient transaction batch. Injected + resources are borrowed unless explicitly transferred via ``owned_resources``. + """ def __init__( self, @@ -339,6 +359,7 @@ def __init__( receipt_handler: ReportingReceiptHandler | None = None, reconciled_billing: bool = False, worker_error_handler: ReportingWorkerErrorHandler | None = None, + owned_resources: Sequence[ReportingServiceResource] = (), ) -> None: effective_clock = clock or (lambda: datetime.now(timezone.utc)) self.store = store @@ -361,9 +382,13 @@ def __init__( ReportingConfigurationGenerationKey, ReportingConfiguration ] = {} self._initialized = False - self._closed = False - self._worker_task: asyncio.Task[None] | None = None + self._configuration_lock = asyncio.Lock() self._turn_lock = asyncio.Lock() + self._lifecycle = _ServiceLifecycle( + self._initialize, + resources=owned_resources, + configuration_error=ReliableReportingConfigurationError, + ) @classmethod def memory( @@ -406,8 +431,14 @@ def postgres( async def configure(self, configuration: ReportingConfiguration) -> None: """Resolve trusted account facts once and freeze this generation's route.""" - if self._closed: - raise RuntimeError("the reporting service is closed") + + async def configure() -> None: + async with self._configuration_lock: + await self._configure(configuration) + + await self._lifecycle.call(configure, before_start=True) + + async def _configure(self, configuration: ReportingConfiguration) -> None: key = configuration.generation_key existing = self._bindings.get(key) if existing is not None: @@ -546,42 +577,52 @@ def validate(self) -> None: ) async def initialize(self) -> None: - if self._closed: - raise RuntimeError("the reporting service is closed") - if self._initialized: - return + """Prepare resources once; start/run_worker opens reporting admission.""" + await self._lifecycle.initialize() + + async def _initialize(self) -> None: self.validate() - await self.store.create_schema() - for configuration in self._pending_configurations.values(): - await self.store.put_configuration(configuration) - self._pending_configurations.clear() - self.sources.freeze() - self._initialized = True + async with self._configuration_lock: + await self.store.create_schema() + for configuration in self._pending_configurations.values(): + await self.store.put_configuration(configuration) + self._pending_configurations.clear() + self.sources.freeze() + self._initialized = True async def start(self) -> None: await self.initialize() - if self._worker_interval is not None and self._worker_task is None: - self._worker_task = asyncio.create_task( - self._worker_loop(), name="adcp-reliable-reporting" - ) + await self._lifecycle.activate( + self._worker_loop if self._worker_interval is not None else None + ) - async def close(self) -> None: - self._closed = True - if self._worker_task is not None: - self._worker_task.cancel() - try: - await self._worker_task - except asyncio.CancelledError: - pass - self._worker_task = None - for component in ( - self._notification_worker, - self._materialization_worker, - self._receipt_handler, - ): - close = getattr(component, "close", None) - if close is not None: - await _resolve(close()) + @property + def state(self) -> ReliableReportingState: + """Current process lifecycle; STOPPING retains all unsettled ownership.""" + return self._lifecycle.state + + @property + def ready(self) -> bool: + """Whether this process admits work (not a durable-tier graph proof).""" + return self._lifecycle.ready + + @property + def failure(self) -> ReliableReportingServiceError | None: + """Sanitized first unexpected failure, if any.""" + return self._lifecycle.failure + + async def close(self, *, timeout: float | None = None) -> None: + """Reject new work, drain admitted operations, then close owned resources. + + A timeout or cancelled waiter leaves the shared shutdown running and the + service STOPPING until settlement. Injected components are borrowed; + transfer lifetime explicitly with ``owned_resources`` to close them here. + """ + await self._lifecycle.close(timeout=timeout) + + async def wait(self) -> None: + """Wait for shutdown and raise a typed failure if supervision failed.""" + await self._lifecycle.wait() async def __aenter__(self) -> ReliableReportingService: await self.start() @@ -592,33 +633,31 @@ async def __aexit__(self, *_exc: object) -> None: async def _worker_loop(self) -> None: assert self._worker_interval is not None - while True: - try: - await self.run_worker() - except Exception as error: - # A composition bug outside the isolated producer/extension - # turns must be visible, but must not silently kill the - # lifecycle-owned scheduler forever. - await self._report_worker_error("service", error) - await asyncio.sleep(self._worker_interval.total_seconds()) - - async def _report_worker_error(self, component: str, error: BaseException) -> None: - logger.error( - "Reliable Reporting worker component %s failed", - component, - exc_info=(type(error), error, error.__traceback__), - ) + while not self._lifecycle.stopping: + await self.run_worker() + await self._lifecycle.wait_for_stop(self._worker_interval.total_seconds()) + + async def _report_worker_error(self, component: FailureComponent) -> None: + error = await self._lifecycle.fail(component) + logger.error("Reliable Reporting worker stopped: %s", error) if self._worker_error_handler is None: return try: await _resolve(self._worker_error_handler(component, error)) except Exception: - logger.exception("Reliable Reporting worker error handler failed") + logger.error("Reliable Reporting worker error handler failed") async def run_worker(self, *, now: datetime | None = None) -> ReliableReportingTurn: """Route one turn across every frozen configuration generation.""" await self.initialize() + await self._lifecycle.activate() + return await self._lifecycle.call(lambda: self._run_worker(now=now)) + + async def _run_worker(self, *, now: datetime | None) -> ReliableReportingTurn: async with self._turn_lock: + # A turn admitted before STOPPING may have waited behind another + # turn. It must not start fresh work after that turn has drained. + self._lifecycle.require_ready() turn = ReliableReportingTurn() for key, binding in sorted( self._bindings.items(), @@ -628,30 +667,32 @@ async def run_worker(self, *, now: datetime | None = None) -> ReliableReportingT item[0].delivery_config_version, ), ): + if self._lifecycle.stopping: + break try: turn.configurations[key] = await binding.producer.run_configuration( binding.configuration, now=now ) - except Exception as error: - turn.configuration_errors[key] = error - await self._report_worker_error( - "configuration:" - f"{key.account_id}:{key.delivery_config_id}@" - f"{key.delivery_config_version}", - error, - ) + except Exception: + await self._report_worker_error("configuration") + assert self.failure is not None + raise self.failure from None extensions = ( ("materialization", self._materialization_worker), ("notification", self._notification_worker), ) for name, extension in extensions: + if self._lifecycle.stopping: + break if extension is not None: try: result = await extension.run_once() - except Exception as error: - turn.extension_errors[name] = error - await self._report_worker_error(name, error) - continue + except Exception: + await self._report_worker_error( + "materialization" if name == "materialization" else "notification" + ) + assert self.failure is not None + raise self.failure from None if result is not None: turn.extension_results.append(result) return turn @@ -677,6 +718,9 @@ def _default_caller(request: Any, context: Any | None) -> ReportingStatusCaller: async def get_reporting_status( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + return await self._lifecycle.call(lambda: self._get_reporting_status(request, context)) + + async def _get_reporting_status(self, request: Any, context: Any | None) -> dict[str, Any]: caller = await self.caller_for(request, context) return await ReportingStatusHandler( self.store, @@ -687,6 +731,9 @@ async def get_reporting_status( async def sync_reporting_status( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + return await self._lifecycle.call(lambda: self._sync_reporting_status(request, context)) + + async def _sync_reporting_status(self, request: Any, context: Any | None) -> dict[str, Any]: if not self._consumer_status_enabled: raise ReliableReportingConfigurationError("consumer status ingest is disabled") caller = await self.caller_for(request, context) @@ -699,6 +746,9 @@ async def sync_reporting_status( async def sync_reporting_receipts( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + return await self._lifecycle.call(lambda: self._sync_reporting_receipts(request, context)) + + async def _sync_reporting_receipts(self, request: Any, context: Any | None) -> dict[str, Any]: if self._receipt_handler is None: raise ReliableReportingConfigurationError("reporting receipt handling is disabled") caller = await self.caller_for(request, context) @@ -709,6 +759,9 @@ async def sync_reporting_receipts( async def get_revision_content( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + return await self._lifecycle.call(lambda: self._get_revision_content(request, context)) + + async def _get_revision_content(self, request: Any, context: Any | None) -> dict[str, Any]: payload = _wire(request) revision_id = payload.get("reporting_revision_id") if not revision_id: @@ -765,10 +818,8 @@ async def get_revision_content( def capability_block(self) -> dict[str, Any]: """Project only components and offerings this service actually installed.""" - if not self._bindings: - raise ReliableReportingConfigurationError( - "at least one configured reporting generation is required for capabilities" - ) + if not self.ready or not self._bindings: + return {} offerings: dict[str, dict[str, Any]] = {} for binding in self._bindings.values(): offering = _thaw(binding.context.capability_offering) @@ -800,8 +851,11 @@ def capability_block(self) -> dict[str, Any]: def inject_capabilities(self, response: Any) -> dict[str, Any]: """Merge the truthful reporting block into a base capability response.""" payload = _wire(response) + block = self.capability_block() + if not block: + return payload media_buy = dict(payload.get("media_buy") or {}) - media_buy["reporting_delivery"] = self.capability_block() + media_buy["reporting_delivery"] = block payload["media_buy"] = media_buy protocols = [str(item) for item in payload.get("supported_protocols") or []] if "media_buy" not in protocols: diff --git a/src/adcp/reporting/service_lifecycle.py b/src/adcp/reporting/service_lifecycle.py new file mode 100644 index 000000000..83613db76 --- /dev/null +++ b/src/adcp/reporting/service_lifecycle.py @@ -0,0 +1,346 @@ +"""Owned lifecycle and admission for the adapter-first reporting service. + +This module owns process lifetime only. Account activation, durable leases and +production-tier dependency proofs remain the responsibility of their reporting +components. No PostgreSQL driver is imported by this lifecycle surface. +""" + +from __future__ import annotations + +import asyncio +import inspect +import math +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from contextlib import asynccontextmanager +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Literal, TypeVar + +from adcp.reporting._settlement import cancel_and_settle, settle_task + +_T = TypeVar("_T") + +__all__ = [ + "ReliableReportingServiceError", + "ReliableReportingShutdownTimeoutError", + "ReliableReportingState", + "ReliableReportingUnavailableError", + "ReportingServiceResource", +] + + +class ReliableReportingState(str, Enum): + NEW = "new" + STARTING = "starting" + INITIALIZED = "initialized" + READY = "ready" + STOPPING = "stopping" + CLOSED = "closed" + FAILED = "failed" + + +class ReliableReportingUnavailableError(RuntimeError): + """The service is not admitting new reporting work.""" + + def __init__(self, state: ReliableReportingState) -> None: + self.state = state + super().__init__(f"reporting service is unavailable ({state.value})") + + +FailureComponent = Literal[ + "startup", "configuration", "materialization", "notification", "service", "shutdown" +] + + +class ReliableReportingServiceError(RuntimeError): + """Closed diagnostic for an unexpected failure; contains no provider body.""" + + def __init__(self, component: FailureComponent) -> None: + if component not in ( + "startup", + "configuration", + "materialization", + "notification", + "service", + "shutdown", + ): + raise ValueError("unknown reporting service failure component") + self.component = component + super().__init__(f"reporting service failed ({component})") + + +class ReliableReportingShutdownTimeoutError(TimeoutError): + """A close waiter timed out; the service still owns the unsettled work.""" + + +ResourceCallback = Callable[[], object | Awaitable[object]] + + +@dataclass(frozen=True) +class ReportingServiceResource: + """Explicit transfer of a resource's lifetime to one service. + + Resources open in declaration order before schema initialization and close + once, in reverse order, after all admitted work settles. Ownership transfers + at construction, so ``close`` must tolerate unopened or partially opened + resources (including a service that never starts). Injected stores, adapters, + workers, senders and pools are otherwise borrowed. + + Use async callbacks for loop-bound resources. Blocking synchronous callbacks + run in a thread and are settled before cleanup can advance. Do not transfer + the same resource to more than one owner. + """ + + close: ResourceCallback = field(repr=False) + open: ResourceCallback | None = field(default=None, repr=False) + + def __post_init__(self) -> None: + if not callable(self.close) or (self.open is not None and not callable(self.open)): + raise TypeError("resource callbacks must be callable") + + +async def _invoke(callback: ResourceCallback) -> None: + if inspect.iscoroutinefunction(callback) or inspect.iscoroutinefunction( + getattr(callback, "__call__", None) + ): + result = callback() + if inspect.isawaitable(result): + await result + else: + + async def run() -> None: + result = await asyncio.to_thread(callback) + if inspect.isawaitable(result): + await result + + await settle_task(asyncio.create_task(run())) + + +class _ServiceLifecycle: + """One loop-owned, locked state machine shared by tasks and mounted RPCs.""" + + def __init__( + self, + initialize: Callable[[], Awaitable[None]], + *, + resources: Sequence[ReportingServiceResource], + configuration_error: type[Exception], + ) -> None: + self.state = ReliableReportingState.NEW + self.failure: ReliableReportingServiceError | None = None + self._initialize = initialize + self._configuration_error = configuration_error + self._resources = tuple(resources) + if any(not isinstance(item, ReportingServiceResource) for item in self._resources): + raise TypeError("owned_resources must contain ReportingServiceResource values") + if any( + item.close == earlier.close + for index, item in enumerate(self._resources) + for earlier in self._resources[:index] + ): + raise ValueError("a resource close callback cannot be transferred twice") + self._lock = asyncio.Lock() + self._initialization: asyncio.Task[None] | None = None + self._worker: asyncio.Task[None] | None = None + self._shutdown: asyncio.Task[None] | None = None + self._admissions: dict[asyncio.Task[Any], int] = {} + self._idle = asyncio.Event() + self._idle.set() + self._stop = asyncio.Event() + self._stopped = asyncio.Event() + + @property + def ready(self) -> bool: + return self.state is ReliableReportingState.READY + + @property + def stopping(self) -> bool: + return self._stop.is_set() + + def require_ready(self) -> None: + if not self.ready: + raise ReliableReportingUnavailableError(self.state) + + async def initialize(self) -> None: + async with self._lock: + if self.state in (ReliableReportingState.INITIALIZED, ReliableReportingState.READY): + return + if self.state is ReliableReportingState.NEW: + self.state = ReliableReportingState.STARTING + self._initialization = asyncio.create_task( + self._open(), name="adcp-reporting-startup" + ) + elif self.state is not ReliableReportingState.STARTING: + raise ReliableReportingUnavailableError(self.state) + task = self._initialization + assert task is not None + try: + await asyncio.shield(task) + except asyncio.CancelledError: + # The opener may hold a transaction or an uninterruptible thread. + # Stop admission and retain it for the shared shutdown task. + await self.request_stop() + raise + except Exception: + await self.close() + raise + if self.state not in (ReliableReportingState.INITIALIZED, ReliableReportingState.READY): + raise ReliableReportingUnavailableError(self.state) + + async def _open(self) -> None: + try: + for resource in self._resources: + if self.stopping: + return + if resource.open is not None: + await _invoke(resource.open) + if self.stopping: + return + await self._initialize() + async with self._lock: + if self.state is ReliableReportingState.STARTING: + self.state = ReliableReportingState.INITIALIZED + except BaseException as error: + failure = await self.fail("startup") + if isinstance(error, self._configuration_error): + raise + raise failure from None + + async def activate(self, run: Callable[[], Awaitable[None]] | None = None) -> None: + async with self._lock: + if self.state not in (ReliableReportingState.INITIALIZED, ReliableReportingState.READY): + raise ReliableReportingUnavailableError(self.state) + if run is not None and self._worker is None: + self._worker = asyncio.create_task( + self._supervise(run), name="adcp-reliable-reporting" + ) + self.state = ReliableReportingState.READY + + async def _supervise(self, run: Callable[[], Awaitable[None]]) -> None: + try: + await run() + except BaseException: + if not self.stopping: + await self.fail("service") + else: + if not self.stopping: + await self.fail("service") + + async def call( + self, operation: Callable[[], Awaitable[_T]], *, before_start: bool = False + ) -> _T: + """Own an admitted task and propagate cancellation to it only once. + + Repeated cancellation of the transport/caller cannot interrupt the + operation's transaction cleanup. Context variables are inherited, but + callers must use low-level store APIs for ambient transaction batches: + high-level service operations own their tasks and transaction lifetime. + """ + + async def admitted() -> _T: + async with self.admit(before_start=before_start): + return await operation() + + task = asyncio.create_task(admitted(), name="adcp-reporting-admission") + try: + return await asyncio.shield(task) + except asyncio.CancelledError: + await cancel_and_settle(task) + raise + + @asynccontextmanager + async def admit(self, *, before_start: bool = False) -> AsyncIterator[None]: + task = asyncio.current_task() + assert task is not None + async with self._lock: + if not self.ready and not ( + before_start + and self.state + in ( + ReliableReportingState.NEW, + ReliableReportingState.STARTING, + ReliableReportingState.INITIALIZED, + ) + ): + raise ReliableReportingUnavailableError(self.state) + self._admissions[task] = self._admissions.get(task, 0) + 1 + self._idle.clear() + try: + yield + finally: + # No await: repeated cancellation cannot interrupt deregistration. + # All admission bookkeeping is confined to this service's loop. + remaining = self._admissions[task] - 1 + if remaining: + self._admissions[task] = remaining + else: + del self._admissions[task] + if not self._admissions: + self._idle.set() + + def _stop_locked(self) -> None: + if self._shutdown is None: + self.state = ReliableReportingState.STOPPING + self._stop.set() + self._shutdown = asyncio.create_task(self._drain(), name="adcp-reporting-shutdown") + + async def request_stop(self) -> None: + async with self._lock: + self._stop_locked() + + async def fail(self, component: FailureComponent) -> ReliableReportingServiceError: + async with self._lock: + if self.failure is None: + self.failure = ReliableReportingServiceError(component) + self._stop_locked() + return self.failure + + async def wait_for_stop(self, seconds: float) -> None: + try: + await asyncio.wait_for(self._stop.wait(), timeout=seconds) + except asyncio.TimeoutError: + pass + + async def close(self, *, timeout: float | None = None) -> None: + if timeout is not None and (not math.isfinite(timeout) or timeout < 0): + raise ValueError("close timeout must be a finite nonnegative number") + current = asyncio.current_task() + if ( + current in self._admissions + or current is self._initialization + or current is self._shutdown + ): + raise RuntimeError("close must be awaited outside an admitted reporting operation") + await self.request_stop() + assert self._shutdown is not None + try: + await asyncio.wait_for(asyncio.shield(self._shutdown), timeout=timeout) + except asyncio.TimeoutError: + raise ReliableReportingShutdownTimeoutError( + "reporting service is still stopping; admitted work has not settled" + ) from None + + async def _drain(self) -> None: + for task in (self._initialization, self._worker): + if task is not None: + await asyncio.gather(task, return_exceptions=True) + await self._idle.wait() + for resource in reversed(self._resources): + try: + await _invoke(resource.close) + except BaseException: + # A failed closer must not strand earlier resources. Preserve + # the first failure, without retaining a resource's raw error. + await self.fail("shutdown") + async with self._lock: + self.state = ( + ReliableReportingState.FAILED + if self.failure is not None + else ReliableReportingState.CLOSED + ) + self._stopped.set() + + async def wait(self) -> None: + await self._stopped.wait() + if self.failure is not None: + raise self.failure from None diff --git a/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py b/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py new file mode 100644 index 000000000..781b2a7f1 --- /dev/null +++ b/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py @@ -0,0 +1,214 @@ +"""Actual service admission/settlement with real PostgreSQL transactions.""" + +from __future__ import annotations + +import asyncio +import threading +from dataclasses import replace +from datetime import timedelta +from typing import Any + +import pytest + +from adcp.reporting.fixtures import redacted_capabilities +from adcp.reporting.ledger.pg import PgReportingLedgerStore +from adcp.reporting.service import ( + ReliableReportingService, + ReliableReportingShutdownTimeoutError, + ReliableReportingState, + ReliableReportingUnavailableError, + ReportingServiceResource, +) +from adcp.reporting.testing import ScriptedReportingAdapter +from tests.test_reliable_reporting_lifecycle import checkpoint +from tests.test_reliable_reporting_service import _account_context, _configuration, _rows + +from ._generation_support import isolated_reporting_pool + + +@pytest.mark.parametrize("autocommit", [False, True]) +@pytest.mark.parametrize("cancel_call", [False, True]) +async def test_stop_settles_public_configuration_transaction_before_owned_cleanup( + autocommit: bool, cancel_call: bool +) -> None: + async with isolated_reporting_pool(autocommit=autocommit) as pool: + inserted = asyncio.Event() + release = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + configurations_at_close: list[int] = [] + + class PausedLedger(PgReportingLedgerStore): + async def put_configuration(self, configuration: Any) -> None: + # Pause after the real ledger mutation but before its real + # transaction commits. No private prepared producer/store. + async with self.transaction(): + await super().put_configuration(configuration) + inserted.set() + try: + await release.wait() + finally: + if cancel_call: + cleanup_started.set() + await cleanup_release.wait() + + configuration = _configuration() + observer = PgReportingLedgerStore(pool=pool) + + async def close_owned() -> None: + retained = await observer.list_configurations(account_id=configuration.account_id) + configurations_at_close.append(len(retained)) + + service = ReliableReportingService( + store=PausedLedger(pool=pool), + account_context=_account_context, + owned_resources=(ReportingServiceResource(close=close_owned),), + ) + service.sources.register("gam", ScriptedReportingAdapter(redacted_capabilities(), [])) + await service.start() + configuring = asyncio.create_task(service.configure(configuration)) + await asyncio.wait_for(inserted.wait(), 5) + try: + assert await observer.list_configurations(account_id=configuration.account_id) == () + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert service.state is ReliableReportingState.STOPPING + assert configurations_at_close == [] + with pytest.raises(ReliableReportingUnavailableError): + await service.configure(replace(configuration, delivery_config_version=2)) + if cancel_call: + configuring.cancel() + await asyncio.wait_for(cleanup_started.wait(), 5) + configuring.cancel() + await checkpoint() + await checkpoint() + assert not configuring.done() + assert configurations_at_close == [] + cleanup_release.set() + with pytest.raises(asyncio.CancelledError): + await configuring + else: + release.set() + await configuring + finally: + release.set() + cleanup_release.set() + await asyncio.gather(configuring, return_exceptions=True) + await asyncio.wait_for(service.close(), 5) + assert configurations_at_close == [0 if cancel_call else 1] + assert pool.closed is False # an injected PostgreSQL pool is borrowed + restarted = ReliableReportingService.postgres(pool=pool, account_context=_account_context) + await restarted.start() + try: + retained = await restarted.store.list_configurations( + account_id=configuration.account_id + ) + assert retained == (() if cancel_call else (configuration,)) + finally: + await restarted.close() + + +async def test_cancelled_public_producer_settles_thread_and_recovers_same_durable_obligation() -> ( + None +): + async with isolated_reporting_pool() as pool: + async with pool.connection() as connection: + row = await (await connection.execute("SELECT clock_timestamp()")).fetchone() + now = row[0] + boundary = now.replace(minute=0, second=0, microsecond=0) + configuration = replace( + _configuration(), + activated_at=boundary - timedelta(hours=3) + timedelta(minutes=20), + deactivated_at=boundary - timedelta(hours=1), + ) + entered = asyncio.Event() + release = threading.Event() + finished = threading.Event() + closes: list[str] = [] + loop = asyncio.get_running_loop() + + class Adapter: + capabilities = redacted_capabilities() + + def fetch_slice(self, _request: Any) -> list[dict[str, Any]]: + loop.call_soon_threadsafe(entered.set) + try: + assert release.wait(10), "test cleanup watchdog" + return _rows(11) + finally: + finished.set() + + async def close(self) -> None: + assert finished.is_set() + closes.append("adapter") + + adapter = Adapter() + service = ReliableReportingService( + store=PgReportingLedgerStore(pool=pool), + account_context=_account_context, + # The semantic producer boundary is sampled from the database; + # PostgreSQL evidence itself uses the store's unmodified DB clock. + clock=lambda: now, + owned_resources=(ReportingServiceResource(close=adapter.close),), + ) + service.sources.register("gam", adapter) + await service.configure(configuration) + turn = asyncio.create_task(service.run_worker(now=now)) + await asyncio.wait_for(entered.wait(), 5) + try: + async with pool.connection() as connection: + rows = await ( + await connection.execute( + "SELECT reporting_obligation_id FROM reporting_obligations" + ) + ).fetchall() + assert len(rows) == 1 + obligation_id = rows[0][0] + turn.cancel() + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert service.state is ReliableReportingState.STOPPING + assert closes == [] + finally: + release.set() + await asyncio.gather(turn, return_exceptions=True) + await asyncio.wait_for(service.close(), 5) + assert closes == ["adapter"] + assert pool.closed is False + assert ( + await service.store.list_revisions( + account_id=configuration.account_id, reporting_obligation_id=obligation_id + ) + == () + ) + + # A fresh service replays the accepted generation through the public + # route and producer. This is restart recovery, not late discovery (the + # durable binding/discovery bridge remains a subsequent B1 slice). + restarted = ReliableReportingService( + store=PgReportingLedgerStore(pool=pool), + account_context=_account_context, + clock=lambda: now, + ) + restarted.sources.register( + "gam", ScriptedReportingAdapter(redacted_capabilities(), [_rows(11)]) + ) + await restarted.configure(configuration) + try: + recovered = await restarted.run_worker(now=now) + assert ( + len(recovered.configurations[configuration.generation_key].revisions_committed) == 1 + ) + revisions = await restarted.store.list_revisions( + account_id=configuration.account_id, reporting_obligation_id=obligation_id + ) + assert len(revisions) == 1 + async with pool.connection() as connection: + rows = await ( + await connection.execute( + "SELECT reporting_obligation_id FROM reporting_obligations" + ) + ).fetchall() + assert rows == [(obligation_id,)] + finally: + await restarted.close() diff --git a/tests/test_reliable_reporting_lifecycle.py b/tests/test_reliable_reporting_lifecycle.py new file mode 100644 index 000000000..c7f6d12fd --- /dev/null +++ b/tests/test_reliable_reporting_lifecycle.py @@ -0,0 +1,752 @@ +"""Lifecycle barriers through the public service and actual inline adapter.""" + +from __future__ import annotations + +import asyncio +import threading +from dataclasses import replace +from datetime import timedelta +from pathlib import Path +from typing import Any + +import pytest + +from adcp.reporting.fixtures import redacted_capabilities, redacted_snapshot_request +from adcp.reporting.inline_source import FileSystemStagingStore, InlineReportingSource +from adcp.reporting.ledger import ( + InMemoryReportingLedgerStore, + ReportingProducer, + ReportingStatusCaller, +) +from adcp.reporting.service import ( + ReliableReportingService, + ReliableReportingServiceError, + ReliableReportingShutdownTimeoutError, + ReliableReportingState, + ReliableReportingUnavailableError, + ReportingServiceResource, +) +from adcp.reporting.testing import ScriptedReportingAdapter +from adcp.server import ADCPHandler +from tests.test_reliable_reporting_service import ( + NOW, + _account_context, + _configuration, + _rows, +) + + +async def checkpoint() -> None: + """Deliver scheduled task/cancellation callbacks, without advancing domain time.""" + loop = asyncio.get_running_loop() + reached = loop.create_future() + loop.call_soon(reached.set_result, None) + await reached + + +class BlockingSchema(InMemoryReportingLedgerStore): + def __init__(self) -> None: + super().__init__() + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.calls = 0 + + async def create_schema(self) -> None: + self.calls += 1 + self.entered.set() + await self.release.wait() + + +async def test_injected_components_are_borrowed_even_across_repeated_close() -> None: + class BorrowedWorker: + closes = 0 + + async def run_once(self) -> None: + return None + + async def close(self) -> None: + self.closes += 1 + + worker = BorrowedWorker() + service = ReliableReportingService.memory( + account_context=_account_context, materialization_worker=worker + ) + await service.start() + await service.close() + await service.close() + assert worker.closes == 0 + + +async def test_concurrent_initialization_has_one_schema_operation() -> None: + store = BlockingSchema() + service = ReliableReportingService(store=store, account_context=_account_context) + first = asyncio.create_task(service.initialize()) + await asyncio.wait_for(store.entered.wait(), 2) + second = asyncio.create_task(service.initialize()) + try: + await checkpoint() + assert store.calls == 1 + finally: + store.release.set() + await asyncio.gather(first, second, return_exceptions=True) + await service.close() + + +async def test_close_waits_for_a_start_already_in_progress() -> None: + store = BlockingSchema() + service = ReliableReportingService(store=store, account_context=_account_context) + starting = asyncio.create_task(service.start()) + await asyncio.wait_for(store.entered.wait(), 2) + closing = asyncio.create_task(service.close()) + try: + await checkpoint() + assert not closing.done(), "schema work still owns the store" + finally: + store.release.set() + await asyncio.gather(starting, closing, return_exceptions=True) + + +async def test_closed_service_rejects_reporting_rpc_admission() -> None: + service = ReliableReportingService.memory( + account_context=_account_context, + caller_resolver=lambda _request, _context: ReportingStatusCaller("account", "buyer"), + ) + await service.start() + await service.close() + with pytest.raises(RuntimeError, match="unavailable|closed"): + await service.get_reporting_status({"view": "periods"}) + + +async def test_cancelled_sync_source_remains_unsettled_until_its_thread_finishes() -> None: + entered = asyncio.Event() + release = threading.Event() + finished = threading.Event() + loop = asyncio.get_running_loop() + + def fetch(_request: Any) -> list[dict[str, Any]]: + loop.call_soon_threadsafe(entered.set) + try: + assert release.wait(5), "test cleanup watchdog" + return [] + finally: + finished.set() + + source = InlineReportingSource( + capabilities=redacted_capabilities(), fetch=fetch, clock=lambda: NOW + ) + task = asyncio.create_task(source.execute(redacted_snapshot_request(), cancel=asyncio.Event())) + await asyncio.wait_for(entered.wait(), 2) + try: + task.cancel() + await checkpoint() + await checkpoint() + assert not task.done(), "cancelled source returned while its sync fetch was still running" + task.cancel() # repeated cancellation must not bypass the settlement barrier + await checkpoint() + assert not task.done() + finally: + release.set() + await asyncio.gather(task, return_exceptions=True) + await asyncio.to_thread(finished.wait, 2) + assert task.cancelled() + assert finished.is_set() + + +async def test_cancelled_filesystem_staging_joins_its_owned_write(tmp_path: Path) -> None: + entered = asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + + class Staging(FileSystemStagingStore): + @staticmethod + def _write(target: Path, payload: bytes) -> None: + loop.call_soon_threadsafe(entered.set) + assert release.wait(5), "test cleanup watchdog" + FileSystemStagingStore._write(target, payload) + + staging = Staging(tmp_path) + values = dict(account_id="account", source_execution_key="attempt", ordinal=0, payload=b"rows") + task = asyncio.create_task(staging.stage(**values)) + await asyncio.wait_for(entered.wait(), 2) + try: + task.cancel() + await checkpoint() + task.cancel() + await checkpoint() + assert not task.done() + finally: + release.set() + await asyncio.gather(task, return_exceptions=True) + reference, generation = await staging.stage(**values) + assert ( + await staging.read( + object_ref=reference, + object_generation=generation, + account_id="account", + source_scope={}, + cancel=asyncio.Event(), + ) + == b"rows" + ) + + +async def test_startup_failure_unwinds_transferred_resources_in_reverse_order() -> None: + events: list[str] = [] + + def resource(name: str, *, fail: bool = False) -> ReportingServiceResource: + async def opening() -> None: + events.append(f"open:{name}") + if fail: + raise RuntimeError("secret-provider-response") + + async def closing() -> None: + events.append(f"close:{name}") + + return ReportingServiceResource(open=opening, close=closing) + + service = ReliableReportingService.memory( + account_context=_account_context, + owned_resources=(resource("pool"), resource("writer", fail=True), resource("sender")), + ) + with pytest.raises(ReliableReportingServiceError, match="startup") as caught: + await service.start() + assert "secret" not in str(caught.value) + assert service.state is ReliableReportingState.FAILED + assert service.capability_block() == {} + await service.close() + await service.close() + assert events == ["open:pool", "open:writer", "close:sender", "close:writer", "close:pool"] + + +async def test_failed_closer_does_not_skip_other_resources_or_repeat_close() -> None: + events: list[str] = [] + + async def close_pool() -> None: + events.append("pool") + + async def close_writer() -> None: + events.append("writer") + raise RuntimeError("secret-close-body") + + service = ReliableReportingService.memory( + account_context=_account_context, + owned_resources=( + ReportingServiceResource(close=close_pool), + ReportingServiceResource(close=close_writer), + ), + ) + await service.start() + await asyncio.gather(service.close(), service.close(), service.close()) + assert events == ["writer", "pool"] + assert service.state is ReliableReportingState.FAILED + with pytest.raises(ReliableReportingServiceError, match="shutdown"): + await service.wait() + assert "secret" not in repr(service.failure) + + +async def test_cancelled_starter_keeps_resources_until_startup_settles() -> None: + store = BlockingSchema() + closed = asyncio.Event() + + async def close_owned() -> None: + closed.set() + + service = ReliableReportingService( + store=store, + account_context=_account_context, + owned_resources=(ReportingServiceResource(close=close_owned),), + ) + starting = asyncio.create_task(service.start()) + await asyncio.wait_for(store.entered.wait(), 2) + starting.cancel() + with pytest.raises(asyncio.CancelledError): + await starting + try: + assert service.state is ReliableReportingState.STOPPING + assert not closed.is_set() + assert service.capability_block() == {} + with pytest.raises(ReliableReportingUnavailableError): + await service.start() + finally: + store.release.set() + await asyncio.wait_for(service.close(), 2) + assert closed.is_set() + assert service.state is ReliableReportingState.CLOSED + + +@pytest.mark.parametrize("background", [False, True]) +async def test_timed_out_and_cancelled_close_retain_live_sync_adapter(background: bool) -> None: + entered = asyncio.Event() + release = threading.Event() + finished = threading.Event() + closed: list[str] = [] + loop = asyncio.get_running_loop() + + class Adapter: + capabilities = redacted_capabilities() + + def fetch_slice(self, _request: Any) -> list[dict[str, Any]]: + loop.call_soon_threadsafe(entered.set) + try: + assert release.wait(5), "test cleanup watchdog" + return _rows(5) + finally: + finished.set() + + async def close(self) -> None: + assert finished.is_set(), "resource was closed with a live provider thread" + closed.append("adapter") + + adapter = Adapter() + service = ReliableReportingService.memory( + account_context=_account_context, + clock=lambda: NOW, + worker_interval=timedelta(days=1) if background else None, + owned_resources=(ReportingServiceResource(close=adapter.close),), + ) + service.sources.register("gam", adapter) + await service.configure(_configuration()) + await service.start() + turn = None if background else asyncio.create_task(service.run_worker()) + await asyncio.wait_for(entered.wait(), 2) + try: + # Cancel the public producer caller as well as a close waiter. The + # threaded adapter must remain owned through both cancellation paths. + if turn is not None: + turn.cancel() + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert service.state is ReliableReportingState.STOPPING + if turn is not None: + turn.cancel() + await checkpoint() + await checkpoint() + assert not turn.done(), "repeated producer cancellation abandoned its source task" + waiter = asyncio.create_task(service.close()) + await checkpoint() + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + assert service.state is ReliableReportingState.STOPPING + assert closed == [] + assert service.capability_block() == {} + with pytest.raises(ReliableReportingUnavailableError): + await service.run_worker() + finally: + release.set() + if turn is not None: + await asyncio.gather(turn, return_exceptions=True) + await asyncio.wait_for(service.close(), 2) + assert finished.is_set() + assert closed == ["adapter"] + assert service.state is ReliableReportingState.CLOSED + await service.close() + assert closed == ["adapter"] + + +@pytest.mark.parametrize( + "method", + [ + "get_reporting_status", + "sync_reporting_status", + "sync_reporting_receipts", + "get_revision_content", + ], +) +async def test_stopping_rejects_all_reporting_rpc_admission_before_authorization( + method: str, +) -> None: + entered = asyncio.Event() + release = asyncio.Event() + resolutions: list[str] = [] + + class Worker: + async def run_once(self) -> None: + entered.set() + await release.wait() + + def caller(_request: Any, _context: Any) -> ReportingStatusCaller: + resolutions.append("called") + return ReportingStatusCaller("account", "buyer") + + service = ReliableReportingService.memory( + account_context=_account_context, caller_resolver=caller, materialization_worker=Worker() + ) + turn = asyncio.create_task(service.run_worker()) + await asyncio.wait_for(entered.wait(), 2) + try: + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + with pytest.raises(ReliableReportingUnavailableError): + await getattr(service, method)({}) + with pytest.raises(ReliableReportingUnavailableError): + await service.configure(_configuration()) + assert resolutions == [] + finally: + release.set() + await turn + await service.close() + + +async def test_cancelled_async_adapter_settles_cleanup_before_owned_close() -> None: + entered = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + events: list[str] = [] + + class Adapter: + capabilities = redacted_capabilities() + + async def fetch_slice(self, _request: Any) -> list[dict[str, Any]]: + entered.set() + try: + await asyncio.Event().wait() + finally: + cleanup_started.set() + await cleanup_release.wait() + events.append("fetch-settled") + return [] + + async def aclose(self) -> None: + events.append("adapter-closed") + + adapter = Adapter() + service = ReliableReportingService.memory( + account_context=_account_context, + clock=lambda: NOW, + owned_resources=(ReportingServiceResource(close=adapter.aclose),), + ) + service.sources.register("gam", adapter) + await service.configure(_configuration()) + turn = asyncio.create_task(service.run_worker()) + await asyncio.wait_for(entered.wait(), 2) + try: + turn.cancel() + await asyncio.wait_for(cleanup_started.wait(), 2) + turn.cancel() + await checkpoint() + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert service.state is ReliableReportingState.STOPPING + assert not turn.done() + assert events == [] + finally: + cleanup_release.set() + await asyncio.gather(turn, return_exceptions=True) + await service.close() + assert events == ["fetch-settled", "adapter-closed"] + + +@pytest.mark.parametrize("fail_worker", [False, True]) +async def test_installed_receipt_admission_drains_before_owned_cleanup(fail_worker: bool) -> None: + entered = asyncio.Event() + release = asyncio.Event() + failed = asyncio.Event() + order: list[str] = [] + + class Receipts: + async def handle(self, request: dict[str, Any], **_caller: Any) -> dict[str, Any]: + entered.set() + await release.wait() + order.append("receipt-settled") + return request + + class Materializer: + async def run_once(self) -> None: + await entered.wait() + raise RuntimeError("secret-worker-body") + + async def close_owned() -> None: + order.append("resource-closed") + + service = ReliableReportingService.memory( + account_context=_account_context, + caller_resolver=lambda _request, _context: ReportingStatusCaller("account", "buyer"), + materialization_worker=Materializer(), + receipt_handler=Receipts(), + worker_interval=timedelta(days=1) if fail_worker else None, + worker_error_handler=lambda _component, _error: failed.set(), + owned_resources=(ReportingServiceResource(close=close_owned),), + ) + handler = service.install(ADCPHandler()) + await service.start() + receipt = asyncio.create_task(handler.sync_reporting_receipts({"request_id": "kept"})) + await asyncio.wait_for(entered.wait(), 2) + try: + if fail_worker: + await asyncio.wait_for(failed.wait(), 2) + assert service.state is ReliableReportingState.STOPPING + assert service.failure is not None + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert order == [] + with pytest.raises(ReliableReportingUnavailableError): + await handler.sync_reporting_receipts({}) + finally: + release.set() + result = await receipt + await service.close() + assert result == {"request_id": "kept"} + assert order == ["receipt-settled", "resource-closed"] + if fail_worker: + with pytest.raises(ReliableReportingServiceError, match="materialization"): + await service.wait() + + +async def test_repeated_rpc_cancellation_cannot_interrupt_transaction_cleanup() -> None: + entered = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + settled = asyncio.Event() + closed = asyncio.Event() + + class Receipts: + async def handle(self, request: dict[str, Any], **_caller: Any) -> dict[str, Any]: + entered.set() + try: + await asyncio.Event().wait() + finally: + cleanup_started.set() + await cleanup_release.wait() + settled.set() + return request + + class Materializer: + async def run_once(self) -> None: + return None + + async def close_owned() -> None: + closed.set() + + service = ReliableReportingService.memory( + account_context=_account_context, + caller_resolver=lambda _request, _context: ReportingStatusCaller("account", "buyer"), + materialization_worker=Materializer(), + receipt_handler=Receipts(), + owned_resources=(ReportingServiceResource(close=close_owned),), + ) + await service.start() + call = asyncio.create_task(service.sync_reporting_receipts({})) + await asyncio.wait_for(entered.wait(), 2) + try: + call.cancel() + await asyncio.wait_for(cleanup_started.wait(), 2) + call.cancel() + await checkpoint() + await checkpoint() + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + assert not call.done(), "repeated cancellation interrupted admitted cleanup" + assert service.state is ReliableReportingState.STOPPING + assert not closed.is_set() + finally: + cleanup_release.set() + await asyncio.gather(call, return_exceptions=True) + await service.close() + assert settled.is_set() + assert closed.is_set() + + +async def test_failure_withdraws_installed_capabilities_and_redacts_diagnostics( + caplog: Any, +) -> None: + entered = asyncio.Event() + release = asyncio.Event() + errors: list[str] = [] + + class Worker: + calls = 0 + + async def run_once(self) -> None: + self.calls += 1 + entered.set() + await release.wait() + raise RuntimeError("secret-provider-body https://destination.test/?token=secret") + + class Handler(ADCPHandler): + async def get_adcp_capabilities(self, params: Any, context: Any = None) -> Any: + return {"supported_protocols": ["media_buy"]} + + worker = Worker() + + def report_error(_component: str, error: BaseException) -> None: + errors.append(str(error)) + raise RuntimeError("secret-error-handler-body") + + service = ReliableReportingService.memory( + account_context=_account_context, + clock=lambda: NOW, + materialization_worker=worker, + worker_interval=timedelta(days=1), + worker_error_handler=report_error, + ) + service.sources.register("gam", ScriptedReportingAdapter(redacted_capabilities(), [_rows(1)])) + await service.configure(_configuration()) + handler = service.install(Handler()) + assert "reporting_delivery" not in (await handler.get_adcp_capabilities({})).get( + "media_buy", {} + ) + await service.start() + await asyncio.wait_for(entered.wait(), 2) + assert (await handler.get_adcp_capabilities({}))["media_buy"]["reporting_delivery"]["supported"] + release.set() + with pytest.raises(ReliableReportingServiceError, match="materialization"): + await asyncio.wait_for(service.wait(), 2) + assert worker.calls == 1 + assert not service.ready + assert "reporting_delivery" not in (await handler.get_adcp_capabilities({})).get( + "media_buy", {} + ) + assert "secret" not in caplog.text + assert errors == ["reporting service failed (materialization)"] + + +async def test_configure_during_startup_is_serialized_without_losing_an_accepted_generation() -> ( + None +): + entered = asyncio.Event() + release = asyncio.Event() + + async def context(configuration: Any) -> Any: + entered.set() + await release.wait() + return _account_context(configuration) + + service = ReliableReportingService.memory(account_context=context) + service.sources.register("gam", ScriptedReportingAdapter(redacted_capabilities(), [_rows(1)])) + configuration = _configuration() + configuring = asyncio.create_task(service.configure(configuration)) + await asyncio.wait_for(entered.wait(), 2) + starting = asyncio.create_task(service.start()) + try: + await checkpoint() + assert service.state is ReliableReportingState.STARTING + finally: + release.set() + await asyncio.gather(configuring, starting) + assert await service.store.list_configurations(account_id=configuration.account_id) == ( + configuration, + ) + await service.close() + + +async def test_retryable_source_outcome_does_not_fail_supervision() -> None: + service = ReliableReportingService.memory(account_context=_account_context, clock=lambda: NOW) + adapter = ScriptedReportingAdapter(redacted_capabilities(), [None, _rows(8)]) + service.sources.register("gam", adapter) + await service.configure(replace(_configuration(), deactivated_at=NOW)) + first = await service.run_worker() + assert any(turn.slices_failed for turn in first.configurations.values()) + assert service.ready + assert service.failure is None + second = await service.run_worker() + assert any(turn.revisions_committed for turn in second.configurations.values()) + await service.close() + + +async def test_low_level_producer_keeps_existing_lease_until_cancelled_sync_work_settles() -> None: + entered = asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + + def fetch(_request: Any) -> list[dict[str, Any]]: + loop.call_soon_threadsafe(entered.set) + assert release.wait(5), "test cleanup watchdog" + return _rows(3) + + configuration = _configuration() + store = InMemoryReportingLedgerStore(clock=lambda: NOW) + await store.put_configuration(configuration) + source = InlineReportingSource( + capabilities=redacted_capabilities(), fetch=fetch, clock=lambda: NOW + ) + producer = ReportingProducer( + source=source, + object_reader=source.staging, + offerings=_account_context(configuration).producer_offerings(), + store=store, + clock=lambda: NOW, + worker_id="first", + ) + running = asyncio.create_task(producer.run_worker()) + await asyncio.wait_for(entered.wait(), 2) + try: + running.cancel() + await checkpoint() + running.cancel() + await checkpoint() + await checkpoint() + assert not running.done() + assert await store.lease_period_close(worker_id="second", now=NOW, lease_seconds=60) is None + finally: + release.set() + await asyncio.gather(running, return_exceptions=True) + lease = await store.lease_period_close(worker_id="second", now=NOW, lease_seconds=60) + assert lease is not None + await store.release_period_close(lease, worker_id="second") + + +async def test_mounted_mcp_status_call_drains_and_stopping_rejects_a_new_call() -> None: + import httpx + from asgi_lifespan import LifespanManager + + from adcp.server import create_mcp_server + from tests.test_mcp_middleware_composition import ( + _call_tool, + _initialize_session, + _parse_event_stream, + ) + + entered = asyncio.Event() + release = asyncio.Event() + resolutions: list[str] = [] + + class Handler(ADCPHandler): + def get_adcp_version(self) -> str: + return "3.2.0-rc.6" + + async def caller(_request: Any, _context: Any) -> ReportingStatusCaller: + resolutions.append("authorized") + entered.set() + await release.wait() + return ReportingStatusCaller("account-redacted", "buyer") + + service = ReliableReportingService.memory( + account_context=_account_context, caller_resolver=caller, clock=lambda: NOW + ) + service.sources.register("gam", ScriptedReportingAdapter(redacted_capabilities(), [])) + await service.configure(_configuration()) + handler = service.install(Handler()) + mcp = create_mcp_server(handler, stateless_http=True, allowed_hosts=["localhost"]) + app = mcp.streamable_http_app() + await service.start() + async with ( + LifespanManager(app), + httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://localhost", + follow_redirects=True, + ) as client, + ): + await _initialize_session(client) + request = {"account": {"account_id": "account-redacted"}, "view": "periods"} + first = asyncio.create_task(_call_tool(client, "get_reporting_status", request)) + started = asyncio.create_task(entered.wait()) + try: + completed, _ = await asyncio.wait( + {first, started}, timeout=5, return_when=asyncio.FIRST_COMPLETED + ) + assert started in completed, (await first).text if first.done() else "no admission" + with pytest.raises(ReliableReportingShutdownTimeoutError): + await service.close(timeout=0) + response = await _call_tool(client, "get_reporting_status", request) + assert _parse_event_stream(response.text)["result"]["isError"] is True + assert resolutions == ["authorized"] + finally: + release.set() + started.cancel() + await asyncio.gather(started, return_exceptions=True) + completed = await first + await service.close() + assert not _parse_event_stream(completed.text)["result"].get("isError", False) + assert service.state is ReliableReportingState.CLOSED diff --git a/tests/test_reliable_reporting_service.py b/tests/test_reliable_reporting_service.py index e62500ff1..2fa7d772a 100644 --- a/tests/test_reliable_reporting_service.py +++ b/tests/test_reliable_reporting_service.py @@ -26,6 +26,8 @@ from adcp.reporting.service import ( ReliableReportingConfigurationError, ReliableReportingService, + ReliableReportingServiceError, + ReliableReportingState, ReportingAccountContext, ) from adcp.reporting.testing import ( @@ -235,6 +237,10 @@ async def test_capability_block_is_schema_valid_and_only_advertises_installed_ti ) service.sources.register("gam", ScriptedReportingAdapter(redacted_capabilities(), [_rows(1)])) await service.configure(_configuration()) + assert service.capability_block() == {} + await service.initialize() + assert service.capability_block() == {} + await service.start() block = service.capability_block() assert block["managed_delivery"] is False @@ -247,6 +253,7 @@ async def test_capability_block_is_schema_valid_and_only_advertises_installed_ti validator = get_named_validator("core/reporting-delivery-capabilities.json") assert validator is not None assert list(validator.iter_errors(block)) == [] + await service.close() async def test_startup_rejects_impossible_component_combinations() -> None: @@ -272,8 +279,7 @@ async def handle(self, request: dict[str, Any], **kwargs: Any) -> dict[str, Any] ).initialize() -async def test_background_worker_reports_an_error_and_recovers_on_the_next_turn() -> None: - recovered = asyncio.Event() +async def test_background_worker_reports_an_error_and_stops_admitting_work() -> None: errors: list[tuple[str, str]] = [] class FlakyWorker: @@ -283,25 +289,29 @@ async def run_once(self) -> str: self.calls += 1 if self.calls == 1: raise RuntimeError("temporary materializer failure") - recovered.set() return "recovered" async def capture(component: str, error: BaseException) -> None: errors.append((component, str(error))) + worker = FlakyWorker() service = ReliableReportingService.memory( account_context=_account_context, - materialization_worker=FlakyWorker(), + materialization_worker=worker, worker_interval=timedelta(milliseconds=1), worker_error_handler=capture, ) await service.start() try: - await asyncio.wait_for(recovered.wait(), timeout=1) + with pytest.raises(ReliableReportingServiceError, match="materialization"): + await asyncio.wait_for(service.wait(), timeout=1) finally: await service.close() - assert errors == [("materialization", "temporary materializer failure")] + assert errors == [("materialization", "reporting service failed (materialization)")] + assert worker.calls == 1 + assert service.state is ReliableReportingState.FAILED + assert not service.ready async def test_configuration_rejects_unregistered_routes_and_mutated_generations() -> None: diff --git a/tests/type_checks/reliable_reporting_lifecycle.py b/tests/type_checks/reliable_reporting_lifecycle.py new file mode 100644 index 000000000..956a43c97 --- /dev/null +++ b/tests/type_checks/reliable_reporting_lifecycle.py @@ -0,0 +1,55 @@ +"""Public ownership/liveness APIs used by a strictly typed adopter.""" + +from typing import Protocol + +from adcp.reporting import ( + ReliableReportingService, + ReliableReportingServiceError, + ReliableReportingShutdownTimeoutError, + ReliableReportingState, + ReliableReportingUnavailableError, + ReportingServiceResource, +) +from adcp.reporting.ledger import ReportingConfiguration, ReportingLedgerStore +from adcp.reporting.service import ReportingContextResolver + + +class AsyncResource(Protocol): + async def open(self) -> None: ... + + async def close(self) -> None: ... + + +def compose( + store: ReportingLedgerStore, + context: ReportingContextResolver, + owned: AsyncResource, +) -> ReliableReportingService: + # The store is borrowed. Only the explicitly transferred resource closes. + return ReliableReportingService( + store=store, + account_context=context, + owned_resources=(ReportingServiceResource(open=owned.open, close=owned.close),), + ) + + +async def serve_one_turn( + service: ReliableReportingService, configuration: ReportingConfiguration +) -> None: + try: + await service.configure(configuration) + await service.run_worker() + except ReliableReportingUnavailableError as unavailable: + state: ReliableReportingState = unavailable.state + assert state is not ReliableReportingState.READY + finally: + try: + await service.close(timeout=5) + except ReliableReportingShutdownTimeoutError: + # Retain the service and its loop until admitted work has settled. + await service.close() + try: + await service.wait() + except ReliableReportingServiceError as failure: + component: str = failure.component + assert component From 84bd392a7aaad754dbe014f88b56dd048861765a Mon Sep 17 00:00:00 2001 From: Brian O'Kelley Date: Sat, 26 Sep 2026 14:53:49 +0000 Subject: [PATCH 2/2] fix(reporting): isolate worker failures and preserve error callbacks --- docs/reliable-reporting-service.md | 17 +++- src/adcp/reporting/service.py | 42 +++++--- ...st_reliable_reporting_service_lifecycle.py | 69 +++++++++++++ tests/test_reliable_reporting_lifecycle.py | 69 +++++++++---- tests/test_reliable_reporting_service.py | 99 ++++++++++++++++--- .../reliable_reporting_lifecycle.py | 6 +- 6 files changed, 251 insertions(+), 51 deletions(-) diff --git a/docs/reliable-reporting-service.md b/docs/reliable-reporting-service.md index 5f3f5e059..73e415ab0 100644 --- a/docs/reliable-reporting-service.md +++ b/docs/reliable-reporting-service.md @@ -136,9 +136,20 @@ both managed delivery and receipts. Capability output is derived from the components actually installed. Unexpected failures are isolated by configuration or extension, logged, and -included on `ReliableReportingTurn`. A configured `worker_error_handler` can -page or emit telemetry; the lifecycle worker continues with its next turn -instead of silently stopping. +included on `ReliableReportingTurn.configuration_errors` or `extension_errors`. +Other configurations and extensions still run, and the background scheduler +retries on its next turn. These failures do not change service readiness. +`worker_error_handler(component, error)` receives the original exception and +the component name: `configuration:{account_id}:{delivery_config_id}@{version}`, +`materialization`, or `notification`. SDK logs contain only the component kind; +adopter callbacks control any further diagnostics. A callback failure is logged +without interrupting the remaining work. + +Failures outside an individual configuration or extension turn stop the +background scheduler, notify the callback as `service` with the original error, +and withdraw readiness. Startup, scheduler, and resource cleanup failures are +terminal: after admitted work settles, `wait()` raises a sanitized +`ReliableReportingServiceError`. Restart those services with a new instance. `run_worker()` is also available for an external scheduler. Calls within one service process are serialized. Ledger writes are convergent, but deployments diff --git a/src/adcp/reporting/service.py b/src/adcp/reporting/service.py index e7226d596..feb467195 100644 --- a/src/adcp/reporting/service.py +++ b/src/adcp/reporting/service.py @@ -45,7 +45,6 @@ ) from adcp.reporting.ledger.store import LedgerConflictError, decode_cursor from adcp.reporting.service_lifecycle import ( - FailureComponent, ReliableReportingServiceError, ReliableReportingShutdownTimeoutError, ReliableReportingState, @@ -608,7 +607,7 @@ def ready(self) -> bool: @property def failure(self) -> ReliableReportingServiceError | None: - """Sanitized first unexpected failure, if any.""" + """Sanitized first terminal lifecycle failure, if any.""" return self._lifecycle.failure async def close(self, *, timeout: float | None = None) -> None: @@ -634,12 +633,21 @@ async def __aexit__(self, *_exc: object) -> None: async def _worker_loop(self) -> None: assert self._worker_interval is not None while not self._lifecycle.stopping: - await self.run_worker() + try: + await self.run_worker() + except Exception as error: + # Configuration and extension failures are isolated inside a + # turn. Only a failure of the scheduler itself stops service. + if not self._lifecycle.stopping: + await self._lifecycle.fail("service") + await self._report_worker_error("service", error) + return await self._lifecycle.wait_for_stop(self._worker_interval.total_seconds()) - async def _report_worker_error(self, component: FailureComponent) -> None: - error = await self._lifecycle.fail(component) - logger.error("Reliable Reporting worker stopped: %s", error) + async def _report_worker_error(self, component: str, error: BaseException) -> None: + # Preserve the adopter callback's detailed identity and original error, + # without writing account identifiers or provider bodies to SDK logs. + logger.error("Reliable Reporting worker component %s failed", component.split(":", 1)[0]) if self._worker_error_handler is None: return try: @@ -673,10 +681,14 @@ async def _run_worker(self, *, now: datetime | None) -> ReliableReportingTurn: turn.configurations[key] = await binding.producer.run_configuration( binding.configuration, now=now ) - except Exception: - await self._report_worker_error("configuration") - assert self.failure is not None - raise self.failure from None + except Exception as error: + turn.configuration_errors[key] = error + await self._report_worker_error( + "configuration:" + f"{key.account_id}:{key.delivery_config_id}@" + f"{key.delivery_config_version}", + error, + ) extensions = ( ("materialization", self._materialization_worker), ("notification", self._notification_worker), @@ -687,12 +699,10 @@ async def _run_worker(self, *, now: datetime | None) -> ReliableReportingTurn: if extension is not None: try: result = await extension.run_once() - except Exception: - await self._report_worker_error( - "materialization" if name == "materialization" else "notification" - ) - assert self.failure is not None - raise self.failure from None + except Exception as error: + turn.extension_errors[name] = error + await self._report_worker_error(name, error) + continue if result is not None: turn.extension_results.append(result) return turn diff --git a/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py b/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py index 781b2a7f1..89a3e7970 100644 --- a/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py +++ b/tests/conformance/reporting/test_reliable_reporting_service_lifecycle.py @@ -26,6 +26,75 @@ from ._generation_support import isolated_reporting_pool +@pytest.mark.parametrize("autocommit", [False, True]) +async def test_configuration_database_error_isolated_and_retried(autocommit: bool) -> None: + async with isolated_reporting_pool(autocommit=autocommit) as pool: + from psycopg import OperationalError + + async with pool.connection() as connection: + row = await (await connection.execute("SELECT clock_timestamp()")).fetchone() + now = row[0] + boundary = now.replace(minute=0, second=0, microsecond=0) + errors: list[tuple[str, BaseException]] = [] + + class FlakyLedger(PgReportingLedgerStore): + failure: OperationalError | None = None + + async def find_obligation(self, **kwargs: Any) -> Any: + if kwargs["account_id"] == "account-a" and self.failure is None: + # Raise a real driver OperationalError inside a transaction; + # its rollback must not prevent the next account from running. + try: + async with pool.connection() as connection: + await connection.execute( + "DO $$ BEGIN RAISE EXCEPTION 'temporary ledger failure' " + "USING ERRCODE = '08006'; END $$" + ) + except OperationalError as error: + self.failure = error + raise + return await super().find_obligation(**kwargs) + + ledger = FlakyLedger(pool=pool) + service = ReliableReportingService( + store=ledger, + account_context=_account_context, + clock=lambda: now, + worker_error_handler=lambda component, error: errors.append((component, error)), + ) + adapter = ScriptedReportingAdapter(redacted_capabilities(), [_rows(11), _rows(12)]) + service.sources.register("gam", adapter) + configurations = [ + replace( + _configuration(account_id=account), + activated_at=boundary - timedelta(hours=3) + timedelta(minutes=20), + deactivated_at=boundary - timedelta(hours=1), + ) + for account in ("account-a", "account-b") + ] + for configuration in configurations: + await service.configure(configuration) + failed_key, healthy_key = (item.generation_key for item in configurations) + try: + first = await service.run_worker() + assert isinstance(ledger.failure, OperationalError) + assert first.configuration_errors == {failed_key: ledger.failure} + assert set(first.configurations) == {healthy_key} + assert len(first.configurations[healthy_key].revisions_committed) == 1 + assert errors == [("configuration:account-a:gam-delivery@1", ledger.failure)] + assert service.ready + assert service.failure is None + + recovered = await service.run_worker() + assert not recovered.configuration_errors + assert len(recovered.configurations[failed_key].revisions_committed) == 1 + assert len(adapter.calls) == 2 + assert service.ready + finally: + await service.close() + await service.wait() + + @pytest.mark.parametrize("autocommit", [False, True]) @pytest.mark.parametrize("cancel_call", [False, True]) async def test_stop_settles_public_configuration_transaction_before_owned_cleanup( diff --git a/tests/test_reliable_reporting_lifecycle.py b/tests/test_reliable_reporting_lifecycle.py index c7f6d12fd..1bab9ece1 100644 --- a/tests/test_reliable_reporting_lifecycle.py +++ b/tests/test_reliable_reporting_lifecycle.py @@ -5,7 +5,7 @@ import asyncio import threading from dataclasses import replace -from datetime import timedelta +from datetime import datetime, timedelta from pathlib import Path from typing import Any @@ -23,6 +23,7 @@ ReliableReportingServiceError, ReliableReportingShutdownTimeoutError, ReliableReportingState, + ReliableReportingTurn, ReliableReportingUnavailableError, ReportingServiceResource, ) @@ -453,13 +454,17 @@ async def handle(self, request: dict[str, Any], **_caller: Any) -> dict[str, Any class Materializer: async def run_once(self) -> None: + return None + + class BrokenScheduler(ReliableReportingService): + async def run_worker(self, *, now: datetime | None = None) -> ReliableReportingTurn: await entered.wait() - raise RuntimeError("secret-worker-body") + raise RuntimeError("secret-scheduler-body") async def close_owned() -> None: order.append("resource-closed") - service = ReliableReportingService.memory( + service = BrokenScheduler.memory( account_context=_account_context, caller_resolver=lambda _request, _context: ReportingStatusCaller("account", "buyer"), materialization_worker=Materializer(), @@ -489,7 +494,7 @@ async def close_owned() -> None: assert result == {"request_id": "kept"} assert order == ["receipt-settled", "resource-closed"] if fail_worker: - with pytest.raises(ReliableReportingServiceError, match="materialization"): + with pytest.raises(ReliableReportingServiceError, match="service"): await service.wait() @@ -547,12 +552,16 @@ async def close_owned() -> None: assert closed.is_set() -async def test_failure_withdraws_installed_capabilities_and_redacts_diagnostics( +@pytest.mark.parametrize("fail_scheduler", [False, True]) +async def test_only_scheduler_failure_withdraws_capabilities_and_sdk_logs_stay_redacted( caplog: Any, + fail_scheduler: bool, ) -> None: entered = asyncio.Event() release = asyncio.Event() - errors: list[str] = [] + reported = asyncio.Event() + errors: list[tuple[str, BaseException]] = [] + failure = RuntimeError("secret-provider-body https://destination.test/?token=secret") class Worker: calls = 0 @@ -561,7 +570,15 @@ async def run_once(self) -> None: self.calls += 1 entered.set() await release.wait() - raise RuntimeError("secret-provider-body https://destination.test/?token=secret") + raise failure + + class Service(ReliableReportingService): + async def run_worker(self, *, now: datetime | None = None) -> ReliableReportingTurn: + if fail_scheduler: + entered.set() + await release.wait() + raise failure + return await super().run_worker(now=now) class Handler(ADCPHandler): async def get_adcp_capabilities(self, params: Any, context: Any = None) -> Any: @@ -569,11 +586,12 @@ async def get_adcp_capabilities(self, params: Any, context: Any = None) -> Any: worker = Worker() - def report_error(_component: str, error: BaseException) -> None: - errors.append(str(error)) + def report_error(component: str, error: BaseException) -> None: + errors.append((component, error)) + reported.set() raise RuntimeError("secret-error-handler-body") - service = ReliableReportingService.memory( + service = Service.memory( account_context=_account_context, clock=lambda: NOW, materialization_worker=worker, @@ -589,16 +607,29 @@ def report_error(_component: str, error: BaseException) -> None: await service.start() await asyncio.wait_for(entered.wait(), 2) assert (await handler.get_adcp_capabilities({}))["media_buy"]["reporting_delivery"]["supported"] - release.set() - with pytest.raises(ReliableReportingServiceError, match="materialization"): - await asyncio.wait_for(service.wait(), 2) - assert worker.calls == 1 - assert not service.ready - assert "reporting_delivery" not in (await handler.get_adcp_capabilities({})).get( - "media_buy", {} - ) + try: + release.set() + await asyncio.wait_for(reported.wait(), 2) + if fail_scheduler: + with pytest.raises(ReliableReportingServiceError, match="service"): + await asyncio.wait_for(service.wait(), 2) + assert not service.ready + assert "secret" not in repr(service.failure) + assert "reporting_delivery" not in (await handler.get_adcp_capabilities({})).get( + "media_buy", {} + ) + else: + assert service.ready + assert service.failure is None + assert (await handler.get_adcp_capabilities({}))["media_buy"]["reporting_delivery"][ + "supported" + ] + finally: + release.set() + await service.close() + assert worker.calls == (0 if fail_scheduler else 1) assert "secret" not in caplog.text - assert errors == ["reporting service failed (materialization)"] + assert errors == [("service" if fail_scheduler else "materialization", failure)] async def test_configure_during_startup_is_serialized_without_losing_an_accepted_generation() -> ( diff --git a/tests/test_reliable_reporting_service.py b/tests/test_reliable_reporting_service.py index 2fa7d772a..c5939e1e7 100644 --- a/tests/test_reliable_reporting_service.py +++ b/tests/test_reliable_reporting_service.py @@ -17,16 +17,17 @@ redacted_snapshot_request, ) from adcp.reporting.ledger import ( + InMemoryReportingLedgerStore, ReportingConfiguration, ReportingConfigurationGenerationKey, ReportingDefinitionBinding, ReportingScheduleSpec, ReportingStatusCaller, ) +from adcp.reporting.ledger.store import LedgerConflictError from adcp.reporting.service import ( ReliableReportingConfigurationError, ReliableReportingService, - ReliableReportingServiceError, ReliableReportingState, ReportingAccountContext, ) @@ -279,8 +280,82 @@ async def handle(self, request: dict[str, Any], **kwargs: Any) -> dict[str, Any] ).initialize() -async def test_background_worker_reports_an_error_and_stops_admitting_work() -> None: - errors: list[tuple[str, str]] = [] +async def test_worker_isolates_configuration_and_extension_errors_and_retries() -> None: + configuration_error = LedgerConflictError("LEASE_LOST", "secret-store-body") + errors: list[tuple[str, BaseException]] = [] + + class FlakyLedger(InMemoryReportingLedgerStore): + failed = False + + async def find_obligation(self, **kwargs: Any) -> Any: + if kwargs["account_id"] == "account-a" and not self.failed: + self.failed = True + raise configuration_error + return await super().find_obligation(**kwargs) + + class FlakyWorker: + def __init__(self, name: str) -> None: + self.name = name + self.error = RuntimeError(f"secret-{name}-body") + self.calls = 0 + + async def run_once(self) -> str: + self.calls += 1 + if self.calls == 1: + raise self.error + return self.name + + materializer = FlakyWorker("materialization") + notifier = FlakyWorker("notification") + service = ReliableReportingService( + store=FlakyLedger(), + account_context=_account_context, + clock=lambda: NOW, + materialization_worker=materializer, + notification_worker=notifier, + notification_attempt_store=object(), + worker_error_handler=lambda component, error: errors.append((component, error)), + ) + adapter = ScriptedReportingAdapter(redacted_capabilities(), [_rows(10), _rows(20)]) + service.sources.register("gam", adapter) + failed_config = _configuration(account_id="account-a") + healthy_config = _configuration(account_id="account-b") + await service.configure(failed_config) + await service.configure(healthy_config) + try: + first = await service.run_worker() + assert first.configuration_errors == {failed_config.generation_key: configuration_error} + assert set(first.configurations) == {healthy_config.generation_key} + assert first.configurations[healthy_config.generation_key].revisions_committed + assert first.extension_errors == { + "materialization": materializer.error, + "notification": notifier.error, + } + assert first.did_work + assert errors == [ + ("configuration:account-a:gam-delivery@1", configuration_error), + ("materialization", materializer.error), + ("notification", notifier.error), + ] + assert service.ready + assert service.failure is None + + recovered = await service.run_worker() + assert not recovered.configuration_errors + assert not recovered.extension_errors + assert recovered.configurations[failed_config.generation_key].revisions_committed + assert recovered.extension_results == ["materialization", "notification"] + assert len(adapter.calls) == 2 + assert service.ready + finally: + await service.close() + await service.wait() + + +async def test_background_worker_reports_an_error_and_recovers_on_the_next_turn() -> None: + errors: list[tuple[str, BaseException]] = [] + failure = RuntimeError("temporary materializer failure") + recovered = asyncio.Event() class FlakyWorker: calls = 0 @@ -288,11 +363,12 @@ class FlakyWorker: async def run_once(self) -> str: self.calls += 1 if self.calls == 1: - raise RuntimeError("temporary materializer failure") + raise failure + recovered.set() return "recovered" async def capture(component: str, error: BaseException) -> None: - errors.append((component, str(error))) + errors.append((component, error)) worker = FlakyWorker() service = ReliableReportingService.memory( @@ -303,15 +379,16 @@ async def capture(component: str, error: BaseException) -> None: ) await service.start() try: - with pytest.raises(ReliableReportingServiceError, match="materialization"): - await asyncio.wait_for(service.wait(), timeout=1) + await asyncio.wait_for(recovered.wait(), timeout=2) + assert errors == [("materialization", failure)] + assert worker.calls >= 2 + assert service.ready + assert service.failure is None finally: await service.close() - assert errors == [("materialization", "reporting service failed (materialization)")] - assert worker.calls == 1 - assert service.state is ReliableReportingState.FAILED - assert not service.ready + await service.wait() + assert service.state is ReliableReportingState.CLOSED async def test_configuration_rejects_unregistered_routes_and_mutated_generations() -> None: diff --git a/tests/type_checks/reliable_reporting_lifecycle.py b/tests/type_checks/reliable_reporting_lifecycle.py index 956a43c97..f02c7a4b9 100644 --- a/tests/type_checks/reliable_reporting_lifecycle.py +++ b/tests/type_checks/reliable_reporting_lifecycle.py @@ -15,9 +15,11 @@ class AsyncResource(Protocol): - async def open(self) -> None: ... + async def open(self) -> None: + pass - async def close(self) -> None: ... + async def close(self) -> None: + pass def compose(