Added unit tests for database, markov, flags, CLI, and version modules.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 9s
CI / Type Checking (push) Failing after 10s
CI / Spelling (push) Successful in 5s

This commit is contained in:
2026-03-17 14:13:31 -04:00
parent 34ae3b3e64
commit 2d13dfb063
8 changed files with 770 additions and 7 deletions
+18
View File
@@ -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()
+94
View File
@@ -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
+245
View File
@@ -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()
+119
View File
@@ -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
+283
View File
@@ -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!"
+8 -5
View File
@@ -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__