iTagPDF / backend /database.py
peyajm
Sync latest backend and pipeline code
fb7b1a4
Raw
History Blame
8.69 kB
"""
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()