Files
govoplan-dataflow/tests/test_service.py

452 lines
15 KiB
Python

from __future__ import annotations
import unittest
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session, sessionmaker
from govoplan_core.auth import ApiPrincipal
from govoplan_core.core.access import PrincipalRef
from govoplan_core.core.dataflows import (
DataflowPublicationTarget,
DataflowRunRequest,
)
from govoplan_core.core.datasources import (
CAPABILITY_DATASOURCE_PUBLICATION,
DatasourceDescriptor,
DatasourceMaterialization,
DatasourcePublicationResult,
)
from govoplan_core.db.base import Base
from govoplan_dataflow.backend.db.models import (
DataflowPipeline,
DataflowPipelineRevision,
DataflowRun,
)
from govoplan_dataflow.backend.schemas import (
GraphEdge,
GraphNode,
GraphPosition,
PipelineCreateRequest,
PipelineGraph,
PipelinePreviewRequest,
PipelineUpdateRequest,
)
from govoplan_dataflow.backend.service import (
DataflowConflictError,
DataflowNotFoundError,
create_pipeline,
delete_pipeline,
get_pipeline,
list_pipelines,
preview_pipeline,
start_pipeline_run,
update_pipeline,
)
def sample_graph(*, minimum: int = 10) -> PipelineGraph:
return PipelineGraph(
nodes=[
GraphNode(
id="source",
type="source.inline",
label="Monthly input",
position=GraphPosition(x=40, y=160),
config={
"source_name": "monthly_files",
"rows": [
{"id": 1, "amount": 5},
{"id": 2, "amount": 15},
{"id": 3, "amount": 25},
],
},
),
GraphNode(
id="filter",
type="filter",
label="Minimum amount",
position=GraphPosition(x=260, y=160),
config={"column": "amount", "operator": "gte", "value": minimum},
),
GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=480, y=160),
config={},
),
],
edges=[
GraphEdge(id="source-filter", source="source", target="filter"),
GraphEdge(id="filter-output", source="filter", target="output"),
],
)
def principal(tenant_id: str = "tenant-1") -> ApiPrincipal:
return ApiPrincipal(
principal=PrincipalRef(
account_id="account-1",
membership_id="membership-1",
tenant_id=tenant_id,
scopes=frozenset(),
),
account=object(),
user=object(),
)
class FakePublicationProvider:
def __init__(self) -> None:
self.requests = []
def publish_rows(self, _session, _principal, *, request):
self.requests.append(request)
descriptor = DatasourceDescriptor(
ref=request.target_datasource_ref or "datasource:output-1",
source_name=request.source_name or "existing_output",
name=request.name or "Existing output",
kind="custom",
mode="static",
shape="tabular",
fingerprint="published-fingerprint",
)
return DatasourcePublicationResult(
ref="publication:publication-1",
status="published",
datasource=descriptor,
materialization=DatasourceMaterialization(
ref="materialization:materialization-1",
datasource_ref=descriptor.ref,
revision=1,
state="published",
fingerprint=descriptor.fingerprint,
),
)
class FakeRegistry:
def __init__(self, publication_provider: FakePublicationProvider) -> None:
self.publication_provider = publication_provider
def has_capability(self, name: str) -> bool:
return name == CAPABILITY_DATASOURCE_PUBLICATION
def capability(self, name: str):
if not self.has_capability(name):
raise KeyError(name)
return self.publication_provider
class DataflowServiceTests(unittest.TestCase):
def setUp(self) -> None:
self.engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(
self.engine,
tables=[
DataflowPipeline.__table__,
DataflowPipelineRevision.__table__,
DataflowRun.__table__,
],
)
self.Session = sessionmaker(bind=self.engine)
self.session: Session = self.Session()
def tearDown(self) -> None:
self.session.close()
Base.metadata.drop_all(
self.engine,
tables=[
DataflowRun.__table__,
DataflowPipelineRevision.__table__,
DataflowPipeline.__table__,
],
)
self.engine.dispose()
def _create(self, *, tenant_id: str = "tenant-1") -> DataflowPipeline:
pipeline = create_pipeline(
self.session,
tenant_id=tenant_id,
actor_id="user-1",
payload=PipelineCreateRequest(
name="Monthly comparison",
description="First governed pipeline",
status="active",
graph=sample_graph(),
editor_mode="graph",
),
)
self.session.commit()
return pipeline
def test_create_and_update_make_immutable_revisions(self) -> None:
pipeline = self._create()
updated = update_pipeline(
self.session,
tenant_id="tenant-1",
pipeline_id=pipeline.id,
actor_id="user-2",
payload=PipelineUpdateRequest(
name="Monthly comparison",
description="Changed threshold",
graph=sample_graph(minimum=20),
editor_mode="graph",
expected_revision=1,
),
)
self.session.commit()
revisions = list(
self.session.scalars(
select(DataflowPipelineRevision)
.where(DataflowPipelineRevision.pipeline_id == pipeline.id)
.order_by(DataflowPipelineRevision.revision)
)
)
self.assertEqual(2, updated.current_revision)
self.assertEqual([1, 2], [item.revision for item in revisions])
self.assertNotEqual(revisions[0].content_hash, revisions[1].content_hash)
self.assertEqual(10, revisions[0].graph["nodes"][1]["config"]["value"])
self.assertEqual(20, revisions[1].graph["nodes"][1]["config"]["value"])
def test_stale_revision_is_rejected(self) -> None:
pipeline = self._create()
with self.assertRaises(DataflowConflictError):
update_pipeline(
self.session,
tenant_id="tenant-1",
pipeline_id=pipeline.id,
actor_id="user-2",
payload=PipelineUpdateRequest(
name="Stale edit",
graph=sample_graph(minimum=30),
editor_mode="graph",
expected_revision=2,
),
)
def test_pipeline_lookup_is_tenant_isolated(self) -> None:
pipeline = self._create()
with self.assertRaises(DataflowNotFoundError):
get_pipeline(
self.session,
tenant_id="tenant-2",
pipeline_id=pipeline.id,
)
def test_list_and_soft_delete_are_tenant_isolated(self) -> None:
first = self._create()
second = self._create(tenant_id="tenant-2")
self.assertEqual([first.id], [item.id for item in list_pipelines(self.session, tenant_id="tenant-1")])
self.assertEqual([second.id], [item.id for item in list_pipelines(self.session, tenant_id="tenant-2")])
delete_pipeline(
self.session,
tenant_id="tenant-1",
pipeline_id=first.id,
actor_id="user-2",
)
self.session.commit()
self.assertEqual([], list_pipelines(self.session, tenant_id="tenant-1"))
with self.assertRaises(DataflowNotFoundError):
get_pipeline(
self.session,
tenant_id="tenant-1",
pipeline_id=first.id,
)
self.assertEqual(
1,
self.session.scalar(
select(func.count())
.select_from(DataflowPipelineRevision)
.where(DataflowPipelineRevision.pipeline_id == first.id)
),
)
def test_saved_preview_records_lineage_but_not_result_rows(self) -> None:
pipeline = self._create()
response = preview_pipeline(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
payload=PipelinePreviewRequest(
pipeline_id=pipeline.id,
preview_node_id="source",
row_limit=1,
),
principal=principal(),
)
self.session.commit()
run = self.session.scalar(select(DataflowRun).where(DataflowRun.id == response.run_id))
self.assertEqual("succeeded", response.status)
self.assertEqual([{"id": 2, "amount": 15}], response.rows)
self.assertEqual(2, response.total_rows)
self.assertTrue(response.truncated)
self.assertIsNotNone(response.node_preview)
self.assertEqual("source", response.node_preview.node_id)
self.assertEqual([{"id": 1, "amount": 5}], response.node_preview.rows)
self.assertEqual(3, response.node_preview.total_rows)
self.assertTrue(response.node_preview.truncated)
self.assertEqual(2, run.output_row_count)
self.assertEqual(3, run.input_row_count)
self.assertEqual(1, len(run.source_fingerprints))
self.assertEqual(3, response.input_row_count)
self.assertEqual(run.source_fingerprints, response.source_fingerprints)
self.assertFalse(hasattr(run, "result_rows"))
self.assertEqual(
1,
self.session.scalar(select(func.count()).select_from(DataflowRun)),
)
def test_failed_preview_keeps_upstream_lineage_and_failed_node_diagnostic(self) -> None:
graph = sample_graph()
graph.nodes[0].config["rows"] = [{"id": 1, "amount": "not-a-number"}]
pipeline = create_pipeline(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
payload=PipelineCreateRequest(
name="Invalid runtime value",
graph=graph,
editor_mode="graph",
),
)
self.session.commit()
response = preview_pipeline(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
payload=PipelinePreviewRequest(pipeline_id=pipeline.id),
principal=principal(),
)
self.session.commit()
run = self.session.scalar(select(DataflowRun).where(DataflowRun.id == response.run_id))
self.assertEqual("failed", response.status)
self.assertEqual(
[("source", "succeeded"), ("filter", "failed")],
[(item.node_id, item.status) for item in response.node_diagnostics],
)
self.assertEqual(1, run.input_row_count)
self.assertEqual(1, len(run.source_fingerprints))
self.assertEqual(1, response.input_row_count)
self.assertEqual(run.source_fingerprints, response.source_fingerprints)
self.assertEqual("preview.execution", run.diagnostics[-1]["code"])
def test_pinned_run_publishes_once_and_replays_idempotently(self) -> None:
pipeline = self._create()
publication_provider = FakePublicationProvider()
request = DataflowRunRequest(
pipeline_ref=f"pipeline:{pipeline.id}",
revision=1,
idempotency_key="monthly-2026-07",
publication=DataflowPublicationTarget(
name="Monthly result",
source_name="monthly_result",
freeze=True,
frozen_label="July 2026",
),
)
first, first_replayed = start_pipeline_run(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
principal=principal(),
registry=FakeRegistry(publication_provider),
request=request,
)
replay, replayed = start_pipeline_run(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
principal=principal(),
registry=FakeRegistry(publication_provider),
request=request,
)
self.session.commit()
self.assertFalse(first_replayed)
self.assertTrue(replayed)
self.assertEqual(first.id, replay.id)
self.assertEqual("succeeded", first.status)
self.assertEqual("datasource:output-1", first.output_datasource_ref)
self.assertEqual(
"materialization:materialization-1",
first.output_materialization_ref,
)
self.assertEqual(
[{"id": 2, "amount": 15}, {"id": 3, "amount": 25}],
list(publication_provider.requests[0].rows),
)
self.assertEqual(1, len(publication_provider.requests))
def test_run_idempotency_key_rejects_changed_parameters(self) -> None:
pipeline = self._create()
publication_provider = FakePublicationProvider()
registry = FakeRegistry(publication_provider)
start_pipeline_run(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
principal=principal(),
registry=registry,
request=DataflowRunRequest(
pipeline_ref=f"pipeline:{pipeline.id}",
revision=1,
idempotency_key="stable-key",
row_limit=100,
),
)
with self.assertRaises(DataflowConflictError):
start_pipeline_run(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
principal=principal(),
registry=registry,
request=DataflowRunRequest(
pipeline_ref=f"pipeline:{pipeline.id}",
revision=1,
idempotency_key="stable-key",
row_limit=10,
),
)
def test_publication_without_datasources_finishes_as_failed_run(self) -> None:
pipeline = self._create()
run, _ = start_pipeline_run(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
principal=principal(),
registry=None,
request=DataflowRunRequest(
pipeline_ref=f"pipeline:{pipeline.id}",
revision=1,
idempotency_key="missing-publisher",
publication=DataflowPublicationTarget(
name="Output",
source_name="output",
),
),
)
self.assertEqual("failed", run.status)
self.assertIn("Datasources publication capability", run.error)
if __name__ == "__main__":
unittest.main()