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)