feat: add local chat provider and CLI
This commit is contained in:
@@ -0,0 +1,236 @@
|
||||
"""SQLite persistence for local chat sessions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
from javis.providers.base import ChatMessage
|
||||
|
||||
|
||||
class SessionStoreError(RuntimeError):
|
||||
"""Base class for expected persistence failures."""
|
||||
|
||||
|
||||
class SessionNotFoundError(SessionStoreError):
|
||||
"""A requested session ID does not exist."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ChatSession:
|
||||
id: str
|
||||
created_at: str
|
||||
updated_at: str
|
||||
provider: str
|
||||
model: str
|
||||
message_count: int = 0
|
||||
|
||||
|
||||
def _utc_now() -> str:
|
||||
return datetime.now(UTC).isoformat(timespec="seconds")
|
||||
|
||||
|
||||
class SQLiteSessionStore:
|
||||
def __init__(self, database_path: Path) -> None:
|
||||
self.database_path = database_path
|
||||
try:
|
||||
self.database_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._initialize()
|
||||
except (OSError, sqlite3.Error) as exc:
|
||||
raise SessionStoreError(
|
||||
f"Die Sitzungsdatenbank kann nicht geöffnet werden: {database_path}"
|
||||
) from exc
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
connection = sqlite3.connect(self.database_path, timeout=10)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys = ON")
|
||||
return connection
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Iterator[sqlite3.Connection]:
|
||||
connection = self._connect()
|
||||
try:
|
||||
with connection:
|
||||
yield connection
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def _initialize(self) -> None:
|
||||
with self._connection() as connection:
|
||||
connection.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id TEXT PRIMARY KEY,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
provider TEXT NOT NULL,
|
||||
model TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_id TEXT NOT NULL,
|
||||
position INTEGER NOT NULL,
|
||||
role TEXT NOT NULL CHECK (role IN ('user', 'assistant')),
|
||||
content TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL,
|
||||
FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE,
|
||||
UNIQUE (session_id, position)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_session_position
|
||||
ON messages(session_id, position);
|
||||
"""
|
||||
)
|
||||
|
||||
def create_session(self, provider: str, model: str) -> ChatSession:
|
||||
session_id = str(uuid.uuid4())
|
||||
now = _utc_now()
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO sessions (id, created_at, updated_at, provider, model)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(session_id, now, now, provider, model),
|
||||
)
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Sitzung konnte nicht gespeichert werden.") from exc
|
||||
return ChatSession(session_id, now, now, provider, model)
|
||||
|
||||
def get_session(self, session_id: str) -> ChatSession:
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT s.*, COUNT(m.id) AS message_count
|
||||
FROM sessions AS s
|
||||
LEFT JOIN messages AS m ON m.session_id = s.id
|
||||
WHERE s.id = ?
|
||||
GROUP BY s.id
|
||||
""",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Sitzung konnte nicht geladen werden.") from exc
|
||||
if row is None:
|
||||
raise SessionNotFoundError(f"Sitzung '{session_id}' wurde nicht gefunden.")
|
||||
return self._session_from_row(row)
|
||||
|
||||
def list_sessions(self) -> list[ChatSession]:
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT s.*, COUNT(m.id) AS message_count
|
||||
FROM sessions AS s
|
||||
LEFT JOIN messages AS m ON m.session_id = s.id
|
||||
GROUP BY s.id
|
||||
ORDER BY s.updated_at DESC, s.created_at DESC
|
||||
"""
|
||||
).fetchall()
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Sitzungen konnten nicht aufgelistet werden.") from exc
|
||||
return [self._session_from_row(row) for row in rows]
|
||||
|
||||
def get_messages(self, session_id: str) -> list[ChatMessage]:
|
||||
self.get_session(session_id)
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT role, content
|
||||
FROM messages
|
||||
WHERE session_id = ?
|
||||
ORDER BY position ASC
|
||||
""",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Nachrichten konnten nicht geladen werden.") from exc
|
||||
return [ChatMessage(row["role"], row["content"]) for row in rows]
|
||||
|
||||
def append_exchange(
|
||||
self,
|
||||
session_id: str,
|
||||
user_content: str,
|
||||
assistant_content: str,
|
||||
) -> None:
|
||||
now = _utc_now()
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
exists = connection.execute(
|
||||
"SELECT 1 FROM sessions WHERE id = ?", (session_id,)
|
||||
).fetchone()
|
||||
if exists is None:
|
||||
raise SessionNotFoundError(f"Sitzung '{session_id}' wurde nicht gefunden.")
|
||||
row = connection.execute(
|
||||
"SELECT COALESCE(MAX(position), -1) + 1 FROM messages WHERE session_id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
next_position = int(row[0])
|
||||
connection.executemany(
|
||||
"""
|
||||
INSERT INTO messages
|
||||
(session_id, position, role, content, created_at)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
(session_id, next_position, "user", user_content, now),
|
||||
(
|
||||
session_id,
|
||||
next_position + 1,
|
||||
"assistant",
|
||||
assistant_content,
|
||||
now,
|
||||
),
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET updated_at = ? WHERE id = ?",
|
||||
(now, session_id),
|
||||
)
|
||||
except SessionNotFoundError:
|
||||
raise
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Nachrichten konnten nicht gespeichert werden.") from exc
|
||||
|
||||
def clear_messages(self, session_id: str) -> None:
|
||||
now = _utc_now()
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
cursor = connection.execute(
|
||||
"DELETE FROM messages WHERE session_id = ?", (session_id,)
|
||||
)
|
||||
exists = connection.execute(
|
||||
"SELECT 1 FROM sessions WHERE id = ?", (session_id,)
|
||||
).fetchone()
|
||||
if exists is None:
|
||||
raise SessionNotFoundError(f"Sitzung '{session_id}' wurde nicht gefunden.")
|
||||
if cursor.rowcount:
|
||||
connection.execute(
|
||||
"UPDATE sessions SET updated_at = ? WHERE id = ?",
|
||||
(now, session_id),
|
||||
)
|
||||
except SessionNotFoundError:
|
||||
raise
|
||||
except sqlite3.Error as exc:
|
||||
raise SessionStoreError("Die Sitzung konnte nicht geleert werden.") from exc
|
||||
|
||||
@staticmethod
|
||||
def _session_from_row(row: sqlite3.Row) -> ChatSession:
|
||||
return ChatSession(
|
||||
id=row["id"],
|
||||
created_at=row["created_at"],
|
||||
updated_at=row["updated_at"],
|
||||
provider=row["provider"],
|
||||
model=row["model"],
|
||||
message_count=row["message_count"],
|
||||
)
|
||||
Reference in New Issue
Block a user