diff options
Diffstat (limited to 'packages/meshbay-node/src/meshbay_node/chat/store.py')
| -rw-r--r-- | packages/meshbay-node/src/meshbay_node/chat/store.py | 119 |
1 files changed, 119 insertions, 0 deletions
diff --git a/packages/meshbay-node/src/meshbay_node/chat/store.py b/packages/meshbay-node/src/meshbay_node/chat/store.py new file mode 100644 index 0000000..1dbcc2b --- /dev/null +++ b/packages/meshbay-node/src/meshbay_node/chat/store.py @@ -0,0 +1,119 @@ +""" +MeshBay Node — SQLite-backed chat message store. + +One database per group. Stores encrypted Sender Keys messages for offline +retrieval and history. Messages are stored as received (ciphertext) — +decryption happens on the client side. +""" + +import logging +import time +from dataclasses import dataclass +from pathlib import Path + +import aiosqlite + +log = logging.getLogger(__name__) + +_SCHEMA = """ +CREATE TABLE IF NOT EXISTS messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + sender_id TEXT NOT NULL, + iteration INTEGER NOT NULL, + payload BLOB NOT NULL, + timestamp REAL NOT NULL, + thread_id TEXT DEFAULT NULL +); +CREATE INDEX IF NOT EXISTS idx_messages_ts ON messages(timestamp); +CREATE INDEX IF NOT EXISTS idx_messages_thread ON messages(thread_id); +""" + + +@dataclass +class StoredMessage: + id: int + sender_id: str + iteration: int + payload: bytes + timestamp: float + thread_id: str | None + + +class ChatStore: + """Async SQLite chat store for one group.""" + + def __init__(self, db_path: Path): + self._db_path = db_path + self._db: aiosqlite.Connection | None = None + + async def open(self) -> None: + self._db_path.parent.mkdir(parents=True, exist_ok=True) + self._db = await aiosqlite.connect(str(self._db_path)) + await self._db.executescript(_SCHEMA) + await self._db.commit() + + async def close(self) -> None: + if self._db: + await self._db.close() + self._db = None + + async def __aenter__(self): + await self.open() + return self + + async def __aexit__(self, *_): + await self.close() + + async def save_message( + self, + sender_id: str, + iteration: int, + payload: bytes, + thread_id: str | None = None, + ) -> int: + """Store a message. Returns the row id.""" + ts = time.time() + cursor = await self._db.execute( + "INSERT INTO messages (sender_id, iteration, payload, timestamp, thread_id) " + "VALUES (?, ?, ?, ?, ?)", + (sender_id, iteration, payload, ts, thread_id), + ) + await self._db.commit() + return cursor.lastrowid + + async def get_messages( + self, + since: float = 0, + limit: int = 100, + ) -> list[StoredMessage]: + """Get messages after a timestamp, most recent last.""" + cursor = await self._db.execute( + "SELECT id, sender_id, iteration, payload, timestamp, thread_id " + "FROM messages WHERE timestamp > ? ORDER BY timestamp ASC LIMIT ?", + (since, limit), + ) + rows = await cursor.fetchall() + return [ + StoredMessage(id=r[0], sender_id=r[1], iteration=r[2], + payload=r[3], timestamp=r[4], thread_id=r[5]) + for r in rows + ] + + async def get_thread(self, thread_id: str, limit: int = 100) -> list[StoredMessage]: + """Get messages in a thread.""" + cursor = await self._db.execute( + "SELECT id, sender_id, iteration, payload, timestamp, thread_id " + "FROM messages WHERE thread_id = ? ORDER BY timestamp ASC LIMIT ?", + (thread_id, limit), + ) + rows = await cursor.fetchall() + return [ + StoredMessage(id=r[0], sender_id=r[1], iteration=r[2], + payload=r[3], timestamp=r[4], thread_id=r[5]) + for r in rows + ] + + async def message_count(self) -> int: + cursor = await self._db.execute("SELECT COUNT(*) FROM messages") + row = await cursor.fetchone() + return row[0] |