summaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/src/meshbay_node/chat/store.py
blob: 1dbcc2bb5836b7d4d4f07588a9bd6f2186a6c57b (plain) (blame)
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
"""
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]