| """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 |
|
|
| |
| def db_init(): |
| sb = _supabase() |
| if sb: return |
| 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() |
|
|
| |
| 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() |
|
|
| |
| 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<?",(now,)) |
| c.commit(); n = cur.rowcount; c.close(); return n |
|
|
| |
| def insert_invoice(user_id: int, data: dict) -> 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 |
|
|