Reworked Markov storage to reduce database size.
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user