from __future__ import annotations import json import sqlite3 import tempfile import unittest from contextlib import closing from pathlib import Path from unittest.mock import patch from datatest.case_factory import customer_demo_cases from datatest.domain import AssertionSpec, TestCaseSpec from datatest.service import DataTestService from datatest.sqlite_source import seed_demo_databases from datatest.validation import ValidationError, validate_case, validate_read_only_sql ROOT = Path(__file__).resolve().parents[1] class DataTestCoreTests(unittest.TestCase): def setUp(self) -> None: self.temporary = tempfile.TemporaryDirectory() self.service = DataTestService(Path(self.temporary.name)) self.service.initialize_demo(ROOT / "examples/requirements/customer_etl.md") def tearDown(self) -> None: self.temporary.cleanup() def test_metadata_is_collected_before_cases(self) -> None: metadata = self.service.latest_metadata("REQ-CUSTOMER-001") self.assertEqual([item["name"] for item in metadata["databases"]["ods"]], ["ods_customer"]) self.assertEqual([item["name"] for item in metadata["databases"]["dwd"]], ["dwd_customer_info"]) self.assertEqual(len(self.service.list_cases("REQ-CUSTOMER-001")), 6) dashboard = self.service.dashboard() self.assertEqual(len(dashboard["metadata"]), 2) self.assertEqual(len(dashboard["case_items"]), 6) self.assertEqual(dashboard["case_items"][0]["requirement_name"], "客户主题 ETL 加工需求") self.assertEqual( Path(dashboard["requirements"][0]["source_path"]).name, "customer_etl.md", ) target = next(item for item in dashboard["metadata"] if item["database_name"] == "dwd") self.assertEqual(target["name"], "dwd_customer_info") self.assertEqual(target["requirement_name"], "客户主题 ETL 加工需求") self.assertEqual(len(target["columns"]), 5) def test_all_demo_cases_follow_naming_rule(self) -> None: cases = self.service.list_cases("REQ-CUSTOMER-001") self.assertTrue(all(item["name"].startswith(f"{item['table_name']}_") for item in cases)) def test_workflow_reset_preserves_document_and_sqlite_data(self) -> None: with closing(sqlite3.connect(self.service.store.source_path)) as connection: source_count = connection.execute("SELECT COUNT(*) FROM ods_customer").fetchone()[0] result = self.service.reset_requirement_workflow("REQ-CUSTOMER-001") self.assertEqual(result["status"], "imported") self.assertEqual(result["removed"]["cases"], 6) requirement = self.service.list_requirements()[0] self.assertEqual(requirement["status"], "imported") self.assertIsNone(requirement["extraction"]) self.assertFalse(requirement["metadata_ready"]) self.assertEqual(self.service.list_cases("REQ-CUSTOMER-001"), []) self.assertEqual( self.service.store.query( "SELECT * FROM etl_tasks WHERE requirement_id = 'REQ-CUSTOMER-001'" ), [], ) with self.assertRaisesRegex(ValueError, "尚未获取 Metadata"): self.service.latest_metadata("REQ-CUSTOMER-001") with closing(sqlite3.connect(self.service.store.source_path)) as connection: self.assertEqual( connection.execute("SELECT COUNT(*) FROM ods_customer").fetchone()[0], source_count, ) def test_valid_demo_run_passes_and_persists_metric(self) -> None: progress_events: list[dict[str, object]] = [] result = self.service.run_cases( "REQ-CUSTOMER-001", batch_id="2026-08-22", biz_date="2026-08-22", progress_callback=progress_events.append, ) self.assertEqual(result["status"], "PASS") self.assertEqual(progress_events[0]["event"], "run_started") self.assertEqual( len([item for item in progress_events if item["event"] == "case_pending"]), 6 ) completed = [item for item in progress_events if item["event"] == "case_completed"] self.assertEqual(len(completed), 6) self.assertTrue(all(item["status"] == "PASS" for item in completed)) self.assertEqual(progress_events[-1]["event"], "run_completed") persisted = self.service.get_run(result["run_id"]) self.assertEqual(len(persisted["results"]), 6) report = self.service.generate_report(result["run_id"]) self.assertTrue(Path(report["report_path"]).exists()) reports = self.service.dashboard()["reports"] self.assertEqual(len(reports), 1) self.assertEqual(reports[0]["requirement_id"], "REQ-CUSTOMER-001") self.assertIn("# ETL 测试报告", reports[0]["content"]) metrics = self.service.store.query("SELECT * FROM metric_snapshots") self.assertEqual(metrics[0]["metric_value"], 3) with self.assertRaisesRegex(ValueError, "只有 FAIL 或 ERROR"): self.service.analyze_failure_with_ai( result["run_id"], "CASE-001", ROOT / "schemas/failure-analysis.schema.json" ) def test_bad_transformation_returns_failure_evidence(self) -> None: with closing(sqlite3.connect(self.service.store.target_path)) as connection, connection: connection.execute( "UPDATE dwd_customer_info SET cust_status = 'BROKEN' WHERE cust_id = 2" ) result = self.service.run_cases( "REQ-CUSTOMER-001", case_ids=["CASE-003"], batch_id="bad-data" ) self.assertEqual(result["status"], "FAIL") self.assertEqual(result["results"][0]["samples"][0]["cust_id"], 2) ai_output = { "summary": "目标状态值不在允许枚举中", "suspected_layer": "target_data", "root_cause": "cust_id=2 的目标状态被写为 BROKEN", "evidence": ["invalid_count=1", "失败样例 cust_id=2"], "recommendations": ["检查状态映射逻辑"], "validation_sql": ["SELECT * FROM dwd.dwd_customer_info WHERE cust_id = 2"], "confidence": "high", } with patch("datatest.service.CodexCLIAdapter") as adapter_class: adapter = adapter_class.return_value adapter.run_structured.return_value = ai_output adapter.input_hash.return_value = "test-hash" analysis = self.service.analyze_failure_with_ai( result["run_id"], "CASE-003", ROOT / "schemas/failure-analysis.schema.json" ) self.assertEqual(analysis["analysis"]["confidence"], "high") dashboard = self.service.dashboard() self.assertEqual(dashboard["result_items"][0]["status"], "FAIL") self.assertEqual(dashboard["failure_analyses"][0]["case_id"], "CASE-003") def test_complex_demo_has_scoped_metadata_large_shapes_and_one_intentional_failure(self) -> None: initialized = self.service.initialize_complex_demo( ROOT / "examples/requirements/customer_risk_complex.md", customer_count=1_000, transaction_count=12_000, ) self.assertEqual(initialized["data"]["customers"], 1_000) self.assertEqual(initialized["data"]["transactions"], 12_000) metadata = self.service.latest_metadata("REQ-RISK-002") self.assertEqual( {item["name"] for item in metadata["databases"]["ods"]}, { "ods_customer_master_full", "ods_account_full", "ods_risk_tag_full", "ods_fx_rate_full", "ods_transaction_inc", }, ) self.assertEqual( {item["name"] for item in metadata["databases"]["dwd"]}, {"dwd_customer_risk_profile_full", "dws_customer_trade_risk_di"}, ) profile = next( item for item in metadata["databases"]["dwd"] if item["name"] == "dwd_customer_risk_profile_full" ) self.assertIn("risk_score", {item["name"] for item in profile["columns"]}) self.assertIn("profile_version", {item["name"] for item in profile["columns"]}) self.assertEqual(len(self.service.list_cases("REQ-RISK-002")), 11) result = self.service.run_cases( "REQ-RISK-002", batch_id="complex-bad-data", biz_date="2026-08-22" ) failed = [item for item in result["results"] if item["status"] != "PASS"] self.assertEqual(result["status"], "FAIL") self.assertEqual(len(failed), 1) self.assertIn("复合风险评分与等级计算一致性", failed[0]["name"]) self.assertEqual(len(failed[0]["samples"]), 1) def test_write_sql_is_rejected(self) -> None: with self.assertRaises(ValidationError): validate_read_only_sql("DELETE FROM dwd.dwd_customer_info") with self.assertRaises(ValidationError): validate_read_only_sql("SELECT 1; DROP TABLE x") with self.assertRaises(ValidationError): validate_read_only_sql("PRAGMA writable_schema = ON") def test_metadata_gate_rejects_unknown_field(self) -> None: metadata = self.service.latest_metadata("REQ-CUSTOMER-001") case = TestCaseSpec( name="dwd_customer_info_未知字段校验", requirement_id="REQ-CUSTOMER-001", requirement_version=1, etl_task_id="TASK-CUSTOMER-001", database_name="dwd", table_name="dwd_customer_info", fields=["not_exists"], category="quality", sql="SELECT COUNT(*) AS count FROM dwd.dwd_customer_info", assertions=[AssertionSpec("greater_than", "count", 0)], ) errors = validate_case(case, metadata) self.assertTrue(any("not_exists" in item for item in errors)) def test_metadata_gate_rejects_unsupported_assertion_type(self) -> None: metadata = self.service.latest_metadata("REQ-CUSTOMER-001") case = TestCaseSpec( name="dwd_customer_info_错误断言类型校验", requirement_id="REQ-CUSTOMER-001", requirement_version=1, etl_task_id="TASK-CUSTOMER-001", database_name="dwd", table_name="dwd_customer_info", fields=["cust_id"], category="quality", sql="SELECT COUNT(*) AS count FROM dwd.dwd_customer_info", assertions=[AssertionSpec("count_equals", "count", 3)], ) errors = validate_case(case, metadata) self.assertIn("第 1 个断言类型不受支持: count_equals", errors) def test_import_to_metadata_generation_review_and_execution_workflow(self) -> None: with tempfile.TemporaryDirectory() as directory: service = DataTestService(Path(directory)) seed_demo_databases(service.store.source_path, service.store.target_path) service.import_requirement( "PROJECT-FLOW", "完整流程", "REQ-FLOW-001", "客户流程需求", ROOT / "examples/requirements/customer_etl.md", ) extraction = { "requirement_name": "客户流程需求", "tasks": [{ "name": "客户加工", "sources": ["ods.ods_customer"], "targets": ["dwd.dwd_customer_info"], "field_mappings": [], "rules": ["过滤 is_deleted = 0"], }], "open_questions": [], } generated = { "cases": [{ "name": "dwd_customer_info_客户数量大于零校验", "database_name": "dwd", "table_name": "dwd_customer_info", "fields": ["cust_id"], "category": "reconciliation", "sql": "SELECT COUNT(*) AS row_count FROM dwd.dwd_customer_info", "assertions": [{"type": "greater_than", "actual": "row_count", "expected": 0}], }] } with patch("datatest.service.CodexCLIAdapter") as adapter_class: adapter = adapter_class.return_value adapter.run_structured.side_effect = [extraction, generated] adapter.input_hash.return_value = "flow-hash" service.parse_requirement_with_ai( "REQ-FLOW-001", ROOT / "schemas/requirement-extraction.schema.json" ) candidate = service.latest_metadata("REQ-FLOW-001") self.assertEqual(candidate["stage"], "candidate") self.assertEqual(candidate["requirement_scope"], [ "dwd.dwd_customer_info", "ods.ods_customer", ]) self.assertEqual( candidate["databases"]["ods"][0]["sample_rows"][0]["customer_id"], 1, ) parse_payload = adapter.run_structured.call_args_list[0].args[1] self.assertIn("database_catalog", parse_payload) self.assertEqual( parse_payload["read_only_database_access"]["command"], "sqlite3 -readonly ", ) self.assertEqual( parse_payload["database_catalog"]["databases"]["dwd"][0]["name"], "dwd_customer_info", ) pending = service.list_requirements()[0] self.assertEqual(pending["status"], "pending_confirmation") self.assertTrue(pending["metadata_ready"]) confirmed = service.confirm_requirement("REQ-FLOW-001") self.assertEqual(service.latest_metadata("REQ-FLOW-001")["stage"], "confirmed") created = service.generate_cases_with_ai( "REQ-FLOW-001", ROOT / "schemas/test-case.schema.json" ) self.assertEqual(confirmed["status"], "metadata_ready") self.assertEqual(created[0]["status"], "draft") with self.assertRaisesRegex(ValueError, "已审核"): service.run_cases("REQ-FLOW-001") service.approve_case(created[0]["id"]) result = service.run_cases("REQ-FLOW-001") self.assertEqual(result["status"], "PASS") self.assertEqual(service.list_requirements()[0]["status"], "ready") def test_agent_exploration_blocks_unknown_database_objects_before_confirmation(self) -> None: with tempfile.TemporaryDirectory() as directory: service = DataTestService(Path(directory)) seed_demo_databases(service.store.source_path, service.store.target_path) service.import_requirement( "PROJECT-FLOW", "完整流程", "REQ-FLOW-UNKNOWN", "未知表需求", ROOT / "examples/requirements/customer_etl.md", ) extraction = { "requirement_name": "未知表需求", "tasks": [{ "name": "错误范围", "sources": ["ods.ods_customer"], "targets": ["dwd.not_existing"], "field_mappings": [], "rules": [], }], "open_questions": [], } with patch("datatest.service.CodexCLIAdapter") as adapter_class: adapter = adapter_class.return_value adapter.run_structured.return_value = extraction adapter.input_hash.return_value = "unknown-hash" service.parse_requirement_with_ai( "REQ-FLOW-UNKNOWN", ROOT / "schemas/requirement-extraction.schema.json" ) self.assertEqual( service.latest_metadata("REQ-FLOW-UNKNOWN")["missing_tables"], ["dwd.not_existing"], ) with self.assertRaisesRegex(ValueError, "数据库中不存在"): service.confirm_requirement("REQ-FLOW-UNKNOWN") def test_codex_case_conversation_resets_modified_case_to_draft(self) -> None: original = self.service.store.query( "SELECT version, spec_json FROM test_cases WHERE id = 'CASE-001'" )[0] spec = json.loads(original["spec_json"]) response = { "assistant_message": "已将数量校验改为显式大于零,请重新审核。", "cases": [{ "case_id": "CASE-001", "name": "dwd_customer_info_目标表客户数量大于零校验", "database_name": "dwd", "table_name": "dwd_customer_info", "fields": ["cust_id"], "category": "reconciliation", "sql": "SELECT COUNT(*) AS row_count FROM dwd.dwd_customer_info", "assertions": [{"type": "greater_than", "actual": "row_count", "expected": 0}], "sample_sql": spec.get("sample_sql"), "sample_limit": 100, }], "removed_case_ids": [], } progress_events: list[dict[str, object]] = [] with patch("datatest.service.CodexCLIAdapter") as adapter_class: adapter = adapter_class.return_value def stream_response( _instruction: str, _payload: dict[str, object], _schema: Path, on_event: object ) -> dict[str, object]: on_event({"type": "thread.started", "thread_id": "demo"}) on_event({"type": "turn.started"}) on_event({"type": "item.completed", "item": { "type": "reasoning", "text": "不应传到 UI 的内部内容", }}) on_event({"type": "turn.completed", "usage": {}}) return response adapter.run_structured_streaming.side_effect = stream_response adapter.input_hash.return_value = "chat-hash" result = self.service.chat_about_cases_with_ai( "REQ-CUSTOMER-001", "把数量校验改成大于零", ROOT / "schemas/case-agent-response.schema.json", progress_events.append, ) self.assertEqual(result["changed_cases"][0]["status"], "draft") changed = self.service.store.query("SELECT * FROM test_cases WHERE id = 'CASE-001'")[0] self.assertEqual(changed["version"], original["version"] + 1) self.assertEqual(changed["status"], "draft") with self.assertRaisesRegex(ValueError, "已审核"): self.service.run_cases("REQ-CUSTOMER-001", case_ids=["CASE-001"]) messages = self.service.dashboard()["case_agent_messages"] self.assertEqual([item["role"] for item in messages], ["user", "assistant"]) self.assertEqual(progress_events[0]["phase"], "context") self.assertEqual(progress_events[-1]["phase"], "persistence") self.assertNotIn( "不应传到 UI 的内部内容", " ".join(str(item) for item in progress_events), ) self.service.approve_case("CASE-001") self.assertEqual( self.service.store.query("SELECT status FROM test_cases WHERE id = 'CASE-001'")[0]["status"], "approved", ) if __name__ == "__main__": unittest.main()