import os import io import shutil from fastapi import FastAPI, UploadFile, File, Form, HTTPException from fastapi.staticfiles import StaticFiles from fastapi.responses import FileResponse from fastapi.middleware.cors import CORSMiddleware from PIL import Image from models import CLIPEmbedder from vector_store import VectorStore app = FastAPI() app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) os.makedirs("static", exist_ok=True) os.makedirs("static/images", exist_ok=True) if not os.path.exists("catalog.json"): import generate_catalog generate_catalog.main() embedder = CLIPEmbedder() store = VectorStore() @app.get("/api/products") async def get_products(): return store.get_all_products() @app.post("/api/search") async def search(q: str = Form(None), file: UploadFile = File(None)): is_image_search = False if file: try: content = await file.read() image = Image.open(io.BytesIO(content)) vector = embedder.get_image_embedding(image) is_image_search = True except Exception as e: raise HTTPException(status_code=400, detail=str(e)) elif q: try: vector = embedder.get_text_embedding(q) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) else: raise HTTPException(status_code=400, detail="Query text or image file is required") # Request 7 results to allow filtering out exact match while keeping 6 related products results = store.search_by_vector(vector, top_k=7) exact_match = None if is_image_search: # Visual exact match based on similarity threshold EXACT_MATCH_THRESHOLD = 0.92 if results and results[0]["score"] >= EXACT_MATCH_THRESHOLD: exact_match = results[0] elif q: # Text exact match based on case-insensitive ID or Title match q_clean = q.strip().lower() for product in store.get_all_products(): if q_clean == product["title"].strip().lower() or q_clean == product["id"].strip().lower(): score = 1.0 for r in results: if r["product"]["id"] == product["id"]: score = r["score"] break exact_match = { "product": product, "score": score } break # Filter out the exact match from the related products list related_products = [] for res in results: if exact_match and res["product"]["id"] == exact_match["product"]["id"]: continue related_products.append(res) return { "exact_match": exact_match, "related_products": related_products[:6] } @app.post("/api/search/similar") async def search_similar(payload: dict): product_id = payload.get("product_id") if not product_id: raise HTTPException(status_code=400, detail="product_id is required") return store.search_similar(product_id) @app.post("/api/index") async def index_product( id: str = Form(...), title: str = Form(...), desc: str = Form(...), price: float = Form(...), category: str = Form(...), file: UploadFile = File(...) ): try: content = await file.read() image = Image.open(io.BytesIO(content)) image_path = f"static/images/{id}.png" with open(image_path, "wb") as buffer: buffer.write(content) vector = embedder.get_image_embedding(image) product = { "id": id, "title": title, "desc": desc, "price": price, "category": category, "image_url": f"/static/images/{id}.png" } store.add_product(product, vector) return {"status": "success", "product": product} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) app.mount("/", StaticFiles(directory="static", html=True), name="static")