Added unit tests for database, markov, flags, CLI, and version modules.
This commit is contained in:
@@ -13,3 +13,21 @@
|
||||
# limitations under the License.
|
||||
|
||||
"""Shared pytest fixtures for the test suite."""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from crabstero.database import Database
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def db(tmp_path: Path) -> AsyncIterator[Database]:
|
||||
"""Yield a Database backed by a temporary SQLite file."""
|
||||
database = await Database.connect(str(tmp_path / "test.db"))
|
||||
yield database
|
||||
await database.close()
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
# Copyright 2026 Logan Fick
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for CLI argument parsing and credential reading.
|
||||
|
||||
Tests cover _parse_args and _read_credential from crabstero.__main__.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from crabstero.__main__ import _parse_args, _read_credential
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class TestParseArgs:
|
||||
"""Argument parsing with defaults and overrides."""
|
||||
|
||||
def test_token_from_arg(self) -> None:
|
||||
"""--token flag sets the token."""
|
||||
args = _parse_args(["--token", "abc"])
|
||||
assert args.token == "abc" # noqa: S105
|
||||
|
||||
def test_database_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Omitting --database-path defaults to 'crabstero.db'."""
|
||||
monkeypatch.delenv("DATABASE_PATH", raising=False)
|
||||
args = _parse_args(["--token", "test"])
|
||||
assert args.database_path == "crabstero.db"
|
||||
|
||||
def test_database_override(self) -> None:
|
||||
"""--database-path overrides the default."""
|
||||
args = _parse_args(["--token", "test", "--database-path", "/custom.db"])
|
||||
assert args.database_path == "/custom.db"
|
||||
|
||||
def test_ingestion_workers_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Omitting --ingestion-workers defaults to 4."""
|
||||
monkeypatch.delenv("INGESTION_WORKERS", raising=False)
|
||||
args = _parse_args(["--token", "test"])
|
||||
assert args.ingestion_workers == 4
|
||||
|
||||
def test_ingestion_workers_override(self) -> None:
|
||||
"""--ingestion-workers overrides the default."""
|
||||
args = _parse_args(["--token", "test", "--ingestion-workers", "8"])
|
||||
assert args.ingestion_workers == 8
|
||||
|
||||
def test_ingest_only_flag(self) -> None:
|
||||
"""--ingest-only sets ingest_only to True."""
|
||||
args = _parse_args(["--token", "test", "--ingest-only"])
|
||||
assert args.ingest_only is True
|
||||
|
||||
def test_missing_token_exits(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Missing token causes SystemExit."""
|
||||
monkeypatch.delenv("TOKEN", raising=False)
|
||||
monkeypatch.delenv("CREDENTIALS_DIRECTORY", raising=False)
|
||||
with pytest.raises(SystemExit):
|
||||
_parse_args([])
|
||||
|
||||
|
||||
class TestReadCredential:
|
||||
"""Systemd credential file reading."""
|
||||
|
||||
def test_reads_credential_file(
|
||||
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Reads and strips the credential value from the file."""
|
||||
(tmp_path / "mytoken").write_text(" secret123 \n")
|
||||
monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path))
|
||||
assert _read_credential("mytoken") == "secret123"
|
||||
|
||||
def test_returns_none_without_env(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Returns None when CREDENTIALS_DIRECTORY is not set."""
|
||||
monkeypatch.delenv("CREDENTIALS_DIRECTORY", raising=False)
|
||||
assert _read_credential("anything") is None
|
||||
|
||||
def test_returns_none_for_missing_file(
|
||||
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Returns None when the credential file does not exist."""
|
||||
monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path))
|
||||
assert _read_credential("nonexistent") is None
|
||||
@@ -0,0 +1,245 @@
|
||||
# Copyright 2026 Logan Fick
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the Database class.
|
||||
|
||||
Tests cover Database.connect (pragmas, schema), markov start word and
|
||||
transition CRUD, image storage, flag CRUD, channel ingestion tracking,
|
||||
and the auto-commit write-batching mechanism.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from crabstero.database import AUTO_COMMIT_WRITE_THRESHOLD, Database
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class TestConnect:
|
||||
"""Database.connect creates a configured SQLite database."""
|
||||
|
||||
async def test_wal_mode(self, db: Database) -> None:
|
||||
"""Journal mode is set to WAL."""
|
||||
async with db._connection.execute("PRAGMA journal_mode") as cursor:
|
||||
row = await cursor.fetchone()
|
||||
assert row is not None
|
||||
assert row[0] == "wal"
|
||||
|
||||
async def test_synchronous_normal(self, db: Database) -> None:
|
||||
"""Synchronous mode is set to NORMAL."""
|
||||
async with db._connection.execute("PRAGMA synchronous") as cursor:
|
||||
row = await cursor.fetchone()
|
||||
assert row is not None
|
||||
assert row[0] == 1
|
||||
|
||||
async def test_schema_creates_tables(self, db: Database) -> None:
|
||||
"""All expected tables exist after connect."""
|
||||
async with db._connection.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type = 'table' ORDER BY name"
|
||||
) as cursor:
|
||||
tables = [row[0] for row in await cursor.fetchall()]
|
||||
assert tables == [
|
||||
"channel_images",
|
||||
"flags",
|
||||
"ingested_channels",
|
||||
"markov_start_words",
|
||||
"markov_transitions",
|
||||
]
|
||||
|
||||
|
||||
class TestMarkovStartWords:
|
||||
"""Start word storage and random retrieval."""
|
||||
|
||||
async def test_add_and_retrieve(self, db: Database) -> None:
|
||||
"""Inserted start word can be retrieved by channel."""
|
||||
await db.add_start_words_batch([(1, 100, "Hello")])
|
||||
await db.commit()
|
||||
result = await db.get_random_start_word(1)
|
||||
assert result == "Hello"
|
||||
|
||||
async def test_returns_none_when_empty(self, db: Database) -> None:
|
||||
"""Returns None for a channel with no start words."""
|
||||
result = await db.get_random_start_word(999)
|
||||
assert result is None
|
||||
|
||||
async def test_empty_batch_stores_nothing(self, db: Database) -> None:
|
||||
"""An empty batch does not store any start words."""
|
||||
await db.add_start_words_batch([])
|
||||
await db.commit()
|
||||
assert await db.get_random_start_word(1) is None
|
||||
|
||||
|
||||
class TestMarkovTransitions:
|
||||
"""Word transition storage and random retrieval."""
|
||||
|
||||
async def test_add_and_retrieve(self, db: Database) -> None:
|
||||
"""Inserted transition can be retrieved by channel and word."""
|
||||
await db.add_transitions_batch([(1, 100, "Hello", "world.")])
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == "world."
|
||||
|
||||
async def test_returns_none_when_empty(self, db: Database) -> None:
|
||||
"""Returns None when no transitions exist for the word."""
|
||||
result = await db.get_random_next_word(1, "nonexistent")
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"completing_word",
|
||||
[
|
||||
pytest.param("world.", id="period"),
|
||||
pytest.param("world!", id="exclamation"),
|
||||
pytest.param("world?", id="question"),
|
||||
pytest.param("world§", id="section-sign"),
|
||||
],
|
||||
)
|
||||
async def test_completing_word_filters_punctuation(
|
||||
self, db: Database, completing_word: str
|
||||
) -> None:
|
||||
"""get_random_completing_next_word only returns sentence-ending words."""
|
||||
await db.add_transitions_batch(
|
||||
[
|
||||
(1, 100, "Hello", "beautiful"),
|
||||
(1, 100, "Hello", completing_word),
|
||||
]
|
||||
)
|
||||
await db.commit()
|
||||
for _ in range(100):
|
||||
result = await db.get_random_completing_next_word(1, "Hello")
|
||||
assert result == completing_word
|
||||
|
||||
async def test_completing_returns_none_without_match(self, db: Database) -> None:
|
||||
"""Returns None when no transitions end with sentence punctuation."""
|
||||
await db.add_transitions_batch([(1, 100, "Hello", "beautiful")])
|
||||
await db.commit()
|
||||
result = await db.get_random_completing_next_word(1, "Hello")
|
||||
assert result is None
|
||||
|
||||
async def test_empty_batch_stores_nothing(self, db: Database) -> None:
|
||||
"""An empty batch does not store any transitions."""
|
||||
await db.add_transitions_batch([])
|
||||
await db.commit()
|
||||
assert await db.get_random_next_word(1, "Hello") is None
|
||||
|
||||
|
||||
class TestImages:
|
||||
"""Image URL storage and random retrieval."""
|
||||
|
||||
async def test_add_and_retrieve(self, db: Database) -> None:
|
||||
"""Inserted image URL can be retrieved by channel."""
|
||||
await db.add_image(1, 100, "https://example.com/cat.png")
|
||||
await db.commit()
|
||||
result = await db.get_random_image(1)
|
||||
assert result == "https://example.com/cat.png"
|
||||
|
||||
async def test_returns_none_when_empty(self, db: Database) -> None:
|
||||
"""Returns None for a channel with no images."""
|
||||
result = await db.get_random_image(999)
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestFlags:
|
||||
"""Flag CRUD operations on entities."""
|
||||
|
||||
async def test_set_and_check(self, db: Database) -> None:
|
||||
"""A set flag is reported as set."""
|
||||
await db.set_flag("channel", "123", "noReply")
|
||||
await db.commit()
|
||||
assert await db.is_flag_set("channel", "123", "noReply") is True
|
||||
|
||||
async def test_unset_flag_is_false(self, db: Database) -> None:
|
||||
"""An unset flag is reported as not set."""
|
||||
assert await db.is_flag_set("channel", "123", "noReply") is False
|
||||
|
||||
async def test_clear_flag(self, db: Database) -> None:
|
||||
"""A cleared flag is no longer reported as set."""
|
||||
await db.set_flag("channel", "123", "noReply")
|
||||
await db.commit()
|
||||
await db.clear_flag("channel", "123", "noReply")
|
||||
await db.commit()
|
||||
assert await db.is_flag_set("channel", "123", "noReply") is False
|
||||
|
||||
async def test_set_idempotent(self, db: Database) -> None:
|
||||
"""Setting the same flag twice does not raise."""
|
||||
await db.set_flag("channel", "123", "noReply")
|
||||
await db.set_flag("channel", "123", "noReply")
|
||||
await db.commit()
|
||||
assert await db.is_flag_set("channel", "123", "noReply") is True
|
||||
|
||||
|
||||
class TestChannelIngestion:
|
||||
"""Channel ingestion tracking."""
|
||||
|
||||
async def test_mark_and_check(self, db: Database) -> None:
|
||||
"""A marked channel is reported as ingested."""
|
||||
await db.mark_channel_ingested(42)
|
||||
await db.commit()
|
||||
assert await db.is_channel_ingested(42) is True
|
||||
|
||||
async def test_not_ingested_by_default(self, db: Database) -> None:
|
||||
"""Unmarked channels are not reported as ingested."""
|
||||
assert await db.is_channel_ingested(42) is False
|
||||
|
||||
async def test_mark_idempotent(self, db: Database) -> None:
|
||||
"""Marking the same channel twice does not raise."""
|
||||
await db.mark_channel_ingested(42)
|
||||
await db.mark_channel_ingested(42)
|
||||
await db.commit()
|
||||
assert await db.is_channel_ingested(42) is True
|
||||
|
||||
|
||||
class TestAutoCommit:
|
||||
"""Write batching and automatic commit behavior."""
|
||||
|
||||
async def test_commits_after_threshold(self, db: Database) -> None:
|
||||
"""Pending writes reset to zero after reaching the write threshold."""
|
||||
for i in range(AUTO_COMMIT_WRITE_THRESHOLD):
|
||||
await db.add_start_words_batch([(1, 100, f"word{i}")])
|
||||
assert db._pending_writes == 0
|
||||
assert await db.get_random_start_word(1) is not None
|
||||
|
||||
async def test_close_commits_pending(self, tmp_path: Path) -> None:
|
||||
"""close() commits any pending writes before closing."""
|
||||
db_path = str(tmp_path / "close_test.db")
|
||||
db = await Database.connect(db_path)
|
||||
await db.add_start_words_batch([(1, 100, "hello")])
|
||||
assert db._pending_writes > 0
|
||||
await db.close()
|
||||
|
||||
# Reopen and verify data persisted.
|
||||
db2 = await Database.connect(db_path)
|
||||
result = await db2.get_random_start_word(1)
|
||||
await db2.close()
|
||||
assert result == "hello"
|
||||
|
||||
async def test_close_without_pending_writes(self, tmp_path: Path) -> None:
|
||||
"""close() succeeds when there are no pending writes."""
|
||||
db = await Database.connect(str(tmp_path / "clean_close.db"))
|
||||
await db.close()
|
||||
|
||||
async def test_flush_timer_commits(
|
||||
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Background flush timer commits pending writes after timeout."""
|
||||
monkeypatch.setattr("crabstero.database.AUTO_COMMIT_TIMEOUT_SECONDS", 0.05)
|
||||
db = await Database.connect(str(tmp_path / "timer_test.db"))
|
||||
await db.add_start_words_batch([(1, 100, "hello")])
|
||||
assert db._pending_writes > 0
|
||||
await asyncio.sleep(0.1)
|
||||
assert db._pending_writes == 0
|
||||
await db.close()
|
||||
@@ -0,0 +1,119 @@
|
||||
# Copyright 2026 Logan Fick
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the flag convenience wrappers and enums.
|
||||
|
||||
Tests cover _entity_id, Flag and EntityType enums, and the high-level
|
||||
set_flag/clear_flag/is_flag_set wrappers from crabstero.flags.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from crabstero.flags import (
|
||||
EntityType,
|
||||
Flag,
|
||||
_entity_id,
|
||||
clear_flag,
|
||||
is_flag_set,
|
||||
set_flag,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from crabstero.database import Database
|
||||
|
||||
|
||||
class _StubSnowflake:
|
||||
"""Minimal stand-in for discord.abc.Snowflake."""
|
||||
|
||||
def __init__(self, *, entity_id: int = 123456789) -> None:
|
||||
self.id = entity_id
|
||||
|
||||
|
||||
class TestEntityId:
|
||||
"""Entity ID extraction from Discord objects and integers."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("entity", "expected"),
|
||||
[
|
||||
pytest.param(42, "42", id="raw-integer"),
|
||||
pytest.param(
|
||||
_StubSnowflake(entity_id=99),
|
||||
"99",
|
||||
id="snowflake-object",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_extraction(self, entity: object, expected: str) -> None:
|
||||
"""Converts the entity to the expected string ID."""
|
||||
assert _entity_id(entity) == expected # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestFlagEnums:
|
||||
"""Flag and EntityType enum values."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("member", "expected"),
|
||||
[
|
||||
pytest.param(Flag.NO_REPLY, "noReply", id="no-reply"),
|
||||
pytest.param(Flag.NO_INGEST, "noIngest", id="no-ingest"),
|
||||
pytest.param(Flag.ALLOW_PINGS, "allowPings", id="allow-pings"),
|
||||
],
|
||||
)
|
||||
def test_flag_values(self, member: Flag, expected: str) -> None:
|
||||
"""Flag enum value matches the expected database string."""
|
||||
assert member.value == expected
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("member", "expected"),
|
||||
[
|
||||
pytest.param(EntityType.CHANNEL, "channel", id="channel"),
|
||||
pytest.param(EntityType.SERVER, "server", id="server"),
|
||||
pytest.param(EntityType.USER, "user", id="user"),
|
||||
],
|
||||
)
|
||||
def test_entity_type_values(self, member: EntityType, expected: str) -> None:
|
||||
"""EntityType enum value matches the expected database string."""
|
||||
assert member.value == expected
|
||||
|
||||
|
||||
class TestSetClearCheck:
|
||||
"""High-level flag set/clear/check cycle through the flags module."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"flag",
|
||||
[
|
||||
pytest.param(Flag.NO_REPLY, id="no-reply"),
|
||||
pytest.param(Flag.NO_INGEST, id="no-ingest"),
|
||||
pytest.param(Flag.ALLOW_PINGS, id="allow-pings"),
|
||||
],
|
||||
)
|
||||
async def test_set_then_check(self, db: Database, flag: Flag) -> None:
|
||||
"""A set flag is reported as set."""
|
||||
await set_flag(db, 1, EntityType.CHANNEL, flag)
|
||||
await db.commit()
|
||||
assert await is_flag_set(db, 1, EntityType.CHANNEL, flag) is True
|
||||
|
||||
async def test_unset_returns_false(self, db: Database) -> None:
|
||||
"""An unset flag is reported as not set."""
|
||||
assert await is_flag_set(db, 1, EntityType.CHANNEL, Flag.NO_REPLY) is False
|
||||
|
||||
async def test_clear_removes_flag(self, db: Database) -> None:
|
||||
"""A cleared flag is no longer reported as set."""
|
||||
await set_flag(db, 1, EntityType.CHANNEL, Flag.NO_REPLY)
|
||||
await db.commit()
|
||||
await clear_flag(db, 1, EntityType.CHANNEL, Flag.NO_REPLY)
|
||||
await db.commit()
|
||||
assert await is_flag_set(db, 1, EntityType.CHANNEL, Flag.NO_REPLY) is False
|
||||
@@ -0,0 +1,283 @@
|
||||
# Copyright 2026 Logan Fick
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for Markov chain ingestion and generation.
|
||||
|
||||
Tests cover is_complete_sentence, _ingest_sentence, ingest, and generate
|
||||
from crabstero.markov.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from crabstero.markov import (
|
||||
DEFAULT_SENTENCE_END,
|
||||
_ingest_sentence,
|
||||
generate,
|
||||
ingest,
|
||||
is_complete_sentence,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from crabstero.database import Database
|
||||
|
||||
|
||||
class TestIsCompleteSentence:
|
||||
"""Sentence completeness detection."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("sentence", "expected"),
|
||||
[
|
||||
pytest.param("Hello.", True, id="period"),
|
||||
pytest.param("Wow!", True, id="exclamation"),
|
||||
pytest.param("Really?", True, id="question"),
|
||||
pytest.param(f"Hello{DEFAULT_SENTENCE_END}", True, id="section-sign"),
|
||||
pytest.param("Hello", False, id="no-punctuation"),
|
||||
pytest.param("", False, id="empty-string"),
|
||||
pytest.param("Hello. ", False, id="trailing-space"),
|
||||
pytest.param("Hello,", False, id="comma"),
|
||||
],
|
||||
)
|
||||
def test_detection(self, sentence: str, expected: bool) -> None:
|
||||
"""Correctly identifies sentence completeness."""
|
||||
assert is_complete_sentence(sentence) is expected
|
||||
|
||||
|
||||
class TestIngestSentence:
|
||||
"""Single sentence ingestion into the Markov chain."""
|
||||
|
||||
async def test_stores_start_word(self, db: Database) -> None:
|
||||
"""First word of the sentence is stored as a start word."""
|
||||
await _ingest_sentence(db, 1, 100, "Hello world.")
|
||||
await db.commit()
|
||||
result = await db.get_random_start_word(1)
|
||||
assert result == "Hello"
|
||||
|
||||
async def test_stores_transitions(self, db: Database) -> None:
|
||||
"""Adjacent words create transitions."""
|
||||
await _ingest_sentence(db, 1, 100, "Hello world.")
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == "world."
|
||||
|
||||
async def test_appends_sentence_end_if_missing(self, db: Database) -> None:
|
||||
"""Unpunctuated sentence gets the default sentence-end marker."""
|
||||
await _ingest_sentence(db, 1, 100, "Hello world")
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == f"world{DEFAULT_SENTENCE_END}"
|
||||
|
||||
async def test_preserves_existing_punctuation(self, db: Database) -> None:
|
||||
"""Already-punctuated sentence keeps its terminator."""
|
||||
await _ingest_sentence(db, 1, 100, "Hello world!")
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == "world!"
|
||||
|
||||
async def test_single_word_stores_nothing(self, db: Database) -> None:
|
||||
"""A single-word sentence produces no start words or transitions."""
|
||||
await _ingest_sentence(db, 1, 100, "Hello.")
|
||||
await db.commit()
|
||||
assert await db.get_random_start_word(1) is None
|
||||
assert await db.get_random_next_word(1, "Hello.") is None
|
||||
|
||||
async def test_stores_all_transitions(self, db: Database) -> None:
|
||||
"""All adjacent word pairs create transitions."""
|
||||
await _ingest_sentence(db, 1, 100, "A B C.")
|
||||
await db.commit()
|
||||
assert await db.get_random_next_word(1, "A") == "B"
|
||||
assert await db.get_random_next_word(1, "B") == "C."
|
||||
|
||||
|
||||
class TestIngest:
|
||||
"""Paragraph ingestion splits into sentences."""
|
||||
|
||||
async def test_single_sentence(self, db: Database) -> None:
|
||||
"""A single sentence paragraph is ingested."""
|
||||
await ingest(db, 1, 100, "Hello world.")
|
||||
await db.commit()
|
||||
assert await db.get_random_start_word(1) == "Hello"
|
||||
|
||||
async def test_multiple_sentences(self, db: Database) -> None:
|
||||
"""Multiple sentences are split and ingested individually."""
|
||||
await ingest(db, 1, 100, "Hello world. Goodbye world!")
|
||||
await db.commit()
|
||||
# Both "Hello" and "Goodbye" should appear as start words.
|
||||
async with db._connection.execute(
|
||||
"SELECT DISTINCT word FROM markov_start_words WHERE channel_id = ?",
|
||||
(1,),
|
||||
) as cursor:
|
||||
start_words = {row[0] for row in await cursor.fetchall()}
|
||||
assert start_words == {"Hello", "Goodbye"}
|
||||
|
||||
async def test_normalizes_whitespace(self, db: Database) -> None:
|
||||
"""Extra spaces and newlines are collapsed."""
|
||||
await ingest(db, 1, 100, "Hello world.\nGoodbye world!")
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == "world."
|
||||
|
||||
async def test_appends_default_end(self, db: Database) -> None:
|
||||
"""Unpunctuated paragraph gets the default sentence-end marker."""
|
||||
await ingest(db, 1, 100, "Hello world")
|
||||
await db.commit()
|
||||
result = await db.get_random_next_word(1, "Hello")
|
||||
assert result == f"world{DEFAULT_SENTENCE_END}"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"paragraph",
|
||||
[
|
||||
pytest.param("", id="empty-string"),
|
||||
pytest.param(" ", id="whitespace-only"),
|
||||
],
|
||||
)
|
||||
async def test_empty_input_stores_nothing(
|
||||
self, db: Database, paragraph: str
|
||||
) -> None:
|
||||
"""Empty or whitespace-only input does not store any data."""
|
||||
await ingest(db, 1, 100, paragraph)
|
||||
await db.commit()
|
||||
assert await db.get_random_start_word(1) is None
|
||||
|
||||
async def test_splits_on_punctuation_followed_by_space(self, db: Database) -> None:
|
||||
"""Punctuation followed by a space splits into separate sentences."""
|
||||
await ingest(db, 1, 100, "Dr. Smith likes cats")
|
||||
await db.commit()
|
||||
# "Dr." splits off as a single-word sentence (stores nothing).
|
||||
# "Smith likes cats" becomes a sentence with "Smith" as start word.
|
||||
assert await db.get_random_start_word(1) == "Smith"
|
||||
|
||||
async def test_no_split_without_space_after_punctuation(self, db: Database) -> None:
|
||||
"""Punctuation not followed by a space keeps words together."""
|
||||
await ingest(db, 1, 100, "Hello.World is here")
|
||||
await db.commit()
|
||||
assert await db.get_random_start_word(1) == "Hello.World"
|
||||
|
||||
|
||||
class TestGenerate:
|
||||
"""Markov chain text generation."""
|
||||
|
||||
async def test_fallback_on_empty_channel(self, db: Database) -> None:
|
||||
"""Returns 'Hello world!' when the channel has no data."""
|
||||
result = await generate(db, 1)
|
||||
assert result == "Hello world!"
|
||||
|
||||
async def test_generates_from_ingested_data(self, db: Database) -> None:
|
||||
"""Generated text uses words from ingested data."""
|
||||
await ingest(db, 1, 100, "The quick brown fox.")
|
||||
await db.commit()
|
||||
result = await generate(db, 1)
|
||||
assert result == "The quick brown fox."
|
||||
|
||||
async def test_strips_section_sign(self, db: Database) -> None:
|
||||
"""The internal section sign marker never appears in output."""
|
||||
await ingest(db, 1, 100, "Hello world")
|
||||
await db.commit()
|
||||
result = await generate(db, 1)
|
||||
assert result == "Hello world"
|
||||
|
||||
async def test_respects_hard_limit(self, db: Database) -> None:
|
||||
"""Output is truncated at hard_limit."""
|
||||
words = [f"w{i}" for i in range(100)]
|
||||
text = " ".join(words)
|
||||
await ingest(db, 1, 100, text)
|
||||
await db.commit()
|
||||
result = await generate(db, 1, soft_limit=10, hard_limit=49)
|
||||
assert result == "w0 w1 w2 w3 w4 w5 w6 w7 w8 w9 w10 w11 w12 w13 w14"
|
||||
|
||||
async def test_prefers_completing_word_after_soft_limit(self, db: Database) -> None:
|
||||
"""After soft_limit, generation prefers completing words."""
|
||||
# Chain: A → B → C → D → {E, "end."}
|
||||
# A, B, C have only non-completing transitions, so past the soft limit
|
||||
# the loop falls back to get_random_next_word for each. D has both a
|
||||
# continuing ("E") and completing ("end.") transition, so
|
||||
# get_random_completing_next_word deterministically picks "end.".
|
||||
await db.add_start_words_batch([(1, 100, "A")])
|
||||
await db.add_transitions_batch(
|
||||
[
|
||||
(1, 100, "A", "B"),
|
||||
(1, 100, "B", "C"),
|
||||
(1, 100, "C", "D"),
|
||||
(1, 100, "D", "E"),
|
||||
(1, 100, "D", "F"),
|
||||
(1, 100, "D", "G"),
|
||||
(1, 100, "D", "H"),
|
||||
(1, 100, "D", "end."),
|
||||
]
|
||||
)
|
||||
await db.commit()
|
||||
for _ in range(100):
|
||||
result = await generate(db, 1, soft_limit=1, hard_limit=1000)
|
||||
assert result == "A B C D end."
|
||||
|
||||
async def test_start_word_already_ends_sentence(self, db: Database) -> None:
|
||||
"""Generation stops immediately when the start word is sentence-ending."""
|
||||
await db.add_start_words_batch([(1, 100, "Yes.")])
|
||||
await db.commit()
|
||||
result = await generate(db, 1)
|
||||
assert result == "Yes."
|
||||
|
||||
async def test_chain_dead_end(self, db: Database) -> None:
|
||||
"""Generation stops when no next word exists (dead-end chain)."""
|
||||
await db.add_start_words_batch([(1, 100, "Hello")])
|
||||
await db.commit()
|
||||
# "Hello" has no transitions, so the loop breaks immediately.
|
||||
result = await generate(db, 1)
|
||||
assert result == "Hello"
|
||||
|
||||
async def test_hard_limit_strips_section_sign(self, db: Database) -> None:
|
||||
"""Section sign at the truncation boundary is stripped."""
|
||||
# Build a chain: "A" -> "B" -> "C§".
|
||||
await db.add_start_words_batch([(1, 100, "A")])
|
||||
await db.add_transitions_batch(
|
||||
[
|
||||
(1, 100, "A", "B"),
|
||||
(1, 100, "B", f"C{DEFAULT_SENTENCE_END}"),
|
||||
]
|
||||
)
|
||||
await db.commit()
|
||||
# hard_limit=5 truncates "A B C§" (length 6) to "A B C".
|
||||
result = await generate(db, 1, soft_limit=100, hard_limit=5)
|
||||
assert result == "A B C"
|
||||
|
||||
async def test_soft_limit_falls_back_to_regular_next_word(
|
||||
self, db: Database
|
||||
) -> None:
|
||||
"""After soft_limit, falls back when no completing word exists."""
|
||||
# "A" -> "B" (no completing transition). Past soft_limit,
|
||||
# get_random_completing_next_word returns None, falls back to "B".
|
||||
await db.add_start_words_batch([(1, 100, "A")])
|
||||
await db.add_transitions_batch([(1, 100, "A", "B")])
|
||||
await db.commit()
|
||||
result = await generate(db, 1, soft_limit=1, hard_limit=1000)
|
||||
assert result == "A B"
|
||||
|
||||
async def test_hard_limit_truncates_mid_word(self, db: Database) -> None:
|
||||
"""Hard limit slices output even when it falls inside a word."""
|
||||
await db.add_start_words_batch([(1, 100, "AB")])
|
||||
await db.add_transitions_batch([(1, 100, "AB", "CDEF")])
|
||||
await db.commit()
|
||||
# "AB CDEF" is 7 chars; hard_limit=5 truncates to "AB CD".
|
||||
result = await generate(db, 1, soft_limit=100, hard_limit=5)
|
||||
assert result == "AB CD"
|
||||
|
||||
async def test_channel_isolation(self, db: Database) -> None:
|
||||
"""Data ingested into one channel does not leak into another."""
|
||||
await ingest(db, 1, 100, "Channel one data.")
|
||||
await ingest(db, 2, 100, "Channel two data.")
|
||||
await db.commit()
|
||||
# Channel 3 has no data; should get the fallback.
|
||||
result = await generate(db, 3)
|
||||
assert result == "Hello world!"
|
||||
@@ -12,12 +12,15 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for version availability."""
|
||||
"""Unit tests for package version availability."""
|
||||
|
||||
import crabstero
|
||||
|
||||
|
||||
def test_version_is_available() -> None:
|
||||
"""Verify __version__ is a non-empty string."""
|
||||
assert isinstance(crabstero.__version__, str)
|
||||
assert crabstero.__version__
|
||||
class TestVersion:
|
||||
"""Package version metadata."""
|
||||
|
||||
def test_version_is_available(self) -> None:
|
||||
"""__version__ is a non-empty string."""
|
||||
assert isinstance(crabstero.__version__, str)
|
||||
assert crabstero.__version__
|
||||
|
||||
Reference in New Issue
Block a user