test_pg_pattern_repository.py 2.9 KB

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