feat: harden datasource materialization payloads
This commit is contained in:
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user