updates on the backend
This commit is contained in:
93
scripts/backfill_barcodes.py
Normal file
93
scripts/backfill_barcodes.py
Normal file
@@ -0,0 +1,93 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Backfill PostgreSQL brand tables with `barcode` and `barcode_type` from seed catalogs.
|
||||
This ensures existing products with barcode details in seed files or DB rows are populated and served to the API and UI.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.vector_store import _connect, list_available_brands, _sanitize_name, resolve_parent_brand, ensure_brand_schema
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def backfill_barcodes() -> None:
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
logger.error("Could not connect to PostgreSQL database.")
|
||||
return
|
||||
|
||||
brands = list_available_brands()
|
||||
logger.info("Ensuring schema for %d brand(s)...", len(brands))
|
||||
for brand in brands:
|
||||
ensure_brand_schema(brand)
|
||||
|
||||
seed_files = sorted(SEED_DIR.glob("*.json")) if SEED_DIR.exists() else []
|
||||
logger.info("Found %d seed file(s) in %s", len(seed_files), SEED_DIR)
|
||||
|
||||
updated_count = 0
|
||||
with conn.cursor() as cur:
|
||||
for seed_file in seed_files:
|
||||
try:
|
||||
data = json.loads(seed_file.read_text(encoding="utf-8-sig"))
|
||||
except Exception as e:
|
||||
logger.warning("Could not read seed file %s: %s", seed_file.name, e)
|
||||
continue
|
||||
|
||||
products = data.get("products", [])
|
||||
brand = data.get("brand") or (products[0].get("brand_name") if products else None)
|
||||
if not brand or not products:
|
||||
continue
|
||||
|
||||
table_name = f"brand_{_sanitize_name(resolve_parent_brand(brand))}"
|
||||
|
||||
for p in products:
|
||||
image_id = p.get("image_id") or p.get("sku") or ""
|
||||
product_name = p.get("product_name") or p.get("title") or ""
|
||||
barcode = str(p.get("barcode") or p.get("Barcode") or "").strip() or None
|
||||
barcode_type = str(p.get("barcode_type") or p.get("Barcode_Type") or "").strip() or None
|
||||
|
||||
if not barcode and not barcode_type:
|
||||
continue
|
||||
|
||||
if image_id:
|
||||
cur.execute(
|
||||
f"""
|
||||
UPDATE {table_name}
|
||||
SET barcode = COALESCE(%s, barcode),
|
||||
barcode_type = COALESCE(%s, barcode_type)
|
||||
WHERE image_id = %s
|
||||
""",
|
||||
(barcode, barcode_type, image_id)
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
updated_count += cur.rowcount
|
||||
elif product_name:
|
||||
cur.execute(
|
||||
f"""
|
||||
UPDATE {table_name}
|
||||
SET barcode = COALESCE(%s, barcode),
|
||||
barcode_type = COALESCE(%s, barcode_type)
|
||||
WHERE product_name = %s
|
||||
""",
|
||||
(barcode, barcode_type, product_name)
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
updated_count += cur.rowcount
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
logger.info("🎉 Barcode backfill complete! Updated %d product row(s).", updated_count)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
backfill_barcodes()
|
||||
111
scripts/backfill_hsn_prices.py
Normal file
111
scripts/backfill_hsn_prices.py
Normal file
@@ -0,0 +1,111 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Backfill PostgreSQL brand tables with `hsn_code`, `final_selling_price`, and `selling_price` from seed catalogs.
|
||||
This ensures existing products with these details in seed files or DB rows are populated and served to the API and UI.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.vector_store import _connect, list_available_brands, _sanitize_name, resolve_parent_brand, ensure_brand_schema
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def backfill_hsn_prices() -> None:
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
logger.error("Could not connect to PostgreSQL database.")
|
||||
return
|
||||
|
||||
brands = list_available_brands()
|
||||
logger.info("Ensuring schema for %d brand(s)...", len(brands))
|
||||
for brand in brands:
|
||||
ensure_brand_schema(brand)
|
||||
|
||||
seed_files = sorted(SEED_DIR.glob("*.json")) if SEED_DIR.exists() else []
|
||||
logger.info("Found %d seed file(s) in %s", len(seed_files), SEED_DIR)
|
||||
|
||||
updated_count = 0
|
||||
with conn.cursor() as cur:
|
||||
for seed_file in seed_files:
|
||||
try:
|
||||
data = json.loads(seed_file.read_text(encoding="utf-8-sig"))
|
||||
except Exception as e:
|
||||
logger.warning("Could not read seed file %s: %s", seed_file.name, e)
|
||||
continue
|
||||
|
||||
products = data.get("products", [])
|
||||
brand = data.get("brand") or (products[0].get("brand_name") if products else None)
|
||||
if not brand or not products:
|
||||
continue
|
||||
|
||||
table_name = f"brand_{_sanitize_name(resolve_parent_brand(brand))}"
|
||||
|
||||
for p in products:
|
||||
image_id = p.get("image_id") or p.get("sku") or ""
|
||||
product_name = p.get("product_name") or p.get("title") or ""
|
||||
hsn = str(p.get("hsn_code") or p.get("HSN_Code") or p.get("hsn") or "").strip() or None
|
||||
|
||||
raw_fsp = p.get("final_selling_price") if "final_selling_price" in p else p.get("Final_Selling_Price")
|
||||
if raw_fsp is None:
|
||||
raw_fsp = p.get("final_price")
|
||||
try:
|
||||
fsp = float(raw_fsp) if raw_fsp is not None and str(raw_fsp).strip() != "" else None
|
||||
except (ValueError, TypeError):
|
||||
fsp = None
|
||||
|
||||
raw_sp = p.get("selling_price") if "selling_price" in p else p.get("Selling_Price")
|
||||
try:
|
||||
sp = float(raw_sp) if raw_sp is not None and str(raw_sp).strip() != "" else None
|
||||
except (ValueError, TypeError):
|
||||
sp = None
|
||||
|
||||
if fsp is None and sp is not None:
|
||||
fsp = sp
|
||||
|
||||
if not hsn and fsp is None and sp is None:
|
||||
continue
|
||||
|
||||
if image_id:
|
||||
cur.execute(
|
||||
f"""
|
||||
UPDATE {table_name}
|
||||
SET hsn_code = COALESCE(%s, hsn_code),
|
||||
final_selling_price = COALESCE(%s, final_selling_price),
|
||||
selling_price = COALESCE(%s, selling_price)
|
||||
WHERE image_id = %s
|
||||
""",
|
||||
(hsn, fsp, sp, image_id)
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
updated_count += cur.rowcount
|
||||
elif product_name:
|
||||
cur.execute(
|
||||
f"""
|
||||
UPDATE {table_name}
|
||||
SET hsn_code = COALESCE(%s, hsn_code),
|
||||
final_selling_price = COALESCE(%s, final_selling_price),
|
||||
selling_price = COALESCE(%s, selling_price)
|
||||
WHERE product_name = %s
|
||||
""",
|
||||
(hsn, fsp, sp, product_name)
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
updated_count += cur.rowcount
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
logger.info("🎉 HSN and Price backfill complete! Updated %d product row(s).", updated_count)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
backfill_hsn_prices()
|
||||
80
scripts/backfill_s3_urls.py
Normal file
80
scripts/backfill_s3_urls.py
Normal file
@@ -0,0 +1,80 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Backfill PostgreSQL product rows that have `image_id` set but empty/null `image_url` or `image_urls`.
|
||||
This updates products in-place in pgvector so that API requests and RAG queries serve image URLs directly from DB
|
||||
without making S3 list_objects_v2 network calls.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.vector_store import _connect, list_available_brands, _sanitize_name, resolve_parent_brand, ensure_brand_schema
|
||||
from app.services.s3_service import s3_service
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def backfill_brand_images() -> None:
|
||||
if not s3_service.enabled:
|
||||
logger.warning("S3 service is not enabled. Skipping backfill.")
|
||||
return
|
||||
|
||||
brands = list_available_brands()
|
||||
logger.info("Found %d brands in PostgreSQL database", len(brands))
|
||||
|
||||
for brand in brands:
|
||||
ensure_brand_schema(brand)
|
||||
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
logger.error("Could not connect to PostgreSQL database.")
|
||||
return
|
||||
|
||||
brands = list_available_brands()
|
||||
logger.info("Found %d brands in PostgreSQL database", len(brands))
|
||||
|
||||
updated_total = 0
|
||||
with conn.cursor() as cur:
|
||||
for brand in brands:
|
||||
table_name = f"brand_{_sanitize_name(resolve_parent_brand(brand))}"
|
||||
try:
|
||||
cur.execute(f"SELECT id, image_id, image_url, image_urls FROM {table_name}")
|
||||
rows = cur.fetchall()
|
||||
except Exception as e:
|
||||
logger.warning("Could not read table %s: %s", table_name, e)
|
||||
conn.rollback()
|
||||
continue
|
||||
|
||||
for row in rows:
|
||||
row_id, image_id, image_url, image_urls = row
|
||||
if not image_id:
|
||||
continue
|
||||
|
||||
db_has_images = bool(image_url or (image_urls and any(image_urls)))
|
||||
if not db_has_images:
|
||||
# Fetch from S3 (or construct primary S3 URL)
|
||||
s3_urls = s3_service.get_product_image_urls(brand, image_id)
|
||||
primary_url = s3_urls[0] if s3_urls else s3_service.get_product_image_url(brand, image_id)
|
||||
if not s3_urls and primary_url:
|
||||
s3_urls = [primary_url]
|
||||
|
||||
if primary_url or s3_urls:
|
||||
cur.execute(
|
||||
f"UPDATE {table_name} SET image_url = %s, image_urls = %s WHERE id = %s",
|
||||
(primary_url, s3_urls, row_id)
|
||||
)
|
||||
updated_total += 1
|
||||
logger.info("Updated product ID %s in %s (image_id=%s) with %d image(s)", row_id, table_name, image_id, len(s3_urls))
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
logger.info("🎉 Backfill complete! Updated %d product(s) in database.", updated_total)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
backfill_brand_images()
|
||||
8
scripts/check_stores.py
Normal file
8
scripts/check_stores.py
Normal file
@@ -0,0 +1,8 @@
|
||||
import sys
|
||||
sys.path.insert(0, ".")
|
||||
from app.services.store_db import list_stores
|
||||
|
||||
stores = list_stores()
|
||||
print("Stores in DB:")
|
||||
for s in stores:
|
||||
print(f" - {s['store_id']}: {s['store_name']} ({s.get('city')}) [{s.get('tier')}]")
|
||||
59
scripts/enrich_nutrition.py
Normal file
59
scripts/enrich_nutrition.py
Normal file
@@ -0,0 +1,59 @@
|
||||
"""
|
||||
CLI entry point for Feature 1/4/15's retrieval pipeline - fetches
|
||||
verified nutrition data (Open Food Facts) for every product in the
|
||||
catalog, computes transparent scores/insights, and persists them.
|
||||
|
||||
Usage:
|
||||
python scripts/enrich_nutrition.py
|
||||
python scripts/enrich_nutrition.py --max-products 200 --no-narrative
|
||||
python scripts/enrich_nutrition.py --force # re-fetch even already-verified products
|
||||
|
||||
Equivalent to POST /api/admin/nutrition-intelligence/enrich.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.nutrition_db import ensure_nutrition_schema # noqa: E402
|
||||
from app.services.nutrition_enrichment_service import enrich_all_products # noqa: E402
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _progress(done: int, total: int) -> None:
|
||||
if total and (done % 25 == 0 or done == total):
|
||||
logger.info(f"Enrichment progress: {done}/{total}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Enrich the catalog with verified nutrition data")
|
||||
parser.add_argument("--force", action="store_true", help="Re-fetch even products already marked 'verified'")
|
||||
parser.add_argument("--no-narrative", action="store_true", help="Skip the LLM narrative step (facts/scores/tags only)")
|
||||
parser.add_argument("--max-products", type=int, default=None, help="Cap the number of products processed (for a quick test run)")
|
||||
args = parser.parse_args()
|
||||
|
||||
ensure_nutrition_schema()
|
||||
result = enrich_all_products(
|
||||
skip_if_verified=not args.force,
|
||||
generate_narrative=not args.no_narrative,
|
||||
progress_cb=_progress,
|
||||
max_products=args.max_products,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Done in {result.duration_seconds}s - "
|
||||
f"{result.verified} verified, {result.partial} partial, {result.unavailable} unavailable "
|
||||
f"of {result.total_products} total products"
|
||||
)
|
||||
if result.errors:
|
||||
logger.warning(f"{len(result.errors)} errors (showing up to 10): {result.errors[:10]}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
144
scripts/export_seed_data.py
Normal file
144
scripts/export_seed_data.py
Normal file
@@ -0,0 +1,144 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Export data FROM pgvector brand tables BACK INTO seed_catalog JSON files.
|
||||
|
||||
This is the REVERSE of seed_sample_data.py. Use it when the database has
|
||||
been populated with better/real data (via direct ingestion or manual fixes)
|
||||
and you want to snapshot that state back into the seed files so that future
|
||||
runs of `seed_sample_data.py` don't regress to stale AI-generated content.
|
||||
|
||||
Usage:
|
||||
python scripts/export_seed_data.py
|
||||
python scripts/export_seed_data.py --only-having-products
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Any
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.vector_store import _connect, _sanitize_name
|
||||
from app.infrastructure.settings import DB_HOST, DB_PORT, DB_NAME
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def export_all(only_having_products: bool = False) -> None:
|
||||
conn = _connect()
|
||||
if not conn:
|
||||
logger.error("Cannot connect to database. Is pgvector running?")
|
||||
sys.exit(1)
|
||||
|
||||
try:
|
||||
with conn, conn.cursor() as cur:
|
||||
cur.execute("""
|
||||
SELECT table_name FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_name LIKE 'brand_%'
|
||||
ORDER BY table_name
|
||||
""")
|
||||
tables = cur.fetchall()
|
||||
|
||||
if not tables:
|
||||
logger.warning("No brand_* tables found in the database.")
|
||||
conn.close()
|
||||
return
|
||||
|
||||
exported_count = 0
|
||||
for (table_name,) in tables:
|
||||
suffix = table_name[len("brand_"):]
|
||||
display_brand = suffix.replace("_", " ").title()
|
||||
|
||||
# Build a friendly filename (strip problematic chars)
|
||||
safe_name = suffix.replace(" ", "_").replace("-", "_").replace("&", "_")
|
||||
safe_name = safe_name.strip("_")
|
||||
filepath = SEED_DIR / f"brand_catalog_{safe_name}.json"
|
||||
|
||||
products = _read_products(conn, table_name)
|
||||
|
||||
if only_having_products and not products:
|
||||
logger.info("Skipping %s (0 products)", table_name)
|
||||
continue
|
||||
|
||||
catalog = {
|
||||
"brand": display_brand,
|
||||
"search_query": None,
|
||||
"generation_timestamp": f"Exported from {DB_HOST}:{DB_PORT}/{DB_NAME} on {datetime.now().isoformat()}",
|
||||
"total_products": len(products),
|
||||
"total_images": sum(len(p.get("image_urls") or []) for p in products),
|
||||
"engine_info": {
|
||||
"gemini_enabled": False,
|
||||
"architecture": "Exported from pgvector database",
|
||||
},
|
||||
"products": products,
|
||||
}
|
||||
|
||||
SEED_DIR.mkdir(parents=True, exist_ok=True)
|
||||
filepath.write_text(
|
||||
json.dumps(catalog, indent=2, ensure_ascii=False, default=str),
|
||||
encoding="utf-8",
|
||||
)
|
||||
logger.info("Exported %d products -> %s", len(products), filepath.name)
|
||||
exported_count += 1
|
||||
|
||||
print(f"\nDone. Exported {exported_count} brand table(s) to {SEED_DIR}")
|
||||
print("Run `python scripts/seed_sample_data.py` to re-import them if needed.")
|
||||
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _read_products(conn, table_name: str) -> List[Dict[str, Any]]:
|
||||
"""Read all rows from a brand table and format them as seed-catalog products."""
|
||||
with conn, conn.cursor() as cur:
|
||||
try:
|
||||
cur.execute(f"SELECT * FROM {table_name} ORDER BY updated_at DESC")
|
||||
colnames = [desc[0] for desc in cur.description]
|
||||
rows = cur.fetchall()
|
||||
except Exception as e:
|
||||
logger.warning("Failed to read %s: %s", table_name, e)
|
||||
return []
|
||||
|
||||
products = []
|
||||
for row in rows:
|
||||
record = dict(zip(colnames, row))
|
||||
product = {
|
||||
"title": record.get("title") or record.get("product_name") or "Unknown",
|
||||
"category": record.get("category") or "Uncategorized",
|
||||
"description": record.get("description") or "",
|
||||
"image_id": record.get("image_id") or "",
|
||||
"price_range": record.get("price_range") or "",
|
||||
"size_variants": list(record.get("size_variants") or []),
|
||||
"providers": list(record.get("providers") or []),
|
||||
"highlights": list(record.get("highlights") or []),
|
||||
"nutrients": list(record.get("nutrients") or []),
|
||||
"search_query": record.get("search_query") or "",
|
||||
}
|
||||
|
||||
if record.get("embedding") is not None:
|
||||
emb = record["embedding"]
|
||||
if hasattr(emb, "__iter__"):
|
||||
product["embedding"] = [float(x) for x in emb]
|
||||
|
||||
products.append(product)
|
||||
|
||||
return products
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Export pgvector brand tables to seed catalog JSON files")
|
||||
parser.add_argument(
|
||||
"--only-having-products", action="store_true",
|
||||
help="Skip brand tables with zero products",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
export_all(only_having_products=args.only_having_products)
|
||||
34
scripts/fast_test.py
Normal file
34
scripts/fast_test.py
Normal file
@@ -0,0 +1,34 @@
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, ".")
|
||||
from app.services.rag_service import retrieve, answer_query
|
||||
from app.services.query_intent import extract_max_price, extract_brand_mention, is_count_query, detected_category
|
||||
|
||||
print("=== Q1: How many products are there in Cadbury? ===")
|
||||
print(" Is Count:", is_count_query("How many products are there in Cadbury?"))
|
||||
print(" Brand Mention:", extract_brand_mention("How many products are there in Cadbury?"))
|
||||
print(" Category:", detected_category("How many products are there in Cadbury?"))
|
||||
|
||||
print("\n=== Q2: Suggest panner less than ₹100 ===")
|
||||
print(" Is Count:", is_count_query("Suggest panner less than ₹100"))
|
||||
print(" Brand Mention:", extract_brand_mention("Suggest panner less than ₹100"))
|
||||
print(" Category:", detected_category("Suggest panner less than ₹100"))
|
||||
print(" Max Price:", extract_max_price("Suggest panner less than ₹100"))
|
||||
q2_prods = retrieve("Suggest panner less than ₹100")
|
||||
for p in q2_prods:
|
||||
print(f" - [{p.brand}] {p.product_name} | Price: {p.price_range} | Category: {p.category}")
|
||||
|
||||
print("\n=== Q3: Recommend Paneer under ₹150 ===")
|
||||
print(" Is Count:", is_count_query("Recommend Paneer under ₹150"))
|
||||
print(" Brand Mention:", extract_brand_mention("Recommend Paneer under ₹150"))
|
||||
print(" Category:", detected_category("Recommend Paneer under ₹150"))
|
||||
print(" Max Price:", extract_max_price("Recommend Paneer under ₹150"))
|
||||
q3_prods = retrieve("Recommend Paneer under ₹150")
|
||||
for p in q3_prods:
|
||||
print(f" - [{p.brand}] {p.product_name} | Price: {p.price_range} | Category: {p.category}")
|
||||
|
||||
print("\n=== Q4: Suggest low sugar biscuits ===")
|
||||
print(" Category:", detected_category("Suggest low sugar biscuits"))
|
||||
q4_prods = retrieve("Suggest low sugar biscuits")
|
||||
for p in q4_prods:
|
||||
print(f" - [{p.brand}] {p.product_name} | Price: {p.price_range} | Category: {p.category}")
|
||||
103
scripts/fix_duplicate_image_ids.py
Normal file
103
scripts/fix_duplicate_image_ids.py
Normal file
@@ -0,0 +1,103 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Fix duplicate image_ids in seed catalog JSON files.
|
||||
|
||||
Products sharing the same `title` (e.g. "Nestle Kitkat" with sizes 30g, 50g, 70g)
|
||||
previously all got the same `image_id` (generated from `title` only). Since the
|
||||
database uses `ON CONFLICT (image_id) DO UPDATE`, only one variant per title
|
||||
survived the upsert.
|
||||
|
||||
This script re-generates `image_id` from `product_name` (or `title` + `size`)
|
||||
so every product variant gets a unique, deterministic image_id.
|
||||
|
||||
Usage:
|
||||
python scripts/fix_duplicate_image_ids.py
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def make_image_id(source: str) -> str:
|
||||
"""Deterministic image_id: sanitize source name + short hash suffix."""
|
||||
sanitized = source.lower()
|
||||
sanitized = ''.join(c if c.isalnum() else '_' for c in sanitized)
|
||||
sanitized = '_'.join(filter(None, sanitized.split('_')))
|
||||
if len(sanitized) > 50:
|
||||
sanitized = sanitized[:50]
|
||||
# Use a hash of the source string so the same source always produces the same ID
|
||||
short_hash = str(uuid.uuid5(uuid.NAMESPACE_DNS, source))[:8]
|
||||
return f"{sanitized}_{short_hash}"
|
||||
|
||||
|
||||
def fix_seed_file(path: Path) -> int:
|
||||
with open(path, encoding="utf-8-sig") as f:
|
||||
data = json.load(f)
|
||||
|
||||
products = data.get("products", [])
|
||||
fixed = 0
|
||||
seen_ids: set[str] = set()
|
||||
|
||||
for product in products:
|
||||
# Determine the source for the image_id: prefer product_name, fall back to title+size, then title
|
||||
source = product.get("product_name") or ""
|
||||
if not source:
|
||||
title = product.get("title", "")
|
||||
size = product.get("size", "")
|
||||
source = f"{title}_{size}" if size else title
|
||||
|
||||
if not source:
|
||||
continue
|
||||
|
||||
new_id = make_image_id(source)
|
||||
|
||||
# Ensure no collisions across all products
|
||||
while new_id in seen_ids:
|
||||
new_id = make_image_id(source + "_" + str(uuid.uuid4())[:4])
|
||||
|
||||
old_id = product.get("image_id", "")
|
||||
if old_id != new_id:
|
||||
logger.debug(f" {old_id} -> {new_id} (source={source!r})")
|
||||
product["image_id"] = new_id
|
||||
seen_ids.add(new_id)
|
||||
fixed += 1
|
||||
|
||||
data["total_products"] = len(products)
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
logger.info(f"Fixed {fixed} products in {path.name}")
|
||||
return fixed
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not SEED_DIR.exists():
|
||||
logger.error("Seed directory not found: %s", SEED_DIR)
|
||||
sys.exit(1)
|
||||
|
||||
files = sorted(SEED_DIR.glob("*.json"))
|
||||
if not files:
|
||||
logger.warning("No JSON files found in %s", SEED_DIR)
|
||||
return
|
||||
|
||||
total = 0
|
||||
for f in files:
|
||||
total += fix_seed_file(f)
|
||||
|
||||
print(f"\nDone. Fixed {total} total products across {len(files)} file(s).")
|
||||
print("Now re-run: python scripts/seed_sample_data.py")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
147
scripts/resanitize_catalog.py
Normal file
147
scripts/resanitize_catalog.py
Normal file
@@ -0,0 +1,147 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Retroactively fix cross-category language in already-generated product
|
||||
descriptions, and recompute embeddings from the cleaned text.
|
||||
|
||||
Why this is needed
|
||||
-------------------
|
||||
This is a one-time (or occasional) *data* fix to go with the *code* fix in
|
||||
`app/services/category_registry.py` / `app/core/catalog_engine.py`. New
|
||||
catalog generations already sanitize descriptions before embedding them,
|
||||
but products ingested BEFORE that change (including the bundled
|
||||
`data/seed_catalogs/*.json` samples) may still have descriptions like:
|
||||
|
||||
"ITC Bingo Korean Style is a crispy, savory biscuit that adds a
|
||||
touch of spice..."
|
||||
|
||||
for a product whose actual category is `Snacks`, not `Biscuits & Cookies`.
|
||||
That leftover word is exactly what let a query like "recommend biscuits
|
||||
with low sugar" surface Bingo snack products - the description text (and
|
||||
therefore its embedding) said "biscuit" even though the product isn't one.
|
||||
|
||||
What this script does
|
||||
----------------------
|
||||
For every product in every `data/seed_catalogs/*.json` file:
|
||||
1. Re-run `sanitize_category_language(description, category)` to strip
|
||||
out any other category's keywords.
|
||||
2. If the description changed, recompute its embedding the same way
|
||||
`catalog_engine.py` does at ingestion time.
|
||||
3. Write the updated JSON back to disk.
|
||||
|
||||
Pass `--apply-to-db` to also push the corrected rows straight into
|
||||
pgvector (equivalent to re-running `seed_sample_data.py` afterwards).
|
||||
|
||||
Usage:
|
||||
python scripts/resanitize_catalog.py # dry-run, prints a diff summary
|
||||
python scripts/resanitize_catalog.py --write # updates the JSON files in place
|
||||
python scripts/resanitize_catalog.py --write --apply-to-db
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.category_registry import sanitize_category_language # noqa: E402
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def _process_file(path: Path, write: bool) -> List[Dict[str, Any]]:
|
||||
"""Returns the list of products whose description changed (for reporting)."""
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
products = data.get("products", [])
|
||||
changed: List[Dict[str, Any]] = []
|
||||
needs_reembed: List[int] = []
|
||||
|
||||
for i, p in enumerate(products):
|
||||
original_desc = p.get("description") or ""
|
||||
category = p.get("category") or "General"
|
||||
cleaned_desc = sanitize_category_language(original_desc, category)
|
||||
if cleaned_desc != original_desc:
|
||||
changed.append({
|
||||
"title": p.get("title") or p.get("product_name"),
|
||||
"category": category,
|
||||
"before": original_desc[:160],
|
||||
"after": cleaned_desc[:160],
|
||||
})
|
||||
p["description"] = cleaned_desc
|
||||
needs_reembed.append(i)
|
||||
|
||||
if needs_reembed and write:
|
||||
# Local import: only pay the torch/sentence-transformers startup
|
||||
# cost when there's actually something to re-embed.
|
||||
from app.services.embeddings_service import embed_texts
|
||||
|
||||
texts = []
|
||||
for i in needs_reembed:
|
||||
p = products[i]
|
||||
title = p.get("title") or ""
|
||||
cat = p.get("category") or ""
|
||||
desc = p.get("description") or ""
|
||||
texts.append(f"{title} [CAT={cat}] :: {desc}")
|
||||
vectors = embed_texts(texts)
|
||||
for i, vec in zip(needs_reembed, vectors):
|
||||
products[i]["embedding"] = vec
|
||||
logger.info("Recomputed %d embedding(s) for %s", len(needs_reembed), path.name)
|
||||
|
||||
if write and changed:
|
||||
path.write_text(json.dumps(data, indent=2, ensure_ascii=False), encoding="utf-8")
|
||||
logger.info("Wrote %d description fix(es) to %s", len(changed), path.name)
|
||||
|
||||
return changed
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--write", action="store_true", help="Write changes to the JSON files (default: dry-run report only)")
|
||||
parser.add_argument("--apply-to-db", action="store_true", help="Also upsert the corrected rows into pgvector (implies --write)")
|
||||
args = parser.parse_args()
|
||||
write = args.write or args.apply_to_db
|
||||
|
||||
if not SEED_DIR.exists():
|
||||
logger.error("Seed directory not found: %s", SEED_DIR)
|
||||
sys.exit(1)
|
||||
|
||||
total_changed = 0
|
||||
for path in sorted(SEED_DIR.glob("*.json")):
|
||||
changed = _process_file(path, write=write)
|
||||
total_changed += len(changed)
|
||||
for c in changed:
|
||||
logger.info(" [%s] %r\n before: %s...\n after: %s...",
|
||||
c["category"], c["title"], c["before"], c["after"])
|
||||
|
||||
if not write:
|
||||
logger.info(
|
||||
"%d description(s) would be changed across %d file(s). "
|
||||
"Re-run with --write to apply (and --apply-to-db to also update pgvector).",
|
||||
total_changed, len(list(SEED_DIR.glob('*.json'))),
|
||||
)
|
||||
return
|
||||
|
||||
logger.info("%d description(s) fixed.", total_changed)
|
||||
|
||||
if args.apply_to_db:
|
||||
from app.services.vector_store import ensure_brand_schema, upsert_brand_products
|
||||
|
||||
for path in sorted(SEED_DIR.glob("*.json")):
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
brand = data.get("brand")
|
||||
products = data.get("products", [])
|
||||
if not brand or not products:
|
||||
continue
|
||||
ensure_brand_schema(brand)
|
||||
upsert_brand_products(brand, products, cleanup=True)
|
||||
logger.info("Upserted %d product(s) for brand %r into pgvector", len(products), brand)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
119
scripts/seed_sample_data.py
Normal file
119
scripts/seed_sample_data.py
Normal file
@@ -0,0 +1,119 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Seed pgvector with the sample catalogs bundled in `data/seed_catalogs/`.
|
||||
|
||||
These JSON files are catalogs produced by an earlier ingestion run
|
||||
(Parle, Cadbury, Brooke Bond, Sunfeast, ...) and already include
|
||||
precomputed 384-dim embeddings, so this script does NOT call Ollama or
|
||||
the embeddings model at all - it just upserts the existing rows straight
|
||||
into pgvector. This is the fastest way to get the React UI + RAG chat
|
||||
showing real results without waiting on web scraping or a local LLM.
|
||||
|
||||
For brands not covered by the bundled samples, use
|
||||
`cli/ingest_brand.py "<brand>"` instead, which runs the full
|
||||
discovery -> images -> embeddings pipeline.
|
||||
|
||||
Usage:
|
||||
python scripts/seed_sample_data.py
|
||||
python scripts/seed_sample_data.py --only Parle Cadbury
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.vector_store import ensure_brand_schema, upsert_brand_products, _connect, _sanitize_name # noqa: E402
|
||||
from app.services.brand_registry import resolve_parent_brand # noqa: E402
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED_DIR = Path(__file__).resolve().parents[1] / "data" / "seed_catalogs"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Seed pgvector with bundled sample catalogs")
|
||||
parser.add_argument(
|
||||
"--only", nargs="*", default=None,
|
||||
help="Optional list of brand names (case-insensitive substring match on filename) to limit seeding to",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-if-seeded", action="store_true",
|
||||
help="Skip seeding if database already contains products",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.skip_if_seeded:
|
||||
from app.services.vector_store import count_products_all_brands
|
||||
try:
|
||||
cnt = count_products_all_brands()
|
||||
if cnt > 0:
|
||||
logger.info("⚡ Database already contains %d products. Skipping sample data seed.", cnt)
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if not SEED_DIR.exists():
|
||||
logger.error("Seed directory not found: %s", SEED_DIR)
|
||||
sys.exit(1)
|
||||
|
||||
files = sorted(SEED_DIR.glob("*.json"))
|
||||
if args.only:
|
||||
wanted = [w.lower() for w in args.only]
|
||||
files = [f for f in files if any(w in f.name.lower() for w in wanted)]
|
||||
|
||||
if not files:
|
||||
logger.warning("No matching seed files found in %s", SEED_DIR)
|
||||
return
|
||||
|
||||
# Aggregate products by resolved (canonical) brand so that multiple
|
||||
# seed files contributing to the same brand table are upserted as a
|
||||
# single batch. This ensures the stale-product cleanup (cleanup=True)
|
||||
# doesn't orphan products from a sibling file.
|
||||
brand_products: dict[str, list[dict]] = defaultdict(list)
|
||||
file_brand_map: dict[str, str] = {}
|
||||
|
||||
for f in files:
|
||||
data = json.loads(f.read_text(encoding="utf-8-sig"))
|
||||
products = data.get("products", [])
|
||||
brand = data.get("brand") or (products[0].get("brand_name") if products else None)
|
||||
if not brand or not products:
|
||||
logger.warning("Skipping %s - no brand/products found", f.name)
|
||||
continue
|
||||
|
||||
resolved = resolve_parent_brand(brand)
|
||||
brand_products[resolved].extend(products)
|
||||
file_brand_map[f.name] = resolved
|
||||
logger.info("Read %d products from %s -> resolved brand '%s'", len(products), f.name, resolved)
|
||||
|
||||
if not brand_products:
|
||||
logger.warning("No valid seed data found in %s", SEED_DIR)
|
||||
return
|
||||
|
||||
# Do not drop existing brand tables to preserve user database state.
|
||||
# Stale table cleanup disabled per environment requirements.
|
||||
|
||||
|
||||
total = 0
|
||||
for resolved_brand, all_products in brand_products.items():
|
||||
logger.info("Seeding %d product(s) for brand '%s'", len(all_products), resolved_brand)
|
||||
table = ensure_brand_schema(resolved_brand)
|
||||
if not table:
|
||||
logger.error("Could not create/verify table for brand '%s' - is pgvector reachable?", resolved_brand)
|
||||
continue
|
||||
upsert_brand_products(resolved_brand, all_products, cleanup=True)
|
||||
logger.info("Seeded %d products for brand '%s' (table=%s)", len(all_products), resolved_brand, table)
|
||||
total += len(all_products)
|
||||
|
||||
print(f"\nDone. Seeded {total} products across {len(brand_products)} brand(s) into pgvector.")
|
||||
print("Start the API with `uvicorn app.main:app --reload` and try a search/chat query.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
64
scripts/seed_store_intelligence.py
Normal file
64
scripts/seed_store_intelligence.py
Normal file
@@ -0,0 +1,64 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Seeds the Multi-Store Intelligence layer (Features 1, 2, 8):
|
||||
- Creates the 5 stores (Store-A .. Store-E)
|
||||
- Assigns each store a random-but-reproducible subset of the existing
|
||||
product catalog, with independent pricing and stock
|
||||
- Simulates realistic order history over the requested date range
|
||||
|
||||
Run from `backend/`:
|
||||
python scripts/seed_store_intelligence.py
|
||||
python scripts/seed_store_intelligence.py --days 120 --seed 7
|
||||
python scripts/seed_store_intelligence.py --no-reset-orders # append instead of replacing order history
|
||||
|
||||
Requires at least one brand already ingested via the existing catalog
|
||||
pipeline (`python scripts/seed_sample_data.py` or POST /api/catalog/generate) -
|
||||
this script only distributes/prices/stocks/simulates orders for
|
||||
products that already exist; it never generates new products itself.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger("seed_store_intelligence")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
parser.add_argument("--days", type=int, default=90, help="Days of order history to simulate (default: 90)")
|
||||
parser.add_argument("--seed", type=int, default=42, help="Random seed for reproducible store/order generation")
|
||||
parser.add_argument("--no-reset-orders", action="store_true", help="Append to existing order history instead of clearing it first")
|
||||
parser.add_argument("--skip-if-seeded", action="store_true", help="Skip if stores are already provisioned")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.skip_if_seeded:
|
||||
from app.services.store_db import list_stores
|
||||
try:
|
||||
stores = list_stores()
|
||||
if stores and len(stores) >= 5:
|
||||
logger.info("⚡ Stores already provisioned (%d stores). Skipping store intelligence seed.", len(stores))
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
from app.services.store_seed_service import run_seed
|
||||
|
||||
logger.info("Seeding store intelligence (days=%d, seed=%d, reset_orders=%s)...", args.days, args.seed, not args.no_reset_orders)
|
||||
result = run_seed(reset_orders=not args.no_reset_orders, days=args.days, seed=args.seed)
|
||||
|
||||
logger.info("Done.")
|
||||
logger.info(" Stores: %d", result["stores"])
|
||||
for store_id, count in result["store_products"].items():
|
||||
logger.info(" %s: %d products", store_id, count)
|
||||
logger.info(" Orders simulated: %d (%d line items)", result["orders"], result["order_items"])
|
||||
logger.info("Next: python scripts/train_ml_models.py")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
30
scripts/test_rag.py
Normal file
30
scripts/test_rag.py
Normal file
@@ -0,0 +1,30 @@
|
||||
import sys
|
||||
from app.services.rag_service import retrieve, answer_query
|
||||
from app.services.query_intent import extract_max_price, extract_brand_mention, is_count_query
|
||||
|
||||
def run_test():
|
||||
queries = [
|
||||
"How many products are there in Cadbury?",
|
||||
"Suggest low sugar biscuits",
|
||||
"Recommend Paneer under ₹150",
|
||||
"Suggest panner less than ₹100",
|
||||
]
|
||||
|
||||
for q in queries:
|
||||
print("="*60)
|
||||
print(f"QUERY: {q}")
|
||||
print(f" Is Count: {is_count_query(q)}")
|
||||
print(f" Brand Mention: {extract_brand_mention(q)}")
|
||||
print(f" Max Price: {extract_max_price(q)}")
|
||||
|
||||
prods = retrieve(q)
|
||||
print(" Retrieved Products:")
|
||||
for p in prods[:5]:
|
||||
print(f" - {p.brand} | {p.product_name} | Price: {p.price_range} | Cat: {p.category}")
|
||||
|
||||
ans = answer_query(q)
|
||||
print(f" ANSWER:\n{ans.answer}")
|
||||
print("="*60)
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_test()
|
||||
12
scripts/test_rag_response.py
Normal file
12
scripts/test_rag_response.py
Normal file
@@ -0,0 +1,12 @@
|
||||
import sys
|
||||
sys.stdout.reconfigure(encoding='utf-8')
|
||||
sys.path.insert(0, ".")
|
||||
from app.services.rag_service import answer_query
|
||||
|
||||
res = answer_query("Recommend a low sugar biscuit")
|
||||
print("=== RAG ANSWER ===")
|
||||
print(res.answer)
|
||||
print("=== SOURCES FOUND ===")
|
||||
print(len(res.sources))
|
||||
for p in res.sources[:3]:
|
||||
print(f" - {p.product_name} ({p.brand}) - {p.price_range}")
|
||||
50
scripts/train_ml_models.py
Normal file
50
scripts/train_ml_models.py
Normal file
@@ -0,0 +1,50 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Trains every ML model in `app/intelligence/` against the seeded store/
|
||||
order data (Feature 9) and refreshes cached outputs (trending
|
||||
snapshots, demand forecasts, discount history).
|
||||
|
||||
Run from `backend/` AFTER scripts/seed_store_intelligence.py:
|
||||
python scripts/train_ml_models.py
|
||||
python scripts/train_ml_models.py --models discount trending # train a subset
|
||||
|
||||
Models trained: discount, trending, popularity, forecast (demand +
|
||||
inventory), store_performance, purchase_propensity. Every trained
|
||||
model is saved to app/intelligence/artifacts/*.joblib and loaded lazily
|
||||
by the API on first use - restart the API process after retraining to
|
||||
pick up new artifacts (see docs/CHANGES.md, "Retraining").
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger("train_ml_models")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
parser.add_argument(
|
||||
"--models", nargs="+", default=None,
|
||||
choices=["discount", "trending", "popularity", "forecast", "store_performance", "purchase_propensity"],
|
||||
help="Subset of models to train (default: all)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
from app.services.ml_training_service import train_all
|
||||
|
||||
logger.info("Training models: %s", args.models or "all")
|
||||
results = train_all(models=args.models)
|
||||
|
||||
logger.info("Training summary:")
|
||||
for name, info in results.items():
|
||||
logger.info(" %s: %s", name, info)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
34
scripts/train_nutrition_models.py
Normal file
34
scripts/train_nutrition_models.py
Normal file
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
CLI entry point for Feature 8/14's ML layer - trains the nutritional-
|
||||
similarity (cosine/KNN) and nutrition-based clustering (KMeans) models
|
||||
over whatever is currently enriched in `nutrition_facts`.
|
||||
|
||||
Run after scripts/enrich_nutrition.py has populated some data.
|
||||
|
||||
Usage:
|
||||
python scripts/train_nutrition_models.py
|
||||
|
||||
Equivalent to POST /api/admin/nutrition-intelligence/train.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.nutrition_enrichment_service import train_all_models # noqa: E402
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
result = train_all_models()
|
||||
logger.info(f"Nutrition similarity index: {result['similarity']}")
|
||||
logger.info(f"Nutrition clustering: {result['clustering']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user