Files
Crabstero/tests/unit/test_markov.py
T
LogalDeveloper 49062159f9
CI / Formatting (push) Failing after 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 4s
Moved project into a src-based layout and reorganized tests.
2026-06-16 14:21:02 -04:00

349 lines
14 KiB
Python

# 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.database import StartWord, Transition
from crabstero.markov import (
DEFAULT_SENTENCE_END,
_ingest_sentence,
_uningest_sentence,
generate,
ingest,
is_complete_sentence,
uningest,
)
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: # noqa: FBT001
"""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.")
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.")
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")
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!")
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.")
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.")
assert await db.get_random_next_word(1, "A") == "B"
assert await db.get_random_next_word(1, "B") == "C."
class TestIngestParagraph:
"""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.")
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!")
# 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!")
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")
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)
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")
# "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")
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 an informational message when the channel has no data."""
result = await generate(db, 1)
expected = (
"I do not have enough data to generate a message yet."
" Chat a bit more so I can learn how this channel talks."
)
assert result == expected
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.")
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")
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)
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_markov_data(
[StartWord(1, 100, "A")],
[
Transition(1, 100, "A", "B"),
Transition(1, 100, "B", "C"),
Transition(1, 100, "C", "D"),
Transition(1, 100, "D", "E"),
Transition(1, 100, "D", "F"),
Transition(1, 100, "D", "G"),
Transition(1, 100, "D", "H"),
Transition(1, 100, "D", "end."),
],
)
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_markov_data([StartWord(1, 100, "Yes.")], [])
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_markov_data([StartWord(1, 100, "Hello")], [])
# "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_markov_data(
[StartWord(1, 100, "A")],
[
Transition(1, 100, "A", "B"),
Transition(1, 100, "B", f"C{DEFAULT_SENTENCE_END}"),
],
)
# 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_markov_data(
[StartWord(1, 100, "A")],
[Transition(1, 100, "A", "B")],
)
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_markov_data(
[StartWord(1, 100, "AB")],
[Transition(1, 100, "AB", "CDEF")],
)
# "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.")
# Channel 3 has no data; should get the fallback.
result = await generate(db, 3)
expected = (
"I do not have enough data to generate a message yet."
" Chat a bit more so I can learn how this channel talks."
)
assert result == expected
class TestUningestSentence:
"""Single sentence uningest from the Markov chain."""
async def test_removes_start_word_and_transition(self, db: Database) -> None:
"""Uningest removes the start word and transition added by ingest."""
await _ingest_sentence(db, 1, 100, "Hello world.")
await _uningest_sentence(db, 1, 100, "Hello world.")
assert await db.get_random_start_word(1) is None
assert await db.get_random_next_word(1, "Hello") is None
async def test_preserves_other_data(self, db: Database) -> None:
"""Uningest only removes data for the specified sentence."""
await _ingest_sentence(db, 1, 100, "Hello world.")
await _ingest_sentence(db, 1, 100, "Goodbye world.")
await _uningest_sentence(db, 1, 100, "Hello world.")
assert await db.get_random_start_word(1) == "Goodbye"
async def test_handles_missing_punctuation(self, db: Database) -> None:
"""Uningest appends the default sentence end, matching ingest behavior."""
await _ingest_sentence(db, 1, 100, "Hello world")
await _uningest_sentence(db, 1, 100, "Hello world")
assert await db.get_random_start_word(1) is None
async def test_single_word_is_noop(self, db: Database) -> None:
"""Uningesting a single-word sentence does not error."""
await _ingest_sentence(db, 1, 100, "Hello.")
await _uningest_sentence(db, 1, 100, "Hello.")
class TestUningestParagraph:
"""Paragraph-level uningest."""
async def test_uningest_multiple_sentences(self, db: Database) -> None:
"""Uningest reverses a multi-sentence paragraph."""
await ingest(db, 1, 100, "Hello world. Goodbye world!")
await uningest(db, 1, 100, "Hello world. Goodbye world!")
assert await db.get_random_start_word(1) is None
async def test_uningest_preserves_duplicate_data(self, db: Database) -> None:
"""Uningesting one copy leaves the other intact."""
await ingest(db, 1, 100, "Hello world.")
await ingest(db, 1, 100, "Hello world.")
await uningest(db, 1, 100, "Hello world.")
assert await db.get_random_start_word(1) == "Hello"
assert await db.get_random_next_word(1, "Hello") == "world."