diff --git a/src/adcp/__init__.py b/src/adcp/__init__.py index b2f338229..c6fcf8557 100644 --- a/src/adcp/__init__.py +++ b/src/adcp/__init__.py @@ -203,6 +203,24 @@ def _resolve_version() -> str: "discover_oauth_metadata", "start_oauth_authorization", ), + "adcp.principal": ( + "DestinationSetupSnapshot", + "PrincipalClient", + "PrincipalConfigurationError", + "PrincipalManager", + "PrincipalSyncOutcome", + "sync_principal_configuration", + ), + "adcp.reporting_inspection": ( + "ControlTotalCalculator", + "HttpsReportingResourceReader", + "ManifestReportingInspector", + "ReportingCredentialProvider", + "ReportingInspectionCode", + "ReportingInspectionError", + "ReportingResourceReader", + "ReportingTrustedOriginPolicy", + ), "adcp.exceptions": ( "AdagentsAccessBlockedError", "AdagentsNotFoundError", @@ -319,6 +337,20 @@ def _resolve_version() -> str: "SyncAgentNotificationConfigsResponse", "GetPrincipalRequest", "GetPrincipalResponse", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", "GetReportingStatusRequest", "GetReportingStatusResponse", "SyncPrincipalRequest", @@ -326,6 +358,11 @@ def _resolve_version() -> str: "SyncReportingReceiptsRequest", "SyncReportingReceiptsResponse", "AccountAuthorization", + "AgentDeclarations", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "AccountReference", "AccountScope", "AccountWithAuthorization", @@ -1117,6 +1154,11 @@ def get_adcp_version() -> str: "McpWebhookPayload", # Account operations "AccountAuthorization", + "AgentDeclarations", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "AccountReference", "AccountWithAuthorization", "GetAccountFinancialsRequest", @@ -1229,6 +1271,34 @@ def get_adcp_version() -> str: "SyncAgentNotificationConfigsResponse", "GetPrincipalRequest", "GetPrincipalResponse", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", + "DestinationSetupSnapshot", + "PrincipalClient", + "PrincipalConfigurationError", + "PrincipalManager", + "PrincipalSyncOutcome", + "sync_principal_configuration", + "ControlTotalCalculator", + "HttpsReportingResourceReader", + "ManifestReportingInspector", + "ReportingCredentialProvider", + "ReportingInspectionCode", + "ReportingInspectionError", + "ReportingResourceReader", + "ReportingTrustedOriginPolicy", "GetReportingStatusRequest", "GetReportingStatusResponse", "SyncPrincipalRequest", @@ -1798,6 +1868,14 @@ def get_adcp_version() -> str: start_oauth_authorization, ) from adcp.observability import get_tracer, inject_trace_headers, is_tracing_available + from adcp.principal import ( + DestinationSetupSnapshot, + PrincipalClient, + PrincipalConfigurationError, + PrincipalManager, + PrincipalSyncOutcome, + sync_principal_configuration, + ) from adcp.property_registry import PropertyRegistry from adcp.registry import RegistryClient from adcp.registry_sync import ( @@ -1806,6 +1884,16 @@ def get_adcp_version() -> str: FileCursorStore, RegistrySync, ) + from adcp.reporting_inspection import ( + ControlTotalCalculator, + HttpsReportingResourceReader, + ManifestReportingInspector, + ReportingCredentialProvider, + ReportingInspectionCode, + ReportingInspectionError, + ReportingResourceReader, + ReportingTrustedOriginPolicy, + ) from adcp.substitution import ( MacroMapping, MacroMappingEntry, @@ -1874,6 +1962,11 @@ def get_adcp_version() -> str: ActivateSignalRequest, ActivateSignalResponse, AdvertiserIndustry, + AgentDeclarations, + AgentNotificationConfig, + AgentNotificationConfigState, + AgentReportingDestination, + AgentReportingDestinationState, # Creative types ArtifactWebhookPayload, AssetContentType, @@ -2096,6 +2189,20 @@ def get_adcp_version() -> str: PriceGuidance, PricingCurrency, PricingModel, + PrincipalAppliedResult, + PrincipalChangedWebhook, + PrincipalConfiguration, + PrincipalCurrentResult, + PrincipalDeclarationsState, + PrincipalKind, + PrincipalReadFailedResult, + PrincipalReadResult, + PrincipalRecognizedResult, + PrincipalState, + PrincipalSyncFailedResult, + PrincipalSyncResult, + PrincipalUnconfiguredResult, + PrincipalValidatedResult, Product, ProductAllowedAction, ProductFilters, diff --git a/src/adcp/principal.py b/src/adcp/principal.py new file mode 100644 index 000000000..273ff0b43 --- /dev/null +++ b/src/adcp/principal.py @@ -0,0 +1,327 @@ +"""Buyer-side orchestration for the AdCP 3.2 principal layer. + +The generated request and response models describe one wire exchange. This +module implements the stateful client workflow adopters otherwise have to +repeat: read the current version, submit a guarded section replacement, and +poll asynchronous destination proof until it reaches an actionable state. +""" + +from __future__ import annotations + +import asyncio +import time +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any, Protocol +from uuid import uuid4 + +from adcp.types import ( + AgentReportingDestinationState, + GetPrincipalRequest, + GetPrincipalResponse, + PrincipalConfiguration, + PrincipalDeclarationsState, + PrincipalState, + SyncPrincipalRequest, + SyncPrincipalResponse, +) +from adcp.types.core import TaskResult + + +class PrincipalClient(Protocol): + """Minimal client surface required by :class:`PrincipalManager`.""" + + async def get_principal(self, request: GetPrincipalRequest) -> TaskResult[GetPrincipalResponse]: + raise NotImplementedError + + async def sync_principal( + self, request: SyncPrincipalRequest + ) -> TaskResult[SyncPrincipalResponse]: + raise NotImplementedError + + +class PrincipalConfigurationError(RuntimeError): + """A principal read, mutation, or destination setup operation failed.""" + + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + + +def _enum(value: object) -> str: + return str(getattr(value, "value", value)) + + +def _failure_message(result: object, fallback: str) -> str: + errors = getattr(result, "errors", None) + if errors: + messages = [str(getattr(error, "message", "")) for error in errors] + detail = "; ".join(message for message in messages if message) + if detail: + return detail + return fallback + + +@dataclass(frozen=True) +class DestinationSetupSnapshot: + """Authoritative setup states from one principal readback.""" + + states: dict[str, AgentReportingDestinationState] + + @property + def validating(self) -> tuple[str, ...]: + return self._with_state("validating") + + @property + def ready(self) -> tuple[str, ...]: + return self._with_state("ready") + + @property + def action_required(self) -> tuple[str, ...]: + return self._with_state("action_required") + + @property + def inactive(self) -> tuple[str, ...]: + return self._with_state("inactive") + + @property + def rejected(self) -> tuple[str, ...]: + return self._with_state("rejected") + + @property + def settled(self) -> bool: + """Whether no destination remains in provider-side validation.""" + return not self.validating + + def _with_state(self, expected: str) -> tuple[str, ...]: + return tuple( + destination_id + for destination_id, state in self.states.items() + if _enum(state.state) == expected + ) + + +@dataclass(frozen=True) +class PrincipalSyncOutcome: + """Applied configuration plus negotiated and setup-state projections.""" + + response: SyncPrincipalResponse + principal_id: str | None + configuration_version: str | None + configuration: PrincipalState | None + destinations: DestinationSetupSnapshot + declarations: PrincipalDeclarationsState | None + + @property + def selected_async_adcp_version(self) -> str | None: + if self.declarations is None: + return None + return self.declarations.selected_async_adcp_version + + +def _destination_snapshot( + configuration: PrincipalState | None, + destination_ids: set[str] | None = None, +) -> DestinationSetupSnapshot: + destinations = configuration.reporting_destinations if configuration else None + states = { + item.destination_id: item + for item in destinations or [] + if destination_ids is None or item.destination_id in destination_ids + } + return DestinationSetupSnapshot(states) + + +class PrincipalManager: + """Coordinate version-fenced principal configuration for one seller.""" + + def __init__(self, client: PrincipalClient) -> None: + self._client = client + + async def read(self) -> GetPrincipalResponse: + """Read principal state and normalize transport/payload failures.""" + task = await self._client.get_principal(GetPrincipalRequest()) + if not task.success or task.data is None: + raise PrincipalConfigurationError( + "PRINCIPAL_READ_FAILED", task.error or "get_principal failed" + ) + response = task.data + if response.result.kind == "failed": + raise PrincipalConfigurationError( + "PRINCIPAL_READ_FAILED", + _failure_message(response.result, "get_principal returned a failed result"), + ) + return response + + async def sync( + self, + configuration: PrincipalConfiguration | Mapping[str, Any], + *, + idempotency_key: str | None = None, + use_version_fence: bool = True, + expected_configuration_version: str | None = None, + expected_principal_kind: object | None = None, + dry_run: bool = False, + wait_for_setup: bool = False, + setup_timeout: float = 300.0, + poll_interval: float = 1.0, + ) -> PrincipalSyncOutcome: + """Replace selected sections, optionally waiting for destination proof. + + With ``use_version_fence`` (the default), this first calls + ``get_principal``. A current configuration version is copied to + ``expected_configuration_version`` and the seller-resolved principal + kind is asserted. Callers targeting a seller that explicitly declares + ``optimistic_concurrency: false`` should disable the fence. + """ + desired = PrincipalConfiguration.model_validate(configuration) + current: GetPrincipalResponse | None = None + if use_version_fence: + current = await self.read() + current_result = current.result + if expected_configuration_version is None and current_result.kind == "current": + expected_configuration_version = current_result.configuration_version + if expected_principal_kind is None: + if current_result.kind == "current": + expected_principal_kind = current_result.principal_kind + elif current_result.kind == "recognized": + expected_principal_kind = current_result.principal_kind + + payload: dict[str, Any] = { + "idempotency_key": idempotency_key or str(uuid4()), + "configuration": desired, + "dry_run": dry_run, + } + if expected_configuration_version is not None: + payload["expected_configuration_version"] = expected_configuration_version + if expected_principal_kind is not None: + payload["expected_principal_kind"] = expected_principal_kind + + task = await self._client.sync_principal(SyncPrincipalRequest.model_validate(payload)) + if not task.success or task.data is None: + raise PrincipalConfigurationError( + "PRINCIPAL_SYNC_FAILED", task.error or "sync_principal failed" + ) + response = task.data + result = response.result + if result.kind == "failed": + raise PrincipalConfigurationError( + "PRINCIPAL_SYNC_FAILED", + _failure_message(result, "sync_principal returned a failed result"), + ) + if result.kind == "validated": + return PrincipalSyncOutcome( + response, None, None, None, DestinationSetupSnapshot({}), None + ) + + destination_ids = { + destination.destination_id for destination in desired.reporting_destinations or [] + } + configuration_state = result.configuration + destination_states = _destination_snapshot(configuration_state, destination_ids) + if wait_for_setup and destination_states.validating: + expected_destination_refs = { + destination_id: state.destination_ref + for destination_id, state in destination_states.states.items() + } + readback, destination_states = await self.wait_for_destinations( + destination_ids, + timeout=setup_timeout, + poll_interval=poll_interval, + expected_destination_refs=expected_destination_refs, + ) + read_result = readback.result + if read_result.kind != "current": + raise PrincipalConfigurationError( + "PRINCIPAL_STATE_LOST", + "get_principal stopped returning current state during setup polling", + ) + configuration_state = read_result.configuration + principal_id = read_result.principal_id + configuration_version = read_result.configuration_version + else: + principal_id = result.principal_id + configuration_version = result.configuration_version + + return PrincipalSyncOutcome( + response, + principal_id, + configuration_version, + configuration_state, + destination_states, + configuration_state.declarations, + ) + + async def wait_for_destinations( + self, + destination_ids: set[str], + *, + timeout: float = 300.0, + poll_interval: float = 1.0, + expected_destination_refs: Mapping[str, str] | None = None, + ) -> tuple[GetPrincipalResponse, DestinationSetupSnapshot]: + """Poll while requested destinations are ``validating``. + + ``ready``, ``action_required``, ``inactive``, and ``rejected`` are + returned to the caller as actionable settled states. The helper never + follows or executes an untrusted setup URL. + """ + if timeout <= 0 or poll_interval <= 0: + raise ValueError("timeout and poll_interval must be positive") + if not destination_ids: + raise ValueError("destination_ids must not be empty") + + deadline = time.monotonic() + timeout + while True: + response = await self.read() + result = response.result + if result.kind != "current": + raise PrincipalConfigurationError( + "PRINCIPAL_NOT_CONFIGURED", + "get_principal did not return current state during setup polling", + ) + snapshot = _destination_snapshot(result.configuration, destination_ids) + missing = destination_ids.difference(snapshot.states) + if missing: + raise PrincipalConfigurationError( + "DESTINATION_STATE_MISSING", + "get_principal omitted requested destination(s): " + ", ".join(sorted(missing)), + ) + if expected_destination_refs: + replaced = [ + destination_id + for destination_id, destination_ref in expected_destination_refs.items() + if snapshot.states[destination_id].destination_ref != destination_ref + ] + if replaced: + raise PrincipalConfigurationError( + "DESTINATION_GENERATION_CHANGED", + "get_principal replaced requested destination generation(s): " + + ", ".join(sorted(replaced)), + ) + if snapshot.settled: + return response, snapshot + remaining = deadline - time.monotonic() + if remaining <= 0: + pending = ", ".join(snapshot.validating) + raise TimeoutError(f"reporting destination setup did not settle in time: {pending}") + await asyncio.sleep(min(poll_interval, remaining)) + + +async def sync_principal_configuration( + client: PrincipalClient, + configuration: PrincipalConfiguration | Mapping[str, Any], + **kwargs: Any, +) -> PrincipalSyncOutcome: + """Functional wrapper for :meth:`PrincipalManager.sync`.""" + return await PrincipalManager(client).sync(configuration, **kwargs) + + +__all__ = [ + "DestinationSetupSnapshot", + "PrincipalClient", + "PrincipalConfigurationError", + "PrincipalManager", + "PrincipalSyncOutcome", + "sync_principal_configuration", +] diff --git a/src/adcp/reporting.py b/src/adcp/reporting.py index 8b881987b..688022e51 100644 --- a/src/adcp/reporting.py +++ b/src/adcp/reporting.py @@ -13,8 +13,9 @@ from collections.abc import Awaitable, Callable, Iterable from dataclasses import dataclass, field from datetime import datetime, timezone +from enum import Enum from math import isfinite -from typing import Protocol, TypeVar +from typing import TYPE_CHECKING, Protocol, TypeVar from uuid import uuid4 from pydantic import BaseModel @@ -24,6 +25,7 @@ GetReportingStatusResponse, ReportingCanonicalContentDigest, ReportingControlTotal, + ReportingDeliveryCapabilities, ReportingMaterialization, ReportingObligation, ReportingReceipt, @@ -33,13 +35,19 @@ ) from adcp.types.core import TaskResult +if TYPE_CHECKING: + from adcp.reporting_inspection import ReportingResourceReader -class ReportingReconciliationClient(Protocol): + +class ReportingStatusClient(Protocol): async def get_reporting_status( self, request: GetReportingStatusRequest ) -> TaskResult[GetReportingStatusResponse]: raise NotImplementedError + +class ReportingReconciliationClient(ReportingStatusClient, Protocol): + async def sync_reporting_receipts( self, request: SyncReportingReceiptsRequest ) -> TaskResult[SyncReportingReceiptsResponse]: @@ -60,6 +68,31 @@ def __init__(self, code: str, message: str) -> None: self.code = code +class ReportingTier(str, Enum): + """Feature tiers advertised by ``media_buy.reporting_delivery``.""" + + CORE = "core" + MANAGED_DELIVERY = "managed_delivery" + RECONCILED_BILLING = "reconciled_billing" + + +def reporting_tiers( + capabilities: ReportingDeliveryCapabilities, +) -> frozenset[ReportingTier]: + """Project reporting capability flags into their cumulative SDK tiers.""" + tiers = {ReportingTier.CORE} + if capabilities.managed_delivery: + tiers.add(ReportingTier.MANAGED_DELIVERY) + if capabilities.reconciled_billing: + if not capabilities.managed_delivery: + raise ReportingReconciliationError( + "INVALID_REPORTING_CAPABILITIES", + "reconciled_billing requires managed_delivery", + ) + tiers.add(ReportingTier.RECONCILED_BILLING) + return frozenset(tiers) + + @dataclass(frozen=True) class ExpectedReportingPeriod: delivery_config_id: str @@ -191,7 +224,7 @@ def _add_immutable( async def load_reporting_ledger( - client: ReportingReconciliationClient, + client: ReportingStatusClient, request: GetReportingStatusRequest, *, max_snapshot_restarts: int = 2, @@ -686,26 +719,75 @@ def evaluate_reporting_ledger( ) +async def reconcile_reporting_core( + client: ReportingStatusClient, + request: GetReportingStatusRequest, + *, + expected_periods: list[ExpectedReportingPeriod], + max_snapshot_restarts: int = 2, + now: datetime | None = None, +) -> ReportingReconciliationResult: + """Reconcile the Core API-delivered tier without destination handling. + + A Core caller only compares obligations, revisions, coverage, finality, and + the reporting clock. Destination materializations, manifests, digests, + and consumer receipts are deliberately rejected rather than accidentally + activating a higher tier. + """ + ledger = await load_reporting_ledger( + client, request, max_snapshot_restarts=max_snapshot_restarts + ) + if any(obligation.destination_ref is not None for obligation in ledger.obligations): + raise ReportingReconciliationError( + "MANAGED_DELIVERY_NOT_ENABLED", + "Core reconciliation received a managed-delivery obligation", + ) + if ledger.materializations or ledger.receipts: + raise ReportingReconciliationError( + "MANAGED_DELIVERY_NOT_ENABLED", + "Core reconciliation received destination or receipt records", + ) + if any( + _enum(obligation.reconciliation_mode) == "consumer_receipt" + for obligation in ledger.obligations + ): + raise ReportingReconciliationError( + "RECONCILED_BILLING_NOT_ENABLED", + "Core reconciliation received a consumer-receipt obligation", + ) + return evaluate_reporting_ledger(ledger, expected_periods=expected_periods, now=now) + + async def reconcile_reporting( client: ReportingReconciliationClient, request: GetReportingStatusRequest, - inspect: Callable[[ReportingInspectionContext], Awaitable[ReportingObservation]], + inspect: Callable[[ReportingInspectionContext], Awaitable[ReportingObservation]] | None = None, *, expected_periods: list[ExpectedReportingPeriod], + resource_reader: ReportingResourceReader | None = None, checkpoint_store: ReportingCheckpointStore | None = None, max_snapshot_restarts: int = 2, max_inspection_attempts: int = 3, inspection_timeout_seconds: float = 30.0, inspection_retry_backoff_seconds: float = 1.0, now: datetime | None = None, + reporting_capabilities: ReportingDeliveryCapabilities | None = None, ) -> ReportingReconciliationResult: """Reconcile a closed ledger, persist observations, and submit receipts. - Each destination inspection is time-bounded. Transient failures retry with - exponential backoff so a hung or overloaded destination cannot block the - reconciliation loop indefinitely. + Pass ``resource_reader`` for the built-in manifest/file inspector, or + ``inspect`` as an advanced adapter for warehouses and native shares. Each + inspection is time-bounded. Typed transient failures retry with exponential + backoff; permanent integrity failures stop immediately. """ + if inspect is not None and resource_reader is not None: + raise ValueError("pass inspect or resource_reader, not both") + if resource_reader is not None: + from adcp.reporting_inspection import ManifestReportingInspector + + inspect = ManifestReportingInspector(resource_reader) + if ( not isinstance(max_inspection_attempts, int) or isinstance(max_inspection_attempts, bool) @@ -720,6 +802,40 @@ async def reconcile_reporting( ledger = await load_reporting_ledger( client, request, max_snapshot_restarts=max_snapshot_restarts ) + if reporting_capabilities is not None: + tiers = reporting_tiers(reporting_capabilities) + if ReportingTier.MANAGED_DELIVERY not in tiers and any( + item.destination_ref is not None for item in ledger.obligations + ): + raise ReportingReconciliationError( + "MANAGED_DELIVERY_NOT_ENABLED", + "reporting obligations require the managed_delivery tier", + ) + if ReportingTier.RECONCILED_BILLING not in tiers and any( + _enum(item.reconciliation_mode) == "consumer_receipt" + or _enum(item.feed_purpose) == "billing" + for item in ledger.obligations + ): + raise ReportingReconciliationError( + "RECONCILED_BILLING_NOT_ENABLED", + "consumer receipts and billing reporting require the reconciled_billing tier", + ) + if ReportingTier.RECONCILED_BILLING not in tiers and any( + item.canonical_content_digest is not None for item in ledger.revisions + ): + raise ReportingReconciliationError( + "RECONCILED_BILLING_NOT_ENABLED", + "canonical-digest reporting requires the reconciled_billing tier", + ) + if ReportingTier.RECONCILED_BILLING not in tiers and any( + item.verification + and _enum(item.verification.verification_profile) == "canonical_digest" + for item in ledger.materializations + ): + raise ReportingReconciliationError( + "RECONCILED_BILLING_NOT_ENABLED", + "canonical-digest verification requires the reconciled_billing tier", + ) submitted: list[ReportingReceipt] = [] for obligation in ledger.obligations: if _enum(obligation.reconciliation_mode) != "consumer_receipt": @@ -746,6 +862,11 @@ async def reconcile_reporting( or checkpoint_is_recorded or not _receipt_targets(receipt, obligation, revision, materialization) ): + if inspect is None: + raise ReportingReconciliationError( + "INSPECTOR_REQUIRED", + "consumer-receipt reconciliation requires inspect or resource_reader", + ) last_error: Exception | None = None observation = None inspection_context = ReportingInspectionContext(obligation, revision, materialization) @@ -756,6 +877,9 @@ async def reconcile_reporting( ) break except Exception as error: # destination SDKs define their own transient errors + if getattr(error, "retryable", None) is False: + code = getattr(error, "code", "INSPECTION_FAILED") + raise ReportingReconciliationError(_enum(code), str(error)) from error last_error = ( TimeoutError( "materialization inspection timed out after " @@ -815,8 +939,12 @@ async def reconcile_reporting( "ReportingReconciliationClient", "ReportingReconciliationError", "ReportingReconciliationResult", + "ReportingStatusClient", + "ReportingTier", "build_reporting_receipt", "evaluate_reporting_ledger", "load_reporting_ledger", "reconcile_reporting", + "reconcile_reporting_core", + "reporting_tiers", ] diff --git a/src/adcp/reporting_inspection.py b/src/adcp/reporting_inspection.py new file mode 100644 index 000000000..cab56c444 --- /dev/null +++ b/src/adcp/reporting_inspection.py @@ -0,0 +1,988 @@ +"""Built-in inspection for immutable reporting file manifests.""" + +from __future__ import annotations + +import asyncio +import base64 +import binascii +import csv +import hashlib +import io +import ipaddress +import json +import zlib +from collections.abc import Awaitable, Callable, Iterable, Mapping +from decimal import Decimal, InvalidOperation +from enum import Enum +from typing import Any, Protocol, cast +from urllib.parse import urljoin, urlsplit + +import httpx +import idna +import jsonschema +import rfc8785 +from pydantic import BaseModel, ValidationError + +from adcp.reporting import ReportingInspectionContext, ReportingObservation +from adcp.signing._bounded_http import ResponseTooLargeError, async_read_limited_bytes +from adcp.signing._idna_canonicalize import canonicalize_host +from adcp.signing.ip_pinned_transport import build_async_ip_pinned_transport +from adcp.signing.jwks import SSRFValidationError +from adcp.types import ( + ReportingCanonicalContentDigest, + ReportingCanonicalizationContract, + ReportingControlTotal, + ReportingFileManifest, + ReportingReportDefinition, +) + + +class ReportingInspectionCode(str, Enum): + RESOURCE_UNAVAILABLE = "RESOURCE_UNAVAILABLE" + UNSAFE_RESOURCE = "UNSAFE_RESOURCE" + RESOURCE_TOO_LARGE = "RESOURCE_TOO_LARGE" + UNEXPECTED_CONTENT_TYPE = "UNEXPECTED_CONTENT_TYPE" + MANIFEST_DIGEST_MISMATCH = "MANIFEST_DIGEST_MISMATCH" + INVALID_MANIFEST = "INVALID_MANIFEST" + MANIFEST_IDENTITY_MISMATCH = "MANIFEST_IDENTITY_MISMATCH" + DUPLICATE_OBJECT = "DUPLICATE_OBJECT" + OBJECT_SIZE_MISMATCH = "OBJECT_SIZE_MISMATCH" + OBJECT_DIGEST_MISMATCH = "OBJECT_DIGEST_MISMATCH" + UNSUPPORTED_FORMAT = "UNSUPPORTED_FORMAT" + UNSUPPORTED_COMPRESSION = "UNSUPPORTED_COMPRESSION" + INVALID_ROWS = "INVALID_ROWS" + ROW_SCHEMA_INVALID = "ROW_SCHEMA_INVALID" + ROW_COUNT_MISMATCH = "ROW_COUNT_MISMATCH" + CONTROL_TOTAL_MISMATCH = "CONTROL_TOTAL_MISMATCH" + UNSUPPORTED_CONTROL_TOTAL = "UNSUPPORTED_CONTROL_TOTAL" + CANONICAL_DIGEST_MISMATCH = "CANONICAL_DIGEST_MISMATCH" + INVALID_CONTRACT = "INVALID_CONTRACT" + + +class ReportingInspectionError(RuntimeError): + def __init__( + self, code: ReportingInspectionCode, message: str, *, retryable: bool = False + ) -> None: + super().__init__(message) + self.code = code + self.retryable = retryable + + +class ReportingResourceReader(Protocol): + """Read a bounded locator; adapters own credentials and provider routing.""" + + async def read(self, locator: str, *, base: str | None = None, max_bytes: int) -> bytes: + raise NotImplementedError + + +ReportingCredentialProvider = Callable[[str], Awaitable[Mapping[str, str]]] +ReportingTrustedOriginPolicy = Callable[[str], bool] +ControlTotalCalculator = Callable[ + [list[dict[str, Any]], ReportingReportDefinition], list[ReportingControlTotal] +] + + +class HttpsReportingResourceReader: + """Redirect-free, DNS-pinned reader for explicitly trusted HTTPS origins. + + ``trusted_origins`` must contain the seller, provider, and/or registry + origins authorized by the selected reporting offering. This deliberately + has no permissive default: ledger-controlled contract URLs are not an + authorization decision. A policy is useful when the trusted origin set is + resolved from an authenticated principal at inspection time. + """ + + def __init__( + self, + credential_provider: ReportingCredentialProvider | None = None, + *, + timeout_seconds: float = 10.0, + trusted_origins: Iterable[str] | None = None, + trusted_origin_policy: ReportingTrustedOriginPolicy | None = None, + ) -> None: + if timeout_seconds <= 0: + raise ValueError("timeout_seconds must be positive") + if trusted_origins is not None and trusted_origin_policy is not None: + raise ValueError("pass trusted_origins or trusted_origin_policy, not both") + self._credential_provider = credential_provider + self._timeout = timeout_seconds + self._trusted_origins = ( + frozenset(_origin(value) for value in trusted_origins) + if trusted_origins is not None + else None + ) + self._trusted_origin_policy = trusted_origin_policy + + async def read( + self, + locator: str, + *, + base: str | None = None, + max_bytes: int, + expected_content_types: frozenset[str] | None = None, + ) -> bytes: + url = urljoin(base, locator) if base else locator + origin = _origin(url) + if base: + if origin != _origin(base): + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting object reference crossed the manifest origin", + ) + origin_text = _origin_text(origin) + if self._trusted_origins is not None: + trusted = origin in self._trusted_origins + elif self._trusted_origin_policy is not None: + trusted = self._trusted_origin_policy(origin_text) + else: + trusted = False + if not trusted: + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + f"reporting resource origin {origin_text!r} is not trusted", + ) + try: + # Host resolution in the pin factory uses socket.getaddrinfo. Keep + # it off the event loop so the reconciler's inspection deadline can + # still cancel the awaiting task during a slow DNS lookup. + transport = await asyncio.to_thread( + build_async_ip_pinned_transport, url, allowed_ports=frozenset({443}) + ) + headers = ( + dict(await self._credential_provider(url)) if self._credential_provider else {} + ) + async with httpx.AsyncClient( + transport=transport, + timeout=self._timeout, + follow_redirects=False, + trust_env=False, + ) as client: + async with client.stream("GET", url, headers=headers) as response: + if 300 <= response.status_code < 400: + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "redirects are not allowed for reporting resources", + ) + if response.status_code != 200: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_UNAVAILABLE, + f"reporting resource returned HTTP {response.status_code}", + retryable=response.status_code >= 500 + or response.status_code in {408, 429}, + ) + if expected_content_types is not None: + content_type = response.headers.get("content-type", "").split(";", 1)[0] + if content_type.strip().lower() not in expected_content_types: + raise ReportingInspectionError( + ReportingInspectionCode.UNEXPECTED_CONTENT_TYPE, + "reporting resource returned an unexpected content type", + ) + return await async_read_limited_bytes(response, limit=max_bytes) + except ReportingInspectionError: + raise + except ResponseTooLargeError as error: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, str(error) + ) from error + except SSRFValidationError as error: + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting resource failed public-network validation", + ) from error + except (httpx.HTTPError, OSError, ValueError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_UNAVAILABLE, + "reporting resource fetch failed", + retryable=True, + ) from error + + +def _origin(url: str) -> tuple[str, str, int]: + """Return a normalized HTTPS origin while rejecting ambiguous locators.""" + parsed = urlsplit(url) + if ( + parsed.scheme != "https" + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + ): + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting resources must be credential-free absolute HTTPS URLs", + ) + try: + # IP literals are never valid reporting resource origins, even when + # publicly routable: safe-fetch policy authorizes DNS names only. + ipaddress.ip_address(parsed.hostname) + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting resource URLs must not use IP-literal hosts", + ) + except ValueError: + pass + try: + hostname = canonicalize_host(parsed.hostname) + port = parsed.port or 443 + except (ValueError, UnicodeError, idna.IDNAError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting resource URL has an invalid host or port", + ) from error + if port != 443: + raise ReportingInspectionError( + ReportingInspectionCode.UNSAFE_RESOURCE, + "reporting resources must use port 443", + ) + return ("https", hostname, port) + + +def _origin_text(origin: tuple[str, str, int]) -> str: + scheme, host, port = origin + return f"{scheme}://{host}" if port == 443 else f"{scheme}://{host}:{port}" + + +def _digest(body: bytes) -> str: + return hashlib.sha256(body).hexdigest() + + +def _dump(value: object) -> object: + if isinstance(value, BaseModel): + return value.model_dump(mode="json", exclude_none=True) + return value + + +def _same(left: object, right: object) -> bool: + return _dump(left) == _dump(right) + + +def _totals_same(left: list[ReportingControlTotal], right: list[ReportingControlTotal]) -> bool: + def normalized(values: list[ReportingControlTotal]) -> list[object]: + return sorted( + (_dump(item) for item in values), + key=lambda item: json.dumps(item, sort_keys=True, separators=(",", ":")), + ) + + return normalized(left) == normalized(right) + + +_MAX_SCHEMA_DEPTH = 64 +_MAX_SCHEMA_NODES = 10_000 + + +def _validate_row_schema(schema: dict[str, Any], expected_dialect: str) -> None: + """Apply the reporting profile's non-executable schema safety contract.""" + if schema.get("$schema") != expected_dialect: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema does not declare the revision schema dialect", + ) + stack: list[tuple[object, int]] = [(schema, 0)] + nodes = 0 + while stack: + value, depth = stack.pop() + nodes += 1 + if nodes > _MAX_SCHEMA_NODES or depth > _MAX_SCHEMA_DEPTH: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema exceeds static safety limits", + ) + if isinstance(value, dict): + if "$dynamicRef" in value or "$recursiveRef" in value or "$vocabulary" in value: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema uses an unsupported reference or vocabulary", + ) + ref = value.get("$ref") + if ref is not None and (not isinstance(ref, str) or not ref.startswith("#")): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema contains a non-local reference", + ) + for key, child in value.items(): + # stdlib/jsonschema evaluates Python backtracking regexes with + # no per-pattern deadline. Until the SDK uses a linear-time + # engine, reject regex-bearing contracts before compilation; + # even a short nested-quantifier pattern can be catastrophic. + if key == "pattern" and isinstance(child, str): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema contains an unsupported regular expression", + ) + if key == "patternProperties" and isinstance(child, dict): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema contains unsupported pattern properties", + ) + stack.append((child, depth + 1)) + elif isinstance(value, list): + stack.extend((child, depth + 1) for child in value) + _reject_local_ref_cycles(schema) + + +def _local_refs(value: object) -> set[str]: + if isinstance(value, dict): + result = {value["$ref"]} if isinstance(value.get("$ref"), str) else set() + for child in value.values(): + result.update(_local_refs(child)) + return result + if isinstance(value, list): + return set().union(*(_local_refs(child) for child in value)) if value else set() + return set() + + +def _resolve_json_pointer(document: dict[str, Any], ref: str) -> object: + value: object = document + if ref == "#": + return value + for token in ref.removeprefix("#/").split("/"): + token = token.replace("~1", "/").replace("~0", "~") + if isinstance(value, dict) and token in value: + value = value[token] + elif isinstance(value, list) and token.isdigit() and int(token) < len(value): + value = value[int(token)] + else: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema contains an unresolved local reference", + ) + return value + + +def _reject_local_ref_cycles(schema: dict[str, Any]) -> None: + """Reject local reference cycles before jsonschema can recursively follow them.""" + targets: dict[str, set[str]] = {} + + def edges(ref: str) -> set[str]: + if ref not in targets: + targets[ref] = _local_refs(_resolve_json_pointer(schema, ref)) + return targets[ref] + + visiting: set[str] = set() + visited: set[str] = set() + + def visit(ref: str) -> None: + if ref in visiting: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "reporting row schema contains a local reference cycle", + ) + if ref in visited: + return + visiting.add(ref) + for child in edges(ref): + visit(child) + visiting.remove(ref) + visited.add(ref) + + for ref in _local_refs(schema): + visit(ref) + + +class _InvalidJsonError(ValueError): + pass + + +def _json_loads_no_duplicates(body: bytes) -> object: + def object_pairs(pairs: list[tuple[str, object]]) -> dict[str, object]: + result: dict[str, object] = {} + for key, value in pairs: + if key in result: + raise _InvalidJsonError(f"duplicate JSON object key {key!r}") + result[key] = value + return result + + def reject_constant(value: str) -> object: + raise _InvalidJsonError(f"non-finite JSON value {value!r}") + + try: + return json.loads(body, object_pairs_hook=object_pairs, parse_constant=reject_constant) + except (json.JSONDecodeError, RecursionError, UnicodeDecodeError) as error: + raise _InvalidJsonError("invalid JSON") from error + + +def _decode_rows( + body: bytes, + format_name: str, + compression: str, + *, + max_decoded_bytes: int, + max_rows: int, +) -> tuple[list[dict[str, Any]], int]: + if compression == "gzip": + try: + decoder = zlib.decompressobj(16 + zlib.MAX_WBITS) + decoded = decoder.decompress(body, max_decoded_bytes + 1) + if ( + len(decoded) > max_decoded_bytes + or decoder.unconsumed_tail + or not decoder.eof + or decoder.unused_data + ): + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "decoded reporting object exceeds the configured byte limit", + ) + body = decoded + except ReportingInspectionError: + raise + except zlib.error as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, "invalid gzip reporting object" + ) from error + elif compression != "none": + raise ReportingInspectionError( + ReportingInspectionCode.UNSUPPORTED_COMPRESSION, + f"built-in reader does not support {compression} compression", + ) + decoded_size = len(body) + try: + text = body.decode("utf-8") + if format_name == "jsonl": + values = [] + for line in text.splitlines(): + if line.strip(): + values.append(_json_loads_no_duplicates(line.encode("utf-8"))) + if len(values) > max_rows: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "reporting object exceeds the configured row limit", + ) + elif format_name == "csv": + parser = csv.DictReader(io.StringIO(text, newline="")) + if parser.fieldnames and len(parser.fieldnames) != len(set(parser.fieldnames)): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, + "CSV reporting object has duplicate column names", + ) + values = [] + for value in parser: + values.append(value) + if len(values) > max_rows: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "reporting object exceeds the configured row limit", + ) + else: + raise ReportingInspectionError( + ReportingInspectionCode.UNSUPPORTED_FORMAT, + f"built-in reader does not support {format_name}; install a provider adapter", + ) + except ReportingInspectionError: + raise + except (_InvalidJsonError, UnicodeDecodeError, json.JSONDecodeError, csv.Error) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, "reporting object is not valid row data" + ) from error + if any(not isinstance(value, dict) for value in values): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, "every reporting row must be an object" + ) + return cast(list[dict[str, Any]], values), decoded_size + + +def _decimal_string(value: Decimal, value_type: str) -> str: + if value_type == "integer": + if value != value.to_integral_value(): + raise ValueError("non-integral value") + return str(int(value)) + rendered = format(value, "f") + if "." in rendered: + rendered = rendered.rstrip("0").rstrip(".") + return rendered or "0" + + +def _default_control_totals( + rows: list[dict[str, Any]], + definition: ReportingReportDefinition, + expected: list[ReportingControlTotal], +) -> list[ReportingControlTotal]: + metrics = {metric.name: metric for metric in definition.metrics} + totals: list[ReportingControlTotal] = [] + for target in expected: + metric = metrics.get(target.name) + if metric is None: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + f"control total {target.name!r} is absent from the report definition", + ) + aggregation = str(metric.aggregation) + expression = metric.source_expression + if aggregation == "count" and expression == "*": + value = Decimal(len(rows)) + else: + if not expression.isidentifier(): + raise ReportingInspectionError( + ReportingInspectionCode.UNSUPPORTED_CONTROL_TOTAL, + f"control total {target.name!r} requires a provider expression adapter", + ) + try: + values = [ + Decimal(str(row[expression])) for row in rows if row.get(expression) is not None + ] + except (InvalidOperation, KeyError, ValueError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, + f"control total source {expression!r} is not numeric", + ) from error + if aggregation == "count": + value = Decimal(len(values)) + elif aggregation == "sum": + value = sum(values, Decimal(0)) + elif aggregation == "min" and values: + value = min(values) + elif aggregation == "max" and values: + value = max(values) + elif aggregation == "average" and values: + value = sum(values, Decimal(0)) / Decimal(len(values)) + else: + raise ReportingInspectionError( + ReportingInspectionCode.UNSUPPORTED_CONTROL_TOTAL, + f"control total {target.name!r} requires a custom calculator", + ) + try: + payload = target.model_dump(mode="json", exclude_none=True) + payload["value"] = _decimal_string(value, str(target.value_type)) + totals.append(ReportingControlTotal.model_validate(payload)) + except (ValueError, ValidationError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.CONTROL_TOTAL_MISMATCH, + f"control total {target.name!r} cannot be represented canonically", + ) from error + return totals + + +class ManifestReportingInspector: + """Retrieve and verify a manifest materialization into an observation.""" + + def __init__( + self, + reader: ReportingResourceReader, + *, + max_manifest_bytes: int = 1024 * 1024, + max_object_bytes: int = 128 * 1024 * 1024, + max_contract_bytes: int = 4 * 1024 * 1024, + max_total_object_bytes: int = 256 * 1024 * 1024, + max_total_decoded_bytes: int = 256 * 1024 * 1024, + max_total_rows: int = 500_000, + max_files: int = 4_096, + control_total_calculator: ControlTotalCalculator | None = None, + ) -> None: + if ( + min( + max_manifest_bytes, + max_object_bytes, + max_contract_bytes, + max_total_object_bytes, + max_total_decoded_bytes, + max_total_rows, + max_files, + ) + <= 0 + ): + raise ValueError("inspection byte limits must be positive") + self._reader = reader + self._max_manifest_bytes = max_manifest_bytes + self._max_object_bytes = max_object_bytes + self._max_contract_bytes = max_contract_bytes + self._max_total_object_bytes = max_total_object_bytes + self._max_total_decoded_bytes = max_total_decoded_bytes + self._max_total_rows = max_total_rows + self._max_files = max_files + self._control_total_calculator = control_total_calculator + + async def __call__(self, context: ReportingInspectionContext) -> ReportingObservation: + materialization = context.materialization + resource = materialization.resource + revision = context.revision + obligation = context.obligation + if resource is None or str(resource.kind) != "manifest" or not resource.manifest_sha256: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_MANIFEST, + "materialization does not expose a digest-pinned manifest", + ) + + manifest_bytes = await self._read( + resource.location, + max_bytes=self._max_manifest_bytes, + expected_content_types=frozenset( + {"application/json", "application/vnd.adcp.reporting-file-manifest+json"} + ), + ) + manifest_digest = _digest(manifest_bytes) + if manifest_digest.lower() != resource.manifest_sha256.lower(): + raise ReportingInspectionError( + ReportingInspectionCode.MANIFEST_DIGEST_MISMATCH, + "manifest digest does not match the resource descriptor", + ) + try: + manifest_value = _json_loads_no_duplicates(manifest_bytes) + manifest = ReportingFileManifest.model_validate(manifest_value) + except (ValidationError, _InvalidJsonError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_MANIFEST, + "manifest is not valid reporting-file-manifest 1.0", + ) from error + if not ( + manifest.reporting_revision_id == revision.reporting_revision_id + and manifest.reporting_obligation_id == obligation.reporting_obligation_id + and manifest.reporting_materialization_id + == materialization.reporting_materialization_id + and _same(manifest.period, revision.period) + and manifest.row_count == revision.row_count + and _totals_same(manifest.control_totals, revision.control_totals) + ): + raise ReportingInspectionError( + ReportingInspectionCode.MANIFEST_IDENTITY_MISMATCH, + "manifest does not match the reporting ledger records", + ) + refs = [entry.object_ref for entry in manifest.files] + if len(refs) != len(set(refs)): + raise ReportingInspectionError( + ReportingInspectionCode.DUPLICATE_OBJECT, + "manifest contains duplicate object references", + ) + if len(manifest.files) > self._max_files: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "manifest exceeds the configured file limit", + ) + verification = materialization.verification + physical = { + item.object_ref: item.value.lower() + for item in (verification.physical_checksums if verification else None) or [] + if str(item.algorithm) == "sha256" + } + if any(physical.get(entry.object_ref) != entry.sha256.lower() for entry in manifest.files): + raise ReportingInspectionError( + ReportingInspectionCode.MANIFEST_IDENTITY_MISMATCH, + "manifest checksums do not match producer verification evidence", + ) + if manifest.total_size_bytes != sum(entry.size_bytes for entry in manifest.files): + raise ReportingInspectionError( + ReportingInspectionCode.OBJECT_SIZE_MISMATCH, + "manifest total_size_bytes does not equal its file entries", + ) + if manifest.row_count != sum(entry.row_count for entry in manifest.files): + raise ReportingInspectionError( + ReportingInspectionCode.ROW_COUNT_MISMATCH, + "manifest row_count does not equal its file entries", + ) + + rows: list[dict[str, Any]] = [] + total_object_bytes = 0 + total_decoded_bytes = 0 + total_rows = 0 + for entry in manifest.files: + if entry.size_bytes > self._max_object_bytes: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "manifest object exceeds the configured byte limit", + ) + total_object_bytes += entry.size_bytes + total_rows += entry.row_count + if ( + total_object_bytes > self._max_total_object_bytes + or total_rows > self._max_total_rows + ): + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "manifest exceeds the configured aggregate inspection limits", + ) + body = await self._read( + entry.object_ref, + base=resource.location, + max_bytes=min(self._max_object_bytes, entry.size_bytes + 1), + expected_content_types=_object_content_types(str(manifest.format)), + ) + if len(body) != entry.size_bytes: + raise ReportingInspectionError( + ReportingInspectionCode.OBJECT_SIZE_MISMATCH, + f"reporting object {entry.object_ref!r} has the wrong size", + ) + if _digest(body).lower() != entry.sha256.lower(): + raise ReportingInspectionError( + ReportingInspectionCode.OBJECT_DIGEST_MISMATCH, + f"reporting object {entry.object_ref!r} failed checksum validation", + ) + file_rows, decoded_bytes = _decode_rows( + body, + str(manifest.format), + str(manifest.compression), + max_decoded_bytes=self._max_object_bytes, + max_rows=self._max_total_rows - len(rows), + ) + total_decoded_bytes += decoded_bytes + if total_decoded_bytes > self._max_total_decoded_bytes: + raise ReportingInspectionError( + ReportingInspectionCode.RESOURCE_TOO_LARGE, + "manifest exceeds the configured aggregate decoded-byte limit", + ) + if len(file_rows) != entry.row_count: + raise ReportingInspectionError( + ReportingInspectionCode.ROW_COUNT_MISMATCH, + f"reporting object {entry.object_ref!r} has the wrong row count", + ) + rows.extend(file_rows) + + schema = await self._read_json_contract( + str(revision.schema_uri), + revision.schema_sha256, + "row schema", + frozenset({"application/schema+json", "application/json"}), + ) + _validate_row_schema(schema, str(revision.schema_dialect)) + try: + jsonschema.Draft202012Validator.check_schema(schema) + validator = jsonschema.Draft202012Validator(schema) + first_error = next( + (error for row in rows for error in validator.iter_errors(row)), None + ) + except (jsonschema.SchemaError, RecursionError, TypeError, ValueError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, "reporting row schema is invalid" + ) from error + if first_error is not None: + raise ReportingInspectionError( + ReportingInspectionCode.ROW_SCHEMA_INVALID, + "a reporting row failed its pinned schema", + ) + + definition_json = await self._read_json_contract( + str(revision.report_definition_uri), + revision.report_definition_sha256, + "report definition", + frozenset({"application/vnd.adcp.reporting-definition+json"}), + ) + try: + definition = ReportingReportDefinition.model_validate(definition_json) + except ValidationError as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "report definition is invalid", + ) from error + if ( + definition.report_definition_id != revision.report_definition_id + or definition.reporting_profile != revision.reporting_profile + ): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "report definition identity does not match the revision", + ) + policy_keys = [ + (policy.finality_policy_id, str(policy.basis)) + for policy in definition.finality_policies + ] + if len({policy.finality_policy_id for policy in definition.finality_policies}) != len( + definition.finality_policies + ): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "report definition has duplicate finality policy identifiers", + ) + if ( + str(revision.finality) == "official" + and ( + revision.finality_policy_id, + str(revision.finality_basis), + ) + not in policy_keys + ): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "official revision finality does not match the report definition", + ) + + observed_totals = ( + self._control_total_calculator(rows, definition) + if self._control_total_calculator + else _default_control_totals(rows, definition, revision.control_totals) + ) + if not _totals_same(observed_totals, revision.control_totals): + raise ReportingInspectionError( + ReportingInspectionCode.CONTROL_TOTAL_MISMATCH, + "recomputed control totals do not match the revision", + ) + + canonical_digest = None + if revision.canonical_content_digest: + canonical_digest = await self._canonical_digest( + rows, revision.canonical_content_digest, revision.schema_sha256 + ) + return ReportingObservation( + row_count=len(rows), + control_totals=list(observed_totals), + canonical_content_digest=canonical_digest, + manifest_sha256=manifest_digest, + ) + + async def _read_json_contract( + self, + uri: str, + expected_digest: str, + description: str, + expected_content_types: frozenset[str], + ) -> dict[str, Any]: + body = await self._read( + uri, + max_bytes=self._max_contract_bytes, + expected_content_types=expected_content_types, + ) + if _digest(body).lower() != expected_digest.lower(): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + f"{description} digest mismatch", + ) + try: + value = _json_loads_no_duplicates(body) + except _InvalidJsonError as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + f"{description} is not valid UTF-8 JSON", + ) from error + if not isinstance(value, dict): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + f"{description} root must be an object", + ) + return value + + async def _read( + self, + locator: str, + *, + max_bytes: int, + base: str | None = None, + expected_content_types: frozenset[str] | None = None, + ) -> bytes: + """Use typed content checks for the built-in HTTPS reader. + + The small resource-reader protocol intentionally remains byte-only so + existing destination adapters stay compatible; those adapters own + equivalent media-type validation for their native transports. + """ + if isinstance(self._reader, HttpsReportingResourceReader): + return await self._reader.read( + locator, + base=base, + max_bytes=max_bytes, + expected_content_types=expected_content_types, + ) + return await self._reader.read(locator, base=base, max_bytes=max_bytes) + + async def _canonical_digest( + self, + rows: list[dict[str, Any]], + expected: ReportingCanonicalContentDigest, + schema_sha256: str, + ) -> ReportingCanonicalContentDigest: + contract_json = await self._read_json_contract( + str(expected.canonicalization_uri), + expected.canonicalization_sha256, + "canonicalization contract", + frozenset({"application/vnd.adcp.reporting-canonicalization+json"}), + ) + try: + contract = ReportingCanonicalizationContract.model_validate(contract_json) + except ValidationError as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "canonicalization contract is invalid", + ) from error + if contract.schema_sha256.lower() != schema_sha256.lower(): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "canonicalization contract targets a different row schema", + ) + keys = [item.root for item in contract.primary_keys] + try: + _validate_golden_vectors(contract, keys) + value = hashlib.sha256(_canonical_rows_bytes(rows, keys)).hexdigest() + except ReportingInspectionError: + raise + except (KeyError, TypeError, ValueError, rfc8785.CanonicalizationError) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, + "reporting rows cannot be canonicalized with the declared primary keys", + ) from error + if value.lower() != expected.value.lower(): + raise ReportingInspectionError( + ReportingInspectionCode.CANONICAL_DIGEST_MISMATCH, + "canonical logical-content digest does not match the revision", + ) + return expected + + +def _object_content_types(format_name: str) -> frozenset[str]: + if format_name == "jsonl": + return frozenset({"application/x-ndjson", "application/ndjson", "application/jsonl"}) + if format_name == "csv": + return frozenset({"text/csv"}) + return frozenset() + + +def _canonical_rows_bytes(rows: list[dict[str, Any]], keys: list[str]) -> bytes: + """Canonicalize rows using the byte ordering required by adcp_jcs_rows_v1.""" + encoded: list[tuple[bytes, bytes]] = [] + for row in rows: + identity = [row[key] for key in keys] + if any(isinstance(value, (dict, list)) for value in identity): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, + "reporting primary keys must be scalar values", + ) + encoded.append((rfc8785.dumps(identity), rfc8785.dumps(row))) + encoded.sort(key=lambda item: item[0]) + if any(left[0] == right[0] for left, right in zip(encoded, encoded[1:])): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_ROWS, + "reporting rows contain duplicate primary keys", + ) + return b"[" + b",".join(row for _identity, row in encoded) + b"]" + + +def _validate_golden_vectors(contract: ReportingCanonicalizationContract, keys: list[str]) -> None: + vectors: list[Any] = [ + contract.golden_vectors.empty_report, + contract.golden_vectors.ordering_encoding, + *(contract.golden_vectors.additional or []), + ] + if len({vector.name for vector in vectors}) != len(vectors): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "canonicalization contract has duplicate golden-vector names", + ) + for vector in vectors: + try: + declared = base64.b64decode(vector.canonical_utf8_base64, validate=True) + reproduced = _canonical_rows_bytes(vector.input_rows, keys) + except ( + binascii.Error, + ValueError, + TypeError, + KeyError, + rfc8785.CanonicalizationError, + ) as error: + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "canonicalization contract has an invalid golden vector", + ) from error + if ( + declared != reproduced + or hashlib.sha256(declared).hexdigest().lower() != vector.sha256.lower() + ): + raise ReportingInspectionError( + ReportingInspectionCode.INVALID_CONTRACT, + "canonicalization contract golden vector does not match adcp_jcs_rows_v1", + ) + + +__all__ = [ + "ControlTotalCalculator", + "HttpsReportingResourceReader", + "ManifestReportingInspector", + "ReportingCredentialProvider", + "ReportingInspectionCode", + "ReportingInspectionError", + "ReportingResourceReader", + "ReportingTrustedOriginPolicy", +] diff --git a/src/adcp/server/__init__.py b/src/adcp/server/__init__.py index 1a34d8a5d..eb536e8d6 100644 --- a/src/adcp/server/__init__.py +++ b/src/adcp/server/__init__.py @@ -130,6 +130,17 @@ async def get_products(params, context=None): register_handler_tools, validate_discovery_set, ) +from adcp.server.principal import ( + DestinationProofHook, + InMemoryPrincipalRecordStore, + NotificationProofHook, + PrincipalChange, + PrincipalChangeEmitter, + PrincipalIdentity, + PrincipalRecord, + PrincipalRecordStore, + PrincipalService, +) from adcp.server.proposal import ProposalBuilder, ProposalNotSupported from adcp.server.responses import ( activate_signal_response, @@ -211,6 +222,16 @@ async def get_products(params, context=None): "GovernanceEnforcementDecision", "make_governance_enforcement_middleware", "SponsoredIntelligenceHandler", + # Principal state primitives + "DestinationProofHook", + "InMemoryPrincipalRecordStore", + "NotificationProofHook", + "PrincipalChange", + "PrincipalChangeEmitter", + "PrincipalIdentity", + "PrincipalRecord", + "PrincipalRecordStore", + "PrincipalService", # Proposal helpers "ProposalBuilder", "ProposalNotSupported", diff --git a/src/adcp/server/a2a_server.py b/src/adcp/server/a2a_server.py index f4c6ce5f2..729921e3a 100644 --- a/src/adcp/server/a2a_server.py +++ b/src/adcp/server/a2a_server.py @@ -399,6 +399,7 @@ async def _call_test_controller( params, context=context, account_resolver=resolver, + validate_schema=True, ) # This skill bypasses ``create_tool_caller`` (the success-path # enhancer site), so apply the enhancer here too — otherwise diff --git a/src/adcp/server/principal.py b/src/adcp/server/principal.py new file mode 100644 index 000000000..ffc29a9f4 --- /dev/null +++ b/src/adcp/server/principal.py @@ -0,0 +1,818 @@ +"""State primitives for implementing the AdCP 3.2 principal layer. + +This module deliberately does not authenticate callers or prove control of +external resources. The server resolves a stable transport identity first, +then passes it here; adopter hooks perform endpoint and provider proof before +transitioning a destination from ``validating`` to ``ready``. +""" + +from __future__ import annotations + +import asyncio +import json +import re +import unicodedata +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, Protocol +from urllib.parse import urlsplit, urlunsplit +from uuid import uuid4 + +from pydantic import BaseModel + +from adcp.server.idempotency import canonical_json_sha256 +from adcp.types import ( + AgentDeclarations, + AgentNotificationConfig, + AgentNotificationConfigState, + AgentReportingDestination, + AgentReportingDestinationState, + GetPrincipalRequest, + GetPrincipalResponse, + PrincipalConfiguration, + PrincipalDeclarationsState, + PrincipalState, + SyncPrincipalRequest, + SyncPrincipalResponse, +) + + +@dataclass(frozen=True) +class PrincipalIdentity: + """Seller-resolved stable identity; never derive this from request args.""" + + subject: str + kind: str + + +@dataclass(frozen=True) +class PrincipalChange: + """State invalidation supplied to an adopter-owned webhook emitter.""" + + principal_id: str + subject: str + changed_at: datetime + reason: str + destination_id: str | None = None + + +PrincipalChangeEmitter = Callable[[PrincipalChange], Awaitable[None]] +DestinationProofHook = Callable[ + [PrincipalIdentity, AgentReportingDestination, AgentReportingDestinationState | None], + Awaitable[AgentReportingDestinationState], +] +NotificationProofHook = Callable[ + [PrincipalIdentity, AgentNotificationConfig], Awaitable[AgentNotificationConfigState] +] + + +@dataclass +class PrincipalRecord: + subject: str + principal_id: str + principal_kind: str + configuration_version: str | None = None + configuration: PrincipalState | None = None + idempotency: dict[str, tuple[str, SyncPrincipalResponse]] = field(default_factory=dict) + # This is an internal revision, separate from the caller-visible configuration + # version. Idempotency writes and seller-driven state transitions need a CAS + # fence even when they intentionally do not change configuration_version. + store_revision: int = 0 + + +class PrincipalRecordStore(Protocol): + async def get(self, subject: str) -> PrincipalRecord | None: + raise NotImplementedError + + async def compare_and_swap( + self, + subject: str, + expected_store_revision: int | None, + record: PrincipalRecord, + ) -> bool: + """Atomically replace one record when its internal revision still matches. + + ``None`` means the record must not yet exist. Implementations must make + this operation atomic across all workers sharing the backing store. + """ + raise NotImplementedError + + +class InMemoryPrincipalRecordStore: + """Process-local reference store suitable for tests and single workers.""" + + def __init__(self) -> None: + self._records: dict[str, PrincipalRecord] = {} + self._lock = asyncio.Lock() + + async def get(self, subject: str) -> PrincipalRecord | None: + record = self._records.get(subject) + return _copy_record(record) if record else None + + async def compare_and_swap( + self, + subject: str, + expected_store_revision: int | None, + record: PrincipalRecord, + ) -> bool: + async with self._lock: + existing = self._records.get(subject) + if expected_store_revision is None: + if existing is not None: + return False + elif existing is None or existing.store_revision != expected_store_revision: + return False + self._records[subject] = _copy_record(record) + return True + + +def _copy_record(record: PrincipalRecord) -> PrincipalRecord: + return PrincipalRecord( + subject=record.subject, + principal_id=record.principal_id, + principal_kind=record.principal_kind, + configuration_version=record.configuration_version, + configuration=( + record.configuration.model_copy(deep=True) if record.configuration else None + ), + idempotency={ + key: (digest, response.model_copy(deep=True)) + for key, (digest, response) in record.idempotency.items() + }, + store_revision=record.store_revision, + ) + + +def _json(value: object) -> str: + if isinstance(value, BaseModel): + value = value.model_dump(mode="json", exclude_none=True) + return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str) + + +def _failed(code: str, message: str, *, context: object | None = None) -> SyncPrincipalResponse: + return SyncPrincipalResponse.model_validate( + { + "status": "rejected", + "result": {"kind": "failed", "errors": [{"code": code, "message": message}]}, + "context": context, + } + ) + + +def _notification_state(config: AgentNotificationConfig) -> AgentNotificationConfigState: + value = config.model_dump(mode="json", exclude_none=True) + authentication = value.get("authentication") + if isinstance(authentication, dict): + authentication.pop("credentials", None) + return AgentNotificationConfigState.model_validate(value) + + +class _PrincipalValidationError(ValueError): + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + + +_MEDIA_BUY_NOTIFICATION_TYPES = frozenset( + {"scheduled", "final", "delayed", "adjusted", "window_update", "impairment"} +) +_ACCOUNT_NOTIFICATION_TYPES = frozenset( + { + "creative.status_changed", + "creative.assignment_changed", + "indicators.changed", + "creative.purged", + "account.status_changed", + "account.change_recorded", + "product.created", + "product.updated", + "product.priced", + "product.removed", + "signal.created", + "signal.updated", + "signal.priced", + "signal.removed", + "wholesale_feed.bulk_change", + "reporting.delivery_ready", + "reporting.status_changed", + } +) +_SECRET_COORDINATE_RE = re.compile( + r"(?:bearer\s+|private[ _-]?key|(?:api|access)[ _-]?key\s*[=:]|" + r"(?:token|secret|password|signature|sig)\s*[=:])", + re.IGNORECASE, +) + + +def _require_unique(values: Sequence[object], attribute: str, section: str) -> None: + seen: set[str] = set() + for index, value in enumerate(values): + key = str(getattr(value, attribute)) + if key in seen: + raise ValueError(f"{section}[{index}].{attribute} duplicates an earlier entry") + seen.add(key) + + +def _validate_webhook_destination_url(url: str, *, field: str) -> Any: + """Lazily invoke webhook validation without creating a facade import cycle.""" + from adcp.webhooks import validate_webhook_destination_url + + return validate_webhook_destination_url(url, field=field) + + +def _normalize_notification_config( + config: AgentNotificationConfig, index: int +) -> AgentNotificationConfig: + event_types = {str(item) for item in config.event_types} + if event_types & _MEDIA_BUY_NOTIFICATION_TYPES: + raise ValueError( + f"notification_configs[{index}].event_types contains a media-buy-anchored event" + ) + if event_types & _ACCOUNT_NOTIFICATION_TYPES and config.all_authorized_accounts is not True: + raise ValueError( + f"notification_configs[{index}].all_authorized_accounts must be true for account events" + ) + if _SECRET_COORDINATE_RE.search(str(config.url)): + raise ValueError(f"notification_configs[{index}].url must not embed a credential or secret") + try: + # This performs registration-time HTTPS, normalized-host, and + # reserved-range checks for active *and* inactive registrations. + validation = _validate_webhook_destination_url( + str(config.url), field=f"notification_configs[{index}].url" + ) + except Exception as error: + raise ValueError(str(error)) from error + parsed = urlsplit(validation.effective_url) + port = parsed.port + netloc = validation.hostname + if port is not None and port != 443: + netloc = f"{netloc}:{port}" + normalized_url = urlunsplit(("https", netloc, parsed.path, parsed.query, "")) + payload = config.model_dump(mode="json", exclude_none=True) + payload["url"] = normalized_url + return AgentNotificationConfig.model_validate(payload) + + +def _normalize_coordinate(value: str, *, field: str) -> str: + """Reject secret-bearing locators and normalize URL-like coordinates.""" + normalized = unicodedata.normalize("NFC", value) + if _SECRET_COORDINATE_RE.search(normalized): + raise ValueError(f"{field} must not contain credentials or a secret") + parsed = urlsplit(normalized) + if not parsed.scheme: + return normalized + if parsed.query or parsed.fragment: + raise ValueError(f"{field} must not include query or fragment components") + # ``abfss://container@account/...`` is a legitimate provider locator; a + # colon in the authority is what distinguishes credential-style userinfo. + authority = parsed.netloc.rsplit("@", 1)[0] if "@" in parsed.netloc else "" + scheme = parsed.scheme.lower() + if authority and (":" in authority or scheme not in {"abfs", "abfss"}): + raise ValueError(f"{field} must not embed userinfo credentials") + if parsed.hostname is None: + raise ValueError(f"{field} must have a host when it is URL-like") + host = parsed.hostname.encode("idna").decode("ascii").lower().rstrip(".") + try: + port = parsed.port + except ValueError as error: + raise ValueError(f"{field} has an invalid port") from error + default_port = {"http": 80, "https": 443}.get(scheme) + netloc = host if port is None or port == default_port else f"{host}:{port}" + if authority: + netloc = f"{authority}@{netloc}" + path = re.sub(r"/{2,}", "/", parsed.path) + return urlunsplit((scheme, netloc, path, "", "")) + + +def _normalize_destination(destination: AgentReportingDestination) -> AgentReportingDestination: + payload = destination.model_dump(mode="json", exclude_none=True) + if "location" in payload: + payload["location"] = _normalize_coordinate( + payload["location"], field="reporting_destinations.location" + ) + recipient = payload.get("recipient") + if isinstance(recipient, dict): + identity = recipient.get("identity") + if isinstance(identity, str): + recipient["identity"] = _normalize_coordinate( + identity, field="reporting_destinations.recipient.identity" + ) + return AgentReportingDestination.model_validate(payload) + + +class PrincipalService: + """Implement atomic principal reads, guarded section sync, and proof state.""" + + def __init__( + self, + store: PrincipalRecordStore | None = None, + *, + destination_proof: DestinationProofHook | None = None, + notification_proof: NotificationProofHook | None = None, + emit_change: PrincipalChangeEmitter | None = None, + accepted_async_adcp_versions: tuple[str, ...] = ("3.2",), + accepted_webhook_signing_algorithms: tuple[str, ...] = (), + accepted_experimental_features: tuple[str, ...] = (), + ) -> None: + self._store = store or InMemoryPrincipalRecordStore() + self._destination_proof = destination_proof + self._notification_proof = notification_proof + self._emit_change = emit_change + self._accepted_declarations = { + "async_adcp_versions": accepted_async_adcp_versions, + "webhook_signing_algorithms": accepted_webhook_signing_algorithms, + "experimental_features": accepted_experimental_features, + } + + def set_declaration_support( + self, + *, + async_adcp_versions: tuple[str, ...] | None = None, + webhook_signing_algorithms: tuple[str, ...] | None = None, + experimental_features: tuple[str, ...] | None = None, + ) -> None: + """Update objective seller support before refreshing affected records.""" + replacements = { + "async_adcp_versions": async_adcp_versions, + "webhook_signing_algorithms": webhook_signing_algorithms, + "experimental_features": experimental_features, + } + for axis, values in replacements.items(): + if values is not None: + self._accepted_declarations[axis] = values + + async def recognize(self, identity: PrincipalIdentity) -> str: + """Materialize a durable recognized principal without configuration.""" + while True: + record = await self._store.get(identity.subject) + if record is not None: + return record.principal_id + record = PrincipalRecord( + identity.subject, f"prin:{uuid4()}", identity.kind, store_revision=1 + ) + if await self._store.compare_and_swap(identity.subject, None, record): + return record.principal_id + + async def get_principal( + self, + identity: PrincipalIdentity, + request: GetPrincipalRequest | Mapping[str, Any] | None = None, + ) -> GetPrincipalResponse: + request = GetPrincipalRequest.model_validate(request or {}) + record = await self._store.get(identity.subject) + if record is None: + return GetPrincipalResponse.model_validate( + {"result": {"kind": "unconfigured"}, "context": request.context} + ) + if record.principal_kind != identity.kind: + return GetPrincipalResponse.model_validate( + { + "status": "rejected", + "result": { + "kind": "failed", + "errors": [ + { + "code": "CONFLICT", + "message": "authenticated principal kind changed", + } + ], + }, + "context": request.context, + } + ) + if record.configuration is None or record.configuration_version is None: + return GetPrincipalResponse.model_validate( + { + "result": { + "kind": "recognized", + "principal_id": record.principal_id, + "principal_kind": record.principal_kind, + }, + "context": request.context, + } + ) + return self._current_response(record, context=request.context) + + async def sync_principal( + self, + identity: PrincipalIdentity, + request: SyncPrincipalRequest | Mapping[str, Any], + ) -> SyncPrincipalResponse: + request = SyncPrincipalRequest.model_validate(request) + if not request.configuration.model_fields_set: + return _failed( + "INVALID_REQUEST", + "configuration must contain at least one section", + context=request.context, + ) + request_digest = canonical_json_sha256(request.model_dump(mode="json", exclude_none=True)) + + while True: + record = await self._store.get(identity.subject) + if record is not None: + replay = record.idempotency.get(request.idempotency_key) + if replay: + digest, response = replay + if digest == request_digest: + return response.model_copy(update={"replayed": True}, deep=True) + return _failed( + "IDEMPOTENCY_CONFLICT", + "idempotency key was reused with a new request", + context=request.context, + ) + if request.expected_principal_kind is not None and ( + str(request.expected_principal_kind) != identity.kind + ): + return _failed( + "CONFLICT", + "expected_principal_kind is stale", + context=request.context, + ) + if request.expected_configuration_version is not None and ( + record is None + or request.expected_configuration_version != record.configuration_version + ): + return _failed( + "CONFLICT", + "expected_configuration_version is stale", + context=request.context, + ) + + previous = record.configuration if record else None + try: + next_state = await self._replace_sections( + identity, + previous, + request.configuration, + dry_run=bool(request.dry_run), + ) + except _PrincipalValidationError as error: + return _failed(error.code, str(error), context=request.context) + except ValueError as error: + return _failed("INVALID_REQUEST", str(error), context=request.context) + changed = _json(previous) != _json(next_state) + cleared = self._submitted_sections_are_empty(request.configuration) + if request.dry_run: + action = ( + "would_clear" + if changed and cleared + else "would_update" if changed else "would_be_unchanged" + ) + return SyncPrincipalResponse.model_validate( + { + "result": {"kind": "validated", "action": action, "dry_run": True}, + "context": request.context, + } + ) + + if record is None: + record = PrincipalRecord(identity.subject, f"prin:{uuid4()}", identity.kind) + expected_store_revision: int | None = None + else: + expected_store_revision = record.store_revision + record.configuration = next_state + if changed or record.configuration_version is None: + record.configuration_version = f"cfg:{uuid4()}" + action = "cleared" if changed and cleared else "updated" if changed else "unchanged" + response = SyncPrincipalResponse.model_validate( + { + "result": { + "kind": "applied", + "action": action, + "dry_run": False, + "principal_id": record.principal_id, + "principal_kind": record.principal_kind, + "configuration_version": record.configuration_version, + "configuration": record.configuration, + }, + "context": request.context, + } + ) + record.idempotency[request.idempotency_key] = (request_digest, response) + record.store_revision = (expected_store_revision or 0) + 1 + if await self._store.compare_and_swap( + identity.subject, expected_store_revision, record + ): + return response + # A different process committed between our read and write. Repeat + # from the authoritative record so stale fences reject and an exact + # same-key retry returns the winner's cached response. + + async def transition_destination( + self, + identity: PrincipalIdentity, + destination_id: str, + state: str, + *, + setup: Mapping[str, Any] | None = None, + issues: list[Mapping[str, Any]] | None = None, + reason: str = "destination_state_changed", + ) -> AgentReportingDestinationState: + """Record an adopter-verified provider transition without version churn.""" + while True: + record = await self._store.get(identity.subject) + if not record or not record.configuration: + raise KeyError(destination_id) + if record.principal_kind != identity.kind: + raise ValueError("authenticated principal kind changed") + destinations = list(record.configuration.reporting_destinations or []) + index = next( + (i for i, item in enumerate(destinations) if item.destination_id == destination_id), + None, + ) + if index is None: + raise KeyError(destination_id) + current = destinations[index] + allowed = { + "validating": {"ready", "action_required", "inactive", "rejected"}, + "action_required": {"validating", "ready", "inactive", "rejected"}, + "ready": {"action_required", "inactive", "rejected"}, + "inactive": {"validating"}, + "rejected": set(), + } + current_state = str(current.state) + if state != current_state and state not in allowed[current_state]: + raise ValueError(f"invalid destination transition {current_state!r} -> {state!r}") + payload = current.model_dump(mode="json", exclude_none=True) + payload.update(state=state) + if setup is not None: + payload["setup"] = dict(setup) + elif state in {"ready", "inactive", "rejected"}: + payload.pop("setup", None) + if issues is not None: + payload["issues"] = issues + destinations[index] = AgentReportingDestinationState.model_validate(payload) + state_payload = record.configuration.model_dump(mode="json", exclude_none=True) + state_payload["reporting_destinations"] = destinations + record.configuration = PrincipalState.model_validate(state_payload) + expected_store_revision = record.store_revision + record.store_revision += 1 + if await self._store.compare_and_swap( + identity.subject, expected_store_revision, record + ): + break + + if self._emit_change: + await self._emit_change( + PrincipalChange( + record.principal_id, + identity.subject, + datetime.now(timezone.utc), + reason, + destination_id, + ) + ) + return destinations[index] + + async def refresh_declarations( + self, identity: PrincipalIdentity + ) -> PrincipalDeclarationsState | None: + """Recompute a persisted declaration intersection after seller changes. + + Like destination proof transitions, this seller-driven state change + emits ``principal.changed`` but does not advance the caller-owned + ``configuration_version``. + """ + changed = False + declarations: PrincipalDeclarationsState | None = None + while True: + record = await self._store.get(identity.subject) + if not record or not record.configuration: + return None + current = record.configuration.declarations + if current is None: + return None + declarations = self._negotiate_declarations(current.declared) + changed = _json(current) != _json(declarations) + if changed: + payload = record.configuration.model_dump(mode="json", exclude_none=True) + payload["declarations"] = declarations + record.configuration = PrincipalState.model_validate(payload) + expected_store_revision = record.store_revision + record.store_revision += 1 + if not await self._store.compare_and_swap( + identity.subject, expected_store_revision, record + ): + continue + break + if changed and self._emit_change: + await self._emit_change( + PrincipalChange( + record.principal_id, + identity.subject, + datetime.now(timezone.utc), + "declarations_intersection_changed", + ) + ) + return declarations + + async def _replace_sections( + self, + identity: PrincipalIdentity, + previous: PrincipalState | None, + desired: PrincipalConfiguration, + *, + dry_run: bool, + ) -> PrincipalState: + payload = previous.model_dump(mode="json", exclude_none=True) if previous else {} + fields = desired.model_fields_set + if "notification_configs" in fields: + configs = desired.notification_configs or [] + _require_unique(configs, "subscriber_id", "notification_configs") + notification_states: list[AgentNotificationConfigState] = [] + for index, config in enumerate(configs): + if config.active and self._notification_proof is None and not dry_run: + raise ValueError( + "active notification configs require a notification_proof hook" + ) + config = _normalize_notification_config(config, index) + notification_state = ( + await self._notification_proof(identity, config) + if self._notification_proof and not dry_run + else _notification_state(config) + ) + if notification_state.subscriber_id != config.subscriber_id: + raise ValueError( + "notification_proof returned a state for a different subscriber" + ) + if _json(notification_state) != _json(_notification_state(config)): + raise ValueError( + "notification_proof returned a state for a different notification contract" + ) + notification_states.append(notification_state) + payload["notification_configs"] = notification_states + + if "reporting_destinations" in fields: + requested_destinations = [ + _normalize_destination(destination) + for destination in desired.reporting_destinations or [] + ] + _require_unique(requested_destinations, "destination_id", "reporting_destinations") + prior = { + item.destination_id: item + for item in (previous.reporting_destinations if previous else None) or [] + } + destination_states: list[AgentReportingDestinationState] = [] + for destination in requested_destinations: + old = prior.pop(destination.destination_id, None) + if old and self._same_destination_contract(old.configuration, destination): + if _json(old.configuration) == _json(destination): + destination_states.append(old) + continue + updated = old.model_dump(mode="json", exclude_none=True) + updated["configuration"] = destination + updated["state"] = "validating" if destination.active else "inactive" + destination_states.append( + AgentReportingDestinationState.model_validate(updated) + ) + continue + if self._destination_proof and not dry_run: + proved = await self._destination_proof(identity, destination, old) + if ( + proved.destination_id != destination.destination_id + or _json(proved.configuration) != _json(destination) + or (not destination.active and str(proved.state) != "inactive") + ): + raise ValueError( + "destination_proof returned a state for a different " + "destination contract" + ) + destination_states.append(proved) + continue + refs = [] + if old: + refs = [old.destination_ref, *(old.prior_destination_refs or [])] + destination_states.append( + AgentReportingDestinationState.model_validate( + { + "destination_id": destination.destination_id, + "destination_ref": f"dest:{uuid4()}", + "prior_destination_refs": refs or None, + "state": "validating" if destination.active else "inactive", + "configuration": destination, + } + ) + ) + retired: list[dict[str, Any]] = [ + item.model_dump(mode="json", exclude_none=True) + for item in (previous.retired_destinations if previous else None) or [] + ] + now = datetime.now(timezone.utc) + for destination_id, old in prior.items(): + retired.append( + { + "destination_id": destination_id, + "destination_refs": [ + old.destination_ref, + *(old.prior_destination_refs or []), + ], + "revoked_at": now, + } + ) + payload["reporting_destinations"] = destination_states + payload["retired_destinations"] = retired + + if "declarations" in fields: + payload["declarations"] = self._negotiate_declarations( + desired.declarations or AgentDeclarations() + ) + state = PrincipalState.model_validate(payload) + active_notifications = [ + config for config in state.notification_configs or [] if config.active is not False + ] + declarations = state.declarations + accepted_algorithms = ( + declarations.accepted.webhook_signing_algorithms if declarations else None + ) + if active_notifications and not accepted_algorithms: + raise _PrincipalValidationError( + "UNSUPPORTED_FEATURE", + "active notification configs require an accepted webhook signing algorithm", + ) + return state + + def _negotiate_declarations(self, declared: AgentDeclarations) -> PrincipalDeclarationsState: + raw = declared.model_dump(mode="json", exclude_none=True) + accepted: dict[str, list[str]] = {} + exclusions: list[dict[str, str]] = [] + for axis, offered in raw.items(): + supported = set(self._accepted_declarations[axis]) + intersection = [value for value in offered if value in supported] + # An omitted field represents an empty set: generated declaration + # fields correctly reject a present empty array. + if intersection: + accepted[axis] = intersection + exclusions.extend( + { + "axis": axis, + "value": value, + "reason": "unsupported by this seller", + } + for value in offered + if value not in supported + ) + selected = next(iter(accepted.get("async_adcp_versions", [])), None) + return PrincipalDeclarationsState.model_validate( + { + "declared": declared, + "accepted": accepted, + "selected_async_adcp_version": selected, + "exclusions": exclusions or None, + } + ) + + @staticmethod + def _submitted_sections_are_empty(configuration: PrincipalConfiguration) -> bool: + fields = configuration.model_fields_set + if not fields: + return False + return ( + all( + getattr(configuration, field) == [] + for field in fields + if field in {"notification_configs", "reporting_destinations"} + ) + and "declarations" not in fields + ) + + @staticmethod + def _same_destination_contract( + previous: AgentReportingDestination, + desired: AgentReportingDestination, + ) -> bool: + old = previous.model_dump(mode="json", exclude_none=True) + new = desired.model_dump(mode="json", exclude_none=True) + old.pop("active", None) + new.pop("active", None) + return bool(old == new) + + @staticmethod + def _current_response( + record: PrincipalRecord, *, context: object | None = None + ) -> GetPrincipalResponse: + return GetPrincipalResponse.model_validate( + { + "result": { + "kind": "current", + "principal_id": record.principal_id, + "principal_kind": record.principal_kind, + "configuration_version": record.configuration_version, + "configuration": record.configuration, + }, + "context": context, + } + ) + + +__all__ = [ + "DestinationProofHook", + "InMemoryPrincipalRecordStore", + "NotificationProofHook", + "PrincipalChange", + "PrincipalChangeEmitter", + "PrincipalIdentity", + "PrincipalRecord", + "PrincipalRecordStore", + "PrincipalService", +] diff --git a/src/adcp/server/test_controller.py b/src/adcp/server/test_controller.py index e862a2871..2e5c081dd 100644 --- a/src/adcp/server/test_controller.py +++ b/src/adcp/server/test_controller.py @@ -101,10 +101,14 @@ def __call__(self, ref: dict[str, Any] | None, *, auth_info: AuthInfo | None = N # Scenario names — must match the AdCP comply_test_controller schema SCENARIOS = [ + "expire_account_change_cursor", "force_creative_status", + "force_creative_purge", "force_account_status", "force_media_buy_status", "force_create_media_buy_arm", + "force_get_products_arm", + "force_get_signals_arm", "force_task_completion", "force_session_status", "simulate_delivery", @@ -115,13 +119,46 @@ def __call__(self, ref: dict[str, Any] | None, *, auth_info: AuthInfo | None = N "seed_creative", "seed_plan", "seed_media_buy", + "seed_account", + "seed_rights_grant", "seed_creative_format", + "seed_measurement_catalog", + "query_upstream_traffic", + "query_provenance_audit_observations", + "force_upstream_unavailable", + "catalog_item_availability_probe", + "compact_product_lifecycle_probe", + "compact_direct_buy_lifecycle_probe", ] _MAX_TASK_ID = 128 _MAX_MESSAGE = 2000 _MAX_RESULT_BYTES = 256 * 1024 # 256 KB soft cap per AdCP 3.0.1 +# Before the dispatcher became signature-driven, these optional arguments +# were always supplied with ``None`` when absent. Preserve that behavior for +# existing store overrides whose signatures made the arguments positional/ +# required even though the wire fields are optional. +_LEGACY_OPTIONAL_SCENARIO_PARAMS: dict[str, tuple[str, ...]] = { + "force_creative_status": ("rejection_reason",), + "force_media_buy_status": ("rejection_reason",), + "force_session_status": ("termination_reason",), + "force_create_media_buy_arm": ("task_id", "message"), + "simulate_delivery": ( + "impressions", + "clicks", + "conversions", + "reported_spend", + ), + "simulate_budget_spend": ("account_id", "media_buy_id"), + "seed_product": ("fixture", "product_id"), + "seed_pricing_option": ("fixture", "product_id", "pricing_option_id"), + "seed_creative": ("fixture", "creative_id"), + "seed_plan": ("fixture", "plan_id"), + "seed_media_buy": ("fixture", "media_buy_id"), + "seed_creative_format": ("fixture", "format_id"), +} + class TestControllerError(Exception): """Typed error for test controller store methods. @@ -168,11 +205,28 @@ class TestControllerStore: Stores that don't declare ``context`` keep working unchanged. """ + async def expire_account_change_cursor( + self, + account_id: str, + *, + account: dict[str, Any] | None = None, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Rotate an account's authorization-scope epoch. + + Returns a state transition describing the previous and current + cursor epochs. ``account`` is the verified sandbox account assertion + carried by the controller request. + """ + raise NotImplementedError + async def force_creative_status( self, creative_id: str, status: str, rejection_reason: str | None = None, + reason_code: str | None = None, + reason_detail: str | None = None, *, context: ToolContext | None = None, ) -> dict[str, Any]: @@ -183,6 +237,22 @@ async def force_creative_status( """ raise NotImplementedError + async def force_creative_purge( + self, + creative_id: str, + purge_kind: str | None = None, + reason_code: str | None = None, + reason_detail: str | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Soft-delete or permanently purge a sandbox creative. + + Returns: + {"previous_state": str, "current_state": str} + """ + raise NotImplementedError + async def force_account_status( self, account_id: str, @@ -266,6 +336,43 @@ async def force_create_media_buy_arm( """ raise NotImplementedError + async def force_get_products_arm( + self, + arm: str, + task_id: str | None = None, + message: str | None = None, + reason: str | None = None, + suggestions: list[str] | None = None, + *, + account: dict[str, Any] | None = None, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Shape the next ``get_products`` call into a deterministic arm. + + ``submitted`` requires ``task_id``; ``rejected`` requires ``reason`` + and may include buyer-facing ``suggestions``. + + Returns: + {"forced": {"arm": str, ...}} + """ + raise NotImplementedError + + async def force_get_signals_arm( + self, + arm: str, + task_id: str, + message: str | None = None, + *, + account: dict[str, Any] | None = None, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Shape the next ``get_signals`` call into the submitted arm. + + Returns: + {"forced": {"arm": "submitted", "task_id": str}} + """ + raise NotImplementedError + async def force_task_completion( self, task_id: str, @@ -315,8 +422,21 @@ async def simulate_delivery( media_buy_id: str, impressions: int | None = None, clicks: int | None = None, + plays: int | None = None, + dooh_metrics: dict[str, Any] | None = None, conversions: int | None = None, + delivery_date: str | None = None, + conversion_value: float | None = None, + commissionable_value: float | None = None, reported_spend: dict[str, Any] | None = None, + reach: float | None = None, + frequency: float | None = None, + reach_window: dict[str, Any] | None = None, + viewability: dict[str, Any] | None = None, + vendor_metric_values: list[dict[str, Any]] | None = None, + vendor_metric_values_by_package: dict[str, list[dict[str, Any]]] | None = None, + not_yet_measurable_vendor_metrics: list[dict[str, Any]] | None = None, + not_yet_measurable_vendor_metrics_by_package: dict[str, list[dict[str, Any]]] | None = None, *, context: ToolContext | None = None, ) -> dict[str, Any]: @@ -413,6 +533,26 @@ async def seed_media_buy( """ raise NotImplementedError + async def seed_account( + self, + account_id: str, + fixture: dict[str, Any] | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Pre-populate an advertiser account fixture.""" + raise NotImplementedError + + async def seed_rights_grant( + self, + rights_id: str, + fixture: dict[str, Any] | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Pre-populate a rights-grant fixture.""" + raise NotImplementedError + async def seed_creative_format( self, fixture: dict[str, Any] | None = None, @@ -430,6 +570,84 @@ async def seed_creative_format( """ raise NotImplementedError + async def seed_measurement_catalog( + self, + vendor: dict[str, Any], + metrics: list[dict[str, Any]], + fixture: dict[str, Any] | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Seed a measurement vendor's metric catalog.""" + raise NotImplementedError + + async def query_upstream_traffic( + self, + since_timestamp: str | None = None, + endpoint_pattern: str | None = None, + limit: int | None = None, + attestation_mode: str | None = None, + identifier_value_digests: list[str] | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Return caller-scoped outbound calls recorded by the sandbox.""" + raise NotImplementedError + + async def query_provenance_audit_observations( + self, + creative_id: str, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Return sandbox audit observations recorded for a creative.""" + raise NotImplementedError + + async def force_upstream_unavailable( + self, + tool: str, + upstream_name: str | None = None, + cache_age_seconds: int | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Mark a tool's upstream dependency unavailable for the session.""" + raise NotImplementedError + + async def catalog_item_availability_probe( + self, + operation: str, + catalog_id: str, + item_id: str, + catalog_generation: str | None = None, + target_time: str | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Apply a deterministic catalog availability probe operation.""" + raise NotImplementedError + + async def compact_product_lifecycle_probe( + self, + operation: str, + product_id: str | None = None, + proposal_id: str | None = None, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Prepare compact proposal lifecycle state or expire a proposal.""" + raise NotImplementedError + + async def compact_direct_buy_lifecycle_probe( + self, + operation: str, + product_id: str, + *, + context: ToolContext | None = None, + ) -> dict[str, Any]: + """Prepare deterministic compact direct-buy lifecycle state.""" + raise NotImplementedError + def _list_scenarios(store: TestControllerStore) -> list[str]: """Detect which scenarios a store actually implements. @@ -501,6 +719,45 @@ def _accepts_context_kwarg(method: Any) -> bool: return _accepts_kwarg(method, "context") +def _accepted_scenario_kwargs(method: Any, params: dict[str, Any]) -> dict[str, Any]: + """Return wire scenario params accepted by a store override. + + Scenario payloads are additive: the protocol can introduce an optional + parameter before the SDK publishes a matching base-class signature. The + dispatcher therefore follows the override's signature instead of keeping + a second, hand-maintained parameter allowlist. Explicit parameters and + ``**kwargs`` are both opt-ins; transport-level ``account`` and ``context`` + are reserved for the dispatcher's separately verified values. + """ + try: + signature = inspect.signature(method) + except (TypeError, ValueError): + return {} + + reserved = {"account", "context"} + parameters = signature.parameters.values() + if any(param.kind == inspect.Parameter.VAR_KEYWORD for param in parameters): + return {name: value for name, value in params.items() if name not in reserved} + + allowed_kinds = { + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.KEYWORD_ONLY, + } + accepted = { + param.name + for param in signature.parameters.values() + if param.kind in allowed_kinds and param.name not in reserved + } + return {name: value for name, value in params.items() if name in accepted} + + +def _require_scenario_params(params: dict[str, Any], *names: str) -> None: + """Raise ``KeyError`` when a scenario omits one of its required params.""" + for name in names: + if name not in params: + raise KeyError(name) + + def _extract_auth_info_from_context(context: ToolContext | None) -> AuthInfo | None: """Pull verified ``AuthInfo`` from the ``ToolContext.metadata``. @@ -726,11 +983,23 @@ def _invoke_account_resolver( return resolver(ref) +def _canonical_validation_error(params: dict[str, Any]) -> dict[str, Any] | None: + """Return a controller error when the canonical request schema rejects input.""" + from adcp.validation.schema_validator import format_issues, validate_request + + validation = validate_request("comply_test_controller", params) + if validation.valid: + return None + return _controller_error("INVALID_PARAMS", format_issues(validation.issues)) + + async def _handle_test_controller( store: TestControllerStore, params: dict[str, Any], context: ToolContext | None = None, account_resolver: _AccountResolver | None = None, + *, + validate_schema: bool = False, ) -> dict[str, Any]: """Dispatch a comply_test_controller request to the store. @@ -747,11 +1016,21 @@ async def _handle_test_controller( observation guard). Phase 1 of the lifecycle-state-and-sandbox- authority proposal — see ``docs/proposals/lifecycle-state-and-sandbox-authority.md``. + + Public MCP and A2A registrations set ``validate_schema=True`` so the + bundled canonical conditional schema is enforced before dispatch. The + default remains false for this private helper's legacy in-process test + callers, which exercise individual dispatch and sandbox-gate branches + with intentionally partial envelopes. """ scenario = params.get("scenario") implemented = _list_scenarios(store) if scenario == "list_scenarios": + if validate_schema: + validation_error = _canonical_validation_error(params) + if validation_error is not None: + return validation_error # Capability probe — exempt from the sandbox gate. Returning the # implemented-scenarios list is a discovery surface every buyer # needs to call regardless of mode. @@ -773,6 +1052,11 @@ async def _handle_test_controller( if gate_response is not None: return gate_response + if validate_schema: + validation_error = _canonical_validation_error(params) + if validation_error is not None: + return validation_error + if scenario not in SCENARIOS: return _controller_error( "UNKNOWN_SCENARIO", @@ -787,6 +1071,8 @@ async def _handle_test_controller( method = getattr(store, scenario) scenario_params = params.get("params", {}) + if not isinstance(scenario_params, dict): + return _controller_error("INVALID_PARAMS", "params must be an object") extra: dict[str, Any] = {} if context is not None and _accepts_context_kwarg(method): @@ -796,33 +1082,28 @@ async def _handle_test_controller( extra["account"] = account try: - if scenario == "force_creative_status": - result = await method( - creative_id=scenario_params["creative_id"], - status=scenario_params["status"], - rejection_reason=scenario_params.get("rejection_reason"), - **extra, - ) + method_kwargs = _accepted_scenario_kwargs(method, scenario_params) + for name in _LEGACY_OPTIONAL_SCENARIO_PARAMS.get(str(scenario), ()): + if name not in method_kwargs and _accepts_kwarg(method, name): + method_kwargs[name] = None + if scenario == "expire_account_change_cursor": + if not isinstance(account, dict) or not isinstance(account.get("account_id"), str): + return _controller_error( + "INVALID_PARAMS", + "account.account_id is required for expire_account_change_cursor", + ) + if _accepts_kwarg(method, "account_id"): + method_kwargs["account_id"] = account["account_id"] + elif scenario == "force_creative_status": + _require_scenario_params(scenario_params, "creative_id", "status") + elif scenario == "force_creative_purge": + _require_scenario_params(scenario_params, "creative_id") elif scenario == "force_account_status": - result = await method( - account_id=scenario_params["account_id"], - status=scenario_params["status"], - **extra, - ) + _require_scenario_params(scenario_params, "account_id", "status") elif scenario == "force_media_buy_status": - result = await method( - media_buy_id=scenario_params["media_buy_id"], - status=scenario_params["status"], - rejection_reason=scenario_params.get("rejection_reason"), - **extra, - ) + _require_scenario_params(scenario_params, "media_buy_id", "status") elif scenario == "force_session_status": - result = await method( - session_id=scenario_params["session_id"], - status=scenario_params["status"], - termination_reason=scenario_params.get("termination_reason"), - **extra, - ) + _require_scenario_params(scenario_params, "session_id", "status") elif scenario == "force_create_media_buy_arm": arm = scenario_params.get("arm") or "" if arm not in ("submitted", "input-required"): @@ -857,12 +1138,95 @@ async def _handle_test_controller( "INVALID_PARAMS", f"message must be a string of at most {_MAX_MESSAGE} characters", ) - result = await method( - arm=arm, - task_id=task_id, - message=message, - **extra, - ) + for name, value in (("arm", arm), ("task_id", task_id), ("message", message)): + if _accepts_kwarg(method, name): + method_kwargs[name] = value + elif scenario == "force_get_products_arm": + arm = scenario_params.get("arm") + if arm not in ("submitted", "rejected"): + return _controller_error( + "INVALID_PARAMS", + "arm must be 'submitted' or 'rejected'", + ) + message = scenario_params.get("message") + if message is not None and ( + not isinstance(message, str) or len(message) > _MAX_MESSAGE + ): + return _controller_error( + "INVALID_PARAMS", + f"message must be a string of at most {_MAX_MESSAGE} characters", + ) + if arm == "submitted": + raw_task_id = scenario_params.get("task_id") + task_id = raw_task_id.strip() if isinstance(raw_task_id, str) else None + if not task_id: + return _controller_error( + "INVALID_PARAMS", + "task_id is required when arm is 'submitted'", + ) + if len(task_id) > _MAX_TASK_ID: + return _controller_error( + "INVALID_PARAMS", + f"task_id must be at most {_MAX_TASK_ID} characters", + ) + if "reason" in scenario_params or "suggestions" in scenario_params: + return _controller_error( + "INVALID_PARAMS", + "reason and suggestions are not allowed when arm is 'submitted'", + ) + if _accepts_kwarg(method, "task_id"): + method_kwargs["task_id"] = task_id + else: + reason = scenario_params.get("reason") + if not isinstance(reason, str) or not reason or len(reason) > _MAX_MESSAGE: + return _controller_error( + "INVALID_PARAMS", + f"reason must be a string of 1 to {_MAX_MESSAGE} characters", + ) + if "task_id" in scenario_params or "message" in scenario_params: + return _controller_error( + "INVALID_PARAMS", + "task_id and message are not allowed when arm is 'rejected'", + ) + suggestions = scenario_params.get("suggestions") + if suggestions is not None and ( + not isinstance(suggestions, list) + or not 1 <= len(suggestions) <= 20 + or any( + not isinstance(item, str) or not item or len(item) > 1000 + for item in suggestions + ) + ): + return _controller_error( + "INVALID_PARAMS", + "suggestions must contain 1 to 20 non-empty strings " + "of at most 1000 characters", + ) + elif scenario == "force_get_signals_arm": + if scenario_params.get("arm") != "submitted": + return _controller_error("INVALID_PARAMS", "arm must be 'submitted'") + message = scenario_params.get("message") + if message is not None and ( + not isinstance(message, str) or len(message) > _MAX_MESSAGE + ): + return _controller_error( + "INVALID_PARAMS", + f"message must be a string of at most {_MAX_MESSAGE} characters", + ) + raw_task_id = scenario_params.get("task_id") + task_id = raw_task_id.strip() if isinstance(raw_task_id, str) else None + if not task_id: + return _controller_error( + "INVALID_PARAMS", + "task_id is required when arm is 'submitted'", + ) + if len(task_id) > _MAX_TASK_ID: + return _controller_error( + "INVALID_PARAMS", + f"task_id must be at most {_MAX_TASK_ID} characters", + ) + if _accepts_kwarg(method, "task_id"): + method_kwargs["task_id"] = task_id elif scenario == "force_task_completion": raw_task_id = scenario_params.get("task_id") task_id = raw_task_id.strip() if isinstance(raw_task_id, str) else None @@ -888,66 +1252,86 @@ async def _handle_test_controller( "INVALID_PARAMS", f"result payload exceeds {_MAX_RESULT_BYTES // 1024} KB limit", ) - result = await method( - task_id=task_id, - result=result_value, - **extra, - ) + if _accepts_kwarg(method, "task_id"): + method_kwargs["task_id"] = task_id + if _accepts_kwarg(method, "result"): + method_kwargs["result"] = result_value elif scenario == "simulate_delivery": - result = await method( - media_buy_id=scenario_params["media_buy_id"], - impressions=scenario_params.get("impressions"), - clicks=scenario_params.get("clicks"), - conversions=scenario_params.get("conversions"), - reported_spend=scenario_params.get("reported_spend"), - **extra, - ) + _require_scenario_params(scenario_params, "media_buy_id") elif scenario == "simulate_budget_spend": - result = await method( - spend_percentage=scenario_params["spend_percentage"], - account_id=scenario_params.get("account_id"), - media_buy_id=scenario_params.get("media_buy_id"), - **extra, - ) - elif scenario == "seed_product": - result = await method( - fixture=scenario_params.get("fixture"), - product_id=scenario_params.get("product_id"), - **extra, - ) - elif scenario == "seed_pricing_option": - result = await method( - fixture=scenario_params.get("fixture"), - product_id=scenario_params.get("product_id"), - pricing_option_id=scenario_params.get("pricing_option_id"), - **extra, - ) - elif scenario == "seed_creative": - result = await method( - fixture=scenario_params.get("fixture"), - creative_id=scenario_params.get("creative_id"), - **extra, - ) - elif scenario == "seed_plan": - result = await method( - fixture=scenario_params.get("fixture"), - plan_id=scenario_params.get("plan_id"), - **extra, - ) - elif scenario == "seed_media_buy": - result = await method( - fixture=scenario_params.get("fixture"), - media_buy_id=scenario_params.get("media_buy_id"), - **extra, - ) - elif scenario == "seed_creative_format": - result = await method( - fixture=scenario_params.get("fixture"), - format_id=scenario_params.get("format_id"), - **extra, - ) - else: + _require_scenario_params(scenario_params, "spend_percentage") + elif scenario == "seed_account": + _require_scenario_params(scenario_params, "account_id") + elif scenario == "seed_rights_grant": + _require_scenario_params(scenario_params, "rights_id") + elif scenario == "seed_measurement_catalog": + _require_scenario_params(scenario_params, "vendor", "metrics") + elif scenario == "query_provenance_audit_observations": + _require_scenario_params(scenario_params, "creative_id") + elif scenario == "force_upstream_unavailable": + _require_scenario_params(scenario_params, "tool") + elif scenario == "catalog_item_availability_probe": + _require_scenario_params(scenario_params, "operation", "catalog_id", "item_id") + operation = scenario_params["operation"] + if operation not in { + "seed_inaccessible_item", + "query_eligibility", + "advance_time", + "recreate_catalog", + }: + return _controller_error( + "INVALID_PARAMS", + "Unsupported catalog_item_availability_probe operation", + ) + if operation in {"query_eligibility", "advance_time", "recreate_catalog"}: + _require_scenario_params(scenario_params, "catalog_generation") + if operation == "advance_time": + _require_scenario_params(scenario_params, "target_time") + elif scenario == "compact_product_lifecycle_probe": + _require_scenario_params(scenario_params, "operation") + operation = scenario_params["operation"] + if operation == "prepare": + _require_scenario_params(scenario_params, "product_id") + if "proposal_id" in scenario_params: + return _controller_error( + "INVALID_PARAMS", + "proposal_id is not allowed for the prepare operation", + ) + elif operation == "expire_proposal": + _require_scenario_params(scenario_params, "proposal_id") + forbidden = {"expires_at", "target_time", "product_id"} & scenario_params.keys() + if forbidden: + return _controller_error( + "INVALID_PARAMS", + f"Fields not allowed for expire_proposal: {', '.join(sorted(forbidden))}", + ) + else: + return _controller_error( + "INVALID_PARAMS", + "operation must be 'prepare' or 'expire_proposal'", + ) + elif scenario == "compact_direct_buy_lifecycle_probe": + _require_scenario_params(scenario_params, "operation", "product_id") + if scenario_params["operation"] != "prepare": + return _controller_error("INVALID_PARAMS", "operation must be 'prepare'") + forbidden = {"proposal_id", "expires_at", "target_time"} & scenario_params.keys() + if forbidden: + return _controller_error( + "INVALID_PARAMS", + f"Fields not allowed for prepare: {', '.join(sorted(forbidden))}", + ) + elif scenario not in { + "seed_product", + "seed_pricing_option", + "seed_creative", + "seed_plan", + "seed_media_buy", + "seed_creative_format", + "query_upstream_traffic", + }: return _controller_error("UNKNOWN_SCENARIO", f"Unknown scenario: {scenario}") + + result = await method(**method_kwargs, **extra) except TestControllerError as e: return _controller_error(e.code, str(e), current_state=e.current_state) except KeyError as e: @@ -1042,6 +1426,11 @@ async def force_account_status(self, account_id, status): from adcp.server.base import ToolContext as _ToolContext from adcp.server.serve import RequestMetadata as _RequestMetadata + from adcp.validation.schema_loader import get_mcp_schema + + controller_schema = get_mcp_schema("comply_test_controller", "request") + if controller_schema is None: + raise RuntimeError("bundled comply_test_controller request schema is unavailable") async def comply_test_controller(**kwargs: Any) -> dict[str, Any]: context: _ToolContext | None = None @@ -1058,6 +1447,7 @@ async def comply_test_controller(**kwargs: Any) -> dict[str, Any]: kwargs, context=context, account_resolver=account_resolver, + validate_schema=True, ) tool = Tool.from_function( @@ -1066,22 +1456,10 @@ async def comply_test_controller(**kwargs: Any) -> dict[str, Any]: description="Compliance test controller. Sandbox only, not for production use.", ) - # Override schema with the proper comply_test_controller inputSchema. - # Derived from SCENARIOS so it can't drift from the dispatcher. - tool.parameters = { - "type": "object", - "properties": { - "account": {"type": "object"}, - "scenario": { - "type": "string", - # Derived from SCENARIOS so the enum never drifts from the dispatcher. - "enum": ["list_scenarios"] + SCENARIOS, - }, - "params": {"type": "object"}, - "context": {"type": "object"}, - }, - "required": ["scenario"], - } + # Advertise the same canonical, self-contained schema enforced above. + # Scenario-specific conditionals therefore evolve with the bundled AdCP + # schema instead of a second hand-maintained dispatcher contract. + tool.parameters = controller_schema # Override fn_metadata with a permissive model class _ControllerArgs(ArgModelBase): diff --git a/src/adcp/types/__init__.py b/src/adcp/types/__init__.py index fa809f39a..b4b7e16e0 100644 --- a/src/adcp/types/__init__.py +++ b/src/adcp/types/__init__.py @@ -395,6 +395,11 @@ "AssetContentType", "AssetType", # Deprecated "AdvertiserIndustry", + "AgentDeclarations", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "AudienceSource", "BrandIdentity", "BrandReference", @@ -633,6 +638,20 @@ "PaginationResponse", "PricingModel", "PrimaryCountry", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", "PropertyIdentifierTypes", "PropertyType", "PublisherDesignatedPreviewProvider", @@ -1197,9 +1216,14 @@ def __dir__() -> list[str]: AdcpProtocol, AdvertiserIndustry, AgentConfig, + AgentDeclarations, AgentDeployment, AgentDestination, + AgentNotificationConfig, + AgentNotificationConfigState, AgentPermissionDeniedDetails, + AgentReportingDestination, + AgentReportingDestinationState, AggregatedTotals, AiTool, Artifact, @@ -1702,6 +1726,20 @@ def __dir__() -> list[str]: PricingModel, PricingOption, PrimaryCountry, + PrincipalAppliedResult, + PrincipalChangedWebhook, + PrincipalConfiguration, + PrincipalCurrentResult, + PrincipalDeclarationsState, + PrincipalKind, + PrincipalReadFailedResult, + PrincipalReadResult, + PrincipalRecognizedResult, + PrincipalState, + PrincipalSyncFailedResult, + PrincipalSyncResult, + PrincipalUnconfiguredResult, + PrincipalValidatedResult, Product, ProductAllocation, ProductAllowedAction, diff --git a/src/adcp/types/_eager.py b/src/adcp/types/_eager.py index 5b6a6a291..51c0fe069 100644 --- a/src/adcp/types/_eager.py +++ b/src/adcp/types/_eager.py @@ -71,7 +71,12 @@ ActivateSignalResponse, AdcpProtocol, AdvertiserIndustry, + AgentDeclarations, + AgentNotificationConfig, + AgentNotificationConfigState, AgentPermissionDeniedDetails, + AgentReportingDestination, + AgentReportingDestinationState, AggregatedTotals, AiTool, Artifact, @@ -303,6 +308,10 @@ PricingCurrency, PricingModel, PrimaryCountry, + PrincipalChangedWebhook, + PrincipalDeclarationsState, + PrincipalKind, + PrincipalState, ProductAllowedAction, ProductCard, ProductCardDetailed, @@ -710,6 +719,16 @@ PostalArea, PreviewRenderingOrigin, PricingOption, + PrincipalAppliedResult, + PrincipalConfiguration, + PrincipalCurrentResult, + PrincipalReadFailedResult, + PrincipalReadResult, + PrincipalRecognizedResult, + PrincipalSyncFailedResult, + PrincipalSyncResult, + PrincipalUnconfiguredResult, + PrincipalValidatedResult, ProductAllocation, ProductDiscoveryProductId, ProductFilterCountry, @@ -1123,6 +1142,11 @@ def __init__(self, *args: object, **kwargs: object) -> None: "ActionNotAllowedDetails", "ActionNotAllowedReason", "AgentPermissionDeniedDetails", + "AgentDeclarations", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "AudienceTooSmallDetails", "BillingNotPermittedForAgentDetails", "BillingNotSupportedDetails", @@ -1632,6 +1656,20 @@ def __init__(self, *args: object, **kwargs: object) -> None: "PricingModel", "PricingOption", "PrimaryCountry", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", "Product", "ProductCard", "ProductCardDetailed", diff --git a/src/adcp/types/aliases.py b/src/adcp/types/aliases.py index 4c27f1300..cad1d06e1 100644 --- a/src/adcp/types/aliases.py +++ b/src/adcp/types/aliases.py @@ -144,6 +144,30 @@ from adcp.types.generated_poc.core.product import TrustedMatch from adcp.types.generated_poc.core.product_allocation import ProductAllocation from adcp.types.generated_poc.core.signal_coverage_forecast import SignalCoverageForecast +from adcp.types.generated_poc.protocol.get_principal_response import ( + Result as PrincipalUnconfiguredResult, +) +from adcp.types.generated_poc.protocol.get_principal_response import ( + Result6 as PrincipalCurrentResult, +) +from adcp.types.generated_poc.protocol.get_principal_response import ( + Result7 as PrincipalRecognizedResult, +) +from adcp.types.generated_poc.protocol.get_principal_response import ( + Result9 as PrincipalReadFailedResult, +) +from adcp.types.generated_poc.protocol.sync_principal_request import ( + Configuration as PrincipalConfiguration, +) +from adcp.types.generated_poc.protocol.sync_principal_response import ( + Result as PrincipalValidatedResult, +) +from adcp.types.generated_poc.protocol.sync_principal_response import ( + Result17 as PrincipalAppliedResult, +) +from adcp.types.generated_poc.protocol.sync_principal_response import ( + Result19 as PrincipalSyncFailedResult, +) from adcp.types.generated_poc.core.vendor_pricing_option import ( VendorPricingOption as VendorPricingOptionUnion, ) @@ -916,6 +940,18 @@ def process_result(result: SyncCatalogResult) -> None: ComplyErrorResponse: TypeAlias = ComplyTestControllerResponse4 """Error - operation failed.""" +# Principal response variants. The generated names are traversal-order +# suffixes, so expose stable semantic names for adopter control flow. +PrincipalReadResult: TypeAlias = ( + PrincipalCurrentResult + | PrincipalRecognizedResult + | PrincipalUnconfiguredResult + | PrincipalReadFailedResult +) +PrincipalSyncResult: TypeAlias = ( + PrincipalAppliedResult | PrincipalValidatedResult | PrincipalSyncFailedResult +) + # ============================================================================ # REQUEST TYPE ALIASES - Operation Variants # ============================================================================ @@ -2586,6 +2622,17 @@ class UnknownGroupAsset(_BaseGroupAsset): "ComplyStateTransitionResponse", "ComplySimulationResponse", "ComplyErrorResponse", + # Principal request and response aliases + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalRecognizedResult", + "PrincipalUnconfiguredResult", + "PrincipalReadFailedResult", + "PrincipalAppliedResult", + "PrincipalValidatedResult", + "PrincipalSyncFailedResult", + "PrincipalReadResult", + "PrincipalSyncResult", # Creative format asset slot aliases (item_type='individual') # Forward-compat fallback arms (novel asset_type values parse as these) "UnknownFormatAsset", diff --git a/src/adcp/types/protocol.py b/src/adcp/types/protocol.py index 5ab51caf6..130a1afaa 100644 --- a/src/adcp/types/protocol.py +++ b/src/adcp/types/protocol.py @@ -51,6 +51,29 @@ "GeneratedTaskStatus", "GetAdcpCapabilitiesRequest", "GetAdcpCapabilitiesResponse", + "GetPrincipalRequest", + "GetPrincipalResponse", + "SyncPrincipalRequest", + "SyncPrincipalResponse", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalRecognizedResult", + "PrincipalUnconfiguredResult", + "PrincipalReadFailedResult", + "PrincipalAppliedResult", + "PrincipalValidatedResult", + "PrincipalSyncFailedResult", + "PrincipalReadResult", + "PrincipalSyncResult", + "PrincipalState", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalChangedWebhook", + "AgentDeclarations", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "WebhookChallenge", "WebhookChallengeResponse", "WebhookResponseType", @@ -81,6 +104,11 @@ # lazily through ``__getattr__`` at runtime. from adcp.types import ( # noqa: F401 AdcpProtocol, + AgentDeclarations, + AgentNotificationConfig, + AgentNotificationConfigState, + AgentReportingDestination, + AgentReportingDestinationState, Authentication, AuthenticationScheme, AuthorizationRequiredDetails, @@ -91,6 +119,8 @@ GeneratedTaskStatus, GetAdcpCapabilitiesRequest, GetAdcpCapabilitiesResponse, + GetPrincipalRequest, + GetPrincipalResponse, GetTaskStatusRequest, GetTaskStatusResponse, ListTasksRequest, @@ -102,6 +132,20 @@ Pagination, PaginationRequest, PaginationResponse, + PrincipalAppliedResult, + PrincipalChangedWebhook, + PrincipalConfiguration, + PrincipalCurrentResult, + PrincipalDeclarationsState, + PrincipalKind, + PrincipalReadFailedResult, + PrincipalReadResult, + PrincipalRecognizedResult, + PrincipalState, + PrincipalSyncFailedResult, + PrincipalSyncResult, + PrincipalUnconfiguredResult, + PrincipalValidatedResult, Protocol, ProtocolEnvelope, ProtocolResponse, @@ -115,6 +159,8 @@ SortApplied, SortDirection, StatusSummary, + SyncPrincipalRequest, + SyncPrincipalResponse, TaskResult, TaskType, WebhookChallenge, diff --git a/tests/conformance/a2a/test_comply_test_controller_artifacts.py b/tests/conformance/a2a/test_comply_test_controller_artifacts.py index 8ea840fe9..22064d676 100644 --- a/tests/conformance/a2a/test_comply_test_controller_artifacts.py +++ b/tests/conformance/a2a/test_comply_test_controller_artifacts.py @@ -51,6 +51,7 @@ async def force_account_status(self, account_id: str, status: str) -> dict[str, async def _send(client: httpx.AsyncClient, scenario_payload: dict[str, Any]) -> dict[str, Any]: """POST a JSON-RPC ``message/send`` for comply_test_controller.""" + scenario_payload = {"account": {"sandbox": True}, **scenario_payload} body = { "jsonrpc": "2.0", "id": "1", diff --git a/tests/fixtures/public_api_snapshot.json b/tests/fixtures/public_api_snapshot.json index d243c5665..54acc8d4a 100644 --- a/tests/fixtures/public_api_snapshot.json +++ b/tests/fixtures/public_api_snapshot.json @@ -72,9 +72,14 @@ "AgentCapabilities", "AgentCompliance", "AgentConfig", + "AgentDeclarations", "AgentDeployment", "AgentDestination", "AgentHealth", + "AgentNotificationConfig", + "AgentNotificationConfigState", + "AgentReportingDestination", + "AgentReportingDestinationState", "AgentStats", "ArtifactWebhookPayload", "AssetContentType", @@ -141,6 +146,7 @@ "ContextObject", "ControlMediaBuyRequest", "ControlMediaBuyResponse", + "ControlTotalCalculator", "CpaPricingOption", "CpcPricingOption", "CpcvPricingOption", @@ -179,6 +185,7 @@ "DeliveryStatus", "Deployment", "Destination", + "DestinationSetupSnapshot", "DevicePlatform", "DeviceType", "DirectoryDiscoveryMethod", @@ -268,6 +275,7 @@ "Gtin", "HtmlContent", "HtmlPreviewRender", + "HttpsReportingResourceReader", "IdempotencyConflictError", "IdempotencyExpiredError", "IdempotencyScopeError", @@ -348,6 +356,7 @@ "LogEventSuccessResponse", "MacroMapping", "MacroMappingEntry", + "ManifestReportingInspector", "MarkdownAsset", "McpWebhookPayload", "MediaBuy", @@ -410,6 +419,24 @@ "PricingCurrency", "PricingModel", "PricingOption", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalClient", + "PrincipalConfiguration", + "PrincipalConfigurationError", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalManager", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncOutcome", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", "Product", "ProductAllowedAction", "ProductFilters", @@ -461,6 +488,11 @@ "ReportPlanOutcomeResponse", "ReportUsageRequest", "ReportUsageResponse", + "ReportingCredentialProvider", + "ReportingInspectionCode", + "ReportingInspectionError", + "ReportingResourceReader", + "ReportingTrustedOriginPolicy", "RequestProposalsRequest", "RequestProposalsResponse", "ResolvedBrand", @@ -664,6 +696,7 @@ "sign_legacy_webhook", "sign_webhook", "start_oauth_authorization", + "sync_principal_configuration", "test_agent", "test_agent_a2a", "test_agent_a2a_no_auth", @@ -728,9 +761,14 @@ "AdcpProtocol", "AdvertiserIndustry", "AgentConfig", + "AgentDeclarations", "AgentDeployment", "AgentDestination", + "AgentNotificationConfig", + "AgentNotificationConfigState", "AgentPermissionDeniedDetails", + "AgentReportingDestination", + "AgentReportingDestinationState", "AggregatedTotals", "AiTool", "Artifact", @@ -1233,6 +1271,20 @@ "PricingModel", "PricingOption", "PrimaryCountry", + "PrincipalAppliedResult", + "PrincipalChangedWebhook", + "PrincipalConfiguration", + "PrincipalCurrentResult", + "PrincipalDeclarationsState", + "PrincipalKind", + "PrincipalReadFailedResult", + "PrincipalReadResult", + "PrincipalRecognizedResult", + "PrincipalState", + "PrincipalSyncFailedResult", + "PrincipalSyncResult", + "PrincipalUnconfiguredResult", + "PrincipalValidatedResult", "Product", "ProductAllocation", "ProductAllowedAction", diff --git a/tests/test_a2a_server.py b/tests/test_a2a_server.py index 5a8ed1b54..a6add5e89 100644 --- a/tests/test_a2a_server.py +++ b/tests/test_a2a_server.py @@ -737,7 +737,7 @@ async def test_execute_test_controller_list_scenarios(): request=MessageSendParams( message=_make_datapart_msg( "comply_test_controller", - {"scenario": "list_scenarios"}, + {"scenario": "list_scenarios", "account": {"sandbox": True}}, ) ) ) @@ -767,6 +767,7 @@ async def test_execute_test_controller_force_account_status(): message=_make_datapart_msg( "comply_test_controller", { + "account": {"sandbox": True}, "scenario": "force_account_status", "params": {"account_id": "acct-1", "status": "suspended"}, }, @@ -800,6 +801,7 @@ async def test_execute_test_controller_error(): message=_make_datapart_msg( "comply_test_controller", { + "account": {"sandbox": True}, "scenario": "force_account_status", "params": {"account_id": "nonexistent", "status": "active"}, }, @@ -826,6 +828,21 @@ async def test_execute_test_controller_error(): assert result["error"] == "NOT_FOUND" +async def test_execute_test_controller_rejects_noncanonical_params(): + """A2A applies the same canonical request validation as MCP.""" + executor = ADCPAgentExecutor(_TestHandler(), test_controller=_TestStore()) + result = await executor._tool_callers["comply_test_controller"]( + { + "account": {"sandbox": True}, + "scenario": "simulate_budget_spend", + "params": {"spend_percentage": 50}, + } + ) + + assert result["success"] is False + assert result["error"] == "INVALID_PARAMS" + + @pytest.mark.skipif( sys.version_info < (3, 11), reason="a2a-sdk starlette integration requires Python 3.11+", diff --git a/tests/test_principal.py b/tests/test_principal.py new file mode 100644 index 000000000..6d08ab6d1 --- /dev/null +++ b/tests/test_principal.py @@ -0,0 +1,234 @@ +from __future__ import annotations + +from collections.abc import Iterable + +import pytest + +from adcp.principal import PrincipalConfigurationError, PrincipalManager +from adcp.types import ( + GetPrincipalResponse, + PrincipalConfiguration, + PrincipalCurrentResult, + PrincipalState, + SyncPrincipalRequest, + SyncPrincipalResponse, +) +from adcp.types.core import TaskResult, TaskStatus + +DESTINATION = { + "pattern": "file_transfer", + "destination_id": "archive", + "active": True, + "provider": {"domain": "object-store.example"}, + "transport": "s3", + "location": "s3://buyer-reporting/adcp/", + "accepted_formats": ["parquet"], + "accepted_verification_profiles": ["manifest_checksums"], +} + + +def _state(destination_state: str = "ready", destination_ref: str = "dest-1") -> dict[str, object]: + return { + "reporting_destinations": [ + { + "destination_id": "archive", + "destination_ref": destination_ref, + "state": destination_state, + "configuration": DESTINATION, + **( + {"setup": {"action": "grant_access", "setup_url": "https://setup.example/"}} + if destination_state == "action_required" + else {} + ), + } + ], + "declarations": { + "declared": {"async_adcp_versions": ["3.2", "3.1"]}, + "accepted": {"async_adcp_versions": ["3.2"]}, + "selected_async_adcp_version": "3.2", + "exclusions": [ + { + "axis": "async_adcp_versions", + "value": "3.1", + "reason": "unsupported by this seller", + } + ], + }, + } + + +def _current( + destination_state: str = "ready", destination_ref: str = "dest-1" +) -> GetPrincipalResponse: + return GetPrincipalResponse.model_validate( + { + "result": { + "kind": "current", + "principal_id": "principal-1", + "principal_kind": "buyer_agent", + "configuration_version": "version-7", + "configuration": _state(destination_state, destination_ref), + } + } + ) + + +def _applied(destination_state: str = "ready") -> SyncPrincipalResponse: + return SyncPrincipalResponse.model_validate( + { + "result": { + "kind": "applied", + "action": "updated", + "dry_run": False, + "principal_id": "principal-1", + "principal_kind": "buyer_agent", + "configuration_version": "version-8", + "configuration": _state(destination_state), + } + } + ) + + +def _completed(data): + return TaskResult(status=TaskStatus.COMPLETED, data=data) + + +class _Client: + def __init__( + self, + reads: Iterable[GetPrincipalResponse], + sync_response: SyncPrincipalResponse | None = None, + ) -> None: + self.reads = iter(reads) + self.sync_response = sync_response or _applied() + self.sync_requests: list[SyncPrincipalRequest] = [] + + async def get_principal(self, request): + return _completed(next(self.reads)) + + async def sync_principal(self, request: SyncPrincipalRequest): + self.sync_requests.append(request) + return _completed(self.sync_response) + + +@pytest.mark.asyncio +async def test_sync_bootstraps_version_and_projects_negotiation() -> None: + client = _Client([_current()]) + + outcome = await PrincipalManager(client).sync( + { + "reporting_destinations": [DESTINATION], + "declarations": {"async_adcp_versions": ["3.2", "3.1"]}, + }, + idempotency_key="principal-sync-key-0001", + ) + + request = client.sync_requests[0] + assert request.expected_configuration_version == "version-7" + assert request.expected_principal_kind == "buyer_agent" + assert outcome.configuration_version == "version-8" + assert outcome.destinations.ready == ("archive",) + assert outcome.selected_async_adcp_version == "3.2" + assert outcome.declarations is not None + assert outcome.declarations.exclusions[0].value == "3.1" + + +@pytest.mark.asyncio +async def test_sync_can_disable_version_fence() -> None: + client = _Client([]) + + await PrincipalManager(client).sync( + {"declarations": {}}, + idempotency_key="principal-sync-key-0002", + use_version_fence=False, + ) + + request = client.sync_requests[0] + assert request.expected_configuration_version is None + assert request.expected_principal_kind is None + + +@pytest.mark.asyncio +async def test_sync_polls_validating_destination_to_action_required() -> None: + client = _Client( + [_current(), _current("validating"), _current("action_required")], + _applied("validating"), + ) + + outcome = await PrincipalManager(client).sync( + {"reporting_destinations": [DESTINATION]}, + idempotency_key="principal-sync-key-0003", + wait_for_setup=True, + poll_interval=0.001, + ) + + assert outcome.destinations.settled + assert outcome.destinations.action_required == ("archive",) + state = outcome.destinations.states["archive"] + assert str(state.setup.setup_url) == "https://setup.example/" + + +@pytest.mark.asyncio +async def test_wait_for_destinations_rejects_missing_readback() -> None: + response = GetPrincipalResponse.model_validate( + { + "result": { + "kind": "current", + "principal_id": "principal-1", + "principal_kind": "buyer_agent", + "configuration_version": "version-7", + "configuration": {}, + } + } + ) + + with pytest.raises(PrincipalConfigurationError) as error: + await PrincipalManager(_Client([response])).wait_for_destinations( + {"archive"}, poll_interval=0.001 + ) + + assert error.value.code == "DESTINATION_STATE_MISSING" + + +@pytest.mark.asyncio +async def test_wait_for_destinations_rejects_a_replaced_generation() -> None: + client = _Client( + [_current(), _current("validating"), _current("ready", "dest-replaced")], + _applied("validating"), + ) + + with pytest.raises(PrincipalConfigurationError) as error: + await PrincipalManager(client).sync( + {"reporting_destinations": [DESTINATION]}, + idempotency_key="principal-sync-key-0005", + wait_for_setup=True, + poll_interval=0.001, + ) + + assert error.value.code == "DESTINATION_GENERATION_CHANGED" + + +@pytest.mark.asyncio +async def test_payload_failure_is_not_treated_as_success() -> None: + failed = SyncPrincipalResponse.model_validate( + { + "result": { + "kind": "failed", + "errors": [{"code": "CONFLICT", "message": "stale configuration"}], + } + } + ) + + with pytest.raises(PrincipalConfigurationError, match="stale configuration"): + await PrincipalManager(_Client([_current()], failed)).sync( + {"declarations": {}}, idempotency_key="principal-sync-key-0004" + ) + + +def test_principal_semantic_types_are_public() -> None: + from adcp import PrincipalCurrentResult as RootPrincipalCurrentResult + from adcp.types.protocol import PrincipalState as PartialPrincipalState + + assert RootPrincipalCurrentResult is PrincipalCurrentResult + assert PartialPrincipalState is PrincipalState + assert PrincipalConfiguration.model_fields["declarations"] diff --git a/tests/test_reporting_reconciliation.py b/tests/test_reporting_reconciliation.py index b80bce4ec..894a75327 100644 --- a/tests/test_reporting_reconciliation.py +++ b/tests/test_reporting_reconciliation.py @@ -1,10 +1,16 @@ from __future__ import annotations import asyncio +import hashlib +import json +import threading from copy import deepcopy from datetime import datetime +from urllib.parse import urljoin +import httpx import pytest +import rfc8785 from adcp.decisioning.capabilities import MediaBuy from adcp.reporting import ( @@ -12,12 +18,28 @@ ReportingInspectionContext, ReportingObservation, ReportingReconciliationError, + ReportingTier, build_reporting_receipt, evaluate_reporting_ledger, load_reporting_ledger, reconcile_reporting, + reconcile_reporting_core, + reporting_tiers, +) +from adcp.reporting_inspection import ( + HttpsReportingResourceReader, + ManifestReportingInspector, + ReportingInspectionCode, + ReportingInspectionError, + _canonical_rows_bytes, + _decode_rows, +) +from adcp.types import ( + ReportingDeliveryCapabilities, + ReportingMaterialization, + ReportingObligation, + ReportingRevision, ) -from adcp.types import ReportingDeliveryCapabilities from adcp.types.core import TaskResult, TaskStatus from adcp.types.generated_poc.core.reporting_canonical_content_digest import ( ReportingCanonicalContentDigest, @@ -296,6 +318,551 @@ async def sync_reporting_receipts( ) +class _CoreClient: + async def get_reporting_status( + self, request: GetReportingStatusRequest + ) -> TaskResult[GetReportingStatusResponse]: + obligation = _obligation("obligation-core") + for field in ( + "destination_ref", + "materialization_count", + "successful_materialization_count", + "receipt_count", + "accepted_receipt_count", + "resource_retained_until", + ): + obligation.pop(field, None) + obligation.update( + reconciliation_mode="delivery_only", + reconciliation_status="not_required", + health="complete", + ) + revision = deepcopy(REVISION) + revision.pop("canonical_content_digest") + response = _response() + response.update( + periods=[obligation], + revisions=[revision], + materializations=[], + receipts=[], + pagination={"has_more": False, "total_count": 2}, + ) + return TaskResult( + status=TaskStatus.COMPLETED, + data=GetReportingStatusResponse.model_validate(response), + ) + + +@pytest.mark.asyncio +async def test_core_reconciliation_needs_no_inspector_or_receipt_client() -> None: + result = await reconcile_reporting_core( + _CoreClient(), + GetReportingStatusRequest.model_validate( + { + "account": {"account_id": "account-1"}, + "view": "periods", + "period": {"start": PERIOD["start"], "end": PERIOD["end"]}, + } + ), + expected_periods=[ + ExpectedReportingPeriod( + "billing-feed", + 1, + "billing-v1", + "billing", + "billing-v1", + ("buy-1", "buy-2"), + PERIOD["start"], + PERIOD["end"], + ) + ], + now=datetime.fromisoformat("2026-09-03T00:00:00+00:00"), + ) + + assert result.definitive + assert result.obligations[0].definitive + + +def test_reporting_tiers_enforce_cumulative_flags() -> None: + core = ReportingDeliveryCapabilities.model_construct( + managed_delivery=False, reconciled_billing=False + ) + assert reporting_tiers(core) == frozenset({ReportingTier.CORE}) + + invalid = ReportingDeliveryCapabilities.model_construct( + managed_delivery=False, reconciled_billing=True + ) + with pytest.raises(ReportingReconciliationError) as error: + reporting_tiers(invalid) + assert error.value.code == "INVALID_REPORTING_CAPABILITIES" + + +class _ResourceReader: + def __init__(self, resources: dict[str, bytes]) -> None: + self.resources = resources + self.calls: list[str] = [] + + async def read(self, locator: str, *, base: str | None = None, max_bytes: int) -> bytes: + resolved = urljoin(base, locator) if base else locator + self.calls.append(resolved) + body = self.resources[resolved] + if len(body) > max_bytes: + raise ValueError("too large") + return body + + +def _manifest_inspection_fixture(): + rows = [ + {"media_buy_id": "buy-2", "date": "2026-08-01", "impressions": 4, "spend": 5}, + {"media_buy_id": "buy-1", "date": "2026-08-01", "impressions": 3, "spend": 7}, + ] + data = b"".join(json.dumps(row, separators=(",", ":")).encode() + b"\n" for row in rows) + schema = { + "$schema": "https://json-schema.org/draft/2020-12/schema", + "type": "object", + "properties": { + "media_buy_id": {"type": "string"}, + "date": {"type": "string"}, + "impressions": {"type": "integer"}, + "spend": {"type": "number"}, + }, + "required": ["media_buy_id", "date", "impressions", "spend"], + "additionalProperties": False, + } + schema_body = json.dumps(schema, separators=(",", ":")).encode() + schema_digest = hashlib.sha256(schema_body).hexdigest() + definition = { + "contract_version": "1.0", + "media_type": "application/vnd.adcp.reporting-definition+json", + "report_definition_id": "billing-v1", + "reporting_profile": "billing-v1", + "grain": "media_buy_day", + "source": { + "provider": {"domain": "seller.example"}, + "system": "test-ledger", + "api_version": "1", + "query_semantics": {}, + }, + "calendar": {"timezone_basis": "utc"}, + "metrics": [ + {"name": "impressions", "source_expression": "impressions", "aggregation": "sum"}, + {"name": "spend", "source_expression": "spend", "aggregation": "sum"}, + ], + "dimensions": ["media_buy_id", "date"], + "restatement_policy": { + "source_requery_duration": "P30D", + "emit_only_on_content_change": True, + }, + "finality_policies": [ + { + "finality_policy_id": "billing-v1-source-final", + "basis": "source_final", + "source_signal": "closed", + } + ], + } + definition_body = json.dumps(definition, separators=(",", ":")).encode() + canonical_contract = { + "contract_version": "1.0", + "media_type": "application/vnd.adcp.reporting-canonicalization+json", + "algorithm": "adcp_jcs_rows_v1", + "schema_sha256": schema_digest, + "primary_keys": ["media_buy_id", "date"], + "golden_vectors": { + "empty_report": { + "name": "empty", + "purpose": "empty_report", + "input_rows": [], + "canonical_utf8_base64": "W10=", + "sha256": "4f53cda18c2baa0c0354bb5f9a3ecbe5ed12ab4d8e11ba873c2f11161202b945", + }, + "ordering_encoding": { + "name": "ordering", + "purpose": "ordering_encoding", + "input_rows": [ + {"date": "2026-08-02", "media_buy_id": "buy-2"}, + {"date": "2026-08-01", "media_buy_id": "buy-1"}, + ], + "canonical_utf8_base64": ( + "W3siZGF0ZSI6IjIwMjYtMDgtMDEiLCJtZWRpYV9idXlfaWQiOiJidXktMSJ9LHsiZGF0ZSI6" + "IjIwMjYtMDgtMDIiLCJtZWRpYV9idXlfaWQiOiJidXktMiJ9XQ==" + ), + "sha256": "ce716e16aef215a0ee62500e669d9a89c6de33afa190f528b195648f1a2f71c8", + }, + }, + } + canonical_body = json.dumps(canonical_contract, separators=(",", ":")).encode() + ordered = sorted(rows, key=lambda row: (row["media_buy_id"], row["date"])) + canonical_digest = hashlib.sha256(rfc8785.dumps(ordered)).hexdigest() + totals = [ + {"name": "impressions", "value": "7", "value_type": "integer"}, + {"name": "spend", "value": "12", "value_type": "decimal", "unit": "USD"}, + ] + manifest = { + "manifest_version": "1.0", + "complete": True, + "reporting_revision_id": "revision-august-official", + "reporting_obligation_id": "obligation-billing", + "reporting_materialization_id": "materialization-billing", + "period": PERIOD, + "format": "jsonl", + "compression": "none", + "files": [ + { + "object_ref": "data.jsonl", + "size_bytes": len(data), + "sha256": hashlib.sha256(data).hexdigest(), + "row_count": 2, + } + ], + "total_size_bytes": len(data), + "row_count": 2, + "control_totals": totals, + "created_at": "2026-09-02T00:00:00Z", + } + manifest_body = json.dumps(manifest, separators=(",", ":")).encode() + revision = deepcopy(REVISION) + revision.update( + row_count=2, + control_totals=totals, + schema_uri="https://contracts.example/rows.json", + schema_sha256=schema_digest, + report_definition_uri="https://contracts.example/definition.json", + report_definition_sha256=hashlib.sha256(definition_body).hexdigest(), + canonical_content_digest={ + "algorithm": "sha256", + "value": canonical_digest, + "canonicalization_id": "rows-v1", + "canonicalization_uri": "https://contracts.example/canonical.json", + "canonicalization_sha256": hashlib.sha256(canonical_body).hexdigest(), + }, + ) + obligation = _obligation() + materialization = _materialization() + materialization.update( + method="file_transfer", + transport="https", + resource={ + "resource_ref": "resource-manifest", + "kind": "manifest", + "location": "https://files.example/manifest.json", + "manifest_version": "1.0", + "manifest_sha256": hashlib.sha256(manifest_body).hexdigest(), + "immutability": "immutable_location", + "expires_at": "2026-12-01T00:00:00Z", + }, + verification={ + "verified_at": "2026-09-02T00:00:05Z", + "verification_path": "destination", + "verification_profile": "canonical_digest", + "row_count": 2, + "control_totals": totals, + "canonical_content_digest": revision["canonical_content_digest"], + "physical_checksums": [ + { + "object_ref": "data.jsonl", + "algorithm": "sha256", + "value": hashlib.sha256(data).hexdigest(), + } + ], + }, + ) + context = ReportingInspectionContext( + ReportingObligation.model_validate(obligation), + ReportingRevision.model_validate(revision), + ReportingMaterialization.model_validate(materialization), + ) + resources = { + "https://files.example/manifest.json": manifest_body, + "https://files.example/data.jsonl": data, + "https://contracts.example/rows.json": schema_body, + "https://contracts.example/definition.json": definition_body, + "https://contracts.example/canonical.json": canonical_body, + } + return context, resources + + +@pytest.mark.asyncio +async def test_builtin_manifest_inspector_verifies_complete_resource() -> None: + context, resources = _manifest_inspection_fixture() + + observation = await ManifestReportingInspector(_ResourceReader(resources))(context) + + assert observation.row_count == 2 + assert observation.manifest_sha256 == context.materialization.resource.manifest_sha256 + assert observation.canonical_content_digest == context.revision.canonical_content_digest + + +@pytest.mark.asyncio +async def test_builtin_manifest_inspector_rejects_corrupt_object() -> None: + context, resources = _manifest_inspection_fixture() + data = resources["https://files.example/data.jsonl"] + resources["https://files.example/data.jsonl"] = b"X" + data[1:] + + with pytest.raises(ReportingInspectionError) as error: + await ManifestReportingInspector(_ResourceReader(resources))(context) + + assert error.value.code == ReportingInspectionCode.OBJECT_DIGEST_MISMATCH + + +@pytest.mark.asyncio +async def test_builtin_manifest_inspector_enforces_aggregate_decoded_byte_limit() -> None: + context, resources = _manifest_inspection_fixture() + + with pytest.raises(ReportingInspectionError) as error: + await ManifestReportingInspector(_ResourceReader(resources), max_total_decoded_bytes=1)( + context + ) + + assert error.value.code == ReportingInspectionCode.RESOURCE_TOO_LARGE + assert "decoded-byte" in str(error.value) + + +@pytest.mark.asyncio +async def test_https_reporting_reader_rejects_unsafe_locators_before_fetch() -> None: + reader = HttpsReportingResourceReader() + + with pytest.raises(ReportingInspectionError) as insecure: + await reader.read("http://files.example/report.json", max_bytes=100) + assert insecure.value.code == ReportingInspectionCode.UNSAFE_RESOURCE + + with pytest.raises(ReportingInspectionError) as cross_origin: + await reader.read( + "https://attacker.example/data.jsonl", + base="https://files.example/manifest.json", + max_bytes=100, + ) + assert cross_origin.value.code == ReportingInspectionCode.UNSAFE_RESOURCE + + +@pytest.mark.asyncio +async def test_https_reporting_reader_requires_trusted_credential_free_dns_origin() -> None: + reader = HttpsReportingResourceReader(trusted_origins=["https://files.example"]) + + for locator in ( + "https://user:password@files.example/report.json", + "https://8.8.8.8/report.json", + "https://other.example/report.json", + ): + with pytest.raises(ReportingInspectionError) as error: + await reader.read(locator, max_bytes=100) + assert error.value.code == ReportingInspectionCode.UNSAFE_RESOURCE + + +@pytest.mark.asyncio +async def test_https_reporting_reader_checks_content_type_and_resolves_off_loop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + factory_threads: list[int] = [] + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, headers={"content-type": "text/plain"}, content=b"report") + + def build_transport(*_: object, **__: object) -> httpx.AsyncBaseTransport: + factory_threads.append(threading.get_ident()) + return httpx.MockTransport(handler) + + monkeypatch.setattr( + "adcp.reporting_inspection.build_async_ip_pinned_transport", build_transport + ) + reader = HttpsReportingResourceReader(trusted_origins=["https://files.example"]) + + with pytest.raises(ReportingInspectionError) as error: + await reader.read( + "https://files.example/report.json", + max_bytes=100, + expected_content_types=frozenset({"application/json"}), + ) + + assert error.value.code == ReportingInspectionCode.UNEXPECTED_CONTENT_TYPE + assert factory_threads[0] != threading.get_ident() + + +def test_builtin_inspector_rejects_duplicate_json_keys_and_uses_jcs_key_ordering() -> None: + with pytest.raises(ReportingInspectionError) as duplicate: + _decode_rows( + b'{"media_buy_id":"buy-1","media_buy_id":"buy-2"}\n', + "jsonl", + "none", + max_decoded_bytes=1024, + max_rows=10, + ) + assert duplicate.value.code == ReportingInspectionCode.INVALID_ROWS + + assert _canonical_rows_bytes([{"id": 2}, {"id": 10}], ["id"]) == (b'[{"id":10},{"id":2}]') + + +@pytest.mark.asyncio +async def test_builtin_inspector_rejects_bad_golden_vector_and_unsafe_schema() -> None: + context, resources = _manifest_inspection_fixture() + canonical_url = "https://contracts.example/canonical.json" + canonical = json.loads(resources[canonical_url]) + canonical["golden_vectors"]["ordering_encoding"]["sha256"] = "0" * 64 + canonical_body = json.dumps(canonical, separators=(",", ":")).encode() + resources[canonical_url] = canonical_body + digest = context.revision.canonical_content_digest + assert digest is not None + updated_digest = digest.model_copy( + update={"canonicalization_sha256": hashlib.sha256(canonical_body).hexdigest()} + ) + updated_context = ReportingInspectionContext( + context.obligation, + context.revision.model_copy(update={"canonical_content_digest": updated_digest}), + context.materialization.model_copy( + update={ + "verification": context.materialization.verification.model_copy( + update={"canonical_content_digest": updated_digest} + ) + } + ), + ) + with pytest.raises(ReportingInspectionError) as bad_vector: + await ManifestReportingInspector(_ResourceReader(resources))(updated_context) + assert bad_vector.value.code == ReportingInspectionCode.INVALID_CONTRACT + + context, resources = _manifest_inspection_fixture() + schema_url = "https://contracts.example/rows.json" + schema = json.loads(resources[schema_url]) + schema["$dynamicRef"] = "#" + schema_body = json.dumps(schema, separators=(",", ":")).encode() + resources[schema_url] = schema_body + unsafe_schema_context = ReportingInspectionContext( + context.obligation, + context.revision.model_copy( + update={"schema_sha256": hashlib.sha256(schema_body).hexdigest()} + ), + context.materialization, + ) + with pytest.raises(ReportingInspectionError) as unsafe_schema: + await ManifestReportingInspector(_ResourceReader(resources))(unsafe_schema_context) + assert unsafe_schema.value.code == ReportingInspectionCode.INVALID_CONTRACT + + context, resources = _manifest_inspection_fixture() + schema_url = "https://contracts.example/rows.json" + schema = json.loads(resources[schema_url]) + schema["properties"]["media_buy_id"]["pattern"] = "^(a+)+$" + schema_body = json.dumps(schema, separators=(",", ":")).encode() + resources[schema_url] = schema_body + regex_schema_context = ReportingInspectionContext( + context.obligation, + context.revision.model_copy( + update={"schema_sha256": hashlib.sha256(schema_body).hexdigest()} + ), + context.materialization, + ) + with pytest.raises(ReportingInspectionError) as regex_schema: + await ManifestReportingInspector(_ResourceReader(resources))(regex_schema_context) + assert regex_schema.value.code == ReportingInspectionCode.INVALID_CONTRACT + + context, resources = _manifest_inspection_fixture() + schema_url = "https://contracts.example/rows.json" + schema = json.loads(resources[schema_url]) + schema["$schema"] = "https://json-schema.org/draft/2019-09/schema" + schema_body = json.dumps(schema, separators=(",", ":")).encode() + resources[schema_url] = schema_body + dialect_mismatch_context = ReportingInspectionContext( + context.obligation, + context.revision.model_copy( + update={"schema_sha256": hashlib.sha256(schema_body).hexdigest()} + ), + context.materialization, + ) + with pytest.raises(ReportingInspectionError) as dialect_mismatch: + await ManifestReportingInspector(_ResourceReader(resources))(dialect_mismatch_context) + assert dialect_mismatch.value.code == ReportingInspectionCode.INVALID_CONTRACT + + context, resources = _manifest_inspection_fixture() + schema_url = "https://contracts.example/rows.json" + schema = json.loads(resources[schema_url]) + schema["$defs"] = {"loop": {"$ref": "#/$defs/loop"}} + schema["allOf"] = [{"$ref": "#/$defs/loop"}] + schema_body = json.dumps(schema, separators=(",", ":")).encode() + resources[schema_url] = schema_body + cyclic_schema_context = ReportingInspectionContext( + context.obligation, + context.revision.model_copy( + update={"schema_sha256": hashlib.sha256(schema_body).hexdigest()} + ), + context.materialization, + ) + with pytest.raises(ReportingInspectionError) as cyclic_schema: + await ManifestReportingInspector(_ResourceReader(resources))(cyclic_schema_context) + assert cyclic_schema.value.code == ReportingInspectionCode.INVALID_CONTRACT + + +@pytest.mark.asyncio +async def test_builtin_inspector_enforces_aggregate_row_limit() -> None: + context, resources = _manifest_inspection_fixture() + + with pytest.raises(ReportingInspectionError) as limit: + await ManifestReportingInspector(_ResourceReader(resources), max_total_rows=1)(context) + assert limit.value.code == ReportingInspectionCode.RESOURCE_TOO_LARGE + + +@pytest.mark.asyncio +async def test_reconciler_requires_reconciled_tier_for_billing_and_canonical_data() -> None: + async def inspect(_: ReportingInspectionContext) -> ReportingObservation: + raise AssertionError("tier validation must run before inspection") + + capabilities = ReportingDeliveryCapabilities.model_construct( + managed_delivery=True, reconciled_billing=False + ) + with pytest.raises(ReportingReconciliationError) as error: + await reconcile_reporting( + _Client(), + GetReportingStatusRequest.model_validate( + {"account": {"account_id": "account-1"}, "view": "periods"} + ), + inspect, + expected_periods=[], + reporting_capabilities=capabilities, + ) + assert error.value.code == "RECONCILED_BILLING_NOT_ENABLED" + + +@pytest.mark.asyncio +async def test_reconciler_uses_builtin_reader_and_does_not_retry_corruption() -> None: + context, resources = _manifest_inspection_fixture() + data = resources["https://files.example/data.jsonl"] + resources["https://files.example/data.jsonl"] = b"X" + data[1:] + reader = _ResourceReader(resources) + response = _response() + response.update( + periods=[context.obligation.model_dump(mode="json", exclude_none=True)], + revisions=[context.revision.model_dump(mode="json", exclude_none=True)], + materializations=[context.materialization.model_dump(mode="json", exclude_none=True)], + pagination={"has_more": False, "total_count": 3}, + ) + + class Client: + async def get_reporting_status(self, request): + return TaskResult( + status=TaskStatus.COMPLETED, + data=GetReportingStatusResponse.model_validate(response), + ) + + async def sync_reporting_receipts(self, request): + raise AssertionError("corrupt data must not produce a receipt") + + with pytest.raises(ReportingReconciliationError) as error: + await reconcile_reporting( + Client(), + GetReportingStatusRequest.model_validate( + { + "account": {"account_id": "account-1"}, + "view": "periods", + "period": {"start": PERIOD["start"], "end": PERIOD["end"]}, + } + ), + expected_periods=[], + resource_reader=reader, + max_inspection_attempts=3, + ) + + assert error.value.code == "OBJECT_DIGEST_MISMATCH" + assert reader.calls.count("https://files.example/data.jsonl") == 1 + + @pytest.mark.asyncio async def test_reconciles_billing_and_records_matching_receipt() -> None: client = _Client() diff --git a/tests/test_response_enhancer.py b/tests/test_response_enhancer.py index 9b488ee63..e40c2a9e2 100644 --- a/tests/test_response_enhancer.py +++ b/tests/test_response_enhancer.py @@ -489,7 +489,7 @@ class _Store(TestControllerStore): # ``list_scenarios`` is the gate-exempt capability probe — exercises # the bypass closure without needing a resolved sandbox account. result = await executor._tool_callers["comply_test_controller"]( - {"scenario": "list_scenarios"} + {"scenario": "list_scenarios", "account": {"sandbox": True}} ) assert result["enhanced"] is True assert result["success"] is True @@ -513,7 +513,11 @@ class _Store(TestControllerStore): response_enhancer=enhancer, ) result = await executor._tool_callers["comply_test_controller"]( - {"scenario": "list_scenarios", "context": {"correlation_id": "abc"}} + { + "scenario": "list_scenarios", + "account": {"sandbox": True}, + "context": {"correlation_id": "abc"}, + } ) assert result["enhanced"] is True assert seen == {"name": "comply_test_controller", "context_present": True} diff --git a/tests/test_server_dx.py b/tests/test_server_dx.py index 33ecf6c9b..eb56abe3d 100644 --- a/tests/test_server_dx.py +++ b/tests/test_server_dx.py @@ -30,13 +30,14 @@ update_media_buy_response, ) from adcp.server.test_controller import ( - TestControllerError as ControllerError, -) -from adcp.server.test_controller import ( + SCENARIOS, TestControllerStore, _handle_test_controller, _list_scenarios, ) +from adcp.server.test_controller import ( + TestControllerError as ControllerError, +) _PACKAGED_ADCP_VERSION = normalize_to_release_precision(get_adcp_spec_version()) @@ -797,6 +798,70 @@ def test_empty_store(self): assert _list_scenarios(store) == [] +def test_store_signatures_cover_documented_scenario_params(): + """Protocol additions must update the public store contract. + + The upstream schema currently records optional parameter ownership in + descriptions rather than a machine-readable per-scenario property map. + Match both spellings it uses for delivery until that metadata is exposed. + """ + import inspect + import json + from pathlib import Path + + from adcp import get_adcp_spec_version + + schema_path = ( + Path(__file__).parents[1] + / "schemas" + / "cache" + / get_adcp_spec_version() + / "compliance" + / "comply-test-controller-request.json" + ) + schema = json.loads(schema_path.read_text()) + properties = schema["properties"]["params"]["properties"] + terms_by_scenario = { + "force_creative_status": ("force_creative_status",), + "simulate_delivery": ("simulate_delivery", "simulated delivery"), + } + + for scenario, terms in terms_by_scenario.items(): + documented = { + name + for name, field_schema in properties.items() + if any(term in json.dumps(field_schema).lower() for term in terms) + } + accepted = set(inspect.signature(getattr(TestControllerStore, scenario)).parameters) + missing = documented - accepted + assert not missing, f"{scenario} is missing documented params: {sorted(missing)}" + + +def test_controller_scenarios_match_packaged_schema(): + """Protocol upgrades cannot silently leave the server scenario list behind.""" + import json + from pathlib import Path + + from adcp import get_adcp_spec_version + + schema_path = ( + Path(__file__).parents[1] + / "schemas" + / "cache" + / get_adcp_spec_version() + / "compliance" + / "comply-test-controller-request.json" + ) + schema = json.loads(schema_path.read_text()) + documented = { + condition["if"]["properties"]["scenario"]["const"] + for condition in schema["allOf"] + if "if" in condition and "scenario" in condition["if"].get("properties", {}) + } + + assert set(SCENARIOS) == documented + + class TestTestControllerError: @pytest.mark.asyncio async def test_error_is_caught(self): @@ -866,6 +931,358 @@ async def test_unimplemented_scenario_returns_error(self): assert result["success"] is False assert result["error"] == "UNKNOWN_SCENARIO" + @pytest.mark.parametrize( + ("scenario", "params", "account", "expected"), + [ + ( + "expire_account_change_cursor", + {}, + {"sandbox": True, "account_id": "acct-1"}, + {"account_id": "acct-1"}, + ), + ( + "force_creative_purge", + {"creative_id": "cr-1", "purge_kind": "soft"}, + None, + {"creative_id": "cr-1", "purge_kind": "soft"}, + ), + ( + "force_get_products_arm", + {"arm": "submitted", "task_id": "task-products"}, + None, + {"arm": "submitted", "task_id": "task-products"}, + ), + ( + "force_get_signals_arm", + {"arm": "submitted", "task_id": "task-signals"}, + None, + {"arm": "submitted", "task_id": "task-signals"}, + ), + ( + "seed_account", + {"account_id": "acct-seed", "fixture": {"status": "active"}}, + None, + {"account_id": "acct-seed"}, + ), + ( + "seed_rights_grant", + {"rights_id": "rights-1", "fixture": {"status": "active"}}, + None, + {"rights_id": "rights-1"}, + ), + ( + "seed_measurement_catalog", + { + "vendor": {"domain": "measurement.example"}, + "metrics": [{"metric_id": "attention"}], + }, + None, + {"vendor": {"domain": "measurement.example"}}, + ), + ( + "query_upstream_traffic", + {"endpoint_pattern": "POST *", "limit": 25}, + None, + {"endpoint_pattern": "POST *", "limit": 25}, + ), + ( + "query_provenance_audit_observations", + {"creative_id": "cr-audit"}, + None, + {"creative_id": "cr-audit"}, + ), + ( + "force_upstream_unavailable", + {"tool": "get_products", "upstream_name": "inventory"}, + None, + {"tool": "get_products", "upstream_name": "inventory"}, + ), + ( + "catalog_item_availability_probe", + { + "operation": "seed_inaccessible_item", + "catalog_id": "catalog-1", + "item_id": "item-1", + }, + None, + {"operation": "seed_inaccessible_item", "catalog_id": "catalog-1"}, + ), + ( + "compact_product_lifecycle_probe", + {"operation": "prepare", "product_id": "product-1"}, + None, + {"operation": "prepare", "product_id": "product-1"}, + ), + ( + "compact_direct_buy_lifecycle_probe", + {"operation": "prepare", "product_id": "product-direct"}, + None, + {"operation": "prepare", "product_id": "product-direct"}, + ), + ], + ) + @pytest.mark.asyncio + async def test_schema_scenarios_dispatch_to_store( + self, + scenario: str, + params: dict[str, Any], + account: dict[str, Any] | None, + expected: dict[str, Any], + ): + received: dict[str, Any] = {} + + async def implementation(self, **kwargs: Any) -> dict[str, Any]: + received.update(kwargs) + return {"simulated": kwargs} + + store_type = type( + "_SchemaScenarioStore", (TestControllerStore,), {scenario: implementation} + ) + request: dict[str, Any] = {"scenario": scenario, "params": params} + if account is not None: + request["account"] = account + + result = await _handle_test_controller(store_type(), request) + + assert result["success"] is True + for name, value in params.items(): + assert received[name] == value + for name, value in expected.items(): + assert received[name] == value + + @pytest.mark.asyncio + async def test_signature_dispatch_preserves_legacy_optional_none_arguments(self): + received: dict[str, Any] = {} + + class LegacyStore(TestControllerStore): + async def simulate_delivery( + self, + media_buy_id: str, + impressions: int | None, + clicks: int | None, + conversions: int | None, + reported_spend: dict[str, Any] | None, + ) -> dict[str, Any]: + received.update( + media_buy_id=media_buy_id, + impressions=impressions, + clicks=clicks, + conversions=conversions, + reported_spend=reported_spend, + ) + return {} + + result = await _handle_test_controller( + LegacyStore(), + {"scenario": "simulate_delivery", "params": {"media_buy_id": "mb-1"}}, + ) + + assert result["success"] is True + assert received == { + "media_buy_id": "mb-1", + "impressions": None, + "clicks": None, + "conversions": None, + "reported_spend": None, + } + + @pytest.mark.parametrize( + ("scenario", "params"), + [ + ("expire_account_change_cursor", {}), + ("force_get_products_arm", {"arm": "rejected"}), + ("force_get_signals_arm", {"arm": "submitted"}), + ( + "catalog_item_availability_probe", + {"operation": "query_eligibility", "catalog_id": "cat-1", "item_id": "item-1"}, + ), + ("compact_product_lifecycle_probe", {"operation": "prepare"}), + ( + "compact_direct_buy_lifecycle_probe", + {"operation": "expire_proposal", "product_id": "product-1"}, + ), + ], + ) + @pytest.mark.asyncio + async def test_schema_scenario_validation_rejects_invalid_params( + self, + scenario: str, + params: dict[str, Any], + ): + async def implementation(self, **kwargs: Any) -> dict[str, Any]: + return {"simulated": kwargs} + + store_type = type( + "_SchemaScenarioStore", (TestControllerStore,), {scenario: implementation} + ) + result = await _handle_test_controller( + store_type(), + {"scenario": scenario, "params": params}, + ) + + assert result["success"] is False + assert result["error"] == "INVALID_PARAMS" + + @pytest.mark.asyncio + async def test_simulate_delivery_dispatches_reporting_metrics(self): + received: dict[str, Any] = {} + + class _DeliveryStore(TestControllerStore): + async def simulate_delivery( + self, + media_buy_id: str, + **metrics: Any, + ) -> dict[str, Any]: + received.update(media_buy_id=media_buy_id, **metrics) + return {"simulated": received} + + delivery_params = { + "media_buy_id": "mb-1", + "impressions": 10_000, + "clicks": 150, + "plays": 240, + "dooh_metrics": {"loop_plays": 240, "screens_used": 12}, + "conversions": 20, + "delivery_date": "2026-09-03", + "conversion_value": 400.0, + "commissionable_value": 250.0, + "reported_spend": {"amount": 150.0, "currency": "USD"}, + "reach": 750.0, + "frequency": 2.5, + "reach_window": { + "kind": "rolling", + "period": {"interval": 7, "unit": "days"}, + }, + "viewability": { + "measurable_impressions": 900, + "viewable_impressions": 600, + }, + "vendor_metric_values": [ + { + "vendor": {"domain": "measurement.example"}, + "metric_id": "attention", + "value": 12.0, + } + ], + "vendor_metric_values_by_package": { + "pkg-1": [ + { + "vendor": {"domain": "measurement.example"}, + "metric_id": "attention", + "value": 12.0, + } + ] + }, + "not_yet_measurable_vendor_metrics": [ + {"vendor": {"domain": "measurement.example"}, "metric_id": "brand_lift"} + ], + "not_yet_measurable_vendor_metrics_by_package": { + "pkg-1": [{"vendor": {"domain": "measurement.example"}, "metric_id": "brand_lift"}] + }, + } + result = await _handle_test_controller( + _DeliveryStore(), + { + "scenario": "simulate_delivery", + "params": delivery_params, + }, + ) + + assert result["success"] is True + assert received == delivery_params + + @pytest.mark.asyncio + async def test_dispatches_future_explicit_param(self): + received: dict[str, Any] = {} + + class _ExtensibleDeliveryStore(TestControllerStore): + async def simulate_delivery( + self, + media_buy_id: str, + future_metric: dict[str, Any] | None = None, + ) -> dict[str, Any]: + received["future_metric"] = future_metric + return {"simulated": {"media_buy_id": media_buy_id}} + + result = await _handle_test_controller( + _ExtensibleDeliveryStore(), + { + "scenario": "simulate_delivery", + "params": {"media_buy_id": "mb-future", "future_metric": {"value": 42}}, + }, + ) + + assert result["success"] is True + assert received == {"future_metric": {"value": 42}} + + @pytest.mark.asyncio + async def test_force_creative_status_dispatches_reason_fields(self): + received: dict[str, Any] = {} + + class _CreativeStore(TestControllerStore): + async def force_creative_status( + self, + creative_id: str, + status: str, + rejection_reason: str | None = None, + reason_code: str | None = None, + reason_detail: str | None = None, + ) -> dict[str, Any]: + received.update( + creative_id=creative_id, + status=status, + rejection_reason=rejection_reason, + reason_code=reason_code, + reason_detail=reason_detail, + ) + return {"previous_state": "pending", "current_state": status} + + result = await _handle_test_controller( + _CreativeStore(), + { + "scenario": "force_creative_status", + "params": { + "creative_id": "cr-1", + "status": "rejected", + "rejection_reason": "Policy violation", + "reason_code": "policy_violation", + "reason_detail": "Blocked by sandbox policy", + }, + }, + ) + + assert result["success"] is True + assert received["reason_code"] == "policy_violation" + assert received["reason_detail"] == "Blocked by sandbox policy" + + @pytest.mark.asyncio + async def test_simulate_delivery_keeps_legacy_overrides_working(self): + class _LegacyDeliveryStore(TestControllerStore): + async def simulate_delivery( + self, + media_buy_id: str, + impressions: int | None = None, + clicks: int | None = None, + conversions: int | None = None, + reported_spend: dict[str, Any] | None = None, + ) -> dict[str, Any]: + return {"simulated": {"media_buy_id": media_buy_id}} + + result = await _handle_test_controller( + _LegacyDeliveryStore(), + { + "scenario": "simulate_delivery", + "params": { + "media_buy_id": "mb-legacy", + "reach": 50.0, + "future_metric": {"value": 42}, + }, + }, + ) + + assert result["success"] is True + assert result["simulated"]["media_buy_id"] == "mb-legacy" + @pytest.mark.asyncio async def test_unknown_scenario(self): store = MinimalStore() @@ -990,6 +1407,91 @@ def test_registers_tool(self): tool_names = [t.name for t in mcp._tool_manager.list_tools()] assert "comply_test_controller" in tool_names + def test_advertises_canonical_request_schema(self): + from mcp.server import MCPServer + + from adcp.server.test_controller import register_test_controller + from adcp.validation.schema_loader import get_mcp_schema + + mcp = MCPServer("test") + register_test_controller(mcp, MinimalStore()) + + tool = mcp._tool_manager._tools["comply_test_controller"] + assert tool.parameters == get_mcp_schema("comply_test_controller", "request") + assert "account" in tool.parameters["required"] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "params", + [ + {"spend_percentage": 50}, + None, + ], + ) + async def test_runtime_rejects_noncanonical_scenario_params(self, params: Any): + from mcp.server import MCPServer + + from adcp.server.test_controller import register_test_controller + + called = False + + class Store(TestControllerStore): + async def simulate_budget_spend( + self, + spend_percentage: float, + account_id: str | None = None, + media_buy_id: str | None = None, + ) -> dict[str, Any]: + nonlocal called + called = True + return {} + + mcp = MCPServer("test") + register_test_controller(mcp, Store()) + fn = mcp._tool_manager._tools["comply_test_controller"].fn + + result = await fn( + account={"sandbox": True}, + scenario="simulate_budget_spend", + params=params, + ) + + assert result["success"] is False + assert result["error"] == "INVALID_PARAMS" + assert called is False + + @pytest.mark.asyncio + async def test_runtime_rejects_invalid_conditional_types_and_enums(self): + from mcp.server import MCPServer + + from adcp.server.test_controller import register_test_controller + + called = False + + class Store(TestControllerStore): + async def force_creative_purge( + self, + creative_id: str, + purge_kind: str | None = None, + ) -> dict[str, Any]: + nonlocal called + called = True + return {} + + mcp = MCPServer("test") + register_test_controller(mcp, Store()) + fn = mcp._tool_manager._tools["comply_test_controller"].fn + + result = await fn( + account={"sandbox": True}, + scenario="force_creative_purge", + params={"creative_id": 123, "purge_kind": "wrong"}, + ) + + assert result["success"] is False + assert result["error"] == "INVALID_PARAMS" + assert called is False + class TestServeWithTestController: def test_serve_accepts_test_controller(self): diff --git a/tests/test_server_principal.py b/tests/test_server_principal.py new file mode 100644 index 000000000..277ac893e --- /dev/null +++ b/tests/test_server_principal.py @@ -0,0 +1,468 @@ +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +import pytest + +from adcp.server import InMemoryPrincipalRecordStore, PrincipalIdentity, PrincipalService + +IDENTITY = PrincipalIdentity("https://buyer.example/agent", "buyer_agent") +DESTINATION = { + "pattern": "file_transfer", + "destination_id": "archive", + "active": True, + "provider": {"domain": "object-store.example"}, + "transport": "s3", + "location": "s3://buyer-reporting/adcp/", + "accepted_formats": ["parquet"], + "accepted_verification_profiles": ["manifest_checksums"], +} + + +def _request(key: str, configuration, **extra): + return { + "idempotency_key": key, + "configuration": configuration, + **extra, + } + + +@pytest.mark.asyncio +async def test_recognized_and_unconfigured_read_variants() -> None: + service = PrincipalService() + assert (await service.get_principal(IDENTITY)).result.kind == "unconfigured" + + principal_id = await service.recognize(IDENTITY) + response = await service.get_principal(IDENTITY) + + assert response.result.kind == "recognized" + assert response.result.principal_id == principal_id + + +@pytest.mark.asyncio +async def test_sync_persists_state_negotiation_and_exact_replay() -> None: + service = PrincipalService( + accepted_async_adcp_versions=("3.2",), + accepted_experimental_features=("protocol.principal",), + ) + request = _request( + "principal-server-key-0001", + { + "reporting_destinations": [DESTINATION], + "declarations": { + "async_adcp_versions": ["3.2", "3.1"], + "experimental_features": ["protocol.principal", "future.feature"], + }, + }, + context={"correlation_id": "principal--apply"}, + ) + + first = await service.sync_principal(IDENTITY, request) + replay_request = dict(request, context={"correlation_id": "principal--replay"}) + replay = await service.sync_principal(IDENTITY, replay_request) + current = await service.get_principal( + IDENTITY, {"context": {"correlation_id": "principal--read"}} + ) + + assert first.model_dump(mode="json", exclude={"replayed"}) == replay.model_dump( + mode="json", exclude={"replayed"} + ) + assert first.replayed is False + assert replay.replayed is True + assert first.result.kind == "applied" + assert first.context.correlation_id == "principal--apply" + assert current.context.correlation_id == "principal--read" + assert current.result.kind == "current" + destination = current.result.configuration.reporting_destinations[0] + assert destination.state == "validating" + declarations = current.result.configuration.declarations + assert declarations.selected_async_adcp_version == "3.2" + assert [item.value for item in declarations.exclusions] == ["3.1", "future.feature"] + + +@pytest.mark.asyncio +async def test_version_fence_rejects_stale_replacement() -> None: + service = PrincipalService() + await service.sync_principal( + IDENTITY, _request("principal-server-key-0002", {"declarations": {}}) + ) + + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0003", + {"declarations": {}}, + expected_configuration_version="stale", + ), + ) + + assert response.result.kind == "failed" + assert response.result.errors[0].code == "CONFLICT" + + +@pytest.mark.asyncio +async def test_dry_run_does_not_call_proof_hook_or_persist() -> None: + calls = 0 + + async def proof(identity, desired, previous): + nonlocal calls + calls += 1 + raise AssertionError("dry run must not issue proof") + + service = PrincipalService(destination_proof=proof) + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0004", + {"reporting_destinations": [DESTINATION]}, + dry_run=True, + ), + ) + + assert response.result.kind == "validated" + assert calls == 0 + assert (await service.get_principal(IDENTITY)).result.kind == "unconfigured" + + +@pytest.mark.asyncio +async def test_destination_transition_emits_invalidation_without_version_churn() -> None: + changes = [] + + async def emit(change): + changes.append(change) + + service = PrincipalService(emit_change=emit) + applied = await service.sync_principal( + IDENTITY, + _request("principal-server-key-0005", {"reporting_destinations": [DESTINATION]}), + ) + version = applied.result.configuration_version + + transitioned = await service.transition_destination(IDENTITY, "archive", "ready") + current = await service.get_principal(IDENTITY) + + assert transitioned.state == "ready" + assert current.result.configuration_version == version + assert changes[0].reason == "destination_state_changed" + assert changes[0].destination_id == "archive" + + +@pytest.mark.asyncio +async def test_suspension_preserves_generation_and_revocation_retires_it() -> None: + service = PrincipalService() + first = await service.sync_principal( + IDENTITY, + _request("principal-server-key-0006", {"reporting_destinations": [DESTINATION]}), + ) + original = first.result.configuration.reporting_destinations[0] + suspended = dict(DESTINATION, active=False) + + second = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0007", + {"reporting_destinations": [suspended]}, + expected_configuration_version=first.result.configuration_version, + ), + ) + inactive = second.result.configuration.reporting_destinations[0] + assert inactive.destination_ref == original.destination_ref + assert inactive.state == "inactive" + + third = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0008", + {"reporting_destinations": []}, + expected_configuration_version=second.result.configuration_version, + ), + ) + assert third.result.configuration.reporting_destinations == [] + retired = third.result.configuration.retired_destinations[0] + assert retired.destination_id == "archive" + assert retired.destination_refs[0].root == original.destination_ref + + +@pytest.mark.asyncio +async def test_active_notification_requires_proof_hook() -> None: + service = PrincipalService() + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0009", + { + "notification_configs": [ + { + "subscriber_id": "events", + "url": "https://buyer.example/webhooks", + "event_types": ["principal.changed"], + "active": True, + } + ] + }, + ), + ) + + assert response.result.kind == "failed" + assert "notification_proof" in response.result.errors[0].message + + +@pytest.mark.asyncio +async def test_seller_declaration_change_emits_without_version_churn() -> None: + changes = [] + + async def emit(change): + changes.append(change) + + service = PrincipalService(emit_change=emit, accepted_async_adcp_versions=("3.2", "3.1")) + applied = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0010", + {"declarations": {"async_adcp_versions": ["3.2", "3.1"]}}, + ), + ) + version = applied.result.configuration_version + service.set_declaration_support(async_adcp_versions=("3.2",)) + + declarations = await service.refresh_declarations(IDENTITY) + current = await service.get_principal(IDENTITY) + + assert [item.root for item in declarations.accepted.async_adcp_versions] == ["3.2"] + assert current.result.configuration_version == version + assert changes[0].reason == "declarations_intersection_changed" + + +@pytest.mark.asyncio +async def test_rejects_duplicate_logical_keys_before_state_is_persisted() -> None: + service = PrincipalService() + destination = dict(DESTINATION, active=False) + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0011", + {"reporting_destinations": [destination, destination]}, + ), + ) + assert response.result.kind == "failed" + assert "duplicates" in response.result.errors[0].message + + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0012", + { + "notification_configs": [ + { + "subscriber_id": "duplicate", + "url": "https://127.0.0.1/hook", + "event_types": ["principal.changed"], + "active": False, + } + ] + * 2 + }, + ), + ) + assert response.result.kind == "failed" + assert "duplicates" in response.result.errors[0].message + + +@pytest.mark.asyncio +async def test_inactive_config_still_enforces_webhook_and_destination_security() -> None: + service = PrincipalService() + bad_webhook = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0013", + { + "notification_configs": [ + { + "subscriber_id": "events", + "url": "https://127.0.0.1/hook", + "event_types": ["principal.changed"], + "active": False, + } + ] + }, + ), + ) + assert bad_webhook.result.kind == "failed" + assert "SSRF" in bad_webhook.result.errors[0].message + + secret_webhook = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0013a", + { + "notification_configs": [ + { + "subscriber_id": "events", + "url": "https://buyer.example/hook?token=not-a-credential", + "event_types": ["principal.changed"], + "active": False, + } + ] + }, + ), + ) + assert secret_webhook.result.kind == "failed" + assert "credential or secret" in secret_webhook.result.errors[0].message + + bad_destination = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0014", + { + "reporting_destinations": [ + dict(DESTINATION, active=False, location="s3://key:password@bucket/path") + ] + }, + ), + ) + assert bad_destination.result.kind == "failed" + assert "credentials" in bad_destination.result.errors[0].message + + userinfo_destination = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0014a", + { + "reporting_destinations": [ + dict(DESTINATION, active=False, location="s3://access-key@bucket/path") + ] + }, + ), + ) + assert userinfo_destination.result.kind == "failed" + assert "userinfo credentials" in userinfo_destination.result.errors[0].message + + +@pytest.mark.asyncio +async def test_rejects_invalid_notification_event_scope_and_missing_acknowledgment() -> None: + service = PrincipalService() + for key, event_type, expected in ( + ("principal-server-key-0015", "scheduled", "media-buy-anchored"), + ("principal-server-key-0016", "creative.status_changed", "all_authorized_accounts"), + ): + response = await service.sync_principal( + IDENTITY, + _request( + key, + { + "notification_configs": [ + { + "subscriber_id": "events", + "url": "https://127.0.0.1/hook", + "event_types": [event_type], + "active": False, + } + ] + }, + ), + ) + assert response.result.kind == "failed" + assert expected in response.result.errors[0].message + + +@pytest.mark.asyncio +async def test_unsupported_declarations_persist_empty_intersection_and_exclusions() -> None: + service = PrincipalService() + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0017", + {"declarations": {"async_adcp_versions": ["3.1"]}}, + ), + ) + + assert response.result.kind == "applied" + declarations = response.result.configuration.declarations + assert declarations.accepted.async_adcp_versions is None + assert declarations.exclusions[0].value == "3.1" + + +@pytest.mark.asyncio +async def test_active_notifications_need_an_accepted_signing_algorithm(monkeypatch) -> None: + def valid_url(url, *, field): + return SimpleNamespace(effective_url=url, hostname="buyer.example") + + monkeypatch.setattr("adcp.server.principal._validate_webhook_destination_url", valid_url) + + async def prove(identity, config): + from adcp.server.principal import _notification_state + + return _notification_state(config) + + service = PrincipalService(notification_proof=prove) + response = await service.sync_principal( + IDENTITY, + _request( + "principal-server-key-0018", + { + "notification_configs": [ + { + "subscriber_id": "events", + "url": "https://buyer.example/hook", + "event_types": ["principal.changed"], + } + ] + }, + ), + ) + assert response.result.kind == "failed" + assert response.result.errors[0].code == "UNSUPPORTED_FEATURE" + + +class _RacingStore: + def __init__(self) -> None: + self._store = InMemoryPrincipalRecordStore() + self._race = False + self._readers = 0 + self._gate = asyncio.Event() + + async def get(self, subject): + record = await self._store.get(subject) + if self._race and record is not None: + self._readers += 1 + if self._readers == 2: + self._gate.set() + await self._gate.wait() + return record + + async def compare_and_swap(self, subject, expected_store_revision, record): + return await self._store.compare_and_swap(subject, expected_store_revision, record) + + +@pytest.mark.asyncio +async def test_shared_store_cas_rejects_one_of_two_concurrent_fenced_writes() -> None: + store = _RacingStore() + first = PrincipalService(store, accepted_async_adcp_versions=("3.1", "3.2")) + second = PrincipalService(store, accepted_async_adcp_versions=("3.1", "3.2")) + initial = await first.sync_principal( + IDENTITY, _request("principal-server-key-0019", {"declarations": {}}) + ) + store._race = True + version = initial.result.configuration_version + one, two = await asyncio.gather( + first.sync_principal( + IDENTITY, + _request( + "principal-server-key-0020", + {"declarations": {"async_adcp_versions": ["3.1"]}}, + expected_configuration_version=version, + ), + ), + second.sync_principal( + IDENTITY, + _request( + "principal-server-key-0021", + {"declarations": {"async_adcp_versions": ["3.2"]}}, + expected_configuration_version=version, + ), + ), + ) + assert {one.result.kind, two.result.kind} == {"applied", "failed"} + failed = next(item for item in (one, two) if item.result.kind == "failed") + assert failed.result.errors[0].code == "CONFLICT" diff --git a/tests/test_test_controller_context.py b/tests/test_test_controller_context.py index 0c5cf6062..e601055cb 100644 --- a/tests/test_test_controller_context.py +++ b/tests/test_test_controller_context.py @@ -109,12 +109,10 @@ def test_accepts_context_kwarg_rejects_positional_only_context(): # the file still parses cleanly without introducing a sigil. ns: dict[str, Any] = {} exec( - textwrap.dedent( - """ + textwrap.dedent(""" async def fn(self, context, /, account_id, status): return {} - """ - ), + """), ns, ) assert _accepts_context_kwarg(ns["fn"]) is False @@ -399,6 +397,7 @@ def build_context(meta: RequestMetadata) -> ToolContext: # FastMCP's tool wrapper takes the function args as kwargs. fn = tool.fn # type: ignore[attr-defined] result = await fn( + account={"sandbox": True}, scenario="force_account_status", params={"account_id": "acc-1", "status": "suspended"}, ) @@ -431,7 +430,7 @@ async def force_account_status(self, account_id: str, status: str) -> dict[str, tool = mcp._tool_manager._tools["comply_test_controller"] fn = tool.fn # type: ignore[attr-defined] - result = await fn(scenario="list_scenarios") + result = await fn(account={"sandbox": True}, scenario="list_scenarios") assert isinstance(result, dict), "must be a dict, not a JSON string" assert result["success"] is True @@ -457,6 +456,7 @@ def bad_factory(meta: RequestMetadata) -> Any: with pytest.raises(TypeError, match="not a ToolContext"): await fn( + account={"sandbox": True}, scenario="force_account_status", params={"account_id": "acc-1", "status": "suspended"}, )