File size: 7,388 Bytes
447450c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e21048b
 
 
 
 
 
 
 
 
 
 
447450c
 
 
e21048b
447450c
 
e21048b
 
 
 
 
 
 
 
 
 
447450c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
SQLite database setup — async aiosqlite/libsql with table creation on startup.
"""

from contextlib import asynccontextmanager
from pathlib import Path
import aiosqlite
from app.config import get_settings


class LibsqlRow(dict):
    """aiosqlite.Row compatible row wrapper for libsql-client."""
    def __init__(self, colnames, row_data):
        super().__init__(zip(colnames, row_data))
        self._list = list(row_data)

    def __getitem__(self, key):
        if isinstance(key, int):
            return self._list[key]
        return super().__getitem__(key)


class LibsqlCursorWrapper:
    """Cursor-like wrapper for remote libsql execution. Supports both awaiting and async with."""
    def __init__(self, connection, client, sql, parameters):
        self.connection = connection
        self.client = client
        self.sql = sql
        self.parameters = parameters
        self.result_set = None
        self.row_factory = connection.row_factory
        self._index = 0

    def __await__(self):
        return self._execute().__await__()

    async def _execute(self):
        if self.result_set is None:
            self.result_set = await self.client.execute(self.sql, self.parameters)
            if self.row_factory == aiosqlite.Row:
                self.row_factory = lambda cursor, row_data: LibsqlRow(self.result_set.columns, row_data)
        return self

    @property
    def rowcount(self) -> int:
        return getattr(self.result_set, "rows_affected", 0) if self.result_set is not None else 0

    async def fetchone(self):
        await self._execute()
        if self._index >= len(self.result_set.rows):
            return None
        row = self.result_set.rows[self._index]
        self._index += 1
        if self.row_factory:
            return self.row_factory(self, row)
        return row

    async def fetchall(self):
        await self._execute()
        rows = self.result_set.rows[self._index:]
        self._index = len(self.result_set.rows)
        if self.row_factory:
            return [self.row_factory(self, r) for r in rows]
        return rows

    async def __aenter__(self):
        await self._execute()
        return self

    async def __aexit__(self, exc_type, exc_val, exc_tb):
        pass


class LibsqlConnectionWrapper:
    """Connection-like wrapper for remote libsql client."""
    def __init__(self, client):
        self.client = client
        self._row_factory = None

    @property
    def row_factory(self):
        return self._row_factory

    @row_factory.setter
    def row_factory(self, factory):
        self._row_factory = factory

    def execute(self, sql: str, parameters: tuple = ()):
        # Return a wrapper that is both awaitable and an async context manager
        return LibsqlCursorWrapper(self, self.client, sql, parameters)

    async def executescript(self, script: str):
        # Split statements by semicolon and execute each
        statements = [stmt.strip() for stmt in script.split(";") if stmt.strip()]
        for stmt in statements:
            await self.client.execute(stmt)

    async def commit(self):
        # libsql auto-commits
        pass

    async def __aenter__(self):
        return self

    async def __aexit__(self, exc_type, exc_val, exc_tb):
        pass


_global_client = None


async def close_db():
    """Close the global database client if active."""
    global _global_client
    if _global_client is not None:
        await _global_client.close()
        _global_client = None


@asynccontextmanager
async def connect_db():
    """Yield a database connection wrapper based on active settings (Turso vs. local sqlite)."""
    global _global_client
    settings = get_settings()
    if settings.TURSO_DATABASE_URL and settings.TURSO_AUTH_TOKEN:
        if _global_client is None:
            url = settings.TURSO_DATABASE_URL
            if url.startswith("libsql://"):
                url = "https://" + url[len("libsql://"):]
            import libsql_client
            _global_client = libsql_client.create_client(
                url=url,
                auth_token=settings.TURSO_AUTH_TOKEN
            )
        yield LibsqlConnectionWrapper(_global_client)
    else:
        db_path = settings.DB_PATH
        Path(db_path).parent.mkdir(parents=True, exist_ok=True)
        async with aiosqlite.connect(db_path) as db:
            yield db


async def get_db_path() -> str:
    settings = get_settings()
    Path(settings.DB_PATH).parent.mkdir(parents=True, exist_ok=True)
    return settings.DB_PATH


async def init_db() -> None:
    """Create all tables if they don't exist."""
    async with connect_db() as db:
        await db.executescript("""
            CREATE TABLE IF NOT EXISTS users (
                username TEXT PRIMARY KEY,
                password_hash TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS user_sessions (
                token TEXT PRIMARY KEY,
                username TEXT NOT NULL REFERENCES users(username) ON DELETE CASCADE,
                expires_at TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS sessions (
                id TEXT PRIMARY KEY,
                user_id TEXT NOT NULL,
                title TEXT,
                created_at TEXT NOT NULL,
                last_active_at TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS messages (
                id TEXT PRIMARY KEY,
                session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
                role TEXT NOT NULL,
                content TEXT NOT NULL,
                intent TEXT,
                confidence TEXT,
                response_id TEXT,
                citations TEXT,
                pipeline_trace TEXT,
                created_at TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS feedback (
                id TEXT PRIMARY KEY,
                response_id TEXT NOT NULL,
                session_id TEXT NOT NULL,
                rating INTEGER NOT NULL,
                comment TEXT,
                created_at TEXT NOT NULL
            );

            CREATE TABLE IF NOT EXISTS user_documents (
                id TEXT PRIMARY KEY,
                user_id TEXT NOT NULL,
                filename TEXT NOT NULL,
                file_type TEXT,
                collection_name TEXT NOT NULL,
                chunk_count INTEGER DEFAULT 0,
                file_size_bytes INTEGER DEFAULT 0,
                uploaded_at TEXT NOT NULL
            );

            CREATE INDEX IF NOT EXISTS idx_messages_session ON messages(session_id);
            CREATE INDEX IF NOT EXISTS idx_sessions_user ON sessions(user_id);
            CREATE INDEX IF NOT EXISTS idx_user_docs_user ON user_documents(user_id);
            CREATE INDEX IF NOT EXISTS idx_feedback_response ON feedback(response_id);
        """)
        
        # Run-time schema migration check
        db.row_factory = aiosqlite.Row
        async with db.execute("PRAGMA table_info(messages)") as cur:
            columns = [row["name"] for row in await cur.fetchall()]
        
        if "citations" not in columns:
            await db.execute("ALTER TABLE messages ADD COLUMN citations TEXT")
        if "pipeline_trace" not in columns:
            await db.execute("ALTER TABLE messages ADD COLUMN pipeline_trace TEXT")
            
        await db.commit()