test_pg_pattern_repository.py 3.3 KB

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