feat: add idempotent producer publication

This commit is contained in:
2026-07-28 13:47:53 +02:00
parent 1cb6228442
commit 8ae7dace99
9 changed files with 530 additions and 4 deletions
+230
View File
@@ -1,5 +1,8 @@
from __future__ import annotations
import hashlib
import json
import re
from collections.abc import Mapping, Sequence
from dataclasses import replace
from typing import Any, cast
@@ -18,6 +21,8 @@ from govoplan_core.core.datasources import (
DatasourceNotFoundError,
DatasourceOrigin,
DatasourceOriginReadRequest,
DatasourcePublicationRequest,
DatasourcePublicationResult,
DatasourceReadRequest,
DatasourceReadResult,
DatasourceStage,
@@ -29,6 +34,7 @@ from govoplan_core.core.datasources import (
from govoplan_core.db.base import utcnow
from govoplan_datasources.backend.db.models import (
DatasourceMaterializationRecord,
DatasourcePublicationRecord,
DatasourceRecord,
DatasourceStageRecord,
)
@@ -53,6 +59,178 @@ class SqlDatasourceProvider:
def __init__(self, *, registry: object | None = None) -> None:
self._registry = registry
def publish_rows(
self,
session: object,
principal: object,
*,
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,
)
)
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,
)
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(),
}
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),
frozen=request.freeze,
frozen_label=request.frozen_label,
source_timestamp=request.source_timestamp,
provenance=publication_provenance,
metadata=dict(request.metadata),
set_current=request.set_current,
)
publication = DatasourcePublicationRecord(
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),
)
db.add(publication)
db.flush()
return DatasourcePublicationResult(
ref=_publication_ref(publication.id),
status=publication.status,
datasource=_datasource_dto(datasource),
materialization=_materialization_dto(materialization),
replayed=False,
)
def list_datasources(
self,
session: object,
@@ -1038,6 +1216,19 @@ def _context(
return session, principal
def _publication_context(
session: object,
principal: object,
) -> tuple[Session, ApiPrincipal]:
if not isinstance(session, Session):
raise TypeError("Datasource providers require a SQLAlchemy session.")
if not isinstance(principal, ApiPrincipal):
raise DatasourceAccessError("A tenant API principal is required.")
if not principal.tenant_id:
raise DatasourceAccessError("A tenant API principal is required.")
return session, principal
def _validate_columns(
schema: Sequence[Mapping[str, object]],
columns: tuple[str, ...],
@@ -1082,6 +1273,10 @@ def _stage_ref(item_id: str) -> str:
return f"stage:{item_id}"
def _publication_ref(item_id: str) -> str:
return f"publication:{item_id}"
def _actor_id(principal: ApiPrincipal) -> str | None:
return principal.account_id or principal.membership_id or principal.identity_id
@@ -1102,6 +1297,41 @@ def _escape_like(value: str) -> str:
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
def _valid_source_name(value: str) -> bool:
return bool(re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,119}", value))
def _publication_request_hash(
request: DatasourcePublicationRequest,
*,
normalized: Sequence[Mapping[str, object]],
fingerprint: str,
) -> str:
payload = {
"producer_module": request.producer_module.strip(),
"producer_run_ref": request.producer_run_ref.strip(),
"target_datasource_ref": request.target_datasource_ref,
"name": request.name,
"source_name": request.source_name,
"description": request.description,
"rows": [dict(row) for row in normalized],
"fingerprint": fingerprint,
"freeze": request.freeze,
"frozen_label": request.frozen_label,
"set_current": request.set_current,
"source_timestamp": request.source_timestamp,
"provenance": dict(request.provenance),
"metadata": dict(request.metadata),
}
encoded = json.dumps(
payload,
sort_keys=True,
separators=(",", ":"),
default=str,
)
return hashlib.sha256(encoded.encode("utf-8")).hexdigest()
__all__ = [
"ADMIN_SCOPE",
"CATALOGUE_READ_SCOPE",