updates on catalog search and suggestions in backend
This commit is contained in:
@@ -5,25 +5,44 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Query
|
||||
|
||||
from app.api.schemas import SearchOut, SourceProductOut
|
||||
from app.services.rag_service import retrieve
|
||||
from app.infrastructure.settings import SEARCH_DEFAULT_TOP_K, SEARCH_MAX_TOP_K
|
||||
from app.services.catalog_search import search_catalog
|
||||
|
||||
router = APIRouter(tags=["search"])
|
||||
|
||||
|
||||
@router.get("/search", response_model=SearchOut)
|
||||
def semantic_search(
|
||||
def catalog_search_endpoint(
|
||||
q: str = Query(..., min_length=1, max_length=500, description="Free-text search query"),
|
||||
brand: Optional[str] = Query(None, description="Restrict search to a single brand"),
|
||||
category: Optional[str] = Query(None, description="Restrict search to a category"),
|
||||
top_k: int = Query(10, ge=1, le=50),
|
||||
top_k: int = Query(SEARCH_DEFAULT_TOP_K, ge=1, le=SEARCH_MAX_TOP_K),
|
||||
offset: int = Query(0, ge=0, description="Pagination offset (brand listings only)"),
|
||||
) -> SearchOut:
|
||||
"""Pure vector similarity search over the catalog - no LLM call, just
|
||||
pgvector ranking. This is what powers the instant search-as-you-type
|
||||
grid in the React 'Search' tab. For a conversational, LLM-generated
|
||||
answer use POST /api/chat instead."""
|
||||
results = retrieve(q, brand=brand, top_k=top_k, category=category)
|
||||
"""Catalog search for the React 'Browse & Search' grid - no LLM call.
|
||||
|
||||
Two behaviours, chosen from the query text:
|
||||
|
||||
* A bare brand name ("Amul", "Colgate", "coke") returns that brand's whole
|
||||
catalog - the same rows as GET /api/brands/{brand}/products.
|
||||
* Anything else ("Amul Butter", "low sugar biscuit") ranks name matches
|
||||
first, then semantically similar products.
|
||||
|
||||
This does NOT share rag_service.retrieve() with /api/chat: that path clamps
|
||||
results to RAG_MAX_TOP_K to protect the LLM prompt budget, which is why a
|
||||
brand search used to come back with only 15 products.
|
||||
|
||||
For a conversational, LLM-generated answer use POST /api/chat instead.
|
||||
"""
|
||||
result = search_catalog(q, brand=brand, category=category, limit=top_k, offset=offset)
|
||||
return SearchOut(
|
||||
query=q,
|
||||
brand=brand,
|
||||
results=[SourceProductOut(**r.to_dict()) for r in results],
|
||||
results=[SourceProductOut(**p.to_dict()) for p in result.products],
|
||||
total=result.total,
|
||||
limit=result.limit,
|
||||
offset=result.offset,
|
||||
match_mode=result.match_mode,
|
||||
detected_brand=result.detected_brand,
|
||||
detected_category=result.detected_category,
|
||||
)
|
||||
|
||||
35
app/api/routers/suggest.py
Normal file
35
app/api/routers/suggest.py
Normal file
@@ -0,0 +1,35 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Query
|
||||
|
||||
from app.api.schemas import SuggestOut, SuggestionOut
|
||||
from app.infrastructure.settings import SUGGEST_DEFAULT_LIMIT
|
||||
from app.services.suggest_service import suggest as suggest_service
|
||||
|
||||
router = APIRouter(tags=["search"])
|
||||
|
||||
|
||||
@router.get("/suggest", response_model=SuggestOut)
|
||||
def search_suggest(
|
||||
q: str = Query(..., min_length=1, max_length=64, description="Partial search text"),
|
||||
limit: int = Query(SUGGEST_DEFAULT_LIMIT, ge=1, le=20),
|
||||
) -> SuggestOut:
|
||||
"""Autocomplete for the catalog search box: brand and category names.
|
||||
|
||||
Answers from in-process caches, so it is safe to call on every keystroke.
|
||||
Product names are deliberately not suggested - brand tables have no
|
||||
trigram index, so that would mean an unindexed scan per keystroke.
|
||||
|
||||
Public, matching GET /api/search.
|
||||
"""
|
||||
results = suggest_service(q, limit=limit)
|
||||
return SuggestOut(
|
||||
query=q,
|
||||
suggestions=[
|
||||
SuggestionOut(
|
||||
type=s.type, value=s.value, label=s.label,
|
||||
sublabel=s.sublabel, score=s.score,
|
||||
)
|
||||
for s in results
|
||||
],
|
||||
)
|
||||
@@ -126,8 +126,34 @@ class AllProductsOut(BaseModel):
|
||||
|
||||
class SearchOut(BaseModel):
|
||||
query: str
|
||||
# Echoes the *request* param, as it always has. The brand inferred from the
|
||||
# query text goes in `detected_brand` instead - repurposing this field would
|
||||
# break any consumer reading it as "the filter I sent".
|
||||
brand: Optional[str] = None
|
||||
results: List[SourceProductOut]
|
||||
# Exact in brand_catalog mode. None in hybrid mode: the merge happens in
|
||||
# Python across N brand tables after per-table LIMITs, so there is no cheap
|
||||
# exact count and inventing one would misreport how much was found.
|
||||
total: Optional[int] = None
|
||||
limit: int = 0
|
||||
offset: int = 0
|
||||
match_mode: str = "hybrid" # "brand_catalog" | "hybrid"
|
||||
detected_brand: Optional[str] = None
|
||||
detected_category: Optional[str] = None
|
||||
|
||||
|
||||
class SuggestionOut(BaseModel):
|
||||
"""One row in the search box's autocomplete dropdown."""
|
||||
type: str # "brand" | "category"
|
||||
value: str # what the search box / filter should use
|
||||
label: str # display text
|
||||
sublabel: Optional[str] = None # e.g. "128 products"
|
||||
score: float = 0.0
|
||||
|
||||
|
||||
class SuggestOut(BaseModel):
|
||||
query: str
|
||||
suggestions: List[SuggestionOut]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -357,3 +357,24 @@ RAG_MAX_CONTEXT_CHARS = int(os.getenv("RAG_MAX_CONTEXT_CHARS", "4000"))
|
||||
# if you want an extra cutoff on top of that.
|
||||
_raw_max_distance = os.getenv("RAG_MAX_DISTANCE", "").strip()
|
||||
RAG_MAX_DISTANCE = float(_raw_max_distance) if _raw_max_distance else None
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Catalog search (GET /api/search) and suggest (GET /api/suggest)
|
||||
# ---------------------------------------------------------------------------
|
||||
# Deliberately separate from RAG_MAX_TOP_K above. That ceiling exists to protect
|
||||
# the LLM prompt budget in /api/chat, and raising it would degrade every chat
|
||||
# answer. /api/search feeds a product grid, which has no such budget - sharing
|
||||
# the constant is what silently truncated every brand search to 15 products.
|
||||
SEARCH_DEFAULT_TOP_K = int(os.getenv("SEARCH_DEFAULT_TOP_K", "60"))
|
||||
# Mirrors the browse endpoints' le=100000 so a brand search can return exactly
|
||||
# the same set as GET /api/brands/{brand}/products.
|
||||
SEARCH_MAX_TOP_K = int(os.getenv("SEARCH_MAX_TOP_K", "100000"))
|
||||
# Separate, much smaller ceiling for the hybrid path: semantic_search multiplies
|
||||
# top_k by 5 per brand table when a category/price filter is present, and every
|
||||
# read does SELECT * (which drags the vector(384) embedding column over the wire).
|
||||
SEARCH_HYBRID_MAX_TOP_K = int(os.getenv("SEARCH_HYBRID_MAX_TOP_K", "100"))
|
||||
SEARCH_LEXICAL_CANDIDATES = int(os.getenv("SEARCH_LEXICAL_CANDIDATES", "200"))
|
||||
|
||||
SUGGEST_DEFAULT_LIMIT = int(os.getenv("SUGGEST_DEFAULT_LIMIT", "8"))
|
||||
SUGGEST_MIN_QUERY_LEN = int(os.getenv("SUGGEST_MIN_QUERY_LEN", "2"))
|
||||
SUGGEST_FUZZY_MIN_RATIO = float(os.getenv("SUGGEST_FUZZY_MIN_RATIO", "0.72"))
|
||||
|
||||
@@ -21,7 +21,7 @@ from fastapi.responses import FileResponse
|
||||
|
||||
from app.infrastructure.persistence import restore_bundled_assets
|
||||
from app.infrastructure.settings import API_CORS_ORIGINS, BRAND_SYNC_INTERVAL_SECONDS
|
||||
from app.api.routers import health, brands, search, chat, catalog, system
|
||||
from app.api.routers import health, brands, search, suggest, chat, catalog, system
|
||||
from app.api.routers import stores, discounts, analytics as store_analytics, trending, recommendations, store_admin
|
||||
from app.api.routers import nutrition, nutrition_admin, upload
|
||||
from app.api.routers import auth, user_products, admin_train, mcp_info
|
||||
@@ -193,6 +193,7 @@ app.include_router(admin_train.router, prefix="/api")
|
||||
app.include_router(system.router, prefix="/api")
|
||||
app.include_router(brands.router, prefix="/api")
|
||||
app.include_router(search.router, prefix="/api")
|
||||
app.include_router(suggest.router, prefix="/api")
|
||||
app.include_router(chat.router, prefix="/api")
|
||||
app.include_router(catalog.router, prefix="/api")
|
||||
app.include_router(stores.router, prefix="/api")
|
||||
|
||||
240
app/services/catalog_search.py
Normal file
240
app/services/catalog_search.py
Normal file
@@ -0,0 +1,240 @@
|
||||
"""Search behind the catalog page's search box (GET /api/search).
|
||||
|
||||
Deliberately separate from rag_service.retrieve(), which serves /api/chat.
|
||||
retrieve() clamps results to RAG_MAX_TOP_K (15) to protect the LLM's prompt
|
||||
budget; sharing that path is what silently truncated every brand search to 15
|
||||
products even though the UI asked for more. A product grid has no prompt
|
||||
budget, so it gets its own entry point and its own ceiling.
|
||||
|
||||
Two shapes are handled differently, which is the whole point of the module:
|
||||
|
||||
"Amul" -> brand_catalog mode: the brand's entire listing, exactly the
|
||||
rows GET /api/brands/{brand}/products returns.
|
||||
"Amul Butter" -> hybrid mode: name matches first (so every size variant is
|
||||
present and ranked together), then semantically similar
|
||||
products.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from app.infrastructure.settings import (
|
||||
SEARCH_DEFAULT_TOP_K,
|
||||
SEARCH_HYBRID_MAX_TOP_K,
|
||||
SEARCH_LEXICAL_CANDIDATES,
|
||||
)
|
||||
from app.services.category_registry import detect_category_from_text
|
||||
from app.services.embeddings_service import embed_texts
|
||||
from app.services.query_intent import (
|
||||
classify_query,
|
||||
extract_attributes,
|
||||
extract_max_price,
|
||||
)
|
||||
from app.services.rag_service import (
|
||||
RetrievedProduct,
|
||||
filter_to_category,
|
||||
rerank_by_attributes,
|
||||
row_to_retrieved_product,
|
||||
)
|
||||
from app.services.vector_store import (
|
||||
_sanitize_name,
|
||||
count_products_by_brand,
|
||||
get_products_by_brand,
|
||||
lexical_search,
|
||||
semantic_search,
|
||||
text_search,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Fallback similarity per lexical tier, used only when the embedding model is
|
||||
# unavailable and a real cosine distance cannot be computed. Mirrors
|
||||
# RetrievedProduct.similarity == 1 - distance/2.
|
||||
_LEX_TIER_SCORE = {0: 1.00, 1: 0.92, 2: 0.80, 3: 0.70}
|
||||
|
||||
# Words that add nothing to a lexical name match.
|
||||
_TERM_STOP_WORDS = {
|
||||
"a", "an", "the", "of", "for", "with", "and", "in", "on",
|
||||
"show", "me", "find", "get", "all", "any", "please",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class CatalogSearchResult:
|
||||
products: List[RetrievedProduct]
|
||||
match_mode: str # "brand_catalog" | "hybrid"
|
||||
detected_brand: Optional[str] = None
|
||||
detected_category: Optional[str] = None
|
||||
total: Optional[int] = None # exact only in brand_catalog mode
|
||||
limit: int = 0
|
||||
offset: int = 0
|
||||
|
||||
|
||||
def _row_key(row: Dict[str, Any]) -> tuple:
|
||||
"""Dedup key for a product row.
|
||||
|
||||
image_id is UNIQUE per brand table but not across them, so the brand has
|
||||
to be part of the key for a cross-brand merge to be correct.
|
||||
"""
|
||||
return (_sanitize_name(str(row.get("brand") or "")), str(row.get("image_id") or ""))
|
||||
|
||||
|
||||
def _tokenize(text: str) -> List[str]:
|
||||
return [t for t in text.lower().split() if t and t not in _TERM_STOP_WORDS]
|
||||
|
||||
|
||||
def _brand_catalog(
|
||||
brand: str, category: Optional[str], limit: int, offset: int
|
||||
) -> CatalogSearchResult:
|
||||
"""The brand's full listing - the same rows the Browse tab shows.
|
||||
|
||||
Category is NOT auto-detected here. detect_category_from_text() has a fuzzy
|
||||
fallback that misfires on brand names ("Colgate" scores as Chocolates,
|
||||
"Milky Mist" as Dairy), and applying that to a brand-only query filters most
|
||||
of the catalog away. Only an explicit category filter applies.
|
||||
"""
|
||||
rows = get_products_by_brand(brand, limit=limit, offset=offset, category=category)
|
||||
total = count_products_by_brand(brand, category=category)
|
||||
|
||||
for row in rows:
|
||||
row["brand"] = row.get("brand") or brand
|
||||
row["distance"] = 0.0
|
||||
|
||||
return CatalogSearchResult(
|
||||
products=[row_to_retrieved_product(r) for r in rows],
|
||||
match_mode="brand_catalog",
|
||||
detected_brand=brand,
|
||||
detected_category=category,
|
||||
total=total,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
|
||||
def _hybrid_rank(
|
||||
semantic_rows: List[Dict[str, Any]],
|
||||
lexical_rows: List[Dict[str, Any]],
|
||||
limit: int,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Merge both arms, name matches first, semantic similarity within each tier."""
|
||||
merged: Dict[tuple, Dict[str, Any]] = {}
|
||||
|
||||
# Semantic rows first: they carry a real cosine distance worth keeping.
|
||||
for row in semantic_rows:
|
||||
merged[_row_key(row)] = row
|
||||
|
||||
for row in lexical_rows:
|
||||
key = _row_key(row)
|
||||
existing = merged.get(key)
|
||||
if existing is not None:
|
||||
# Same product from both arms: keep the real distance, take the tier.
|
||||
existing["lex_tier"] = row.get("lex_tier", 3)
|
||||
else:
|
||||
if row.get("distance") is None:
|
||||
tier = row.get("lex_tier", 3)
|
||||
row["distance"] = 2.0 * (1.0 - _LEX_TIER_SCORE.get(tier, 0.7))
|
||||
merged[key] = row
|
||||
|
||||
ordered = sorted(
|
||||
merged.values(),
|
||||
key=lambda r: (
|
||||
r.get("lex_tier", 9), # any name match outranks semantic-only
|
||||
r.get("distance", 9.0), # then vector similarity within the tier
|
||||
len(str(r.get("product_name") or r.get("title") or "")),
|
||||
str(r.get("product_name") or ""),
|
||||
),
|
||||
)
|
||||
return ordered[:limit]
|
||||
|
||||
|
||||
def _hybrid(
|
||||
query: str,
|
||||
brand: Optional[str],
|
||||
category: Optional[str],
|
||||
residual: str,
|
||||
limit: int,
|
||||
) -> CatalogSearchResult:
|
||||
"""Name matches first, then semantically similar products."""
|
||||
# Detect the category from the residual, never the whole query: the brand
|
||||
# token itself can fuzzy-match a category keyword ("Colgate" -> Chocolates).
|
||||
target_category = category or detect_category_from_text(residual or query)
|
||||
max_price = extract_max_price(query)
|
||||
hybrid_limit = min(limit, SEARCH_HYBRID_MAX_TOP_K)
|
||||
|
||||
try:
|
||||
vectors = embed_texts([query])
|
||||
except Exception as e: # noqa: BLE001 - fall back to lexical-only
|
||||
logger.warning("Embedding model failed: %s. Lexical-only search.", e)
|
||||
vectors = None
|
||||
|
||||
query_embedding = vectors[0] if vectors else None
|
||||
|
||||
terms = _tokenize(residual or query)
|
||||
lexical_rows: List[Dict[str, Any]] = []
|
||||
if terms:
|
||||
lexical_rows = lexical_search(
|
||||
terms,
|
||||
brand=brand,
|
||||
limit=SEARCH_LEXICAL_CANDIDATES,
|
||||
category=target_category,
|
||||
max_price=max_price,
|
||||
query_embedding=query_embedding,
|
||||
exact_phrase=" ".join(terms),
|
||||
)
|
||||
lexical_rows = filter_to_category(lexical_rows, target_category)
|
||||
|
||||
semantic_rows: List[Dict[str, Any]] = []
|
||||
if query_embedding is not None:
|
||||
semantic_rows = semantic_search(
|
||||
query_embedding=query_embedding,
|
||||
brand=brand,
|
||||
top_k=hybrid_limit,
|
||||
category=target_category,
|
||||
max_price=max_price,
|
||||
)
|
||||
semantic_rows = filter_to_category(semantic_rows, target_category)
|
||||
|
||||
if not semantic_rows and not lexical_rows:
|
||||
semantic_rows = filter_to_category(
|
||||
text_search(query, brand=brand, top_k=hybrid_limit,
|
||||
category=target_category, max_price=max_price),
|
||||
target_category,
|
||||
)
|
||||
|
||||
rows = _hybrid_rank(semantic_rows, lexical_rows, limit)
|
||||
products = [row_to_retrieved_product(r) for r in rows]
|
||||
|
||||
attrs = extract_attributes(query)
|
||||
if attrs and products:
|
||||
products = rerank_by_attributes(products, attrs)
|
||||
|
||||
return CatalogSearchResult(
|
||||
products=products,
|
||||
match_mode="hybrid",
|
||||
detected_brand=brand,
|
||||
detected_category=target_category,
|
||||
total=None, # no cheap exact count across N tables; do not invent one
|
||||
limit=limit,
|
||||
offset=0,
|
||||
)
|
||||
|
||||
|
||||
def search_catalog(
|
||||
query: str,
|
||||
brand: Optional[str] = None,
|
||||
category: Optional[str] = None,
|
||||
limit: int = SEARCH_DEFAULT_TOP_K,
|
||||
offset: int = 0,
|
||||
) -> CatalogSearchResult:
|
||||
"""Entry point for GET /api/search."""
|
||||
shape = classify_query(query, explicit_brand=brand)
|
||||
# An explicit filter from the sidebar always wins over one inferred from the
|
||||
# query text - same precedence retrieve() uses.
|
||||
effective_brand = brand or shape.brand
|
||||
|
||||
if shape.kind == "brand_only" and effective_brand:
|
||||
return _brand_catalog(effective_brand, category, limit, offset)
|
||||
|
||||
return _hybrid(query, effective_brand, category, shape.residual, limit)
|
||||
@@ -16,6 +16,7 @@ from __future__ import annotations
|
||||
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from app.services.category_registry import detect_category_from_text # noqa: F401 (re-exported)
|
||||
@@ -253,3 +254,134 @@ def extract_brand_mention(query: str) -> Optional[str]:
|
||||
def detected_category(query: str) -> Optional[str]:
|
||||
"""Thin wrapper kept for readability at call sites in rag_service."""
|
||||
return detect_category_from_text(query)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Search-box query classification
|
||||
# ---------------------------------------------------------------------------
|
||||
# Used by catalog_search (GET /api/search) to tell three shapes apart:
|
||||
#
|
||||
# "Amul" -> the user wants the brand's whole catalog
|
||||
# "Amul Butter" -> the user wants to narrow *within* a brand
|
||||
# "low sugar biscuit" -> no brand involved
|
||||
#
|
||||
# This deliberately does NOT consult brand_registry.BRAND_ALIASES. That table
|
||||
# maps sub-brands to their storage parent and contains entries like
|
||||
# "amul butter" -> "amul", so resolve_parent_brand() cannot tell "Amul" from
|
||||
# "Amul Butter" - using it here would dump the entire Amul catalog for a query
|
||||
# that asked for butter. Only the brand_search_map from _brand_index() (real
|
||||
# brand names plus hand-tuned synonyms like "coke"/"hul") is authoritative for
|
||||
# the "is this query nothing but a brand?" question.
|
||||
|
||||
# Words that carry no product meaning, so stripping them still leaves a
|
||||
# brand-only query. Deliberately tiny and explicit: no product noun may ever
|
||||
# appear here, or "Amul Butter" would collapse to "Amul".
|
||||
_BRAND_QUERY_FILLERS = {
|
||||
"show", "me", "all", "the", "a", "an", "products", "product",
|
||||
"items", "item", "from", "of", "by", "brand", "list", "catalog",
|
||||
"catalogue", "please", "everything", "in", "for",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QueryShape:
|
||||
"""How a search-box query relates to the brand catalog."""
|
||||
kind: str # "brand_only" | "brand_plus" | "generic"
|
||||
brand: Optional[str] = None # canonical display name, e.g. "Colgate-Palmolive"
|
||||
residual: str = "" # query minus the brand surface, e.g. "butter"
|
||||
|
||||
|
||||
def _normalize_query(text: str) -> str:
|
||||
"""Lowercase, drop trailing punctuation, collapse whitespace."""
|
||||
if not text:
|
||||
return ""
|
||||
lowered = text.lower().strip().rstrip("?!.,")
|
||||
return re.sub(r"\s+", " ", lowered).strip()
|
||||
|
||||
|
||||
def _lookup_variants(norm: str) -> list:
|
||||
"""Spelling variants to try against the brand map, most literal first.
|
||||
|
||||
Covers "coca-cola" / "coca cola" / "cocacola" and "p&g" / "pg" without
|
||||
needing a row in the map for each.
|
||||
"""
|
||||
variants = [norm]
|
||||
spaced = re.sub(r"[-._&]+", " ", norm)
|
||||
spaced = re.sub(r"\s+", " ", spaced).strip()
|
||||
squashed = re.sub(r"[-._&\s]+", "", norm)
|
||||
for candidate in (spaced, squashed):
|
||||
if candidate and candidate not in variants:
|
||||
variants.append(candidate)
|
||||
return variants
|
||||
|
||||
|
||||
def _strip_fillers(norm: str) -> str:
|
||||
kept = [w for w in norm.split() if w not in _BRAND_QUERY_FILLERS]
|
||||
return " ".join(kept)
|
||||
|
||||
|
||||
def _is_whole_query_a_category_word(norm: str) -> bool:
|
||||
"""True when the query is itself a category keyword, e.g. "butter".
|
||||
|
||||
Safety valve for a brand whose name is also a product word: without it, a
|
||||
brand literally named "Butter" would make the bare query "butter" list that
|
||||
brand's whole catalog instead of searching for butter.
|
||||
"""
|
||||
if not norm:
|
||||
return False
|
||||
from app.services.category_registry import CATEGORY_REGISTRY
|
||||
|
||||
for entry in CATEGORY_REGISTRY:
|
||||
for keyword in entry["keywords"]:
|
||||
if norm == keyword.lower():
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def classify_query(query: str, explicit_brand: Optional[str] = None) -> QueryShape:
|
||||
"""Classify a search-box query as brand-only, brand-plus-terms, or generic."""
|
||||
norm = _normalize_query(query)
|
||||
if not norm:
|
||||
# An empty query under an explicit brand filter is still "show me that brand".
|
||||
return QueryShape("brand_only", explicit_brand, "") if explicit_brand \
|
||||
else QueryShape("generic", None, "")
|
||||
|
||||
_known, brand_map = _brand_index()
|
||||
|
||||
# 1. The whole query (optionally minus filler words) IS a brand.
|
||||
stripped = _strip_fillers(norm) or norm
|
||||
for candidate in (norm, stripped):
|
||||
if _is_whole_query_a_category_word(candidate):
|
||||
break
|
||||
for variant in _lookup_variants(candidate):
|
||||
hit = brand_map.get(variant)
|
||||
if hit:
|
||||
return QueryShape("brand_only", hit, "")
|
||||
|
||||
# 2. A brand is mentioned alongside other terms -> narrow within that brand.
|
||||
brand = extract_brand_mention(query)
|
||||
if brand:
|
||||
residual = _residual_after_brand(norm, brand, brand_map)
|
||||
if not _strip_fillers(residual):
|
||||
# e.g. "show me products of amul" - the leftovers were all filler.
|
||||
return QueryShape("brand_only", brand, "")
|
||||
return QueryShape("brand_plus", brand, residual)
|
||||
|
||||
return QueryShape("generic", None, norm)
|
||||
|
||||
|
||||
def _residual_after_brand(norm: str, brand: str, brand_map: Dict[str, str]) -> str:
|
||||
"""Remove the matched brand surface form from `norm`, leaving the rest.
|
||||
|
||||
Only surfaces that resolve to `brand` are removed, longest first, so
|
||||
"cadbury dairy milk" keeps "dairy milk" rather than losing it to a
|
||||
sub-brand alias.
|
||||
"""
|
||||
surfaces = [alias for alias, canonical in brand_map.items() if canonical == brand]
|
||||
if brand.lower() not in surfaces:
|
||||
surfaces.append(brand.lower())
|
||||
for surface in sorted(surfaces, key=len, reverse=True):
|
||||
pattern = r"\b" + re.escape(surface) + r"\b"
|
||||
if re.search(pattern, norm):
|
||||
return re.sub(r"\s+", " ", re.sub(pattern, " ", norm, count=1)).strip()
|
||||
return norm
|
||||
|
||||
@@ -349,3 +349,16 @@ def answer_query(
|
||||
answer=answer_text, sources=products, query=query, brand=target_brand or brand,
|
||||
detected_category=target_category,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public aliases for catalog_search (GET /api/search)
|
||||
# ---------------------------------------------------------------------------
|
||||
# catalog_search runs its own retrieval so that raising the result ceiling for
|
||||
# the product grid cannot affect the chat prompt budget enforced by retrieve()
|
||||
# above. It reuses these helpers rather than reimplementing them, so the two
|
||||
# endpoints can never disagree about image URL resolution, price coercion, or
|
||||
# category safety.
|
||||
row_to_retrieved_product = _row_to_retrieved_product
|
||||
filter_to_category = _filter_to_category
|
||||
rerank_by_attributes = _rerank_by_attributes
|
||||
|
||||
211
app/services/suggest_service.py
Normal file
211
app/services/suggest_service.py
Normal file
@@ -0,0 +1,211 @@
|
||||
"""Autocomplete for the catalog search box (GET /api/suggest).
|
||||
|
||||
Suggests brands and categories, not product names. Brand tables have no
|
||||
pg_trgm or tsvector index, so a product-name lookup would mean an unindexed
|
||||
ILIKE scan across every brand table on every keystroke. Brands and categories
|
||||
answer from memory instead:
|
||||
|
||||
* brands come from query_intent._brand_index(), which already merges the live
|
||||
brand list with the hand-tuned synonym table ("coke", "hul", "colgate") and
|
||||
caches for 5 minutes. It is invalidated on ingest, so a newly ingested brand
|
||||
becomes suggestible without any extra wiring here.
|
||||
* categories come from category_registry, whose keyword lists already carry
|
||||
common misspellings ("choclate", "biskut"), intersected with the categories
|
||||
actually present in the catalog.
|
||||
|
||||
Both of the user-facing examples resolve on the prefix tier:
|
||||
"Cavin" -> "Cavinkare", "Colgate" -> "Colgate-Palmolive".
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import difflib
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from app.infrastructure.settings import (
|
||||
SUGGEST_DEFAULT_LIMIT,
|
||||
SUGGEST_FUZZY_MIN_RATIO,
|
||||
SUGGEST_MIN_QUERY_LEN,
|
||||
)
|
||||
from app.services.category_registry import CATEGORY_REGISTRY
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Score floor per match quality. A shorter surface breaks ties, so "Cavinkare"
|
||||
# beats a longer brand that also starts with the same letters.
|
||||
_TIER_EXACT = 1.00
|
||||
_TIER_PREFIX = 0.90
|
||||
_TIER_WORD_PREFIX = 0.75
|
||||
_TIER_SUBSTRING = 0.55
|
||||
_FUZZY_WEIGHT = 0.40
|
||||
|
||||
# Fuzzy matching is a last resort: it only runs on longer queries, and only
|
||||
# when the precise tiers did not already fill the dropdown.
|
||||
_FUZZY_MIN_QUERY_LEN = 4
|
||||
|
||||
|
||||
@dataclass
|
||||
class Suggestion:
|
||||
type: str # "brand" | "category"
|
||||
value: str
|
||||
label: str
|
||||
sublabel: Optional[str] = None
|
||||
score: float = 0.0
|
||||
|
||||
|
||||
def _normalize(text: str) -> str:
|
||||
"""Fold case and treat -, &, . and _ as spaces so surfaces compare fairly."""
|
||||
lowered = (text or "").lower().strip()
|
||||
lowered = re.sub(r"[-&._]+", " ", lowered)
|
||||
return re.sub(r"\s+", " ", lowered).strip()
|
||||
|
||||
|
||||
def _score_surface(q: str, surface: str) -> Optional[float]:
|
||||
"""Score one candidate string against the query, or None for no match."""
|
||||
if not surface:
|
||||
return None
|
||||
if surface == q:
|
||||
return _TIER_EXACT
|
||||
if surface.startswith(q):
|
||||
return _TIER_PREFIX
|
||||
if any(word.startswith(q) for word in surface.split()):
|
||||
return _TIER_WORD_PREFIX
|
||||
if q in surface:
|
||||
return _TIER_SUBSTRING
|
||||
return None
|
||||
|
||||
|
||||
def _fuzzy_score(q: str, surface: str) -> Optional[float]:
|
||||
ratio = difflib.SequenceMatcher(None, q, surface).ratio()
|
||||
if ratio >= SUGGEST_FUZZY_MIN_RATIO:
|
||||
return _FUZZY_WEIGHT * ratio
|
||||
return None
|
||||
|
||||
|
||||
def _brand_surfaces(canonical: str, brand_map: Dict[str, str]) -> List[str]:
|
||||
"""Every string a user might type to mean this brand.
|
||||
|
||||
For "Colgate-Palmolive" that is the full name, the alias "colgate", and the
|
||||
individual words - so "palmoliv" finds it too.
|
||||
"""
|
||||
surfaces = {_normalize(canonical)}
|
||||
for alias, target in brand_map.items():
|
||||
if target == canonical:
|
||||
surfaces.add(_normalize(alias))
|
||||
for word in re.split(r"[-\s&]+", canonical):
|
||||
if len(word) >= 3:
|
||||
surfaces.add(_normalize(word))
|
||||
return [s for s in surfaces if s]
|
||||
|
||||
|
||||
def _brand_counts() -> Dict[str, int]:
|
||||
"""Product counts per brand, best-effort - the dropdown works without them."""
|
||||
try:
|
||||
from app.services.vector_store import get_brand_overview
|
||||
|
||||
return {row["display_name"]: row.get("product_count", 0) for row in get_brand_overview()}
|
||||
except Exception as e: # noqa: BLE001 - a DB blip must not break autocomplete
|
||||
logger.debug("Brand counts unavailable for suggest: %s", e)
|
||||
return {}
|
||||
|
||||
|
||||
def _suggest_brands(q: str, limit: int) -> List[Suggestion]:
|
||||
from app.services.query_intent import _brand_index
|
||||
|
||||
known, brand_map = _brand_index()
|
||||
counts = _brand_counts()
|
||||
|
||||
best: Dict[str, tuple] = {} # canonical -> (score, surface)
|
||||
for canonical in known:
|
||||
for surface in _brand_surfaces(canonical, brand_map):
|
||||
score = _score_surface(q, surface)
|
||||
if score is None:
|
||||
continue
|
||||
adjusted = score - 0.001 * len(surface)
|
||||
if canonical not in best or adjusted > best[canonical][0]:
|
||||
best[canonical] = (adjusted, surface)
|
||||
|
||||
# Fuzzy matching is only for rescuing a typo, so it runs only when nothing
|
||||
# matched precisely. Otherwise "colgat" would list Coca-Cola alongside the
|
||||
# Colgate-Palmolive the user obviously meant.
|
||||
if not best and len(q) >= _FUZZY_MIN_QUERY_LEN:
|
||||
for canonical in known:
|
||||
if canonical in best:
|
||||
continue
|
||||
for surface in _brand_surfaces(canonical, brand_map):
|
||||
score = _fuzzy_score(q, surface)
|
||||
if score is None:
|
||||
continue
|
||||
adjusted = score - 0.001 * len(surface)
|
||||
if canonical not in best or adjusted > best[canonical][0]:
|
||||
best[canonical] = (adjusted, surface)
|
||||
|
||||
suggestions = []
|
||||
for canonical, (score, _surface) in best.items():
|
||||
count = counts.get(canonical)
|
||||
suggestions.append(Suggestion(
|
||||
type="brand",
|
||||
value=canonical,
|
||||
label=canonical,
|
||||
sublabel=f"{count} products" if count else None,
|
||||
score=round(score, 4),
|
||||
))
|
||||
suggestions.sort(key=lambda s: (-s.score, s.label))
|
||||
return suggestions[:limit]
|
||||
|
||||
|
||||
def _live_categories() -> Optional[set]:
|
||||
"""Categories actually present in the catalog, or None if unknown."""
|
||||
try:
|
||||
from app.services.vector_store import list_all_categories
|
||||
|
||||
live = list_all_categories()
|
||||
return {_normalize(c) for c in live} if live else None
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("Live categories unavailable for suggest: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _suggest_categories(q: str, limit: int) -> List[Suggestion]:
|
||||
live = _live_categories()
|
||||
|
||||
best: Dict[str, float] = {}
|
||||
for entry in CATEGORY_REGISTRY:
|
||||
category = str(entry["category"])
|
||||
if live is not None and _normalize(category) not in live:
|
||||
# Don't offer a category nothing in the catalog is filed under.
|
||||
continue
|
||||
surfaces = [_normalize(category)] + [_normalize(k) for k in entry["keywords"]]
|
||||
for surface in surfaces:
|
||||
score = _score_surface(q, surface)
|
||||
if score is None:
|
||||
continue
|
||||
adjusted = score - 0.001 * len(surface)
|
||||
if category not in best or adjusted > best[category]:
|
||||
best[category] = adjusted
|
||||
|
||||
suggestions = [
|
||||
Suggestion(type="category", value=c, label=c, sublabel="Category", score=round(s, 4))
|
||||
for c, s in best.items()
|
||||
]
|
||||
suggestions.sort(key=lambda s: (-s.score, s.label))
|
||||
return suggestions[:limit]
|
||||
|
||||
|
||||
def suggest(query: str, limit: int = SUGGEST_DEFAULT_LIMIT) -> List[Suggestion]:
|
||||
"""Brand and category suggestions for a partial search query."""
|
||||
q = _normalize(query)
|
||||
if len(q) < SUGGEST_MIN_QUERY_LEN:
|
||||
return []
|
||||
|
||||
brands = _suggest_brands(q, limit)
|
||||
categories = _suggest_categories(q, limit)
|
||||
|
||||
# Rank by match quality, not by kind: a strong category match ("choc" ->
|
||||
# Chocolates) must not sit below a weak brand match. Brands only win ties,
|
||||
# since a bare word is more often a brand than a category.
|
||||
merged = brands + categories
|
||||
merged.sort(key=lambda s: (-s.score, s.type != "brand", s.label))
|
||||
return merged[:limit]
|
||||
@@ -465,6 +465,7 @@ def upsert_brand_products(brand: str, products: List[Dict[str, Any]], cleanup: b
|
||||
# one place that has to invalidate the derived views: the brand cards'
|
||||
# counts, and the brand list that chat/search intent parsing scopes on.
|
||||
invalidate_brand_overview_cache()
|
||||
invalidate_all_categories_cache()
|
||||
try:
|
||||
from app.services.query_intent import invalidate_brand_mention_cache
|
||||
invalidate_brand_mention_cache()
|
||||
@@ -1028,6 +1029,174 @@ def text_search(
|
||||
return results[:top_k]
|
||||
|
||||
|
||||
ALL_CATEGORIES_TTL_SECONDS = int(os.getenv("ALL_CATEGORIES_TTL_SECONDS", "300"))
|
||||
|
||||
_ALL_CATEGORIES_CACHE: Dict[str, Any] = {"at": 0.0, "data": None}
|
||||
|
||||
|
||||
def invalidate_all_categories_cache() -> None:
|
||||
"""Drop the cached cross-brand category list so the next read is fresh."""
|
||||
_ALL_CATEGORIES_CACHE.update(at=0.0, data=None)
|
||||
|
||||
|
||||
def list_all_categories(force_refresh: bool = False) -> List[str]:
|
||||
"""Every distinct category across all brand tables.
|
||||
|
||||
list_categories_for_brand() only answers for one brand, and the sidebar
|
||||
leaves its category list empty until a brand is selected, so the search
|
||||
suggester needs this cross-brand view. Cached because it runs per
|
||||
keystroke and touches every brand table.
|
||||
"""
|
||||
now = time.monotonic()
|
||||
if (
|
||||
not force_refresh
|
||||
and _ALL_CATEGORIES_CACHE["data"] is not None
|
||||
and now - _ALL_CATEGORIES_CACHE["at"] < ALL_CATEGORIES_TTL_SECONDS
|
||||
):
|
||||
return _ALL_CATEGORIES_CACHE["data"]
|
||||
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
return []
|
||||
|
||||
categories: set = set()
|
||||
try:
|
||||
with conn, conn.cursor() as cur:
|
||||
for suffix in _list_brand_table_suffixes(cur):
|
||||
table_name = f"brand_{suffix}"
|
||||
try:
|
||||
cur.execute(
|
||||
f"SELECT DISTINCT category FROM {table_name} "
|
||||
f"WHERE category IS NOT NULL AND category <> ''"
|
||||
)
|
||||
categories.update(r[0] for r in cur.fetchall() if r[0])
|
||||
except Exception as e: # noqa: BLE001 - a stale table must not break suggest
|
||||
logger.warning("list_all_categories skipped %s: %s", table_name, e)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("Failed to list categories: %s", e)
|
||||
return []
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
result = sorted(categories)
|
||||
_ALL_CATEGORIES_CACHE.update(at=now, data=result)
|
||||
return result
|
||||
|
||||
|
||||
def lexical_search(
|
||||
terms: List[str],
|
||||
brand: Optional[str] = None,
|
||||
limit: int = 200,
|
||||
category: Optional[str] = None,
|
||||
max_price: Optional[float] = None,
|
||||
query_embedding: Optional[List[float]] = None,
|
||||
exact_phrase: Optional[str] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Name-anchored search that returns a relevance tier per row.
|
||||
|
||||
Distinct from text_search() above, which stays exactly as it is because
|
||||
retrieve() depends on it as a zero-result fallback. Three differences that
|
||||
matter for the search box:
|
||||
|
||||
1. `terms` are ANDed, one clause each - "amul butter" must match rows
|
||||
containing both words, not either.
|
||||
2. `description` is NOT searched. Matching description is why a query for
|
||||
butter surfaces unrelated products that merely mention butter in their
|
||||
marketing copy.
|
||||
3. Each row carries `lex_tier` (0 = the name IS the query, 1 = the name
|
||||
starts with it, 2 = the name contains it, 3 = matched on terms only),
|
||||
which the caller uses to rank exact matches above semantic ones.
|
||||
|
||||
When `query_embedding` is given, a real cosine distance comes back in the
|
||||
same round trip, so a lexical hit reports a truthful similarity instead of
|
||||
a synthetic constant.
|
||||
"""
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
return []
|
||||
|
||||
terms = [t for t in (terms or []) if t]
|
||||
if not terms:
|
||||
conn.close()
|
||||
return []
|
||||
|
||||
phrase = (exact_phrase or " ".join(terms)).strip().lower()
|
||||
results: List[Dict[str, Any]] = []
|
||||
embedding_str = (
|
||||
"[" + ",".join(map(str, query_embedding)) + "]" if query_embedding else None
|
||||
)
|
||||
|
||||
try:
|
||||
with conn, conn.cursor() as cur:
|
||||
if brand:
|
||||
tables = [(brand, _table_name(brand))]
|
||||
else:
|
||||
tables = [(name, f"brand_{name}") for name in _list_brand_table_suffixes(cur)]
|
||||
|
||||
for brand_label, table_name in tables:
|
||||
if not _table_exists(cur, table_name):
|
||||
continue
|
||||
|
||||
params: List[Any] = []
|
||||
distance_select = ""
|
||||
if embedding_str is not None:
|
||||
distance_select = ", embedding <=> %s::vector AS distance"
|
||||
params.append(embedding_str)
|
||||
|
||||
# Tier the match by how closely the product NAME matches.
|
||||
tier_sql = (
|
||||
"CASE"
|
||||
" WHEN lower(coalesce(product_name, title, '')) = %s THEN 0"
|
||||
" WHEN lower(coalesce(product_name, title, '')) LIKE %s THEN 1"
|
||||
" WHEN coalesce(product_name, '') ILIKE %s"
|
||||
" OR coalesce(title, '') ILIKE %s THEN 2"
|
||||
" ELSE 3 END AS lex_tier"
|
||||
)
|
||||
params.extend([phrase, f"{phrase}%", f"%{phrase}%", f"%{phrase}%"])
|
||||
|
||||
where_clauses = []
|
||||
for term in terms:
|
||||
where_clauses.append(
|
||||
"(coalesce(product_name, '') ILIKE %s OR coalesce(title, '') ILIKE %s)"
|
||||
)
|
||||
params.extend([f"%{term}%", f"%{term}%"])
|
||||
|
||||
if category:
|
||||
where_clauses.append("category ILIKE %s")
|
||||
params.append(f"%{category}%")
|
||||
|
||||
sql = (
|
||||
f"SELECT *{distance_select}, {tier_sql} FROM {table_name} "
|
||||
f"WHERE {' AND '.join(where_clauses)} "
|
||||
f"ORDER BY lex_tier ASC, length(coalesce(product_name, title, '')) ASC, "
|
||||
f"updated_at DESC LIMIT %s"
|
||||
)
|
||||
params.append(limit)
|
||||
|
||||
try:
|
||||
cur.execute(sql, params)
|
||||
except Exception as e: # noqa: BLE001 - a stale table must not break search
|
||||
logger.warning("Lexical search failed for table %s: %s", table_name, e)
|
||||
continue
|
||||
|
||||
colnames = [desc[0] for desc in cur.description]
|
||||
for row in cur.fetchall():
|
||||
record = dict(zip(colnames, row))
|
||||
record["brand"] = record.get("brand") or brand_label
|
||||
|
||||
if max_price is not None:
|
||||
min_p, _max_p = parse_price_range(record.get("price_range"))
|
||||
if min_p is not None and min_p > max_price:
|
||||
continue
|
||||
|
||||
results.append(record)
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
results.sort(key=lambda r: (r.get("lex_tier", 9), r.get("distance", 9.0)))
|
||||
return results[:limit]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -87,5 +87,23 @@ def test_openapi_schema_lists_all_routers(client) -> None:
|
||||
resp = client.get("/openapi.json")
|
||||
assert resp.status_code == 200
|
||||
paths = resp.json()["paths"]
|
||||
for expected in ("/api/health", "/api/brands", "/api/search", "/api/chat", "/api/catalog/generate"):
|
||||
for expected in ("/api/health", "/api/brands", "/api/search", "/api/suggest",
|
||||
"/api/chat", "/api/catalog/generate"):
|
||||
assert expected in paths, f"missing route: {expected}"
|
||||
|
||||
|
||||
def test_suggest_works_without_a_database(client) -> None:
|
||||
"""Autocomplete must degrade, not 500, when the catalog is unreachable.
|
||||
|
||||
This suite runs with no reachable database, so this pins that the search
|
||||
box keeps suggesting from the static brand table instead of erroring - a
|
||||
failing suggest call must never block someone from typing and searching.
|
||||
"""
|
||||
resp = client.get("/api/suggest", params={"q": "cavin"})
|
||||
assert resp.status_code == 200
|
||||
labels = [s["label"] for s in resp.json()["suggestions"]]
|
||||
assert "Cavinkare" in labels
|
||||
|
||||
|
||||
def test_suggest_requires_a_query(client) -> None:
|
||||
assert client.get("/api/suggest").status_code == 422
|
||||
|
||||
91
tests/test_catalog_search_rank.py
Normal file
91
tests/test_catalog_search_rank.py
Normal file
@@ -0,0 +1,91 @@
|
||||
"""Hybrid merge/dedup/order for the catalog search box.
|
||||
|
||||
Pure-Python over hand-built rows, so no database is involved.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from app.services.catalog_search import _hybrid_rank, _row_key
|
||||
|
||||
|
||||
def _row(brand, image_id, name, **extra):
|
||||
row = {"brand": brand, "image_id": image_id, "product_name": name}
|
||||
row.update(extra)
|
||||
return row
|
||||
|
||||
|
||||
def test_name_matches_outrank_semantic_only_matches():
|
||||
semantic = [
|
||||
_row("Amul", "s1", "Amul Cheese Cubes", distance=0.10),
|
||||
_row("Amul", "s2", "Amul Fresh Cream", distance=0.12),
|
||||
]
|
||||
lexical = [
|
||||
_row("Amul", "l1", "Amul Butter 100 g", distance=0.40, lex_tier=1),
|
||||
_row("Amul", "l2", "Amul Butter 500 g", distance=0.45, lex_tier=1),
|
||||
]
|
||||
|
||||
ranked = _hybrid_rank(semantic, lexical, limit=10)
|
||||
names = [r["product_name"] for r in ranked]
|
||||
|
||||
# Both butter variants come first even though their vector distance is worse.
|
||||
assert names[:2] == ["Amul Butter 100 g", "Amul Butter 500 g"]
|
||||
assert all(r.get("lex_tier", 9) == 1 for r in ranked[:2])
|
||||
assert all(r.get("lex_tier", 9) == 9 for r in ranked[2:])
|
||||
|
||||
|
||||
def test_lower_lex_tier_wins_within_the_lexical_block():
|
||||
lexical = [
|
||||
_row("Amul", "l2", "Amul Butter Salted", distance=0.20, lex_tier=2),
|
||||
_row("Amul", "l0", "Amul Butter", distance=0.90, lex_tier=0),
|
||||
_row("Amul", "l1", "Amul Butter 500 g", distance=0.50, lex_tier=1),
|
||||
]
|
||||
ranked = _hybrid_rank([], lexical, limit=10)
|
||||
assert [r["lex_tier"] for r in ranked] == [0, 1, 2]
|
||||
|
||||
|
||||
def test_distance_breaks_ties_inside_a_tier():
|
||||
lexical = [
|
||||
_row("Amul", "b", "Amul Butter 500 g", distance=0.50, lex_tier=1),
|
||||
_row("Amul", "a", "Amul Butter 100 g", distance=0.20, lex_tier=1),
|
||||
]
|
||||
ranked = _hybrid_rank([], lexical, limit=10)
|
||||
assert [r["image_id"] for r in ranked] == ["a", "b"]
|
||||
|
||||
|
||||
def test_dedup_keeps_the_real_semantic_distance_and_takes_the_tier():
|
||||
semantic = [_row("Amul", "same", "Amul Butter 100 g", distance=0.11)]
|
||||
lexical = [_row("Amul", "same", "Amul Butter 100 g", distance=None, lex_tier=1)]
|
||||
|
||||
ranked = _hybrid_rank(semantic, lexical, limit=10)
|
||||
|
||||
assert len(ranked) == 1, "the same product must not appear twice"
|
||||
assert ranked[0]["distance"] == 0.11, "the real vector distance must survive"
|
||||
assert ranked[0]["lex_tier"] == 1, "the lexical tier must be applied"
|
||||
|
||||
|
||||
def test_same_image_id_in_different_brands_is_not_deduped():
|
||||
# image_id is UNIQUE per brand table, not globally.
|
||||
semantic = [
|
||||
_row("Amul", "dup", "Amul Butter", distance=0.10),
|
||||
_row("Nestle", "dup", "Nestle Butter", distance=0.20),
|
||||
]
|
||||
ranked = _hybrid_rank(semantic, [], limit=10)
|
||||
assert len(ranked) == 2
|
||||
assert _row_key(semantic[0]) != _row_key(semantic[1])
|
||||
|
||||
|
||||
def test_lexical_only_rows_get_a_synthetic_distance_matching_their_tier():
|
||||
lexical = [
|
||||
_row("Amul", "t0", "Amul Butter", distance=None, lex_tier=0),
|
||||
_row("Amul", "t3", "Butter Amul Pack", distance=None, lex_tier=3),
|
||||
]
|
||||
ranked = _hybrid_rank([], lexical, limit=10)
|
||||
by_id = {r["image_id"]: r for r in ranked}
|
||||
|
||||
# similarity == 1 - distance/2, so an exact name match must read as 1.0.
|
||||
assert by_id["t0"]["distance"] == 0.0
|
||||
assert by_id["t3"]["distance"] > by_id["t0"]["distance"]
|
||||
|
||||
|
||||
def test_limit_is_respected():
|
||||
semantic = [_row("Amul", f"s{i}", f"Product {i}", distance=i / 100) for i in range(50)]
|
||||
assert len(_hybrid_rank(semantic, [], limit=10)) == 10
|
||||
80
tests/test_search_intent.py
Normal file
80
tests/test_search_intent.py
Normal file
@@ -0,0 +1,80 @@
|
||||
"""Query-shape classification behind the catalog search box.
|
||||
|
||||
The case that matters most is "Amul Butter". brand_registry.BRAND_ALIASES
|
||||
contains "amul butter" -> "amul", so any classifier built on
|
||||
resolve_parent_brand() would call it brand-only and list Amul's entire
|
||||
catalog - the opposite of what the user asked for. These tests pin the
|
||||
distinction.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services import query_intent
|
||||
from app.services.query_intent import classify_query
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _static_brand_index(monkeypatch):
|
||||
"""Pin the brand index so the tests don't depend on a live database."""
|
||||
known = list(query_intent.KNOWN_BRANDS)
|
||||
mapping = dict(query_intent.BRAND_SEARCH_MAP)
|
||||
monkeypatch.setattr(query_intent, "_brand_index", lambda: (known, mapping))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("query,brand", [
|
||||
("Amul", "Amul"),
|
||||
("amul", "Amul"),
|
||||
("AMUL", "Amul"),
|
||||
("Cavinkare", "Cavinkare"),
|
||||
("Milky Mist", "Milky Mist"),
|
||||
("Lion Dates", "Lion Dates"),
|
||||
("Colgate", "Colgate-Palmolive"),
|
||||
("colgate-palmolive", "Colgate-Palmolive"),
|
||||
("coke", "Coca-Cola"),
|
||||
("coca cola", "Coca-Cola"),
|
||||
("hul", "Hindustan Unilever"),
|
||||
("pepsi", "Pepsico"),
|
||||
("show me amul products", "Amul"),
|
||||
("all products from Cadbury", "Cadbury"),
|
||||
("Amul?", "Amul"),
|
||||
])
|
||||
def test_brand_only_queries(query, brand):
|
||||
shape = classify_query(query)
|
||||
assert shape.kind == "brand_only", f"{query!r} -> {shape}"
|
||||
assert shape.brand == brand
|
||||
|
||||
|
||||
@pytest.mark.parametrize("query,brand,residual_contains", [
|
||||
("Amul Butter", "Amul", "butter"),
|
||||
("amul butter", "Amul", "butter"),
|
||||
("Amul Cheese Slices", "Amul", "cheese"),
|
||||
("Cadbury Dairy Milk", "Cadbury", "dairy milk"),
|
||||
("Colgate toothpaste", "Colgate-Palmolive", "toothpaste"),
|
||||
])
|
||||
def test_brand_plus_queries_keep_their_product_terms(query, brand, residual_contains):
|
||||
shape = classify_query(query)
|
||||
assert shape.kind == "brand_plus", f"{query!r} -> {shape}"
|
||||
assert shape.brand == brand
|
||||
assert residual_contains in shape.residual
|
||||
|
||||
|
||||
@pytest.mark.parametrize("query", [
|
||||
"low sugar biscuit",
|
||||
"something nobody sells",
|
||||
"",
|
||||
" ",
|
||||
])
|
||||
def test_generic_queries(query):
|
||||
assert classify_query(query).kind == "generic"
|
||||
|
||||
|
||||
def test_explicit_brand_with_empty_query_is_brand_only():
|
||||
shape = classify_query("", explicit_brand="Amul")
|
||||
assert shape.kind == "brand_only"
|
||||
assert shape.brand == "Amul"
|
||||
|
||||
|
||||
def test_bare_category_word_is_not_treated_as_a_brand():
|
||||
# Guards a future brand literally named after a product word.
|
||||
assert classify_query("butter").kind != "brand_only"
|
||||
110
tests/test_suggest_service.py
Normal file
110
tests/test_suggest_service.py
Normal file
@@ -0,0 +1,110 @@
|
||||
"""Autocomplete ranking for the catalog search box.
|
||||
|
||||
Runs without a database: the brand index is monkeypatched and the live
|
||||
category/count lookups degrade to empty, which is also how the endpoint
|
||||
behaves in tests/test_api.py where no database is reachable.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services import query_intent, suggest_service
|
||||
from app.services.suggest_service import suggest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _static_brand_index(monkeypatch):
|
||||
known = list(query_intent.KNOWN_BRANDS)
|
||||
mapping = dict(query_intent.BRAND_SEARCH_MAP)
|
||||
monkeypatch.setattr(query_intent, "_brand_index", lambda: (known, mapping))
|
||||
# No database in unit tests: counts and live categories are unavailable.
|
||||
monkeypatch.setattr(suggest_service, "_brand_counts", lambda: {})
|
||||
monkeypatch.setattr(suggest_service, "_live_categories", lambda: None)
|
||||
|
||||
|
||||
def _labels(query, **kwargs):
|
||||
return [s.label for s in suggest(query, **kwargs)]
|
||||
|
||||
|
||||
def test_partial_brand_suggests_the_full_name():
|
||||
# The two examples from the original request.
|
||||
assert "Cavinkare" in _labels("Cavin")
|
||||
assert "Colgate-Palmolive" in _labels("Colgate")
|
||||
|
||||
|
||||
def test_partial_match_is_case_insensitive():
|
||||
assert "Cavinkare" in _labels("cAvIn")
|
||||
|
||||
|
||||
def test_alias_resolves_to_the_canonical_brand():
|
||||
assert "Coca-Cola" in _labels("coke")
|
||||
assert "Hindustan Unilever" in _labels("hul")
|
||||
assert "Pepsico" in _labels("pepsi")
|
||||
|
||||
|
||||
def test_word_prefix_matches_a_later_word_in_the_name():
|
||||
assert "Colgate-Palmolive" in _labels("palmoliv")
|
||||
|
||||
|
||||
def test_exact_match_outranks_a_mere_prefix():
|
||||
results = suggest("amul")
|
||||
assert results[0].label == "Amul"
|
||||
|
||||
|
||||
def test_typo_still_finds_the_brand():
|
||||
assert "Colgate-Palmolive" in _labels("colgat")
|
||||
|
||||
|
||||
def test_unknown_text_returns_nothing():
|
||||
assert suggest("zzzzqqq") == []
|
||||
|
||||
|
||||
def test_query_below_minimum_length_is_ignored():
|
||||
assert suggest("c") == []
|
||||
|
||||
|
||||
def test_limit_is_respected():
|
||||
assert len(suggest("a", limit=3)) <= 3
|
||||
assert len(suggest("ca", limit=2)) <= 2
|
||||
|
||||
|
||||
def test_categories_are_suggested_by_keyword():
|
||||
labels = _labels("choc")
|
||||
assert "Chocolates" in labels
|
||||
|
||||
|
||||
def test_category_keyword_misspelling_still_matches():
|
||||
# category_registry carries deliberate misspellings.
|
||||
assert "Biscuits & Cookies" in _labels("biskut")
|
||||
|
||||
|
||||
def test_a_brand_wins_a_tie_against_a_category():
|
||||
# Ranking is by match quality; kind only breaks ties, because a bare word
|
||||
# is more often reaching for a brand.
|
||||
results = suggest("ca")
|
||||
for a, b in zip(results, results[1:]):
|
||||
if a.score == b.score and a.type != b.type:
|
||||
assert a.type == "brand"
|
||||
|
||||
|
||||
def test_a_strong_category_match_outranks_a_weak_brand_match():
|
||||
# "choc" must lead with Chocolates, not with a fuzzy brand near-miss.
|
||||
results = suggest("choc")
|
||||
assert results, "expected at least one suggestion"
|
||||
assert results[0].label == "Chocolates"
|
||||
|
||||
|
||||
def test_fuzzy_matches_do_not_pollute_a_precise_match():
|
||||
# "colgat" is a prefix of Colgate-Palmolive, so no fuzzy noise should appear.
|
||||
labels = _labels("colgat")
|
||||
assert labels[0] == "Colgate-Palmolive"
|
||||
assert "Coca-Cola" not in labels
|
||||
|
||||
|
||||
def test_suggestions_are_sorted_by_descending_score_within_each_type():
|
||||
# Brands are grouped ahead of categories on purpose, so the score ordering
|
||||
# holds within each group rather than across the whole list.
|
||||
results = suggest("co")
|
||||
for kind in ("brand", "category"):
|
||||
scores = [s.score for s in results if s.type == kind]
|
||||
assert scores == sorted(scores, reverse=True), kind
|
||||
Reference in New Issue
Block a user