| 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) |
|
|