""" SQLite database for user accounts and session metadata. Uses Python's built-in sqlite3; all DB calls are wrapped with fastapi.concurrency.run_in_threadpool so they don't block the event loop. """ import sqlite3 import shutil from pathlib import Path from typing import Optional from config import DB_PATH, SESSIONS_DIR SESSION_LIMIT = 999 # effectively unlimited; sessions kept until user deletes manually def _connect() -> sqlite3.Connection: DB_PATH.parent.mkdir(parents=True, exist_ok=True) conn = sqlite3.connect(str(DB_PATH), check_same_thread=False) conn.row_factory = sqlite3.Row return conn def init_db() -> None: """Create tables if they don't exist. Called once at startup.""" conn = _connect() try: conn.executescript(""" CREATE TABLE IF NOT EXISTS users ( id TEXT PRIMARY KEY, google_id TEXT UNIQUE NOT NULL, email TEXT NOT NULL, name TEXT, picture TEXT, created_at TEXT NOT NULL DEFAULT (datetime('now')) ); CREATE TABLE IF NOT EXISTS pipeline_sessions ( session_id TEXT PRIMARY KEY, user_id TEXT NOT NULL REFERENCES users(id), status TEXT NOT NULL DEFAULT 'idle', pdf_name TEXT, last_stage INTEGER NOT NULL DEFAULT 0, created_at TEXT NOT NULL DEFAULT (datetime('now')), updated_at TEXT NOT NULL DEFAULT (datetime('now')) ); CREATE TABLE IF NOT EXISTS user_events ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id TEXT NOT NULL REFERENCES users(id), session_id TEXT, event_type TEXT NOT NULL, meta TEXT, created_at TEXT NOT NULL DEFAULT (datetime('now')) ); """) # Migration: add columns that may not exist in older databases for migration in [ "ALTER TABLE pipeline_sessions ADD COLUMN last_stage INTEGER NOT NULL DEFAULT 0", "ALTER TABLE users ADD COLUMN openai_key TEXT", ]: try: conn.execute(migration) conn.commit() except sqlite3.OperationalError: pass # column already exists finally: conn.close() # --------------------------------------------------------------------------- # Users # --------------------------------------------------------------------------- def upsert_user(google_id: str, email: str, name: str, picture: str) -> dict: """Insert or update user; return the user row as a dict.""" import uuid conn = _connect() try: row = conn.execute( "SELECT * FROM users WHERE google_id = ?", (google_id,) ).fetchone() if row: conn.execute( "UPDATE users SET name=?, picture=? WHERE google_id=?", (name, picture, google_id), ) conn.commit() row = conn.execute( "SELECT * FROM users WHERE google_id = ?", (google_id,) ).fetchone() else: user_id = str(uuid.uuid4()) conn.execute( "INSERT INTO users (id, google_id, email, name, picture) VALUES (?,?,?,?,?)", (user_id, google_id, email, name, picture), ) conn.commit() row = conn.execute( "SELECT * FROM users WHERE id = ?", (user_id,) ).fetchone() return dict(row) finally: conn.close() def get_user_by_id(user_id: str) -> Optional[dict]: conn = _connect() try: row = conn.execute("SELECT * FROM users WHERE id = ?", (user_id,)).fetchone() return dict(row) if row else None finally: conn.close() def get_user_openai_key(user_id: str) -> Optional[str]: conn = _connect() try: row = conn.execute("SELECT openai_key FROM users WHERE id = ?", (user_id,)).fetchone() return row["openai_key"] if row else None finally: conn.close() def set_user_openai_key(user_id: str, key: str) -> None: conn = _connect() try: conn.execute("UPDATE users SET openai_key = ? WHERE id = ?", (key, user_id)) conn.commit() finally: conn.close() # --------------------------------------------------------------------------- # Pipeline sessions # --------------------------------------------------------------------------- def save_session_record(session_id: str, user_id: str, pdf_name: str, status: str = "idle") -> None: conn = _connect() try: conn.execute( """INSERT OR REPLACE INTO pipeline_sessions (session_id, user_id, status, pdf_name, updated_at) VALUES (?, ?, ?, ?, datetime('now'))""", (session_id, user_id, status, pdf_name), ) conn.commit() finally: conn.close() def update_session_status(session_id: str, status: str, last_stage: int = None) -> None: conn = _connect() try: if last_stage is not None: conn.execute( "UPDATE pipeline_sessions SET status=?, last_stage=?, updated_at=datetime('now') WHERE session_id=?", (status, last_stage, session_id), ) else: conn.execute( "UPDATE pipeline_sessions SET status=?, updated_at=datetime('now') WHERE session_id=?", (status, session_id), ) conn.commit() finally: conn.close() def get_sessions_for_user(user_id: str) -> list[dict]: conn = _connect() try: rows = conn.execute( """SELECT session_id, status, pdf_name, last_stage, created_at, updated_at FROM pipeline_sessions WHERE user_id=? ORDER BY updated_at DESC""", (user_id,), ).fetchall() return [dict(r) for r in rows] finally: conn.close() def get_session_record(session_id: str) -> Optional[dict]: conn = _connect() try: row = conn.execute( "SELECT * FROM pipeline_sessions WHERE session_id=?", (session_id,) ).fetchone() return dict(row) if row else None finally: conn.close() def delete_session_record(session_id: str) -> None: conn = _connect() try: conn.execute("DELETE FROM pipeline_sessions WHERE session_id=?", (session_id,)) conn.commit() finally: conn.close() def get_sessions_beyond_limit(user_id: str, limit: int = SESSION_LIMIT) -> list[str]: """Return session_ids (oldest first) that exceed the per-user limit.""" conn = _connect() try: rows = conn.execute( """SELECT session_id FROM pipeline_sessions WHERE user_id=? ORDER BY updated_at DESC""", (user_id,), ).fetchall() ids = [r["session_id"] for r in rows] return ids[limit:] # everything beyond the limit, oldest last finally: conn.close() def enforce_session_limit(user_id: str) -> list[str]: """ Delete sessions beyond SESSION_LIMIT for the user. Returns list of deleted session_ids so caller can remove disk data. """ to_delete = get_sessions_beyond_limit(user_id) for sid in to_delete: delete_session_record(sid) # Remove session directory from disk session_dir = SESSIONS_DIR / sid if session_dir.exists(): shutil.rmtree(session_dir, ignore_errors=True) return to_delete # --------------------------------------------------------------------------- # User events # --------------------------------------------------------------------------- def log_event(user_id: str, event_type: str, session_id: str = None, meta: dict = None) -> None: """Record a user interaction event.""" import json conn = _connect() try: conn.execute( "INSERT INTO user_events (user_id, session_id, event_type, meta) VALUES (?,?,?,?)", (user_id, session_id, event_type, json.dumps(meta) if meta else None), ) conn.commit() finally: conn.close() def get_events_for_user(user_id: str, limit: int = 500) -> list[dict]: """Return recent events for a user, newest first.""" conn = _connect() try: rows = conn.execute( "SELECT * FROM user_events WHERE user_id=? ORDER BY created_at DESC LIMIT ?", (user_id, limit), ).fetchall() return [dict(r) for r in rows] finally: conn.close()