349 lines
14 KiB
Python
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."
|