test_mysql_demand_content_sink.py 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226
  1. import json
  2. import sys
  3. import types
  4. import unittest
  5. from unittest.mock import patch
  6. class _FakeCursor:
  7. def __init__(self):
  8. self.statements = []
  9. self.lastrowid = 321
  10. def __enter__(self):
  11. return self
  12. def __exit__(self, exc_type, exc, tb):
  13. return False
  14. def execute(self, sql, params=None):
  15. self.statements.append((sql, params))
  16. return 1
  17. def fetchone(self):
  18. return None
  19. class _FakeConnection:
  20. def __init__(self):
  21. self.cursor_obj = _FakeCursor()
  22. self.committed = False
  23. self.rolled_back = False
  24. self.closed = False
  25. def cursor(self):
  26. return self.cursor_obj
  27. def commit(self):
  28. self.committed = True
  29. def rollback(self):
  30. self.rolled_back = True
  31. def close(self):
  32. self.closed = True
  33. def _install_fake_pymysql():
  34. fake_pymysql = types.ModuleType("pymysql")
  35. fake_pymysql.cursors = types.SimpleNamespace(DictCursor=object)
  36. fake_pymysql.connect = lambda **kwargs: _FakeConnection()
  37. sys.modules["pymysql"] = fake_pymysql
  38. def _valid_row():
  39. evidence_pack = {
  40. "pattern_source_system": "pg_pattern_v2",
  41. "source_kind": "pattern_itemset",
  42. "case_id_type": "post_id",
  43. "source_certainty": "db_validated",
  44. "validation_status": "passed",
  45. "source_post_id": "p1",
  46. "pattern_execution_id": 581,
  47. "mining_config_id": 7,
  48. "evidence_sources": [
  49. {
  50. "source_kind": "pattern_itemset",
  51. "source_tool": "get_itemset_detail",
  52. "itemset_ids": [11],
  53. "mining_config_ids": [7],
  54. "source_terms": ["x"],
  55. "source_post_id": "p1",
  56. "matched_post_count": 2,
  57. "matched_post_ids": ["p1", "p2"],
  58. "support": 0.5,
  59. "absolute_support": 2,
  60. }
  61. ],
  62. "itemset_ids": [11],
  63. "itemset_items": [{"itemset_id": 11, "category_id": 22, "category_path": "A>B"}],
  64. "category_bindings": [{"category_id": 22, "category_path": "A>B"}],
  65. "element_bindings": [{"post_id": "p1", "category_id": 22, "element_name": "x"}],
  66. "support": 0.5,
  67. "absolute_support": 2,
  68. "matched_post_ids": ["p1", "p2"],
  69. "video_ids": ["p1", "p2"],
  70. "case_ids": ["p1", "p2"],
  71. "seed_terms": ["x"],
  72. "query_seed_points": [
  73. {
  74. "point_text": "搜索点",
  75. "point_type": "灵感点",
  76. "coverage_post_count": 2,
  77. "rank": 1,
  78. }
  79. ],
  80. "demand_scope": {
  81. "scope_source": "manual_cli",
  82. "merge_leve2": "PG Pattern V2 需求测试",
  83. "platform": "piaoquan",
  84. "pattern_execution_id": 581,
  85. },
  86. "filtered_absolute_support": 2,
  87. "scoped_post_count": 2,
  88. "trace_id": "trace-1",
  89. }
  90. return {
  91. "merge_leve2": "PG Pattern V2 需求测试",
  92. "name": "测试需求",
  93. "reason": "原因",
  94. "suggestion": "建议",
  95. "score": 1.0,
  96. "dt": "20260606",
  97. "ext_data": {"evidence_pack": evidence_pack},
  98. }
  99. class MySQLDemandContentSinkTest(unittest.TestCase):
  100. def setUp(self):
  101. _install_fake_pymysql()
  102. def test_insert_updates_evidence_pack_with_real_mysql_id(self):
  103. from examples.demand import mysql_demand_content_sink as sink
  104. fake_conn = _FakeConnection()
  105. with patch.object(sink, "_connect", return_value=fake_conn):
  106. result = sink.write_demand_content_rows([_valid_row()], run_label="batch01")
  107. self.assertEqual(result.inserted_count, 1)
  108. self.assertEqual(result.inserted_ids, [321])
  109. self.assertTrue(fake_conn.committed)
  110. update_statements = [
  111. params
  112. for sql, params in fake_conn.cursor_obj.statements
  113. if sql.strip().startswith("UPDATE demand_content")
  114. ]
  115. self.assertEqual(len(update_statements), 1)
  116. updated_ext_data = json.loads(update_statements[0][0])
  117. self.assertEqual(updated_ext_data["run_label"], "batch01")
  118. self.assertEqual(updated_ext_data["evidence_pack"]["demand_content_id"], 321)
  119. def test_rejects_rows_without_required_pg_evidence(self):
  120. from examples.demand import mysql_demand_content_sink as sink
  121. row = _valid_row()
  122. row["ext_data"]["evidence_pack"]["pattern_source_system"] = "mysql_topic_pattern"
  123. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  124. with self.assertRaisesRegex(ValueError, "pattern_source_system"):
  125. sink.write_demand_content_rows([row], run_label="batch01")
  126. def test_rejects_rows_with_mismatched_case_or_video_ids(self):
  127. from examples.demand import mysql_demand_content_sink as sink
  128. row = _valid_row()
  129. row["ext_data"]["evidence_pack"]["case_ids"] = ["p1"]
  130. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  131. with self.assertRaisesRegex(ValueError, "case_ids"):
  132. sink.write_demand_content_rows([row], run_label="batch01")
  133. def test_rejects_rows_when_support_is_not_covered_by_posts(self):
  134. from examples.demand import mysql_demand_content_sink as sink
  135. row = _valid_row()
  136. row["ext_data"]["evidence_pack"]["absolute_support"] = 3
  137. row["ext_data"]["evidence_pack"].pop("filtered_absolute_support")
  138. row["ext_data"]["evidence_pack"].pop("scoped_post_count")
  139. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  140. with self.assertRaisesRegex(ValueError, "absolute_support"):
  141. sink.write_demand_content_rows([row], run_label="batch01")
  142. def test_rejects_rows_without_query_seed_points(self):
  143. from examples.demand import mysql_demand_content_sink as sink
  144. row = _valid_row()
  145. row["ext_data"]["evidence_pack"].pop("query_seed_points")
  146. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  147. with self.assertRaisesRegex(ValueError, "query_seed_points"):
  148. sink.write_demand_content_rows([row], run_label="batch01")
  149. def test_rejects_invalid_query_seed_point_type(self):
  150. from examples.demand import mysql_demand_content_sink as sink
  151. row = _valid_row()
  152. row["ext_data"]["evidence_pack"]["query_seed_points"][0]["point_type"] = "关键点"
  153. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  154. with self.assertRaisesRegex(ValueError, "point_type"):
  155. sink.write_demand_content_rows([row], run_label="batch01")
  156. def test_rejects_mismatched_demand_scope(self):
  157. from examples.demand import mysql_demand_content_sink as sink
  158. row = _valid_row()
  159. row["ext_data"]["evidence_pack"]["demand_scope"]["merge_leve2"] = "其他品类"
  160. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  161. with self.assertRaisesRegex(ValueError, "demand_scope"):
  162. sink.write_demand_content_rows([row], run_label="batch01")
  163. def test_rejects_non_piaoquan_demand_scope_platform(self):
  164. from examples.demand import mysql_demand_content_sink as sink
  165. row = _valid_row()
  166. row["ext_data"]["evidence_pack"]["demand_scope"]["platform"] = "xiaohongshu"
  167. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  168. with self.assertRaisesRegex(ValueError, "platform"):
  169. sink.write_demand_content_rows([row], run_label="batch01")
  170. def test_accepts_multiple_itemset_ids(self):
  171. from examples.demand import mysql_demand_content_sink as sink
  172. row = _valid_row()
  173. row["ext_data"]["evidence_pack"]["itemset_ids"] = [11, 12]
  174. row["ext_data"]["evidence_pack"]["evidence_sources"][0]["itemset_ids"] = [11, 12]
  175. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  176. result = sink.write_demand_content_rows([row], run_label="batch01")
  177. self.assertEqual(result.inserted_count, 1)
  178. if __name__ == "__main__":
  179. unittest.main()