File size: 10,235 Bytes
0e61be5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
"""
Unit tests for WikiData loader.

Note: These tests require the data files to be present. Tests will be
skipped if data files are missing.
"""

import numpy as np
import pytest

from src.config import EMBEDDING_DIM, validate_data_files

# Skip all tests if data files are missing
pytestmark = pytest.mark.skipif(
    not all(validate_data_files().values()),
    reason="Data files not available",
)


@pytest.fixture(scope="module")
def wiki_data():
    """Load wiki_data once for all tests in this module."""
    from src.data import wiki_data

    # Trigger loading
    _ = wiki_data.article_count()
    return wiki_data


class TestCoreAccessors:
    """Test core accessor methods."""

    def test_article_count_positive(self, wiki_data):
        """Article count should be positive."""
        assert wiki_data.article_count() > 0

    def test_traversable_count_positive(self, wiki_data):
        """Traversable count should be positive."""
        assert wiki_data.traversable_count() > 0

    def test_traversable_less_than_total(self, wiki_data):
        """Traversable count should be less than total (not all have links)."""
        assert wiki_data.traversable_count() <= wiki_data.article_count()

    def test_get_title_valid_index(self, wiki_data):
        """Should return title for valid index."""
        title = wiki_data.get_title(0)
        assert isinstance(title, str)

    def test_get_title_invalid_index_raises(self, wiki_data):
        """Should raise IndexError for invalid index."""
        with pytest.raises(IndexError):
            wiki_data.get_title(-1)
        with pytest.raises(IndexError):
            wiki_data.get_title(999_999_999)

    def test_get_index_known_article(self, wiki_data):
        """Should return index for known article."""
        idx = wiki_data.get_index("Albert Einstein")
        assert idx is not None
        assert isinstance(idx, int)
        assert idx >= 0

    def test_get_index_unknown_article(self, wiki_data):
        """Should return None for unknown article."""
        idx = wiki_data.get_index("ZZZZZ_NOT_A_REAL_ARTICLE_12345")
        assert idx is None

    def test_has_article_true(self, wiki_data):
        """Should return True for existing article."""
        assert wiki_data.has_article("Albert Einstein") is True

    def test_has_article_false(self, wiki_data):
        """Should return False for non-existing article."""
        assert wiki_data.has_article("ZZZZZ_NOT_A_REAL_ARTICLE_12345") is False

    def test_is_traversable_true(self, wiki_data):
        """Should return True for article with outgoing links."""
        # Albert Einstein should have links
        assert wiki_data.is_traversable("Albert Einstein") is True

    def test_is_traversable_false_nonexistent(self, wiki_data):
        """Should return False for non-existent article."""
        assert wiki_data.is_traversable("ZZZZZ_NOT_A_REAL_ARTICLE_12345") is False


class TestEmbeddingAccessors:
    """Test embedding accessor methods."""

    def test_get_embedding_valid(self, wiki_data):
        """Should return embedding for valid article."""
        emb = wiki_data.get_embedding("Albert Einstein")
        assert emb is not None
        assert isinstance(emb, np.ndarray)
        assert emb.shape == (EMBEDDING_DIM,)

    def test_get_embedding_invalid(self, wiki_data):
        """Should return None for invalid article."""
        emb = wiki_data.get_embedding("ZZZZZ_NOT_A_REAL_ARTICLE_12345")
        assert emb is None

    def test_get_embedding_by_idx(self, wiki_data):
        """Should return embedding by index."""
        idx = wiki_data.get_index("Albert Einstein")
        emb = wiki_data.get_embedding_by_idx(idx)
        assert isinstance(emb, np.ndarray)
        assert emb.shape == (EMBEDDING_DIM,)

    def test_get_normalized_embedding(self, wiki_data):
        """Should return normalized embedding (L2 norm ~= 1)."""
        emb = wiki_data.get_normalized_embedding("Albert Einstein")
        assert emb is not None
        norm = np.linalg.norm(emb)
        assert 0.99 < norm < 1.01  # Should be approximately 1

    def test_get_embeddings_batch(self, wiki_data):
        """Should return batch of embeddings."""
        titles = ["Albert Einstein", "Physics", "Mathematics"]
        batch = wiki_data.get_embeddings_batch(titles)
        assert isinstance(batch, np.ndarray)
        assert batch.shape == (3, EMBEDDING_DIM)

    def test_get_embeddings_batch_partial(self, wiki_data):
        """Should handle batch with some invalid titles."""
        titles = ["Albert Einstein", "INVALID_ARTICLE", "Physics"]
        batch = wiki_data.get_embeddings_batch(titles)
        assert batch.shape[0] == 2  # Only 2 valid


class TestSimilarityFunctions:
    """Test similarity computation methods."""

    def test_similarity_valid(self, wiki_data):
        """Should return similarity for valid articles."""
        sim = wiki_data.similarity("Physics", "Mathematics")
        assert sim is not None
        assert isinstance(sim, float)
        assert -1 <= sim <= 1  # Cosine similarity range

    def test_similarity_invalid(self, wiki_data):
        """Should return None if either article is invalid."""
        assert wiki_data.similarity("Physics", "INVALID") is None
        assert wiki_data.similarity("INVALID", "Physics") is None

    def test_similarity_self(self, wiki_data):
        """Similarity of article with itself should be ~1."""
        sim = wiki_data.similarity("Albert Einstein", "Albert Einstein")
        assert sim is not None
        assert sim > 0.99

    def test_similarity_semantic_ordering(self, wiki_data):
        """Related articles should be more similar than unrelated."""
        sim_related = wiki_data.similarity("Physics", "Mathematics")
        sim_unrelated = wiki_data.similarity("Physics", "Pizza")
        assert sim_related is not None
        assert sim_unrelated is not None
        assert sim_related > sim_unrelated

    def test_rank_by_similarity(self, wiki_data):
        """Should rank candidates by similarity."""
        candidates = ["Mathematics", "Pizza", "Chemistry"]
        ranked = wiki_data.rank_by_similarity(candidates, "Physics")
        assert len(ranked) == 3
        assert all(isinstance(r, tuple) and len(r) == 2 for r in ranked)
        # Check sorted descending
        sims = [r[1] for r in ranked]
        assert sims == sorted(sims, reverse=True)

    def test_rank_by_similarity_empty(self, wiki_data):
        """Should return empty list for invalid target."""
        ranked = wiki_data.rank_by_similarity(["Physics"], "INVALID")
        assert ranked == []

    def test_nearest_neighbors(self, wiki_data):
        """Should return nearest neighbors."""
        neighbors = wiki_data.nearest_neighbors("Physics", k=5)
        assert len(neighbors) == 5
        assert all(isinstance(n, tuple) and len(n) == 2 for n in neighbors)
        # Should not include self
        assert all(n[0] != "Physics" for n in neighbors)

    def test_nearest_neighbors_invalid(self, wiki_data):
        """Should return empty list for invalid article."""
        neighbors = wiki_data.nearest_neighbors("INVALID", k=5)
        assert neighbors == []


class TestGraphAccessors:
    """Test graph accessor methods."""

    def test_get_links_valid(self, wiki_data):
        """Should return links for valid traversable article."""
        links = wiki_data.get_links("Albert Einstein")
        assert isinstance(links, list)
        assert len(links) > 0
        assert all(isinstance(link, str) for link in links)

    def test_get_links_invalid(self, wiki_data):
        """Should return empty list for invalid article."""
        links = wiki_data.get_links("INVALID_ARTICLE_12345")
        assert links == []

    def test_get_links_by_idx(self, wiki_data):
        """Should return link indices."""
        idx = wiki_data.get_index("Albert Einstein")
        if wiki_data.is_traversable("Albert Einstein"):
            links = wiki_data.get_links_by_idx(idx)
            assert isinstance(links, list)
            assert all(isinstance(i, int) for i in links)

    def test_get_popularity(self, wiki_data):
        """Should return popularity (inbound link count)."""
        # Popular articles should have high popularity
        pop = wiki_data.get_popularity("United States")
        assert isinstance(pop, int)
        assert pop >= 0
        # United States should be linked to by many articles
        # Note: This test may be slow first time due to on-demand calculation

    def test_get_popularity_invalid(self, wiki_data):
        """Should return 0 for invalid article."""
        pop = wiki_data.get_popularity("INVALID_ARTICLE_12345")
        assert pop == 0


class TestEdgeCases:
    """Test edge cases and error handling."""

    def test_empty_string_title(self, wiki_data):
        """Should handle empty string gracefully."""
        assert wiki_data.get_index("") is None or wiki_data.get_index("") is not None
        # Just shouldn't crash

    def test_unicode_title(self, wiki_data):
        """Should handle unicode titles."""
        # Try some unicode - may or may not exist
        wiki_data.get_index("日本")  # Japan in Japanese
        wiki_data.get_index("Москва")  # Moscow in Russian
        # Just shouldn't crash

    def test_special_characters(self, wiki_data):
        """Should handle special characters in titles."""
        wiki_data.get_index("C++")
        wiki_data.get_index("AC/DC")
        # Just shouldn't crash


class TestValidation:
    """Test validation and stats methods."""

    def test_validate_all_pass(self, wiki_data):
        """All validation checks should pass."""
        validation = wiki_data.validate()
        assert all(validation.values()), f"Failed checks: {validation}"

    def test_stats_complete(self, wiki_data):
        """Stats should include all expected keys."""
        stats = wiki_data.stats()
        expected_keys = [
            "total_articles",
            "traversable_articles",
            "embedding_dim",
            "embedding_dtype",
            "faiss_vectors",
        ]
        for key in expected_keys:
            assert key in stats, f"Missing key: {key}"