Files
govoplan-datasources/src/govoplan_datasources/backend/payloads.py

375 lines
11 KiB
Python

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