drowzeys's picture
v1.0 alpha: keys-Auto Receipts Studio (iPhone / may add Autonomous Lamp Skill)
2edb151 verified
Raw
History Blame Contribute Delete
4.72 kB
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)