summaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/tests/test_chat_store.py
blob: 2533921aa74c82e94e8b76c7ad777acc88ade5e6 (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
"""
Tests for the SQLite-backed chat message store.
"""


import pytest
import pytest_asyncio
from meshbay_node.chat.store import ChatStore


@pytest_asyncio.fixture
async def store(tmp_path):
    s = ChatStore(db_path=tmp_path / "test_chat.db")
    await s.open()
    yield s
    await s.close()


@pytest.mark.asyncio
async def test_save_and_retrieve(store):
    row_id = await store.save_message(
        sender_id="alice", iteration=0, payload=b"hello",
    )
    assert row_id == 1

    msgs = await store.get_messages()
    assert len(msgs) == 1
    assert msgs[0].sender_id == "alice"
    assert msgs[0].iteration == 0
    assert msgs[0].payload == b"hello"
    assert msgs[0].thread_id is None


@pytest.mark.asyncio
async def test_message_count(store):
    assert await store.message_count() == 0
    await store.save_message("alice", 0, b"msg1")
    await store.save_message("bob", 1, b"msg2")
    assert await store.message_count() == 2


@pytest.mark.asyncio
async def test_get_messages_since(store):
    await store.save_message("alice", 0, b"old")
    all_msgs = await store.get_messages()
    cutoff = all_msgs[0].timestamp
    await store.save_message("bob", 1, b"new")

    msgs = await store.get_messages(since=cutoff)
    assert len(msgs) == 1
    assert msgs[0].sender_id == "bob"


@pytest.mark.asyncio
async def test_thread_messages(store):
    await store.save_message("alice", 0, b"root", thread_id="t1")
    await store.save_message("bob", 1, b"reply", thread_id="t1")
    await store.save_message("carol", 2, b"other")

    thread = await store.get_thread("t1")
    assert len(thread) == 2
    assert thread[0].sender_id == "alice"
    assert thread[1].sender_id == "bob"


@pytest.mark.asyncio
async def test_message_ordering(store):
    for i in range(5):
        await store.save_message(f"user-{i}", i, f"msg-{i}".encode())

    msgs = await store.get_messages()
    assert len(msgs) == 5
    for i, m in enumerate(msgs):
        assert m.sender_id == f"user-{i}"


@pytest.mark.asyncio
async def test_limit(store):
    for i in range(10):
        await store.save_message("alice", i, f"msg-{i}".encode())

    msgs = await store.get_messages(limit=3)
    assert len(msgs) == 3


@pytest.mark.asyncio
async def test_context_manager(tmp_path):
    async with ChatStore(db_path=tmp_path / "ctx_test.db") as store:
        await store.save_message("alice", 0, b"test")
        assert await store.message_count() == 1