From f3288dbd2824aebf0fbfdc57e9a5cc6f670b542e Mon Sep 17 00:00:00 2001 From: Albrecht Degering Date: Wed, 29 Jul 2026 19:38:34 +0200 Subject: [PATCH] feat: harden datasource materialization payloads --- src/govoplan_datasources/backend/db/models.py | 102 ++++ src/govoplan_datasources/backend/manifest.py | 4 + ...a2b8c3e7_v0114_materialization_payloads.py | 297 +++++++++ src/govoplan_datasources/backend/payloads.py | 374 ++++++++++++ src/govoplan_datasources/backend/service.py | 569 ++++++++++++------ src/govoplan_datasources/backend/tabular.py | 34 +- tests/test_lifecycle.py | 133 ++++ tests/test_migrations.py | 4 +- ...st_postgres_materialization_concurrency.py | 122 ++++ 9 files changed, 1460 insertions(+), 179 deletions(-) create mode 100644 src/govoplan_datasources/backend/migrations/versions/d5f0a2b8c3e7_v0114_materialization_payloads.py create mode 100644 src/govoplan_datasources/backend/payloads.py create mode 100644 tests/test_postgres_materialization_concurrency.py diff --git a/src/govoplan_datasources/backend/db/models.py b/src/govoplan_datasources/backend/db/models.py index 0103045..76637d4 100644 --- a/src/govoplan_datasources/backend/db/models.py +++ b/src/govoplan_datasources/backend/db/models.py @@ -122,6 +122,17 @@ class DatasourceMaterializationRecord(Base, TimestampMixin): default=list, nullable=False, ) + payload_id: Mapped[str | None] = mapped_column( + ForeignKey("datasource_payloads.id", ondelete="RESTRICT"), + nullable=True, + index=True, + ) + payload_checksum: Mapped[str | None] = mapped_column( + String(64), + nullable=True, + ) + # Kept nullable-in-practice for rolling upgrades. New materializations use + # immutable payload rows and leave this legacy field empty. rows: Mapped[list[dict[str, Any]]] = mapped_column(JSON, default=list, nullable=False) fingerprint: Mapped[str] = mapped_column(String(64), nullable=False, index=True) row_count: Mapped[int] = mapped_column(Integer, nullable=False) @@ -147,6 +158,95 @@ class DatasourceMaterializationRecord(Base, TimestampMixin): created_by: Mapped[str | None] = mapped_column(String(255), nullable=True, index=True) datasource: Mapped[DatasourceRecord] = relationship(back_populates="materializations") + payload: Mapped["DatasourcePayloadRecord | None"] = relationship() + + +class DatasourcePayloadRecord(Base, TimestampMixin): + __tablename__ = "datasource_payloads" + __table_args__ = ( + Index( + "ix_datasource_payloads_tenant_backend_state", + "tenant_id", + "backend", + "state", + ), + Index( + "ix_datasource_payloads_checksum", + "tenant_id", + "checksum", + ), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid) + tenant_id: Mapped[str] = mapped_column(String(36), nullable=False, index=True) + backend: Mapped[str] = mapped_column( + String(30), + default="database_rows", + nullable=False, + ) + state: Mapped[str] = mapped_column( + String(30), + default="staging", + nullable=False, + index=True, + ) + locator: Mapped[str | None] = mapped_column(String(1000), nullable=True) + media_type: Mapped[str] = mapped_column( + String(150), + default="application/x-ndjson", + nullable=False, + ) + checksum: Mapped[str] = mapped_column(String(64), nullable=False) + row_count: Mapped[int] = mapped_column(Integer, nullable=False) + byte_count: Mapped[int] = mapped_column(Integer, nullable=False) + checkpoint_: Mapped[dict[str, Any]] = mapped_column( + "checkpoint", + JSON, + default=dict, + nullable=False, + ) + metadata_: Mapped[dict[str, Any]] = mapped_column( + "metadata", + JSON, + default=dict, + nullable=False, + ) + created_by: Mapped[str | None] = mapped_column( + String(255), + nullable=True, + index=True, + ) + + rows: Mapped[list["DatasourcePayloadRowRecord"]] = relationship( + back_populates="payload", + cascade="all, delete-orphan", + order_by="DatasourcePayloadRowRecord.row_index", + ) + + +class DatasourcePayloadRowRecord(Base): + __tablename__ = "datasource_payload_rows" + __table_args__ = ( + Index( + "ix_datasource_payload_rows_payload_window", + "payload_id", + "row_index", + ), + ) + + payload_id: Mapped[str] = mapped_column( + ForeignKey("datasource_payloads.id", ondelete="CASCADE"), + primary_key=True, + ) + row_index: Mapped[int] = mapped_column(Integer, primary_key=True) + row_: Mapped[dict[str, Any]] = mapped_column( + "row", + JSON, + nullable=False, + ) + checksum: Mapped[str] = mapped_column(String(64), nullable=False) + + payload: Mapped[DatasourcePayloadRecord] = relationship(back_populates="rows") class DatasourceStageRecord(Base, TimestampMixin): @@ -268,6 +368,8 @@ class DatasourcePublicationRecord(Base, TimestampMixin): __all__ = [ "DatasourceMaterializationRecord", + "DatasourcePayloadRecord", + "DatasourcePayloadRowRecord", "DatasourcePublicationRecord", "DatasourceRecord", "DatasourceStageRecord", diff --git a/src/govoplan_datasources/backend/manifest.py b/src/govoplan_datasources/backend/manifest.py index 92fcc8a..7572774 100644 --- a/src/govoplan_datasources/backend/manifest.py +++ b/src/govoplan_datasources/backend/manifest.py @@ -244,6 +244,8 @@ manifest = ModuleManifest( datasource_models.DatasourcePublicationRecord, datasource_models.DatasourceStageRecord, datasource_models.DatasourceMaterializationRecord, + datasource_models.DatasourcePayloadRowRecord, + datasource_models.DatasourcePayloadRecord, datasource_models.DatasourceRecord, label="Datasources", ), @@ -256,6 +258,8 @@ manifest = ModuleManifest( persistent_table_uninstall_guard( datasource_models.DatasourceRecord, datasource_models.DatasourceMaterializationRecord, + datasource_models.DatasourcePayloadRecord, + datasource_models.DatasourcePayloadRowRecord, datasource_models.DatasourceStageRecord, datasource_models.DatasourcePublicationRecord, label="Datasources", diff --git a/src/govoplan_datasources/backend/migrations/versions/d5f0a2b8c3e7_v0114_materialization_payloads.py b/src/govoplan_datasources/backend/migrations/versions/d5f0a2b8c3e7_v0114_materialization_payloads.py new file mode 100644 index 0000000..797afad --- /dev/null +++ b/src/govoplan_datasources/backend/migrations/versions/d5f0a2b8c3e7_v0114_materialization_payloads.py @@ -0,0 +1,297 @@ +"""v0.1.14 immutable datasource materialization payloads + +Revision ID: d5f0a2b8c3e7 +Revises: c4e9f1a7d2b6 +Create Date: 2026-07-29 00:00:00.000000 +""" +from __future__ import annotations + +import hashlib +import json +import uuid + +from alembic import op +import sqlalchemy as sa + + +revision = "d5f0a2b8c3e7" +down_revision = "c4e9f1a7d2b6" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "datasource_payloads", + sa.Column("id", sa.String(length=36), nullable=False), + sa.Column("tenant_id", sa.String(length=36), nullable=False), + sa.Column("backend", sa.String(length=30), nullable=False), + sa.Column("state", sa.String(length=30), nullable=False), + sa.Column("locator", sa.String(length=1000), nullable=True), + sa.Column("media_type", sa.String(length=150), nullable=False), + sa.Column("checksum", sa.String(length=64), nullable=False), + sa.Column("row_count", sa.Integer(), nullable=False), + sa.Column("byte_count", sa.Integer(), nullable=False), + sa.Column("checkpoint", sa.JSON(), nullable=False), + sa.Column("metadata", sa.JSON(), nullable=False), + sa.Column("created_by", sa.String(length=255), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id", name=op.f("pk_datasource_payloads")), + ) + op.create_index( + op.f("ix_datasource_payloads_tenant_id"), + "datasource_payloads", + ["tenant_id"], + unique=False, + ) + op.create_index( + op.f("ix_datasource_payloads_state"), + "datasource_payloads", + ["state"], + unique=False, + ) + op.create_index( + op.f("ix_datasource_payloads_created_by"), + "datasource_payloads", + ["created_by"], + unique=False, + ) + op.create_index( + "ix_datasource_payloads_tenant_backend_state", + "datasource_payloads", + ["tenant_id", "backend", "state"], + unique=False, + ) + op.create_index( + "ix_datasource_payloads_checksum", + "datasource_payloads", + ["tenant_id", "checksum"], + unique=False, + ) + op.create_table( + "datasource_payload_rows", + sa.Column("payload_id", sa.String(length=36), nullable=False), + sa.Column("row_index", sa.Integer(), nullable=False), + sa.Column("row", sa.JSON(), nullable=False), + sa.Column("checksum", sa.String(length=64), nullable=False), + sa.ForeignKeyConstraint( + ["payload_id"], + ["datasource_payloads.id"], + name=op.f( + "fk_datasource_payload_rows_payload_id_datasource_payloads" + ), + ondelete="CASCADE", + ), + sa.PrimaryKeyConstraint( + "payload_id", + "row_index", + name=op.f("pk_datasource_payload_rows"), + ), + ) + op.create_index( + "ix_datasource_payload_rows_payload_window", + "datasource_payload_rows", + ["payload_id", "row_index"], + unique=False, + ) + with op.batch_alter_table( + "datasource_materializations", + schema=None, + ) as batch_op: + batch_op.add_column( + sa.Column("payload_id", sa.String(length=36), nullable=True) + ) + batch_op.add_column( + sa.Column("payload_checksum", sa.String(length=64), nullable=True) + ) + batch_op.create_foreign_key( + op.f( + "fk_datasource_materializations_payload_id_datasource_payloads" + ), + "datasource_payloads", + ["payload_id"], + ["id"], + ondelete="RESTRICT", + ) + batch_op.create_index( + op.f("ix_datasource_materializations_payload_id"), + ["payload_id"], + unique=False, + ) + _backfill_legacy_payloads() + + +def downgrade() -> None: + _restore_legacy_rows() + with op.batch_alter_table( + "datasource_materializations", + schema=None, + ) as batch_op: + batch_op.drop_index( + op.f("ix_datasource_materializations_payload_id") + ) + batch_op.drop_constraint( + op.f( + "fk_datasource_materializations_payload_id_datasource_payloads" + ), + type_="foreignkey", + ) + batch_op.drop_column("payload_checksum") + batch_op.drop_column("payload_id") + op.drop_index( + "ix_datasource_payload_rows_payload_window", + table_name="datasource_payload_rows", + ) + op.drop_table("datasource_payload_rows") + op.drop_index( + "ix_datasource_payloads_checksum", + table_name="datasource_payloads", + ) + op.drop_index( + "ix_datasource_payloads_tenant_backend_state", + table_name="datasource_payloads", + ) + op.drop_index( + op.f("ix_datasource_payloads_created_by"), + table_name="datasource_payloads", + ) + op.drop_index( + op.f("ix_datasource_payloads_state"), + table_name="datasource_payloads", + ) + op.drop_index( + op.f("ix_datasource_payloads_tenant_id"), + table_name="datasource_payloads", + ) + op.drop_table("datasource_payloads") + + +def _backfill_legacy_payloads() -> None: + connection = op.get_bind() + metadata = sa.MetaData() + materializations = sa.Table( + "datasource_materializations", + metadata, + autoload_with=connection, + ) + payloads = sa.Table( + "datasource_payloads", + metadata, + autoload_with=connection, + ) + payload_rows = sa.Table( + "datasource_payload_rows", + metadata, + autoload_with=connection, + ) + rows = connection.execute( + sa.select( + materializations.c.id, + materializations.c.tenant_id, + materializations.c.rows, + materializations.c.created_by, + materializations.c.created_at, + materializations.c.updated_at, + ).where(materializations.c.payload_id.is_(None)) + ).mappings() + for materialization in rows: + row_payload = list(materialization["rows"] or []) + encoded = _encoded_rows(row_payload) + payload_id = str(uuid.uuid4()) + checksum = hashlib.sha256(encoded).hexdigest() + connection.execute( + payloads.insert().values( + id=payload_id, + tenant_id=materialization["tenant_id"], + backend="database_rows", + state="ready", + locator=None, + media_type="application/x-ndjson", + checksum=checksum, + row_count=len(row_payload), + byte_count=len(encoded), + checkpoint={}, + metadata={"migrated_from": materialization["id"]}, + created_by=materialization["created_by"], + created_at=materialization["created_at"], + updated_at=materialization["updated_at"], + ) + ) + if row_payload: + connection.execute( + payload_rows.insert(), + [ + { + "payload_id": payload_id, + "row_index": index, + "row": row, + "checksum": hashlib.sha256( + json.dumps( + row, + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") + ).hexdigest(), + } + for index, row in enumerate(row_payload) + ], + ) + connection.execute( + materializations.update() + .where(materializations.c.id == materialization["id"]) + .values( + payload_id=payload_id, + payload_checksum=checksum, + row_count=len(row_payload), + byte_count=len(encoded), + rows=[], + ) + ) + + +def _restore_legacy_rows() -> None: + connection = op.get_bind() + metadata = sa.MetaData() + materializations = sa.Table( + "datasource_materializations", + metadata, + autoload_with=connection, + ) + payload_rows = sa.Table( + "datasource_payload_rows", + metadata, + autoload_with=connection, + ) + rows = connection.execute( + sa.select( + materializations.c.id, + materializations.c.payload_id, + ).where(materializations.c.payload_id.is_not(None)) + ).mappings() + for materialization in rows: + payload = list( + connection.execute( + sa.select(payload_rows.c.row) + .where( + payload_rows.c.payload_id + == materialization["payload_id"] + ) + .order_by(payload_rows.c.row_index.asc()) + ).scalars() + ) + connection.execute( + materializations.update() + .where(materializations.c.id == materialization["id"]) + .values(rows=payload) + ) + + +def _encoded_rows(rows: list[dict[str, object]]) -> bytes: + return json.dumps( + rows, + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") diff --git a/src/govoplan_datasources/backend/payloads.py b/src/govoplan_datasources/backend/payloads.py new file mode 100644 index 0000000..0570549 --- /dev/null +++ b/src/govoplan_datasources/backend/payloads.py @@ -0,0 +1,374 @@ +from __future__ import annotations + +import hashlib +import json +from collections.abc import Mapping, Sequence +from typing import Protocol, runtime_checkable + +from sqlalchemy import delete, func, insert, select +from sqlalchemy.orm import Session + +from govoplan_core.core.datasources import ( + DatasourceUnavailableError, + DatasourceValidationError, +) +from govoplan_datasources.backend.db.models import ( + DatasourceMaterializationRecord, + DatasourcePayloadRecord, + DatasourcePayloadRowRecord, +) +from govoplan_datasources.backend.tabular import ( + encoded_size, + payload_checksum, + payload_row_checksum, +) + +DATABASE_ROWS_BACKEND = "database_rows" +PAYLOAD_INSERT_BATCH_SIZE = 500 + + +@runtime_checkable +class DatasourcePayloadBackend(Protocol): + """Extension point for object, file, and streaming-checkpoint payloads.""" + + backend: str + + def read_rows( + self, + session: Session, + payload: DatasourcePayloadRecord, + *, + offset: int, + limit: int, + ) -> Sequence[Mapping[str, object]]: ... + + def verify( + self, + session: Session, + payload: DatasourcePayloadRecord, + ) -> None: ... + + def delete( + self, + session: Session, + payload: DatasourcePayloadRecord, + ) -> None: ... + + +class DatabaseRowsPayloadBackend: + backend = DATABASE_ROWS_BACKEND + + def read_rows( + self, + session: Session, + payload: DatasourcePayloadRecord, + *, + offset: int, + limit: int, + ) -> Sequence[Mapping[str, object]]: + statement = ( + select( + DatasourcePayloadRowRecord.row_, + DatasourcePayloadRowRecord.checksum, + ) + .where(DatasourcePayloadRowRecord.payload_id == payload.id) + .order_by(DatasourcePayloadRowRecord.row_index.asc()) + .offset(max(0, int(offset))) + .limit(max(1, int(limit))) + ) + rows: list[dict[str, object]] = [] + for row, checksum in session.execute(statement): + normalized = dict(row) + if payload_row_checksum(normalized) != checksum: + raise DatasourceUnavailableError( + "Datasource payload row checksum verification failed." + ) + rows.append(normalized) + return tuple(rows) + + def verify( + self, + session: Session, + payload: DatasourcePayloadRecord, + ) -> None: + persisted_count = int( + session.scalar( + select(func.count()) + .select_from(DatasourcePayloadRowRecord) + .where(DatasourcePayloadRowRecord.payload_id == payload.id) + ) + or 0 + ) + if persisted_count != payload.row_count: + raise DatasourceUnavailableError( + "Datasource payload row count does not match its metadata." + ) + + def delete( + self, + session: Session, + payload: DatasourcePayloadRecord, + ) -> None: + session.execute( + delete(DatasourcePayloadRowRecord).where( + DatasourcePayloadRowRecord.payload_id == payload.id + ) + ) + + +class PayloadBackendRegistry: + def __init__( + self, + backends: Sequence[DatasourcePayloadBackend] = (), + ) -> None: + self._backends: dict[str, DatasourcePayloadBackend] = { + DATABASE_ROWS_BACKEND: DatabaseRowsPayloadBackend(), + } + for backend in backends: + self.register(backend) + + def register(self, backend: DatasourcePayloadBackend) -> None: + name = str(backend.backend).strip() + if not name: + raise ValueError("Datasource payload backends require a name.") + self._backends[name] = backend + + def require(self, name: str) -> DatasourcePayloadBackend: + backend = self._backends.get(name) + if backend is None: + raise DatasourceUnavailableError( + f"Datasource payload backend {name!r} is not available." + ) + return backend + + +def create_database_rows_payload( + session: Session, + *, + tenant_id: str, + rows: Sequence[Mapping[str, object]], + actor_id: str | None, + metadata: Mapping[str, object] | None = None, +) -> DatasourcePayloadRecord: + row_payload = tuple(dict(row) for row in rows) + checksum = payload_checksum(row_payload) + byte_count = encoded_size(row_payload) + payload = DatasourcePayloadRecord( + tenant_id=tenant_id, + backend=DATABASE_ROWS_BACKEND, + state="staging", + checksum=checksum, + row_count=len(row_payload), + byte_count=byte_count, + checkpoint_={}, + metadata_=dict(metadata or {}), + created_by=actor_id, + ) + session.add(payload) + session.flush() + for start in range(0, len(row_payload), PAYLOAD_INSERT_BATCH_SIZE): + batch = row_payload[start : start + PAYLOAD_INSERT_BATCH_SIZE] + session.execute( + insert(DatasourcePayloadRowRecord), + [ + { + "payload_id": payload.id, + "row_index": start + index, + "row_": row, + "checksum": payload_row_checksum(row), + } + for index, row in enumerate(batch) + ], + ) + payload.state = "ready" + session.add(payload) + session.flush() + return payload + + +def create_external_payload_reference( + session: Session, + *, + tenant_id: str, + backend: str, + locator: str, + checksum: str, + row_count: int, + byte_count: int, + actor_id: str | None, + media_type: str = "application/octet-stream", + checkpoint: Mapping[str, object] | None = None, + metadata: Mapping[str, object] | None = None, +) -> DatasourcePayloadRecord: + backend_name = backend.strip() + locator_value = locator.strip() + if backend_name == DATABASE_ROWS_BACKEND or not backend_name: + raise DatasourceValidationError( + "External payload references require a non-database backend." + ) + if not locator_value or len(locator_value) > 1000: + raise DatasourceValidationError( + "External payload references require a locator of at most 1000 characters." + ) + try: + valid_checksum = len(checksum) == 64 and int(checksum, 16) >= 0 + except ValueError: + valid_checksum = False + if not valid_checksum: + raise DatasourceValidationError( + "External payload references require a SHA-256 checksum." + ) + if row_count < 0 or byte_count < 0: + raise DatasourceValidationError( + "External payload sizes cannot be negative." + ) + payload = DatasourcePayloadRecord( + tenant_id=tenant_id, + backend=backend_name, + state="ready", + locator=locator_value, + media_type=media_type.strip() or "application/octet-stream", + checksum=checksum.casefold(), + row_count=row_count, + byte_count=byte_count, + checkpoint_=dict(checkpoint or {}), + metadata_=dict(metadata or {}), + created_by=actor_id, + ) + session.add(payload) + session.flush() + return payload + + +def payload_for_materialization( + session: Session, + materialization: DatasourceMaterializationRecord, +) -> DatasourcePayloadRecord | None: + if materialization.payload_id is None: + return None + payload = session.get(DatasourcePayloadRecord, materialization.payload_id) + if ( + payload is None + or payload.tenant_id != materialization.tenant_id + or payload.state != "ready" + ): + raise DatasourceUnavailableError( + "Datasource materialization payload is unavailable." + ) + if ( + materialization.payload_checksum != payload.checksum + or materialization.row_count != payload.row_count + or materialization.byte_count != payload.byte_count + ): + raise DatasourceUnavailableError( + "Datasource materialization metadata does not match its payload." + ) + return payload + + +def validate_payload_size( + payload: DatasourcePayloadRecord, + *, + expected_byte_count: int, +) -> None: + if payload.byte_count != expected_byte_count: + raise DatasourceValidationError( + "Datasource payload byte count changed while it was being materialized." + ) + + +def verify_payload_integrity( + session: Session, + payload: DatasourcePayloadRecord, + *, + registry: PayloadBackendRegistry | None = None, +) -> None: + backends = registry or PayloadBackendRegistry() + backend = backends.require(payload.backend) + backend.verify(session, payload) + if payload.backend != DATABASE_ROWS_BACKEND: + return + statement = ( + select(DatasourcePayloadRowRecord.row_) + .where(DatasourcePayloadRowRecord.payload_id == payload.id) + .order_by(DatasourcePayloadRowRecord.row_index.asc()) + ) + digest = hashlib.sha256() + digest.update(b"[") + byte_count = 2 + row_count = 0 + for row in session.scalars(statement).yield_per(500): + encoded = json.dumps( + dict(row), + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") + if row_count: + digest.update(b",") + byte_count += 1 + digest.update(encoded) + byte_count += len(encoded) + row_count += 1 + digest.update(b"]") + if ( + digest.hexdigest() != payload.checksum + or byte_count != payload.byte_count + or row_count != payload.row_count + ): + raise DatasourceUnavailableError( + "Datasource payload checksum or byte count verification failed." + ) + + +def mark_unreferenced_payload_for_deletion( + session: Session, + payload: DatasourcePayloadRecord, +) -> bool: + references = int( + session.scalar( + select(func.count()) + .select_from(DatasourceMaterializationRecord) + .where( + DatasourceMaterializationRecord.payload_id == payload.id + ) + ) + or 0 + ) + if references: + return False + payload.state = "deleting" + session.add(payload) + session.flush() + return True + + +def finalize_payload_deletion( + session: Session, + payload: DatasourcePayloadRecord, + *, + registry: PayloadBackendRegistry | None = None, +) -> None: + if payload.state != "deleting": + raise DatasourceValidationError( + "Payload deletion must be staged before it is finalized." + ) + backends = registry or PayloadBackendRegistry() + backends.require(payload.backend).delete(session, payload) + session.delete(payload) + session.flush() + + +__all__ = [ + "DATABASE_ROWS_BACKEND", + "DatabaseRowsPayloadBackend", + "DatasourcePayloadBackend", + "PayloadBackendRegistry", + "create_database_rows_payload", + "create_external_payload_reference", + "payload_for_materialization", + "finalize_payload_deletion", + "mark_unreferenced_payload_for_deletion", + "validate_payload_size", + "verify_payload_integrity", +] diff --git a/src/govoplan_datasources/backend/service.py b/src/govoplan_datasources/backend/service.py index 382fbdb..fbac34e 100644 --- a/src/govoplan_datasources/backend/service.py +++ b/src/govoplan_datasources/backend/service.py @@ -4,7 +4,7 @@ import hashlib import json import re from collections.abc import Mapping, Sequence -from dataclasses import replace +from dataclasses import dataclass, replace from typing import Any, cast from sqlalchemy import func, or_, select @@ -34,10 +34,18 @@ from govoplan_core.core.datasources import ( from govoplan_core.db.base import utcnow from govoplan_datasources.backend.db.models import ( DatasourceMaterializationRecord, + DatasourcePayloadRecord, DatasourcePublicationRecord, DatasourceRecord, DatasourceStageRecord, ) +from govoplan_datasources.backend.payloads import ( + DatasourcePayloadBackend, + PayloadBackendRegistry, + create_database_rows_payload, + payload_for_materialization, + validate_payload_size, +) from govoplan_datasources.backend.tabular import ( MAX_READ_ROWS, MAX_STAGE_ROWS, @@ -55,9 +63,27 @@ STAGE_WRITE_SCOPE = "datasources:stage:write" ADMIN_SCOPE = "datasources:source:admin" +@dataclass(frozen=True, slots=True) +class _PreparedPublication: + producer_module: str + producer_run_ref: str + idempotency_key: str + rows: tuple[dict[str, Any], ...] + schema: tuple[DatasourceField, ...] + fingerprint: str + byte_count: int + request_hash: str + + class SqlDatasourceProvider: - def __init__(self, *, registry: object | None = None) -> None: + def __init__( + self, + *, + registry: object | None = None, + payload_backends: Sequence[DatasourcePayloadBackend] = (), + ) -> None: self._registry = registry + self._payload_backends = PayloadBackendRegistry(payload_backends) def publish_rows( self, @@ -67,162 +93,49 @@ class SqlDatasourceProvider: request: DatasourcePublicationRequest, ) -> DatasourcePublicationResult: db, api_principal = _publication_context(session, principal) - producer_module = request.producer_module.strip() - producer_run_ref = request.producer_run_ref.strip() - idempotency_key = request.idempotency_key.strip() - if not producer_module or len(producer_module) > 100: - raise DatasourceValidationError( - "A producer module of at most 100 characters is required." - ) - if not producer_run_ref or len(producer_run_ref) > 500: - raise DatasourceValidationError( - "A producer run reference of at most 500 characters is required." - ) - if not idempotency_key or len(idempotency_key) > 255: - raise DatasourceValidationError( - "An idempotency key of at most 255 characters is required." - ) - normalized = normalize_rows(request.rows) - schema = infer_schema(normalized) - fingerprint = fingerprint_rows(normalized, schema) - request_hash = _publication_request_hash( - request, - normalized=normalized, - fingerprint=fingerprint, - ) - existing = db.scalar( - select(DatasourcePublicationRecord).where( - DatasourcePublicationRecord.tenant_id - == api_principal.tenant_id, - DatasourcePublicationRecord.producer_module == producer_module, - DatasourcePublicationRecord.idempotency_key == idempotency_key, - ) + prepared = _prepare_publication(request) + existing = _existing_publication_result( + db, + tenant_id=api_principal.tenant_id, + producer_module=prepared.producer_module, + idempotency_key=prepared.idempotency_key, + request_hash=prepared.request_hash, ) if existing is not None: - if existing.request_hash != request_hash: - raise DatasourceValidationError( - "The publication idempotency key was already used with " - "different output." - ) - datasource = db.get(DatasourceRecord, existing.datasource_id) - materialization = db.get( - DatasourceMaterializationRecord, - existing.materialization_id, - ) - if ( - datasource is None - or materialization is None - or datasource.tenant_id != api_principal.tenant_id - or materialization.tenant_id != api_principal.tenant_id - ): - raise DatasourceUnavailableError( - "The prior publication result is no longer available." - ) - return DatasourcePublicationResult( - ref=_publication_ref(existing.id), - status=existing.status, - datasource=_datasource_dto(datasource), - materialization=_materialization_dto(materialization), - replayed=True, - ) + return existing - datasource = None - if request.target_datasource_ref: - datasource = _required_datasource( - db, - tenant_id=api_principal.tenant_id, - datasource_ref=request.target_datasource_ref, - ) - if datasource.mode == "live" or datasource.shape != "tabular": - raise DatasourceValidationError( - "Produced rows require a static or cached tabular datasource." - ) - else: - name = str(request.name or "").strip() - source_name = str(request.source_name or "").strip() - if not name: - raise DatasourceValidationError( - "A datasource name is required for a new publication target." - ) - if not _valid_source_name(source_name): - raise DatasourceValidationError( - "Datasource keys must start with a letter or underscore and " - "contain only letters, numbers, and underscores." - ) - _ensure_source_name_available( - db, - tenant_id=api_principal.tenant_id, - source_name=source_name, - ) - datasource = DatasourceRecord( - tenant_id=api_principal.tenant_id, - source_name=source_name, - name=name, - description=_clean_optional(request.description), - kind="custom", - mode="static", - shape="tabular", - status="active", - provider=producer_module, - provider_ref=producer_run_ref, - schema_version=1, - schema_=[field_payload(field) for field in schema], - fingerprint=fingerprint, - row_count=len(normalized), - byte_count=encoded_size(normalized), - provenance_={ - **dict(request.provenance), - "producer_module": producer_module, - "producer_run_ref": producer_run_ref, - }, - metadata_=dict(request.metadata), - created_by=_actor_id(api_principal), - updated_by=_actor_id(api_principal), - ) - db.add(datasource) - db.flush() - - publication_provenance = { - **dict(request.provenance), - "producer_module": producer_module, - "producer_run_ref": producer_run_ref, - "idempotency_key": idempotency_key, - "published_at": utcnow().isoformat(), - } + actor_id = _actor_id(api_principal) + datasource = _publication_target( + db, + tenant_id=api_principal.tenant_id, + actor_id=actor_id, + request=request, + prepared=prepared, + ) materialization = _append_materialization( db, datasource=datasource, - rows=normalized, - schema=[field_payload(field) for field in schema], - fingerprint=fingerprint, - byte_count=encoded_size(normalized), - actor_id=_actor_id(api_principal), + rows=prepared.rows, + schema=[field_payload(field) for field in prepared.schema], + fingerprint=prepared.fingerprint, + byte_count=prepared.byte_count, + actor_id=actor_id, frozen=request.freeze, frozen_label=request.frozen_label, source_timestamp=request.source_timestamp, - provenance=publication_provenance, + provenance=_publication_provenance(request, prepared), metadata=dict(request.metadata), set_current=request.set_current, ) - publication = DatasourcePublicationRecord( + publication = _create_publication_record( + db, tenant_id=api_principal.tenant_id, - producer_module=producer_module, - producer_run_ref=producer_run_ref, - idempotency_key=idempotency_key, - request_hash=request_hash, - datasource_id=datasource.id, - materialization_id=materialization.id, - status="published", - details_={ - "fingerprint": fingerprint, - "row_count": len(normalized), - "set_current": request.set_current, - "frozen": request.freeze, - }, - created_by=_actor_id(api_principal), + actor_id=actor_id, + datasource=datasource, + materialization=materialization, + request=request, + prepared=prepared, ) - db.add(publication) - db.flush() return DatasourcePublicationResult( ref=_publication_ref(publication.id), status=publication.status, @@ -293,13 +206,12 @@ class SqlDatasourceProvider: offset = max(0, int(request.offset)) columns = tuple(dict.fromkeys(request.columns)) - direct_live = request.consistency == "live" - implicit_live = ( + read_live = request.consistency == "live" or ( item.mode == "live" and not request.materialization_ref and request.consistency == "current" ) - if direct_live or implicit_live: + if read_live: if not item.provider_ref: raise DatasourceUnavailableError( "This datasource has no live origin." @@ -322,10 +234,29 @@ class SqlDatasourceProvider: ) return result - materialization = _selected_materialization( + return self._read_materialized( db, item=item, request=request, + limit=limit, + offset=offset, + columns=columns, + ) + + def _read_materialized( + self, + session: Session, + *, + item: DatasourceRecord, + request: DatasourceReadRequest, + limit: int, + offset: int, + columns: tuple[str, ...], + ) -> DatasourceReadResult: + materialization = _selected_materialization( + session, + item=item, + request=request, ) if materialization is None: message = ( @@ -342,7 +273,18 @@ class SqlDatasourceProvider: "The datasource fingerprint changed; refresh the consuming definition." ) _validate_columns(materialization.schema_, columns) - window = materialization.rows[offset : offset + limit] + payload = payload_for_materialization(session, materialization) + if payload is None: + window = materialization.rows[offset : offset + limit] + else: + backend = self._payload_backends.require(payload.backend) + backend.verify(session, payload) + window = backend.read_rows( + session, + payload, + offset=offset, + limit=limit, + ) rows = tuple(_select_columns(row, columns) for row in window) descriptor = _datasource_dto(item) return DatasourceReadResult( @@ -698,10 +640,11 @@ class SqlDatasourceProvider: raise DatasourceUnavailableError( "The datasource has no current state to freeze." ) + current_payload = payload_for_materialization(db, current) materialization = _append_materialization( db, datasource=item, - rows=current.rows, + rows=current.rows if current_payload is None else (), schema=current.schema_, fingerprint=current.fingerprint, byte_count=current.byte_count, @@ -715,6 +658,7 @@ class SqlDatasourceProvider: }, metadata=dict(current.metadata_), set_current=False, + reusable_payload=current_payload, ) return _materialization_dto(materialization) @@ -918,21 +862,28 @@ def _append_materialization( provenance: Mapping[str, object] | None = None, metadata: Mapping[str, object] | None = None, set_current: bool, + reusable_payload: DatasourcePayloadRecord | None = None, ) -> DatasourceMaterializationRecord: - revision = ( - session.scalar( - select(func.max(DatasourceMaterializationRecord.revision)).where( - DatasourceMaterializationRecord.datasource_id == datasource.id - ) + datasource = _lock_datasource_for_materialization(session, datasource) + revision = _allocate_materialization_revision(session, datasource) + payload = reusable_payload or create_database_rows_payload( + session, + tenant_id=datasource.tenant_id, + rows=rows, + actor_id=actor_id, + metadata={ + "datasource_id": datasource.id, + "fingerprint": fingerprint, + }, + ) + if payload.tenant_id != datasource.tenant_id or payload.state != "ready": + raise DatasourceValidationError( + "Only a ready payload from the same tenant can be materialized." ) - or 0 - ) + 1 - schema_payload = [dict(field) for field in schema] - schema_changed = datasource.schema_ != schema_payload - schema_version = ( - datasource.schema_version + 1 - if schema_changed - else datasource.schema_version + validate_payload_size(payload, expected_byte_count=byte_count) + schema_payload, schema_version = _materialization_schema( + datasource, + schema, ) materialization = DatasourceMaterializationRecord( tenant_id=datasource.tenant_id, @@ -941,10 +892,12 @@ def _append_materialization( state="published", schema_version=max(1, int(schema_version or 1)), schema_=schema_payload, - rows=[dict(row) for row in rows], + payload_id=payload.id, + payload_checksum=payload.checksum, + rows=[], fingerprint=fingerprint, - row_count=len(rows), - byte_count=byte_count, + row_count=payload.row_count, + byte_count=payload.byte_count, frozen_at=utcnow() if frozen else None, frozen_label=_clean_optional(frozen_label), source_timestamp=source_timestamp, @@ -955,17 +908,85 @@ def _append_materialization( session.add(materialization) session.flush() if set_current: - datasource.current_materialization_id = materialization.id - datasource.schema_ = schema_payload - datasource.schema_version = schema_version - datasource.fingerprint = fingerprint - datasource.row_count = materialization.row_count - datasource.byte_count = materialization.byte_count - datasource.updated_by = actor_id + _apply_current_materialization( + datasource, + materialization=materialization, + schema=schema_payload, + schema_version=schema_version, + actor_id=actor_id, + ) session.flush() return materialization +def _lock_datasource_for_materialization( + session: Session, + datasource: DatasourceRecord, +) -> DatasourceRecord: + locked = session.scalar( + select(DatasourceRecord) + .where( + DatasourceRecord.id == datasource.id, + DatasourceRecord.tenant_id == datasource.tenant_id, + DatasourceRecord.deleted_at.is_(None), + ) + .with_for_update() + .execution_options(populate_existing=True) + ) + if locked is None: + raise DatasourceUnavailableError( + "The datasource is no longer available for materialization." + ) + return locked + + +def _allocate_materialization_revision( + session: Session, + datasource: DatasourceRecord, +) -> int: + """Allocate under the datasource row lock held by the caller.""" + + return int( + session.scalar( + select(func.max(DatasourceMaterializationRecord.revision)).where( + DatasourceMaterializationRecord.datasource_id == datasource.id + ) + ) + or 0 + ) + 1 + + +def _materialization_schema( + datasource: DatasourceRecord, + schema: Sequence[Mapping[str, object]], +) -> tuple[list[dict[str, object]], int]: + schema_payload = [dict(field) for field in schema] + schema_changed = datasource.schema_ != schema_payload + schema_version = ( + datasource.schema_version + 1 + if schema_changed + else datasource.schema_version + ) + return schema_payload, int(schema_version or 1) + + +def _apply_current_materialization( + datasource: DatasourceRecord, + *, + materialization: DatasourceMaterializationRecord, + schema: list[dict[str, object]], + schema_version: int, + actor_id: str | None, +) -> None: + datasource.current_materialization_id = materialization.id + datasource.schema_ = schema + datasource.schema_version = schema_version + datasource.fingerprint = materialization.fingerprint + datasource.row_count = materialization.row_count + datasource.byte_count = materialization.byte_count + datasource.updated_by = actor_id + + def _selected_materialization( session: Session, *, @@ -1295,6 +1316,208 @@ def _valid_source_name(value: str) -> bool: return bool(re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,119}", value)) +def _prepare_publication( + request: DatasourcePublicationRequest, +) -> _PreparedPublication: + producer_module = request.producer_module.strip() + producer_run_ref = request.producer_run_ref.strip() + idempotency_key = request.idempotency_key.strip() + _validate_publication_identity( + producer_module=producer_module, + producer_run_ref=producer_run_ref, + idempotency_key=idempotency_key, + ) + normalized = normalize_rows(request.rows) + schema = infer_schema(normalized) + fingerprint = fingerprint_rows(normalized, schema) + return _PreparedPublication( + producer_module=producer_module, + producer_run_ref=producer_run_ref, + idempotency_key=idempotency_key, + rows=normalized, + schema=schema, + fingerprint=fingerprint, + byte_count=encoded_size(normalized), + request_hash=_publication_request_hash( + request, + normalized=normalized, + fingerprint=fingerprint, + ), + ) + + +def _validate_publication_identity( + *, + producer_module: str, + producer_run_ref: str, + idempotency_key: str, +) -> None: + values = ( + (producer_module, 100, "A producer module"), + (producer_run_ref, 500, "A producer run reference"), + (idempotency_key, 255, "An idempotency key"), + ) + for value, maximum, label in values: + if not value or len(value) > maximum: + raise DatasourceValidationError( + f"{label} of at most {maximum} characters is required." + ) + + +def _existing_publication_result( + session: Session, + *, + tenant_id: str, + producer_module: str, + idempotency_key: str, + request_hash: str, +) -> DatasourcePublicationResult | None: + publication = session.scalar( + select(DatasourcePublicationRecord).where( + DatasourcePublicationRecord.tenant_id == tenant_id, + DatasourcePublicationRecord.producer_module == producer_module, + DatasourcePublicationRecord.idempotency_key == idempotency_key, + ) + ) + if publication is None: + return None + if publication.request_hash != request_hash: + raise DatasourceValidationError( + "The publication idempotency key was already used with different output." + ) + datasource = session.get(DatasourceRecord, publication.datasource_id) + materialization = session.get( + DatasourceMaterializationRecord, + publication.materialization_id, + ) + if ( + datasource is None + or materialization is None + or datasource.tenant_id != tenant_id + or materialization.tenant_id != tenant_id + ): + raise DatasourceUnavailableError( + "The prior publication result is no longer available." + ) + return DatasourcePublicationResult( + ref=_publication_ref(publication.id), + status=publication.status, + datasource=_datasource_dto(datasource), + materialization=_materialization_dto(materialization), + replayed=True, + ) + + +def _publication_target( + session: Session, + *, + tenant_id: str, + actor_id: str | None, + request: DatasourcePublicationRequest, + prepared: _PreparedPublication, +) -> DatasourceRecord: + if request.target_datasource_ref: + datasource = _required_datasource( + session, + tenant_id=tenant_id, + datasource_ref=request.target_datasource_ref, + ) + if datasource.mode == "live" or datasource.shape != "tabular": + raise DatasourceValidationError( + "Produced rows require a static or cached tabular datasource." + ) + return datasource + name = str(request.name or "").strip() + source_name = str(request.source_name or "").strip() + if not name: + raise DatasourceValidationError( + "A datasource name is required for a new publication target." + ) + if not _valid_source_name(source_name): + raise DatasourceValidationError( + "Datasource keys must start with a letter or underscore and contain " + "only letters, numbers, and underscores." + ) + _ensure_source_name_available( + session, + tenant_id=tenant_id, + source_name=source_name, + ) + datasource = DatasourceRecord( + tenant_id=tenant_id, + source_name=source_name, + name=name, + description=_clean_optional(request.description), + kind="custom", + mode="static", + shape="tabular", + status="active", + provider=prepared.producer_module, + provider_ref=prepared.producer_run_ref, + schema_version=1, + schema_=[field_payload(field) for field in prepared.schema], + fingerprint=prepared.fingerprint, + row_count=len(prepared.rows), + byte_count=prepared.byte_count, + provenance_={ + **dict(request.provenance), + "producer_module": prepared.producer_module, + "producer_run_ref": prepared.producer_run_ref, + }, + metadata_=dict(request.metadata), + created_by=actor_id, + updated_by=actor_id, + ) + session.add(datasource) + session.flush() + return datasource + + +def _publication_provenance( + request: DatasourcePublicationRequest, + prepared: _PreparedPublication, +) -> dict[str, object]: + return { + **dict(request.provenance), + "producer_module": prepared.producer_module, + "producer_run_ref": prepared.producer_run_ref, + "idempotency_key": prepared.idempotency_key, + "published_at": utcnow().isoformat(), + } + + +def _create_publication_record( + session: Session, + *, + tenant_id: str, + actor_id: str | None, + datasource: DatasourceRecord, + materialization: DatasourceMaterializationRecord, + request: DatasourcePublicationRequest, + prepared: _PreparedPublication, +) -> DatasourcePublicationRecord: + publication = DatasourcePublicationRecord( + tenant_id=tenant_id, + producer_module=prepared.producer_module, + producer_run_ref=prepared.producer_run_ref, + idempotency_key=prepared.idempotency_key, + request_hash=prepared.request_hash, + datasource_id=datasource.id, + materialization_id=materialization.id, + status="published", + details_={ + "fingerprint": prepared.fingerprint, + "row_count": len(prepared.rows), + "set_current": request.set_current, + "frozen": request.freeze, + }, + created_by=actor_id, + ) + session.add(publication) + session.flush() + return publication + + def _publication_request_hash( request: DatasourcePublicationRequest, *, diff --git a/src/govoplan_datasources/backend/tabular.py b/src/govoplan_datasources/backend/tabular.py index 497366f..5841267 100644 --- a/src/govoplan_datasources/backend/tabular.py +++ b/src/govoplan_datasources/backend/tabular.py @@ -111,11 +111,33 @@ def fingerprint_rows( def encoded_size(rows: Sequence[Mapping[str, object]]) -> int: - return len( - json.dumps(rows, sort_keys=True, separators=(",", ":"), default=str).encode( - "utf-8" - ) - ) + return len(_encoded_rows(rows)) + + +def payload_checksum(rows: Sequence[Mapping[str, object]]) -> str: + return hashlib.sha256(_encoded_rows(rows)).hexdigest() + + +def payload_row_checksum(row: Mapping[str, object]) -> str: + return hashlib.sha256(_encoded_row(row)).hexdigest() + + +def _encoded_rows(rows: Sequence[Mapping[str, object]]) -> bytes: + return json.dumps( + rows, + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") + + +def _encoded_row(row: Mapping[str, object]) -> bytes: + return json.dumps( + row, + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") def field_payload(field: DatasourceField) -> dict[str, object]: @@ -194,5 +216,7 @@ __all__ = [ "fingerprint_rows", "infer_schema", "normalize_rows", + "payload_checksum", + "payload_row_checksum", "parse_csv_rows", ] diff --git a/tests/test_lifecycle.py b/tests/test_lifecycle.py index efa4caa..d3f7e97 100644 --- a/tests/test_lifecycle.py +++ b/tests/test_lifecycle.py @@ -17,11 +17,14 @@ from govoplan_core.core.datasources import ( DatasourcePublicationRequest, DatasourceReadRequest, DatasourceStageInput, + DatasourceUnavailableError, DatasourceValidationError, ) from govoplan_core.db.base import Base, utcnow from govoplan_datasources.backend.db.models import ( DatasourceMaterializationRecord, + DatasourcePayloadRecord, + DatasourcePayloadRowRecord, DatasourcePublicationRecord, DatasourceRecord, DatasourceStageRecord, @@ -32,6 +35,12 @@ from govoplan_datasources.backend.service import ( STAGE_WRITE_SCOPE, SqlDatasourceProvider, ) +from govoplan_datasources.backend.payloads import ( + create_database_rows_payload, + finalize_payload_deletion, + mark_unreferenced_payload_for_deletion, + verify_payload_integrity, +) def principal( @@ -137,6 +146,8 @@ class DatasourceLifecycleTests(unittest.TestCase): self.engine, tables=[ DatasourceRecord.__table__, + DatasourcePayloadRecord.__table__, + DatasourcePayloadRowRecord.__table__, DatasourceMaterializationRecord.__table__, DatasourceStageRecord.__table__, DatasourcePublicationRecord.__table__, @@ -157,6 +168,8 @@ class DatasourceLifecycleTests(unittest.TestCase): DatasourceStageRecord.__table__, DatasourcePublicationRecord.__table__, DatasourceMaterializationRecord.__table__, + DatasourcePayloadRowRecord.__table__, + DatasourcePayloadRecord.__table__, DatasourceRecord.__table__, ], ) @@ -235,6 +248,18 @@ class DatasourceLifecycleTests(unittest.TestCase): current.datasource.schema_version, ) self.assertEqual(frozen.ref, frozen_result.materialization.ref) + first_record = self.session.get( + DatasourceMaterializationRecord, + first.ref.removeprefix("materialization:"), + ) + frozen_record = self.session.get( + DatasourceMaterializationRecord, + frozen.ref.removeprefix("materialization:"), + ) + self.assertIsNotNone(first_record) + self.assertIsNotNone(frozen_record) + self.assertEqual(first_record.payload_id, frozen_record.payload_id) + self.assertEqual([], first_record.rows) def test_live_reads_origin_and_cached_refresh_is_explicit(self) -> None: live = self.provider.register_origin( @@ -445,6 +470,114 @@ class DatasourceLifecycleTests(unittest.TestCase): ), ) + def test_payload_preview_is_paged_and_metadata_mismatch_is_rejected(self) -> None: + stage = self.provider.create_stage( + self.session, + principal(), + stage=DatasourceStageInput( + name="Paged", + source_name="paged", + kind="upload", + mode="static", + shape="tabular", + rows=tuple({"id": index} for index in range(20)), + ), + ) + datasource, materialization = self.provider.promote_stage( + self.session, + principal(), + stage_ref=stage.ref, + ) + self.session.commit() + + result = self.provider.read_datasource( + self.session, + principal(), + request=DatasourceReadRequest( + datasource_ref=datasource.ref, + offset=7, + limit=3, + ), + ) + self.assertEqual([{"id": 7}, {"id": 8}, {"id": 9}], list(result.rows)) + record = self.session.get( + DatasourceMaterializationRecord, + materialization.ref.removeprefix("materialization:"), + ) + self.assertIsNotNone(record) + self.assertEqual([], record.rows) + self.assertEqual( + 20, + self.session.query(DatasourcePayloadRowRecord) + .filter(DatasourcePayloadRowRecord.payload_id == record.payload_id) + .count(), + ) + + payload = self.session.get(DatasourcePayloadRecord, record.payload_id) + self.assertIsNotNone(payload) + verify_payload_integrity(self.session, payload) + payload.row_count += 1 + self.session.flush() + with self.assertRaises(DatasourceUnavailableError): + self.provider.read_datasource( + self.session, + principal(), + request=DatasourceReadRequest(datasource_ref=datasource.ref), + ) + + def test_payload_deletion_is_staged_and_reference_safe(self) -> None: + payload = create_database_rows_payload( + self.session, + tenant_id="tenant-1", + rows=({"id": 1}, {"id": 2}), + actor_id="account-1", + ) + self.session.commit() + + self.assertTrue( + mark_unreferenced_payload_for_deletion(self.session, payload) + ) + self.assertEqual("deleting", payload.state) + self.session.commit() + + finalize_payload_deletion(self.session, payload) + self.session.commit() + self.assertEqual(0, self.session.query(DatasourcePayloadRecord).count()) + self.assertEqual( + 0, + self.session.query(DatasourcePayloadRowRecord).count(), + ) + + def test_rolled_back_materialization_leaves_no_payload_rows(self) -> None: + stage = self.provider.create_stage( + self.session, + principal(), + stage=DatasourceStageInput( + name="Rollback", + source_name="rollback", + kind="upload", + mode="static", + shape="tabular", + rows=({"id": 1},), + ), + ) + self.provider.promote_stage( + self.session, + principal(), + stage_ref=stage.ref, + ) + self.session.rollback() + + self.assertEqual(0, self.session.query(DatasourcePayloadRecord).count()) + self.assertEqual( + 0, + self.session.query(DatasourcePayloadRowRecord).count(), + ) + self.assertEqual( + 0, + self.session.query(DatasourceMaterializationRecord).count(), + ) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_migrations.py b/tests/test_migrations.py index 02b44c5..d1c6467 100644 --- a/tests/test_migrations.py +++ b/tests/test_migrations.py @@ -24,13 +24,15 @@ class DatasourceMigrationTests(unittest.TestCase): try: with engine.connect() as connection: self.assertIn( - "c4e9f1a7d2b6", + "d5f0a2b8c3e7", set(MigrationContext.configure(connection).get_current_heads()), ) self.assertEqual( { "datasource_catalogue", "datasource_materializations", + "datasource_payload_rows", + "datasource_payloads", "datasource_publications", "datasource_stages", }, diff --git a/tests/test_postgres_materialization_concurrency.py b/tests/test_postgres_materialization_concurrency.py new file mode 100644 index 0000000..d170455 --- /dev/null +++ b/tests/test_postgres_materialization_concurrency.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +import os +import threading +import unittest +import uuid +from concurrent.futures import ThreadPoolExecutor + +from sqlalchemy import create_engine, text +from sqlalchemy.orm import Session + +from govoplan_core.db.base import Base +from govoplan_datasources.backend.db.models import ( + DatasourceMaterializationRecord, + DatasourcePayloadRecord, + DatasourcePayloadRowRecord, + DatasourceRecord, +) +from govoplan_datasources.backend.service import _append_materialization +from govoplan_datasources.backend.tabular import ( + encoded_size, + field_payload, + fingerprint_rows, + infer_schema, +) + + +@unittest.skipUnless( + os.environ.get("GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL"), + "set GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL for PostgreSQL concurrency checks", +) +class DatasourceMaterializationPostgresTests(unittest.TestCase): + def setUp(self) -> None: + database_url = os.environ["GOVOPLAN_DATASOURCES_TEST_POSTGRES_URL"] + self.schema = f"datasource_revision_race_{uuid.uuid4().hex}" + self.admin_engine = create_engine(database_url) + with self.admin_engine.begin() as connection: + connection.execute(text(f'CREATE SCHEMA "{self.schema}"')) + self.engine = create_engine( + database_url, + connect_args={"options": f"-c search_path={self.schema}"}, + ) + self.tables = [ + DatasourceRecord.__table__, + DatasourcePayloadRecord.__table__, + DatasourcePayloadRowRecord.__table__, + DatasourceMaterializationRecord.__table__, + ] + Base.metadata.create_all(self.engine, tables=self.tables) + with Session(self.engine) as session: + session.add( + DatasourceRecord( + id="datasource-1", + tenant_id="tenant-1", + source_name="concurrent", + name="Concurrent", + kind="custom", + mode="static", + shape="tabular", + ) + ) + session.commit() + + def tearDown(self) -> None: + try: + Base.metadata.drop_all( + self.engine, + tables=list(reversed(self.tables)), + ) + finally: + self.engine.dispose() + with self.admin_engine.begin() as connection: + connection.execute(text(f'DROP SCHEMA "{self.schema}"')) + self.admin_engine.dispose() + + def test_competing_publishers_receive_distinct_monotonic_revisions(self) -> None: + barrier = threading.Barrier(2) + + def publish(value: int) -> int: + rows = ({"id": value},) + schema = infer_schema(rows) + with Session(self.engine) as session: + datasource = session.get(DatasourceRecord, "datasource-1") + self.assertIsNotNone(datasource) + barrier.wait(timeout=10) + materialization = _append_materialization( + session, + datasource=datasource, + rows=rows, + schema=[field_payload(field) for field in schema], + fingerprint=fingerprint_rows(rows, schema), + byte_count=encoded_size(rows), + actor_id=f"worker-{value}", + set_current=True, + ) + revision = materialization.revision + session.commit() + return revision + + with ThreadPoolExecutor(max_workers=2) as executor: + revisions = tuple(executor.map(publish, (1, 2))) + + self.assertEqual((1, 2), tuple(sorted(revisions))) + with Session(self.engine) as session: + self.assertEqual( + [1, 2], + list( + session.scalars( + DatasourceMaterializationRecord.__table__.select() + .with_only_columns( + DatasourceMaterializationRecord.revision + ) + .order_by( + DatasourceMaterializationRecord.revision.asc() + ) + ) + ), + ) + + +if __name__ == "__main__": + unittest.main()