import json import sys import types import unittest from unittest.mock import patch class _FakeCursor: def __init__(self): self.statements = [] self.lastrowid = 321 def __enter__(self): return self def __exit__(self, exc_type, exc, tb): return False def execute(self, sql, params=None): self.statements.append((sql, params)) return 1 def fetchone(self): return None class _FakeConnection: def __init__(self): self.cursor_obj = _FakeCursor() self.committed = False self.rolled_back = False self.closed = False def cursor(self): return self.cursor_obj def commit(self): self.committed = True def rollback(self): self.rolled_back = True def close(self): self.closed = True def _install_fake_pymysql(): fake_pymysql = types.ModuleType("pymysql") fake_pymysql.cursors = types.SimpleNamespace(DictCursor=object) fake_pymysql.connect = lambda **kwargs: _FakeConnection() sys.modules["pymysql"] = fake_pymysql def _valid_row(): evidence_pack = { "pattern_source_system": "pg_pattern_v2", "source_kind": "pattern_itemset", "case_id_type": "post_id", "source_certainty": "db_validated", "validation_status": "passed", "source_post_id": "p1", "pattern_execution_id": 581, "mining_config_id": 7, "evidence_sources": [ { "source_kind": "pattern_itemset", "source_tool": "get_itemset_detail", "itemset_ids": [11], "mining_config_ids": [7], "source_terms": ["x"], "source_post_id": "p1", "matched_post_count": 2, "matched_post_ids": ["p1", "p2"], "support": 0.5, "absolute_support": 2, } ], "itemset_ids": [11], "itemset_items": [{"itemset_id": 11, "category_id": 22, "category_path": "A>B"}], "category_bindings": [{"category_id": 22, "category_path": "A>B"}], "element_bindings": [{"post_id": "p1", "category_id": 22, "element_name": "x"}], "support": 0.5, "absolute_support": 2, "matched_post_ids": ["p1", "p2"], "video_ids": ["p1", "p2"], "case_ids": ["p1", "p2"], "seed_terms": ["x"], "query_seed_points": [ { "point_text": "搜索点", "point_type": "灵感点", "coverage_post_count": 2, "rank": 1, } ], "demand_scope": { "scope_source": "manual_cli", "merge_leve2": "PG Pattern V2 需求测试", "platform": "piaoquan", "pattern_execution_id": 581, }, "filtered_absolute_support": 2, "scoped_post_count": 2, "trace_id": "trace-1", } return { "merge_leve2": "PG Pattern V2 需求测试", "name": "测试需求", "reason": "原因", "suggestion": "建议", "score": 1.0, "dt": "20260606", "ext_data": {"evidence_pack": evidence_pack}, } class MySQLDemandContentSinkTest(unittest.TestCase): def setUp(self): _install_fake_pymysql() def test_insert_updates_evidence_pack_with_real_mysql_id(self): from examples.demand import mysql_demand_content_sink as sink fake_conn = _FakeConnection() with patch.object(sink, "_connect", return_value=fake_conn): result = sink.write_demand_content_rows([_valid_row()], run_label="batch01") self.assertEqual(result.inserted_count, 1) self.assertEqual(result.inserted_ids, [321]) self.assertTrue(fake_conn.committed) update_statements = [ params for sql, params in fake_conn.cursor_obj.statements if sql.strip().startswith("UPDATE demand_content") ] self.assertEqual(len(update_statements), 1) updated_ext_data = json.loads(update_statements[0][0]) self.assertEqual(updated_ext_data["run_label"], "batch01") self.assertEqual(updated_ext_data["evidence_pack"]["demand_content_id"], 321) def test_rejects_rows_without_required_pg_evidence(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["pattern_source_system"] = "mysql_topic_pattern" with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "pattern_source_system"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_rows_with_mismatched_case_or_video_ids(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["case_ids"] = ["p1"] with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "case_ids"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_rows_when_support_is_not_covered_by_posts(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["absolute_support"] = 3 row["ext_data"]["evidence_pack"].pop("filtered_absolute_support") row["ext_data"]["evidence_pack"].pop("scoped_post_count") with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "absolute_support"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_rows_without_query_seed_points(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"].pop("query_seed_points") with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "query_seed_points"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_invalid_query_seed_point_type(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["query_seed_points"][0]["point_type"] = "关键点" with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "point_type"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_mismatched_demand_scope(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["demand_scope"]["merge_leve2"] = "其他品类" with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "demand_scope"): sink.write_demand_content_rows([row], run_label="batch01") def test_rejects_non_piaoquan_demand_scope_platform(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["demand_scope"]["platform"] = "xiaohongshu" with patch.object(sink, "_connect", return_value=_FakeConnection()): with self.assertRaisesRegex(ValueError, "platform"): sink.write_demand_content_rows([row], run_label="batch01") def test_accepts_multiple_itemset_ids(self): from examples.demand import mysql_demand_content_sink as sink row = _valid_row() row["ext_data"]["evidence_pack"]["itemset_ids"] = [11, 12] row["ext_data"]["evidence_pack"]["evidence_sources"][0]["itemset_ids"] = [11, 12] with patch.object(sink, "_connect", return_value=_FakeConnection()): result = sink.write_demand_content_rows([row], run_label="batch01") self.assertEqual(result.inserted_count, 1) if __name__ == "__main__": unittest.main()