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