ocado-grocy/ocado_grocy/grocy.py

304 lines
15 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Small Grocy API client and resumable importer."""
from decimal import Decimal
from dataclasses import replace
import hashlib
import html
import json
import re
from pathlib import Path
import sqlite3
import requests
from .product_metadata import read_metadata
from .receipt import ImportError, number, receipt_date
def normalized(name):
return " ".join(name.casefold().split())
class Grocy:
def __init__(self, url, api_key, timeout=30):
self.url = url.rstrip("/")
if not self.url.endswith("/api"):
self.url += "/api"
self.session = requests.Session()
self.session.headers.update({"GROCY-API-KEY": api_key,
"User-Agent": "ocado-grocy/0.1"})
self.timeout = timeout
def request(self, method, path, data=None):
try:
response = self.session.request(method, self.url + path, json=data,
timeout=self.timeout, allow_redirects=False)
except requests.RequestException as exc:
raise ImportError(f"Grocy {method} {path} failed ({type(exc).__name__}); "
"check connectivity before retrying") from exc
if not 200 <= response.status_code < 300:
raise ImportError(f"Grocy {method} {path}: HTTP {response.status_code}")
try:
return response.json() if response.content else None
except ValueError as exc:
raise ImportError(f"Grocy {method} {path} returned non-JSON data") from exc
def get(self, path):
return self.request("GET", path)
def post(self, path, data):
return self.request("POST", path, data)
def named_object(self, entity, name, **fields):
matches = [x for x in self.get(f"/objects/{entity}")
if normalized(x["name"]) == normalized(name)]
if len(matches) > 1:
raise ImportError(f"Multiple Grocy {entity} named {name!r}")
if matches:
return int(matches[0]["id"])
return int(self.post(f"/objects/{entity}", {"name": name, **fields})["created_object_id"])
class Journal:
"""Commit intent before each write; ambiguous writes require explicit reconciliation."""
def __init__(self, path: Path):
path.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
self.db = sqlite3.connect(path)
self.db.execute("""CREATE TABLE IF NOT EXISTS imports (
server TEXT, order_id TEXT, line INTEGER, fingerprint TEXT,
status TEXT, product_id INTEGER, PRIMARY KEY(server, order_id, line))""")
self.db.commit()
def check(self, server, order, line, fingerprint):
row = self.db.execute("SELECT fingerprint, status FROM imports WHERE server=? AND order_id=? AND line=?",
(server, order, line)).fetchone()
if not row:
return False
if row[0] != fingerprint:
raise ImportError(f"Order {order} line {line} changed since a previous import; reconcile it manually")
if row[1] != "done":
raise ImportError(f"Order {order} line {line} has an uncertain previous write. "
"Check Grocy stock/logs and reconcile the journal before retrying (see README).")
return True
def record(self, server, order, line, fingerprint, status, product_id):
with self.db:
self.db.execute("INSERT OR REPLACE INTO imports VALUES (?, ?, ?, ?, ?, ?)",
(server, order, line, fingerprint, status, product_id))
def marker(item):
return f"Ocado product ID: {item.product_id}" if item.product_id else ""
def description(item):
details = {"Ocado product ID": item.product_id, "Barcode": item.barcode,
"Product URL": item.url, **item.metadata}
return "<p>" + html.escape(marker(item) or "Imported from Ocado") + "</p><pre>" + html.escape(
json.dumps(details, ensure_ascii=False, indent=2, default=str)) + "</pre>"
def match_product(item, products, barcodes, mappings):
override = mappings.get(item.product_id or item.name)
if override is not None:
matches = [p for p in products if int(p["id"]) == int(override["product_id"])]
if not matches:
raise ImportError(f"Mapped Grocy product does not exist: {item.name}")
return matches[0], override
matches = []
if item.product_id:
matches = [p for p in products if f"<p>{html.escape(marker(item))}</p>" in (p.get("description") or "")]
if not matches and item.barcode:
ids = {int(b["product_id"]) for b in barcodes if b["barcode"] == item.barcode}
matches = [p for p in products if int(p["id"]) in ids]
if not matches:
matches = [p for p in products if normalized(p["name"]) == normalized(item.name)]
if len(matches) > 1:
raise ImportError(f"Ambiguous existing product: {item.name}; add a config product mapping")
return (matches[0] if matches else None), {}
def item_fingerprint(receipt, item):
return hashlib.sha256(json.dumps({
"name": item.name, "product_id": item.product_id, "quantity": str(item.quantity.normalize()),
"total": str(item.total.normalize()), "best_before": item.best_before,
"date": receipt.purchased_date}, sort_keys=True).encode()).hexdigest()
def align_recorded_lines(receipt, server, journal):
"""Restore journal positions when Ocado reorders equal-expiry receipt rows."""
records = journal.db.execute(
"SELECT line, fingerprint FROM imports WHERE server=? AND order_id=? ORDER BY line",
(server, receipt.order_id)).fetchall()
remaining = list(receipt.items)
aligned = [None] * len(remaining)
for line, fingerprint in records:
candidates = [i for i,item in enumerate(remaining) if item_fingerprint(receipt,item) == fingerprint]
if not candidates or not 1 <= line <= len(aligned):
raise ImportError(f"Order {receipt.order_id} line {line} changed since its import; review before continuing")
aligned[line - 1] = remaining.pop(candidates[0])
rest = iter(remaining)
return replace(receipt, items=[item if item is not None else next(rest) for item in aligned])
def refresh_imported_products(receipt, client, journal, *, seen=None):
"""Refresh details of journal-linked products, including excluded perishables."""
seen = seen if seen is not None else set()
receipt = align_recorded_lines(receipt, client.url, journal)
records = journal.db.execute(
"SELECT line, product_id FROM imports WHERE server=? AND order_id=? AND status='done' ORDER BY line",
(client.url, receipt.order_id)).fetchall()
count = 0
for line, product_id in records:
if product_id in seen:
continue
if not 1 <= line <= len(receipt.items):
raise ImportError(f"Receipt line {line} no longer exists; cannot refresh order {receipt.order_id}")
item = receipt.items[line - 1]
product = client.get(f"/objects/products/{product_id}")
old = product.get("description") or ""
expected = f"<p>{html.escape(marker(item))}</p>" if item.product_id else ""
if expected and expected not in old:
raise ImportError(f"Product identity changed for order {receipt.order_id}, line {line}; review before refreshing")
if not expected and normalized(product["name"]) != normalized(item.name):
raise ImportError(f"Product name changed for order {receipt.order_id}, line {line}")
metadata = read_metadata(old)
for key, value in item.metadata.items():
if value not in (None, "", {}, []):
if isinstance(value, dict) and isinstance(metadata.get(key), dict):
metadata[key] = {**metadata[key], **value}
else:
metadata[key] = value
barcode = item.barcode or metadata.pop("Barcode", "")
url = item.url or metadata.pop("Product URL", "")
for key in ("Ocado product ID", "Barcode", "Product URL"):
metadata.pop(key, None)
updated = description(replace(item, metadata=metadata, barcode=barcode, url=url))
if old != updated:
client.request("PUT", f"/objects/products/{product_id}", {"description": updated})
count += 1
seen.add(product_id)
return count
def import_receipt(receipt, client, journal, config, echo=print, *, include_lines=None):
if not receipt.purchased_date:
raise ImportError("No purchase/delivery date found; supply --purchased-date YYYY-MM-DD")
products = client.get("/objects/products")
barcodes = client.get("/objects/product_barcodes")
plans = []
# Validate every existing product/conversion before any mutations.
for line, item in enumerate(receipt.items, 1):
if include_lines is not None and line not in include_lines:
continue
fingerprint = item_fingerprint(receipt, item)
if journal.check(client.url, receipt.order_id, line, fingerprint):
echo(f"Skipped already imported: {item.name}")
continue
product, override = match_product(item, products, barcodes, config.get("products", {}))
factor = number(override.get("stock_per_purchase", 1), "stock conversion")
if product:
if product.get("enable_tare_weight_handling") or product.get("no_own_stock"):
raise ImportError(f"Unsupported tare-weight or parent-only product: {item.name}")
if "stock_per_purchase" not in override and product["qu_id_purchase"] != product["qu_id_stock"]:
details = client.get(f'/stock/products/{product["id"]}')
factor = number(details.get("qu_conversion_factor_purchase_to_stock"), "stock conversion")
if factor <= 0:
raise ImportError(f"Stock conversion must be positive: {item.name}")
plans.append((line, item, fingerprint, product, factor))
if not plans:
return 0
location = client.named_object("locations", config.get("location", "Ocado imports"))
unit = client.named_object("quantity_units", config.get("quantity_unit", "Pack"), name_plural="Packs")
store = client.named_object("shopping_locations", "Ocado")
count = 0
for line, item, fingerprint, product, factor in plans:
if product is None:
# Recheck the local list: the same product can occur on multiple receipt lines.
product, _ = match_product(item, products, barcodes, config.get("products", {}))
if product is None:
payload = {"name": item.name, "description": description(item), "location_id": location,
"shopping_location_id": store, "qu_id_purchase": unit, "qu_id_stock": unit,
"qu_id_consume": unit, "qu_id_price": unit}
product_id = int(client.post("/objects/products", payload)["created_object_id"])
product = {"id": product_id, **payload}
products.append(product)
if item.barcode:
client.post("/objects/product_barcodes", {"product_id": product_id, "barcode": item.barcode,
"qu_id": unit, "amount": 1})
product_id = int(product["id"])
amount = item.quantity * factor
note = f"Ocado {receipt.purchased_date}"
if not item.best_before:
note += " Expiry unknown."
payload = {"amount": float(amount), "price": float(item.total / amount),
"best_before_date": item.best_before or "2999-12-31",
"purchased_date": receipt.purchased_date, "transaction_type": "purchase",
"shopping_location_id": store, "stock_label_type": 1, "note": note}
journal.record(client.url, receipt.order_id, line, fingerprint, "pending", product_id)
client.post(f"/stock/products/{product_id}/add", payload)
journal.record(client.url, receipt.order_id, line, fingerprint, "done", product_id)
echo(f"Imported {item.quantity} × {item.name}: £{item.total:.2f}")
count += 1
return count
def update_existing_notes(client, journal):
"""Edit only surviving journal-linked stock entries; never book new stock."""
imported = {(str(order), int(line), int(product)) for order, line, product in journal.db.execute(
"SELECT order_id, line, product_id FROM imports WHERE server=? AND status='done'", (client.url,))}
changed = {}
preserved_fields = ("amount", "best_before_date", "price", "open", "location_id",
"shopping_location_id", "purchased_date")
for row in client.get("/objects/stock"):
old = row.get("note") or ""
match = re.match(r"Ocado order (\d+), line (\d+)[.;]", old)
dated = re.fullmatch(r"Ocado (\d{4}-\d{2}-\d{2}), line (\d+)\.(?: Expiry unknown\.)?", old)
if match:
linked = (match[1], int(match[2]), int(row["product_id"])) in imported
elif dated:
linked = any(line == int(dated[2]) and product == int(row["product_id"])
for _, line, product in imported)
else:
linked = False
if not linked:
continue
try:
current = client.get(f"/stock/entry/{row['id']}")
except ImportError:
if not any(entry["id"] == row["id"] for entry in client.get("/objects/stock")):
continue
raise
if current is None or current.get("note") != old:
continue
row = current
purchased = receipt_date(row["purchased_date"])
if not purchased:
raise ImportError(f"Stock entry {row['id']} has no purchase date")
note = f"Ocado {purchased}"
if "Expiry unknown" in old:
note += " Expiry unknown."
payload = {field: row[field] for field in preserved_fields}
try:
client.request("PUT", f"/stock/entry/{row['id']}", {**payload, "note": note})
except ImportError:
if not any(entry["id"] == row["id"] for entry in client.get("/objects/stock")):
continue
raise
changed[row["id"]] = {**payload, "note": note}
for row in client.get("/objects/stock"):
if row["id"] in changed:
if row["note"] != changed[row["id"]]["note"]:
raise ImportError(f"Verification failed for stock entry {row['id']}")
return len(changed)
def imported_products(client, state, order_ids=()):
"""Live products from completed journal lines, optionally scoped to orders."""
query = "SELECT DISTINCT product_id FROM imports WHERE server=? AND status='done'"
params = [client.url]
if order_ids:
query += ' AND order_id IN (' + ','.join('?' for _ in order_ids) + ')'
params.extend(order_ids)
with sqlite3.connect(f'{Path(state).resolve().as_uri()}?mode=ro', uri=True) as db:
ids = {int(row[0]) for row in db.execute(query, params)}
return [p for p in client.get('/objects/products') if int(p['id']) in ids]