from __future__ import annotations import copy import os import pytest from sqlalchemy import create_engine, text from sqlalchemy.orm import Session from app.core.data_rules.schema_resolver import SchemaResolver from tests.core.data_rules.test_contracts import ( valid_dataflow_spec, valid_rule_spec, valid_standard_spec, ) pytestmark = pytest.mark.integration class FakeMetadataCatalog: def load_schema(self, schema_ref): return { "source_revision": "integration:1", "fields": [ { "name": "customer_id", "type": "string", "nullable": False, }, { "name": "name", "type": "string", "nullable": True, }, { "name": "mobile", "type": "string", "nullable": True, }, ], } @pytest.fixture() def database_url(): value = os.environ.get("TEST_DATABASE_URL") if not value: pytest.skip("TEST_DATABASE_URL is not configured") return value def test_postgres_rule_standard_and_production_line_release_is_atomic( database_url, ): from app.core.data_rules.release import ProductionLineReleaseService from app.core.data_rules.repository import DataRuleRepository engine = create_engine(database_url) with engine.connect() as connection: transaction = connection.begin() try: actor = connection.execute( text( "SELECT id::text FROM public.users " "WHERE status = 'active' ORDER BY created_at LIMIT 1" ) ).scalar_one() session = Session(bind=connection) repository = DataRuleRepository(session) quality_spec = valid_rule_spec() quality_spec["steps"] = [copy.deepcopy(quality_spec["steps"][1])] quality_created = repository.create_rule_version( rule_spec=quality_spec, source_text="手机号必须为11位数字", category="standard_clause", created_by=actor, ) quality = repository.publish_rule_version( version_id=quality_created["id"], published_by=actor, ) transform_spec = valid_rule_spec() transform_created = repository.create_rule_version( rule_spec=transform_spec, source_text="清洗姓名并校验手机号", category="flow_scoped", created_by=actor, ) transform = repository.publish_rule_version( version_id=transform_created["id"], published_by=actor, ) standard_spec = valid_standard_spec(quality["id"]) standard_created = repository.create_standard_version( standard_spec=standard_spec, source_text="客户手机号遵循统一格式", created_by=actor, ) standard = repository.publish_standard_version( version_id=standard_created["id"], published_by=actor, ) flow = valid_dataflow_spec(standard["id"], transform["id"]) released = ProductionLineReleaseService( repository, schema_resolver=SchemaResolver(FakeMetadataCatalog(), repository), ).release( dataflow_uid=flow["dataflow_uid"], dataflow_spec=flow, source_text="清洗客户数据并执行客户数据标准", created_by=actor, ) assert released["status"] == "released" assert released["package"]["standard_version_ids"] == [standard["id"]] assert released["package"]["rule_version_ids"] == sorted( [quality["id"], transform["id"]] ) counts = connection.execute( text( "SELECT " "(SELECT COUNT(*) FROM public.dataflow_component_bindings " " WHERE dataflow_version_id = CAST(:id AS uuid)) AS bindings, " "(SELECT COUNT(*) FROM public.rule_execution_plans p " " JOIN public.dataflow_component_bindings b " " ON b.id = p.component_binding_id " " WHERE b.dataflow_version_id = CAST(:id AS uuid) " " AND p.status = 'compiled') AS plans" ), {"id": released["id"]}, ).one() assert counts.bindings == 2 assert counts.plans == 2 finally: transaction.rollback() engine.dispose()