test_pg_pattern_repository.py 3.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798
  1. import unittest
  2. from unittest.mock import patch
  3. from examples.demand.pg_pattern_repository import (
  4. _expand_platform_filters,
  5. _is_element_type_dimension,
  6. normalize_platform,
  7. query_elements,
  8. query_seed_points_for_itemsets,
  9. )
  10. class PgPatternRepositoryTest(unittest.TestCase):
  11. def test_topic_itemset_need_dimension_does_not_filter_element_type(self):
  12. self.assertFalse(_is_element_type_dimension("需求"))
  13. def test_real_element_dimensions_filter_element_type(self):
  14. self.assertTrue(_is_element_type_dimension("实质"))
  15. self.assertTrue(_is_element_type_dimension("形式"))
  16. self.assertTrue(_is_element_type_dimension("意图"))
  17. def test_platform_normalization_keeps_piaoquan_filter_aliases(self):
  18. self.assertEqual(normalize_platform("票圈"), "piaoquan")
  19. self.assertEqual(normalize_platform("piaoquan"), "piaoquan")
  20. self.assertEqual(_expand_platform_filters(["票圈"]), ["piaoquan", "票圈"])
  21. def test_query_seed_points_filters_and_ranks(self):
  22. rows = [
  23. {
  24. "id": 1,
  25. "post_id": "p1",
  26. "source_table": "post_decode_topic_point_element",
  27. "source_element_id": 11,
  28. "point_type": "灵感点",
  29. "point_text": "A",
  30. "element_type": "实质",
  31. "name": "A",
  32. "category_id": 10,
  33. "category_path": "/A",
  34. "topic_point_id": 101,
  35. "matched_itemset_item_id": 1001,
  36. "matched_category_id": 10,
  37. },
  38. {
  39. "id": 2,
  40. "post_id": "p2",
  41. "source_table": "post_decode_topic_point_element",
  42. "source_element_id": 12,
  43. "point_type": "灵感点",
  44. "point_text": "A",
  45. "element_type": "实质",
  46. "name": "A",
  47. "category_id": 10,
  48. "category_path": "/A",
  49. "topic_point_id": 102,
  50. "matched_itemset_item_id": 1001,
  51. "matched_category_id": 10,
  52. },
  53. {
  54. "id": 3,
  55. "post_id": "p1",
  56. "source_table": "post_decode_topic_point_element",
  57. "source_element_id": 13,
  58. "point_type": "目的点",
  59. "point_text": "B",
  60. "element_type": "实质",
  61. "name": "B",
  62. "category_id": 10,
  63. "category_path": "/B",
  64. "topic_point_id": 103,
  65. "matched_itemset_item_id": 1002,
  66. "matched_category_id": 10,
  67. },
  68. {
  69. "id": 4,
  70. "post_id": "p1",
  71. "point_type": "关键点",
  72. "point_text": "C",
  73. },
  74. ]
  75. with patch("examples.demand.pg_pattern_repository._fetch_all", return_value=rows):
  76. result = query_seed_points_for_itemsets(581, [1], ["p1", "p2"], top_k=10)
  77. self.assertEqual([item["point_text"] for item in result], ["A", "B"])
  78. self.assertEqual(result[0]["coverage_post_count"], 2)
  79. self.assertEqual(result[0]["rank"], 1)
  80. self.assertEqual(result[1]["rank"], 2)
  81. def test_query_elements_empty_post_ids_does_not_fallback_to_global(self):
  82. with patch("examples.demand.pg_pattern_repository._fetch_all") as fetch_all:
  83. result = query_elements(581, post_ids=[], merge_leve2="历史名人", platform="piaoquan")
  84. self.assertEqual(result, [])
  85. fetch_all.assert_not_called()
  86. if __name__ == "__main__":
  87. unittest.main()