"""Unit tests for Weaviate product retrieval limits.""" from __future__ import annotations import unittest from types import SimpleNamespace from typing import Any from pipeline.stages.products import ProductCandidate from pipeline.stages.vector import search_products def _candidate(product_id: str) -> ProductCandidate: return ProductCandidate( product_id=product_id, product_name=product_id.upper(), registration_number=f"reg-{product_id}", ) def _hit(product_id: str, distance: float = 0.1) -> SimpleNamespace: return SimpleNamespace( properties={"product_id": product_id, "product_name": product_id.upper()}, metadata=SimpleNamespace(distance=distance), ) class FakeNearTextQuery: def __init__(self, hits_by_query: dict[str, list[Any]]) -> None: self.hits_by_query = hits_by_query self.calls: list[dict[str, Any]] = [] def near_text( self, query: str, limit: int, filters: Any, return_metadata: Any, return_properties: list[str], ) -> SimpleNamespace: self.calls.append({"query": query, "limit": limit}) hits = list(self.hits_by_query.get(query, [])) return SimpleNamespace(objects=hits[:limit]) class FakeClient: def __init__(self, hits_by_query: dict[str, list[Any]]) -> None: self.query = FakeNearTextQuery(hits_by_query) self.collections = SimpleNamespace( get=lambda _name: SimpleNamespace(query=self.query) ) class SearchProductsLimitTests(unittest.TestCase): def test_one_query_keeps_top_four(self) -> None: candidates = [_candidate(f"p{i}") for i in range(1, 7)] hits = [_hit(f"p{i}", distance=0.1 * i) for i in range(1, 7)] client = FakeClient({"q1": hits}) results = search_products(client, ["q1"], candidates) self.assertEqual([item.product_id for item in results], ["p1", "p2", "p3", "p4"]) self.assertEqual([item.query_index for item in results], [0, 0, 0, 0]) self.assertEqual(client.query.calls, [{"query": "q1", "limit": 4}]) def test_two_queries_keep_two_each(self) -> None: candidates = [_candidate(f"p{i}") for i in range(1, 7)] client = FakeClient( { "q1": [_hit("p1"), _hit("p2"), _hit("p3")], "q2": [_hit("p4"), _hit("p5"), _hit("p6")], } ) results = search_products(client, ["q1", "q2"], candidates) self.assertEqual( [(item.product_id, item.query_index) for item in results], [("p1", 0), ("p2", 0), ("p4", 1), ("p5", 1)], ) self.assertEqual( client.query.calls, [{"query": "q1", "limit": 2}, {"query": "q2", "limit": 4}], ) def test_overlapping_hits_keep_first_occurrence(self) -> None: candidates = [_candidate(f"p{i}") for i in range(1, 6)] client = FakeClient( { "q1": [_hit("p1"), _hit("p2"), _hit("p3")], "q2": [_hit("p2"), _hit("p1"), _hit("p4"), _hit("p5")], } ) results = search_products(client, ["q1", "q2"], candidates) self.assertEqual( [(item.product_id, item.query_index) for item in results], [("p1", 0), ("p2", 0), ("p4", 1), ("p5", 1)], ) self.assertEqual( client.query.calls, [{"query": "q1", "limit": 2}, {"query": "q2", "limit": 4}], ) if __name__ == "__main__": unittest.main()