feat: harden datasource materialization payloads

This commit is contained in:
2026-07-29 19:38:34 +02:00
parent 5f0847b8ca
commit f3288dbd28
9 changed files with 1460 additions and 179 deletions

View File

@@ -122,6 +122,17 @@ class DatasourceMaterializationRecord(Base, TimestampMixin):
default=list, default=list,
nullable=False, 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) rows: Mapped[list[dict[str, Any]]] = mapped_column(JSON, default=list, nullable=False)
fingerprint: Mapped[str] = mapped_column(String(64), nullable=False, index=True) fingerprint: Mapped[str] = mapped_column(String(64), nullable=False, index=True)
row_count: Mapped[int] = mapped_column(Integer, nullable=False) 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) created_by: Mapped[str | None] = mapped_column(String(255), nullable=True, index=True)
datasource: Mapped[DatasourceRecord] = relationship(back_populates="materializations") 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): class DatasourceStageRecord(Base, TimestampMixin):
@@ -268,6 +368,8 @@ class DatasourcePublicationRecord(Base, TimestampMixin):
__all__ = [ __all__ = [
"DatasourceMaterializationRecord", "DatasourceMaterializationRecord",
"DatasourcePayloadRecord",
"DatasourcePayloadRowRecord",
"DatasourcePublicationRecord", "DatasourcePublicationRecord",
"DatasourceRecord", "DatasourceRecord",
"DatasourceStageRecord", "DatasourceStageRecord",

View File

@@ -244,6 +244,8 @@ manifest = ModuleManifest(
datasource_models.DatasourcePublicationRecord, datasource_models.DatasourcePublicationRecord,
datasource_models.DatasourceStageRecord, datasource_models.DatasourceStageRecord,
datasource_models.DatasourceMaterializationRecord, datasource_models.DatasourceMaterializationRecord,
datasource_models.DatasourcePayloadRowRecord,
datasource_models.DatasourcePayloadRecord,
datasource_models.DatasourceRecord, datasource_models.DatasourceRecord,
label="Datasources", label="Datasources",
), ),
@@ -256,6 +258,8 @@ manifest = ModuleManifest(
persistent_table_uninstall_guard( persistent_table_uninstall_guard(
datasource_models.DatasourceRecord, datasource_models.DatasourceRecord,
datasource_models.DatasourceMaterializationRecord, datasource_models.DatasourceMaterializationRecord,
datasource_models.DatasourcePayloadRecord,
datasource_models.DatasourcePayloadRowRecord,
datasource_models.DatasourceStageRecord, datasource_models.DatasourceStageRecord,
datasource_models.DatasourcePublicationRecord, datasource_models.DatasourcePublicationRecord,
label="Datasources", label="Datasources",

View File

@@ -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")

View File

@@ -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",
]

View File

@@ -4,7 +4,7 @@ import hashlib
import json import json
import re import re
from collections.abc import Mapping, Sequence from collections.abc import Mapping, Sequence
from dataclasses import replace from dataclasses import dataclass, replace
from typing import Any, cast from typing import Any, cast
from sqlalchemy import func, or_, select 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_core.db.base import utcnow
from govoplan_datasources.backend.db.models import ( from govoplan_datasources.backend.db.models import (
DatasourceMaterializationRecord, DatasourceMaterializationRecord,
DatasourcePayloadRecord,
DatasourcePublicationRecord, DatasourcePublicationRecord,
DatasourceRecord, DatasourceRecord,
DatasourceStageRecord, 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 ( from govoplan_datasources.backend.tabular import (
MAX_READ_ROWS, MAX_READ_ROWS,
MAX_STAGE_ROWS, MAX_STAGE_ROWS,
@@ -55,9 +63,27 @@ STAGE_WRITE_SCOPE = "datasources:stage:write"
ADMIN_SCOPE = "datasources:source:admin" 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: 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._registry = registry
self._payload_backends = PayloadBackendRegistry(payload_backends)
def publish_rows( def publish_rows(
self, self,
@@ -67,162 +93,49 @@ class SqlDatasourceProvider:
request: DatasourcePublicationRequest, request: DatasourcePublicationRequest,
) -> DatasourcePublicationResult: ) -> DatasourcePublicationResult:
db, api_principal = _publication_context(session, principal) db, api_principal = _publication_context(session, principal)
producer_module = request.producer_module.strip() prepared = _prepare_publication(request)
producer_run_ref = request.producer_run_ref.strip() existing = _existing_publication_result(
idempotency_key = request.idempotency_key.strip() db,
if not producer_module or len(producer_module) > 100: tenant_id=api_principal.tenant_id,
raise DatasourceValidationError( producer_module=prepared.producer_module,
"A producer module of at most 100 characters is required." idempotency_key=prepared.idempotency_key,
) request_hash=prepared.request_hash,
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,
)
) )
if existing is not None: if existing is not None:
if existing.request_hash != request_hash: return existing
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,
)
datasource = None actor_id = _actor_id(api_principal)
if request.target_datasource_ref: datasource = _publication_target(
datasource = _required_datasource( db,
db, tenant_id=api_principal.tenant_id,
tenant_id=api_principal.tenant_id, actor_id=actor_id,
datasource_ref=request.target_datasource_ref, request=request,
) prepared=prepared,
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(),
}
materialization = _append_materialization( materialization = _append_materialization(
db, db,
datasource=datasource, datasource=datasource,
rows=normalized, rows=prepared.rows,
schema=[field_payload(field) for field in schema], schema=[field_payload(field) for field in prepared.schema],
fingerprint=fingerprint, fingerprint=prepared.fingerprint,
byte_count=encoded_size(normalized), byte_count=prepared.byte_count,
actor_id=_actor_id(api_principal), actor_id=actor_id,
frozen=request.freeze, frozen=request.freeze,
frozen_label=request.frozen_label, frozen_label=request.frozen_label,
source_timestamp=request.source_timestamp, source_timestamp=request.source_timestamp,
provenance=publication_provenance, provenance=_publication_provenance(request, prepared),
metadata=dict(request.metadata), metadata=dict(request.metadata),
set_current=request.set_current, set_current=request.set_current,
) )
publication = DatasourcePublicationRecord( publication = _create_publication_record(
db,
tenant_id=api_principal.tenant_id, tenant_id=api_principal.tenant_id,
producer_module=producer_module, actor_id=actor_id,
producer_run_ref=producer_run_ref, datasource=datasource,
idempotency_key=idempotency_key, materialization=materialization,
request_hash=request_hash, request=request,
datasource_id=datasource.id, prepared=prepared,
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),
) )
db.add(publication)
db.flush()
return DatasourcePublicationResult( return DatasourcePublicationResult(
ref=_publication_ref(publication.id), ref=_publication_ref(publication.id),
status=publication.status, status=publication.status,
@@ -293,13 +206,12 @@ class SqlDatasourceProvider:
offset = max(0, int(request.offset)) offset = max(0, int(request.offset))
columns = tuple(dict.fromkeys(request.columns)) columns = tuple(dict.fromkeys(request.columns))
direct_live = request.consistency == "live" read_live = request.consistency == "live" or (
implicit_live = (
item.mode == "live" item.mode == "live"
and not request.materialization_ref and not request.materialization_ref
and request.consistency == "current" and request.consistency == "current"
) )
if direct_live or implicit_live: if read_live:
if not item.provider_ref: if not item.provider_ref:
raise DatasourceUnavailableError( raise DatasourceUnavailableError(
"This datasource has no live origin." "This datasource has no live origin."
@@ -322,10 +234,29 @@ class SqlDatasourceProvider:
) )
return result return result
materialization = _selected_materialization( return self._read_materialized(
db, db,
item=item, item=item,
request=request, 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: if materialization is None:
message = ( message = (
@@ -342,7 +273,18 @@ class SqlDatasourceProvider:
"The datasource fingerprint changed; refresh the consuming definition." "The datasource fingerprint changed; refresh the consuming definition."
) )
_validate_columns(materialization.schema_, columns) _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) rows = tuple(_select_columns(row, columns) for row in window)
descriptor = _datasource_dto(item) descriptor = _datasource_dto(item)
return DatasourceReadResult( return DatasourceReadResult(
@@ -698,10 +640,11 @@ class SqlDatasourceProvider:
raise DatasourceUnavailableError( raise DatasourceUnavailableError(
"The datasource has no current state to freeze." "The datasource has no current state to freeze."
) )
current_payload = payload_for_materialization(db, current)
materialization = _append_materialization( materialization = _append_materialization(
db, db,
datasource=item, datasource=item,
rows=current.rows, rows=current.rows if current_payload is None else (),
schema=current.schema_, schema=current.schema_,
fingerprint=current.fingerprint, fingerprint=current.fingerprint,
byte_count=current.byte_count, byte_count=current.byte_count,
@@ -715,6 +658,7 @@ class SqlDatasourceProvider:
}, },
metadata=dict(current.metadata_), metadata=dict(current.metadata_),
set_current=False, set_current=False,
reusable_payload=current_payload,
) )
return _materialization_dto(materialization) return _materialization_dto(materialization)
@@ -918,21 +862,28 @@ def _append_materialization(
provenance: Mapping[str, object] | None = None, provenance: Mapping[str, object] | None = None,
metadata: Mapping[str, object] | None = None, metadata: Mapping[str, object] | None = None,
set_current: bool, set_current: bool,
reusable_payload: DatasourcePayloadRecord | None = None,
) -> DatasourceMaterializationRecord: ) -> DatasourceMaterializationRecord:
revision = ( datasource = _lock_datasource_for_materialization(session, datasource)
session.scalar( revision = _allocate_materialization_revision(session, datasource)
select(func.max(DatasourceMaterializationRecord.revision)).where( payload = reusable_payload or create_database_rows_payload(
DatasourceMaterializationRecord.datasource_id == datasource.id 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 validate_payload_size(payload, expected_byte_count=byte_count)
) + 1 schema_payload, schema_version = _materialization_schema(
schema_payload = [dict(field) for field in schema] datasource,
schema_changed = datasource.schema_ != schema_payload schema,
schema_version = (
datasource.schema_version + 1
if schema_changed
else datasource.schema_version
) )
materialization = DatasourceMaterializationRecord( materialization = DatasourceMaterializationRecord(
tenant_id=datasource.tenant_id, tenant_id=datasource.tenant_id,
@@ -941,10 +892,12 @@ def _append_materialization(
state="published", state="published",
schema_version=max(1, int(schema_version or 1)), schema_version=max(1, int(schema_version or 1)),
schema_=schema_payload, schema_=schema_payload,
rows=[dict(row) for row in rows], payload_id=payload.id,
payload_checksum=payload.checksum,
rows=[],
fingerprint=fingerprint, fingerprint=fingerprint,
row_count=len(rows), row_count=payload.row_count,
byte_count=byte_count, byte_count=payload.byte_count,
frozen_at=utcnow() if frozen else None, frozen_at=utcnow() if frozen else None,
frozen_label=_clean_optional(frozen_label), frozen_label=_clean_optional(frozen_label),
source_timestamp=source_timestamp, source_timestamp=source_timestamp,
@@ -955,17 +908,85 @@ def _append_materialization(
session.add(materialization) session.add(materialization)
session.flush() session.flush()
if set_current: if set_current:
datasource.current_materialization_id = materialization.id _apply_current_materialization(
datasource.schema_ = schema_payload datasource,
datasource.schema_version = schema_version materialization=materialization,
datasource.fingerprint = fingerprint schema=schema_payload,
datasource.row_count = materialization.row_count schema_version=schema_version,
datasource.byte_count = materialization.byte_count actor_id=actor_id,
datasource.updated_by = actor_id )
session.flush() session.flush()
return materialization 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( def _selected_materialization(
session: Session, 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)) 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( def _publication_request_hash(
request: DatasourcePublicationRequest, request: DatasourcePublicationRequest,
*, *,

View File

@@ -111,11 +111,33 @@ def fingerprint_rows(
def encoded_size(rows: Sequence[Mapping[str, object]]) -> int: def encoded_size(rows: Sequence[Mapping[str, object]]) -> int:
return len( return len(_encoded_rows(rows))
json.dumps(rows, sort_keys=True, separators=(",", ":"), default=str).encode(
"utf-8"
) 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]: def field_payload(field: DatasourceField) -> dict[str, object]:
@@ -194,5 +216,7 @@ __all__ = [
"fingerprint_rows", "fingerprint_rows",
"infer_schema", "infer_schema",
"normalize_rows", "normalize_rows",
"payload_checksum",
"payload_row_checksum",
"parse_csv_rows", "parse_csv_rows",
] ]

View File

@@ -17,11 +17,14 @@ from govoplan_core.core.datasources import (
DatasourcePublicationRequest, DatasourcePublicationRequest,
DatasourceReadRequest, DatasourceReadRequest,
DatasourceStageInput, DatasourceStageInput,
DatasourceUnavailableError,
DatasourceValidationError, DatasourceValidationError,
) )
from govoplan_core.db.base import Base, utcnow from govoplan_core.db.base import Base, utcnow
from govoplan_datasources.backend.db.models import ( from govoplan_datasources.backend.db.models import (
DatasourceMaterializationRecord, DatasourceMaterializationRecord,
DatasourcePayloadRecord,
DatasourcePayloadRowRecord,
DatasourcePublicationRecord, DatasourcePublicationRecord,
DatasourceRecord, DatasourceRecord,
DatasourceStageRecord, DatasourceStageRecord,
@@ -32,6 +35,12 @@ from govoplan_datasources.backend.service import (
STAGE_WRITE_SCOPE, STAGE_WRITE_SCOPE,
SqlDatasourceProvider, SqlDatasourceProvider,
) )
from govoplan_datasources.backend.payloads import (
create_database_rows_payload,
finalize_payload_deletion,
mark_unreferenced_payload_for_deletion,
verify_payload_integrity,
)
def principal( def principal(
@@ -137,6 +146,8 @@ class DatasourceLifecycleTests(unittest.TestCase):
self.engine, self.engine,
tables=[ tables=[
DatasourceRecord.__table__, DatasourceRecord.__table__,
DatasourcePayloadRecord.__table__,
DatasourcePayloadRowRecord.__table__,
DatasourceMaterializationRecord.__table__, DatasourceMaterializationRecord.__table__,
DatasourceStageRecord.__table__, DatasourceStageRecord.__table__,
DatasourcePublicationRecord.__table__, DatasourcePublicationRecord.__table__,
@@ -157,6 +168,8 @@ class DatasourceLifecycleTests(unittest.TestCase):
DatasourceStageRecord.__table__, DatasourceStageRecord.__table__,
DatasourcePublicationRecord.__table__, DatasourcePublicationRecord.__table__,
DatasourceMaterializationRecord.__table__, DatasourceMaterializationRecord.__table__,
DatasourcePayloadRowRecord.__table__,
DatasourcePayloadRecord.__table__,
DatasourceRecord.__table__, DatasourceRecord.__table__,
], ],
) )
@@ -235,6 +248,18 @@ class DatasourceLifecycleTests(unittest.TestCase):
current.datasource.schema_version, current.datasource.schema_version,
) )
self.assertEqual(frozen.ref, frozen_result.materialization.ref) 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: def test_live_reads_origin_and_cached_refresh_is_explicit(self) -> None:
live = self.provider.register_origin( 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__": if __name__ == "__main__":
unittest.main() unittest.main()

View File

@@ -24,13 +24,15 @@ class DatasourceMigrationTests(unittest.TestCase):
try: try:
with engine.connect() as connection: with engine.connect() as connection:
self.assertIn( self.assertIn(
"c4e9f1a7d2b6", "d5f0a2b8c3e7",
set(MigrationContext.configure(connection).get_current_heads()), set(MigrationContext.configure(connection).get_current_heads()),
) )
self.assertEqual( self.assertEqual(
{ {
"datasource_catalogue", "datasource_catalogue",
"datasource_materializations", "datasource_materializations",
"datasource_payload_rows",
"datasource_payloads",
"datasource_publications", "datasource_publications",
"datasource_stages", "datasource_stages",
}, },

View File

@@ -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()