| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201 |
- 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,
- "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_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]
- with patch.object(sink, "_connect", return_value=_FakeConnection()):
- with self.assertRaisesRegex(ValueError, "exactly one"):
- sink.write_demand_content_rows([row], run_label="batch01")
- if __name__ == "__main__":
- unittest.main()
|