from __future__ import annotations from datetime import timedelta import unittest from sqlalchemy import create_engine, select from sqlalchemy.orm import Session, sessionmaker from govoplan_core.auth import ApiPrincipal from govoplan_core.core.access import ( CAPABILITY_AUTH_AUTOMATION_PRINCIPAL_PROVIDER, PrincipalRef, ) from govoplan_core.core.automation import AutomationPrincipalResolution from govoplan_core.core.dataflows import ( DataflowPublicationTarget, DataflowRunRequest, ) from govoplan_core.core.datasources import ( CAPABILITY_DATASOURCE_PUBLICATION, DatasourceDescriptor, DatasourceMaterialization, DatasourcePublicationResult, ) from govoplan_core.core.recovery import ( RecoveryCheckpoint, RecoveryOperation, RecoveryStatus, ) from govoplan_core.core.runtime_coordination import ( DistributedLease, RuntimeIdentity, bind_process_runtime_identity, ) from govoplan_core.db.base import Base, utcnow from govoplan_dataflow.backend.db.models import ( DataflowPipeline, DataflowPipelineDeployment, DataflowPipelineRevision, DataflowRun, ) from govoplan_dataflow.backend.run_worker import ( dispatch_pending_runs, purge_expired_runs, ) from govoplan_dataflow.backend.recovery import begin_dataflow_run_recovery from govoplan_dataflow.backend.schemas import ( GraphEdge, GraphNode, GraphPosition, PipelineCreateRequest, PipelineGraph, ) from govoplan_dataflow.backend.service import ( cancel_pipeline_run, create_pipeline, start_pipeline_run, ) def _principal() -> ApiPrincipal: return ApiPrincipal( principal=PrincipalRef( account_id="account-1", membership_id="membership-1", tenant_id="tenant-1", scopes=frozenset({"dataflow:pipeline:run"}), ), account=object(), user=object(), ) def _runtime_identity() -> RuntimeIdentity: return RuntimeIdentity( installation_id="dataflow-worker-tests", node_id="worker-node", incarnation="worker-incarnation", role="worker", software_version="test", composition_hash="b" * 64, ) def _graph() -> PipelineGraph: return PipelineGraph( nodes=[ GraphNode( id="source", type="source.inline", label="Input", position=GraphPosition(x=0, y=0), config={ "source_name": "input_rows", "rows": [{"id": 1}, {"id": 2}], }, ), GraphNode( id="output", type="output", label="Output", position=GraphPosition(x=200, y=0), config={}, ), ], edges=[ GraphEdge( id="source-output", source="source", target="output", ) ], ) class _AutomationProvider: def resolve_automation_principal(self, _session, *, request): return AutomationPrincipalResolution( allowed=True, principal=_principal(), granted_scopes=request.grant_scopes, provenance={"status": "rechecked"}, ) class _PublicationProvider: def __init__(self) -> None: self.requests = [] def publish_rows(self, _session, _principal, *, request): self.requests.append(request) descriptor = DatasourceDescriptor( ref="datasource:worker-output", source_name="worker_output", name="Worker output", kind="custom", mode="static", shape="tabular", fingerprint="worker-output-fingerprint", ) return DatasourcePublicationResult( ref="publication:worker-output", status="published", datasource=descriptor, materialization=DatasourceMaterialization( ref="materialization:worker-output", datasource_ref=descriptor.ref, revision=1, state="published", fingerprint=descriptor.fingerprint, ), ) class _Registry: def __init__( self, publication_provider: _PublicationProvider | None = None, ) -> None: self.provider = _AutomationProvider() self.publication_provider = publication_provider def has_capability(self, name: str) -> bool: return name == CAPABILITY_AUTH_AUTOMATION_PRINCIPAL_PROVIDER or ( name == CAPABILITY_DATASOURCE_PUBLICATION and self.publication_provider is not None ) def capability(self, name: str): if not self.has_capability(name): raise KeyError(name) if name == CAPABILITY_DATASOURCE_PUBLICATION: return self.publication_provider return self.provider class DataflowRunWorkerTests(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite:///:memory:") Base.metadata.create_all( self.engine, tables=[ DistributedLease.__table__, RecoveryOperation.__table__, RecoveryCheckpoint.__table__, DataflowPipeline.__table__, DataflowPipelineRevision.__table__, DataflowRun.__table__, DataflowPipelineDeployment.__table__, ], ) self.Session = sessionmaker(bind=self.engine) self.session: Session = self.Session() bind_process_runtime_identity(_runtime_identity()) self.pipeline = create_pipeline( self.session, tenant_id="tenant-1", actor_id="account-1", payload=PipelineCreateRequest( name="Worker flow", status="active", graph=_graph(), editor_mode="graph", ), ) self.session.commit() def tearDown(self) -> None: bind_process_runtime_identity(None) self.session.close() Base.metadata.drop_all( self.engine, tables=[ DataflowPipelineDeployment.__table__, DataflowRun.__table__, DataflowPipelineRevision.__table__, DataflowPipeline.__table__, RecoveryCheckpoint.__table__, RecoveryOperation.__table__, DistributedLease.__table__, ], ) self.engine.dispose() def _queue( self, key: str, *, publication: bool = False, ) -> DataflowRun: run, replayed = start_pipeline_run( self.session, tenant_id="tenant-1", actor_id="account-1", principal=_principal(), registry=_Registry(), request=DataflowRunRequest( pipeline_ref=f"pipeline:{self.pipeline.id}", revision=1, idempotency_key=key, publication=( DataflowPublicationTarget( name="Worker output", source_name="worker_output", ) if publication else None ), ), defer_execution=True, ) self.session.commit() self.assertFalse(replayed) return run def test_worker_claims_authorizes_and_executes_queued_run(self) -> None: run = self._queue("worker-success") self.assertEqual("queued", run.status) self.assertIsNone(run.started_at) result = dispatch_pending_runs( self.session, registry=_Registry(), worker_id="test-worker", ) self.session.refresh(run) self.assertEqual(1, result["claimed"]) self.assertEqual(1, result["succeeded"]) self.assertEqual("succeeded", run.status) self.assertEqual(2, run.output_row_count) self.assertEqual(1, run.attempts) self.assertEqual(100, run.progress_percent) self.assertIsNone(run.worker_id) self.assertEqual( "rechecked", run.authorization_["last_resolution"]["status"], ) def test_cancelled_queued_run_is_never_claimed(self) -> None: run = self._queue("worker-cancel") cancel_pipeline_run( self.session, tenant_id="tenant-1", run_ref=f"dataflow-run:{run.id}", ) self.session.commit() result = dispatch_pending_runs( self.session, registry=_Registry(), ) self.assertEqual(0, result["claimed"]) self.session.refresh(run) self.assertEqual("cancelled", run.status) def test_expired_worker_lease_is_recovered(self) -> None: run = self._queue("worker-recover") run.status = "running" run.attempts = 1 run.worker_id = "lost-worker" run.lease_expires_at = utcnow() - timedelta(minutes=1) self.session.commit() result = dispatch_pending_runs( self.session, registry=_Registry(), ) self.session.refresh(run) self.assertEqual(1, result["recovered"]) self.assertEqual("succeeded", run.status) self.assertEqual(2, run.attempts) def test_retention_purges_payload_but_keeps_run_evidence(self) -> None: run = self._queue("worker-retention") dispatch_pending_runs(self.session, registry=_Registry()) run.retention_until = utcnow() - timedelta(seconds=1) self.session.commit() result = purge_expired_runs(self.session) self.session.refresh(run) self.assertEqual(1, result["purged"]) self.assertEqual({}, run.request_) self.assertEqual("succeeded", run.status) self.assertEqual(2, run.output_row_count) self.assertIsNotNone(run.purged_at) def test_stale_prepublication_attempt_is_safely_retried(self) -> None: provider = _PublicationProvider() registry = _Registry(provider) run = self._queue("stale-before-publication", publication=True) run.status = "running" run.attempts = 1 run.worker_id = "lost-worker" run.lease_expires_at = utcnow() + timedelta(minutes=5) self.session.commit() begin_dataflow_run_recovery( self.session, run=run, lease_ttl_seconds=120, ) self.session.commit() lease = self.session.scalar( select(DistributedLease).where( DistributedLease.resource_key == f"dataflow:run:{run.id}" ) ) assert lease is not None lease.expires_at = utcnow() - timedelta(minutes=1) run.lease_expires_at = utcnow() - timedelta(minutes=1) self.session.commit() result = dispatch_pending_runs(self.session, registry=registry) self.session.refresh(run) self.assertEqual(1, result["recovered"]) self.assertEqual("succeeded", run.status) self.assertEqual(2, run.attempts) self.assertEqual(1, len(provider.requests)) statuses = set( self.session.scalars( select(RecoveryOperation.status).where( RecoveryOperation.resource_id == run.id ) ) ) self.assertEqual( {RecoveryStatus.RECOVERED.value, RecoveryStatus.SUCCEEDED.value}, statuses, ) def test_stale_publication_attempt_is_not_retried_blindly(self) -> None: provider = _PublicationProvider() registry = _Registry(provider) run = self._queue("stale-after-publication", publication=True) run.status = "running" run.attempts = 1 run.worker_id = "lost-worker" run.lease_expires_at = utcnow() + timedelta(minutes=5) run.output_row_count = 1 self.session.commit() recovery = begin_dataflow_run_recovery( self.session, run=run, lease_ttl_seconds=120, ) recovery.prepare_publication( self.session, run=run, rows=({"id": 1},), ) lease = self.session.scalar( select(DistributedLease).where( DistributedLease.resource_key == f"dataflow:run:{run.id}" ) ) assert lease is not None lease.expires_at = utcnow() - timedelta(minutes=1) run.lease_expires_at = utcnow() - timedelta(minutes=1) self.session.commit() result = dispatch_pending_runs(self.session, registry=registry) self.session.refresh(run) self.assertEqual(1, result["recovered"]) self.assertEqual(1, result["outcome_unknown"]) self.assertEqual(0, result["claimed"]) self.assertEqual("outcome_unknown", run.status) self.assertEqual(1, run.attempts) self.assertEqual([], provider.requests) operation = self.session.scalar( select(RecoveryOperation).where( RecoveryOperation.resource_id == run.id ) ) assert operation is not None self.assertEqual(RecoveryStatus.OUTCOME_UNKNOWN.value, operation.status) if __name__ == "__main__": unittest.main()