"""OmniParse AI — Database layer. SQLite fallback + optional Supabase. All queries parameterized.""" import sqlite3, json, secrets as _secrets from datetime import datetime, timezone from config import SQLITE_PATH, USE_SUPABASE, SUPABASE_URL, SUPABASE_KEY _sb = None def _supabase(): global _sb if _sb is None and USE_SUPABASE: try: from supabase import create_client as sc _sb = sc(SUPABASE_URL, SUPABASE_KEY) except Exception as e: print(f"[WARN] Supabase init failed: {e}") return None return _sb def _sqlite(): c = sqlite3.connect(SQLITE_PATH) c.row_factory = sqlite3.Row c.execute("PRAGMA journal_mode=WAL") c.execute("PRAGMA foreign_keys=ON") return c # ── Init ──────────────────────────────────────────────────────────────────── def db_init(): sb = _supabase() if sb: return # tables created in Supabase SQL editor c = _sqlite() c.executescript(""" CREATE TABLE IF NOT EXISTS users ( id INTEGER PRIMARY KEY AUTOINCREMENT, email TEXT UNIQUE NOT NULL COLLATE NOCASE, name TEXT NOT NULL, password_hash TEXT NOT NULL, plan TEXT NOT NULL DEFAULT 'free', stripe_cid TEXT, api_key TEXT UNIQUE, email_verified INTEGER NOT NULL DEFAULT 0, is_admin INTEGER NOT NULL DEFAULT 0, created_at TEXT NOT NULL, updated_at TEXT ); CREATE TABLE IF NOT EXISTS sessions ( token TEXT PRIMARY KEY, user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, expires_at TEXT NOT NULL, created_at TEXT NOT NULL ); CREATE TABLE IF NOT EXISTS invoices ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, filename TEXT, vendor TEXT, inv_number TEXT, inv_date TEXT, due_date TEXT, amount REAL, vat_amount REAL, total REAL, currency TEXT DEFAULT 'USD', status TEXT DEFAULT 'done', is_duplicate INTEGER DEFAULT 0, confidence REAL, raw_json TEXT, created_at TEXT NOT NULL ); CREATE INDEX IF NOT EXISTS ix_inv_user ON invoices(user_id); CREATE INDEX IF NOT EXISTS ix_inv_date ON invoices(created_at); CREATE INDEX IF NOT EXISTS ix_ses_user ON sessions(user_id); CREATE INDEX IF NOT EXISTS ix_ses_exp ON sessions(expires_at); """) c.commit(); c.close() # ── Users ─────────────────────────────────────────────────────────────────── def create_user(email: str, name: str, password_hash: str, plan: str = "free") -> dict | None: email = email.strip().lower() api_key = "op_live_" + _secrets.token_hex(20) now = datetime.now(timezone.utc).isoformat() sb = _supabase() if sb: ex = sb.table("users").select("id").eq("email", email).execute() if ex.data: return None r = sb.table("users").insert({"email":email,"name":name,"password_hash":password_hash, "plan":plan,"api_key":api_key,"created_at":now}).execute() return r.data[0] if r.data else None c = _sqlite() try: cur = c.execute("INSERT INTO users(email,name,password_hash,plan,api_key,created_at) VALUES(?,?,?,?,?,?)", (email, name, password_hash, plan, api_key, now)) c.commit(); uid = cur.lastrowid row = c.execute("SELECT * FROM users WHERE id=?", (uid,)).fetchone() return dict(row) except sqlite3.IntegrityError: return None finally: c.close() def get_user_by_email(email: str) -> dict | None: email = email.strip().lower() sb = _supabase() if sb: r = sb.table("users").select("*").eq("email",email).execute() return r.data[0] if r.data else None c = _sqlite() row = c.execute("SELECT * FROM users WHERE email=? COLLATE NOCASE", (email,)).fetchone() c.close(); return dict(row) if row else None def get_user_by_id(uid: int) -> dict | None: sb = _supabase() if sb: r = sb.table("users").select("*").eq("id",uid).execute() return r.data[0] if r.data else None c = _sqlite() row = c.execute("SELECT * FROM users WHERE id=?", (uid,)).fetchone() c.close(); return dict(row) if row else None def update_user(uid: int, fields: dict) -> None: fields["updated_at"] = datetime.now(timezone.utc).isoformat() sb = _supabase() if sb: sb.table("users").update(fields).eq("id",uid).execute(); return c = _sqlite() cols = ", ".join(f"{k}=?" for k in fields) vals = list(fields.values()) + [uid] c.execute(f"UPDATE users SET {cols} WHERE id=?", vals) c.commit(); c.close() def delete_user(uid: int) -> None: sb = _supabase() if sb: sb.table("invoices").delete().eq("user_id",uid).execute() sb.table("sessions").delete().eq("user_id",uid).execute() sb.table("users").delete().eq("id",uid).execute(); return c = _sqlite() c.execute("DELETE FROM invoices WHERE user_id=?",(uid,)) c.execute("DELETE FROM sessions WHERE user_id=?",(uid,)) c.execute("DELETE FROM users WHERE id=?",(uid,)) c.commit(); c.close() # ── Sessions ──────────────────────────────────────────────────────────────── def create_session(user_id: int, token: str, expires_at: str) -> None: now = datetime.now(timezone.utc).isoformat() sb = _supabase() if sb: sb.table("sessions").insert({"token":token,"user_id":user_id,"expires_at":expires_at,"created_at":now}).execute() return c = _sqlite() c.execute("INSERT INTO sessions(token,user_id,expires_at,created_at) VALUES(?,?,?,?)",(token,user_id,expires_at,now)) c.commit(); c.close() def get_session(token: str) -> dict | None: if not token: return None sb = _supabase() if sb: r = sb.table("sessions").select("*").eq("token",token).execute() if not r.data: return None session = r.data[0] else: c = _sqlite() row = c.execute("SELECT * FROM sessions WHERE token=?",(token,)).fetchone() c.close() if not row: return None session = dict(row) try: expires = datetime.fromisoformat(session["expires_at"]) if expires.tzinfo is None: expires = expires.replace(tzinfo=timezone.utc) if expires < datetime.now(timezone.utc): delete_session(token); return None except (ValueError, KeyError): return None return session def delete_session(token: str) -> None: if not token: return sb = _supabase() if sb: sb.table("sessions").delete().eq("token",token).execute(); return c = _sqlite() c.execute("DELETE FROM sessions WHERE token=?",(token,)) c.commit(); c.close() def cleanup_sessions() -> int: now = datetime.now(timezone.utc).isoformat() sb = _supabase() if sb: r = sb.table("sessions").delete().lt("expires_at",now).execute() return len(r.data) if r.data else 0 c = _sqlite() cur = c.execute("DELETE FROM sessions WHERE expires_at dict: now = datetime.now(timezone.utc).isoformat() row = { "user_id": user_id, "filename": data.get("filename"), "vendor": data.get("vendor"), "inv_number": data.get("invoice_number"), "inv_date": data.get("invoice_date"), "due_date": data.get("due_date"), "amount": data.get("amount"), "vat_amount": data.get("vat_amount"), "total": data.get("total"), "currency": data.get("currency","USD"), "status": data.get("status","done"), "is_duplicate": bool(data.get("is_duplicate", False)), "confidence": data.get("confidence"), "raw_json": json.dumps(data, ensure_ascii=False, default=str), "created_at": now, } sb = _supabase() if sb: r = sb.table("invoices").insert(row).execute() return r.data[0] if r.data else row c = _sqlite() cols = ", ".join(row.keys()) phs = ", ".join("?" for _ in row) c.execute(f"INSERT INTO invoices({cols}) VALUES({phs})", tuple(row.values())) c.commit() iid = c.execute("SELECT last_insert_rowid()").fetchone()[0] res = c.execute("SELECT * FROM invoices WHERE id=?",(iid,)).fetchone() c.close(); return dict(res) def get_invoices(user_id: int) -> list[dict]: sb = _supabase() if sb: r = sb.table("invoices").select("*").eq("user_id",user_id).order("created_at",desc=True).execute() return r.data or [] c = _sqlite() rows = c.execute("SELECT * FROM invoices WHERE user_id=? ORDER BY created_at DESC",(user_id,)).fetchall() c.close(); return [dict(r) for r in rows] def delete_invoice(user_id: int, invoice_id: int) -> None: sb = _supabase() if sb: sb.table("invoices").delete().eq("id",invoice_id).eq("user_id",user_id).execute(); return c = _sqlite() c.execute("DELETE FROM invoices WHERE id=? AND user_id=?",(invoice_id,user_id)) c.commit(); c.close() def count_invoices_this_month(user_id: int) -> int: invs = get_invoices(user_id); now = datetime.now(timezone.utc); n = 0 for inv in invs: try: cr = datetime.fromisoformat(inv["created_at"]) if cr.year == now.year and cr.month == now.month: n += 1 except (ValueError, KeyError): pass return n def check_duplicate(user_id: int, vendor: str, total: float) -> bool: if not vendor or total is None: return False invs = get_invoices(user_id); now = datetime.now(timezone.utc) vl = vendor.strip().lower() for inv in invs: try: cr = datetime.fromisoformat(inv["created_at"]) except: continue if cr.year == now.year and cr.month == now.month: iv = (inv.get("vendor") or "").strip().lower() it = inv.get("total") if iv == vl and it is not None and abs(float(it)-float(total))<0.01: return True return False