test_mysql_demand_content_sink.py 7.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201
  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. "itemset_ids": [11],
  49. "itemset_items": [{"itemset_id": 11, "category_id": 22, "category_path": "A>B"}],
  50. "category_bindings": [{"category_id": 22, "category_path": "A>B"}],
  51. "element_bindings": [{"post_id": "p1", "category_id": 22, "element_name": "x"}],
  52. "support": 0.5,
  53. "absolute_support": 2,
  54. "matched_post_ids": ["p1", "p2"],
  55. "video_ids": ["p1", "p2"],
  56. "case_ids": ["p1", "p2"],
  57. "seed_terms": ["x"],
  58. "query_seed_points": [
  59. {
  60. "point_text": "搜索点",
  61. "point_type": "灵感点",
  62. "coverage_post_count": 2,
  63. "rank": 1,
  64. }
  65. ],
  66. "demand_scope": {
  67. "scope_source": "manual_cli",
  68. "merge_leve2": "PG Pattern V2 需求测试",
  69. "platform": "piaoquan",
  70. "pattern_execution_id": 581,
  71. },
  72. "filtered_absolute_support": 2,
  73. "scoped_post_count": 2,
  74. "trace_id": "trace-1",
  75. }
  76. return {
  77. "merge_leve2": "PG Pattern V2 需求测试",
  78. "name": "测试需求",
  79. "reason": "原因",
  80. "suggestion": "建议",
  81. "score": 1.0,
  82. "dt": "20260606",
  83. "ext_data": {"evidence_pack": evidence_pack},
  84. }
  85. class MySQLDemandContentSinkTest(unittest.TestCase):
  86. def setUp(self):
  87. _install_fake_pymysql()
  88. def test_insert_updates_evidence_pack_with_real_mysql_id(self):
  89. from examples.demand import mysql_demand_content_sink as sink
  90. fake_conn = _FakeConnection()
  91. with patch.object(sink, "_connect", return_value=fake_conn):
  92. result = sink.write_demand_content_rows([_valid_row()], run_label="batch01")
  93. self.assertEqual(result.inserted_count, 1)
  94. self.assertEqual(result.inserted_ids, [321])
  95. self.assertTrue(fake_conn.committed)
  96. update_statements = [
  97. params
  98. for sql, params in fake_conn.cursor_obj.statements
  99. if sql.strip().startswith("UPDATE demand_content")
  100. ]
  101. self.assertEqual(len(update_statements), 1)
  102. updated_ext_data = json.loads(update_statements[0][0])
  103. self.assertEqual(updated_ext_data["run_label"], "batch01")
  104. self.assertEqual(updated_ext_data["evidence_pack"]["demand_content_id"], 321)
  105. def test_rejects_rows_without_required_pg_evidence(self):
  106. from examples.demand import mysql_demand_content_sink as sink
  107. row = _valid_row()
  108. row["ext_data"]["evidence_pack"]["pattern_source_system"] = "mysql_topic_pattern"
  109. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  110. with self.assertRaisesRegex(ValueError, "pattern_source_system"):
  111. sink.write_demand_content_rows([row], run_label="batch01")
  112. def test_rejects_rows_with_mismatched_case_or_video_ids(self):
  113. from examples.demand import mysql_demand_content_sink as sink
  114. row = _valid_row()
  115. row["ext_data"]["evidence_pack"]["case_ids"] = ["p1"]
  116. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  117. with self.assertRaisesRegex(ValueError, "case_ids"):
  118. sink.write_demand_content_rows([row], run_label="batch01")
  119. def test_rejects_rows_when_support_is_not_covered_by_posts(self):
  120. from examples.demand import mysql_demand_content_sink as sink
  121. row = _valid_row()
  122. row["ext_data"]["evidence_pack"]["absolute_support"] = 3
  123. row["ext_data"]["evidence_pack"].pop("filtered_absolute_support")
  124. row["ext_data"]["evidence_pack"].pop("scoped_post_count")
  125. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  126. with self.assertRaisesRegex(ValueError, "absolute_support"):
  127. sink.write_demand_content_rows([row], run_label="batch01")
  128. def test_rejects_rows_without_query_seed_points(self):
  129. from examples.demand import mysql_demand_content_sink as sink
  130. row = _valid_row()
  131. row["ext_data"]["evidence_pack"].pop("query_seed_points")
  132. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  133. with self.assertRaisesRegex(ValueError, "query_seed_points"):
  134. sink.write_demand_content_rows([row], run_label="batch01")
  135. def test_rejects_invalid_query_seed_point_type(self):
  136. from examples.demand import mysql_demand_content_sink as sink
  137. row = _valid_row()
  138. row["ext_data"]["evidence_pack"]["query_seed_points"][0]["point_type"] = "关键点"
  139. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  140. with self.assertRaisesRegex(ValueError, "point_type"):
  141. sink.write_demand_content_rows([row], run_label="batch01")
  142. def test_rejects_mismatched_demand_scope(self):
  143. from examples.demand import mysql_demand_content_sink as sink
  144. row = _valid_row()
  145. row["ext_data"]["evidence_pack"]["demand_scope"]["merge_leve2"] = "其他品类"
  146. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  147. with self.assertRaisesRegex(ValueError, "demand_scope"):
  148. sink.write_demand_content_rows([row], run_label="batch01")
  149. def test_rejects_multiple_itemset_ids(self):
  150. from examples.demand import mysql_demand_content_sink as sink
  151. row = _valid_row()
  152. row["ext_data"]["evidence_pack"]["itemset_ids"] = [11, 12]
  153. with patch.object(sink, "_connect", return_value=_FakeConnection()):
  154. with self.assertRaisesRegex(ValueError, "exactly one"):
  155. sink.write_demand_content_rows([row], run_label="batch01")
  156. if __name__ == "__main__":
  157. unittest.main()