From dba15b6b0dd8063d38c2df4196491ae64d9720ae Mon Sep 17 00:00:00 2001 From: Brian O'Kelley Date: Tue, 22 Sep 2026 13:58:39 +0000 Subject: [PATCH] fix(reporting): restore progress for late materializer accounts Reset the materializer sampling continuation only after a committed account turn moves its durable served_at rank. Keep continuation across lock misses so a busy prefix cannot block later eligible accounts. Cover continuous producer activity, both account orders, size-one pools, more than two busy pages, concurrent workers, commit rollback and fencing. Exercise typed public enrollment, real HTTP/PG delivery, 503-row exact reads, receipts, reconciliation and restart under continuing catch-up load. Include the store regression in installed production distributions. --- src/adcp/reporting/materializer/pg.py | 7 +- .../reporting/_late_account_server.py | 393 ++++++++++++++++ .../reporting/_late_account_support.py | 103 +++++ .../reporting/_production_packaging.py | 1 + .../test_reporting_materializer_progress.py | 311 +++++++++++++ ...test_reporting_production_late_accounts.py | 433 ++++++++++++++++++ 6 files changed, 1247 insertions(+), 1 deletion(-) create mode 100644 tests/conformance/reporting/_late_account_server.py create mode 100644 tests/conformance/reporting/_late_account_support.py create mode 100644 tests/conformance/reporting/test_reporting_materializer_progress.py create mode 100644 tests/conformance/reporting/test_reporting_production_late_accounts.py diff --git a/src/adcp/reporting/materializer/pg.py b/src/adcp/reporting/materializer/pg.py index 642d5a0a2..d638af0e7 100644 --- a/src/adcp/reporting/materializer/pg.py +++ b/src/adcp/reporting/materializer/pg.py @@ -340,7 +340,12 @@ async def claim_materialization( ) result = await self._claim_account_on(connection, account_id, keys, lease_seconds) await self._schedule_account_on(connection, account_id) - return result + # A committed account turn moves its durable served_at rank. Start + # there again so continuously due peers cannot hide newly enrolled + # accounts behind this pre-update cursor. Lock misses above retain + # the continuation, allowing the next bounded page past a busy prefix. + self._materializer_sample_after = None + return result return ReportingMaterializerTurn("idle") async def _claim_account_on( diff --git a/tests/conformance/reporting/_late_account_server.py b/tests/conformance/reporting/_late_account_server.py new file mode 100644 index 000000000..97bdd9482 --- /dev/null +++ b/tests/conformance/reporting/_late_account_server.py @@ -0,0 +1,393 @@ +"""Public production support starts empty; all enrollment arrives over real HTTP.""" + +import argparse +import importlib.metadata +import json +import os +from datetime import datetime, timezone +from pathlib import Path + +from psycopg_pool import AsyncConnectionPool + +import adcp +from adcp.decisioning.capabilities import Account as AccountCapabilities +from adcp.reporting.ledger import ( + ProducerOfferings, + ReportingConfiguration, + ReportingDestinationBinding, + ReportingProducer, + ReportingScheduleSpec, +) +from adcp.reporting.materializer import ( + ReportingDestinationIO, + ReportingMaterializerService, + ReportingRevisionVerifierRegistry, + ReportingWriterCapability, +) +from adcp.reporting.production import ( + PgReportingProductionStore, + ReportingConfigurationAdmission, + ReportingProductionConfigurationTask, + ReportingProductionOffering, + ReportingProductionSupport, +) +from adcp.reporting.projection import PgReportingStatusProjection +from adcp.reporting.receipts import ReportingReceiptError +from adcp.server import serve +from adcp.server.auth import BearerTokenAuth, Principal, auth_context_factory +from adcp.types import ReportingDeliveryOffering + +from ._generation_support import END, START +from ._late_account_support import ACCOUNTS, AccountSource, CurrencyDestination, verifier_for + +CONSUMERS = {"usd": "urn:buyer:usd", "eur": "urn:buyer:eur"} + + +class WireCapture: + def __init__(self, app, path): + self.app, self.path = app, Path(path) + + async def __call__(self, scope, receive, send): + if scope["type"] != "http" or scope.get("method") != "POST": + return await self.app(scope, receive, send) + request, response = bytearray(), bytearray() + status = None + + async def incoming(): + message = await receive() + if message["type"] == "http.request": + request.extend(message.get("body", b"")) + return message + + async def outgoing(message): + nonlocal status + if message["type"] == "http.response.start": + status = message["status"] + if message["type"] == "http.response.body": + response.extend(message.get("body", b"")) + await send(message) + + try: + await self.app(scope, incoming, outgoing) + finally: + assert len(request) < 2_000_000 and len(response) < 4_000_000 + with self.path.open("a") as output: + output.write( + json.dumps( + { + "path": scope["path"], + "status": status, + "request_utf8": request.decode(), + "response_utf8": response.decode(), + } + ) + + "\n" + ) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, required=True) + parser.add_argument("--root", type=Path, required=True) + parser.add_argument("--schema", required=True) + parser.add_argument("--notifications", action="store_true") + args = parser.parse_args() + root = args.root + pool = AsyncConnectionPool( + os.environ["ADCP_PG_TEST_URL"], + kwargs={ + "autocommit": True, + "application_name": args.schema, + "options": f"-csearch_path={args.schema} -cstatement_timeout=15000", + }, + min_size=1, + max_size=1, + open=False, + ) + store = PgReportingProductionStore(pool=pool, notifications=args.notifications) + capability = ReportingWriterCapability( + "warehouse_materialization", + "fixture-sql", + "jsonl", + "canonical_digest", + "destination", + "immutable_location", + "sha256", + "conditional_create", + ) + verifiers = {currency: verifier_for(currency, capability) for currency in ("USD", "EUR")} + writer = CurrencyDestination( + root / "destination.sqlite", tuple(v.key for v in verifiers.values()) + ) + registry = ReportingRevisionVerifierRegistry(tuple(verifiers.values())) + sources, producers, offerings = {}, {}, {} + # This adopter exposes an immutable historical dataset observed before + # startup. Its source observation remains that timestamp on every fetch; + # production turns and PostgreSQL lease scheduling use their real clocks. + source_observed_at = datetime.now(timezone.utc) + for currency, verifier in verifiers.items(): + key = verifier.key + source = AccountSource( + key, + root / ("source-" + currency), + clock=lambda: source_observed_at, + official=True, + ) + producer = ReportingProducer( + source=source, + store=store, + offerings=ProducerOfferings( + official_offering_id=source.source_id, + publication_namespace=source.capabilities.offerings[0].publication_namespace, + source_scope=source.capabilities.source_scope, + ), + object_reader=source.reader, + revision_verifier=verifier, + # Real clock, bounded catch-up: two actual publications per producer + # turn keep pending work behind a materializer that serves one turn. + # Public reads pin a delivered period while later periods provide + # continuing production load; no scheduler call is injected. + max_periods_per_turn=2, + currency_resolver=lambda config, obligation: ACCOUNTS[config.account_id][0], + ) + profile = { + "id": key.reporting_profile, + "version": key.definition.schema_version, + "schema_uri": key.definition.schema_uri, + "schema_sha256": key.definition.schema_sha256, + "schema_dialect": key.definition.schema_dialect, + "schema_ref_policy": key.definition.schema_ref_policy, + "grain": "row", + "primary_keys": ["row_id"], + "canonicalization_id": key.canonicalization.canonicalization_id, + "canonicalization_uri": key.canonicalization.canonicalization_uri, + "canonicalization_sha256": key.canonicalization.canonicalization_sha256, + } + public_offering = ReportingDeliveryOffering.model_validate( + { + "offering_id": "official-" + currency, + "feed_purpose": "billing", + "report_definition_id": key.report_definition_id, + "report_definition_uri": key.definition.report_definition_uri, + "report_definition_sha256": key.definition.report_definition_sha256, + "reporting_profile": profile, + "schedule": {"period_duration": "PT1H", "alignment": "utc", "delivery_sla": "PT1H"}, + "supported_finality": ["official"], + "reconciliation_mode": "consumer_receipt", + "method": writer.delivery_methods[0].wire(), + } + ) + sources[currency], producers[currency] = source, producer + offerings[currency] = ReportingProductionOffering( + public_offering, producer, key, source.source_id + ) + + def config_for(account): + key = verifiers[ACCOUNTS[account][0]].key + return ReportingConfiguration( + "shared-config", + 1, + account, + key.report_definition_id, + key.reporting_profile, + "billing", + ReportingScheduleSpec("PT1H", "PT1H", period_anchor=START), + "official", + activated_at=START, + media_buy_ids=("shared-media-buy",), + definition=key.definition, + ) + + def binding_for(config): + return ReportingDestinationBinding( + config.generation_key, + CONSUMERS[config.account_id], + "shared-destination", + "trusted-" + config.account_id, + capability.method, + capability.transport, + capability.verification_profile, + "consumer_receipt", + "billing", + 400, + START, + capability.format, + ("fixture-sql-v1",), + "delivered", + ) + + def wire_for(config, binding): + offering = offerings[ACCOUNTS[config.account_id][0]] + key = verifiers[ACCOUNTS[config.account_id][0]].key + return { + "delivery_config_id": config.delivery_config_id, + "delivery_config_version": 1, + "offering_id": offering.offering_id, + "active": True, + "feed_purpose": "billing", + "report_definition_id": key.report_definition_id, + "reporting_profile": key.reporting_profile, + "scope": {"media_buy_ids": list(config.media_buy_ids)}, + "coverage_requirement": "full", + "required_finality": "official", + "reconciliation_mode": "consumer_receipt", + "schedule": offering.configuration_schedule(config), + "method": writer.configuration_binding(binding).wire(), + } + + async def resolve_account(reference, context, consumer): + account = reference.get("account_id") + if ( + account not in ACCOUNTS + or reference != {"account_id": account} + or consumer != CONSUMERS[account] + ): + raise ReportingReceiptError("UNAUTHORIZED") + return account + + async def account_task(request, context, admit): + records = [] + for entry in request["accounts"]: + account = await resolve_account(entry["account"], context, context.caller_identity) + config = config_for(account) + binding = binding_for(config) + currency = ACCOUNTS[account][0] + sources[currency].bind_generation(config) + writer.grant(binding) + states = [] + for supplied in entry.get("reporting_delivery_configs", []): + await admit( + ReportingConfigurationAdmission( + offerings[currency].offering_id, + config, + binding, + configuration_wire=supplied, + ) + ) + states.append( + { + "configuration": supplied, + "state": "ready", + "destination_ref": binding.destination_ref, + "validated_at": END.isoformat(), + "activated_at": START.isoformat(), + "current_coverage": { + "status": "full", + "evaluated_at": END.isoformat(), + "media_buy_ids": list(config.media_buy_ids), + "fully_covered_media_buy_ids": list(config.media_buy_ids), + "partially_covered_media_buy_ids": [], + "unsupported_media_buy_ids": [], + "unknown_media_buy_ids": [], + "package_ids": [], + "covered_package_ids": [], + "unsupported_package_ids": [], + "unknown_package_ids": [], + "limitations": [], + }, + } + ) + records.append( + { + "account_id": account, + "brand": {"domain": account + ".advertiser.example.test"}, + "operator": "buyer.example.test", + "action": "unchanged", + "status": "active", + "billing": "operator", + "timezone": "UTC", + "reporting_delivery_configs": states, + } + ) + return {"accounts": records} + + projection = PgReportingStatusProjection( + store, + consumer_status_enabled=True, + revision_ownership=True, + escalation=producers["USD"].escalation, + ) + support = ReportingProductionSupport( + ReportingMaterializerService(store, ReportingDestinationIO(registry, writer), writer), + projection, + offerings=tuple(offerings.values()), + configuration_task=ReportingProductionConfigurationTask( + account_task, + AccountCapabilities(supported_billing=["operator"], require_operator_auth=True), + ), + resolve_account=resolve_account, + poll_seconds=0.02, + ) + # Only external source/provider configuration is restored after restart. + # No revision, receipt, snapshot or successful work is seeded here. + templates = {} + for account in ACCOUNTS: + config = config_for(account) + sources[ACCOUNTS[account][0]].bind_generation(config) + binding = binding_for(config) + writer.grant(binding) + templates[account] = wire_for(config, binding) + + async def startup(): + await pool.open(wait=True) + await store.create_schema() + async with pool.connection() as connection: + initial = await ( + await connection.execute("SELECT count(*) FROM reporting_configurations") + ).fetchone() + await support.start() + (root / "ready.json").write_text( + json.dumps( + { + "templates": templates, + "initial_configurations": initial[0], + "pool_size": 1, + "source_observed_at": source_observed_at.isoformat(), + "python": __import__("sys").version, + "adcp_file": adcp.__file__, + "pydantic": importlib.metadata.version("pydantic"), + "mcp": importlib.metadata.version("mcp"), + } + ) + ) + + async def shutdown(): + await support.aclose() + await pool.close() + (root / "stopped.json").write_text( + json.dumps( + { + "stopped": True, + "source_requests": { + currency: len(source.requests) for currency, source in sources.items() + }, + "destination_sessions_closed": all( + p.opens == p.closes for p in writer.readers.values() + ), + } + ) + ) + + tokens = { + account + "-test-token": Principal(caller_identity=consumer, tenant_id="shared-tenant") + for account, consumer in CONSUMERS.items() + } + serve( + support.handler, + name="late-account-production", + transport="both", + host="127.0.0.1", + port=args.port, + auth=BearerTokenAuth(validate_token=tokens.get), + context_factory=auth_context_factory, + allowed_hosts=["127.0.0.1", "localhost"], + public_url=f"http://127.0.0.1:{args.port}", + on_startup=[startup], + on_shutdown=[shutdown], + stateless_http=True, + asgi_middleware=[(WireCapture, {"path": str(root / "wire.jsonl")})], + ) + + +if __name__ == "__main__": + main() diff --git a/tests/conformance/reporting/_late_account_support.py b/tests/conformance/reporting/_late_account_support.py new file mode 100644 index 000000000..0fd2dfa24 --- /dev/null +++ b/tests/conformance/reporting/_late_account_support.py @@ -0,0 +1,103 @@ +"""Independent currency contracts and persistent adopter I/O for public progress tests.""" + +import base64 +import hashlib +import json +from dataclasses import replace + +import rfc8785 + +from adcp.reporting.inline_source import InlineFetchResult +from adcp.reporting.materializer import ReportingRevisionVerifier, reference_verifier +from adcp.reporting.materializer.contracts import failure + +from ._production_support import Source, SQLiteDestination + +ACCOUNTS = {"usd": ("USD", 503), "eur": ("EUR", 3)} + + +def rows_for(account): + currency, count = ACCOUNTS[account] + return [ + { + "row_id": f"{number:06d}", + "impressions": number % 3, + "spend": "1.25", + "currency": currency, + "details": {"active": True, "values": [1, None, "é", "e\u0301"]}, + } + for number in range(count) + ] + + +class AccountSource(Source): + async def fetch(self, request): + account = str(request.identity.account_id) + self.requests.append(request) + return InlineFetchResult( + rows_for(account), data_through=request.period.end, currency=ACCOUNTS[account][0] + ) + + +class CurrencyDestination(SQLiteDestination): + """The same SQLite destination, with one independently pinned session per currency.""" + + def __init__(self, path, keys): + super().__init__(path, keys[0]) + self.readers = {key: SQLiteDestination(path, key) for key in keys} + + def resolve(self, request, *, phase, context): + provider = self.readers.get(request.verification_key) + if provider is None: + raise failure("AUTHORIZATION_DENIED") + return provider.resolve(request, phase=phase, context=context) + + +def verifier_for(currency, capability): + base = reference_verifier(capability) + if currency == "USD": + return base + assert currency == "EUR" + definition = json.loads(base.definition_bytes) + definition["report_definition_id"] = "reference-report-eur-v1" + for metric in definition["metrics"]: + if metric.get("unit") == "USD": + metric["unit"] = "EUR" + schema = json.loads(base.schema_bytes) + schema["properties"]["currency"]["const"] = "EUR" + schema["properties"]["spend"]["x-adcp-control-total"]["unit"] = "EUR" + + def encode(value): + return (json.dumps(value, ensure_ascii=False, indent=2) + "\n").encode() + + schema_bytes = encode(schema) + schema_hash = hashlib.sha256(schema_bytes).hexdigest() + contract = json.loads(base.canonicalization_bytes) + contract["schema_sha256"] = schema_hash + for vector in contract["golden_vectors"].values(): + for row in vector["input_rows"]: + row["currency"] = "EUR" + canonical = rfc8785.dumps(sorted(vector["input_rows"], key=lambda row: row["row_id"])) + vector["canonical_utf8_base64"] = base64.b64encode(canonical).decode() + vector["sha256"] = hashlib.sha256(canonical).hexdigest() + definition_bytes, contract_bytes = encode(definition), encode(contract) + key = replace( + base.key, + report_definition_id=definition["report_definition_id"], + definition=replace( + base.key.definition, + report_definition_uri="https://contracts.example.test/reference-eur-definition.json", + report_definition_sha256=hashlib.sha256(definition_bytes).hexdigest(), + schema_uri="https://contracts.example.test/reference-eur-schema.json", + schema_sha256=schema_hash, + monetary_metric_units=(("spend", "EUR"),), + monetary_control_total_units=(("spend", "EUR"),), + ), + canonicalization=replace( + base.key.canonicalization, + canonicalization_id="reference-eur-jcs-rows-v1", + canonicalization_uri="https://contracts.example.test/reference-eur-canonicalization.json", + canonicalization_sha256=hashlib.sha256(contract_bytes).hexdigest(), + ), + ) + return ReportingRevisionVerifier(key, definition_bytes, schema_bytes, contract_bytes) diff --git a/tests/conformance/reporting/_production_packaging.py b/tests/conformance/reporting/_production_packaging.py index 5bebe42f8..d0316d13e 100644 --- a/tests/conformance/reporting/_production_packaging.py +++ b/tests/conformance/reporting/_production_packaging.py @@ -257,6 +257,7 @@ def installed_production(root, python, wheel, source, *, label, driver_absent): for name in ( "tests/conformance/reporting/test_reporting_tier_projection.py", "tests/conformance/reporting/test_reporting_schedule_schema.py", + "tests/conformance/reporting/test_reporting_materializer_progress.py", "tests/test_reporting_revision_ownership.py", "tests/test_reporting_capability_models.py", "tests/test_reporting_scope_models.py", diff --git a/tests/conformance/reporting/test_reporting_materializer_progress.py b/tests/conformance/reporting/test_reporting_materializer_progress.py new file mode 100644 index 000000000..4b3c37cdd --- /dev/null +++ b/tests/conformance/reporting/test_reporting_materializer_progress.py @@ -0,0 +1,311 @@ +"""Account progress while peers stay due, including bounded busy-account scans.""" + +import asyncio +import json +import os +import secrets +from contextlib import asynccontextmanager + +import pytest + +from adcp.reporting.materializer import ( + PgReportingMaterializerStore, + ReportingMaterializerLease, + ReportingWriterError, +) + +from ._durable_materializer_support import durable_case + + +@asynccontextmanager +async def progress_pool(*, autocommit=False): + """Each scenario owns a schema and a size-one pool on the real database.""" + url = os.environ.get("ADCP_PG_TEST_URL") + if not url: + pytest.skip("ADCP_PG_TEST_URL not set — requires real PostgreSQL") + psycopg = pytest.importorskip("psycopg") + psycopg_pool = pytest.importorskip("psycopg_pool") + schema = "adcp_materializer_progress_" + secrets.token_hex(6) + async with await psycopg.AsyncConnection.connect(url, autocommit=True) as admin: + await admin.execute( + psycopg.sql.SQL("CREATE SCHEMA {}").format(psycopg.sql.Identifier(schema)) + ) + try: + async with psycopg_pool.AsyncConnectionPool( + url, + kwargs={ + "options": f"-csearch_path={schema} -cstatement_timeout=15000", + "autocommit": autocommit, + }, + min_size=1, + max_size=1, + open=False, + ) as pool: + await pool.wait(timeout=10) + yield pool, schema + finally: + await admin.execute( + psycopg.sql.SQL("DROP SCHEMA {} CASCADE").format(psycopg.sql.Identifier(schema)) + ) + + +async def wake(pool, account): + # This is the installed function invoked by real producer/binding triggers. + # The store-level regression enrolls due accounts, never successful work. + async with pool.connection() as connection: + await connection.execute("SELECT reporting_materializer_wake(%s)", (account,)) + + +async def served(pool): + async with pool.connection() as connection: + return dict( + await ( + await connection.execute( + "SELECT account_id,served_at::text FROM reporting_materializer_accounts" + " ORDER BY served_at,account_id" + ) + ).fetchall() + ) + + +@pytest.mark.parametrize( + "first,late,keep_waking_first,restart_at_enrollment,expected_late", + [ + ("account-a", "account-z", True, False, 15), + ("account-z", "account-a", True, False, 15), + ("account-a", "account-z", False, False, 30), + ("account-a", "account-z", True, True, 15), + ], + ids=["continuous", "reverse-order", "producer-stops-control", "fresh-store-control"], +) +async def test_late_due_account_is_served_while_the_first_stays_continuously_due( + first, late, keep_waking_first, restart_at_enrollment, expected_late +): + async with progress_pool() as (pool, _): + store = PgReportingMaterializerStore(pool=pool) + await store.create_schema() + await wake(pool, first) + trace = [] + for turn in range(40): + if turn == 10: + await wake(pool, late) + if restart_at_enrollment: + store = PgReportingMaterializerStore(pool=pool) + # Waking from turn zero is essential: a gap before late enrollment + # could reset the old cursor and hide the regression. + if keep_waking_first: + await wake(pool, first) + if turn >= 10: + await wake(pool, late) + before = await served(pool) + await store.claim_materialization(keys=()) + after = await served(pool) + moved = [account for account in after if before.get(account) != after[account]] + assert len(moved) <= 1 + trace.append({"turn": turn, "served": moved}) + late_turns = [row["turn"] for row in trace if late in row["served"]] + print( + json.dumps( + { + "late_account_progress": { + "first": first, + "late": late, + "continuous_first": keep_waking_first, + "restart_at_enrollment": restart_at_enrollment, + "trace": trace, + "late_turns": late_turns, + "pool_size": 1, + } + } + ), + flush=True, + ) + assert late_turns and late_turns[0] <= 11, "late account was excluded from due sampling" + assert len(late_turns) == expected_late + + +def observe_samples(monkeypatch): + psycopg = pytest.importorskip("psycopg") + original = psycopg.AsyncConnection.execute + samples = [] + + async def execute(connection, query, *args, **kwargs): + result = await original(connection, query, *args, **kwargs) + if isinstance(query, str) and query.startswith( + "SELECT account_id,served_at::text FROM reporting_materializer_accounts WHERE due_at" + ): + assert query.endswith("LIMIT 16") + samples.append(result.rowcount) + return result + + monkeypatch.setattr(psycopg.AsyncConnection, "execute", execute) + return samples + + +async def bounded_claim(store, samples, *, keys=()): + before = len(samples) + result = await asyncio.wait_for(store.claim_materialization(keys=keys), 10) + sizes = samples[before:] + assert 1 <= len(sizes) <= 2 and 0 <= sum(sizes) <= 16, sizes + return result + + +@pytest.mark.parametrize("notifications", [False, True]) +async def test_more_than_two_busy_pages_progress_then_recover_after_unlock( + notifications, monkeypatch +): + async with progress_pool(autocommit=True) as (pool, schema): + from psycopg import AsyncConnection + + store = PgReportingMaterializerStore(pool=pool, notifications=notifications) + await store.create_schema() + busy = [f"busy-{number:02}" for number in range(33)] + available = "zz-available" + for account in [*busy, available]: + await wake(pool, account) + samples = observe_samples(monkeypatch) + async with ( + await AsyncConnection.connect( + os.environ["ADCP_PG_TEST_URL"], + options=f"-csearch_path={schema} -cstatement_timeout=15000", + autocommit=True, + ) as holder, + holder.transaction(), + ): + for account in busy: + await store._lock_account(holder, account) + before = await served(pool) + for _ in range(2): + assert (await bounded_claim(store, samples)).state == "idle" + assert await served(pool) == before + assert samples == [16, 16] + await bounded_claim(store, samples) + after = await served(pool) + assert after[available] != before[available] + assert {a: after[a] for a in busy} == {a: before[a] for a in busy} + assert samples[-1] == 2 + # A successful account turn restarts at the oldest durable rank; + # subsequent busy turns must still walk through all three pages. + for expected in (16, 16, 1): + await bounded_claim(store, samples) + assert samples[-1] == expected + await bounded_claim(store, samples) + assert samples[-2:] == [0, 16] + recovered = set() + for _ in range(len(busy) + 3): + before = await served(pool) + await bounded_claim(store, samples) + after = await served(pool) + recovered.update(a for a in busy if after[a] != before[a]) + if recovered == set(busy): + break + assert recovered == set(busy) + assert all(value != "-infinity" for value in (await served(pool)).values()) + + +@pytest.mark.parametrize("notifications", [False, True]) +async def test_commit_failure_rolls_back_rank_and_reservation_then_retries(notifications): + async with progress_pool(autocommit=True) as (pool, _): + store = PgReportingMaterializerStore(pool=pool, notifications=notifications) + await store.create_schema() + case = await durable_case(store) + before = await served(pool) + async with pool.connection() as connection: + # A deferred database error fails COMMIT after all reservation and + # scheduling statements executed. No SDK operation is replaced. + await connection.execute( + "CREATE FUNCTION progress_fail_commit() RETURNS trigger LANGUAGE plpgsql AS $$" + " BEGIN RAISE EXCEPTION 'private-progress-commit-fault'; END $$" + ) + await connection.execute( + "CREATE CONSTRAINT TRIGGER progress_fail_commit" + " AFTER UPDATE ON reporting_materializer_accounts" + " DEFERRABLE INITIALLY DEFERRED FOR EACH ROW" + " EXECUTE FUNCTION progress_fail_commit()" + ) + try: + with pytest.raises(ReportingWriterError) as error: + await store.claim_materialization(keys=case.keys) + assert error.value.failure.code == "RESOURCE_UNAVAILABLE" + assert "private-progress" not in str(error.value) + assert await served(pool) == before + async with pool.connection() as connection: + assert await ( + await connection.execute("SELECT count(*) FROM reporting_materializer_work") + ).fetchone() == (0,) + finally: + async with pool.connection() as connection: + await connection.execute( + "DROP TRIGGER progress_fail_commit ON reporting_materializer_accounts" + ) + await connection.execute("DROP FUNCTION progress_fail_commit()") + lease = await case.claim() + assert isinstance(lease, ReportingMaterializerLease) + assert lease.attempt.attempt == 1 + assert await store.renew_materialization(lease, lease_seconds=30) + assert await served(pool) != before + + +async def test_two_workers_skip_uncommitted_peer_and_reserve_distinct_accounts(monkeypatch): + async with progress_pool(autocommit=True) as (pool, schema): + from psycopg import AsyncConnection + from psycopg_pool import AsyncConnectionPool + + store = PgReportingMaterializerStore(pool=pool) + await store.create_schema() + first = await durable_case(store, account="account-a") + second = await durable_case(store, account="account-z") + async with AsyncConnectionPool( + os.environ["ADCP_PG_TEST_URL"], + kwargs={"options": f"-csearch_path={schema}", "autocommit": True}, + min_size=1, + max_size=1, + open=False, + ) as peer_pool: + await peer_pool.wait(timeout=10) + peer = PgReportingMaterializerStore(pool=peer_pool) + holding, release = asyncio.Event(), asyncio.Event() + original = AsyncConnection.execute + winner = None + + async def execute(connection, query, *args, **kwargs): + result = await original(connection, query, *args, **kwargs) + if ( + asyncio.current_task() is winner + and isinstance(query, str) + and query.startswith("UPDATE reporting_materializer_accounts SET served_at=") + ): + # Only pause after the real update, under the real locks. + holding.set() + await asyncio.wait_for(release.wait(), 10) + return result + + monkeypatch.setattr(AsyncConnection, "execute", execute) + winner = asyncio.create_task(store.claim_materialization(keys=first.keys)) + try: + await asyncio.wait_for(holding.wait(), 10) + other = await asyncio.wait_for(peer.claim_materialization(keys=second.keys), 10) + assert isinstance(other, ReportingMaterializerLease) + assert other.scope.principal.account_id == "account-z" + release.set() + selected = await asyncio.wait_for(winner, 10) + finally: + release.set() + if not winner.done(): + winner.cancel() + await asyncio.gather(winner, return_exceptions=True) + assert isinstance(selected, ReportingMaterializerLease) + assert selected.scope.principal.account_id == "account-a" + assert selected.request.external_id != other.request.external_id + await store.authorize_materialization(selected) + await peer.authorize_materialization(other) + assert await store.renew_materialization(selected, lease_seconds=30) + assert await peer.renew_materialization(other, lease_seconds=30) + for current in (store, peer): + assert not isinstance( + await current.claim_materialization(keys=first.keys), ReportingMaterializerLease + ) + async with pool.connection() as connection: + assert await ( + await connection.execute("SELECT count(*) FROM reporting_materializer_work") + ).fetchone() == (2,) diff --git a/tests/conformance/reporting/test_reporting_production_late_accounts.py b/tests/conformance/reporting/test_reporting_production_late_accounts.py new file mode 100644 index 000000000..c68246c10 --- /dev/null +++ b/tests/conformance/reporting/test_reporting_production_late_accounts.py @@ -0,0 +1,433 @@ +"""Typed public enrollment, autonomous late-account delivery and durable restart.""" + +import asyncio +import hashlib +import json +import os +import signal +import socket +import sqlite3 +import subprocess +import sys +import time +from contextlib import asynccontextmanager +from pathlib import Path + +import pytest +import rfc8785 +from pydantic import TypeAdapter + +from adcp import ADCPClient, AgentConfig +from adcp.reporting import ( + ExpectedReportingPeriod, + ReportingObservation, + load_reporting_ledger, + reconcile_reporting, +) +from adcp.reporting.materializer import ReportingWriterCapability, reference_digest +from adcp.types import ( + GetAdcpCapabilitiesRequest, + GetMediaBuyDeliveryRequest, + GetReportingStatusRequest, + ReportingCanonicalContentDigest, + ReportingControlTotal, + SyncAccountsRequest, +) + +from ._late_account_support import ACCOUNTS, rows_for, verifier_for +from .test_reporting_materializer_progress import progress_pool, served +from .test_reporting_production_scope import assert_access_denied + + +@asynccontextmanager +async def running_server(root, schema, index, *, notifications): + fixture_root = Path(__file__).resolve().parents[3] + launcher = ( + "import sys; sys.path.insert(0,sys.argv.pop(1)); " + "from tests.conformance.reporting._late_account_server import main; main()" + ) + (root / "ready.json").unlink(missing_ok=True) + (root / "stopped.json").unlink(missing_ok=True) + with socket.socket() as listener: + listener.bind(("127.0.0.1", 0)) + port = listener.getsockname()[1] + command = [ + sys.executable, + "-I", + "-c", + launcher, + str(fixture_root), + "--port", + str(port), + "--root", + str(root), + "--schema", + schema, + ] + if notifications: + command.append("--notifications") + with ( + (root / f"server-{index}.stdout").open("xb") as output, + (root / f"server-{index}.stderr").open("xb") as errors, + ): + child = subprocess.Popen( + command, + cwd=root, + stdout=output, + stderr=errors, + start_new_session=True, + ) + try: + for _ in range(1800): + assert child.poll() is None, (root / f"server-{index}.stderr").read_text() + if (root / "ready.json").exists(): + try: + _, writer = await asyncio.open_connection("127.0.0.1", port) + except OSError: + pass + else: + writer.close() + await writer.wait_closed() + break + await asyncio.sleep(0.05) + else: + raise AssertionError("public production startup exceeded 90 seconds") + ready = json.loads((root / "ready.json").read_text()) + (root / f"ready-{index}.json").write_text(json.dumps(ready)) + yield f"http://127.0.0.1:{port}/mcp/", ready + finally: + if child.poll() is None: + os.killpg(child.pid, signal.SIGTERM) + try: + await asyncio.to_thread(child.wait, 25) + except subprocess.TimeoutExpired: + os.killpg(child.pid, signal.SIGKILL) + await asyncio.to_thread(child.wait, 5) + cleanup = { + "pid": child.pid, + "exit": child.returncode, + "reaped": child.poll() is not None, + } + if (root / "stopped.json").exists(): + cleanup.update(json.loads((root / "stopped.json").read_text())) + (root / f"cleanup-{index}.json").write_text(json.dumps(cleanup)) + assert cleanup.get("stopped") and cleanup["destination_sessions_closed"], cleanup + + +def client_for(uri, account): + return ADCPClient( + AgentConfig( + id="late-" + account, + agent_uri=uri, + protocol="mcp", + auth_token=account + "-test-token", + auth_header="Authorization", + auth_type="bearer", + ), + adcp_version="3.2.0-rc.4", + ) + + +def data(result): + assert result.success and result.data is not None, result + return result.data.model_dump(mode="json", exclude_none=True) + + +async def delivered_and_reconciled(client, root, ready, account, *, replay=False, period=None): + currency, count = ACCOUNTS[account] + capability = ReportingWriterCapability( + "warehouse_materialization", + "fixture-sql", + "jsonl", + "canonical_digest", + "destination", + "immutable_location", + "sha256", + "conditional_create", + ) + verifier = verifier_for(currency, capability) + selectors = {"account": {"account_id": account}, "view": "periods"} + if period is not None: + selectors["period"] = period + request = GetReportingStatusRequest.model_validate(selectors) + started = time.monotonic() + ledger = None + calls = 0 + # Every readiness observation is an actual typed RPC. The scheduler and + # producer use real clocks; no worker is stepped, stopped or manually invoked. + while time.monotonic() - started < 45 and calls < 150: + ledger = await load_reporting_ledger(client, request) + calls += 1 + if any( + str(getattr(m.status, "value", m.status)) == "delivered" + for m in ledger.materializations + ): + break + else: + raise AssertionError({"account": account, "materialization_missing_after_calls": calls}) + readiness_seconds = time.monotonic() - started + if period is None: + # Catch-up production is deliberately active throughout this test. Pin + # a publicly delivered period, rather than assuming which candidate the + # materializer chooses first among that account's outstanding periods. + delivered = next( + m + for m in ledger.materializations + if str(getattr(m.status, "value", m.status)) == "delivered" + ) + obligation = next( + o + for o in ledger.obligations + if o.reporting_obligation_id == delivered.reporting_obligation_id + ) + period = {name: getattr(obligation.period, name).isoformat() for name in ("start", "end")} + selectors["period"] = period + request = GetReportingStatusRequest.model_validate(selectors) + ledger = await load_reporting_ledger(client, request) + assert len(ledger.obligations) == len(ledger.revisions) == len(ledger.materializations) == 1 + revision = ledger.revisions[0] + materialization = ledger.materializations[0] + assert revision.account_id == ledger.obligations[0].account_id == account + assert revision.row_count == count + params = { + "account": {"account_id": account}, + "reporting_revision_id": revision.reporting_revision_id, + "pagination": {"max_results": 100}, + } + exact, binding, pages = [], None, 0 + while True: + page = data( + await client.get_media_buy_delivery(GetMediaBuyDeliveryRequest.model_validate(params)) + ) + if binding is None: + binding = page["reporting_revision_binding"] + else: + assert binding == page["reporting_revision_binding"] + exact.extend(page["reporting_rows"]) + pages += 1 + if not page["pagination"]["has_more"]: + break + params["pagination"]["cursor"] = page["pagination"]["cursor"] + assert pages < 10 + assert exact == rows_for(account) + assert pages == (6 if count == 503 else 1) + digest_input = {k: binding[k] for k in ("reporting_revision_id", "row_count", "control_totals")} + digest_input["reporting_rows"] = exact + digest = hashlib.sha256(rfc8785.dumps(digest_input)).hexdigest() + assert digest == binding["content_sha256"] == revision.revision_content_sha256 + + async def inspect(context): + with sqlite3.connect(root / "destination.sqlite") as connection: + records = [ + json.loads(row[0]) for row in connection.execute("SELECT content FROM artifacts") + ] + matched = [ + record + for record in records + if record["revision"] == context.revision.reporting_revision_id + and record["resource"]["location"] == context.materialization.resource.location + ] + assert len(matched) == 1 + rows = [json.loads(row) for row in matched[0]["rows"]] + assert rows == exact + _, totals = verifier.canonicalize(rows) + return ReportingObservation( + len(rows), + [TypeAdapter(ReportingControlTotal).validate_python(t.to_wire()) for t in totals], + ReportingCanonicalContentDigest.model_validate( + reference_digest(verifier, rows).to_wire() + ), + ) + + template = ready["templates"][account] + expected = [ + ExpectedReportingPeriod( + "shared-config", + 1, + template["report_definition_id"], + "billing", + template["reporting_profile"], + ("shared-media-buy",), + period["start"], + period["end"], + ) + ] + outcome = await reconcile_reporting( + client, request, inspect, expected_periods=expected, inspection_retry_backoff_seconds=0 + ) + assert outcome.definitive, [(o.definitive, o.reasons) for o in outcome.obligations] + assert len(outcome.submitted_receipts) == (0 if replay else 1) + repeat = await reconcile_reporting( + client, request, inspect, expected_periods=expected, inspection_retry_backoff_seconds=0 + ) + assert repeat.definitive and not repeat.submitted_receipts + other = "eur" if account == "usd" else "usd" + assert_access_denied( + await client.get_reporting_status( + GetReportingStatusRequest.model_validate( + { + "account": {"account_id": other}, + "view": "periods", + } + ) + ) + ) + assert_access_denied( + await client.get_media_buy_delivery( + GetMediaBuyDeliveryRequest.model_validate( + { + "account": {"account_id": other}, + "reporting_revision_id": revision.reporting_revision_id, + } + ) + ) + ) + return { + "account": account, + "currency": currency, + "rows": count, + "pages": pages, + "revision": revision.reporting_revision_id, + "digest": digest, + "receipt_ids": [receipt.reporting_receipt_id for receipt in repeat.ledger.receipts], + "materialization_id": materialization.reporting_materialization_id, + "period": period, + "readiness_calls": calls, + "readiness_seconds": round(readiness_seconds, 3), + "completed_seconds": round(time.monotonic() - started, 3), + } + + +async def ongoing_first_turns(pool, account): + """Observe real committed turns after first-account RPCs have finished. + + A one-period fixture can settle or encounter a busy-lock wrap during those + RPCs, accidentally passing with the old cursor. Catch-up publications keep + pending work due. These plain MVCC reads neither take account locks nor + wake, advance or execute either SDK worker. + """ + observed = [] + deadline = time.monotonic() + 30 + previous = (await served(pool))[account] + while time.monotonic() < deadline: + async with pool.connection() as connection: + row = await ( + await connection.execute( + "SELECT served_at::text,due_at<=clock_timestamp()," + " (SELECT count(*) FROM reporting_materializer_candidates c" + " WHERE c.account_id=a.account_id AND c.due_at<=clock_timestamp())" + " FROM reporting_materializer_accounts a WHERE account_id=%s", + (account,), + ) + ).fetchone() + if row[0] != previous: + observed.append({"served_at": row[0], "due": row[1], "due_candidates": row[2]}) + previous = row[0] + if len(observed) >= 3 and all( + r["due"] and r["due_candidates"] >= 2 for r in observed[-3:] + ): + return observed[-3:] + await asyncio.sleep(0.02) + raise AssertionError({"continuous_first_work_not_established": observed}) + + +@pytest.mark.parametrize("first", ["usd", "eur"], ids=["late-sorts-first", "late-sorts-last"]) +@pytest.mark.parametrize("notifications", [False, True]) +async def test_late_account_progresses_via_typed_public_support_and_survives_restart( + first, notifications, tmp_path +): + order = (first, "eur" if first == "usd" else "usd") + async with progress_pool(autocommit=True) as (pool, schema): + results = [] + continuing_work = None + for index in range(2): + async with running_server(tmp_path, schema, index, notifications=notifications) as ( + uri, + ready, + ): + assert ready["pool_size"] == 1 + assert ready["initial_configurations"] == (0 if index == 0 else 2) + current = {} + for account in order: + if index == 0 and account != first: + continuing_work = await ongoing_first_turns(pool, first) + async with client_for(uri, account) as client: + caps = data( + await client.get_adcp_capabilities(GetAdcpCapabilitiesRequest()) + ) + assert caps["media_buy"]["reporting_delivery"]["managed_delivery"] + if index == 0: + # The first remains active throughout late admission; + # no worker restart, due-gap workaround or manual turn. + first_before = (await served(pool)).get(first) + onboarding = SyncAccountsRequest.model_validate( + { + "idempotency_key": "late-account-" + account, + "accounts": [ + { + "account": {"account_id": account}, + "reporting_delivery_configs": [ + ready["templates"][account] + ], + } + ], + } + ) + admitted = data(await client.sync_accounts(onboarding)) + assert admitted["accounts"][0]["account_id"] == account + assert ( + admitted["accounts"][0]["reporting_delivery_configs"][0]["state"] + == "ready" + ) + current[account] = await delivered_and_reconciled( + client, + tmp_path, + ready, + account, + replay=bool(index), + period=results[0][account]["period"] if index else None, + ) + if index == 0 and account != first: + assert (await served(pool))[first] != first_before + results.append(current) + async with pool.connection() as connection: + assert await ( + await connection.execute( + "SELECT count(*) FROM pg_stat_activity WHERE application_name=%s", + (schema,), + ) + ).fetchone() == (0,) + for account in order: + for field in ("revision", "digest", "receipt_ids", "materialization_id"): + assert results[0][account][field] == results[1][account][field] + assert results[0]["usd"]["revision"] != results[0]["eur"]["revision"] + assert results[0]["usd"]["materialization_id"] != results[0]["eur"]["materialization_id"] + wire = (tmp_path / "wire.jsonl").read_text() + assert "-test-token" not in wire + for line in wire.splitlines(): + request = json.loads(json.loads(line)["request_utf8"]) + if request.get("method") != "tools/call": + continue + task, arguments = request["params"]["name"], request["params"]["arguments"] + if task == "sync_accounts": + assert arguments["accounts"][0]["reporting_delivery_configs"][0]["scope"] == { + "media_buy_ids": ["shared-media-buy"], + } + if task == "get_media_buy_delivery": + assert "include_package_daily_breakdown" not in arguments + assert "include_window_breakdown" not in arguments + print( + json.dumps( + { + "public_late_account_progress": { + "order": order, + "notifications": notifications, + "runs": results, + "typed_onboarding_and_exact_reads": True, + "autonomous_worker": True, + "observed_due_first_turns_before_late_admission": continuing_work, + } + } + ), + flush=True, + )