- Adjusted the maximum number of recommended products from 6 to 4 in README.md to reflect changes in the product retrieval process. - Modified vector.py to limit the number of products retrieved per query to 4 for single queries and 2 for multiple queries, ensuring better deduplication and efficiency in product searches.
109 lines
3.5 KiB
Python
109 lines
3.5 KiB
Python
"""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()
|