384 lines
19 KiB
Python
384 lines
19 KiB
Python
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 <database_path> <SQL>",
|
|
)
|
|
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()
|