diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index dc83b1f3e..95a71709e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -764,7 +764,7 @@ jobs: runs-on: ubuntu-latest permissions: contents: read - timeout-minutes: 35 + timeout-minutes: 45 services: postgres: image: postgres:16 @@ -786,7 +786,8 @@ jobs: run: pip install -e ".[dev,pg]" - name: Run complete production and projection conformance shell: bash - timeout-minutes: 30 + # A passing suite plus harness cleanup can exceed 30 minutes. + timeout-minutes: 40 env: ADCP_PG_TEST_URL: postgresql://postgres@localhost:5432/adcp_production_test run: | diff --git a/docs/reporting-production.md b/docs/reporting-production.md index 581106c5c..5ad752bed 100644 --- a/docs/reporting-production.md +++ b/docs/reporting-production.md @@ -398,3 +398,15 @@ pre-`e16eb8cf` hardening, and pre-`34c8f6d9` production. In particular, old A wh is not supported, and historical false notification-readiness results do not become healthy through later source integration. These notes do not qualify simultaneous old autonomous writers or release/activation acceptance. + +### Source revocation + +`ReportingProductionSource.configuration_binding(configuration)` is also the +live authorization callback for each dispatch and publication. The SDK checks +again under the account lock before sealing or committing a fetched result. +`None` discards the result without a success checkpoint and stops further source +work for that account in the current turn; restoration resumes on the next turn. +Replay and restart perform fresh checks. The SDK does not cancel an in-flight +fetch for revocation: revocation takes effect at the next dispatch or publish. +Adopters own any caching or latency inside the callback. See the +[release guide](reporting-release-notes.md#source-authorization-d5-option-a). diff --git a/docs/reporting-release-notes.md b/docs/reporting-release-notes.md index 8ddb718e5..99203788e 100644 --- a/docs/reporting-release-notes.md +++ b/docs/reporting-release-notes.md @@ -46,6 +46,25 @@ writes do not authorize concurrent incompatible autonomous materializers, status projectors or clock sweepers. Follow the component-specific migration procedures before starting those workers. +## Source authorization (D5 Option A) + +The adopter's existing `ReportingProductionSource.configuration_binding(configuration)` +is the authority for source work. Returning `None` withdraws authorization. +The SDK calls it immediately before each source dispatch and again under the +account lock before sealing or publishing the result. If authorization has been +withdrawn, the result is discarded without publication or a success checkpoint, +and the account gets no further source work during that turn. A restored +binding resumes work on the next turn. Recovery and replay after a restart make +fresh checks; an earlier successful check grants no authority to publish later. + +An in-flight fetch is allowed to finish: revocation takes effect at the next dispatch or publish. +Adopters own any caching or latency inside their callback. There is no additional +authorization adapter, persisted authorization grant, freshness budget, epoch, +compare-and-swap protocol, or settlement-only mode; the proposed F/L/P values +are not part of this contract. Existing frozen generation mappings still protect +historical scope. Buyer exposure retains the existing per-request feed +reauthorization and per-session destination authorization. + ## Current and historical protocol versions Live reporting mounts and callers use AdCP `3.2-rc.6`, whose bundle spelling is diff --git a/examples/reporting_service_production.py b/examples/reporting_service_production.py new file mode 100644 index 000000000..1b0103da9 --- /dev/null +++ b/examples/reporting_service_production.py @@ -0,0 +1,62 @@ +"""Own an existing B2 graph through the service lifecycle. + +This explicit bridge still requires production providers, fixed source profiles +and the typed account task. It does not implement the separate adapter-first +factory, durable acquisition-envelope, or database-time fencing contracts. +""" + +from collections.abc import Mapping +from typing import Any + +from adcp.reporting.ledger import ReportingProducer +from adcp.reporting.materializer import ReportingMaterializerService +from adcp.reporting.production import ( + ReportingProductionConfigurationTask, + ReportingProductionHandler, + ReportingProductionOffering, + ReportingProductionSourceRegistry, + ReportingProductionSupport, +) +from adcp.reporting.projection import InMemoryReportingStatusProjection, PgReportingStatusProjection +from adcp.reporting.receipts import ReceiptAccountResolver +from adcp.reporting.service import ReliableReportingService, ReportingContextResolver +from adcp.server import ADCPHandler + + +def compose_service( + *, + materializer: ReportingMaterializerService, + projection: InMemoryReportingStatusProjection | PgReportingStatusProjection, + offerings: tuple[ReportingProductionOffering, ...], + producers: Mapping[str, ReportingProducer], + account_context: ReportingContextResolver, + configuration_task: ReportingProductionConfigurationTask, + resolve_account: ReceiptAccountResolver, + application: ADCPHandler[Any], +) -> tuple[ReliableReportingService, ReportingProductionHandler]: + """Register before mounting; start/close own workers, while pools stay borrowed. + + Each stable name identifies one fixed execution profile. Account resolution + must return that profile's currency, scope, metric set and source offerings. + Recovery uses persisted generation facts and the provider's live account + binding; it does not call account_context or enumerate accounts. + + Mount the returned exact handler on MCP and A2A, then use service.start and + service.close as lifespan hooks. Retain the same event loop until shutdown + settles; configure owned pools with ReportingServiceResource when needed. + Reporting admission comes from the typed sync_accounts task, never configure. + """ + registry = ReportingProductionSourceRegistry(account_context=account_context) + for name, producer in producers.items(): + registry.register(name, producer) + support = ReportingProductionSupport( + materializer, + projection, + offerings=offerings, + configuration_task=configuration_task, + resolve_account=resolve_account, + source_registry=registry, + ) + service = ReliableReportingService.from_production(support) + service.install(application) + return service, support.handler diff --git a/pyproject.toml b/pyproject.toml index f869cdc9a..85869fdf9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -202,6 +202,7 @@ adcp = [ "reporting/feed/*.json", "reporting/projection/*.json", "reporting/production/*.json", + "reporting/production/*.sql", # PREVIEW: vendored sync_reporting_status schemas. They are the runtime # validator for the wire conditionals codegen cannot express, so the wheel # must carry them. Removed with the rest of _preview/ at rc.2. diff --git a/src/adcp/reporting/_inline_storage_schema.py b/src/adcp/reporting/_inline_storage_schema.py new file mode 100644 index 000000000..9653fa552 --- /dev/null +++ b/src/adcp/reporting/_inline_storage_schema.py @@ -0,0 +1,131 @@ +"""Required catalog objects from reporting_inline_storage.sql on PostgreSQL 16. + +Generated from a clean migration; unrelated adopter objects are permitted. +""" + +REQUIRED_OBJECTS: dict[str, dict[str, str | bool]] = { + "column:reporting_inline_objects.account_id": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "column:reporting_inline_objects.payload": { + "enabled": True, + "fingerprint": "554c34e416bd4546469b42d5773bc7c59b2104366ffa5259a2838c64cd826e57", + }, + "column:reporting_inline_objects.payload_sha256": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "column:reporting_inline_seals.account_id": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "column:reporting_inline_seals.byte_count": { + "enabled": True, + "fingerprint": "64be57437fdc0a07a97985c2aa058031f8082db7251bdb4d5afa1a9b088de97a", + }, + "column:reporting_inline_seals.manifest": { + "enabled": True, + "fingerprint": "554c34e416bd4546469b42d5773bc7c59b2104366ffa5259a2838c64cd826e57", + }, + "column:reporting_inline_seals.manifest_sha256": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "column:reporting_inline_seals.source_execution_key": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "column:reporting_inline_seals.staged_commit_ref": { + "enabled": True, + "fingerprint": "fdc4b549138bf90a5af5ed00b87a86ff9a0cbea553e7f017b4d95510708f6f88", + }, + "constraint:reporting_inline_objects.reporting_inline_objects_account": { + "enabled": True, + "fingerprint": "ff3e5aec84611ad080f664147be02d92ab106b3e2c9407bc63b5396bc0c8006f", + }, + "constraint:reporting_inline_objects.reporting_inline_objects_bytes": { + "enabled": True, + "fingerprint": "5cf3d75f67d8c5c6631ba87b9d6d301a4f807f9ce619bc77e3b81f43707d6669", + }, + "constraint:reporting_inline_objects.reporting_inline_objects_digest": { + "enabled": True, + "fingerprint": "3a303f02341c258a572f9eaf1e6869cdfd7547bc2932c44c41a8d1a7eed280a2", + }, + "constraint:reporting_inline_objects.reporting_inline_objects_pk": { + "enabled": True, + "fingerprint": "52c73c955e80fe2c8de482afc62851429635aa24f3a10cd4e15925eaa5a05bb3", + }, + "constraint:reporting_inline_objects.reporting_inline_objects_size": { + "enabled": True, + "fingerprint": "7c4fb6f6c0ff3c05be6ae3986c8a849e886e06e7088f5d09fef81b9f3b326ed8", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_account": { + "enabled": True, + "fingerprint": "ff3e5aec84611ad080f664147be02d92ab106b3e2c9407bc63b5396bc0c8006f", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_account_binding": { + "enabled": True, + "fingerprint": "ae8a06d397d61baba67cc4e77001596682a09ebecac0babaa0268c79ac7f242d", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_bytes": { + "enabled": True, + "fingerprint": "2cf54d2016b1d93e6cc2e8f6538ed87fc5b04b980048275210661926eeb49296", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_count": { + "enabled": True, + "fingerprint": "1e87e7574019dc01b1b39584ef6a420e82330e7089490b928f04f5bce4482f7c", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_digest": { + "enabled": True, + "fingerprint": "6d150c5fa94598b18c9adf65800d6ddd58479f97cb9b845b3d0f7a791aafced4", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_key": { + "enabled": True, + "fingerprint": "fadbcdf37ecabe857d9111775f689255637a0cef3028301593c285b69a0551a2", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_key_binding": { + "enabled": True, + "fingerprint": "897155db2cfdde3b4ed88ea3066c43b8938b9961e9bcca08eac4acbe71e5d50b", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_pk": { + "enabled": True, + "fingerprint": "123e2c394bd92bfab2862432ed6191484dd5438ca9b9f6372108b59c901b4248", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_ref": { + "enabled": True, + "fingerprint": "6ef63ffe9db91b8fa1fe7f75817e958e343cfbc0f8d9134a188fc65ce7d39a88", + }, + "constraint:reporting_inline_seals.reporting_inline_seals_size": { + "enabled": True, + "fingerprint": "b6c20825c79f2f4cbee21b439842ace6958549b83913f82d958de73980754ab2", + }, + "function:reporting_inline_immutable()": { + "enabled": True, + "fingerprint": "f2baed2e6158fc173155a71c5dedd3423b32d5923582600ae3c7062086c8ba41", + }, + "index:reporting_inline_objects.reporting_inline_objects_pk": { + "enabled": True, + "fingerprint": "ba41f90f12a63ad244a676e7d7ae37aacca579d32f6a1ced16a5676c420083d0", + }, + "index:reporting_inline_seals.reporting_inline_seals_pk": { + "enabled": True, + "fingerprint": "0816d86030df87c7d77ed0df4cec81b2a30e61d959fb53063750e4c2934623aa", + }, + "table:reporting_inline_objects": { + "enabled": True, + "fingerprint": "1f824779ff80f110344420b019786663d8c9beaad230da90e0439795e734ccda", + }, + "table:reporting_inline_seals": { + "enabled": True, + "fingerprint": "1f824779ff80f110344420b019786663d8c9beaad230da90e0439795e734ccda", + }, + "trigger:reporting_inline_objects.reporting_inline_objects_immutable": { + "enabled": True, + "fingerprint": "da630163871019e8457c2d54b41f23e7f182ad231e539d4f57360ea45c39a1c1", + }, + "trigger:reporting_inline_seals.reporting_inline_seals_immutable": { + "enabled": True, + "fingerprint": "201a920cf5c5343295735a14da5eee8def6fd80bd72826628ba3706ebc75d9d2", + }, +} diff --git a/src/adcp/reporting/_source_authorization.py b/src/adcp/reporting/_source_authorization.py new file mode 100644 index 000000000..e206fc236 --- /dev/null +++ b/src/adcp/reporting/_source_authorization.py @@ -0,0 +1,89 @@ +"""Turn-local source revocation and the SDK inline publication boundary. + +Only denials are remembered, for the remainder of a scheduling turn. Every +dispatch and publication calls the adopter again; no authorization is cached. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable, Iterator +from contextlib import AbstractAsyncContextManager, asynccontextmanager, contextmanager +from contextvars import ContextVar +from typing import TYPE_CHECKING, NoReturn, TypeAlias + +if TYPE_CHECKING: + from adcp.reporting.inline_source import ReportingSealStore, SealedSlice + +InlineSealPublisher: TypeAlias = Callable[[str, "SealedSlice"], Awaitable["SealedSlice"]] + +_REVOKED_ACCOUNTS: ContextVar[set[str] | None] = ContextVar( + "reporting_source_revoked_accounts", default=None +) +_INLINE_PUBLICATION: ContextVar[ + tuple[ + str, + Callable[[ReportingSealStore], AbstractAsyncContextManager[InlineSealPublisher | None]], + ] + | None +] = ContextVar("reporting_inline_publication", default=None) + + +@contextmanager +def source_turn() -> Iterator[None]: + """Share denials across configurations in this turn, never across turns.""" + if _REVOKED_ACCOUNTS.get() is not None: + yield + return + token = _REVOKED_ACCOUNTS.set(set()) + try: + yield + finally: + _REVOKED_ACCOUNTS.reset(token) + + +def account_revoked(account_id: str) -> bool: + return account_id in (_REVOKED_ACCOUNTS.get() or ()) + + +def source_revoked(account_id: str) -> NoReturn: + from adcp.reporting.production.contracts import _SourceAuthorizationRevokedError + + revoked = _REVOKED_ACCOUNTS.get() + if revoked is not None: + revoked.add(account_id) + raise _SourceAuthorizationRevokedError() + + +def require_account_work(account_id: str) -> None: + if account_revoked(account_id): + source_revoked(account_id) + + +@contextmanager +def bind_inline_publication( + account_id: str, + guard: Callable[[ReportingSealStore], AbstractAsyncContextManager[InlineSealPublisher | None]], +) -> Iterator[None]: + """Carry the producer's lock and live check into its inline executor task.""" + token = _INLINE_PUBLICATION.set((account_id, guard)) + try: + yield + finally: + _INLINE_PUBLICATION.reset(token) + + +@asynccontextmanager +async def inline_publication( + account_id: str, seals: ReportingSealStore +) -> AsyncIterator[InlineSealPublisher]: + async def publish(key: str, sealed: SealedSlice) -> SealedSlice: + return await seals.put(account_id=account_id, source_execution_key=key, sealed=sealed) + + bound = _INLINE_PUBLICATION.get() + if bound is None: + yield publish + return + if bound[0] != account_id: + raise ValueError("source publication must belong to the dispatched account") + async with bound[1](seals) as owned_publication: + yield owned_publication or publish diff --git a/src/adcp/reporting/inline_source.py b/src/adcp/reporting/inline_source.py index bc2aca2a0..3efc0160f 100644 --- a/src/adcp/reporting/inline_source.py +++ b/src/adcp/reporting/inline_source.py @@ -74,6 +74,7 @@ from pydantic import TypeAdapter, ValidationError from adcp.reporting._settlement import settle_task +from adcp.reporting._source_authorization import inline_publication from adcp.reporting.currency import ( ReportingCurrencyError, validate_currency, @@ -893,36 +894,36 @@ async def _publish( row_count=len(result.rows), ) - manifest = self._seal_manifest( - request, - observed_at=observed_at, - data_through=data_through, - staged=staged, - constituents=constituents, - cells=cells, - coverage_status=coverage_status, - explicit_zero=explicit_zero, - row_count=len(result.rows), - control_totals=control_totals, - warnings=warnings, - provisional_until=result.provisional_until, - ) + async with inline_publication(request.identity.account_id, self._seals) as publish_seal: + manifest = self._seal_manifest( + request, + observed_at=observed_at, + data_through=data_through, + staged=staged, + constituents=constituents, + cells=cells, + coverage_status=coverage_status, + explicit_zero=explicit_zero, + row_count=len(result.rows), + control_totals=control_totals, + warnings=warnings, + provisional_until=result.provisional_until, + ) - manifest_bytes = encode_source_batch_manifest_v1(manifest) - reference = source_batch_manifest_reference_v1( - f"{self._staged_commit_prefix}.{manifest.publication_id}", manifest_bytes - ) - winner = await self._seals.put( - account_id=request.identity.account_id, - source_execution_key=request.identity.source_execution_key, - sealed=SealedSlice(reference=reference, manifest_bytes=manifest_bytes), - ) - # A concurrent worker may have sealed first; its bytes are the - # publication, and returning ours instead would make the same key - # resolve two ways. - return ReportingSourceExecutorResult.completed( - request=request, manifest=winner.reference, manifest_bytes=winner.manifest_bytes - ) + manifest_bytes = encode_source_batch_manifest_v1(manifest) + reference = source_batch_manifest_reference_v1( + f"{self._staged_commit_prefix}.{manifest.publication_id}", manifest_bytes + ) + winner = await publish_seal( + request.identity.source_execution_key, + SealedSlice(reference=reference, manifest_bytes=manifest_bytes), + ) + # A concurrent worker may have sealed first; its bytes are the + # publication, and returning ours instead would make the same key + # resolve two ways. + return ReportingSourceExecutorResult.completed( + request=request, manifest=winner.reference, manifest_bytes=winner.manifest_bytes + ) def _derive_statuses( self, diff --git a/src/adcp/reporting/inline_storage.py b/src/adcp/reporting/inline_storage.py new file mode 100644 index 000000000..f32d15cfa --- /dev/null +++ b/src/adcp/reporting/inline_storage.py @@ -0,0 +1,593 @@ +"""Borrowed-pool PostgreSQL storage for inline objects and committed replay seals. + +Create either store with an open ``psycopg_pool.AsyncConnectionPool`` and call +``create_schema()`` before use; both stores share the same additive migration. +``check_ready()`` audits required catalog objects without changing the schema. +Operations also audit it within their short transaction. Pools, credentials, +timeouts, database durability, backups and retention remain adopter-owned. +Install ``adcp[pg]`` to construct a store. Importing this module needs no driver. +Catalog fingerprints are qualified on PostgreSQL 16. Other server majors are +unqualified; catalog formatting differences may require separate qualification. + +An authenticated caller supplies ``account_id``. A reference is not a credential; +``source_scope`` does not add authorization. Staging deduplicates exact bytes +within an account, independently of execution key or ordinal. Objects are capped +at 16 MiB (a constructor may lower that cap), manifests at the source contract's +1 MiB cap. No garbage collector or destructive retention operation is supplied. +Payload hashing yields between chunks of at most 1 MiB; it starts no threads. + +Staging and seal insertion are separate commits. A committed seal enables exact +replay after process restart, but a crash before the seal commits can refetch, +even if an object already committed. These stores do not reserve requests before +dispatch, recover a lost provider answer, bind unseen request/routing facts, own +source leases, or compose a production service. The producer's conformance check +still validates a replay against its complete frozen request. Seal validation +does not establish that referenced objects exist or share this storage backend. +A seal write has no lease token and does not fence a stale source owner. + +Private connection participants accept detached backend preparations and return +tentative references or a neutral ``SealedSlice``. Their caller owns the same-task +READ COMMITTED transaction, table locks and schema audit, account-lock ordering, +cancellation/error boundary, commit/rollback and connection return. They acquire +no connection, start no task/transaction, and provide no grant or commit proof. + +Cancellation settles the borrowed connection before propagating a fresh, redacted +``CancelledError``. Cancellation or a resource error may follow a committed write; +resume with the same account/key or content identity to discover its outcome. +An active AnyIO cancellation scope retains only a fixed framework marker so its +timeout/containment semantics survive; caller text and exception chains do not. +The read cancellation event is checked before and after the transaction; cancel +the task to interrupt a pending database operation. An object above a store's +lower configured read cap fails closed with ``INTEGRITY_FAILED``. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import re +from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Mapping +from contextlib import asynccontextmanager +from dataclasses import dataclass +from functools import wraps +from importlib.resources import files +from typing import TYPE_CHECKING, Any, ClassVar, Generic, Literal, ParamSpec, TypeVar + +from anyio import CancelScope, current_effective_deadline + +from adcp.reporting._inline_storage_schema import REQUIRED_OBJECTS +from adcp.reporting.inline_source import SealedSlice +from adcp.reporting.outbox._schema import schema_objects +from adcp.reporting.source import ( + SOURCE_BATCH_MANIFEST_MAX_BYTES_V1, + SourceBatchManifestReferenceV1, + parse_verified_source_batch_manifest_v1, +) + +if TYPE_CHECKING: + from psycopg_pool import AsyncConnectionPool + +__all__ = [ + "INLINE_STAGING_MAX_BYTES", + "InlineStorageError", + "PgReportingSealStore", + "PgReportingStagingStore", +] + +INLINE_STAGING_MAX_BYTES = 16_777_216 +_KEY = re.compile(r"[A-Za-z0-9_.:-]{8,255}") +_DIGEST = re.compile(r"[a-f0-9]{64}") +_Code = Literal[ + "INVALID_INPUT", "NOT_FOUND", "INTEGRITY_FAILED", "SCHEMA_UNREADY", "RESOURCE_UNAVAILABLE" +] +_MESSAGES: dict[_Code, str] = { + "INVALID_INPUT": "inline storage input is invalid", + "NOT_FOUND": "inline object is unavailable within the supplied account", + "INTEGRITY_FAILED": "inline storage integrity verification failed", + "SCHEMA_UNREADY": "inline storage schema is not ready", + "RESOURCE_UNAVAILABLE": ( + "inline storage operation did not confirm an outcome; resume the same identity" + ), +} + + +class InlineStorageError(RuntimeError): + """Closed diagnostic; resource failures can have an unknown commit outcome.""" + + def __init__(self, code: _Code) -> None: + self.code = code + super().__init__(_MESSAGES[code]) + + +P = ParamSpec("P") +T = TypeVar("T") + + +@dataclass(frozen=True, repr=False) +class _Completed(Generic[T]): + value: T + + +@dataclass(frozen=True, repr=False) +class _Failed: + # None represents cancellation; outcomes never contain exception objects. + code: _Code | None + + +def _cancellation_marker(error: asyncio.CancelledError) -> str | None: + # A caller-controlled prefix alone is not evidence of a cancelled scope. + if current_effective_deadline() != float("-inf"): + return None + seen: set[int] = set() + current: BaseException | None = error + while isinstance(current, asyncio.CancelledError) and len(seen) < 16: + if id(current) in seen: + break + seen.add(id(current)) + if current.args and type(current.args[0]) is str: + for prefix in ("Cancelled by cancel scope ", "Cancelled via cancel scope "): + if current.args[0].startswith(prefix): + return prefix + "[redacted]" + current = current.__context__ + return None + + +async def _payload_digest(payload: bytes) -> str: + digest = hashlib.sha256() + view = memoryview(payload) + chunk_bytes = 1_048_576 + for offset in range(0, len(view), chunk_bytes): + digest.update(view[offset : offset + chunk_bytes]) + if offset + chunk_bytes < len(view): + await asyncio.sleep(0) + return digest.hexdigest() + + +def _redact(method: Callable[P, Awaitable[T]]) -> Callable[P, Coroutine[Any, Any, T]]: + @wraps(method) + async def guarded(*args: P.args, **kwargs: P.kwargs) -> T: + async def run() -> _Completed[T] | _Failed: + # A raw exception crossing a Task boundary can remain active in + # Python 3.10's pure-Python Task wakeup frame. Return closed data so + # even that runner cannot attach a driver error to our public error. + try: + return _Completed(await method(*args, **kwargs)) + except InlineStorageError as error: + return _Failed(error.code) + except asyncio.CancelledError: + return _Failed(None) + except Exception: + return _Failed("RESOURCE_UNAVAILABLE") + + operation = asyncio.create_task(run()) + marker = None + try: + outcome = await asyncio.shield(operation) + except asyncio.CancelledError as error: + # Deliver cancellation once to the owned operation. AnyIO's repeated + # level cancellation must not interrupt the driver's query settlement + # and rollback. Repeated explicit Task.cancel() affects only this + # waiter; keep joining until the borrowed resource has been returned. + operation.cancel() + with CancelScope(shield=True): + while not operation.done(): + try: + await asyncio.shield(operation) + except asyncio.CancelledError: + continue + if not operation.cancelled(): + operation.result() + # Also leave an ambient Task wakeup exception when cancellation + # races an already-completed operation (no join await needed). + while True: + try: + await asyncio.sleep(0) + break + except asyncio.CancelledError: + continue + marker = _cancellation_marker(error) + outcome = _Failed(None) + # Raise outside the handler: raw database/input exception contexts must + # not survive even for callers that inspect __context__ directly. + if isinstance(outcome, _Completed): + return outcome.value + if outcome.code is None: + if marker is not None: + raise asyncio.CancelledError(marker) + raise asyncio.CancelledError + raise InlineStorageError(outcome.code) + + return guarded + + +def _account(value: str) -> None: + if type(value) is not str or not 1 <= len(value) <= 255 or "\x00" in value: + raise InlineStorageError("INVALID_INPUT") + # PostgreSQL text requires valid Unicode scalar values. + try: + value.encode("utf-8") + except UnicodeError: + raise InlineStorageError("INVALID_INPUT") from None + + +def _identity(account_id: str, key: str) -> None: + _account(account_id) + if type(key) is not str or _KEY.fullmatch(key) is None: + raise InlineStorageError("INVALID_INPUT") + + +def _reference(account_id: str, digest: str) -> str: + scope = hashlib.sha256(account_id.encode("utf-8")).hexdigest() + return f"pg-inline-v1.{scope}.{digest}" + + +def _prepared_identity(account_id: str, key: str) -> None: + invalid = False + try: + _identity(account_id, key) + except InlineStorageError: + invalid = True + if invalid: + raise InlineStorageError("INVALID_INPUT") + + +def _object_inputs( + account_id: str, key: str, ordinal: int, payload: bytes, max_payload_bytes: int +) -> None: + _prepared_identity(account_id, key) + if type(ordinal) is not int or not 0 <= ordinal < 100_000: + raise InlineStorageError("INVALID_INPUT") + if type(payload) is not bytes or len(payload) > max_payload_bytes: + raise InlineStorageError("INVALID_INPUT") + + +@dataclass(frozen=True, slots=True, repr=False) +class PreparedBackendObjectV1: + """Private immutable inputs; the claimed digest still needs stored-byte verification.""" + + account_id: str + source_execution_key: str + ordinal: int + payload: bytes + payload_sha256: str + + def __post_init__(self) -> None: + self._validate(INLINE_STAGING_MAX_BYTES) + + def _validate(self, max_payload_bytes: int) -> None: + _object_inputs( + self.account_id, + self.source_execution_key, + self.ordinal, + self.payload, + max_payload_bytes, + ) + if type(self.payload_sha256) is not str or _DIGEST.fullmatch(self.payload_sha256) is None: + raise InlineStorageError("INVALID_INPUT") + + def __repr__(self) -> str: + return "PreparedBackendObjectV1()" + + +@dataclass(frozen=True, slots=True, repr=False) +class PreparedBackendSealV1: + """Private detached scalars/bytes, not admission provenance or a committed seal.""" + + account_id: str + source_execution_key: str + staged_commit_ref: str + manifest_sha256: str + byte_count: int + manifest_bytes: bytes + + def __post_init__(self) -> None: + self._validate() + + def _validate(self) -> None: + _prepared_identity(self.account_id, self.source_execution_key) + if ( + type(self.staged_commit_ref) is not str + or not 1 <= len(self.staged_commit_ref) <= 255 + or type(self.manifest_sha256) is not str + or _DIGEST.fullmatch(self.manifest_sha256) is None + or type(self.byte_count) is not int + or not 1 <= self.byte_count <= SOURCE_BATCH_MANIFEST_MAX_BYTES_V1 + or type(self.manifest_bytes) is not bytes + or len(self.manifest_bytes) != self.byte_count + ): + raise InlineStorageError("INVALID_INPUT") + + def __repr__(self) -> str: + return "PreparedBackendSealV1()" + + +@dataclass(frozen=True, repr=False) +class _StoredSeal(SealedSlice): + def __repr__(self) -> str: + return "SealedSlice()" + + def __eq__(self, other: object) -> bool: + if not isinstance(other, SealedSlice): + return NotImplemented + return self.reference == other.reference and self.manifest_bytes == other.manifest_bytes + + +def _seal(account_id: str, key: str, sealed: SealedSlice, code: _Code) -> SealedSlice: + try: + if not isinstance(sealed, SealedSlice) or type(sealed.manifest_bytes) is not bytes: + raise ValueError + # Revalidate even model_construct/model_copy inputs; never normalize the + # retained bytes or trust the caller's mutable instance. + reference = SourceBatchManifestReferenceV1.model_validate(sealed.reference.model_dump()) + raw = sealed.manifest_bytes + manifest = parse_verified_source_batch_manifest_v1(reference, raw) + if ( + manifest.identity.account_id != account_id + or manifest.identity.source_execution_key != key + ): + raise ValueError + return _StoredSeal(reference=reference, manifest_bytes=raw) + except Exception: + failure = InlineStorageError(code) + raise failure + + +class _PgStorage: + is_durable: ClassVar[bool] = True + + def __init__(self, *, pool: AsyncConnectionPool) -> None: + try: + from psycopg_pool import AsyncConnectionPool as Pool + except ImportError: + pass + else: + if not isinstance(pool, Pool): + raise InlineStorageError("INVALID_INPUT") + self._pool = pool + return + raise ImportError("PostgreSQL inline storage requires adcp[pg]") + + def __repr__(self) -> str: + return f"{type(self).__name__}()" + + async def _ready_on(self, connection: Any) -> None: + installed = await schema_objects(connection) + if not REQUIRED_OBJECTS or any( + installed.get(key, {}).get(field) != value[field] + for key, value in REQUIRED_OBJECTS.items() + for field in ("fingerprint", "enabled") + ): + raise InlineStorageError("SCHEMA_UNREADY") + + @asynccontextmanager + async def _transaction(self) -> AsyncIterator[Any]: + from psycopg.errors import UndefinedTable + + async with self._pool.connection() as connection, connection.transaction(): + # DO NOTHING's conflict winner must be visible to the next SELECT, + # regardless of a pool's default isolation level. + await connection.execute("SET TRANSACTION ISOLATION LEVEL READ COMMITTED") + try: + await connection.execute( + "LOCK TABLE reporting_inline_objects, reporting_inline_seals " + "IN ACCESS SHARE MODE" + ) + except UndefinedTable: + raise InlineStorageError("SCHEMA_UNREADY") from None + await self._ready_on(connection) + yield connection + + @_redact + async def create_schema(self) -> None: + """Install the standalone additive migration and verify it atomically.""" + async with self._pool.connection() as connection, connection.transaction(): + await connection.execute("SET TRANSACTION ISOLATION LEVEL READ COMMITTED") + await connection.execute( + files("adcp.reporting.ledger").joinpath("reporting_inline_storage.sql").read_text() + ) + await self._ready_on(connection) + + @_redact + async def check_ready(self) -> None: + """Verify required schema objects, including constraints and write guards.""" + async with self._pool.connection() as connection, connection.transaction(): + await self._ready_on(connection) + + +class PgReportingStagingStore(_PgStorage): + """Immutable account-qualified content storage using a borrowed async pool.""" + + def __init__( + self, *, pool: AsyncConnectionPool, max_payload_bytes: int = INLINE_STAGING_MAX_BYTES + ) -> None: + if ( + type(max_payload_bytes) is not int + or not 1 <= max_payload_bytes <= INLINE_STAGING_MAX_BYTES + ): + raise InlineStorageError("INVALID_INPUT") + super().__init__(pool=pool) + self._max_payload_bytes = max_payload_bytes + + async def _read_on(self, connection: Any, account_id: str, digest: str) -> bytes: + row = await ( + await connection.execute( + "SELECT CASE WHEN octet_length(payload) <= %s THEN payload END" + " FROM reporting_inline_objects WHERE account_id = %s AND payload_sha256 = %s", + (self._max_payload_bytes, account_id, digest), + ) + ).fetchone() + if row is None: + raise InlineStorageError("NOT_FOUND") + payload = row[0] + if type(payload) is not bytes or await _payload_digest(payload) != digest: + raise InlineStorageError("INTEGRITY_FAILED") + return payload + + async def _prepare_object( + self, *, account_id: str, source_execution_key: str, ordinal: int, payload: bytes + ) -> PreparedBackendObjectV1: + _object_inputs(account_id, source_execution_key, ordinal, payload, self._max_payload_bytes) + digest = await _payload_digest(payload) + return PreparedBackendObjectV1(account_id, source_execution_key, ordinal, payload, digest) + + async def _stage_on( + self, connection: Any, prepared: PreparedBackendObjectV1 + ) -> tuple[str, str]: + """Stage on the caller's audited transaction; the returned identity is tentative. + + Caller owns account-lock ordering and the full transaction/error boundary. + No checkout, task, transaction, commit, drain or fence is owned here. + """ + if type(prepared) is not PreparedBackendObjectV1: + raise InlineStorageError("INVALID_INPUT") + prepared._validate(self._max_payload_bytes) + await connection.execute( + "INSERT INTO reporting_inline_objects (account_id, payload_sha256, payload)" + " VALUES (%s, %s, %s) ON CONFLICT (account_id, payload_sha256) DO NOTHING", + (prepared.account_id, prepared.payload_sha256, prepared.payload), + ) + if ( + await self._read_on(connection, prepared.account_id, prepared.payload_sha256) + != prepared.payload + ): + raise InlineStorageError("INTEGRITY_FAILED") + return _reference(prepared.account_id, prepared.payload_sha256), prepared.payload_sha256 + + @_redact + async def stage( + self, *, account_id: str, source_execution_key: str, ordinal: int, payload: bytes + ) -> tuple[str, str]: + prepared = await self._prepare_object( + account_id=account_id, + source_execution_key=source_execution_key, + ordinal=ordinal, + payload=payload, + ) + async with self._transaction() as connection: + return await self._stage_on(connection, prepared) + + @_redact + async def read( + self, + *, + object_ref: str, + object_generation: str, + account_id: str, + source_scope: Mapping[str, Any], + cancel: asyncio.Event, + ) -> bytes: + _account(account_id) + if type(object_generation) is not str or _DIGEST.fullmatch(object_generation) is None: + raise InlineStorageError("INVALID_INPUT") + if type(object_ref) is not str or object_ref != _reference(account_id, object_generation): + raise InlineStorageError("NOT_FOUND") + if cancel.is_set(): + raise asyncio.CancelledError + async with self._transaction() as connection: + payload = await self._read_on(connection, account_id, object_generation) + if cancel.is_set(): + raise asyncio.CancelledError + return payload + + +class PgReportingSealStore(_PgStorage): + """First committed seal wins for each trusted account and execution key.""" + + async def _get_on(self, connection: Any, account_id: str, key: str) -> SealedSlice | None: + row = await ( + await connection.execute( + "SELECT staged_commit_ref, manifest_sha256, byte_count," + " CASE WHEN octet_length(manifest) <= %s THEN manifest END" + " FROM reporting_inline_seals WHERE account_id = %s AND source_execution_key = %s", + (SOURCE_BATCH_MANIFEST_MAX_BYTES_V1, account_id, key), + ) + ).fetchone() + if row is None: + return None + reference: SourceBatchManifestReferenceV1 | None + try: + reference = SourceBatchManifestReferenceV1( + staged_commit_ref=row[0], manifest_sha256=row[1], byte_count=row[2] + ) + except Exception: + reference = None + if reference is None: + raise InlineStorageError("INTEGRITY_FAILED") + return _seal( + account_id, + key, + SealedSlice(reference=reference, manifest_bytes=row[3]), + "INTEGRITY_FAILED", + ) + + @_redact + async def get(self, *, account_id: str, source_execution_key: str) -> SealedSlice | None: + _identity(account_id, source_execution_key) + async with self._transaction() as connection: + return await self._get_on(connection, account_id, source_execution_key) + + def _prepare_seal( + self, *, account_id: str, source_execution_key: str, sealed: SealedSlice + ) -> PreparedBackendSealV1: + _prepared_identity(account_id, source_execution_key) + candidate = _seal(account_id, source_execution_key, sealed, "INVALID_INPUT") + reference = candidate.reference + return PreparedBackendSealV1( + account_id, + source_execution_key, + reference.staged_commit_ref, + reference.manifest_sha256, + reference.byte_count, + candidate.manifest_bytes, + ) + + async def _put_on(self, connection: Any, prepared: PreparedBackendSealV1) -> SealedSlice: + """Return the verified stored winner on the caller's audited transaction. + + The neutral value proves no admission, referenced-object existence or + commit. Caller owns account locks, cancellation, settlement and any fence. + """ + if type(prepared) is not PreparedBackendSealV1: + raise InlineStorageError("INVALID_INPUT") + prepared._validate() + # Build a fresh reference, then revalidate it inside the existing closed + # candidate validator. No mutable caller model is retained or trusted. + reference = SourceBatchManifestReferenceV1.model_construct( + staged_commit_ref=prepared.staged_commit_ref, + manifest_sha256=prepared.manifest_sha256, + byte_count=prepared.byte_count, + ) + candidate = _seal( + prepared.account_id, + prepared.source_execution_key, + SealedSlice(reference=reference, manifest_bytes=prepared.manifest_bytes), + "INVALID_INPUT", + ) + await connection.execute( + "INSERT INTO reporting_inline_seals (account_id, source_execution_key," + " staged_commit_ref, manifest_sha256, byte_count, manifest)" + " VALUES (%s, %s, %s, %s, %s, %s)" + " ON CONFLICT (account_id, source_execution_key) DO NOTHING", + ( + prepared.account_id, + prepared.source_execution_key, + candidate.reference.staged_commit_ref, + candidate.reference.manifest_sha256, + candidate.reference.byte_count, + candidate.manifest_bytes, + ), + ) + winner = await self._get_on(connection, prepared.account_id, prepared.source_execution_key) + if winner is None: + raise InlineStorageError("INTEGRITY_FAILED") + return winner + + @_redact + async def put( + self, *, account_id: str, source_execution_key: str, sealed: SealedSlice + ) -> SealedSlice: + prepared = self._prepare_seal( + account_id=account_id, source_execution_key=source_execution_key, sealed=sealed + ) + async with self._transaction() as connection: + return await self._put_on(connection, prepared) diff --git a/src/adcp/reporting/ledger/pg.py b/src/adcp/reporting/ledger/pg.py index cf49a78cb..e4404139a 100644 --- a/src/adcp/reporting/ledger/pg.py +++ b/src/adcp/reporting/ledger/pg.py @@ -139,6 +139,8 @@ if TYPE_CHECKING: from psycopg_pool import AsyncConnectionPool + from adcp.reporting._source_authorization import InlineSealPublisher + from adcp.reporting.inline_source import ReportingSealStore, SealedSlice from adcp.reporting.ledger.status_projection import ReportingStatusSnapshot try: @@ -226,6 +228,46 @@ async def transaction(self) -> AsyncIterator[PgReportingLedgerStore]: finally: _BOUND_CONNECTION.reset(token) + @asynccontextmanager + async def _source_publication( + self, account_id: str, *, seals: ReportingSealStore | None = None + ) -> AsyncIterator[InlineSealPublisher | None]: + from psycopg import Error + + from adcp.reporting.inline_storage import InlineStorageError, PgReportingSealStore + + if isinstance(seals, PgReportingSealStore) and seals._pool is self._pool: + # The backend's audited READ COMMITTED transaction owns both the + # account lock and seal write. A second pool checkout would deadlock + # with a size-one pool, and would separate the publication boundary. + failure = None + try: + async with seals._transaction() as connection: + await self._lock_account(connection, account_id) + owner = asyncio.current_task() + + async def publish(key: str, sealed: SealedSlice) -> SealedSlice: + if asyncio.current_task() is not owner: + raise InlineStorageError("INVALID_INPUT") + prepared = seals._prepare_seal( + account_id=account_id, source_execution_key=key, sealed=sealed + ) + return await seals._put_on(connection, prepared) + + yield publish + except InlineStorageError as error: + failure = error.code + except Error: + failure = "RESOURCE_UNAVAILABLE" + if failure is not None: + raise InlineStorageError(failure) + return + # Keep the live authorization check and ledger commit on the same + # account transaction, including commits replayed after a restart. + async with self.transaction(), self._connection() as connection: + await self._lock_account(connection, account_id) + yield None + async def create_schema(self) -> None: """Create or upgrade the ledger atomically, serializing concurrent boots. @@ -1002,7 +1044,9 @@ async def commit_provisional_observation( raise LedgerConflictError( "HISTORY_UNAVAILABLE", "observation revision is missing" ) - return retained_revision + # Reuse the immutable publication replay checks while holding + # the observation's account lock and transaction. + return await self.commit_revision(revision, rows) reserved = await ( await connection.execute( "SELECT payload FROM reporting_provisional_acquisitions" diff --git a/src/adcp/reporting/ledger/producer.py b/src/adcp/reporting/ledger/producer.py index f0ec5044b..86c7dee23 100644 --- a/src/adcp/reporting/ledger/producer.py +++ b/src/adcp/reporting/ledger/producer.py @@ -32,12 +32,19 @@ import hashlib import inspect import logging -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence +from contextlib import asynccontextmanager from dataclasses import dataclass, field, replace from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, TypeAlias from adcp.reporting._settlement import cancel_and_settle +from adcp.reporting._source_authorization import ( + bind_inline_publication, + require_account_work, + source_revoked, + source_turn, +) from adcp.reporting.canonical_json import canonical_json_utf8_v1 from adcp.reporting.currency import ( ReportingCurrencyError, @@ -86,6 +93,9 @@ ) if TYPE_CHECKING: + from adcp.reporting._source_authorization import InlineSealPublisher + from adcp.reporting.inline_source import ReportingSealStore + from adcp.reporting.inline_storage import _Code as InlineStorageErrorCode from adcp.reporting.ledger.producer_progress import ReportingProducerProgress from adcp.reporting.materializer.verification import ReportingRevisionVerifier @@ -101,6 +111,11 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class _InlineStorageFailure: + code: InlineStorageErrorCode + + CurrencyResolver: TypeAlias = Callable[ [ReportingConfiguration, ReportingObligationRecord], Awaitable[str] | str ] @@ -354,6 +369,12 @@ async def run_worker(self) -> WorkerTurn: Returns immediately with an empty turn when nothing is leasable, so a caller can back off rather than spin. """ + with source_turn(): + return await self._run_leased_turn() + + async def _run_leased_turn(self) -> WorkerTurn: + from adcp.reporting.production.contracts import _SourceAuthorizationRevokedError + now = self._clock() leased = await self._store.lease_period_close( worker_id=self._worker_id, now=now, lease_seconds=self._lease_seconds @@ -375,8 +396,13 @@ async def run_worker(self) -> WorkerTurn: ) if configuration is None: return turn + require_account_work(configuration.account_id) await self._close_elapsed_periods(configuration, turn, now=now) await self._acquire_pending(configuration, turn, now=now) + except _SourceAuthorizationRevokedError: + # Keep pending acquisitions retryable. Authorization may be + # restored on the next turn; it is not a history failure. + return turn finally: await self._store.release_period_close(leased, worker_id=self._worker_id) return turn @@ -397,10 +423,17 @@ async def run_configuration( Most adopters should continue using :meth:`run_worker`, whose store lease chooses a configuration automatically. """ + from adcp.reporting.production.contracts import _SourceAuthorizationRevokedError + boundary = now or self._clock() turn = WorkerTurn() - await self._close_elapsed_periods(configuration, turn, now=boundary) - await self._acquire_pending(configuration, turn, now=boundary) + with source_turn(): + try: + require_account_work(configuration.account_id) + await self._close_elapsed_periods(configuration, turn, now=boundary) + await self._acquire_pending(configuration, turn, now=boundary) + except _SourceAuthorizationRevokedError: + return turn return turn # -- step 1: obligations before reports ------------------------------ @@ -838,13 +871,7 @@ async def acquire_obligation( f"this producer declares no source offering for {finality} reporting", ) - from adcp.reporting.ledger.producer_progress import ReportingProducerProgress - - constituents = ( - await self._store.producer_constituents(configuration, obligation) - if isinstance(self._store, ReportingProducerProgress) - else None - ) + constituents = await self._admitted_source_constituents(configuration, obligation) checkpoint_store = self._restatement_store() if track_settling else None checkpoint = ( @@ -895,7 +922,11 @@ async def acquire_obligation( "snapshot" if request.publication_class == "PROVISIONAL_SNAPSHOT" else "official" ) cancel = asyncio.Event() - execution = asyncio.create_task(self._source.execute(request, cancel=cancel)) + execution = asyncio.create_task( + self._execute_source( + configuration, request, admitted_constituents=constituents, cancel=cancel + ) + ) try: result = await asyncio.wait_for( asyncio.shield(execution), @@ -915,6 +946,13 @@ async def acquire_obligation( self._note_escalation(obligation, turn, now=now) return None + if isinstance(result, _InlineStorageFailure): + from adcp.reporting.inline_storage import InlineStorageError + + # Return closed data from the executor task: a raw driver exception + # must not survive through a context manager or Task wakeup frame. + raise InlineStorageError(result.code) + if not result.ok: error = result.error assert error is not None @@ -933,6 +971,8 @@ async def acquire_obligation( rows = await self._read_rows(request, manifest) # ``now`` freezes dispatch/lease/cutoff decisions, not publication. # A conforming source can observe finality while acquisition is running. + # Every successful scheduled read, even unchanged content, reaches the + # atomic revision/observation/checkpoint commit with this fresh anchor. published_at = self._clock() if _utc(published_at) < _utc(now): raise LedgerConflictError( @@ -949,6 +989,132 @@ async def acquire_obligation( acquisition=acquisition, ) + async def _admitted_source_constituents( + self, configuration: ReportingConfiguration, obligation: ReportingObligationRecord + ) -> tuple[ReportingConstituent, ...] | None: + from adcp.reporting.ledger.producer_progress import ReportingProducerProgress + from adcp.reporting.production.contracts import _SourceAuthorizationRevokedError + + if not isinstance(self._store, ReportingProducerProgress): + return None + try: + return await self._store.producer_constituents(configuration, obligation) + except _SourceAuthorizationRevokedError: + source_revoked(configuration.account_id) + + def _check_source_authorization( + self, + configuration: ReportingConfiguration, + offering_id: str, + admitted_constituents: tuple[ReportingConstituent, ...] | None, + ) -> None: + from adcp.reporting.materializer.contracts import failure + from adcp.reporting.production.contracts import ( + ReportingProductionSource, + ReportingProductionSourceBinding, + _SourceAuthorizationRevokedError, + ) + + require_account_work(configuration.account_id) + if isinstance(self._source, ReportingProductionSource): + try: + binding = self._source.configuration_binding(configuration) + if binding is None: + source_revoked(configuration.account_id) + if type(binding) is not ReportingProductionSourceBinding: + raise failure("BINDING_MISMATCH") + binding.check(configuration, self._source.capabilities, offering_id) + if ( + admitted_constituents is not None + and binding.constituents() != admitted_constituents + ): + raise failure("BINDING_MISMATCH") + except _SourceAuthorizationRevokedError: + raise + except Exception: + raise failure("BINDING_MISMATCH") from None + + @asynccontextmanager + async def _source_publication( + self, + configuration: ReportingConfiguration, + offering_id: str, + *, + admitted_constituents: tuple[ReportingConstituent, ...] | None, + seals: ReportingSealStore | None = None, + ) -> AsyncIterator[InlineSealPublisher | None]: + from adcp.reporting.production.contracts import ReportingProductionSource + + if not isinstance(self._source, ReportingProductionSource): + yield None + return + publication = getattr(self._store, "_source_publication", None) + if publication is None: + raise TypeError("production source publication requires an SDK account lock") + async with publication(configuration.account_id, seals=seals) as publish_seal: + self._check_source_authorization(configuration, offering_id, admitted_constituents) + yield publish_seal + + async def _execute_source( + self, + configuration: ReportingConfiguration, + request: ReportingSourceSliceRequestV1, + *, + admitted_constituents: tuple[ReportingConstituent, ...] | None, + cancel: asyncio.Event, + ) -> ReportingSourceExecutorResult | _InlineStorageFailure: + from adcp.reporting.inline_storage import InlineStorageError + + try: + with bind_inline_publication( + configuration.account_id, + lambda seals: self._source_publication( + configuration, + request.offering_id, + admitted_constituents=admitted_constituents, + seals=seals, + ), + ): + # Check inside the executing task: scheduling it is not dispatch. + # Revocation never cancels a fetch that has already started. + self._check_source_authorization( + configuration, request.offering_id, admitted_constituents + ) + return await self._source.execute(request, cancel=cancel) + except InlineStorageError as error: + return _InlineStorageFailure(error.code) + + @asynccontextmanager + async def _revision_publication( + self, obligation: ReportingObligationRecord, offering_id: str + ) -> AsyncIterator[None]: + from adcp.reporting.production.contracts import ReportingProductionSource + + if not isinstance(self._source, ReportingProductionSource): + yield + return + configuration = next( + ( + candidate + for candidate in await self._store.list_configurations( + account_id=obligation.account_id, + delivery_config_ids=[obligation.delivery_config_id], + ) + if candidate.generation_key == obligation.generation_key + ), + None, + ) + if configuration is None: + raise LedgerConflictError("HISTORY_UNAVAILABLE", "source generation is unavailable") + # Read immutable admission scope before taking the publication lock. It + # is not an authorization grant: the callback is checked again inside + # the lock against this exact mapping, including on replay. + constituents = await self._admitted_source_constituents(configuration, obligation) + async with self._source_publication( + configuration, offering_id, admitted_constituents=constituents + ): + yield + def _restatement_store(self) -> RestatementCheckpointStore: explicit = all( inspect.getattr_static(self._store, name, None) is not None @@ -1088,6 +1254,8 @@ async def commit_revision_from_manifest( if prior is not None: # Still reconstruct and verify the supplied content below. Merely # finding the ID must not bypass immutable-content validation. + # On replay this retained parent wins over the acquisition's + # predecessor, just as the retained creation time wins over now. supersedes = prior.supersedes_reporting_revision_id if ( _utc(manifest.acquired_at) > _utc(now) @@ -1138,34 +1306,37 @@ async def commit_revision_from_manifest( source_publication_id=manifest.publication_id, source_manifest_sha256=manifest.content_fingerprint.split(":", 1)[-1], ) - if self._revision_verifier is not None: - from adcp.reporting.materializer.publication import verified_publication - - revision = verified_publication(self._revision_verifier, obligation, revision, rows) - if acquisition is None: - committed = await self._store.commit_revision(revision, rows) - else: - checked_at = max(_utc(now), _utc(manifest.acquired_at)) - boundary = manifest.finality_evidence.provisional_until or ( - _utc(obligation.period.end) + acquisition.policy.window - ) - next_due = ( - min(checked_at + acquisition.policy.cadence, _utc(boundary)) - if finality == "snapshot" and checked_at < _utc(boundary) - else None - ) - committed = await self._observation_store().commit_provisional_observation( - ProvisionalObservation( - acquisition, - revision_id, - checked_at, - _utc(boundary), - next_due, - manifest.model_dump_json(), - ), - revision, - rows, - ) + async with self._revision_publication(obligation, manifest.offering_id): + if self._revision_verifier is not None: + from adcp.reporting.materializer.publication import verified_publication + + revision = verified_publication(self._revision_verifier, obligation, revision, rows) + if acquisition is None: + committed = await self._store.commit_revision(revision, rows) + else: + # ``now`` is the publication anchor, sampled after row acquisition. + # The bounds check above already rejects acquired_at after now. + checked_at = _utc(now) + boundary = manifest.finality_evidence.provisional_until or ( + _utc(obligation.period.end) + acquisition.policy.window + ) + next_due = ( + min(checked_at + acquisition.policy.cadence, _utc(boundary)) + if finality == "snapshot" and checked_at < _utc(boundary) + else None + ) + committed = await self._observation_store().commit_provisional_observation( + ProvisionalObservation( + acquisition, + revision_id, + checked_at, + _utc(boundary), + next_due, + manifest.model_dump_json(), + ), + revision, + rows, + ) turn.revisions_committed.append(committed.reporting_revision_id) return committed diff --git a/src/adcp/reporting/ledger/reporting_inline_storage.sql b/src/adcp/reporting/ledger/reporting_inline_storage.sql new file mode 100644 index 000000000..175f42312 --- /dev/null +++ b/src/adcp/reporting/ledger/reporting_inline_storage.sql @@ -0,0 +1,60 @@ +-- Standalone, additive storage for inline-source objects and committed replay seals. +-- The caller must execute this migration and readiness validation in one transaction. +DO $migration$ +BEGIN + PERFORM pg_advisory_xact_lock(hashtext('adcp.reporting.inline_storage.schema')); + + CREATE TABLE IF NOT EXISTS reporting_inline_objects ( + account_id TEXT COLLATE "C" NOT NULL, + payload_sha256 TEXT COLLATE "C" NOT NULL, + payload BYTEA NOT NULL, + CONSTRAINT reporting_inline_objects_pk PRIMARY KEY (account_id, payload_sha256), + CONSTRAINT reporting_inline_objects_account CHECK (char_length(account_id) BETWEEN 1 AND 255), + CONSTRAINT reporting_inline_objects_digest CHECK (payload_sha256 ~ '^[a-f0-9]{64}$'), + CONSTRAINT reporting_inline_objects_size CHECK (octet_length(payload) <= 16777216), + CONSTRAINT reporting_inline_objects_bytes CHECK (payload_sha256 = encode(sha256(payload), 'hex')) + ); + + CREATE TABLE IF NOT EXISTS reporting_inline_seals ( + account_id TEXT COLLATE "C" NOT NULL, + source_execution_key TEXT COLLATE "C" NOT NULL, + staged_commit_ref TEXT COLLATE "C" NOT NULL, + manifest_sha256 TEXT COLLATE "C" NOT NULL, + byte_count INTEGER NOT NULL, + manifest BYTEA NOT NULL, + CONSTRAINT reporting_inline_seals_pk PRIMARY KEY (account_id, source_execution_key), + CONSTRAINT reporting_inline_seals_account CHECK (char_length(account_id) BETWEEN 1 AND 255), + CONSTRAINT reporting_inline_seals_key CHECK (source_execution_key ~ '^[A-Za-z0-9_.:-]{8,255}$'), + CONSTRAINT reporting_inline_seals_ref CHECK (staged_commit_ref ~ '^[A-Za-z0-9][A-Za-z0-9_.-]{0,254}$'), + CONSTRAINT reporting_inline_seals_digest CHECK (manifest_sha256 ~ '^[a-f0-9]{64}$'), + CONSTRAINT reporting_inline_seals_size CHECK (byte_count BETWEEN 1 AND 1048576), + CONSTRAINT reporting_inline_seals_count CHECK (byte_count = octet_length(manifest)), + CONSTRAINT reporting_inline_seals_bytes CHECK (manifest_sha256 = encode(sha256(manifest), 'hex')), + CONSTRAINT reporting_inline_seals_account_binding CHECK ( + (convert_from(manifest, 'UTF8')::jsonb #>> '{identity,account_id}') IS NOT DISTINCT FROM account_id), + CONSTRAINT reporting_inline_seals_key_binding CHECK ( + (convert_from(manifest, 'UTF8')::jsonb #>> '{identity,source_execution_key}') IS NOT DISTINCT FROM source_execution_key) + ); + + IF to_regprocedure('reporting_inline_immutable()') IS NULL THEN + CREATE FUNCTION reporting_inline_immutable() RETURNS TRIGGER LANGUAGE plpgsql AS $function$ + BEGIN + RAISE EXCEPTION USING ERRCODE = '23514', MESSAGE = 'reporting inline storage is immutable'; + END + $function$; + END IF; + + IF NOT EXISTS (SELECT FROM pg_trigger WHERE tgrelid = 'reporting_inline_objects'::regclass + AND tgname = 'reporting_inline_objects_immutable') THEN + CREATE TRIGGER reporting_inline_objects_immutable + BEFORE UPDATE OR DELETE OR TRUNCATE ON reporting_inline_objects + FOR EACH STATEMENT EXECUTE FUNCTION reporting_inline_immutable(); + END IF; + IF NOT EXISTS (SELECT FROM pg_trigger WHERE tgrelid = 'reporting_inline_seals'::regclass + AND tgname = 'reporting_inline_seals_immutable') THEN + CREATE TRIGGER reporting_inline_seals_immutable + BEFORE UPDATE OR DELETE OR TRUNCATE ON reporting_inline_seals + FOR EACH STATEMENT EXECUTE FUNCTION reporting_inline_immutable(); + END IF; +END +$migration$; diff --git a/src/adcp/reporting/ledger/reporting_provisional_observations.sql b/src/adcp/reporting/ledger/reporting_provisional_observations.sql index 6622ea718..5206ffc65 100644 --- a/src/adcp/reporting/ledger/reporting_provisional_observations.sql +++ b/src/adcp/reporting/ledger/reporting_provisional_observations.sql @@ -1,5 +1,8 @@ -- Additive private scheduling metadata. Original revision rows and manifests -- remain readable by historical SDKs; no existing payload is rewritten. +-- Apply after reporting_ledger_reconciliation.sql installs the prerequisite +-- unique indexes reporting_obligations_account_identity and +-- reporting_revisions_account_identity for the account-bound foreign keys. CREATE TABLE IF NOT EXISTS reporting_provisional_acquisitions ( account_id TEXT COLLATE "C" NOT NULL, reporting_obligation_id TEXT COLLATE "C" NOT NULL, diff --git a/src/adcp/reporting/ledger/store.py b/src/adcp/reporting/ledger/store.py index 8155a951f..5cd6358fe 100644 --- a/src/adcp/reporting/ledger/store.py +++ b/src/adcp/reporting/ledger/store.py @@ -77,6 +77,7 @@ from adcp.reporting.ledger.provisional import ProvisionalAcquisition, ProvisionalObservation if TYPE_CHECKING: + from adcp.reporting.inline_source import ReportingSealStore from adcp.reporting.ledger.status_projection import ReportingStatusSnapshot from adcp.reporting.outbox.memory import NotificationState @@ -773,6 +774,14 @@ async def transaction(self) -> AsyncIterator[InMemoryReportingLedgerStore]: async with self._mutation(): yield self + @asynccontextmanager + async def _source_publication( + self, account_id: str, *, seals: ReportingSealStore | None = None + ) -> AsyncIterator[None]: + # The memory mutation lock serializes all accounts, including seals. + async with self._mutation(): + yield + def _record_notification(self, event: ReportingDomainEvent) -> None: if self._notification_state is not None: self._notification_state.enqueue(event) @@ -1135,7 +1144,13 @@ async def commit_provisional_observation( or existing.revision_id != observation.revision_id ): raise LedgerConflictError("OBSERVATION_CONFLICT", "observation replay differs") - return self._revisions[existing.revision_id] + if existing.revision_id not in self._revisions: + raise LedgerConflictError( + "HISTORY_UNAVAILABLE", "observation revision is missing" + ) + # A matching observation identity does not make conflicting + # publication content an idempotent replay. + return await self.commit_revision(revision, rows) retained = self._provisional_acquisitions.get(key) if retained != acquisition: raise LedgerConflictError("OBSERVATION_CONFLICT", "acquisition was not reserved") diff --git a/src/adcp/reporting/outbox/_schema.py b/src/adcp/reporting/outbox/_schema.py index 07f05406d..50ccb4887 100644 --- a/src/adcp/reporting/outbox/_schema.py +++ b/src/adcp/reporting/outbox/_schema.py @@ -33,12 +33,33 @@ def _digest(value: object) -> str: async def schema_objects(connection: Any) -> dict[str, dict[str, Any]]: - tables = await ( - await connection.execute( + return await _catalog_objects(connection) + + +async def _catalog_objects( + connection: Any, + *, + table_names: tuple[str, ...] | None = None, + function_identity: tuple[str, str] | None = None, +) -> dict[str, dict[str, Any]]: + """Read fresh catalog state with the same fingerprints for either scope.""" + table_query = ( + "SELECT c.oid, c.relname, c.relkind, c.relpersistence FROM pg_class c" + " JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = current_schema()" + " AND c.relkind IN ('r','p') AND starts_with(c.relname, 'reporting_')" + " ORDER BY c.relname" + ) + if table_names is not None: + table_query = ( "SELECT c.oid, c.relname, c.relkind, c.relpersistence FROM pg_class c" " JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = current_schema()" " AND c.relkind IN ('r','p') AND starts_with(c.relname, 'reporting_')" - " ORDER BY c.relname" + " AND c.relname = ANY(%s::text[]) ORDER BY c.relname" + ) + tables = await ( + await connection.execute( + table_query, + (list(table_names),) if table_names is not None else None, ) ).fetchall() result: dict[str, dict[str, Any]] = {} @@ -49,7 +70,7 @@ def remember(key: str, value: object, *, enabled: bool = True) -> None: names = {oid: name for oid, name, _, _ in tables} for _, name, kind, persistence in tables: remember(f"table:{name}", (kind, persistence)) - # Capture each object kind for the entire schema in one query. These are + # Capture each object kind for the selected tables in one query. These are # still fresh catalog reads on the caller's connection; no cache can hide # DDL drift. Round trips no longer grow with the number of reporting tables. oids = list(names) @@ -100,16 +121,25 @@ def remember(key: str, value: object, *, enabled: bool = True) -> None: ).fetchall() for oid, *row in triggers: remember(f"trigger:{names[oid]}.{row[0]}", row[1:], enabled=row[1] != "D") - functions = await ( - await connection.execute( + function_query = ( + "SELECT p.proname, pg_get_function_identity_arguments(p.oid), p.prosrc," + " l.lanname, p.provolatile, p.proisstrict, p.prosecdef, p.proconfig," + " pg_get_function_result(p.oid), p.proparallel, p.proleakproof FROM pg_proc p" + " JOIN pg_namespace n ON n.oid = p.pronamespace JOIN pg_language l ON l.oid = p.prolang" + " WHERE n.nspname = current_schema() AND starts_with(p.proname, 'reporting_')" + " ORDER BY p.proname, pg_get_function_identity_arguments(p.oid)" + ) + if function_identity is not None: + function_query = ( "SELECT p.proname, pg_get_function_identity_arguments(p.oid), p.prosrc," " l.lanname, p.provolatile, p.proisstrict, p.prosecdef, p.proconfig," " pg_get_function_result(p.oid), p.proparallel, p.proleakproof FROM pg_proc p" " JOIN pg_namespace n ON n.oid = p.pronamespace JOIN pg_language l ON l.oid = p.prolang" " WHERE n.nspname = current_schema() AND starts_with(p.proname, 'reporting_')" + " AND p.proname=%s AND pg_get_function_identity_arguments(p.oid)=%s" " ORDER BY p.proname, pg_get_function_identity_arguments(p.oid)" ) - ).fetchall() + functions = await (await connection.execute(function_query, function_identity)).fetchall() for row in functions: remember(f"function:{row[0]}({row[1]})", row[2:]) return result @@ -161,7 +191,14 @@ async def validate_provisional_schema(connection: Any) -> None: from adcp.reporting.ledger.store import LedgerConflictError try: - installed = await schema_objects(connection) + installed = await _catalog_objects( + connection, + table_names=( + "reporting_provisional_acquisitions", + "reporting_provisional_observations", + ), + function_identity=("reporting_provisional_immutable", ""), + ) except Exception: raise LedgerConflictError( "PROVISIONAL_SCHEMA_UNREADY", "provisional_schema_unready:catalog_unavailable" diff --git a/src/adcp/reporting/production/__init__.py b/src/adcp/reporting/production/__init__.py index c68445505..6c4d0fa2f 100644 --- a/src/adcp/reporting/production/__init__.py +++ b/src/adcp/reporting/production/__init__.py @@ -33,6 +33,7 @@ ReportingProductionDestination, ReportingProductionSupport, ) +from adcp.reporting.production.source_registry import ReportingProductionSourceRegistry if TYPE_CHECKING: from adcp.reporting.production.pg import PgReportingProductionOutbox, PgReportingProductionStore @@ -54,6 +55,7 @@ "ReportingProductionSigning", "ReportingProductionSource", "ReportingProductionSourceBinding", + "ReportingProductionSourceRegistry", "ReportingProductionSupport", "production_notification_workers", ] diff --git a/src/adcp/reporting/production/configuration.py b/src/adcp/reporting/production/configuration.py index fd81ee1af..40505eeee 100644 --- a/src/adcp/reporting/production/configuration.py +++ b/src/adcp/reporting/production/configuration.py @@ -214,9 +214,7 @@ async def admit(value: ReportingConfigurationAdmission) -> None: raise LedgerConflictError( "CONFIGURATION_GENERATION_IMMUTABLE", "configuration identity conflicts" ) - await support.store.admit_production_configuration( - value.configuration, value.binding, offering_id=value.offering_id - ) + await support._admit_configuration(value) # The caller already entered the migrated, drained production # lifecycle. Complete this account's versioned baseline before # echoing ready; adopters need no account-enumeration worker or diff --git a/src/adcp/reporting/production/contracts.py b/src/adcp/reporting/production/contracts.py index 3e2e1f6d9..e4f6cf309 100644 --- a/src/adcp/reporting/production/contracts.py +++ b/src/adcp/reporting/production/contracts.py @@ -12,7 +12,13 @@ from adcp.reporting.ledger.delivery_models import ReportingDestinationBinding from adcp.reporting.ledger.models import ReportingConfiguration, ReportingConfigurationGenerationKey from adcp.reporting.ledger.store import _config_payload -from adcp.reporting.materializer.contracts import ReportingWriterCapability, failure +from adcp.reporting.materializer.contracts import ( + ReportingWriterCapability, + ReportingWriterError, + ReportingWriterFailure, + failure, +) +from adcp.reporting.production.source_registry import ReportingProductionSourceContext from adcp.reporting.source import ( MediaBuyConstituentV1, ReportingConstituent, @@ -127,6 +133,9 @@ class ReportingProductionSourceBinding: capabilities_sha256: str media_buy_products: tuple[tuple[str, str], ...] configuration_sha256: str = field(kw_only=True) + service_context: ReportingProductionSourceContext | None = field( + default=None, kw_only=True, repr=False + ) def __post_init__(self) -> None: from adcp.reporting.evidence import reporting_identifier, sha256_value @@ -135,6 +144,11 @@ def __post_init__(self) -> None: raise ValueError("source binding requires an exact configuration generation") sha256_value(self.capabilities_sha256) sha256_value(self.configuration_sha256) + if ( + self.service_context is not None + and type(self.service_context) is not ReportingProductionSourceContext + ): + raise ValueError("source binding requires an exact service source context") pairs = tuple(tuple(pair) for pair in self.media_buy_products) if any(len(pair) != 2 for pair in pairs): raise ValueError("source binding requires media-buy/product pairs") @@ -170,7 +184,7 @@ def for_configuration( def document(self) -> dict[str, Any]: key = self.generation_key - return { + document = { "account_id": key.account_id, "delivery_config_id": key.delivery_config_id, "delivery_config_version": key.delivery_config_version, @@ -178,6 +192,12 @@ def document(self) -> dict[str, Any]: "configuration_sha256": self.configuration_sha256, "media_buy_products": [list(pair) for pair in self.media_buy_products], } + if self.service_context is not None: + document["service_context"] = self.service_context.document() + document["service_context_sha256"] = hashlib.sha256( + self.service_context._wire + ).hexdigest() + return document def check( self, @@ -207,15 +227,36 @@ def constituents(self) -> tuple[ReportingConstituent, ...]: ) +class _SourceAuthorizationRevokedError(ReportingWriterError): + """Stop this account's source work for the current scheduling turn.""" + + def __init__(self) -> None: + super().__init__(ReportingWriterFailure("BINDING_MISMATCH")) + + @runtime_checkable class ReportingProductionSource(ReportingSourceExecutor, Protocol): """A source with an authenticated generation mapping available before I/O. Discovery may precede any account binding. Admission and every source turn require the applicable binding, obtained from the source's trusted account - configuration. Returning ``None`` withdraws authorization for new work. + configuration. The SDK calls ``configuration_binding`` immediately before + dispatch and again under the account lock before sealing or publishing the + result, including replay after a restart. Returning ``None`` discards that + result without a success checkpoint and stops this account's remaining + work for the turn. Restoring the binding permits work on the next turn. + + An in-flight fetch is allowed to finish: revocation takes effect at the next + dispatch or publish. Adopters own any caching or latency inside their + callback; the SDK does not retain an authorization grant. """ def configuration_binding( self, configuration: ReportingConfiguration - ) -> ReportingProductionSourceBinding | None: ... + ) -> ReportingProductionSourceBinding | None: + """Return the current trusted binding, or ``None`` to deny source work. + + Adopters own callback caching and latency. In-flight fetches finish; + revocation takes effect at the next dispatch or publish. + """ + raise NotImplementedError diff --git a/src/adcp/reporting/production/handler.py b/src/adcp/reporting/production/handler.py index 6a77df295..6effc9bd2 100644 --- a/src/adcp/reporting/production/handler.py +++ b/src/adcp/reporting/production/handler.py @@ -3,7 +3,10 @@ from __future__ import annotations import hashlib -from typing import TYPE_CHECKING, Any +import inspect +from collections.abc import Callable +from functools import wraps +from typing import TYPE_CHECKING, Any, TypeVar, cast from adcp.decisioning.context import RequestContext from adcp.exceptions import ADCPTaskError @@ -18,7 +21,8 @@ ReportingReceiptHandler, _consumer, ) -from adcp.server.base import NotImplementedResponse, ToolContext +from adcp.server.base import ADCPHandler, NotImplementedResponse, ToolContext +from adcp.server.mcp_tools import ADCP_TOOL_DEFINITIONS, get_tools_for_handler from adcp.server.responses import capabilities_response from adcp.types import ( Error, @@ -47,17 +51,42 @@ def _task_error(task: str, code: str, message: str) -> ADCPTaskError: return ADCPTaskError(operation=task, errors=[Error(code=code, message=message)]) +_Method = TypeVar("_Method", bound=Callable[..., Any]) +_REPORTING_TASKS = { + "get_adcp_capabilities", + "get_reporting_status", + "get_media_buy_delivery", + "sync_accounts", + "sync_reporting_receipts", + "sync_reporting_status", +} +_APPLICATION_TASKS = {tool["name"] for tool in ADCP_TOOL_DEFINITIONS} - _REPORTING_TASKS +_APPLICATION_METHODS = { + "build_creative": "build_creative_legacy", + "preview_creative": "preview_creative_legacy", + "list_creative_formats": "list_creative_formats_legacy", +} + + +def _admitted(method: _Method) -> _Method: + @wraps(method) + async def call(self: ReportingProductionHandler, *args: Any, **kwargs: Any) -> Any: + lifecycle = self.production._service_lifecycle + aggregate = ( + method.__name__ == "get_media_buy_delivery" + and "reporting_revision_id" not in _request(args[0] if args else kwargs["params"]) + ) + if lifecycle is None or aggregate: + return await method(self, *args, **kwargs) + return await lifecycle.call(lambda: method(self, *args, **kwargs)) + + return cast(_Method, call) + + class ReportingProductionHandler(ReportingReceiptHandler): """Mount the same instance on MCP/A2A; the support owns its lifecycle.""" - advertised_tools = { - "get_adcp_capabilities", - "get_reporting_status", - "get_media_buy_delivery", - "sync_accounts", - "sync_reporting_receipts", - "sync_reporting_status", - } + advertised_tools = _REPORTING_TASKS | _APPLICATION_TASKS def __init__( self, @@ -68,6 +97,8 @@ def __init__( adcp_version: str | None = None, ) -> None: self.production = production + self._application: ADCPHandler[Any] | None = None + self._application_tools: frozenset[str] = frozenset() super().__init__( production.store, resolve_account=resolve_account, @@ -76,6 +107,37 @@ def __init__( adcp_version=adcp_version, ) + def bind_application(self, application: ADCPHandler[Any]) -> None: + """Delegate ordinary tasks before mounting this exact SDK handler. + + Reporting/configuration collisions are refused. Supply the typed B2 + configuration task separately; aggregate delivery and base capabilities + may coexist and are dispatched explicitly by their protected methods. + """ + if application is self or not isinstance(application, ADCPHandler): + raise ValueError("application must be a separate ADCPHandler") + if self._application is application: + return + if ( + self._application is not None + or self.production._mounts + or self.production._task is not None + ): + raise ValueError("application delegation must be fixed before mounting") + names = {t["name"] for t in get_tools_for_handler(application, _include_schemas=False)} + if names & (_REPORTING_TASKS - {"get_adcp_capabilities", "get_media_buy_delivery"}): + raise ValueError("application reporting handlers conflict with production ownership") + version = getattr(application, "get_adcp_version", None) + if callable(version) and version() != self.get_adcp_version(): + raise ValueError("application and production protocol versions differ") + self._application, self._application_tools = application, frozenset(names) + + async def _delegate(self, task: str, params: Any, context: ToolContext | None) -> Any: + if self._application is None or task not in self._application_tools: + return self._not_supported(task) + result = getattr(self._application, _APPLICATION_METHODS.get(task, task))(params, context) + return await result if inspect.isawaitable(result) else result + def advertised_tools_for_instance(self) -> set[str]: names = { "get_adcp_capabilities", @@ -87,8 +149,9 @@ def advertised_tools_for_instance(self) -> set[str]: names.add("sync_reporting_receipts") if self._feed_consumer_status_enabled: names.add("sync_reporting_status") - return names + return names | set(self._application_tools & _APPLICATION_TASKS) + @_admitted async def get_reporting_status( self, params: GetReportingStatusRequest | dict[str, Any], @@ -96,6 +159,7 @@ async def get_reporting_status( ) -> dict[str, Any] | NotImplementedResponse: return await super().get_reporting_status(params, context) + @_admitted async def sync_reporting_receipts( self, params: SyncReportingReceiptsRequest | dict[str, Any], @@ -135,11 +199,36 @@ async def get_adcp_capabilities( adcp_version=self.production._protocol_version, supported_versions=[self.production._protocol_version], ) + if self._application is not None and "get_adcp_capabilities" in self._application_tools: + protocol = response["adcp"] + result = await self._delegate("get_adcp_capabilities", params, context) + base = {} if isinstance(result, NotImplementedResponse) else _request(result) + if (base.get("media_buy") or {}).get("reporting_delivery") is not None: + raise ValueError( + "application reporting capabilities conflict with production ownership" + ) + response.update(base) + response["adcp_version"] = self.production._protocol_version + response["adcp"] = { + **response["adcp"], + "major_versions": protocol["major_versions"], + "supported_versions": protocol["supported_versions"], + } + response["supported_protocols"] = list( + dict.fromkeys([*base.get("supported_protocols", ()), "media_buy"]) + ) response["account"] = self.production.configuration_task.account_capabilities() reporting = await self.production.reporting_delivery() if reporting: - response["media_buy"] = {"reporting_delivery": reporting} - response["experimental_features"] = ["media_buy.reporting_delivery"] + response["media_buy"] = { + **response.get("media_buy", {}), + "reporting_delivery": reporting, + } + response["experimental_features"] = list( + dict.fromkeys( + [*response.get("experimental_features", ()), "media_buy.reporting_delivery"] + ) + ) if any( reporting.get(k) for k in ("ledger_notification", "status_notification", "readiness_notification") @@ -153,6 +242,7 @@ async def get_adcp_capabilities( response["identity"] = signing_identity(self.production) return response + @_admitted async def sync_accounts( self, params: SyncAccountsRequest | dict[str, Any], @@ -174,6 +264,7 @@ async def sync_accounts( ) raise _task_error("sync_accounts", code, message) + @_admitted async def sync_reporting_status( self, params: SyncReportingStatusRequest | dict[str, Any], @@ -198,6 +289,7 @@ async def sync_reporting_status( code, message = "REPORTING_STATUS_UNAVAILABLE", "reporting status is unavailable" raise _task_error("sync_reporting_status", code, message) + @_admitted async def get_media_buy_delivery( self, params: GetMediaBuyDeliveryRequest | dict[str, Any], @@ -205,7 +297,10 @@ async def get_media_buy_delivery( ) -> dict[str, Any] | NotImplementedResponse: request = _request(params) if "reporting_revision_id" not in request: - return self._not_supported("get_media_buy_delivery") + return cast( + dict[str, Any] | NotImplementedResponse, + await self._delegate("get_media_buy_delivery", params, context), + ) try: from adcp.validation.schema_loader import get_named_validator @@ -320,3 +415,24 @@ async def get_media_buy_delivery( except Exception: code, message = "REPORTING_CONTENT_UNAVAILABLE", "reporting content is unavailable" raise _task_error("get_media_buy_delivery", code, message) + + +def _application_method(task: str) -> Callable[..., Any]: + async def delegated( + self: ReportingProductionHandler, params: Any, context: ToolContext | None = None + ) -> Any: + return await self._delegate(task, params, context) + + delegated.__name__ = task + return delegated + + +# Concrete SDK methods on the exact handler class, never an adopter subclass. +# The per-instance inventory admits only tools actually implemented by the +# delegate. Protected production methods are excluded from this fixed set. +for _task in _APPLICATION_TASKS: + setattr( + ReportingProductionHandler, + _APPLICATION_METHODS.get(_task, _task), + _application_method(_task), + ) diff --git a/src/adcp/reporting/production/memory.py b/src/adcp/reporting/production/memory.py index b804a23bc..85a78f483 100644 --- a/src/adcp/reporting/production/memory.py +++ b/src/adcp/reporting/production/memory.py @@ -30,6 +30,7 @@ from adcp.reporting.materializer.work import MaterializerContext, ReportingMaterializerLease from adcp.reporting.outbox.memory import InMemoryReportingOutbox, NotificationState from adcp.reporting.production.contracts import ReportingProductionSourceBinding +from adcp.reporting.production.source_registry import ReportingProductionSourceContext from adcp.reporting.projection.memory import InMemoryReportingProjectionStore from adcp.reporting.source import ReportingConstituent @@ -106,6 +107,7 @@ async def admit_production_configuration( binding: ReportingDestinationBinding, *, offering_id: str, + service_context: ReportingProductionSourceContext | None = None, ) -> None: offering = self._owner()._configuration_offering( configuration, binding, offering_id=offering_id @@ -113,15 +115,25 @@ async def admit_production_configuration( async with self._mutation(): await self.put_configuration(configuration) await self.put_destination_binding(binding) - self._enroll(configuration, offering._producer_key, binding) + self._enroll( + configuration, offering._producer_key, binding, service_context=service_context + ) def _enroll( self, configuration: ReportingConfiguration, producer_key: str, destination: ReportingDestinationBinding, + *, + service_context: ReportingProductionSourceContext | None = None, ) -> None: - binding = self._owner()._source_binding(configuration, producer_key) + previous_context = self._production_source_bindings.get(configuration.generation_key) + binding = self._owner()._source_binding( + configuration, + producer_key, + service_context=service_context, + document=previous_context.document() if previous_context is not None else None, + ) previous = self._production_generations.setdefault( configuration.generation_key, producer_key ) @@ -148,24 +160,23 @@ def _destination_document(self, binding: ReportingDestinationBinding) -> dict[st return dict(json.loads(raw)) if raw is not None else None def _configuration_lease_eligible(self, configuration: ReportingConfiguration) -> bool: - if configuration.account_id not in self._production_accounts: - return False - producer_key = self._production_generations.get(configuration.generation_key) - binding = self._production_source_bindings.get(configuration.generation_key) - if producer_key not in self._owner()._producer_keys() or binding is None: - return False try: - self._owner()._check_source_binding(configuration, producer_key, binding.document()) + self._check_source_generation(configuration) except Exception: return False return True def _check_source_generation(self, configuration: ReportingConfiguration) -> None: + producer_key = self._production_generations.get(configuration.generation_key) + binding = self._production_source_bindings.get(configuration.generation_key) if ( - not self._configuration_lease_eligible(configuration) + configuration.account_id not in self._production_accounts + or producer_key not in self._owner()._producer_keys() + or binding is None or self._configurations.get(configuration.generation_key) != configuration ): raise LedgerConflictError("HISTORY_UNAVAILABLE", "producer generation is unavailable") + self._owner()._check_source_binding(configuration, producer_key, binding.document()) def _wake_obligation(self, account_id: str, obligation_id: str) -> None: super()._wake_obligation(account_id, obligation_id) diff --git a/src/adcp/reporting/production/offerings.py b/src/adcp/reporting/production/offerings.py index 932ffc8f6..fdd426810 100644 --- a/src/adcp/reporting/production/offerings.py +++ b/src/adcp/reporting/production/offerings.py @@ -8,6 +8,7 @@ from datetime import datetime from typing import Any +from adcp.reporting._source_authorization import require_account_work from adcp.reporting._timestamp import aware_timestamp from adcp.reporting.canonical_json import canonical_json_utf8_v1 from adcp.reporting.ledger.delivery_models import ReportingDestinationBinding @@ -22,6 +23,7 @@ from adcp.reporting.production.contracts import ( ReportingProductionSource, ReportingProductionSourceBinding, + _SourceAuthorizationRevokedError, ) from adcp.reporting.source import ( AuthoritativeOfferingV1, @@ -229,15 +231,24 @@ def check_source(self, *, effective: bool = False) -> ReportingSourceCapabilitie def source_binding( self, configuration: ReportingConfiguration ) -> ReportingProductionSourceBinding: + require_account_work(configuration.account_id) try: capabilities = self.check_source(effective=True) source = self.producer._source assert isinstance(source, ReportingProductionSource) binding = source.configuration_binding(configuration) + if binding is None: + # Discovery can probe several generations before selecting a + # dispatch. Refuse this candidate without revoking a turn that + # may select another one; the producer's dispatch/publication + # checks own the account-wide stop after a live denial. + raise _SourceAuthorizationRevokedError() if type(binding) is not ReportingProductionSourceBinding: raise failure("BINDING_MISMATCH") binding.check(configuration, capabilities, self.source_offering_id) return binding + except _SourceAuthorizationRevokedError: + raise except Exception: raise failure("BINDING_MISMATCH") from None diff --git a/src/adcp/reporting/production/pg.py b/src/adcp/reporting/production/pg.py index 9d25821e0..cbcaf814c 100644 --- a/src/adcp/reporting/production/pg.py +++ b/src/adcp/reporting/production/pg.py @@ -37,6 +37,7 @@ key_for, ) from adcp.reporting.outbox.status_pg import PgReportingStatusOutbox +from adcp.reporting.production.source_registry import ReportingProductionSourceContext from adcp.reporting.projection.pg import PgReportingProjectionStore from adcp.reporting.source import ReportingConstituent @@ -166,6 +167,7 @@ async def admit_production_configuration( binding: ReportingDestinationBinding, *, offering_id: str, + service_context: ReportingProductionSourceContext | None = None, ) -> None: offering = self._owner()._configuration_offering( configuration, binding, offering_id=offering_id @@ -176,7 +178,13 @@ async def admit_production_configuration( try: await self.put_configuration(configuration) await self.put_destination_binding(binding) - await self._enroll_on(connection, configuration, offering._producer_key, binding) + await self._enroll_on( + connection, + configuration, + offering._producer_key, + binding, + service_context=service_context, + ) finally: _ADMISSION_CONNECTION.reset(token) @@ -186,10 +194,28 @@ async def _enroll_on( configuration: ReportingConfiguration, producer_key: str, destination: ReportingDestinationBinding, + *, + service_context: ReportingProductionSourceContext | None = None, ) -> None: key = configuration.generation_key identity = (key.account_id, key.delivery_config_id, key.delivery_config_version) - binding = self._owner()._source_binding(configuration, producer_key).document() + previous_context = await ( + await connection.execute( + "SELECT source_binding FROM reporting_production_generations" + " WHERE account_id=%s AND delivery_config_id=%s AND delivery_config_version=%s", + identity, + ) + ).fetchone() + binding = ( + self._owner() + ._source_binding( + configuration, + producer_key, + service_context=service_context, + document=previous_context[0] if previous_context is not None else None, + ) + .document() + ) await connection.execute( "INSERT INTO reporting_production_generations VALUES(%s,%s,%s,%s,%s::jsonb)" " ON CONFLICT DO NOTHING", @@ -568,6 +594,9 @@ async def create_schema(self) -> None: "reporting_production.sql", ): await connection.execute(root.joinpath(name).read_text()) + await connection.execute( + files("adcp.reporting.production").joinpath("service_context.sql").read_text() + ) async def materializer_ready(self) -> bool: owner = self._owner() diff --git a/src/adcp/reporting/production/schema.py b/src/adcp/reporting/production/schema.py index b3cf06712..7212a3ffd 100644 --- a/src/adcp/reporting/production/schema.py +++ b/src/adcp/reporting/production/schema.py @@ -11,11 +11,21 @@ from adcp.reporting.projection.schema import validate_projection_schema -async def validate_production_schema(connection: Any, *, notifications: bool = False) -> None: +async def validate_production_schema( + connection: Any, *, notifications: bool = False, service_context: bool = False +) -> None: try: required = json.loads( files("adcp.reporting.production").joinpath("required_schema.json").read_text() ) + if service_context: + required.update( + json.loads( + files("adcp.reporting.production") + .joinpath("service_context_schema.json") + .read_text() + ) + ) actual = await schema_objects(connection) if not required or any(actual.get(k) != v for k, v in required.items()): raise ValueError diff --git a/src/adcp/reporting/production/service.py b/src/adcp/reporting/production/service.py index 1f486f9cf..185343910 100644 --- a/src/adcp/reporting/production/service.py +++ b/src/adcp/reporting/production/service.py @@ -8,10 +8,11 @@ import weakref from collections.abc import Mapping from contextvars import ContextVar -from dataclasses import dataclass +from dataclasses import dataclass, replace from datetime import timedelta from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable +from adcp.reporting._source_authorization import source_turn from adcp.reporting.ledger.delivery_models import ReportingDestinationBinding from adcp.reporting.ledger.models import ReportingConfiguration from adcp.reporting.ledger.notification_models import ReportingNotificationError @@ -40,6 +41,10 @@ ) from adcp.reporting.production.memory import InMemoryReportingProductionStore from adcp.reporting.production.offerings import ReportingProductionOffering +from adcp.reporting.production.source_registry import ( + ReportingProductionSourceContext, + ReportingProductionSourceRegistry, +) from adcp.reporting.projection.memory import InMemoryReportingStatusProjection from adcp.reporting.receipts.handler import ReceiptAccountResolver @@ -49,6 +54,7 @@ from adcp.reporting.outbox.worker import ReportingNotificationWorker from adcp.reporting.production.pg import PgReportingProductionStore from adcp.reporting.projection.pg import PgReportingStatusProjection + from adcp.reporting.service_lifecycle import _ServiceLifecycle from adcp.server.base import ADCPHandler @@ -166,6 +172,7 @@ def __init__( notification_workers: tuple[ReportingNotificationWorker, ...] = (), poll_seconds: float = 0.25, adcp_version: str | None = None, + source_registry: ReportingProductionSourceRegistry | None = None, ) -> None: from adcp.reporting.outbox.worker import ReportingNotificationWorker from adcp.reporting.production.handler import ReportingProductionHandler @@ -203,6 +210,12 @@ def __init__( raise ReportingNotificationError("reporting_production_owner_conflict") self.materializer, self.projection = materializer, projection self.offerings, self.configuration_task = offerings, configuration_task + if source_registry is not None: + if type(source_registry) is not ReportingProductionSourceRegistry: + raise ValueError("production requires an exact source registry") + source_registry.freeze(offerings) + self.source_registry = source_registry + self._service_lifecycle: _ServiceLifecycle | None = None self.automated_recovery_window, self.status_retention_days = ( automated_recovery_window, status_retention_days, @@ -259,6 +272,7 @@ def _component_ids(self) -> tuple[int, ...]: self.projection, self.projection.outbox, self.configuration_task, + self.source_registry, *self.offerings, *self.notification_workers, ) @@ -380,6 +394,8 @@ def _offering_ready(self, offering: ReportingProductionOffering) -> bool: ): return False offering.check_source() + if self.source_registry is not None: + self.source_registry._registration(offering) return True except Exception: return False @@ -466,7 +482,12 @@ def validate_configuration(self, value: ReportingConfigurationAdmission) -> None value.check(self) def _source_binding( - self, configuration: ReportingConfiguration, producer_key: str + self, + configuration: ReportingConfiguration, + producer_key: str, + *, + document: dict[str, Any] | None = None, + service_context: ReportingProductionSourceContext | None = None, ) -> ReportingProductionSourceBinding: self._assert_components() matches = [ @@ -476,7 +497,39 @@ def _source_binding( ] if len(matches) != 1: raise failure("BINDING_MISMATCH") - return matches[0].source_binding(configuration) + offering = matches[0] + binding = offering.source_binding(configuration) + if self.source_registry is not None: + context = self.source_registry.recover( + configuration, + offering, + offering.check_source(effective=True), + ( + service_context.document() + if service_context is not None + else (document or {}).get("service_context") + ), + ) + binding = replace(binding, service_context=context) + elif service_context is not None or (document or {}).get("service_context") is not None: + raise failure("BINDING_MISMATCH") + return binding + + async def _admit_configuration(self, value: ReportingConfigurationAdmission) -> None: + context = None + if self.source_registry is not None: + offering = self._configuration_offering( + value.configuration, value.binding, offering_id=value.offering_id + ) + context = await self.source_registry.resolve( + value.configuration, offering, offering.check_source(effective=True) + ) + await self.store.admit_production_configuration( + value.configuration, + value.binding, + offering_id=value.offering_id, + service_context=context, + ) def _check_source_binding( self, @@ -484,7 +537,12 @@ def _check_source_binding( producer_key: str, document: dict[str, Any] | None, ) -> ReportingProductionSourceBinding: - binding = self._source_binding(configuration, producer_key) + try: + binding = self._source_binding(configuration, producer_key, document=document) + except (ValueError, TypeError): + # Runtime incompatibility follows the existing scoped refusal path: + # unavailable legacy/profile facts cannot stop unrelated work. + raise failure("BINDING_MISMATCH") from None if binding.document() != document: raise failure("BINDING_MISMATCH") return binding @@ -538,7 +596,9 @@ async def scan() -> bool: try: async with pool.connection() as connection: await validate_production_schema( - connection, notifications=self.notifications_enabled + connection, + notifications=self.notifications_enabled, + service_context=self.source_registry is not None, ) if epoch == self._schema_epoch and pool is store._pool: self._schema_positive = identity @@ -582,13 +642,14 @@ async def _run(self) -> None: while not self._stop.is_set(): # Each producer leases a generation through the reviewed fair # indexed acquisition path; no account enumeration is required. - for producer in dict.fromkeys(o.producer for o in self.offerings): - boundary = "producer" - token = self._producer_turn.set(producer) - try: - await producer.run_worker() - finally: - self._producer_turn.reset(token) + with source_turn(): + for producer in dict.fromkeys(o.producer for o in self.offerings): + boundary = "producer" + token = self._producer_turn.set(producer) + try: + await producer.run_worker() + finally: + self._producer_turn.reset(token) boundary = "materializer" await self.materializer.run_once() boundary = "projection" diff --git a/src/adcp/reporting/production/service_context.sql b/src/adcp/reporting/production/service_context.sql new file mode 100644 index 000000000..a43222f55 --- /dev/null +++ b/src/adcp/reporting/production/service_context.sql @@ -0,0 +1,58 @@ +-- Optional fixed-profile service context. Existing B2 documents stay valid. +DO $service_context$ +BEGIN + PERFORM pg_advisory_xact_lock(hashtext('adcp.reporting.schema'), hashtext(current_schema())); + CREATE OR REPLACE FUNCTION reporting_production_service_context_guard() + RETURNS TRIGGER LANGUAGE plpgsql AS $guard$ + DECLARE + facts JSONB := NEW.source_binding->'service_context'; + zone TEXT; + configuration_hash TEXT; + BEGIN + IF facts IS NULL AND NOT NEW.source_binding ? 'service_context_sha256' THEN + RETURN NEW; + END IF; + SELECT account_timezone,reporting_payload_sha256(jsonb_build_object( + 'account_id',account_id,'report_definition_id',report_definition_id, + 'reporting_profile',reporting_profile,'feed_purpose',feed_purpose, + 'required_finality',required_finality,'account_timezone',account_timezone, + 'authoritative_party',authoritative_party,'media_buy_ids',media_buy_ids, + 'definition',definition,'schedule',schedule)) + INTO zone,configuration_hash FROM reporting_configurations + WHERE account_id=NEW.account_id AND delivery_config_id=NEW.delivery_config_id + AND delivery_config_version=NEW.delivery_config_version; + IF jsonb_typeof(facts) IS DISTINCT FROM 'object' + OR NEW.source_binding->>'service_context_sha256' + IS DISTINCT FROM reporting_payload_sha256(facts) + OR facts->'version' IS DISTINCT FROM '1'::jsonb + OR facts->>'account_id' IS DISTINCT FROM NEW.account_id + OR facts->>'account_timezone' IS DISTINCT FROM zone + OR NEW.source_binding->>'configuration_sha256' IS DISTINCT FROM configuration_hash + OR NEW.source_binding->>'account_id' IS DISTINCT FROM NEW.account_id + OR NEW.source_binding->>'delivery_config_id' IS DISTINCT FROM NEW.delivery_config_id + OR NEW.source_binding->'delivery_config_version' + IS DISTINCT FROM to_jsonb(NEW.delivery_config_version) + OR facts->>'capabilities_sha256' + IS DISTINCT FROM NEW.source_binding->>'capabilities_sha256' + OR facts->>'currency' IS NULL OR facts->>'currency' !~ '^[A-Z]{3}$' + OR facts->>'adapter' IS NULL OR facts->>'adapter' !~ '^[A-Za-z0-9_.:-]{1,128}$' + OR NOT facts ?& ARRAY['version','adapter','account_id','account_timezone', + 'capabilities_sha256','currency','offering_id','source_offering_id','snapshot_offering_id', + 'official_offering_id','publication_namespace','requested_metrics', + 'requested_dimensions','source_scope','slice_timeout_microseconds'] + OR (SELECT count(*) FROM jsonb_object_keys(facts)) <> 15 + THEN + RAISE EXCEPTION 'reporting service context is inconsistent'; + END IF; + RETURN NEW; + END + $guard$; + IF NOT EXISTS (SELECT 1 FROM pg_trigger + WHERE tgrelid='reporting_production_generations'::regclass + AND tgname='reporting_production_service_context_consistent') THEN + CREATE TRIGGER reporting_production_service_context_consistent BEFORE INSERT + ON reporting_production_generations FOR EACH ROW + EXECUTE FUNCTION reporting_production_service_context_guard(); + END IF; +END +$service_context$; diff --git a/src/adcp/reporting/production/service_context_schema.json b/src/adcp/reporting/production/service_context_schema.json new file mode 100644 index 000000000..10c4ce70a --- /dev/null +++ b/src/adcp/reporting/production/service_context_schema.json @@ -0,0 +1,10 @@ +{ + "function:reporting_production_service_context_guard()": { + "enabled": true, + "fingerprint": "db567ebfe99bcdaf3083d972c5976156bf1d51371c33fd65add9ac036ea7076e" + }, + "trigger:reporting_production_generations.reporting_production_service_context_consistent": { + "enabled": true, + "fingerprint": "6e099313085dfab338cc98b9964295733fcfb68972d33a9d04cdf29603ceaacd" + } +} diff --git a/src/adcp/reporting/production/source_registry.py b/src/adcp/reporting/production/source_registry.py new file mode 100644 index 000000000..cd5973eb3 --- /dev/null +++ b/src/adcp/reporting/production/source_registry.py @@ -0,0 +1,167 @@ +"""Fixed execution profiles and durable, secret-free service admission facts. + +This is an opt-in bridge to existing production producers. It does not discover +accounts, supply a new scheduler, or replace an adapter's live authorization. +Profiles are fixed before startup; heterogeneous profiles require distinct B2 +offerings/producers. Retry deadlines remain owned by the acquisition protocol. +""" + +from __future__ import annotations + +import inspect +import json +import re +from dataclasses import dataclass, field +from datetime import timedelta +from typing import TYPE_CHECKING, Any + +from adcp.reporting.canonical_json import canonical_json_utf8_v1 + +if TYPE_CHECKING: + from adcp.reporting.ledger.models import ReportingConfiguration + from adcp.reporting.ledger.producer import ProducerOfferings, ReportingProducer + from adcp.reporting.production.offerings import ReportingProductionOffering + from adcp.reporting.service import ReportingContextResolver + from adcp.reporting.source import ReportingSourceCapabilitiesV1 + + +def _profile(value: ProducerOfferings) -> dict[str, Any]: + return { + "snapshot_offering_id": value.snapshot_offering_id, + "official_offering_id": value.official_offering_id, + "publication_namespace": value.publication_namespace, + "requested_metrics": list(value.requested_metrics), + "requested_dimensions": list(value.requested_dimensions), + "currency": value.currency, + "source_scope": value.source_scope, + # Policy, not a source execution key or an absolute retry deadline. + "slice_timeout_microseconds": value.slice_timeout // timedelta(microseconds=1), + } + + +@dataclass(frozen=True) +class ReportingProductionSourceContext: + """Canonical version-one facts; credentials must never be placed in scope. + + Construction is internal to the registry. Persisted input is accepted only + by comparison with its closed expected document, not by a permissive codec. + """ + + _wire: bytes = field(repr=False) + + def document(self) -> dict[str, Any]: + return dict(json.loads(self._wire)) + + +@dataclass(frozen=True) +class _Registration: + adapter: str + producer: ReportingProducer = field(repr=False, compare=False) + profile: bytes = field(repr=False) + + +class ReportingProductionSourceRegistry: + """Bind stable adapter names to the actual fixed production producers. + + ``account_context`` runs only at authenticated configuration admission, + outside the storage transaction. Recovery validates the selected persisted + generation against these profiles and the live production source binding; + it never calls the resolver or rebuilds an in-memory account catalog. + """ + + def __init__(self, *, account_context: ReportingContextResolver) -> None: + self._resolve = account_context + self._entries: dict[str, _Registration] = {} + self._frozen = False + + def register(self, adapter: str, producer: ReportingProducer) -> None: + if self._frozen: + raise ValueError("production source registry is frozen") + if not re.fullmatch(r"[A-Za-z0-9_.:-]{1,128}", adapter): + raise ValueError("adapter must be a stable identifier") + if adapter in self._entries or any(v.producer is producer for v in self._entries.values()): + raise ValueError("production source profile must have one unambiguous registration") + self._entries[adapter] = _Registration( + adapter, producer, canonical_json_utf8_v1(_profile(producer._offerings)) + ) + + def freeze(self, offerings: tuple[ReportingProductionOffering, ...]) -> None: + for offering in offerings: + self._registration(offering) + if any( + not any(entry.producer is offering.producer for offering in offerings) + for entry in self._entries.values() + ): + raise ValueError("registered production source profile has no offering") + self._frozen = True + + def _registration(self, offering: ReportingProductionOffering) -> _Registration: + matches = [v for v in self._entries.values() if v.producer is offering.producer] + if len(matches) != 1: + raise ValueError("production source profile is not registered") + entry = matches[0] + if entry.profile != canonical_json_utf8_v1(_profile(entry.producer._offerings)): + raise ValueError("production source profile changed") + return entry + + def _expected( + self, + configuration: ReportingConfiguration, + offering: ReportingProductionOffering, + capabilities: ReportingSourceCapabilitiesV1, + ) -> dict[str, Any]: + entry = self._registration(offering) + return { + "version": 1, + "adapter": entry.adapter, + "account_id": configuration.account_id, + "account_timezone": configuration.account_timezone, + "offering_id": offering.offering_id, + "source_offering_id": offering.source_offering_id, + "capabilities_sha256": capabilities.capabilities_sha256, + **json.loads(entry.profile), + } + + async def resolve( + self, + configuration: ReportingConfiguration, + offering: ReportingProductionOffering, + capabilities: ReportingSourceCapabilitiesV1, + ) -> ReportingProductionSourceContext: + from adcp.reporting.service import ReportingAccountContext, _thaw + + context = self._resolve(configuration) + if inspect.isawaitable(context): + context = await context + expected = self._expected(configuration, offering, capabilities) + if ( + type(context) is not ReportingAccountContext + or context.account_id != configuration.account_id + or context.account_timezone != configuration.account_timezone + or context.adapter != expected["adapter"] + or canonical_json_utf8_v1(_profile(context.producer_offerings())) + != self._registration(offering).profile + or ( + context.capability_offering + and canonical_json_utf8_v1(_thaw(context.capability_offering)) + != canonical_json_utf8_v1(offering.wire()) + ) + ): + raise ValueError("resolved account context does not match its fixed production profile") + return ReportingProductionSourceContext(canonical_json_utf8_v1(expected)) + + def recover( + self, + configuration: ReportingConfiguration, + offering: ReportingProductionOffering, + capabilities: ReportingSourceCapabilitiesV1, + document: dict[str, Any] | None, + ) -> ReportingProductionSourceContext: + if document is None: + raise ValueError("legacy service source context requires explicit migration") + expected = canonical_json_utf8_v1(self._expected(configuration, offering, capabilities)) + if expected != canonical_json_utf8_v1(document): + raise ValueError( + "persisted account context does not match its fixed production profile" + ) + return ReportingProductionSourceContext(expected) diff --git a/src/adcp/reporting/service.py b/src/adcp/reporting/service.py index feb467195..ebcb4bc3a 100644 --- a/src/adcp/reporting/service.py +++ b/src/adcp/reporting/service.py @@ -17,11 +17,12 @@ from dataclasses import dataclass, field from datetime import datetime, timedelta, timezone from types import MappingProxyType -from typing import Any, Protocol, runtime_checkable +from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from pydantic import BaseModel +from adcp.reporting._source_authorization import source_turn from adcp.reporting.inline_source import ( InlineFetchResult, InlineReportingSource, @@ -60,6 +61,9 @@ ReportingSourceSliceRequestV1, ) +if TYPE_CHECKING: + from adcp.reporting.production.service import ReportingProductionSupport + __all__ = [ "AdapterRegistration", "ReliableReportingConfigurationError", @@ -362,6 +366,7 @@ def __init__( ) -> None: effective_clock = clock or (lambda: datetime.now(timezone.utc)) self.store = store + self._production: ReportingProductionSupport | None = None self.sources = ReportingAdapterRegistry(clock=effective_clock) self._context_resolver = account_context self._caller_resolver = caller_resolver or self._default_caller @@ -389,6 +394,48 @@ def __init__( configuration_error=ReliableReportingConfigurationError, ) + @classmethod + def from_production( + cls, + production: ReportingProductionSupport, + *, + owned_resources: Sequence[ReportingServiceResource] = (), + ) -> ReliableReportingService: + """Own a preassembled B2 graph with durable fixed-profile admission. + + Mount the exact handler returned by ``install(application)`` before + startup. This explicit bridge requires the real production components + and a source registry; it is not an adapter-first component factory. + Production tasks use their existing indexed leases. Pools and providers + stay borrowed unless explicitly transferred in ``owned_resources``. + """ + from adcp.reporting.production.service import ReportingProductionSupport + + if ( + type(production) is not ReportingProductionSupport + or production.source_registry is None + or production._task is not None + or production._closed + or production._service_lifecycle is not None + ): + raise ReliableReportingConfigurationError( + "requires an unstarted registered production graph" + ) + + def unavailable(_: ReportingConfiguration) -> ReportingAccountContext: + raise ReliableReportingConfigurationError( + "production requires typed sync_accounts admission" + ) + + service = cls( + store=production.store, + account_context=unavailable, + owned_resources=(*owned_resources, ReportingServiceResource(close=production.aclose)), + ) + service._production = production + production._service_lifecycle = service._lifecycle + return service + @classmethod def memory( cls, @@ -430,6 +477,10 @@ def postgres( async def configure(self, configuration: ReportingConfiguration) -> None: """Resolve trusted account facts once and freeze this generation's route.""" + if self._production is not None: + raise ReliableReportingConfigurationError( + "production requires typed sync_accounts admission" + ) async def configure() -> None: async with self._configuration_lock: @@ -558,6 +609,10 @@ def _validate_offerings( def validate(self) -> None: """Fail fast on tier combinations the configured components cannot honor.""" + if self._production is not None and self.sources.names: + raise ReliableReportingConfigurationError( + "production adapters must belong to its frozen source registry" + ) if self._worker_interval is not None and self._worker_interval <= timedelta(0): raise ReliableReportingConfigurationError("worker_interval must be greater than zero") if self._notification_worker is not None and self._notification_attempt_store is None: @@ -587,14 +642,34 @@ async def _initialize(self) -> None: await self.store.put_configuration(configuration) self._pending_configurations.clear() self.sources.freeze() + if self._production is not None: + await self._production.start() self._initialized = True async def start(self) -> None: await self.initialize() await self._lifecycle.activate( - self._worker_loop if self._worker_interval is not None else None + self._monitor_production + if self._production is not None + else (self._worker_loop if self._worker_interval is not None else None) ) + async def _monitor_production(self) -> None: + assert self._production is not None + while not self._lifecycle.stopping: + self._production._assert_components() + stop = asyncio.create_task(self._lifecycle.wait_for_stop(60)) + workers = [ + task + for task in (self._production._task, self._production._notification_task) + if task is not None + ] + try: + await asyncio.wait([stop, *workers], return_when=asyncio.FIRST_COMPLETED) + finally: + stop.cancel() + await asyncio.gather(stop, return_exceptions=True) + @property def state(self) -> ReliableReportingState: """Current process lifecycle; STOPPING retains all unsettled ownership.""" @@ -657,6 +732,8 @@ async def _report_worker_error(self, component: str, error: BaseException) -> No async def run_worker(self, *, now: datetime | None = None) -> ReliableReportingTurn: """Route one turn across every frozen configuration generation.""" + if self._production is not None: + raise ReliableReportingConfigurationError("production workers are owned by start/close") await self.initialize() await self._lifecycle.activate() return await self._lifecycle.call(lambda: self._run_worker(now=now)) @@ -667,28 +744,29 @@ async def _run_worker(self, *, now: datetime | None) -> ReliableReportingTurn: # 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(), - key=lambda item: ( - item[0].account_id, - item[0].delivery_config_id, - 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, - ) + with source_turn(): + for key, binding in sorted( + self._bindings.items(), + key=lambda item: ( + item[0].account_id, + item[0].delivery_config_id, + 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, + ) extensions = ( ("materialization", self._materialization_worker), ("notification", self._notification_worker), @@ -728,6 +806,8 @@ def _default_caller(request: Any, context: Any | None) -> ReportingStatusCaller: async def get_reporting_status( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + if self._production is not None: + return _wire(await self._production.handler.get_reporting_status(request, context)) 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]: @@ -741,6 +821,8 @@ async def _get_reporting_status(self, request: Any, context: Any | None) -> dict async def sync_reporting_status( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + if self._production is not None: + return _wire(await self._production.handler.sync_reporting_status(request, context)) 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]: @@ -756,6 +838,8 @@ async def _sync_reporting_status(self, request: Any, context: Any | None) -> dic async def sync_reporting_receipts( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + if self._production is not None: + return _wire(await self._production.handler.sync_reporting_receipts(request, context)) 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]: @@ -769,6 +853,8 @@ async def _sync_reporting_receipts(self, request: Any, context: Any | None) -> d async def get_revision_content( self, request: Any, context: Any | None = None ) -> dict[str, Any]: + if self._production is not None: + return _wire(await self._production.handler.get_media_buy_delivery(request, context)) 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]: @@ -879,6 +965,9 @@ def inject_capabilities(self, response: Any) -> dict[str, Any]: def install(self, platform: Any) -> Any: """Install ready-to-use reporting handlers on an existing platform instance.""" + if self._production is not None: + self._production.handler.bind_application(platform) + return self._production.handler from adcp.server import ADCPHandler from adcp.server.mcp_tools import get_tools_for_handler diff --git a/tests/conformance/reporting/_inline_storage_worker.py b/tests/conformance/reporting/_inline_storage_worker.py new file mode 100644 index 000000000..227c2347e --- /dev/null +++ b/tests/conformance/reporting/_inline_storage_worker.py @@ -0,0 +1,115 @@ +"""Independent pool-of-one processes used by the inline storage contract tests.""" + +from __future__ import annotations + +import argparse +import asyncio +import base64 +import json +import os +from datetime import datetime, timezone +from pathlib import Path + +from psycopg_pool import AsyncConnectionPool + +from adcp.reporting.fixtures import redacted_capabilities, redacted_snapshot_request +from adcp.reporting.inline_source import InlineReportingSource +from adcp.reporting.inline_storage import PgReportingSealStore, PgReportingStagingStore + + +async def run(mode: str, schema: str, output: Path, counter: Path, value: int) -> None: + async with AsyncConnectionPool( + os.environ["ADCP_INLINE_TEST_DSN"], + kwargs={"options": f"-c search_path={schema}"}, + min_size=1, + max_size=1, + open=False, + ) as pool: + staging = PgReportingStagingStore(pool=pool) + seals = PgReportingSealStore(pool=pool) + if mode == "create": + await staging.create_schema() + output.write_text("{}") + return + if mode == "stage": + reference = await staging.stage( + account_id="shared-account", + source_execution_key=f"execution-{value}", + ordinal=value, + payload=b"same-content", + ) + output.write_text(json.dumps(reference)) + return + + async def fetch(_request: object) -> list[dict[str, object]]: + if mode == "replay": + raise AssertionError("a committed seal must replay without dispatch") + with counter.open("a") as stream: + stream.write("dispatch\n") + # Borrowing the sole connection inside the provider proves no + # storage transaction is held across the external source call. + async with pool.connection() as connection: + await connection.execute("SELECT 1") + # Both processes must reach dispatch before either may publish; + # they offer different valid manifests for the same execution key. + for _ in range(3000): + if counter.read_text().count("dispatch\n") == 2: + break + await asyncio.sleep(0.01) + else: + raise AssertionError("the other source process never reached dispatch") + return [ + { + "media_buy_id": "media-buy-redacted", + "campaign_id": "campaign-redacted-1", + "impressions": value, + "spend": "1.25", + } + ] + + source = InlineReportingSource( + capabilities=redacted_capabilities(), + fetch=fetch, + staging=staging, + seals=seals, + clock=lambda: datetime(2026, 11, 6, 12, tzinfo=timezone.utc), + ) + request = redacted_snapshot_request() + result = await source.execute(request, cancel=asyncio.Event()) + assert result.manifest_bytes is not None + from adcp.reporting.conformance import validate_reporting_source_execution + + manifest = await validate_reporting_source_execution( + capabilities=source.capabilities, + request=request, + result=result, + object_reader=staging, + ) + obj = manifest.objects[0] + payload = await staging.read( + object_ref=obj.object_ref, + object_generation=obj.object_generation, + account_id=request.identity.account_id, + source_scope={}, + cancel=asyncio.Event(), + ) + output.write_text( + json.dumps( + { + "manifest": base64.b64encode(result.manifest_bytes).decode(), + "payload": base64.b64encode(payload).decode(), + "reference": result.response.manifest.model_dump(), + } + ) + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("mode", choices=["create", "stage", "publish", "replay"]) + parser.add_argument("schema") + parser.add_argument("output", type=Path) + parser.add_argument("counter", type=Path) + parser.add_argument("--value", type=int, default=10) + args = parser.parse_args() + asyncio.run(run(args.mode, args.schema, args.output, args.counter, args.value)) diff --git a/tests/conformance/reporting/_production_context_source.py b/tests/conformance/reporting/_production_context_source.py new file mode 100644 index 000000000..255cf7812 --- /dev/null +++ b/tests/conformance/reporting/_production_context_source.py @@ -0,0 +1,54 @@ +"""Independent provider authorization persisted across the service restart.""" + +import json +import sqlite3 + +from adcp.reporting.ledger.models import ReportingConfigurationGenerationKey +from adcp.reporting.production.contracts import ReportingProductionSourceBinding + +from ._production_support import Source + + +class DurableBindingSource(Source): + def __init__(self, key, path, rows=None, **kwargs): + super().__init__(key, path, rows, **kwargs) + self.binding_path = path.with_suffix(".bindings") + with sqlite3.connect(self.binding_path) as db: + db.execute( + "CREATE TABLE IF NOT EXISTS bindings (account TEXT, config TEXT, version INTEGER," + " document TEXT, PRIMARY KEY(account,config,version))" + ) + + def bind_generation(self, configuration, **kwargs): + binding = super().bind_generation(configuration, **kwargs) + key = configuration.generation_key + with sqlite3.connect(self.binding_path) as db: + db.execute( + "INSERT OR IGNORE INTO bindings VALUES (?,?,?,?)", + ( + key.account_id, + key.delivery_config_id, + key.delivery_config_version, + json.dumps(binding.document()), + ), + ) + return binding + + def configuration_binding(self, configuration): + key = configuration.generation_key + with sqlite3.connect(self.binding_path) as db: + row = db.execute( + "SELECT document FROM bindings WHERE account=? AND config=? AND version=?", + (key.account_id, key.delivery_config_id, key.delivery_config_version), + ).fetchone() + if row is None: + return None + value = json.loads(row[0]) + return ReportingProductionSourceBinding( + ReportingConfigurationGenerationKey( + value["account_id"], value["delivery_config_id"], value["delivery_config_version"] + ), + value["capabilities_sha256"], + tuple(tuple(pair) for pair in value["media_buy_products"]), + configuration_sha256=value["configuration_sha256"], + ) diff --git a/tests/conformance/reporting/_production_context_worker.py b/tests/conformance/reporting/_production_context_worker.py new file mode 100644 index 000000000..f9252ab1d --- /dev/null +++ b/tests/conformance/reporting/_production_context_worker.py @@ -0,0 +1,55 @@ +"""Fresh process: rebuild the graph, then let the actual indexed worker recover.""" + +import asyncio +import json +import os +import sys +from pathlib import Path + +from psycopg_pool import AsyncConnectionPool + +from ._production_context_source import DurableBindingSource +from ._production_support import production_harness +from .test_reliable_reporting_production_admission import registry_factory, service_factory +from .test_reporting_production_bindings import source_documents + + +async def main(): + selected = {"context_change": {"adapter": "resolver-must-not-run"}} + async with AsyncConnectionPool( + os.environ["ADCP_PG_TEST_URL"], + open=False, + min_size=1, + max_size=1, + kwargs={ + "autocommit": True, + "options": "-csearch_path=" + os.environ["ADCP_CONTEXT_TEST_SCHEMA"], + }, + ) as pool: + await pool.wait() + async with production_harness( + "postgres", + Path(sys.argv[1]), + count=0, + source_publication=True, + existing_pool=pool, + source_factory=DurableBindingSource, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + requests = h.production.offerings[0].producer._source.requests + records = await source_documents(h) + print( + json.dumps( + { + "resolved": len(selected.get("resolved", ())), + "local_bindings": len(h.service._bindings), + "requests": [r.identity.account_id for r in requests], + "contexts": [r[3]["service_context_sha256"] for r in records], + "ready": h.service.ready, + } + ) + ) + + +asyncio.run(main()) diff --git a/tests/conformance/reporting/_production_support.py b/tests/conformance/reporting/_production_support.py index c6e78668d..db2faf139 100644 --- a/tests/conformance/reporting/_production_support.py +++ b/tests/conformance/reporting/_production_support.py @@ -374,6 +374,8 @@ async def production_harness( poll_seconds=60, source_factory=Source, adcp_version=None, + source_registry_factory=None, + service_factory=None, ): from contextlib import AsyncExitStack @@ -632,12 +634,14 @@ async def authorize(account, context, consumer): notification_workers=workers, poll_seconds=poll_seconds, adcp_version=adcp_version, + source_registry=source_registry_factory(offerings) if source_registry_factory else None, ) h.production, h.projection, h.item = support, projection, item + h.service = service_factory(support) if service_factory else None h.source_clock = source_clock h.mount = create_mcp_server(support.handler) try: - await support.start() + await (h.service.start() if h.service else support.start()) yield h finally: - await support.aclose() + await (h.service.close() if h.service else support.aclose()) diff --git a/tests/conformance/reporting/test_reliable_reporting_production_admission.py b/tests/conformance/reporting/test_reliable_reporting_production_admission.py new file mode 100644 index 000000000..f84ca1c8c --- /dev/null +++ b/tests/conformance/reporting/test_reliable_reporting_production_admission.py @@ -0,0 +1,591 @@ +"""Actual mounted typed admission and indexed recovery through the service bridge.""" + +import asyncio +import copy +import hashlib +import json +import os +import sys +from dataclasses import replace + +import pytest + +from adcp.reporting.canonical_json import canonical_json_utf8_v1 +from adcp.reporting.production import ( + ReportingConfigurationAdmission, + ReportingProductionSourceRegistry, +) +from adcp.reporting.production.handler import ReportingProductionHandler +from adcp.reporting.service import ReliableReportingService, ReportingAccountContext +from adcp.server.base import ADCPHandler +from adcp.server.mcp_tools import get_tools_for_handler +from adcp.server.responses import capabilities_response + +from ._generation_support import obligation_for +from ._production_context_source import DurableBindingSource +from ._production_support import production_harness +from ._production_transport import MountedProduction +from .test_reporting_production_bindings import source_documents +from .test_reporting_production_configuration import state_for, wire_configuration +from .test_reporting_production_lock_order import source_turn + + +class Application(ADCPHandler): + async def get_products(self, params, context=None): + return {"products": []} + + async def get_media_buy_delivery(self, params, context=None): + return {"media_buy_deliveries": [], "aggregated_totals": {"impressions": 17}} + + async def get_adcp_capabilities(self, params, context=None): + return capabilities_response(["media_buy"], sandbox=True) + + +def registry_factory(selected): + def make(offerings): + def resolve(config): + selected.setdefault("resolved", []).append(config.generation_key) + p = offerings[0].producer._offerings + context = ReportingAccountContext( + config.account_id, + "adapter", + p.currency, + p.source_scope, + account_timezone=config.account_timezone, + snapshot_offering_id=p.snapshot_offering_id, + official_offering_id=p.official_offering_id, + publication_namespace=p.publication_namespace, + requested_metrics=p.requested_metrics, + requested_dimensions=p.requested_dimensions, + slice_timeout=p.slice_timeout, + ) + return replace(context, **selected.get("context_change", {})) + + registry = ReportingProductionSourceRegistry(account_context=resolve) + registry.register("adapter", offerings[0].producer) + return registry + + return make + + +def service_factory(support): + service = ReliableReportingService.from_production(support) + handler = service.install(Application()) + assert handler is support.handler and type(handler) is ReportingProductionHandler + return service + + +async def account_task(selected, request, context, admit): + h = selected["h"] + wire = request["accounts"][0]["reporting_delivery_configs"][0] + await admit( + ReportingConfigurationAdmission( + h.production.offerings[0].offering_id, + h.item.config, + h.item.binding, + configuration_wire=wire, + ) + ) + return { + "accounts": [ + { + "account_id": h.item.config.account_id, + "brand": {"domain": "advertiser.example.test"}, + "operator": "buyer.example.test", + "action": "unchanged", + "status": "active", + "billing": "operator", + "timezone": "UTC", + "reporting_delivery_configs": [state_for(h, wire)], + } + ] + } + + +async def test_delegate_cannot_replace_production_wire_versions(tmp_path): + class LegacyApplication(Application): + async def get_adcp_capabilities(self, params, context=None): + return capabilities_response( + ["signals"], + major_versions=[2], + adcp_version="2.5", + supported_versions=["2.5"], + build_version="1.2.3", + sandbox=True, + idempotency={"supported": False}, + ) + + application = LegacyApplication() + assert not hasattr(application, "get_adcp_version") + + def install(support): + service = ReliableReportingService.from_production(support) + service.install(application) + return service + + async with production_harness( + "memory", + tmp_path / "delegated-version.sqlite", + count=0, + source_publication=True, + source_registry_factory=registry_factory({}), + service_factory=install, + ) as h: + response = await h.production.handler.get_adcp_capabilities({}) + assert response["adcp_version"] == h.production._protocol_version + assert response["adcp"]["supported_versions"] == [h.production._protocol_version] + assert response["adcp"]["major_versions"] == [3] + assert response["adcp"]["build_version"] == "1.2.3" + assert response["sandbox"] is True + assert response["supported_protocols"] == ["signals", "media_buy"] + assert await h.production.handler.get_products({}) == {"products": []} + mounted = MountedProduction(h) + mounted.authorize(h.item) + async with mounted.client() as client: + for transport in ("mcp", "a2a-0.3", "a2a-1.0"): + _, served = await mounted.call( + client, "get_adcp_capabilities", {}, transport=transport + ) + assert served.get("status") == "completed", served + assert served["adcp_version"] == h.production._protocol_version + assert served["adcp"]["supported_versions"] == [h.production._protocol_version] + assert served["adcp"]["major_versions"] == [3] + assert served["adcp"]["build_version"] == "1.2.3" + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_late_account_mounted_admission_freezes_context_and_uses_indexed_lease( + backend, tmp_path +): + selected = {} + + async def handle(request, context, admit): + return await account_task(selected, request, context, admit) + + async with production_harness( + backend, + tmp_path / "destination.sqlite", + count=0, + source_publication=True, + account_handler=handle, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + assert h.service.ready and not h.service._bindings + assert not await source_documents(h) + late = replace(h.item.config, account_id="late-account") + h.item.config = late + h.item.binding = replace(h.item.binding, generation_key=late.generation_key) + h.item.obligation = obligation_for(late) + h.item.writer.grant(h.item.binding) + producer = h.production.offerings[0].producer + producer._source.bind_generation(late) + selected["h"] = h + mounted = MountedProduction(h) + mounted.authorize(h.item) + request = { + "idempotency_key": "service-context-admission-0001", + "accounts": [ + { + "account": {"account_id": late.account_id}, + "reporting_delivery_configs": [wire_configuration(h)], + } + ], + } + async with mounted.client() as client: + for transport in ("mcp", "a2a-0.3", "a2a-1.0"): + _, response = await mounted.call( + client, "sync_accounts", request, transport=transport + ) + assert response.get("status") == "completed", response + frozen = await source_documents(h) + assert len(frozen) == 1 + doc = frozen[0][3] + assert doc["service_context"]["account_id"] == late.account_id + assert ( + doc["service_context_sha256"] + == hashlib.sha256(canonical_json_utf8_v1(doc["service_context"])).hexdigest() + ) + for change in ( + {"currency": "EUR"}, + {"adapter": "missing"}, + {"requested_metrics": ("impressions",)}, + ): + selected["context_change"] = change + _, denied = await mounted.call(client, "sync_accounts", request) + assert denied.get("status") != "completed" + assert await source_documents(h) == frozen + selected.pop("context_change") + result = await source_turn(h.production) + assert result.leased.account_id == late.account_id + assert len(result.revisions_committed) == 1 + assert producer._source.requests[-1].identity.account_id == late.account_id + names = { + t["name"] for t in get_tools_for_handler(h.production.handler, _include_schemas=False) + } + assert "get_products" in names and "create_media_buy" not in names + assert await h.production.handler.get_products({}) == {"products": []} + assert (await h.production.handler.get_media_buy_delivery({}))["aggregated_totals"][ + "impressions" + ] == 17 + before = len(selected["resolved"]) + assert ( + h.production._check_source_binding( + late, h.production.offerings[0]._producer_key, doc + ).document() + == doc + ) + assert len(selected["resolved"]) == before + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_failed_context_admission_rolls_back_configuration_and_destination(backend, tmp_path): + selected = {"context_change": {"currency": "EUR"}} + async with production_harness( + backend, + tmp_path / "rollback.sqlite", + count=0, + source_publication=True, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + new = replace(h.item.config, delivery_config_version=2) + binding = replace(h.item.binding, generation_key=new.generation_key) + h.item.writer.grant(binding) + h.production.offerings[0].producer._source.bind_generation(new) + old = h.item.config + h.item.config, h.item.binding = new, binding + value = ReportingConfigurationAdmission( + h.production.offerings[0].offering_id, + new, + binding, + configuration_wire=wire_configuration(h), + ) + with pytest.raises(ValueError, match="profile"): + await h.production._admit_configuration(value) + assert not await source_documents(h) + assert new.generation_key not in { + c.generation_key for c in await h.store.list_configurations(account_id=new.account_id) + } + assert old.generation_key != new.generation_key + + +async def test_context_insert_guard_rejects_tampering_and_keeps_legacy_documents(tmp_path): + raise_exception = pytest.importorskip("psycopg").errors.RaiseException + + from adcp.reporting.production.schema import validate_production_schema + + selected = {} + async with production_harness( + "postgres", + tmp_path / "guard.sqlite", + count=0, + source_publication=True, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + offering = h.production.offerings[0] + context = await h.production.source_registry.resolve( + h.item.config, offering, offering.check_source(effective=True) + ) + document = h.production._source_binding( + h.item.config, offering._producer_key, service_context=context + ).document() + identity = ( + h.item.config.account_id, + h.item.config.delivery_config_id, + h.item.config.delivery_config_version, + ) + async with h.pool.connection() as connection: + await validate_production_schema(connection, service_context=True) + for field, value in ( + ("service_context_sha256", "0" * 64), + ("account_id", "other"), + ("configuration_sha256", "0" * 64), + ("delivery_config_version", 9), + ): + mutated = {**document, field: value} + with pytest.raises(raise_exception, match="context is inconsistent"): + async with connection.transaction(): + await connection.execute( + "INSERT INTO reporting_production_generations" + " VALUES(%s,%s,%s,%s,%s::jsonb)", + (*identity, offering._producer_key, json.dumps(mutated)), + ) + for field, value in ( + ("account_id", "other"), + ("account_timezone", "Europe/Paris"), + ("version", 2), + ("capabilities_sha256", "0" * 64), + ): + mutated = copy.deepcopy(document) + mutated["service_context"][field] = value + mutated["service_context_sha256"] = hashlib.sha256( + canonical_json_utf8_v1(mutated["service_context"]) + ).hexdigest() + with pytest.raises(raise_exception, match="context is inconsistent"): + async with connection.transaction(): + await connection.execute( + "INSERT INTO reporting_production_generations" + " VALUES(%s,%s,%s,%s,%s::jsonb)", + (*identity, offering._producer_key, json.dumps(mutated)), + ) + malformed = copy.deepcopy(document) + del malformed["capabilities_sha256"] + del malformed["service_context"]["capabilities_sha256"] + malformed["service_context"]["unknown"] = 42 + malformed["service_context_sha256"] = hashlib.sha256( + canonical_json_utf8_v1(malformed["service_context"]) + ).hexdigest() + with pytest.raises(raise_exception, match="context is inconsistent"): + async with connection.transaction(): + await connection.execute( + "INSERT INTO reporting_production_generations" + " VALUES(%s,%s,%s,%s,%s::jsonb)", + (*identity, offering._producer_key, json.dumps(malformed)), + ) + await connection.execute( + "INSERT INTO reporting_production_generations VALUES(%s,%s,%s,%s,%s::jsonb)", + (*identity, offering._producer_key, json.dumps(document)), + ) + for statement in ( + "UPDATE reporting_production_generations SET source_binding=source_binding" + ' || \'{"service_context_sha256":"bad"}\'', + "DELETE FROM reporting_production_generations", + ): + with pytest.raises(pytest.importorskip("psycopg").errors.CheckViolation): + async with connection.transaction(): + await connection.execute(statement) + before = (await source_documents(h))[0][3] + await h.store.create_schema() + await validate_production_schema(connection, service_context=True) + assert (await source_documents(h))[0][3] == before + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_enrollment_fault_rolls_back_whole_admission(backend, tmp_path, monkeypatch): + selected = {} + async with production_harness( + backend, + tmp_path / "atomic.sqlite", + count=0, + source_publication=True, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + new = replace(h.item.config, delivery_config_version=2) + binding = replace(h.item.binding, generation_key=new.generation_key) + h.item.writer.grant(binding) + h.production.offerings[0].producer._source.bind_generation(new) + h.item.config, h.item.binding = new, binding + value = ReportingConfigurationAdmission( + h.production.offerings[0].offering_id, + new, + binding, + configuration_wire=wire_configuration(h), + ) + if backend == "memory": + original = h.store._enroll + + def broken(*args, **kwargs): + original(*args, **kwargs) + raise RuntimeError("injected after context enrollment") + + monkeypatch.setattr(h.store, "_enroll", broken) + else: + original = h.store._enroll_on + + async def broken(*args, **kwargs): + await original(*args, **kwargs) + raise RuntimeError("injected after context enrollment") + + monkeypatch.setattr(h.store, "_enroll_on", broken) + with pytest.raises(RuntimeError, match="injected"): + await h.production._admit_configuration(value) + assert not await source_documents(h) + assert new.generation_key not in { + c.generation_key for c in await h.store.list_configurations(account_id=new.account_id) + } + + +async def test_service_failure_stops_admission_but_preserves_ordinary_delegate(tmp_path): + from adcp.reporting.service import ( + ReliableReportingServiceError, + ReliableReportingUnavailableError, + ) + + async with production_harness( + "memory", + tmp_path / "lifetime.sqlite", + count=0, + source_registry_factory=registry_factory({}), + service_factory=service_factory, + ) as h: + task = h.production._task + task.cancel() + await asyncio.gather(task, return_exceptions=True) + with pytest.raises(ReliableReportingServiceError): + await asyncio.wait_for(h.service.wait(), 2) + assert not h.service.ready + with pytest.raises(ReliableReportingUnavailableError): + await h.production.handler.sync_accounts({}) + assert await h.production.handler.get_products({}) == {"products": []} + assert (await h.production.handler.get_media_buy_delivery({}))["aggregated_totals"][ + "impressions" + ] == 17 + + +async def test_fresh_process_recovers_late_account_without_context_resolver(tmp_path): + selected = {} + + async def handle(request, context, admit): + return await account_task(selected, request, context, admit) + + path = tmp_path / "restart.sqlite" + async with production_harness( + "postgres", + path, + count=0, + source_publication=True, + source_factory=DurableBindingSource, + account_handler=handle, + source_registry_factory=registry_factory(selected), + service_factory=service_factory, + ) as h: + late = replace(h.item.config, account_id="late-after-empty-start") + h.item.config, h.item.obligation = late, obligation_for(late) + h.item.binding = replace(h.item.binding, generation_key=late.generation_key) + h.item.writer.grant(h.item.binding) + h.production.offerings[0].producer._source.bind_generation(late) + selected["h"] = h + mounted = MountedProduction(h) + mounted.authorize(h.item) + async with mounted.client() as client: + _, result = await mounted.call( + client, + "sync_accounts", + { + "idempotency_key": "fresh-process-admission-0001", + "accounts": [ + { + "account": {"account_id": late.account_id}, + "reporting_delivery_configs": [wire_configuration(h)], + } + ], + }, + ) + assert result.get("status") == "completed", result + document = (await source_documents(h))[0][3] + await h.service.close() + async with h.pool.connection() as connection: + schema = (await (await connection.execute("SELECT current_schema()")).fetchone())[0] + env = {**os.environ, "ADCP_CONTEXT_TEST_SCHEMA": schema} + worker = await asyncio.create_subprocess_exec( + sys.executable, + "-m", + "tests.conformance.reporting._production_context_worker", + str(path), + env=env, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + out, err = await asyncio.wait_for(worker.communicate(), 90) + assert worker.returncode == 0, err.decode() + recovered = json.loads(out) + assert recovered == { + "resolved": 0, + "local_bindings": 0, + "requests": [late.account_id], + "contexts": [document["service_context_sha256"]], + "ready": True, + } + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_legacy_generation_is_never_guessed_or_backfilled(backend, tmp_path): + path = tmp_path / "legacy.sqlite" + async with production_harness(backend, path, count=0, source_publication=True) as original: + await original.production.activate(account_id=original.item.config.account_id) + before = await source_documents(original) + assert "service_context" not in before[0][3] + await original.production.aclose() + async with production_harness( + backend, + path, + count=0, + source_publication=True, + existing_store=original.store if backend == "memory" else None, + existing_pool=original.pool, + source_registry_factory=registry_factory({}), + service_factory=service_factory, + ) as h: + offering = h.production.offerings[0] + from adcp.reporting.materializer.contracts import ReportingWriterError + + with pytest.raises(ReportingWriterError) as refused: + h.production._check_source_binding( + h.item.config, offering._producer_key, before[0][3] + ) + assert refused.value.failure.code == "BINDING_MISMATCH" + assert offering.producer._source.requests == [] + value = ReportingConfigurationAdmission( + offering.offering_id, + h.item.config, + h.item.binding, + configuration_wire=wire_configuration(h), + ) + from adcp.reporting.ledger.notification_models import ReportingNotificationError + + with pytest.raises(ReportingNotificationError, match="source_conflict"): + await h.production._admit_configuration(value) + assert await source_documents(h) == before + + +async def test_delegate_collision_and_actual_mount_mutation_fail_closed(tmp_path): + class Conflicting(Application): + async def sync_accounts(self, params, context=None): + return {"accounts": []} + + def refused(support): + service = ReliableReportingService.from_production(support) + with pytest.raises(ValueError, match="conflict"): + service.install(Conflicting()) + service.install(Application()) + return service + + async with production_harness( + "memory", + tmp_path / "mount.sqlite", + count=0, + source_registry_factory=registry_factory({}), + service_factory=refused, + ) as h: + assert type(h.production.handler) is ReportingProductionHandler + assert h.production._mounted() + h.production.handler.sync_accounts = Application().sync_accounts + from adcp.reporting.ledger.notification_models import ReportingNotificationError + + with pytest.raises(ReportingNotificationError): + h.production._assert_components() + + +async def test_production_refuses_an_unused_local_adapter_registry(tmp_path): + from adcp.reporting.service import ReliableReportingConfigurationError + + def unused(support): + service = service_factory(support) + producer = support.offerings[0].producer + service.sources.register_executor( + "ignored", producer._source, object_reader=producer._object_reader + ) + return service + + with pytest.raises(ReliableReportingConfigurationError, match="frozen source registry"): + async with production_harness( + "memory", + tmp_path / "unused.sqlite", + count=0, + source_registry_factory=registry_factory({}), + service_factory=unused, + ): + pytest.fail("an unused adapter must not be silently accepted") diff --git a/tests/conformance/reporting/test_reporting_core_lifecycle.py b/tests/conformance/reporting/test_reporting_core_lifecycle.py index 505e330a8..3950a30d7 100644 --- a/tests/conformance/reporting/test_reporting_core_lifecycle.py +++ b/tests/conformance/reporting/test_reporting_core_lifecycle.py @@ -586,10 +586,10 @@ async def _restate( ) -> Any: """Ask for a new observation of an already-satisfied snapshot obligation. - A settled period is left alone by an ordinary worker turn, so a - restatement is explicit. The observation ordinal advances with the - committed revision count, which is what gives this acquisition a distinct - source execution key rather than replaying the sealed original. + Worker turns also refresh provisional snapshots on their declared cadence. + This explicit restatement uses the next observation ordinal, giving the + acquisition a distinct source execution key rather than replaying a + previously sealed snapshot. """ producer, _ = _producer(store, source, now=now) restated = await producer.acquire_obligation( diff --git a/tests/conformance/reporting/test_reporting_inline_storage.py b/tests/conformance/reporting/test_reporting_inline_storage.py new file mode 100644 index 000000000..c061578a4 --- /dev/null +++ b/tests/conformance/reporting/test_reporting_inline_storage.py @@ -0,0 +1,1371 @@ +"""Bounded storage, real database winners and committed-seal restart recovery. + +Set ADCP_PG_TEST_URL to an expendable database; each case owns a schema. +No provider, external endpoint, production data or service factory is involved. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import inspect +import json +import os +import subprocess +import sys +import traceback +import uuid +from contextlib import asynccontextmanager +from dataclasses import FrozenInstanceError, replace +from datetime import datetime, timezone +from pathlib import Path + +import anyio +import pytest + +from adcp.reporting.conformance import validate_reporting_source_execution +from adcp.reporting.fixtures import redacted_capabilities, redacted_snapshot_request +from adcp.reporting.inline_source import InlineReportingSource, SealedSlice +from adcp.reporting.inline_storage import ( + INLINE_STAGING_MAX_BYTES, + InlineStorageError, + PgReportingSealStore, + PgReportingStagingStore, + PreparedBackendObjectV1, + PreparedBackendSealV1, +) +from adcp.reporting.source import source_batch_manifest_reference_v1 + +NOW = datetime(2026, 11, 6, 12, tzinfo=timezone.utc) +ROW = { + "media_buy_id": "media-buy-redacted", + "campaign_id": "campaign-redacted-1", + "impressions": 10, + "spend": "1.25", +} + + +@pytest.fixture +async def database(): + dsn = os.environ.get("ADCP_PG_TEST_URL") + if not dsn: + pytest.skip("ADCP_PG_TEST_URL is not configured") + psycopg = pytest.importorskip("psycopg") + pools = pytest.importorskip("psycopg_pool") + schema = "inline_test_" + uuid.uuid4().hex + async with await psycopg.AsyncConnection.connect(dsn, autocommit=True) as admin: + await admin.execute( + psycopg.sql.SQL("CREATE SCHEMA {}").format(psycopg.sql.Identifier(schema)) + ) + try: + async with pools.AsyncConnectionPool( + dsn, + kwargs={"options": f"-c search_path={schema}"}, + min_size=1, + max_size=1, + open=False, + ) as pool: + yield pool, schema, dsn + finally: + await admin.execute( + psycopg.sql.SQL("DROP SCHEMA {} CASCADE").format(psycopg.sql.Identifier(schema)) + ) + + +@pytest.fixture +async def stores(database): + pool, _, _ = database + staging = PgReportingStagingStore(pool=pool) + seals = PgReportingSealStore(pool=pool) + await staging.create_schema() + return staging, seals, pool + + +def source(staging, seals, fetch=None): + return InlineReportingSource( + capabilities=redacted_capabilities(), + fetch=fetch or (lambda _: [ROW]), + staging=staging, + seals=seals, + clock=lambda: NOW, + ) + + +async def sample_seal(value=10): + result = await InlineReportingSource( + capabilities=redacted_capabilities(), + fetch=lambda _: [{**ROW, "impressions": value}], + clock=lambda: NOW, + ).execute(redacted_snapshot_request(), cancel=asyncio.Event()) + assert result.manifest_bytes is not None + return SealedSlice(reference=result.response.manifest, manifest_bytes=result.manifest_bytes) + + +def test_driver_absent_import_and_constructor() -> None: + script = """ +import importlib.abc, sys +class MissingDriver(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path=None, target=None): + if fullname.split('.')[0] in {'psycopg', 'psycopg_pool'}: + raise ModuleNotFoundError('driver deliberately absent') +sys.meta_path.insert(0, MissingDriver()) +from adcp.reporting.inline_storage import PgReportingStagingStore, PgReportingSealStore +for cls in [PgReportingStagingStore, PgReportingSealStore]: + try: + cls(pool=object()) + except ImportError as error: + assert str(error) == 'PostgreSQL inline storage requires adcp[pg]' + assert error.__context__ is None + else: + raise AssertionError('construction requires the optional driver') +""" + completed = subprocess.run( + [sys.executable, "-c", script], capture_output=True, text=True, timeout=60 + ) + assert completed.returncode == 0, completed.stderr + + +@pytest.mark.parametrize("limit", [True, 0, -1, INLINE_STAGING_MAX_BYTES + 1, "secret"]) +def test_invalid_limits_are_closed_before_driver_or_pool_access(limit) -> None: + with pytest.raises(InlineStorageError) as caught: + PgReportingStagingStore(pool=object(), max_payload_bytes=limit) + assert caught.value.code == "INVALID_INPUT" + assert "secret" not in repr(caught.value) + + +async def test_readiness_missing_drift_and_unrelated_ddl(database) -> None: + pool, _, _ = database + store = PgReportingStagingStore(pool=pool) + with pytest.raises(InlineStorageError, match="schema is not ready"): + await store.check_ready() + await store.create_schema() + await store.create_schema() + async with pool.connection() as conn: + await conn.execute("CREATE TABLE adopter_extra (value TEXT)") + await store.check_ready() + async with pool.connection() as conn: + await conn.execute( + "ALTER TABLE reporting_inline_objects DROP CONSTRAINT reporting_inline_objects_bytes" + ) + for operation in (store.check_ready, store.create_schema): + with pytest.raises(InlineStorageError) as caught: + await operation() + assert caught.value.code == "SCHEMA_UNREADY" + assert caught.value.__context__ is None + + +@pytest.mark.parametrize( + "ddl", + [ + "ALTER TABLE reporting_inline_objects DISABLE TRIGGER reporting_inline_objects_immutable", + "ALTER TABLE reporting_inline_seals ALTER COLUMN manifest DROP NOT NULL", + "ALTER TABLE reporting_inline_seals DROP CONSTRAINT reporting_inline_seals_pk", + "ALTER TABLE reporting_inline_objects SET UNLOGGED", + ], +) +async def test_readiness_rejects_schema_drift(stores, ddl) -> None: + store, _, pool = stores + async with pool.connection() as conn: + await conn.execute(ddl) + with pytest.raises(InlineStorageError) as caught: + await store.check_ready() + assert caught.value.code == "SCHEMA_UNREADY" + + +@pytest.mark.parametrize("autocommit", [False, True]) +async def test_borrowed_pool_with_serializable_default(database, autocommit) -> None: + from psycopg_pool import AsyncConnectionPool + + _, schema, dsn = database + async with AsyncConnectionPool( + dsn, + min_size=1, + max_size=1, + open=False, + kwargs={ + "autocommit": autocommit, + "options": f"-c search_path={schema} -c default_transaction_isolation=serializable", + }, + ) as pool: + staging = PgReportingStagingStore(pool=pool) + await staging.create_schema() + args = dict(account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"bytes") + assert await staging.stage(**args) == await staging.stage(**args) + async with pool.connection() as conn: + assert (await (await conn.execute("SHOW default_transaction_isolation")).fetchone())[ + 0 + ] == "serializable" + + +async def test_account_content_identity_exact_bytes_and_bounds(stores) -> None: + staging, _, pool = stores + first = await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"\x00\xff\n" + ) + same = await staging.stage( + account_id="a", source_execution_key="execution-2", ordinal=99999, payload=b"\x00\xff\n" + ) + other = await staging.stage( + account_id="b", source_execution_key="execution-1", ordinal=0, payload=b"\x00\xff\n" + ) + assert first == same and other[0] != first[0] and other[1] == first[1] + assert first[1] == hashlib.sha256(b"\x00\xff\n").hexdigest() + assert "execution" not in first[0] + assert ( + await staging.read( + object_ref=first[0], + object_generation=first[1], + account_id="a", + source_scope={"ignored": "secret"}, + cancel=asyncio.Event(), + ) + == b"\x00\xff\n" + ) + with pytest.raises(InlineStorageError) as caught: + await staging.read( + object_ref=first[0], + object_generation=first[1], + account_id="b", + source_scope={}, + cancel=asyncio.Event(), + ) + assert caught.value.code == "NOT_FOUND" + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 2 + bounded = PgReportingStagingStore(pool=pool, max_payload_bytes=3) + await bounded.stage( + account_id="a", source_execution_key="execution-3", ordinal=0, payload=b"123" + ) + with pytest.raises(InlineStorageError) as caught: + await bounded.stage( + account_id="a", source_execution_key="execution-3", ordinal=0, payload=b"1234" + ) + assert caught.value.code == "INVALID_INPUT" + empty = await staging.stage( + account_id="a", source_execution_key="execution-4", ordinal=0, payload=b"" + ) + assert ( + await staging.read( + object_ref=empty[0], + object_generation=empty[1], + account_id="a", + source_scope={}, + cancel=asyncio.Event(), + ) + == b"" + ) + + +async def test_hard_payload_cap_and_content_winner_verification(stores, monkeypatch) -> None: + staging, _, pool = stores + args = dict(account_id="a", source_execution_key="execution-1", ordinal=0) + with pytest.raises(InlineStorageError) as caught: + await staging.stage(**args, payload=b"x" * (INLINE_STAGING_MAX_BYTES + 1)) + assert caught.value.code == "INVALID_INPUT" + original = staging._read_on + + async def inconsistent_winner(*args, **kwargs): + await original(*args, **kwargs) + return b"different-winning-content" + + with monkeypatch.context() as patch: + patch.setattr(staging, "_read_on", inconsistent_winner) + with pytest.raises(InlineStorageError) as caught: + await staging.stage(**args, payload=b"candidate") + assert caught.value.code == "INTEGRITY_FAILED" + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 0 + + +@pytest.mark.parametrize( + "update", + [ + {"account_id": ""}, + {"account_id": "a" * 256}, + {"account_id": "\ud800"}, + {"account_id": "a\x00"}, + {"source_execution_key": "short"}, + {"source_execution_key": "invalid/key"}, + {"ordinal": True}, + {"ordinal": -1}, + {"ordinal": 100000}, + {"payload": bytearray(b"secret")}, + ], +) +async def test_staging_invalid_input_is_redacted(stores, update) -> None: + staging, _, _ = stores + arguments = dict( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"payload" + ) + arguments.update(update) + with pytest.raises(InlineStorageError) as caught: + await staging.stage(**arguments) + assert caught.value.code == "INVALID_INPUT" and caught.value.__context__ is None + assert "secret" not in repr(caught.value) + + +async def test_seal_winner_validation_and_detachment(stores) -> None: + _, seals, _ = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + first, second = await sample_seal(10), await sample_seal(11) + winner = await seals.put(**args, sealed=first) + assert winner == first and first == winner + assert winner.reference == first.reference and winner.reference is not first.reference + assert winner.manifest_bytes == first.manifest_bytes + assert (await seals.put(**args, sealed=second)).manifest_bytes == first.manifest_bytes + assert (await seals.get(**args)).manifest_bytes == first.manifest_bytes + assert ( + await seals.get( + account_id="other-account", source_execution_key=identity.source_execution_key + ) + is None + ) + assert "redacted" in repr(winner) and "manifest_version" not in repr(winner) + with pytest.raises(Exception): + winner.reference.byte_count = 1 + object.__setattr__(winner.reference, "byte_count", 1) + assert (await seals.get(**args)).reference.byte_count == len(first.manifest_bytes) + + +@pytest.mark.parametrize( + "fault", ["digest", "count", "bytes", "canonical", "oversize", "account", "key", "constructed"] +) +async def test_seal_refuses_invalid_candidate_even_when_winner_exists(stores, fault) -> None: + _, seals, _ = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + good = await sample_seal() + await seals.put(**args, sealed=good) + raw, ref = good.manifest_bytes, good.reference + if fault == "digest": + ref = ref.model_copy(update={"manifest_sha256": "0" * 64}) + elif fault == "count": + ref = ref.model_copy(update={"byte_count": 1}) + elif fault == "bytes": + raw = b"secret" + elif fault == "canonical": + raw = json.dumps(json.loads(raw), indent=2).encode() + ref = source_batch_manifest_reference_v1("reference", raw) + elif fault == "oversize": + raw = b"x" * (1048576 + 1) + elif fault in {"account", "key"}: + args["account_id" if fault == "account" else "source_execution_key"] = "other-identity" + else: + ref = ref.model_copy(update={"encoding": "secret"}) + with pytest.raises(InlineStorageError) as caught: + await seals.put(**args, sealed=SealedSlice(reference=ref, manifest_bytes=raw)) + assert caught.value.code == "INVALID_INPUT" and caught.value.__context__ is None + assert "secret" not in str(caught.value) + + +async def test_corrupt_seal_refused_with_unchanged_valid_schema(stores) -> None: + _, seals, pool = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + good = await sample_seal() + await seals.put(**args, sealed=good) + # An administrative corruption preserves all SQL checks and restores the + # immutable trigger, but destroys the source manifest's semantic validity. + raw = json.dumps( + { + "identity": { + "account_id": identity.account_id, + "source_execution_key": identity.source_execution_key, + } + } + ).encode() + async with pool.connection() as conn, conn.transaction(): + await conn.execute( + "ALTER TABLE reporting_inline_seals DISABLE TRIGGER reporting_inline_seals_immutable" + ) + await conn.execute( + "UPDATE reporting_inline_seals SET manifest=%s, manifest_sha256=%s, byte_count=%s", + (raw, hashlib.sha256(raw).hexdigest(), len(raw)), + ) + await conn.execute( + "ALTER TABLE reporting_inline_seals ENABLE TRIGGER reporting_inline_seals_immutable" + ) + await seals.check_ready() + for operation in (lambda: seals.get(**args), lambda: seals.put(**args, sealed=good)): + with pytest.raises(InlineStorageError) as caught: + await operation() + assert caught.value.code == "INTEGRITY_FAILED" + assert caught.value.__context__ is None and caught.value.__cause__ is None + assert "account-redacted" not in repr(caught.value) + + +@pytest.mark.parametrize( + "sql", + [ + "DELETE FROM reporting_inline_objects", + "UPDATE reporting_inline_objects SET payload=payload", + "TRUNCATE reporting_inline_objects", + "DELETE FROM reporting_inline_seals", + ], +) +async def test_sql_write_guards_are_immutable(stores, sql) -> None: + _, _, pool = stores + with pytest.raises(Exception, match="storage is immutable"): + async with pool.connection() as conn: + await conn.execute(sql) + async with pool.connection() as conn: + assert (await (await conn.execute("SELECT 1")).fetchone())[0] == 1 + + +async def test_fault_before_commit_rolls_back_and_after_commit_resumes(stores, monkeypatch) -> None: + staging, _, pool = stores + args = dict( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"secret-payload" + ) + original_read = staging._read_on + + async def fault(*args, **kwargs): + await original_read(*args, **kwargs) + await args[0].execute("SELECT 'secret-dsn-and-query'::integer") + + with monkeypatch.context() as patch: + patch.setattr(staging, "_read_on", fault) + with pytest.raises(InlineStorageError) as caught: + await staging.stage(**args) + assert caught.value.code == "RESOURCE_UNAVAILABLE" and caught.value.__context__ is None + assert "secret" not in repr(caught.value) + repr(staging) + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 0 + original_transaction = staging._transaction + + @asynccontextmanager + async def fail_after_commit(): + async with original_transaction() as conn: + yield conn + raise RuntimeError("secret-commit-acknowledgement-lost") + + with monkeypatch.context() as patch: + patch.setattr(staging, "_transaction", fail_after_commit) + with pytest.raises(InlineStorageError) as caught: + await staging.stage(**args) + assert caught.value.code == "RESOURCE_UNAVAILABLE" and caught.value.__context__ is None + assert "secret" not in repr(caught.value) + await staging.stage(**args) + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 1 + + +async def test_cancellation_while_waiting_for_sql_settles_pool_one(stores, database) -> None: + import psycopg + + staging, _, pool = stores + _, schema, dsn = database + async with await psycopg.AsyncConnection.connect( + dsn, options=f"-c search_path={schema}" + ) as blocker: + await blocker.execute("LOCK TABLE reporting_inline_objects IN ACCESS EXCLUSIVE MODE") + sdk_errors = [] + + async def stage_and_capture(): + try: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"bytes" + ) + except asyncio.CancelledError as error: + sdk_errors.append(error) + raise + + task = asyncio.create_task(stage_and_capture()) + async with await psycopg.AsyncConnection.connect(dsn, autocommit=True) as observer: + for _ in range(200): + row = await ( + await observer.execute( + "SELECT count(*) FROM pg_stat_activity WHERE wait_event_type='Lock' " + "AND query LIKE 'LOCK TABLE reporting_inline_objects,%'" + ) + ).fetchone() + if row[0]: + break + await asyncio.sleep(0.01) + else: + pytest.fail("store never reached the blocked SQL statement") + task.cancel("secret-cancellation-message") + done, pending = await asyncio.wait({task}, timeout=5) + if pending: + await blocker.rollback() + task.cancel() + await asyncio.gather(task, return_exceptions=True) + pytest.fail("canceled storage operation did not settle within five seconds") + assert done == {task} + with pytest.raises(asyncio.CancelledError) as caught: + await task + # Inspect the SDK exception before Python 3.10's Task boundary adds + # another empty CancelledError around it. Neither boundary may retain + # the caller's cancellation message or a driver exception. + assert len(sdk_errors) == 1 + assert sdk_errors[0].__context__ is None and sdk_errors[0].__cause__ is None + assert not sdk_errors[0].args + error = caught.value + seen = [] + while error is not None: + assert type(error) is asyncio.CancelledError and not error.args + assert error.__cause__ is None and error not in seen + seen.append(error) + error = error.__context__ + # The borrowed pool is still open, its only connection reusable, and no + # detached operation commits after the canceled call returns. + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 0 + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"bytes" + ) + + +async def test_read_cancel_event(stores) -> None: + staging, _, _ = stores + ref, generation = await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"bytes" + ) + cancel = asyncio.Event() + cancel.set() + with pytest.raises(asyncio.CancelledError): + await staging.read( + object_ref=ref, + object_generation=generation, + account_id="a", + source_scope={}, + cancel=cancel, + ) + + +async def worker(database, tmp_path, mode, name, value=10): + _, schema, dsn = database + output = tmp_path / f"{name}.json" + process = await asyncio.create_subprocess_exec( + sys.executable, + str(Path(__file__).with_name("_inline_storage_worker.py")), + mode, + schema, + str(output), + str(tmp_path / "dispatches"), + "--value", + str(value), + env={**os.environ, "ADCP_INLINE_TEST_DSN": dsn}, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + try: + stdout, stderr = await asyncio.wait_for(process.communicate(), 90) + except BaseException: + if process.returncode is None: + process.kill() + await process.wait() + raise + assert process.returncode == 0, (stdout.decode(), stderr.decode()) + return json.loads(output.read_text()) + + +async def test_independent_process_schema_boot_race(database, tmp_path) -> None: + await asyncio.gather(*(worker(database, tmp_path, "create", f"boot-{i}") for i in range(2))) + await PgReportingStagingStore(pool=database[0]).check_ready() + + +async def test_independent_process_staging_deduplicates(stores, database, tmp_path) -> None: + outcomes = await asyncio.gather( + *(worker(database, tmp_path, "stage", f"stage-{i}", i) for i in range(2)) + ) + assert outcomes[0] == outcomes[1] + async with database[0].connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 1 + + +async def test_fresh_process_replay_without_fetch_and_concurrent_winners( + stores, database, tmp_path +) -> None: + outcomes = await asyncio.gather( + *(worker(database, tmp_path, "publish", f"publish-{i}", i + 10) for i in range(2)) + ) + assert outcomes[0] == outcomes[1] + before = (tmp_path / "dispatches").read_bytes() + replay = await worker(database, tmp_path, "replay", "restart") + assert replay == outcomes[0] + assert (tmp_path / "dispatches").read_bytes() == before + assert before.count(b"dispatch\n") == 2 + + +async def test_committed_seal_replays_after_cancel_without_dispatch(stores) -> None: + staging, seals, _ = stores + calls = [] + + def fetch(_): + calls.append(1) + return [ROW] + + inline = source(staging, seals, fetch) + request = redacted_snapshot_request() + original = await inline.execute(request, cancel=asyncio.Event()) + cancel = asyncio.Event() + cancel.set() + replay = await inline.execute(request, cancel=cancel) + assert replay.manifest_bytes == original.manifest_bytes + assert len(calls) == 1 + + +async def test_changed_request_remains_a_conformance_responsibility(stores) -> None: + staging, seals, _ = stores + request = redacted_snapshot_request() + inline = source(staging, seals) + await inline.execute(request, cancel=asyncio.Event()) + changed = request.model_copy(update={"currency": "EUR"}) + replay = await inline.execute(changed, cancel=asyncio.Event()) + with pytest.raises(Exception, match="currency"): + await validate_reporting_source_execution( + capabilities=inline.capabilities, request=changed, result=replay, object_reader=staging + ) + + +async def test_stage_before_seal_failure_can_refetch(stores, monkeypatch) -> None: + staging, seals, pool = stores + calls = [] + + def fetch(_): + calls.append(1) + return [ROW] + + inline = source(staging, seals, fetch) + request = redacted_snapshot_request() + + async def fail_seal(**kwargs): + raise InlineStorageError("RESOURCE_UNAVAILABLE") + + with monkeypatch.context() as patch: + patch.setattr(seals, "put", fail_seal) + with pytest.raises(InlineStorageError): + await inline.execute(request, cancel=asyncio.Event()) + async with pool.connection() as conn: + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_objects")).fetchone() + )[0] == 1 + assert ( + await (await conn.execute("SELECT count(*) FROM reporting_inline_seals")).fetchone() + )[0] == 0 + await inline.execute(request, cancel=asyncio.Event()) + assert len(calls) == 2 + + +@pytest.mark.parametrize("operation", ["stage", "read", "get", "put"]) +async def test_unmigrated_operations_are_schema_unready(database, operation) -> None: + pool, _, _ = database + staging, seals = PgReportingStagingStore(pool=pool), PgReportingSealStore(pool=pool) + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + digest = hashlib.sha256(b"payload").hexdigest() + reference = f"pg-inline-v1.{hashlib.sha256(identity.account_id.encode()).hexdigest()}.{digest}" + calls = { + "stage": lambda: staging.stage(**args, ordinal=0, payload=b"payload"), + "read": lambda: staging.read( + account_id=identity.account_id, + object_ref=reference, + object_generation=digest, + source_scope={}, + cancel=asyncio.Event(), + ), + "get": lambda: seals.get(**args), + "put": lambda: seals.put(**args, sealed=sealed), + } + sealed = await sample_seal() + with pytest.raises(InlineStorageError) as caught: + await calls[operation]() + assert caught.value.code == "SCHEMA_UNREADY" + assert caught.value.__context__ is None and caught.value.__cause__ is None + await staging.create_schema() + await staging.check_ready() + + +async def test_catalog_additive_metadata_does_not_change_readiness(stores, monkeypatch) -> None: + import adcp.reporting.inline_storage as module + + staging, _, _ = stores + original = module.schema_objects + + async def with_metadata(connection): + return { + key: {**value, "adopter_metadata": "secret"} + for key, value in (await original(connection)).items() + } + + monkeypatch.setattr(module, "schema_objects", with_metadata) + await staging.check_ready() + await staging.stage(account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"x") + + +@pytest.mark.parametrize("field", ["fingerprint", "enabled"]) +async def test_catalog_required_fields_still_fail_closed(stores, monkeypatch, field) -> None: + import adcp.reporting.inline_storage as module + + staging, _, _ = stores + original = module.schema_objects + + async def with_missing_field(connection): + objects = await original(connection) + objects[next(iter(module.REQUIRED_OBJECTS))].pop(field) + return objects + + monkeypatch.setattr(module, "schema_objects", with_missing_field) + with pytest.raises(InlineStorageError) as caught: + await staging.check_ready() + assert caught.value.code == "SCHEMA_UNREADY" and caught.value.__context__ is None + + +def assert_framework_cancellation(error) -> None: + assert error.args in ( + ("Cancelled by cancel scope [redacted]",), + ("Cancelled via cancel scope [redacted]",), + ) + assert error.__context__ is None and error.__cause__ is None + + +@pytest.mark.parametrize("kind", ["fail_after", "move_on_after"]) +@pytest.mark.parametrize("blocked_at", ["pool", "sql"]) +async def test_anyio_deadline_contains_redacted_cancellation( + stores, database, kind, blocked_at, caplog +) -> None: + import psycopg + + staging, _, pool = stores + _, schema, dsn = database + calls = 0 + + async def operation(): + nonlocal calls + try: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"x" + ) + except asyncio.CancelledError as error: + calls += 1 + assert_framework_cancellation(error) + raise + + async def timed(): + if kind == "fail_after": + with pytest.raises(TimeoutError): + with anyio.fail_after(0.025) as scope: + await operation() + else: + with anyio.move_on_after(0.025) as scope: + await operation() + assert scope.cancel_called and scope.cancelled_caught + + async with pool.connection() as connection: + original_backend = connection.info.backend_pid + # Repeat using the same size-one pool to expose unsettled borrowers. + for _ in range(2): + if blocked_at == "pool": + async with pool.connection(): + await timed() + else: + async with await psycopg.AsyncConnection.connect( + dsn, options=f"-c search_path={schema}" + ) as blocker: + await blocker.execute( + "LOCK TABLE reporting_inline_objects IN ACCESS EXCLUSIVE MODE" + ) + await timed() + await blocker.rollback() + async with pool.connection() as connection: + assert (await (await connection.execute("SELECT 1")).fetchone())[0] == 1 + assert connection.info.backend_pid == original_backend + assert calls == 2 + assert not [record for record in caplog.records if record.name.startswith("psycopg")] + + +async def test_anyio_nested_deadline_reaches_owning_scope(stores) -> None: + staging, _, pool = stores + async with pool.connection(): + with pytest.raises(TimeoutError): + with anyio.fail_after(0.025) as outer: + with anyio.move_on_after(1) as inner: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"x" + ) + assert outer.cancelled_caught and not inner.cancelled_caught + async with pool.connection() as connection: + assert (await (await connection.execute("SELECT 1")).fetchone())[0] == 1 + + +async def test_anyio_cancel_reason_and_task_name_are_not_exposed(stores) -> None: + staging, _, pool = stores + task = asyncio.current_task() + previous = task.get_name() + task.set_name("secret-task-name") + captured = False + try: + async with pool.connection(): + with anyio.CancelScope() as scope: + if "reason" in inspect.signature(scope.cancel).parameters: + scope.cancel("secret-cancellation-reason") + else: + scope.cancel() + try: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"x" + ) + except asyncio.CancelledError as error: + captured = True + assert_framework_cancellation(error) + assert "secret" not in repr(error) + raise + assert captured and scope.cancelled_caught + finally: + task.set_name(previous) + async with pool.connection() as connection: + assert (await (await connection.execute("SELECT 1")).fetchone())[0] == 1 + + +@pytest.mark.parametrize("prefix", ["Cancelled by cancel scope ", "Cancelled via cancel scope "]) +async def test_forged_framework_marker_without_cancelled_scope_stays_argless( + stores, monkeypatch, prefix +) -> None: + staging, _, _ = stores + + @asynccontextmanager + async def forged_transaction(): + raise asyncio.CancelledError(prefix + "secret-provider-text") + yield + + monkeypatch.setattr(staging, "_transaction", forged_transaction) + with anyio.CancelScope(): + with pytest.raises(asyncio.CancelledError) as caught: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"x" + ) + assert not caught.value.args + assert caught.value.__context__ is None and caught.value.__cause__ is None + + +async def test_large_payload_hash_yields_and_matches_exact_bytes(stores, monkeypatch) -> None: + import adcp.reporting.inline_storage as module + + staging, _, _ = stores + payload = b"x" * INLINE_STAGING_MAX_BYTES + original_digest = module._payload_digest + checked = 0 + + async def observe_digest(payload): + nonlocal checked + advanced = asyncio.Event() + asyncio.get_running_loop().call_soon(advanced.set) + result = await original_digest(payload) + assert advanced.is_set(), "a scheduled callback must run during payload hashing" + checked += 1 + return result + + monkeypatch.setattr(module, "_payload_digest", observe_digest) + reference, digest = await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=payload + ) + assert digest == hashlib.sha256(payload).hexdigest() + assert ( + await staging.read( + account_id="a", + object_ref=reference, + object_generation=digest, + source_scope={}, + cancel=asyncio.Event(), + ) + == payload + ) + assert checked == 3 + + +@pytest.mark.parametrize("during", ["before_transaction", "winner_verification"]) +async def test_cancellation_during_payload_hashing_settles_without_partial_write( + stores, monkeypatch, during +) -> None: + staging, _, pool = stores + payload = b"x" * INLINE_STAGING_MAX_BYTES + import adcp.reporting.inline_storage as module + + original_digest = module._payload_digest + original_transaction = staging._transaction + entered = False + canceled_in_hash = False + digest_calls = 0 + + @asynccontextmanager + async def track_transaction(): + nonlocal entered + entered = True + async with original_transaction() as connection: + yield connection + + async def hash_with_cancellation(payload): + nonlocal canceled_in_hash, digest_calls + digest_calls += 1 + if digest_calls == (1 if during == "before_transaction" else 2): + # The winning row is already read inside the write transaction; cancel + # at the real cooperative hash checkpoint before commit. + asyncio.get_running_loop().call_soon( + asyncio.current_task().cancel, "secret-hash-cancel" + ) + canceled_in_hash = True + return await original_digest(payload) + + monkeypatch.setattr(staging, "_transaction", track_transaction) + monkeypatch.setattr(module, "_payload_digest", hash_with_cancellation) + with pytest.raises(asyncio.CancelledError) as caught: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=payload + ) + assert not caught.value.args and caught.value.__context__ is None + assert entered == (during == "winner_verification") + assert canceled_in_hash + async with pool.connection() as connection: + assert ( + await ( + await connection.execute("SELECT count(*) FROM reporting_inline_objects") + ).fetchone() + )[0] == 0 + + +@pytest.mark.parametrize("phase", ["in_transaction", "after_commit"]) +async def test_repeated_task_cancel_waits_for_owned_settlement( + stores, monkeypatch, caplog, phase +) -> None: + staging, _, pool = stores + entered, cancel_seen, release, settled = (asyncio.Event() for _ in range(4)) + cancellations = 0 + original = staging._transaction + loop = asyncio.get_running_loop() + async with pool.connection() as connection: + backend_pid = connection.info.backend_pid + + async def hold(): + nonlocal cancellations + entered.set() + try: + await release.wait() + except asyncio.CancelledError: + cancellations += 1 + cancel_seen.set() + await release.wait() + raise + + @asynccontextmanager + async def delayed_settlement(): + try: + async with original() as connection: + assert asyncio.get_running_loop() is loop + assert connection.info.backend_pid == backend_pid + yield connection + if phase == "in_transaction": + await hold() + if phase == "after_commit": + await hold() + finally: + settled.set() + + monkeypatch.setattr(staging, "_transaction", delayed_settlement) + + async def call(): + try: + await staging.stage( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"payload" + ) + except asyncio.CancelledError as error: + assert not error.args and error.__context__ is None and error.__cause__ is None + assert settled.is_set() + raise + + acquired = asyncio.Event() + + async def borrower(): + async with pool.connection() as connection: + acquired.set() + assert connection.info.backend_pid == backend_pid + + operation = asyncio.create_task(call()) + waiting = None + try: + await asyncio.wait_for(entered.wait(), 5) + operation.cancel("secret-first-cancel") + await asyncio.wait_for(cancel_seen.wait(), 5) + waiting = asyncio.create_task(borrower()) + for reason in ("secret-second-cancel", "secret-third-cancel"): + operation.cancel(reason) + await asyncio.sleep(0) + assert not operation.done() and not settled.is_set() + assert acquired.is_set() == (phase == "after_commit") + release.set() + _, pending = await asyncio.wait({operation, waiting}, timeout=5) + assert not pending + assert operation.cancelled() and waiting.exception() is None + assert cancellations == 1 and settled.is_set() + finally: + release.set() + if not operation.done(): + operation.cancel() + await asyncio.gather(operation, *([waiting] if waiting else []), return_exceptions=True) + async with pool.connection() as connection: + assert connection.info.backend_pid == backend_pid + count = ( + await ( + await connection.execute("SELECT count(*) FROM reporting_inline_objects") + ).fetchone() + )[0] + assert count == (1 if phase == "after_commit" else 0) + assert not [record for record in caplog.records if record.name.startswith("psycopg")] + + +@pytest.mark.parametrize("factory", ["native", "pure_python"]) +@pytest.mark.parametrize("outcome", ["driver_error", "completed_cancel_race"]) +async def test_task_boundary_never_transports_raw_exception_context( + stores, monkeypatch, factory, outcome +) -> None: + staging, _, pool = stores + loop = asyncio.get_running_loop() + previous_factory = loop.get_task_factory() + marker = "private-" + uuid.uuid4().hex + original_read = staging._read_on + original_transaction = staging._transaction + caller = None + committed = False + + async def driver_error(*args, **kwargs): + await original_read(*args, **kwargs) + query = "SELECT '" + marker + "'::integer" + await args[0].execute(query) + + @asynccontextmanager + async def completed_cancel_race(): + nonlocal committed + async with original_transaction() as connection: + yield connection + committed = True + # Schedule cancellation before the task's done callback wakes its caller. + # The operation has returned its borrowed connection and will be done by + # the time the caller receives cancellation: joining it need not suspend. + loop.call_soon(caller.cancel, marker) + + async def call_public_boundary(): + nonlocal caller + caller = asyncio.current_task() + try: + await staging.stage( + account_id="a", + source_execution_key="execution-1", + ordinal=0, + payload=marker.encode(), + ) + except (InlineStorageError, asyncio.CancelledError) as error: + assert error.__context__ is None and error.__cause__ is None + chain = "".join(traceback.format_exception(type(error), error, error.__traceback__)) + assert marker not in chain + if outcome == "driver_error": + assert isinstance(error, InlineStorageError) + assert error.code == "RESOURCE_UNAVAILABLE" + else: + assert isinstance(error, asyncio.CancelledError) and not error.args + assert committed + return + raise AssertionError("the public call must fail or propagate cancellation") + + def make_task(loop, coroutine, **kwargs): + if factory == "pure_python": + return asyncio.tasks._PyTask(coroutine, loop=loop) + return asyncio.Task(coroutine, loop=loop) + + if outcome == "driver_error": + monkeypatch.setattr(staging, "_read_on", driver_error) + else: + monkeypatch.setattr(staging, "_transaction", completed_cancel_race) + task = None + try: + loop.set_task_factory(make_task) + task = asyncio.create_task(call_public_boundary()) + _, pending = await asyncio.wait({task}, timeout=5) + assert not pending + task.result() + finally: + loop.set_task_factory(previous_factory) + if task is not None and not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + async with pool.connection() as connection: + count = ( + await ( + await connection.execute("SELECT count(*) FROM reporting_inline_objects") + ).fetchone() + )[0] + assert count == (1 if outcome == "completed_cancel_race" else 0) + + +def test_prepared_object_is_immutable_and_redacted() -> None: + payload = b"private-prepared-payload" + prepared = PreparedBackendObjectV1( + "private-account", "execution-1", 0, payload, hashlib.sha256(payload).hexdigest() + ) + assert prepared.payload is payload + assert repr(prepared) == "PreparedBackendObjectV1()" + with pytest.raises(FrozenInstanceError): + prepared.ordinal = 1 + + +@pytest.mark.parametrize( + "fault", ["account", "key", "bool", "negative", "ordinal_cap", "mutable", "digest", "cap"] +) +def test_prepared_object_rejects_invalid_inputs_without_raw_context(fault) -> None: + args = dict( + account_id="a", + source_execution_key="execution-1", + ordinal=0, + payload=b"bytes", + payload_sha256=hashlib.sha256(b"bytes").hexdigest(), + ) + if fault == "account": + args["account_id"] = "private-\ud800" + elif fault == "key": + args["source_execution_key"] = "private/key" + elif fault in {"bool", "negative", "ordinal_cap"}: + args["ordinal"] = {"bool": True, "negative": -1, "ordinal_cap": 100_000}[fault] + elif fault == "mutable": + args["payload"] = bytearray(b"private-bytes") + elif fault == "digest": + args["payload_sha256"] = "private-digest" + else: + args["payload"] = b"x" * (INLINE_STAGING_MAX_BYTES + 1) + with pytest.raises(InlineStorageError) as caught: + PreparedBackendObjectV1(**args) + assert caught.value.code == "INVALID_INPUT" + assert caught.value.__context__ is None and caught.value.__cause__ is None + assert "private" not in repr(caught.value) + + +async def test_prepared_seal_detaches_reference_and_preserves_neutral_winner(stores) -> None: + _, seals, _ = stores + identity = redacted_snapshot_request().identity + original = await sample_seal() + prepared = seals._prepare_seal( + account_id=identity.account_id, + source_execution_key=identity.source_execution_key, + sealed=original, + ) + assert type(prepared) is PreparedBackendSealV1 + assert repr(prepared) == "PreparedBackendSealV1()" + with pytest.raises(FrozenInstanceError): + prepared.byte_count = 1 + object.__setattr__(original.reference, "byte_count", 1) + assert prepared.byte_count == len(original.manifest_bytes) + async with seals._transaction() as connection: + winner = await seals._put_on(connection, prepared) + assert isinstance(winner, SealedSlice) + assert winner.reference.byte_count == prepared.byte_count + assert winner.manifest_bytes == prepared.manifest_bytes + + +@pytest.mark.parametrize("finish", ["commit", "rollback"]) +async def test_participants_share_caller_connection_task_and_transaction( + stores, monkeypatch, finish +) -> None: + staging, seals, pool = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + await staging.stage(**args, ordinal=0, payload=b"preexisting") + old_object = await staging._prepare_object(**args, ordinal=0, payload=b"preexisting") + new_object = await staging._prepare_object(**args, ordinal=1, payload=b"new-object") + first = seals._prepare_seal(**args, sealed=await sample_seal(10)) + second = seals._prepare_seal(**args, sealed=await sample_seal(11)) + assert first.manifest_sha256 != second.manifest_sha256 + owner = asyncio.current_task() + original_transaction = staging._transaction + original_read, original_get = staging._read_on, seals._get_on + seen = [] + + def forbidden_checkout(*args, **kwargs): + raise AssertionError("participants must use the supplied transaction") + + async def read_on(connection, *args): + assert connection is owned_connection and asyncio.current_task() is owner + seen.append("object") + return await original_read(connection, *args) + + async def get_on(connection, *args): + assert connection is owned_connection and asyncio.current_task() is owner + seen.append("seal") + return await original_get(connection, *args) + + class CallerRollbackError(Exception): + pass + + try: + # This fixture owns its account. A composed caller must additionally + # hold its account-order lock; these neutral participants do not do so. + async with original_transaction() as owned_connection: + with monkeypatch.context() as patch: + patch.setattr(pool, "connection", forbidden_checkout) + patch.setattr(staging, "_transaction", forbidden_checkout) + patch.setattr(seals, "_transaction", forbidden_checkout) + patch.setattr(staging, "_read_on", read_on) + patch.setattr(seals, "_get_on", get_on) + await staging._stage_on(owned_connection, old_object) + ref, generation = await staging._stage_on(owned_connection, new_object) + winner = await seals._put_on(owned_connection, first) + assert await seals._put_on(owned_connection, second) == winner + assert winner.manifest_bytes == first.manifest_bytes + assert winner.reference.manifest_sha256 == first.manifest_sha256 + isolation = await ( + await owned_connection.execute("SHOW transaction_isolation") + ).fetchone() + assert isolation[0] == "read committed" + if finish == "rollback": + raise CallerRollbackError + except CallerRollbackError: + assert finish == "rollback" + assert seen == ["object", "object", "seal", "seal"] + async with pool.connection() as connection: + count = ( + await ( + await connection.execute("SELECT count(*) FROM reporting_inline_objects") + ).fetchone() + )[0] + assert count == (2 if finish == "commit" else 1) + stored = await seals.get(**args) + if finish == "commit": + assert stored == winner + assert ( + await staging.read( + account_id=identity.account_id, + object_ref=ref, + object_generation=generation, + source_scope={}, + cancel=asyncio.Event(), + ) + == b"new-object" + ) + else: + assert stored is None + + +@pytest.mark.parametrize("kind", ["stage", "put"]) +async def test_public_writes_delegate_inside_existing_owned_transaction( + stores, monkeypatch, kind +) -> None: + staging, seals, _ = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + store = staging if kind == "stage" else seals + participant_name = "_stage_on" if kind == "stage" else "_put_on" + original_transaction = store._transaction + original_participant = getattr(store, participant_name) + caller = asyncio.current_task() + delegated = [] + settled = False + connection = owner = None + + @asynccontextmanager + async def transaction(): + nonlocal connection, owner, settled + owner = asyncio.current_task() + assert owner is not caller + async with original_transaction() as connection: + yield connection + settled = True + + async def participant(supplied_connection, prepared): + assert supplied_connection is connection and asyncio.current_task() is owner + assert not settled + delegated.append(prepared) + return await original_participant(supplied_connection, prepared) + + monkeypatch.setattr(store, "_transaction", transaction) + monkeypatch.setattr(store, participant_name, participant) + if kind == "stage": + result = await staging.stage(**args, ordinal=0, payload=b"delegated") + assert result[1] == hashlib.sha256(b"delegated").hexdigest() + assert type(delegated[0]) is PreparedBackendObjectV1 + else: + candidate = await sample_seal() + result = await seals.put(**args, sealed=candidate) + assert result == candidate + assert type(delegated[0]) is PreparedBackendSealV1 + assert settled and len(delegated) == 1 + + +@pytest.mark.parametrize("fault", ["digest", "canonical", "account", "cap"]) +async def test_seal_participant_revalidates_prepared_candidate_before_sql(stores, fault) -> None: + _, seals, _ = stores + identity = redacted_snapshot_request().identity + args = dict(account_id=identity.account_id, source_execution_key=identity.source_execution_key) + candidate = await sample_seal() + await seals.put(**args, sealed=candidate) + prepared = seals._prepare_seal(**args, sealed=candidate) + if fault == "digest": + prepared = replace(prepared, manifest_sha256="0" * 64) + elif fault == "canonical": + raw = json.dumps(json.loads(prepared.manifest_bytes), indent=2).encode() + prepared = replace( + prepared, + manifest_bytes=raw, + byte_count=len(raw), + manifest_sha256=hashlib.sha256(raw).hexdigest(), + ) + elif fault == "account": + prepared = replace(prepared, account_id="other-account") + else: + # Frozen records deter normal mutation; the participant still validates + # a deliberately forged instance rather than treating it as authority. + object.__setattr__(prepared, "byte_count", 1048577) + + class NoSql: + async def execute(self, *args, **kwargs): + raise AssertionError("invalid candidate must be refused before SQL") + + with pytest.raises(InlineStorageError) as caught: + await seals._put_on(NoSql(), prepared) + assert caught.value.code == "INVALID_INPUT" + assert caught.value.__context__ is None and caught.value.__cause__ is None + assert (await seals.get(**args)).manifest_bytes == candidate.manifest_bytes + + +async def test_stage_participant_enforces_local_cap_and_stored_byte_winner( + stores, monkeypatch +) -> None: + staging, _, pool = stores + prepared = await staging._prepare_object( + account_id="a", source_execution_key="execution-1", ordinal=0, payload=b"bytes" + ) + smaller = PgReportingStagingStore(pool=pool, max_payload_bytes=4) + + class NoSql: + async def execute(self, *args, **kwargs): + raise AssertionError("local payload cap must be checked before SQL") + + with pytest.raises(InlineStorageError) as caught: + await smaller._stage_on(NoSql(), prepared) + assert caught.value.code == "INVALID_INPUT" + original_read = staging._read_on + + async def wrong_winner(*args): + await original_read(*args) + return b"different-bytes" + + monkeypatch.setattr(staging, "_read_on", wrong_winner) + with pytest.raises(InlineStorageError) as caught: + async with staging._transaction() as connection: + await staging._stage_on(connection, prepared) + assert caught.value.code == "INTEGRITY_FAILED" + async with pool.connection() as connection: + assert ( + await ( + await connection.execute("SELECT count(*) FROM reporting_inline_objects") + ).fetchone() + )[0] == 0 diff --git a/tests/conformance/reporting/test_reporting_prepared_observation_adapter.py b/tests/conformance/reporting/test_reporting_prepared_observation_adapter.py index 25aa31299..52e6aa044 100644 --- a/tests/conformance/reporting/test_reporting_prepared_observation_adapter.py +++ b/tests/conformance/reporting/test_reporting_prepared_observation_adapter.py @@ -1,9 +1,13 @@ """A decorating publisher composes the entire atomic observation contract.""" import asyncio +from dataclasses import replace +from datetime import timedelta import pytest +from adcp.reporting.ledger import LedgerConflictError + from ._reliable_support import ScriptedSource, complete_fetch, configuration, reliable_factory from .test_reporting_evidence_currency_integration import frozen_slice @@ -73,3 +77,88 @@ async def test_prepared_ordinary_revision_keeps_the_existing_commit_seam(backend } assert await h.store.get_provisional_observation(**identity) is None assert await h.store.get_restatement_checkpoint(**identity) is None + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_prepared_observation_replay_validates_managed_content(backend): + async with reliable_factory(backend, notifications=True) as h: + await h.store.put_configuration(configuration("eur")) + script = ScriptedSource(eur=[complete_fetch]) + producer = h.producer(h.source(script.async_fetch)) + await producer.run_worker() + (request,) = script.requests + identity = { + "account_id": "eur", + "reporting_obligation_id": request.identity.reporting_obligation_id, + } + (revision,) = await h.store.list_revisions(**identity) + observation = await h.store.get_provisional_observation(**identity) + checkpoint = await h.store.get_restatement_checkpoint(**identity) + assert observation is not None and checkpoint is not None + assert revision.managed_control_totals is not None + assert revision.canonical_content_digest is not None + rows = ( + await h.store.read_revision_rows( + account_id="eur", reporting_revision_id=revision.reporting_revision_id + ) + ).rows + assert len(rows) == revision.row_count and rows + + async def effects(): + if h.blobs.pool is not None: + async with h.blobs.pool.connection() as connection: + return await ( + await connection.execute( + "SELECT (SELECT count(*) FROM reporting_ledger_changes)," + " (SELECT count(*) FROM reporting_notification_events)" + ) + ).fetchone() + return ( + len(h.store._changes), + len(h.store._notification_state.events), + len(h.store._notification_state.boundaries), + ) + + original_effects = await effects() + changed_rows = ( + dict(rows[0], impressions=rows[0]["impressions"] + 1), + *rows[1:], + ) + for submitted_revision, submitted_rows, code, publisher in ( + (revision, changed_rows, "REVISION_IMMUTABLE", producer._store), + (revision, (), "ROW_COUNT_MISMATCH", producer._store), + ( + replace(revision, revision_content_sha256="0" * 64), + rows, + "REVISION_CONTENT_MISMATCH", + h.store, + ), + ): + with pytest.raises(LedgerConflictError) as conflict: + await publisher.commit_provisional_observation( + observation, submitted_revision, submitted_rows + ) + assert conflict.value.code == code + assert await h.store.list_revisions(**identity) == (revision,) + assert ( + await h.store.read_revision_rows( + account_id="eur", reporting_revision_id=revision.reporting_revision_id + ) + ).rows == rows + assert await h.store.get_provisional_observation(**identity) == observation + assert await h.store.get_restatement_checkpoint(**identity) == checkpoint + assert await effects() == original_effects + + h.clock.advance(timedelta(hours=2)) + later = replace( + observation, + checked_at=h.clock(), + next_due_at=h.clock() + timedelta(hours=1), + ) + assert ( + await producer._store.commit_provisional_observation(later, revision, rows) == revision + ) + assert await h.store.list_revisions(**identity) == (revision,) + assert await h.store.get_provisional_observation(**identity) == observation + assert await h.store.get_restatement_checkpoint(**identity) == checkpoint + assert await effects() == original_effects diff --git a/tests/conformance/reporting/test_reporting_production_migration.py b/tests/conformance/reporting/test_reporting_production_migration.py index 6b08f019f..5c06050c8 100644 --- a/tests/conformance/reporting/test_reporting_production_migration.py +++ b/tests/conformance/reporting/test_reporting_production_migration.py @@ -25,12 +25,18 @@ def manifests(): - return { + required = { package: json.loads( files("adcp.reporting." + package).joinpath("required_schema.json").read_text() ) for package in ("materializer", "receipts", "feed", "projection", "production") } + required["production"].update( + json.loads( + files("adcp.reporting.production").joinpath("service_context_schema.json").read_text() + ) + ) + return required def original_rows(image, before): diff --git a/tests/conformance/reporting/test_reporting_provisional_catalog_scope.py b/tests/conformance/reporting/test_reporting_provisional_catalog_scope.py new file mode 100644 index 000000000..7c580c8bd --- /dev/null +++ b/tests/conformance/reporting/test_reporting_provisional_catalog_scope.py @@ -0,0 +1,127 @@ +"""Fresh provisional readiness examines only its exact catalog identities.""" + +import json +from importlib.resources import files + +import pytest + +from adcp.reporting.ledger import LedgerConflictError +from adcp.reporting.outbox._schema import ( + PROVISIONAL_REQUIRED_OBJECTS, + REQUIRED_OBJECTS, + _catalog_objects, + schema_objects, + validate_provisional_schema, +) +from adcp.reporting.outbox.status_schema import REQUIRED_STATUS_OBJECTS +from adcp.reporting.receipts import PgReportingReceiptStore + +from ._generation_support import isolated_reporting_pool + +TABLES = ("reporting_provisional_acquisitions", "reporting_provisional_observations") +FUNCTION = ("reporting_provisional_immutable", "") + + +class CatalogProbe: + def __init__(self, connection): + self.connection = connection + self.calls = [] + + async def execute(self, query, parameters=None): + rows = await (await self.connection.execute(query, parameters)).fetchall() + self.calls.append((query, parameters, rows)) + + class Cursor: + async def fetchall(self): + return rows + + return Cursor() + + def assert_narrow_scope(self): + assert len(self.calls) == 6 + assert self.calls[0][1] == (list(TABLES),) + assert {row[1] for row in self.calls[0][2]} == set(TABLES) + oids = [row[0] for row in self.calls[0][2]] + for _, parameters, _ in self.calls[1:5]: + assert parameters == (oids,) + assert self.calls[5][1] == FUNCTION + assert [(row[0], row[1]) for row in self.calls[5][2]] == [FUNCTION] + assert [len(rows) for _, _, rows in self.calls] == [2, 10, 8, 4, 2, 1] + + +@pytest.mark.parametrize("autocommit", [False, True]) +async def test_exact_catalog_scope_preserves_full_contract_and_excludes_overload(autocommit): + expected = { + **REQUIRED_OBJECTS, + **{ + key: value + for key, value in REQUIRED_STATUS_OBJECTS.items() + if "reporting_issue_waiver_bindings" in key + }, + **PROVISIONAL_REQUIRED_OBJECTS, + } + for package in ("materializer", "receipts"): + expected.update( + json.loads( + files("adcp.reporting." + package).joinpath("required_schema.json").read_text() + ) + ) + assert len(expected) == 790 and len(PROVISIONAL_REQUIRED_OBJECTS) == 27 + async with isolated_reporting_pool(autocommit=autocommit) as pool: + await PgReportingReceiptStore(pool=pool).create_schema() + async with pool.connection() as connection: + full = CatalogProbe(connection) + assert await schema_objects(full) == expected + assert len(full.calls) == 6 + assert sum(len(rows) for _, _, rows in full.calls) == 790 + await connection.execute( + "CREATE TABLE reporting_unrelated_catalog_scope (value integer);" + " CREATE FUNCTION reporting_provisional_immutable(integer) RETURNS integer" + " LANGUAGE sql IMMUTABLE AS 'SELECT $1'" + ) + default = CatalogProbe(connection) + objects = await schema_objects(default) + assert {key: objects[key] for key in expected} == expected + assert set(objects) - set(expected) == { + "table:reporting_unrelated_catalog_scope", + "column:reporting_unrelated_catalog_scope.value", + "function:reporting_provisional_immutable(integer)", + } + narrow = CatalogProbe(connection) + assert ( + await _catalog_objects(narrow, table_names=TABLES, function_identity=FUNCTION) + == PROVISIONAL_REQUIRED_OBJECTS + ) + narrow.assert_narrow_scope() + protected = CatalogProbe(connection) + await validate_provisional_schema(protected) + protected.assert_narrow_scope() + print( + json.dumps( + { + "catalog_scope": { + "autocommit": autocommit, + "full_queries": len(full.calls), + "full_rows": [len(rows) for _, _, rows in full.calls], + "narrow_queries": len(protected.calls), + "narrow_rows": [len(rows) for _, _, rows in protected.calls], + "default_extra_objects": len(objects) - len(expected), + } + } + ), + flush=True, + ) + # A prior successful validation cannot mask later DDL. The overload + # must neither satisfy nor perturb the exact zero-argument identity. + await connection.execute("DROP FUNCTION reporting_provisional_immutable() CASCADE") + with pytest.raises( + LedgerConflictError, + match=r"missing:function:reporting_provisional_immutable\(\)", + ) as failure: + await validate_provisional_schema(connection) + assert failure.value.code == "PROVISIONAL_SCHEMA_UNREADY" + assert ( + await ( + await connection.execute("SELECT reporting_provisional_immutable(7)") + ).fetchone() + ) == (7,) diff --git a/tests/conformance/reporting/test_reporting_publication_time.py b/tests/conformance/reporting/test_reporting_publication_time.py index a10347a37..0012030dd 100644 --- a/tests/conformance/reporting/test_reporting_publication_time.py +++ b/tests/conformance/reporting/test_reporting_publication_time.py @@ -120,6 +120,7 @@ async def setup( definition=key.definition, ) await store.put_configuration(config) + source.bind_generation(config, product_id=key.report_definition_id) def complete_read(): if not real: diff --git a/tests/conformance/reporting/test_reporting_rc6_rolling.py b/tests/conformance/reporting/test_reporting_rc6_rolling.py index 5995d566d..4417000ee 100644 --- a/tests/conformance/reporting/test_reporting_rc6_rolling.py +++ b/tests/conformance/reporting/test_reporting_rc6_rolling.py @@ -59,12 +59,22 @@ def accepted_b24(request, tmp_path_factory): cwd=root, timeout=180, ) + historical_paths = set( + subprocess.check_output( + ["git", "ls-tree", "-r", "--name-only", B24, "--", "src/adcp"], + cwd=ROOT, + text=True, + ).splitlines() + ) modules = {} with zipfile.ZipFile(wheel) as archive: for name in production_modules(): member = name.replace(".", "/") + ".py" - if member not in archive.namelist(): + if "src/" + member not in historical_paths: member = name.replace(".", "/") + "/__init__.py" + # New production modules are absent from the pinned B2.4 source and wheel. + if "src/" + member not in historical_paths: + continue raw = archive.read(member) assert raw == subprocess.check_output(["git", "show", f"{B24}:src/{member}"], cwd=ROOT) modules[name] = hashlib.sha256(raw).hexdigest() diff --git a/tests/conformance/reporting/test_reporting_source_authorization.py b/tests/conformance/reporting/test_reporting_source_authorization.py new file mode 100644 index 000000000..10a976735 --- /dev/null +++ b/tests/conformance/reporting/test_reporting_source_authorization.py @@ -0,0 +1,517 @@ +"""D5 live source authorization at dispatch, seal, publication and recovery.""" + +import asyncio +from contextlib import AsyncExitStack, asynccontextmanager +from dataclasses import replace +from functools import partial + +import pytest + +from adcp.reporting.ledger import InMemoryReportingLedgerStore +from adcp.reporting.materializer import ReportingWriterError, reference_verifier +from adcp.reporting.service import ReliableReportingService, ReportingAccountContext + +from ._generation_support import END, configuration, isolated_reporting_pool +from ._materializer_support import reference_rows +from ._production_support import production_harness +from .test_reporting_production_lock_order import source_turn +from .test_reporting_production_settling import SettlingSource, revisions + + +class RevocableSource(SettlingSource): + def __init__(self, *args, authorized=True, **kwargs): + super().__init__(*args, **kwargs) + self.authorized = authorized + self.dispatched = [] + self.started = asyncio.Event() + self.release = asyncio.Event() + self.release.set() + self.cancelled = False + + def configuration_binding(self, configuration): + return super().configuration_binding(configuration) if self.authorized else None + + async def execute(self, request, *, cancel, heartbeat=None): + self.dispatched.append(request) + try: + return await super().execute(request, cancel=cancel, heartbeat=heartbeat) + finally: + self.cancelled |= cancel.is_set() + + async def fetch(self, request): + self.started.set() + try: + await asyncio.wait_for(self.release.wait(), 10) + except asyncio.CancelledError: + self.cancelled = True + raise + return await super().fetch(request) + + +async def assert_unpublished(h): + assert not await revisions(h) + identity = { + "account_id": h.item.config.account_id, + "reporting_obligation_id": h.item.obligation.reporting_obligation_id, + } + assert await h.store.get_provisional_observation(**identity) is None + assert await h.store.get_restatement_checkpoint(**identity) is None + + +def harness(backend, path, *, source_factory=RevocableSource, **kwargs): + return production_harness( + backend, + path, + count=1, + source_publication=True, + periods=1, + source_factory=source_factory, + **kwargs, + ) + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_revoked_before_dispatch_resumes_on_next_authorized_turn(backend, tmp_path): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + source.authorized = False + revoked = await source_turn(h.production) + assert not revoked.revisions_committed + assert source.dispatched == [] + await assert_unpublished(h) + + source.authorized = True + restored = await source_turn(h.production) + assert len(restored.revisions_committed) == 1 + assert len(source.requests) == 1 + checkpoint = await h.store.get_restatement_checkpoint( + account_id=h.item.config.account_id, + reporting_obligation_id=h.item.obligation.reporting_obligation_id, + ) + assert checkpoint is not None and checkpoint.next_observation == 1 + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_dispatch_rechecks_after_acquisition_reservation(backend, tmp_path, monkeypatch): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + reserve = h.store.reserve_provisional_acquisition + + async def revoke_after_reservation(acquisition): + result = await reserve(acquisition) + source.authorized = False + return result + + monkeypatch.setattr(h.store, "reserve_provisional_acquisition", revoke_after_reservation) + revoked = await source_turn(h.production) + assert not revoked.revisions_committed + assert source.dispatched == [] + await assert_unpublished(h) + + monkeypatch.setattr(h.store, "reserve_provisional_acquisition", reserve) + source.authorized = True + restored = await source_turn(h.production) + assert len(restored.revisions_committed) == 1 + assert len(source.dispatched) == 1 + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_remapping_after_reservation_prevents_dispatch(backend, tmp_path, monkeypatch): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + reserve = h.store.reserve_provisional_acquisition + + async def remap_after_reservation(acquisition): + result = await reserve(acquisition) + source.bind_generation(h.item.config, product_id="catalog-5820") + return result + + monkeypatch.setattr(h.store, "reserve_provisional_acquisition", remap_after_reservation) + with pytest.raises(ReportingWriterError, match="BINDING_MISMATCH"): + await source_turn(h.production) + assert source.dispatched == [] + await assert_unpublished(h) + + monkeypatch.setattr(h.store, "reserve_provisional_acquisition", reserve) + source.bind_generation(h.item.config, product_id="catalog-7391") + restored = await source_turn(h.production) + assert len(restored.revisions_committed) == 1 + assert len(source.dispatched) == 1 + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +@pytest.mark.parametrize("boundary", ["fetch", "seal_lock", "ledger_lock"]) +async def test_remapping_during_publication_refuses_result_and_allows_retry( + backend, boundary, tmp_path, monkeypatch +): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + publication = h.store._source_publication + publications = 0 + + @asynccontextmanager + async def remap_under_lock(account_id, **kwargs): + nonlocal publications + async with publication(account_id, **kwargs) as publish_seal: + publications += 1 + if (boundary, publications) in (("seal_lock", 1), ("ledger_lock", 2)): + source.bind_generation(h.item.config, product_id="catalog-5820") + yield publish_seal + + monkeypatch.setattr(h.store, "_source_publication", remap_under_lock) + source.release.clear() + running = asyncio.create_task(source_turn(h.production)) + try: + await asyncio.wait_for(source.started.wait(), 10) + if boundary == "fetch": + source.bind_generation(h.item.config, product_id="catalog-5820") + assert not running.done() + finally: + source.release.set() + with pytest.raises(ReportingWriterError, match="BINDING_MISMATCH"): + await asyncio.wait_for(running, 10) + assert not source.cancelled + assert len(source.requests) == 1 + await assert_unpublished(h) + request = source.dispatched[0] + seal = await source.inline._seals.get( + account_id=request.identity.account_id, + source_execution_key=request.identity.source_execution_key, + ) + assert (seal is not None) == (boundary == "ledger_lock") + + monkeypatch.setattr(h.store, "_source_publication", publication) + source.bind_generation(h.item.config, product_id="catalog-7391") + source.rows = reference_rows(2) + restored = await source_turn(h.production) + assert len(restored.revisions_committed) == 1 + assert source.dispatched[1].identity == request.identity + # A seal made before remapping can replay only after fresh authorization. + # Refusal at the seal boundary must acquire the changed rows on retry. + expected_rows = 1 if boundary == "ledger_lock" else 2 + assert (await revisions(h))[0].row_count == expected_rows + assert len(source.requests) == expected_rows + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_inflight_revocation_discards_result_without_cancelling_fetch(backend, tmp_path): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + source.release.clear() + running = asyncio.create_task(source_turn(h.production)) + try: + await asyncio.wait_for(source.started.wait(), 10) + source.authorized = False + assert not running.done() + finally: + source.release.set() + revoked = await asyncio.wait_for(running, 10) + assert not revoked.revisions_committed and not revoked.slices_failed + assert len(source.requests) == 1 + assert not source.cancelled + await assert_unpublished(h) + request = source.dispatched[0] + assert ( + await source.inline._seals.get( + account_id=request.identity.account_id, + source_execution_key=request.identity.source_execution_key, + ) + is None + ) + + # Changed rows prove the rejected result did not become a replay seal. + source.rows = reference_rows(2) + source.authorized = True + restored = await source_turn(h.production) + assert len(restored.revisions_committed) == 1 + assert len(source.requests) == 2 + assert (await revisions(h))[0].row_count == 2 + assert source.dispatched[1].identity == request.identity + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_seal_authorization_checked_after_obtaining_account_lock( + backend, tmp_path, monkeypatch +): + async with harness(backend, tmp_path / "destination.sqlite") as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + source.release.clear() + running = asyncio.create_task(source_turn(h.production)) + waiting = asyncio.Event() + publication = h.store._source_publication + + @asynccontextmanager + async def observe_publication(account_id, **kwargs): + waiting.set() + async with publication(account_id, **kwargs) as publish_seal: + yield publish_seal + + try: + await asyncio.wait_for(source.started.wait(), 10) + monkeypatch.setattr(h.store, "_source_publication", observe_publication) + async with publication(h.item.config.account_id): + source.release.set() + await asyncio.wait_for(waiting.wait(), 10) + assert not running.done() + source.authorized = False + finally: + source.release.set() + revoked = await asyncio.wait_for(running, 10) + assert not revoked.revisions_committed + await assert_unpublished(h) + request = source.dispatched[0] + assert ( + await source.inline._seals.get( + account_id=request.identity.account_id, + source_execution_key=request.identity.source_execution_key, + ) + is None + ) + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_revoked_before_ledger_commit_replay_rechecks_after_restart( + backend, tmp_path, monkeypatch +): + path = tmp_path / "destination.sqlite" + async with harness(backend, path) as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + read = source.reader.read + + async def revoke_after_read(**kwargs): + result = await read(**kwargs) + source.authorized = False + return result + + monkeypatch.setattr(source.reader, "read", revoke_after_read) + revoked = await source_turn(h.production) + assert not revoked.revisions_committed + await assert_unpublished(h) + request = source.dispatched[0] + sealed = await source.inline._seals.get( + account_id=request.identity.account_id, + source_execution_key=request.identity.source_execution_key, + ) + assert sealed is not None, "seal preceded revocation, ledger publication did not" + await h.production.aclose() + + existing = {"existing_store": h.store} if h.pool is None else {"existing_pool": h.pool} + async with harness( + backend, path, source_factory=partial(RevocableSource, authorized=False), **existing + ) as fresh: + resumed = fresh.production.offerings[0].producer._source + await source_turn(fresh.production) + assert resumed.dispatched == [] + await assert_unpublished(fresh) + + resumed.authorized = True + restored = await source_turn(fresh.production) + assert len(restored.revisions_committed) == 1 + assert len(resumed.dispatched) == 1 + assert resumed.dispatched[0].identity == request.identity + assert resumed.requests == [], "the existing seal was replayed with fresh authority" + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_revocation_stops_account_for_turn_and_other_accounts_continue(backend, tmp_path): + class WithdrawOnceSource(RevocableSource): + withdrawn = False + refuse_next_check = False + + async def fetch(self, request): + result = await super().fetch(request) + if request.identity.account_id == "acct_a" and not self.withdrawn: + self.withdrawn = self.refuse_next_check = True + return result + + def configuration_binding(self, configuration): + if self.refuse_next_check: + self.refuse_next_check = False + return None + return super().configuration_binding(configuration) + + key = reference_verifier().key + source = WithdrawOnceSource( + key, + tmp_path / "source", + reference_rows(1), + clock=lambda: END, + close_officially=False, + product_ids=(key.report_definition_id,), + ) + + def account_context(config): + return ReportingAccountContext( + config.account_id, + "source", + "USD", + source.capabilities.source_scope, + snapshot_offering_id=source.source_id, + capability_offering={ + "offering_id": source.source_id, + "feed_purpose": config.feed_purpose, + "report_definition_id": config.report_definition_id, + "supported_finality": ["snapshot"], + }, + publication_namespace=source.capabilities.offerings[0].publication_namespace, + ) + + async with AsyncExitStack() as stack: + if backend == "memory": + store = InMemoryReportingLedgerStore(clock=lambda: END) + else: + from adcp.reporting.ledger.pg import PgReportingLedgerStore + + pool = await stack.enter_async_context(isolated_reporting_pool(autocommit=True)) + store = PgReportingLedgerStore(pool=pool, clock=lambda: END) + service = ReliableReportingService( + store=store, account_context=account_context, clock=lambda: END + ) + stack.push_async_callback(service.close) + service.sources.register_executor("source", source, object_reader=source.reader) + configs = [] + for account_id, config_id in (("acct_a", "first"), ("acct_a", "later"), ("acct_b", "only")): + config = replace( + configuration(account_id), + delivery_config_id=config_id, + definition=key.definition, + report_definition_id=key.report_definition_id, + reporting_profile=key.reporting_profile, + ) + source.bind_generation(config, product_id=key.report_definition_id) + await service.configure(config) + configs.append(config) + + turn = await service.run_worker() + assert not turn.configuration_errors + assert [request.identity.account_id for request in source.dispatched] == [ + "acct_a", + "acct_b", + ] + assert not turn.configurations[configs[0].generation_key].revisions_committed + assert not turn.configurations[configs[1].generation_key].revisions_committed + assert len(turn.configurations[configs[2].generation_key].revisions_committed) == 1 + + restored = await service.run_worker() + assert not restored.configuration_errors + assert len(restored.configurations[configs[0].generation_key].revisions_committed) == 1 + assert len(restored.configurations[configs[1].generation_key].revisions_committed) == 1 + + +@pytest.mark.parametrize("interruption", ["revoked", "storage_error", "cancelled"]) +@pytest.mark.parametrize("autocommit", [False, True]) +async def test_postgres_inline_publication_shares_size_one_pool( + tmp_path, monkeypatch, interruption, autocommit +): + pools = pytest.importorskip("psycopg_pool") + psycopg = pytest.importorskip("psycopg") + + from adcp.reporting.inline_source import InlineReportingSource + from adcp.reporting.inline_storage import ( + InlineStorageError, + PgReportingSealStore, + PgReportingStagingStore, + ) + + async with isolated_reporting_pool(autocommit=autocommit) as isolated: + async with pools.AsyncConnectionPool( + isolated.conninfo, + kwargs=isolated.kwargs, + min_size=1, + max_size=1, + timeout=2, + open=False, + ) as pool: + staging = PgReportingStagingStore(pool=pool) + seals = PgReportingSealStore(pool=pool) + await staging.create_schema() + + def factory(*args, **kwargs): + source = RevocableSource(*args, **kwargs) + source.inline = InlineReportingSource( + capabilities=source.capabilities, + fetch=source.fetch, + staging=staging, + seals=seals, + constituent_of=lambda row, request: ( + request.coverage.constituents[0].constituent_id + ), + clock=kwargs["clock"], + ) + source.reader = staging + return source + + async with harness( + "postgres", + tmp_path / "destination.sqlite", + existing_pool=pool, + source_factory=factory, + ) as h: + await h.production.activate(account_id=h.item.config.account_id) + source = h.production.offerings[0].producer._source + source.authorized = False + denied = await source_turn(h.production) + assert not denied.revisions_committed and not source.dispatched + await assert_unpublished(h) + + source.authorized = True + put_on = seals._put_on + inserted = asyncio.Event() + + async def interrupted_put(connection, prepared): + result = await put_on(connection, prepared) + if interruption == "storage_error": + raise psycopg.OperationalError("private-storage-detail") + if interruption == "cancelled": + inserted.set() + await asyncio.Event().wait() + return result + + monkeypatch.setattr(seals, "_put_on", interrupted_put) + source.release.clear() + running = asyncio.create_task(source_turn(h.production)) + try: + await asyncio.wait_for(source.started.wait(), 10) + if interruption == "revoked": + source.authorized = False + finally: + source.release.set() + if interruption == "cancelled": + await asyncio.wait_for(inserted.wait(), 10) + running.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(running, 10) + elif interruption == "storage_error": + with pytest.raises(InlineStorageError) as caught: + await asyncio.wait_for(running, 10) + assert caught.value.code == "RESOURCE_UNAVAILABLE" + assert caught.value.__context__ is None and caught.value.__cause__ is None + assert "private-storage-detail" not in str(caught.value) + else: + revoked = await asyncio.wait_for(running, 10) + assert not revoked.revisions_committed + await assert_unpublished(h) + request = source.dispatched[0] + assert ( + await seals.get( + account_id=request.identity.account_id, + source_execution_key=request.identity.source_execution_key, + ) + is None + ) + + monkeypatch.setattr(seals, "_put_on", put_on) + source.authorized = True + source.rows = reference_rows(2) + restored = await asyncio.wait_for(source_turn(h.production), 10) + assert len(restored.revisions_committed) == 1 + assert (await revisions(h))[0].row_count == 2 + assert source.cancelled == (interruption == "cancelled") + assert len(source.requests) == 2 diff --git a/tests/test_reporting_production_source_registry.py b/tests/test_reporting_production_source_registry.py new file mode 100644 index 000000000..40cfeefd4 --- /dev/null +++ b/tests/test_reporting_production_source_registry.py @@ -0,0 +1,76 @@ +from dataclasses import replace +from datetime import timedelta +from types import SimpleNamespace + +import pytest + +from adcp.reporting.ledger.producer import ProducerOfferings +from adcp.reporting.production.source_registry import ReportingProductionSourceRegistry +from adcp.reporting.service import ReportingAccountContext + + +def test_registry_freeze_rejects_producers_without_an_offering(): + producer = SimpleNamespace(_offerings=ProducerOfferings()) + unused = SimpleNamespace(_offerings=ProducerOfferings()) + assert unused == producer and unused is not producer + registry = ReportingProductionSourceRegistry(account_context=lambda _: None) + registry.register("used", producer) + registry.register("unused", unused) + offering = SimpleNamespace(producer=producer) + + with pytest.raises(ValueError, match="offering"): + registry.freeze((offering,)) + assert not registry._frozen + + registry.freeze((offering, SimpleNamespace(producer=unused))) + with pytest.raises(ValueError, match="frozen"): + registry.register("later", SimpleNamespace(_offerings=ProducerOfferings())) + + +@pytest.mark.asyncio +async def test_registry_freezes_profile_and_refuses_replaced_currency(): + producer = SimpleNamespace(_offerings=ProducerOfferings(official_offering_id="official")) + context = ReportingAccountContext( + "account", "adapter", "USD", {}, official_offering_id="official" + ) + calls = [] + + def resolve(configuration): + calls.append(configuration) + return context + + registry = ReportingProductionSourceRegistry(account_context=resolve) + registry.register("adapter", producer) + configuration = SimpleNamespace(account_id="account", account_timezone="UTC") + offering = SimpleNamespace( + producer=producer, offering_id="public", source_offering_id="official" + ) + capabilities = SimpleNamespace(capabilities_sha256="a" * 64) + frozen = await registry.resolve(configuration, offering, capabilities) + assert frozen.document()["currency"] == "USD" + assert registry.recover(configuration, offering, capabilities, frozen.document()) == frozen + assert len(calls) == 1 + producer._offerings = replace(producer._offerings, currency="EUR") + with pytest.raises(ValueError, match="profile"): + registry.recover(configuration, offering, capabilities, frozen.document()) + + +@pytest.mark.asyncio +async def test_registry_missing_context_is_explicit_and_budget_is_frozen(): + producer = SimpleNamespace(_offerings=ProducerOfferings(official_offering_id="official")) + context = ReportingAccountContext( + "account", "adapter", "USD", {}, official_offering_id="official" + ) + registry = ReportingProductionSourceRegistry(account_context=lambda _: context) + registry.register("adapter", producer) + configuration = SimpleNamespace(account_id="account", account_timezone="UTC") + offering = SimpleNamespace( + producer=producer, offering_id="public", source_offering_id="official" + ) + capabilities = SimpleNamespace(capabilities_sha256="a" * 64) + with pytest.raises(ValueError, match="migration"): + registry.recover(configuration, offering, capabilities, None) + frozen = await registry.resolve(configuration, offering, capabilities) + producer._offerings = replace(producer._offerings, slice_timeout=timedelta(seconds=1)) + with pytest.raises(ValueError, match="profile"): + registry.recover(configuration, offering, capabilities, frozen.document()) diff --git a/tests/test_reporting_provisional_observations.py b/tests/test_reporting_provisional_observations.py index eede7ae67..5d8906f6b 100644 --- a/tests/test_reporting_provisional_observations.py +++ b/tests/test_reporting_provisional_observations.py @@ -13,8 +13,10 @@ InMemoryReportingLedgerStore, PgReportingLedgerStore, ReportingProducer, + revision_content_sha256, ) from adcp.reporting.ledger.store import LedgerConflictError +from adcp.reporting.source import SourceBatchManifestV1 from tests.conformance.reporting._generation_support import isolated_reporting_pool from tests.test_reporting_settling import ( ACCOUNT, @@ -61,6 +63,30 @@ async def latest(store): return observation +async def _publication_effect_counts(store): + if isinstance(store, PgReportingLedgerStore): + async with store._pool.connection() as connection: + return await ( + await connection.execute( + "SELECT (SELECT count(*) FROM reporting_revisions)," + " (SELECT count(*) FROM reporting_revision_rows)," + " (SELECT count(*) FROM reporting_provisional_observations)," + " (SELECT count(*) FROM reporting_restatement_checkpoints)," + " (SELECT count(*) FROM reporting_ledger_changes)," + " (SELECT count(*) FROM reporting_notification_events)" + ) + ).fetchone() + return ( + len(store._revisions), + sum(len(rows) for rows in store._rows.values()), + len(store._provisional_observations), + len(store._restatement_checkpoints), + len(store._changes), + len(store._notification_state.events), + len(store._notification_state.boundaries), + ) + + async def test_adapter_without_window_rereads_with_sdk_default(make_harness): producer, store, fetch, clock = await make_harness(_capabilities(restatement_window=None)) first = await producer.run_worker() @@ -443,6 +469,133 @@ async def test_generation_and_replay_identity_cannot_be_rebound(make_harness): assert await latest(store) == observation +async def test_recorded_observation_replay_validates_content_and_keeps_original_state(make_harness): + producer, store, fetch, clock = await make_harness( + _capabilities(restatement_window="P3D", restatement_cadence="PT1H") + ) + await producer.run_worker() + parent = (await _revisions(store))[0] + clock[0] += timedelta(hours=1) + fetch.impressions = 27 + await producer.run_worker() + revision = (await _revisions(store))[1] + observation = await latest(store) + assert revision.supersedes_reporting_revision_id == parent.reporting_revision_id + rows = ( + await store.read_revision_rows( + account_id=ACCOUNT, reporting_revision_id=revision.reporting_revision_id + ) + ).rows + assert len(rows) == revision.row_count == 1 + checkpoint = await store.get_restatement_checkpoint( + account_id=ACCOUNT, reporting_obligation_id=observation.acquisition.obligation_id + ) + effects = await _publication_effect_counts(store) + changed_rows = (dict(rows[0], impressions=rows[0]["impressions"] + 1),) + changed_revision = replace( + revision, + revision_content_sha256=revision_content_sha256( + reporting_revision_id=revision.reporting_revision_id, + row_count=revision.row_count, + control_totals=revision.control_totals, + reporting_rows=changed_rows, + ), + ) + assert changed_revision.revision_content_sha256 != revision.revision_content_sha256 + for submitted_revision, submitted_rows, code in ( + (changed_revision, changed_rows, "REVISION_IMMUTABLE"), + (revision, (), "ROW_COUNT_MISMATCH"), + ): + with pytest.raises(LedgerConflictError) as conflict: + await store.commit_provisional_observation( + observation, submitted_revision, submitted_rows + ) + assert conflict.value.code == code + assert (await _revisions(store)) == (parent, revision) + assert ( + await store.read_revision_rows( + account_id=ACCOUNT, reporting_revision_id=revision.reporting_revision_id + ) + ).rows == rows + assert await latest(store) == observation + assert ( + await store.get_restatement_checkpoint( + account_id=ACCOUNT, reporting_obligation_id=observation.acquisition.obligation_id + ) + == checkpoint + ) + assert await _publication_effect_counts(store) == effects + + # This is the producer's acquisition-supplied branch, which reconstructs + # a new digest for the same publication ID before calling the observation store. + manifest = SourceBatchManifestV1.model_validate_json(observation.manifest_json) + obligation = await _only_obligation(store) + with pytest.raises(LedgerConflictError) as conflict: + await producer.commit_revision_from_manifest( + obligation, + manifest, + rows=changed_rows, + finality="snapshot", + now=clock[0], + acquisition=observation.acquisition, + ) + assert conflict.value.code == "REVISION_IMMUTABLE" + assert await latest(store) == observation + assert ( + await store.get_restatement_checkpoint( + account_id=ACCOUNT, reporting_obligation_id=observation.acquisition.obligation_id + ) + == checkpoint + ) + assert await _publication_effect_counts(store) == effects + + clock[0] += timedelta(hours=2) + replay_store = ( + PgReportingLedgerStore(pool=store._pool, clock=lambda: clock[0], notifications=True) + if isinstance(store, PgReportingLedgerStore) + else store + ) + later_observation = replace( + observation, + checked_at=clock[0], + next_due_at=clock[0] + timedelta(hours=1), + ) + assert ( + await replay_store.commit_provisional_observation(later_observation, revision, rows) + == revision + ) + assert ( + await producer.commit_revision_from_manifest( + obligation, + manifest, + rows=rows, + finality="snapshot", + now=clock[0], + acquisition=observation.acquisition, + ) + == revision + ) + assert revision.created_at < clock[0] + assert revision.supersedes_reporting_revision_id == parent.reporting_revision_id + assert observation.acquisition.ordinal == 1 + assert observation.checked_at < clock[0] + assert observation.next_due_at is not None + assert await latest(replay_store) == observation + assert ( + await replay_store.get_restatement_checkpoint( + account_id=ACCOUNT, reporting_obligation_id=observation.acquisition.obligation_id + ) + == checkpoint + ) + assert (await _revisions(replay_store)) == (parent, revision) + assert ( + await replay_store.read_revision_rows( + account_id=ACCOUNT, reporting_revision_id=revision.reporting_revision_id + ) + ).rows == rows + assert await _publication_effect_counts(replay_store) == effects + + async def test_notification_failure_rolls_back_revision_checkpoint_and_history( make_harness, monkeypatch ): diff --git a/tests/type_checks/reliable_reporting_production_context.py b/tests/type_checks/reliable_reporting_production_context.py new file mode 100644 index 000000000..e2ed6ab92 --- /dev/null +++ b/tests/type_checks/reliable_reporting_production_context.py @@ -0,0 +1,24 @@ +"""Installed public fixed-profile production bridge typing.""" + +from typing import Any + +from adcp.reporting.ledger import ReportingProducer +from adcp.reporting.production import ReportingProductionSourceRegistry, ReportingProductionSupport +from adcp.reporting.service import ReliableReportingService, ReportingContextResolver +from adcp.server import ADCPHandler + + +def register( + context: ReportingContextResolver, producer: ReportingProducer +) -> ReportingProductionSourceRegistry: + registry = ReportingProductionSourceRegistry(account_context=context) + registry.register("fixed-provider-profile", producer) + return registry + + +def compose( + production: ReportingProductionSupport, application: ADCPHandler[Any] +) -> ReliableReportingService: + service = ReliableReportingService.from_production(production) + service.install(application) + return service diff --git a/tests/type_checks/reporting_inline_storage.py b/tests/type_checks/reporting_inline_storage.py new file mode 100644 index 000000000..04e0a4856 --- /dev/null +++ b/tests/type_checks/reporting_inline_storage.py @@ -0,0 +1,28 @@ +"""Borrowed pool composition with the existing public source protocols.""" + +from psycopg_pool import AsyncConnectionPool + +from adcp.reporting.inline_source import ( + InlineFetch, + InlineReportingSource, + ReportingSealStore, + ReportingStagingStore, +) +from adcp.reporting.inline_storage import PgReportingSealStore, PgReportingStagingStore +from adcp.reporting.source import ReportingSourceCapabilitiesV1, ReportingSourceExecutor + + +async def configure( + pool: AsyncConnectionPool, + capabilities: ReportingSourceCapabilitiesV1, + fetch: InlineFetch, +) -> ReportingSourceExecutor: + staging = PgReportingStagingStore(pool=pool, max_payload_bytes=1_048_576) + seals = PgReportingSealStore(pool=pool) + await staging.create_schema() + await seals.check_ready() + objects: ReportingStagingStore = staging + replay: ReportingSealStore = seals + return InlineReportingSource( + capabilities=capabilities, fetch=fetch, staging=objects, seals=replay + )