from __future__ import annotations import sqlite3 from app.config import Settings from app.db import find_catalog_by_sku, knn from app.embed import catalog_passage_text, embed_texts, line_query_text, vendor_query_text from app.schemas import LineItem, MatchBand, MatchHit, ReceiptExtract from backends.base import EmbedBackend def distance_to_similarity(distance: float) -> float: return 1.0 - distance def band_for(similarity: float, *, auto: float, review: float) -> MatchBand: if similarity >= auto: return MatchBand.auto if similarity >= review: return MatchBand.review return MatchBand.unmatched def unmatched(reason: str) -> MatchHit: return MatchHit(similarity=0.0, band=MatchBand.unmatched, reason=reason) def match_line_item( con: sqlite3.Connection, settings: Settings, embed: EmbedBackend, extract: ReceiptExtract, item: LineItem, *, k: int = 5, ) -> MatchHit: if item.sku: row = find_catalog_by_sku(con, item.sku) if row is not None: return MatchHit( catalog_id=int(row["id"]), sku=row["sku"], vendor=row["vendor"], description=row["description"], similarity=1.0, band=MatchBand.exact, reason="exact sku", ) catalog_count = con.execute("SELECT COUNT(*) AS n FROM catalog").fetchone()["n"] if catalog_count == 0: return unmatched("empty catalog") query_vec = embed_texts( embed, [line_query_text(extract, item)], input_type="query", settings=settings )[0] hits = knn(con, "catalog_vec", "catalog_id", query_vec, k=k) if not hits: return unmatched("no vectors") catalog_id, distance = hits[0] similarity = distance_to_similarity(distance) row = con.execute("SELECT * FROM catalog WHERE id = ?", (catalog_id,)).fetchone() return MatchHit( catalog_id=catalog_id, sku=None if row is None else row["sku"], vendor=None if row is None else row["vendor"], description=None if row is None else row["description"], similarity=similarity, band=band_for(similarity, auto=settings.sku_auto, review=settings.sku_review), reason="knn", ) def match_vendor( con: sqlite3.Connection, settings: Settings, embed: EmbedBackend, vendor: str, *, k: int = 5, ) -> MatchHit: if not vendor.strip(): return unmatched("no vendor") exact = con.execute( "SELECT * FROM catalog WHERE vendor = ? COLLATE NOCASE LIMIT 1", (vendor,) ).fetchone() if exact is not None: return MatchHit( catalog_id=int(exact["id"]), sku=exact["sku"], vendor=exact["vendor"], description=exact["description"], similarity=1.0, band=MatchBand.exact, reason="exact vendor", ) query_vec = embed_texts( embed, [vendor_query_text(vendor)], input_type="query", settings=settings )[0] hits = knn(con, "catalog_vec", "catalog_id", query_vec, k=k) if not hits: return unmatched("no vectors") catalog_id, distance = hits[0] similarity = distance_to_similarity(distance) row = con.execute("SELECT * FROM catalog WHERE id = ?", (catalog_id,)).fetchone() return MatchHit( catalog_id=catalog_id, sku=None if row is None else row["sku"], vendor=None if row is None else row["vendor"], description=None if row is None else row["description"], similarity=similarity, band=band_for( similarity, auto=settings.vendor_auto, review=settings.vendor_review ), reason="vendor knn", ) def match_receipt( con: sqlite3.Connection, settings: Settings, embed: EmbedBackend, extract: ReceiptExtract, ) -> list[MatchHit]: hits = [match_line_item(con, settings, embed, extract, item) for item in extract.line_items] if extract.vendor: hits.append(match_vendor(con, settings, embed, extract.vendor)) return hits def embed_catalog_row( con: sqlite3.Connection, settings: Settings, embed: EmbedBackend, catalog_id: int, ) -> None: from app.db import upsert_vector row = con.execute("SELECT * FROM catalog WHERE id = ?", (catalog_id,)).fetchone() if row is None: return text = catalog_passage_text( vendor=row["vendor"], sku=row["sku"], description=row["description"], size=row["size"], ) vec = embed_texts(embed, [text], input_type="passage", settings=settings)[0] upsert_vector(con, "catalog_vec", "catalog_id", catalog_id, vec)