File size: 8,690 Bytes
fb7b1a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
"""
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()