diff --git a/backend/app/routers/products.py b/backend/app/routers/products.py index bbcabfb..ae68d24 100644 --- a/backend/app/routers/products.py +++ b/backend/app/routers/products.py @@ -1,3 +1,6 @@ +import difflib +import re + from fastapi import APIRouter, BackgroundTasks, Depends, File, HTTPException, UploadFile, status from fastapi.responses import Response from sqlalchemy.orm import Session, joinedload, selectinload @@ -9,6 +12,7 @@ from ..models import ( Barcode, BaseUnit, Category, + CategoryTracking, Group, Location, Movement, @@ -25,6 +29,9 @@ from ..schemas import ( BarcodeCreate, LocationMinStockIn, LookupResult, + MatchCandidate, + MatchLine, + MatchRequest, ProductCreate, ProductOut, ProductUpdate, @@ -35,6 +42,7 @@ from ..schemas import ( from ..services import images from ..services.categories import suggest_category from .categories import descendant_ids +from .settings import get_receipt_match_threshold from ..services.conversion import ConversionError, resolve_product_unit from ..services.fields import FieldError, apply_field_values from ..services.group_codes import detach as detach_group_code, sync as sync_group_code @@ -42,6 +50,76 @@ from ..services.stock import removal_stats router = APIRouter(prefix="/products", tags=["products"]) +# --------------------------------------------------------------------------- +# Kassenzettel-Abgleich: eine OCR-Zeile den ähnlichsten Lebensmitteln zuordnen. +# --------------------------------------------------------------------------- +_UMLAUTE = str.maketrans({"ä": "ae", "ö": "oe", "ü": "ue", "ß": "ss"}) + + +def _normalisieren(text: str) -> str: + """Klein, Umlaute aufgelöst, nur Buchstaben/Ziffern – Preise/Sonderzeichen weg.""" + text = (text or "").lower().translate(_UMLAUTE) + return re.sub(r"[^a-z0-9]+", " ", text).strip() + + +def _match_score(zeile: str, name: str) -> int: + """Ähnlichkeit 0–100 zwischen Kassenzeile und Produktname. + + Kombiniert die Gesamt-Ähnlichkeit mit einem Token-/Präfix-Abgleich, damit + abgekürzte Kassennamen ("MÜHLEN SCHNITZ") auf den vollen Namen passen. + """ + a, b = _normalisieren(zeile), _normalisieren(name) + if not a or not b: + return 0 + gesamt = difflib.SequenceMatcher(None, a, b).ratio() + a_tok, b_tok = a.split(), b.split() + bester = [] + for t in a_tok: + m = 0.0 + for u in b_tok: + if u.startswith(t) or t.startswith(u): + m = max(m, min(len(t), len(u)) / max(len(t), len(u))) + else: + m = max(m, difflib.SequenceMatcher(None, t, u).ratio()) + bester.append(m) + token = sum(bester) / len(bester) if bester else 0.0 + return round(100 * max(gesamt, token)) + + +@router.post("/match", response_model=list[MatchLine]) +def match_receipt( + payload: MatchRequest, + db: Session = Depends(get_db), + _: User = Depends(get_current_user), +) -> list[MatchLine]: + """Kassenzettel-Zeilen den ähnlichsten LEBENSMITTELN zuordnen (Score 0–100). + Gegenstände (auch Verbrauchsgegenstände) bleiben außen vor; je Zeile die besten + Treffer über dem Schwellwert (aus Request oder Einstellung).""" + schwelle = payload.threshold if payload.threshold is not None else get_receipt_match_threshold(db) + schwelle = max(0, min(100, schwelle)) + lebensmittel = [ + p + for p in db.query(Product).options(joinedload(Product.category)).all() + if product_tracking(db, p) != CategoryTracking.object.value + ] + ergebnis: list[MatchLine] = [] + for zeile in payload.lines: + text = (zeile or "").strip() + treffer = [(p, _match_score(text, p.name)) for p in lebensmittel] + treffer = sorted( + (ps for ps in treffer if ps[1] >= schwelle), + key=lambda ps: ps[1], + reverse=True, + ) + ergebnis.append(MatchLine( + text=text, + candidates=[ + MatchCandidate(product_id=p.id, name=p.name, brand=p.brand, score=s) + for p, s in treffer[:5] + ], + )) + return ergebnis + @router.get("", response_model=list[ProductOut]) def list_products( diff --git a/backend/app/routers/settings.py b/backend/app/routers/settings.py index 418df99..3c4f53f 100644 --- a/backend/app/routers/settings.py +++ b/backend/app/routers/settings.py @@ -10,6 +10,10 @@ from ..schemas import SettingOut router = APIRouter(prefix="/settings", tags=["settings"]) EXPIRY_WARNING_KEY = "expiry_warning_days" +# Ab wie viel Prozent Übereinstimmung ein Artikel beim Kassenzettel-Scan als +# Treffer vorgeschlagen wird. +RECEIPT_THRESHOLD_KEY = "receipt_match_threshold" +RECEIPT_THRESHOLD_DEFAULT = 45 def get_expiry_warning_days(db: Session) -> int: @@ -22,6 +26,17 @@ def get_expiry_warning_days(db: Session) -> int: return get_settings().expiry_warning_days_default +def get_receipt_match_threshold(db: Session) -> int: + """Schwellwert (0–100) für Kassenzettel-Treffer; auf sinnvollen Bereich geklemmt.""" + row = db.get(Setting, RECEIPT_THRESHOLD_KEY) + if row is None: + return RECEIPT_THRESHOLD_DEFAULT + try: + return max(0, min(100, int(row.value))) + except ValueError: + return RECEIPT_THRESHOLD_DEFAULT + + @router.get("", response_model=list[SettingOut]) def list_settings( db: Session = Depends(get_db), _: User = Depends(get_current_user) @@ -29,6 +44,7 @@ def list_settings( rows = db.query(Setting).all() known = {r.key: r.value for r in rows} known.setdefault(EXPIRY_WARNING_KEY, str(get_expiry_warning_days(db))) + known.setdefault(RECEIPT_THRESHOLD_KEY, str(get_receipt_match_threshold(db))) return [SettingOut(key=k, value=v) for k, v in known.items()] diff --git a/backend/app/schemas.py b/backend/app/schemas.py index 76881c8..de724c4 100644 --- a/backend/app/schemas.py +++ b/backend/app/schemas.py @@ -487,6 +487,25 @@ class LotSplit(BaseModel): location_id: str | None = None +# ---- Kassenzettel-Abgleich ---- +class MatchRequest(BaseModel): + """Kassenzettel-Zeilen (OCR) → beste Lebensmittel-Treffer je Zeile.""" + lines: list[str] + threshold: int | None = None # 0–100; ohne Angabe gilt die Einstellung + + +class MatchCandidate(BaseModel): + product_id: int + name: str + brand: str | None = None + score: int # 0–100, Ähnlichkeit zur Kassenzeile + + +class MatchLine(BaseModel): + text: str + candidates: list[MatchCandidate] = [] + + class CheckInResponse(BaseModel): lot: LotOut product_stock: float diff --git a/backend/tests/test_receipt_match.py b/backend/tests/test_receipt_match.py new file mode 100644 index 0000000..695e101 --- /dev/null +++ b/backend/tests/test_receipt_match.py @@ -0,0 +1,60 @@ +"""Kassenzettel-Abgleich: OCR-Zeilen den ähnlichsten Lebensmitteln zuordnen.""" + +import pytest + +from app.models import BaseUnit, Product, Role, User +from app.routers.products import match_receipt +from app.schemas import MatchRequest + + +@pytest.fixture() +def user(db): + person = User(username="tester", password_hash="x", role=Role.admin) + db.add(person) + db.commit() + db.refresh(person) + return person + + +def test_abgekuerzte_kassenzeile_trifft_lebensmittel(db, user): + p = Product( + name="Vegane Mühlen-Schnitzel auf Basis von Soja", + base_unit=BaseUnit.gram, package_size=180, + ) + db.add(p) + db.commit() + db.refresh(p) + + res = match_receipt( + MatchRequest(lines=["MÜHLEN SCHNITZ", "voelliger unsinn xyz"]), + db=db, _=user, + ) + # Abgekürzte Kassenzeile findet den vollen Namen mit hohem Score. + assert res[0].text == "MÜHLEN SCHNITZ" + assert res[0].candidates + assert res[0].candidates[0].product_id == p.id + assert res[0].candidates[0].score >= 70 + # Unsinnszeile bleibt ohne Treffer über dem Schwellwert. + assert res[1].candidates == [] + + +def test_nur_lebensmittel_keine_gegenstaende(db, user): + # Einzelstück-Flag -> tracking "object" -> darf nicht vorgeschlagen werden. + obj = Product(name="Powerbank Anker", base_unit=BaseUnit.piece, individual=True) + db.add(obj) + db.commit() + + res = match_receipt(MatchRequest(lines=["POWERBANK ANKER"], threshold=10), db=db, _=user) + assert res[0].candidates == [] + + +def test_schwellwert_filtert(db, user): + p = Product(name="Basmati Reis", base_unit=BaseUnit.gram, package_size=1) + db.add(p) + db.commit() + + # Nur teilweise passende Zeile: bei sehr hohem Schwellwert kein Treffer. + hart = match_receipt(MatchRequest(lines=["reis lose"], threshold=99), db=db, _=user) + weich = match_receipt(MatchRequest(lines=["reis lose"], threshold=20), db=db, _=user) + assert hart[0].candidates == [] + assert weich[0].candidates and weich[0].candidates[0].product_id == p.id