summaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/src/meshbay_node/chat/store.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/meshbay-node/src/meshbay_node/chat/store.py')
-rw-r--r--packages/meshbay-node/src/meshbay_node/chat/store.py119
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]