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

View File

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

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 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,
*,

View File

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

View File

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

View File

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

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