AI_Agro_Support/tests/test_vector.py
Arsham Mirehvandi 253a6d853d Update product retrieval limits in README and vector search logic
- 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.
2026-09-16 23:18:23 +02:00

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()