返回 last30days-skill
test_entity_extract.py
根目录 / tests / test_entity_extract.py
1 """Tests for entity_extract module."""
2
3 import unittest
4
5 # Add lib to path
6
7 from lib import entity_extract
8
9
10 class TestExtractXHandles(unittest.TestCase):
11 def test_basic_author_handle(self):
12 items = [{"author_handle": "techguru", "text": ""}]
13 result = entity_extract._extract_x_handles(items)
14 self.assertEqual(result, ["techguru"])
15
16 def test_mentions_in_text(self):
17 items = [{"text": "Great thread by @airesearcher and @mldev"}]
18 result = entity_extract._extract_x_handles(items)
19 self.assertIn("airesearcher", result)
20 self.assertIn("mldev", result)
21
22 def test_generic_handles_filtered(self):
23 items = [
24 {"author_handle": "@openai", "text": ""},
25 {"author_handle": "@elonmusk", "text": ""},
26 {"author_handle": "realexpert", "text": ""},
27 ]
28 result = entity_extract._extract_x_handles(items)
29 self.assertEqual(result, ["realexpert"])
30
31 def test_case_normalization(self):
32 items = [{"author_handle": "@CamelCase", "text": ""}]
33 result = entity_extract._extract_x_handles(items)
34 self.assertEqual(result, ["camelcase"])
35
36 def test_frequency_ranking(self):
37 items = [
38 {"author_handle": "popular", "text": ""},
39 {"author_handle": "popular", "text": ""},
40 {"author_handle": "popular", "text": ""},
41 {"author_handle": "rare", "text": ""},
42 ]
43 result = entity_extract._extract_x_handles(items)
44 self.assertEqual(result[0], "popular")
45
46 def test_leading_at_stripped(self):
47 items = [{"author_handle": "@withatsign", "text": ""}]
48 result = entity_extract._extract_x_handles(items)
49 self.assertEqual(result, ["withatsign"])
50
51 def test_empty_input(self):
52 result = entity_extract._extract_x_handles([])
53 self.assertEqual(result, [])
54
55 def test_mixed_items(self):
56 items = [
57 {"author_handle": "poster1", "text": "Check @mentioned"},
58 {"text": "No author here"},
59 {"author_handle": "", "text": ""},
60 ]
61 result = entity_extract._extract_x_handles(items)
62 self.assertIn("poster1", result)
63 self.assertIn("mentioned", result)
64 self.assertEqual(len(result), 2)
65
66
67 class TestExtractXHashtags(unittest.TestCase):
68 def test_basic_hashtag(self):
69 items = [{"text": "Exciting news #AI"}]
70 result = entity_extract._extract_x_hashtags(items)
71 self.assertEqual(result, ["#ai"])
72
73 def test_multiple_tags(self):
74 items = [{"text": "#Python and #MachineLearning are trending"}]
75 result = entity_extract._extract_x_hashtags(items)
76 self.assertIn("#python", result)
77 self.assertIn("#machinelearning", result)
78
79 def test_frequency_ranking(self):
80 items = [
81 {"text": "#ai is great"},
82 {"text": "#ai again"},
83 {"text": "#rare tag"},
84 ]
85 result = entity_extract._extract_x_hashtags(items)
86 self.assertEqual(result[0], "#ai")
87
88 def test_single_char_tag_filtered(self):
89 items = [{"text": "#X is not enough chars but #AI is"}]
90 result = entity_extract._extract_x_hashtags(items)
91 # #X is only 1 char, filtered by \w{2,30} regex
92 self.assertNotIn("#x", result)
93 self.assertIn("#ai", result)
94
95 def test_empty_input(self):
96 result = entity_extract._extract_x_hashtags([])
97 self.assertEqual(result, [])
98
99
100 class TestExtractSubreddits(unittest.TestCase):
101 def test_basic_subreddit_field(self):
102 items = [{"subreddit": "MachineLearning"}]
103 result = entity_extract._extract_subreddits(items)
104 self.assertEqual(result, ["MachineLearning"])
105
106 def test_cross_ref_in_comment_insights(self):
107 items = [{"subreddit": "AI", "comment_insights": ["Check out r/localLLaMA for more"]}]
108 result = entity_extract._extract_subreddits(items)
109 self.assertIn("localLLaMA", result)
110
111 def test_cross_ref_in_top_comments(self):
112 items = [{"subreddit": "tech", "top_comments": [{"excerpt": "Also see r/programming"}]}]
113 result = entity_extract._extract_subreddits(items)
114 self.assertIn("programming", result)
115
116 def test_frequency_ranking(self):
117 items = [
118 {"subreddit": "popular"},
119 {"subreddit": "popular"},
120 {"subreddit": "rare"},
121 ]
122 result = entity_extract._extract_subreddits(items)
123 self.assertEqual(result[0], "popular")
124
125 def test_leading_r_slash_stripped(self):
126 items = [{"subreddit": "r/stripped"}]
127 result = entity_extract._extract_subreddits(items)
128 self.assertEqual(result, ["stripped"])
129
130 def test_empty_input(self):
131 result = entity_extract._extract_subreddits([])
132 self.assertEqual(result, [])
133
134
135 class TestExtractEntities(unittest.TestCase):
136 def test_integration(self):
137 reddit = [{"subreddit": "AI", "comment_insights": ["r/localLLaMA"]}]
138 x = [{"author_handle": "researcher", "text": "#deeplearning @colleague"}]
139 result = entity_extract.extract_entities(reddit, x)
140 self.assertIn("researcher", result["x_handles"])
141 self.assertIn("#deeplearning", result["x_hashtags"])
142 self.assertIn("AI", result["reddit_subreddits"])
143
144 def test_max_limits(self):
145 x = [
146 {"author_handle": f"user{i}", "text": ""}
147 for i in range(10)
148 ]
149 result = entity_extract.extract_entities([], x, max_handles=2)
150 self.assertLessEqual(len(result["x_handles"]), 2)
151
152 def test_empty_inputs(self):
153 result = entity_extract.extract_entities([], [])
154 self.assertEqual(result["x_handles"], [])
155 self.assertEqual(result["x_hashtags"], [])
156 self.assertEqual(result["reddit_subreddits"], [])
157
158 def test_return_keys(self):
159 result = entity_extract.extract_entities([], [])
160 self.assertSetEqual(set(result.keys()), {"x_handles", "x_hashtags", "reddit_subreddits"})
161
162 if __name__ == "__main__":
163 unittest.main()
164
164 lines PYTHON