Reworked Markov storage to reduce database size.
CI / Formatting (push) Failing after 38s
CI / Linting (push) Successful in 8s
CI / Tests (push) Successful in 35s
CI / Type Checking (push) Failing after 12s
CI / Spelling (push) Successful in 7s

This commit is contained in:
2026-06-29 12:29:19 -04:00
parent d23c11cf11
commit 1dd9db9151
5 changed files with 934 additions and 73 deletions
+495 -5
View File
@@ -18,8 +18,11 @@ Tests cover Database.connect (pragmas, schema), markov start word and
transition CRUD, image storage, flag CRUD, and channel ingestion tracking.
"""
from typing import TYPE_CHECKING
import logging
from collections import Counter
from typing import TYPE_CHECKING, Any
import aiosqlite
import pytest
from crabstero.database import ChannelImage, Database, StartWord, Transition
@@ -53,6 +56,370 @@ class TestConnect:
"markov_transitions",
]
async def test_schema_version(self, db: Database) -> None:
"""Fresh databases are marked with the current schema version."""
async with db._connection.execute("PRAGMA user_version") as cursor:
row = await cursor.fetchone()
assert row is not None
assert row[0] == 1
async def test_connection_is_closed_on_setup_failure(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Failed setup closes the connection before re-raising."""
db_path = str(tmp_path / "future.db")
future = await aiosqlite.connect(db_path, isolation_level=None)
try:
await future.execute("PRAGMA user_version = 2")
finally:
await future.close()
opened_connections: list[aiosqlite.Connection] = []
real_connect = aiosqlite.connect
def spy_connect(*args: Any, **kwargs: Any) -> aiosqlite.Connection:
connection = real_connect(*args, **kwargs)
opened_connections.append(connection)
return connection
monkeypatch.setattr(aiosqlite, "connect", spy_connect)
with pytest.raises(RuntimeError, match="newer than supported"):
await Database.connect(db_path)
assert len(opened_connections) == 1
opened_connection = opened_connections[0]
try:
with pytest.raises(ValueError, match="no active connection"):
await opened_connection.execute("SELECT 1")
finally:
await opened_connection.close()
async def test_markov_schema_uses_counted_rows(self, db: Database) -> None:
"""Markov tables use occurrence counters and optimized primary keys."""
async with db._connection.execute(
"PRAGMA table_info(markov_start_words)",
) as cursor:
start_columns = await cursor.fetchall()
async with db._connection.execute(
"PRAGMA table_info(markov_transitions)",
) as cursor:
transition_columns = await cursor.fetchall()
assert [(row[1], row[5]) for row in start_columns] == [
("channel_id", 1),
("user_id", 3),
("word", 2),
("occurrences", 0),
]
assert [(row[1], row[5]) for row in transition_columns] == [
("channel_id", 1),
("user_id", 4),
("word", 2),
("next_word", 3),
("occurrences", 0),
]
async with db._connection.execute(
"SELECT name, sql FROM sqlite_master"
" WHERE type = 'table'"
" AND name IN ('markov_start_words', 'markov_transitions')",
) as cursor:
table_sql = {
str(row[0]): str(row[1]) for row in await cursor.fetchall()
}
assert "occurrences INTEGER NOT NULL CHECK (occurrences > 0)" in table_sql[
"markov_start_words"
]
assert "WITHOUT ROWID" in table_sql["markov_start_words"]
assert "WITHOUT ROWID" in table_sql["markov_transitions"]
async def test_markov_indexes_exist(self, db: Database) -> None:
"""Expected explicit Markov indexes exist on fresh databases."""
async with db._connection.execute(
"SELECT name FROM sqlite_master"
" WHERE type = 'index'"
" AND tbl_name IN ('markov_start_words', 'markov_transitions')"
" ORDER BY name",
) as cursor:
indexes = [row[0] for row in await cursor.fetchall()]
assert indexes == [
"idx_start_user",
"idx_transitions_user",
]
async def test_old_markov_schema_is_converted(
self,
tmp_path: Path,
caplog: pytest.LogCaptureFixture,
) -> None:
"""Old duplicate Markov rows are converted into counted rows."""
db_path = str(tmp_path / "legacy.db")
legacy = await aiosqlite.connect(db_path, isolation_level=None)
try:
await legacy.executescript(
"""
CREATE TABLE markov_start_words (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL
);
CREATE INDEX idx_start_channel
ON markov_start_words(channel_id, word);
CREATE INDEX idx_start_user ON markov_start_words(user_id);
CREATE TABLE markov_transitions (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL,
next_word TEXT NOT NULL
);
CREATE INDEX idx_transitions_channel_word
ON markov_transitions(channel_id, word);
CREATE INDEX idx_transitions_user ON markov_transitions(user_id);
CREATE TABLE channel_images (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
url TEXT NOT NULL
);
CREATE INDEX idx_images_channel ON channel_images(channel_id);
CREATE INDEX idx_images_user ON channel_images(user_id);
CREATE TABLE flags (
entity_type TEXT NOT NULL,
entity_id TEXT NOT NULL,
flag_name TEXT NOT NULL,
PRIMARY KEY (entity_type, entity_id, flag_name)
);
CREATE TABLE ingested_channels (
channel_id INTEGER NOT NULL PRIMARY KEY
);
INSERT INTO markov_start_words
(channel_id, user_id, word)
VALUES
(1, 100, 'Hello'),
(1, 100, 'Hello'),
(1, 200, 'Hello'),
(2, 100, 'Hello');
INSERT INTO markov_transitions
(channel_id, user_id, word, next_word)
VALUES
(1, 100, 'Hello', 'world.'),
(1, 100, 'Hello', 'world.'),
(1, 100, 'Hello', 'friend.'),
(1, 200, 'Hello', 'world.');
INSERT INTO channel_images
(channel_id, user_id, url)
VALUES
(1, 100, 'https://example.com/a.png');
INSERT INTO flags
(entity_type, entity_id, flag_name)
VALUES
('user', '100', 'allowPings');
INSERT INTO ingested_channels (channel_id) VALUES (1);
""",
)
finally:
await legacy.close()
with caplog.at_level(logging.INFO, logger="crabstero.database"):
migrated = await Database.connect(db_path)
assert "Migrating legacy Markov tables to counted occurrences" in caplog.text
assert "Legacy Markov table migration complete." in caplog.text
assert "Post-migration database vacuum complete." in caplog.text
try:
async with migrated._connection.execute(
"SELECT channel_id, user_id, word, occurrences"
" FROM markov_start_words"
" ORDER BY channel_id, user_id, word",
) as cursor:
start_rows = [tuple(row) for row in await cursor.fetchall()]
assert start_rows == [
(1, 100, "Hello", 2),
(1, 200, "Hello", 1),
(2, 100, "Hello", 1),
]
async with migrated._connection.execute(
"SELECT channel_id, user_id, word, next_word, occurrences"
" FROM markov_transitions"
" ORDER BY channel_id, user_id, word, next_word",
) as cursor:
transition_rows = [tuple(row) for row in await cursor.fetchall()]
assert transition_rows == [
(1, 100, "Hello", "friend.", 1),
(1, 100, "Hello", "world.", 2),
(1, 200, "Hello", "world.", 1),
]
assert await migrated.get_random_image(1) == "https://example.com/a.png"
assert await migrated.is_flag_set("user", "100", "allowPings") is True
assert await migrated.is_channel_ingested(1) is True
async with migrated._connection.execute("PRAGMA user_version") as cursor:
row = await cursor.fetchone()
assert row is not None
assert row[0] == 1
async with migrated._connection.execute(
"SELECT name FROM sqlite_master"
" WHERE type = 'table'"
" AND name IN ("
" 'markov_start_words_v0',"
" 'markov_transitions_v0'"
" )",
) as cursor:
assert await cursor.fetchall() == []
async with migrated._connection.execute(
"SELECT name FROM sqlite_master"
" WHERE type = 'index'"
" AND tbl_name IN ('markov_start_words', 'markov_transitions')"
" ORDER BY name",
) as cursor:
indexes = [row[0] for row in await cursor.fetchall()]
assert indexes == [
"idx_start_user",
"idx_transitions_user",
]
finally:
await migrated.close()
reopened = await Database.connect(db_path)
try:
async with reopened._connection.execute(
"SELECT channel_id, user_id, word, occurrences"
" FROM markov_start_words"
" ORDER BY channel_id, user_id, word",
) as cursor:
assert [tuple(row) for row in await cursor.fetchall()] == start_rows
finally:
await reopened.close()
async def test_failed_post_migration_vacuum_is_retried(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
"""A migrated database is not marked current until vacuum succeeds."""
db_path = str(tmp_path / "legacy.db")
legacy = await aiosqlite.connect(db_path, isolation_level=None)
try:
await legacy.executescript(
"""
CREATE TABLE markov_start_words (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL
);
CREATE TABLE markov_transitions (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL,
next_word TEXT NOT NULL
);
INSERT INTO markov_start_words
(channel_id, user_id, word)
VALUES
(1, 100, 'Hello'),
(1, 100, 'Hello');
INSERT INTO markov_transitions
(channel_id, user_id, word, next_word)
VALUES
(1, 100, 'Hello', 'world.'),
(1, 100, 'Hello', 'world.');
""",
)
finally:
await legacy.close()
real_execute = aiosqlite.Connection.execute
vacuum_failures = 0
def fail_first_vacuum(
connection: aiosqlite.Connection,
sql: str,
parameters: Any | None = None,
) -> Any:
nonlocal vacuum_failures
if sql.strip().upper() == "VACUUM" and vacuum_failures == 0:
vacuum_failures += 1
raise RuntimeError("simulated vacuum failure")
return real_execute(connection, sql, parameters)
monkeypatch.setattr(aiosqlite.Connection, "execute", fail_first_vacuum)
with caplog.at_level(logging.INFO, logger="crabstero.database"):
migrated = await Database.connect(db_path)
assert (
"Database vacuum failed after legacy Markov table migration"
in caplog.text
)
assert vacuum_failures == 1
try:
async with migrated._connection.execute("PRAGMA user_version") as cursor:
row = await cursor.fetchone()
assert row is not None
assert row[0] == 0
finally:
await migrated.close()
caplog.clear()
with caplog.at_level(logging.INFO, logger="crabstero.database"):
reopened = await Database.connect(db_path)
assert "Retrying post-migration database vacuum" in caplog.text
assert "Post-migration database vacuum complete." in caplog.text
try:
async with reopened._connection.execute("PRAGMA user_version") as cursor:
row = await cursor.fetchone()
assert row is not None
assert row[0] == 1
finally:
await reopened.close()
async def test_unsupported_counted_markov_schema_is_rejected(
self,
tmp_path: Path,
) -> None:
"""Counted Markov tables must match the optimized schema layout."""
db_path = str(tmp_path / "unsupported.db")
counted = await aiosqlite.connect(db_path, isolation_level=None)
try:
await counted.executescript(
"""
CREATE TABLE markov_start_words (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL,
occurrences INTEGER NOT NULL CHECK (occurrences > 0),
PRIMARY KEY (channel_id, user_id, word)
);
CREATE TABLE markov_transitions (
channel_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
word TEXT NOT NULL,
next_word TEXT NOT NULL,
occurrences INTEGER NOT NULL CHECK (occurrences > 0),
PRIMARY KEY (channel_id, user_id, word, next_word)
);
PRAGMA user_version = 1;
""",
)
finally:
await counted.close()
with pytest.raises(RuntimeError, match="unsupported counted Markov"):
await Database.connect(db_path)
class TestAddMarkovData:
"""Markov start word and transition storage via add_markov_data."""
@@ -83,6 +450,32 @@ class TestAddMarkovData:
await db.add_markov_data([], [])
assert await db.get_random_start_word(1) is None
async def test_duplicate_inputs_increment_occurrences(self, db: Database) -> None:
"""Duplicate inputs in one call increment one counted row."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "Hello")],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "world."),
],
)
async with db._connection.execute(
"SELECT occurrences FROM markov_start_words"
" WHERE channel_id = 1 AND user_id = 100 AND word = 'Hello'",
) as cursor:
row = await cursor.fetchone()
assert row == (2,)
async with db._connection.execute(
"SELECT occurrences FROM markov_transitions"
" WHERE channel_id = 1"
" AND user_id = 100"
" AND word = 'Hello'"
" AND next_word = 'world.'",
) as cursor:
row = await cursor.fetchone()
assert row == (2,)
class TestMarkovReadMethods:
"""Markov read methods return None for out-of-scope or missing data."""
@@ -183,6 +576,54 @@ class TestMarkovReadMethods:
"friend.",
}
async def test_same_candidate_aggregates_across_users(self, db: Database) -> None:
"""Identical candidates contributed by multiple users are still one choice."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 200, "Hello")],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 200, "Hello", "world."),
],
)
for _ in range(20):
assert await db.get_random_start_word(1) == "Hello"
assert await db.get_random_next_word(1, "Hello") == "world."
async def test_equal_weights_sample_evenly(self, db: Database) -> None:
"""Weighted random reads use one random threshold per query."""
await db.add_markov_data(
[
StartWord(1, 100, "Alpha"),
StartWord(1, 100, "Beta"),
StartWord(1, 100, "Gamma"),
],
[
Transition(1, 100, "seed", "alpha"),
Transition(1, 100, "seed", "beta"),
Transition(1, 100, "seed", "gamma"),
],
)
samples = 1200
start_counts: Counter[str] = Counter()
transition_counts: Counter[str] = Counter()
for _ in range(samples):
start_word = await db.get_random_start_word(1)
next_word = await db.get_random_next_word(1, "seed")
assert start_word is not None
assert next_word is not None
start_counts[start_word] += 1
transition_counts[next_word] += 1
expected = samples / 3
tolerance = samples * 0.075
assert set(start_counts) == {"Alpha", "Beta", "Gamma"}
assert set(transition_counts) == {"alpha", "beta", "gamma"}
for counts in (start_counts, transition_counts):
for count in counts.values():
assert expected - tolerance <= count <= expected + tolerance
class TestRemoveMarkovData:
"""Markov data removal via remove_markov_data."""
@@ -209,6 +650,44 @@ class TestRemoveMarkovData:
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") == "world."
async def test_duplicate_removals_decrement_occurrences(self, db: Database) -> None:
"""Duplicate removals in one call decrement counted rows."""
await db.add_markov_data(
[
StartWord(1, 100, "Hello"),
StartWord(1, 100, "Hello"),
StartWord(1, 100, "Hello"),
],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "world."),
],
)
await db.remove_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "Hello")],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "world."),
],
)
async with db._connection.execute(
"SELECT occurrences FROM markov_start_words"
" WHERE channel_id = 1 AND user_id = 100 AND word = 'Hello'",
) as cursor:
row = await cursor.fetchone()
assert row == (1,)
async with db._connection.execute(
"SELECT occurrences FROM markov_transitions"
" WHERE channel_id = 1"
" AND user_id = 100"
" AND word = 'Hello'"
" AND next_word = 'world.'",
) as cursor:
row = await cursor.fetchone()
assert row == (1,)
async def test_removes_last_start_word(self, db: Database) -> None:
"""Removing the only start word leaves the table empty for that channel."""
await db.add_markov_data([StartWord(1, 100, "Hello")], [])
@@ -514,9 +993,9 @@ class TestTransactionRollback:
async with db._transaction():
await db._connection.execute(
"INSERT INTO markov_start_words"
" (channel_id, user_id, word)"
" VALUES (?, ?, ?)",
(1, 100, "should_not_persist"),
" (channel_id, user_id, word, occurrences)"
" VALUES (?, ?, ?, ?)",
(1, 100, "should_not_persist", 1),
)
msg = "simulated failure"
raise RuntimeError(msg)
@@ -549,10 +1028,15 @@ class TestForgetUser:
async def test_preserves_other_users(self, db: Database) -> None:
"""Data belonging to other users is not affected."""
await db.add_markov_data(
[StartWord(1, 100, "Gone"), StartWord(1, 200, "Keep")],
[
StartWord(1, 100, "Gone"),
StartWord(1, 200, "Keep"),
StartWord(1, 200, "Keep"),
],
[
Transition(1, 100, "Gone", "away."),
Transition(1, 200, "Keep", "this."),
Transition(1, 200, "Keep", "this."),
],
)
await db.add_images(
@@ -567,6 +1051,12 @@ class TestForgetUser:
assert await db.get_random_start_word(1) == "Keep"
assert await db.get_random_next_word(1, "Keep") == "this."
async with db._connection.execute(
"SELECT occurrences FROM markov_start_words"
" WHERE channel_id = 1 AND user_id = 200 AND word = 'Keep'",
) as cursor:
row = await cursor.fetchone()
assert row == (2,)
assert await db.get_random_image(1) == "https://example.com/stay.png"
assert await db.is_flag_set(EntityType.USER, "200", Flag.ALLOW_PINGS) is True