From fbb1356e475971f94d6db484f5bb4fab37885a63 Mon Sep 17 00:00:00 2001 From: sriram Date: Thu, 20 Aug 2026 13:03:16 +0530 Subject: [PATCH] updates on catalog search and suggestions in backend --- app/api/routers/search.py | 37 +++-- app/api/routers/suggest.py | 35 +++++ app/api/schemas.py | 26 ++++ app/infrastructure/settings.py | 21 +++ app/main.py | 3 +- app/services/catalog_search.py | 240 ++++++++++++++++++++++++++++++ app/services/query_intent.py | 132 ++++++++++++++++ app/services/rag_service.py | 13 ++ app/services/suggest_service.py | 211 ++++++++++++++++++++++++++ app/services/vector_store.py | 169 +++++++++++++++++++++ tests/test_api.py | 20 ++- tests/test_catalog_search_rank.py | 91 +++++++++++ tests/test_search_intent.py | 80 ++++++++++ tests/test_suggest_service.py | 110 ++++++++++++++ 14 files changed, 1177 insertions(+), 11 deletions(-) create mode 100644 app/api/routers/suggest.py create mode 100644 app/services/catalog_search.py create mode 100644 app/services/suggest_service.py create mode 100644 tests/test_catalog_search_rank.py create mode 100644 tests/test_search_intent.py create mode 100644 tests/test_suggest_service.py diff --git a/app/api/routers/search.py b/app/api/routers/search.py index f772e88..714b835 100644 --- a/app/api/routers/search.py +++ b/app/api/routers/search.py @@ -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, ) diff --git a/app/api/routers/suggest.py b/app/api/routers/suggest.py new file mode 100644 index 0000000..4d31b23 --- /dev/null +++ b/app/api/routers/suggest.py @@ -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 + ], + ) diff --git a/app/api/schemas.py b/app/api/schemas.py index 1395229..ee510ce 100644 --- a/app/api/schemas.py +++ b/app/api/schemas.py @@ -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] # --------------------------------------------------------------------------- diff --git a/app/infrastructure/settings.py b/app/infrastructure/settings.py index 7c6f0a7..1e07cd6 100644 --- a/app/infrastructure/settings.py +++ b/app/infrastructure/settings.py @@ -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")) diff --git a/app/main.py b/app/main.py index 43b5822..c5c5e5c 100644 --- a/app/main.py +++ b/app/main.py @@ -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") diff --git a/app/services/catalog_search.py b/app/services/catalog_search.py new file mode 100644 index 0000000..40ec19b --- /dev/null +++ b/app/services/catalog_search.py @@ -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) diff --git a/app/services/query_intent.py b/app/services/query_intent.py index c04733e..c41337e 100644 --- a/app/services/query_intent.py +++ b/app/services/query_intent.py @@ -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 diff --git a/app/services/rag_service.py b/app/services/rag_service.py index c34a652..4ca64e9 100644 --- a/app/services/rag_service.py +++ b/app/services/rag_service.py @@ -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 diff --git a/app/services/suggest_service.py b/app/services/suggest_service.py new file mode 100644 index 0000000..29b8494 --- /dev/null +++ b/app/services/suggest_service.py @@ -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] diff --git a/app/services/vector_store.py b/app/services/vector_store.py index 952ea64..37098e3 100644 --- a/app/services/vector_store.py +++ b/app/services/vector_store.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_api.py b/tests/test_api.py index f9649aa..93f5bd2 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -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 diff --git a/tests/test_catalog_search_rank.py b/tests/test_catalog_search_rank.py new file mode 100644 index 0000000..61c3da7 --- /dev/null +++ b/tests/test_catalog_search_rank.py @@ -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 diff --git a/tests/test_search_intent.py b/tests/test_search_intent.py new file mode 100644 index 0000000..4fc4e21 --- /dev/null +++ b/tests/test_search_intent.py @@ -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" diff --git a/tests/test_suggest_service.py b/tests/test_suggest_service.py new file mode 100644 index 0000000..2156a3a --- /dev/null +++ b/tests/test_suggest_service.py @@ -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