File size: 4,716 Bytes
2edb151 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | 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)
|