import unittest from unittest.mock import patch from examples.demand.pg_pattern_repository import ( _expand_platform_filters, _is_element_type_dimension, normalize_platform, query_elements, query_seed_points_for_itemsets, ) class PgPatternRepositoryTest(unittest.TestCase): def test_topic_itemset_need_dimension_does_not_filter_element_type(self): self.assertFalse(_is_element_type_dimension("需求")) def test_real_element_dimensions_filter_element_type(self): self.assertTrue(_is_element_type_dimension("实质")) self.assertTrue(_is_element_type_dimension("形式")) self.assertTrue(_is_element_type_dimension("意图")) def test_platform_normalization_keeps_piaoquan_filter_aliases(self): self.assertEqual(normalize_platform("票圈"), "piaoquan") self.assertEqual(normalize_platform("piaoquan"), "piaoquan") self.assertEqual(_expand_platform_filters(["票圈"]), ["piaoquan", "票圈"]) def test_query_seed_points_filters_and_ranks(self): rows = [ { "id": 1, "post_id": "p1", "source_table": "post_decode_topic_point_element", "source_element_id": 11, "point_type": "灵感点", "point_text": "A", "element_type": "实质", "name": "A", "category_id": 10, "category_path": "/A", "topic_point_id": 101, "matched_itemset_item_id": 1001, "matched_category_id": 10, }, { "id": 2, "post_id": "p2", "source_table": "post_decode_topic_point_element", "source_element_id": 12, "point_type": "灵感点", "point_text": "A", "element_type": "实质", "name": "A", "category_id": 10, "category_path": "/A", "topic_point_id": 102, "matched_itemset_item_id": 1001, "matched_category_id": 10, }, { "id": 3, "post_id": "p1", "source_table": "post_decode_topic_point_element", "source_element_id": 13, "point_type": "目的点", "point_text": "B", "element_type": "实质", "name": "B", "category_id": 10, "category_path": "/B", "topic_point_id": 103, "matched_itemset_item_id": 1002, "matched_category_id": 10, }, { "id": 4, "post_id": "p1", "point_type": "关键点", "point_text": "C", }, ] with patch("examples.demand.pg_pattern_repository._fetch_all", return_value=rows): result = query_seed_points_for_itemsets(581, [1], ["p1", "p2"], top_k=10) self.assertEqual([item["point_text"] for item in result], ["A", "B"]) self.assertEqual(result[0]["coverage_post_count"], 2) self.assertEqual(result[0]["rank"], 1) self.assertEqual(result[1]["rank"], 2) def test_query_elements_empty_post_ids_does_not_fallback_to_global(self): with patch("examples.demand.pg_pattern_repository._fetch_all") as fetch_all: result = query_elements(581, post_ids=[], merge_leve2="历史名人", platform="piaoquan") self.assertEqual(result, []) fetch_all.assert_not_called() if __name__ == "__main__": unittest.main()