Expanded integration coverage and enforced test categories.
Audit / Dependencies (push) Successful in 9s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 1m7s
CI / Type Checking (push) Successful in 12s
CI / Spelling (push) Successful in 8s

This commit is contained in:
2026-06-18 16:10:40 -04:00
parent 49062159f9
commit 5bbb0bbde5
22 changed files with 2433 additions and 550 deletions
+15
View File
@@ -0,0 +1,15 @@
# 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.
"""CLI integration tests for Crabstero."""
@@ -0,0 +1,89 @@
# 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.
"""Integration tests for systemd notification socket behavior."""
import logging
import os
import socket
import sys
from typing import TYPE_CHECKING
import pytest
from crabstero.cli import _sd_notify
if TYPE_CHECKING:
from pathlib import Path
pytestmark = pytest.mark.skipif(
os.name != "posix" or not hasattr(socket, "AF_UNIX"),
reason="systemd notification sockets require Unix-domain socket support",
)
class TestSystemdNotifySocket:
"""Systemd notification helper behavior against real Unix-domain sockets."""
def test_sd_notify_sends_datagram(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""_sd_notify sends the payload to NOTIFY_SOCKET."""
socket_path = tmp_path / "notify.sock"
with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as server:
server.bind(str(socket_path))
server.settimeout(1)
monkeypatch.setenv("NOTIFY_SOCKET", str(socket_path))
_sd_notify("READY=1")
assert server.recv(1024) == b"READY=1"
def test_sd_notify_ignores_socket_errors(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
"""_sd_notify logs and suppresses notification socket send failures."""
monkeypatch.setenv("NOTIFY_SOCKET", str(tmp_path / "missing.sock"))
with caplog.at_level(logging.DEBUG, logger="crabstero"):
_sd_notify("READY=1")
assert "Could not send systemd notification" in caplog.text
@pytest.mark.skipif(
sys.platform != "linux",
reason="abstract Unix sockets are Linux-specific",
)
def test_sd_notify_sends_to_abstract_socket(
self,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""@-prefixed NOTIFY_SOCKET values target Linux abstract sockets."""
socket_name = f"crabstero-notify-{os.getpid()}"
with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as server:
server.bind(f"\0{socket_name}")
server.settimeout(1)
monkeypatch.setenv("NOTIFY_SOCKET", f"@{socket_name}")
_sd_notify("READY=1")
assert server.recv(1024) == b"READY=1"
+15
View File
@@ -0,0 +1,15 @@
# 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.
"""Database integration tests for Crabstero."""
+35
View File
@@ -0,0 +1,35 @@
# 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.
"""Fixtures used only by database-bound integration tests."""
from typing import TYPE_CHECKING
import pytest
from crabstero.database import Database
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from pathlib import Path
@pytest.fixture
async def db(tmp_path: Path) -> AsyncGenerator[Database]:
"""Yield an isolated file-backed database for SQLite integration tests."""
database = await Database.connect(str(tmp_path / "crabstero.db"))
try:
yield database
finally:
await database.close()
+645
View File
@@ -0,0 +1,645 @@
# 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.
"""Integration tests for the Database class.
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 pytest
from crabstero.database import ChannelImage, Database, StartWord, Transition
from crabstero.flags import EntityType, Flag
if TYPE_CHECKING:
from pathlib import Path
class TestConnect:
"""Database.connect creates a configured SQLite database."""
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 TestAddMarkovData:
"""Markov start word and transition storage via add_markov_data."""
async def test_stores_start_word(self, db: Database) -> None:
"""Inserted start word can be retrieved by channel."""
await db.add_markov_data([StartWord(1, 100, "Hello")], [])
result = await db.get_random_start_word(1)
assert result == "Hello"
async def test_stores_transition(self, db: Database) -> None:
"""Inserted transition can be retrieved by channel and word."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
result = await db.get_random_next_word(1, "Hello")
assert result == "world."
async def test_stores_both_in_single_call(self, db: Database) -> None:
"""Start words and transitions are stored in a single call."""
await db.add_markov_data(
[StartWord(1, 100, "Hello")],
[Transition(1, 100, "Hello", "world.")],
)
assert await db.get_random_start_word(1) == "Hello"
assert await db.get_random_next_word(1, "Hello") == "world."
async def test_empty_lists_is_noop(self, db: Database) -> None:
"""Empty lists do not store any data."""
await db.add_markov_data([], [])
assert await db.get_random_start_word(1) is None
class TestMarkovReadMethods:
"""Markov read methods return None for out-of-scope or missing data."""
async def test_returns_none_when_empty(self, db: Database) -> None:
"""Returns None when no data has been stored."""
assert await db.get_random_start_word(1) is None
assert await db.get_random_next_word(1, "nonexistent") is None
assert await db.get_random_completing_next_word(1, "nonexistent") is None
async def test_start_word_scoped_to_channel(self, db: Database) -> None:
"""A start word in one channel is not returned for another channel."""
await db.add_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(2) is None
async def test_next_word_scoped_to_channel(self, db: Database) -> None:
"""A transition in one channel is not returned for another channel."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(2, "Hello") is None
async def test_completing_next_word_scoped_to_channel(self, db: Database) -> None:
"""Completing transition is not returned for another channel."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_completing_next_word(2, "Hello") is None
async def test_next_word_scoped_to_word(self, db: Database) -> None:
"""A transition for one word is not returned when querying a different word."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Goodbye") is None
async def test_completing_next_word_scoped_to_word(self, db: Database) -> None:
"""Completing transition is not returned for a different word."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_completing_next_word(1, "Goodbye") 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\u00a7", 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_markov_data(
[],
[
Transition(1, 100, "Hello", "beautiful"),
Transition(1, 100, "Hello", completing_word),
],
)
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_markov_data([], [Transition(1, 100, "Hello", "beautiful")])
result = await db.get_random_completing_next_word(1, "Hello")
assert result is None
async def test_start_word_pooled_across_users(self, db: Database) -> None:
"""Start words from different users are visible in the same channel query."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 200, "Goodbye")],
[],
)
assert await db.get_random_start_word(1) in {"Hello", "Goodbye"}
async def test_next_word_pooled_across_users(self, db: Database) -> None:
"""Transitions from different users are visible in the same channel query."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 200, "Hello", "friend."),
],
)
assert await db.get_random_next_word(1, "Hello") in {"world.", "friend."}
async def test_completing_next_word_pooled_across_users(self, db: Database) -> None:
"""Completing transitions from different users are visible."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 200, "Hello", "friend."),
],
)
assert await db.get_random_completing_next_word(1, "Hello") in {
"world.",
"friend.",
}
class TestRemoveMarkovData:
"""Markov data removal via remove_markov_data."""
async def test_removes_one_start_word(self, db: Database) -> None:
"""Removes exactly one matching start word row."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "Hello")],
[],
)
await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
# One copy should remain.
assert await db.get_random_start_word(1) == "Hello"
async def test_removes_one_transition(self, db: Database) -> None:
"""Removes exactly one matching transition row."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "world."),
],
)
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") == "world."
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")], [])
await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) is None
async def test_removes_last_transition(self, db: Database) -> None:
"""Removing the only transition leaves no next word."""
await db.add_markov_data([], [Transition(1, 100, "Hello", "world.")])
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") is None
async def test_no_match_is_noop(self, db: Database) -> None:
"""Removing a non-existent row does not raise."""
await db.remove_markov_data(
[StartWord(1, 100, "nope")],
[Transition(1, 100, "nope", "nah")],
)
async def test_start_word_removal_scoped_to_channel(self, db: Database) -> None:
"""Removing a start word in one channel leaves another channel intact."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(2, 200, "Hello")],
[],
)
await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) is None
assert await db.get_random_start_word(2) == "Hello"
async def test_transition_removal_scoped_to_channel(self, db: Database) -> None:
"""Removing a transition in one channel leaves another channel intact."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(2, 200, "Hello", "world."),
],
)
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") is None
assert await db.get_random_next_word(2, "Hello") == "world."
async def test_start_word_removal_scoped_to_user(self, db: Database) -> None:
"""Removing a start word for one user leaves another user."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 200, "Hello")],
[],
)
await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) == "Hello"
async def test_transition_removal_scoped_to_user(self, db: Database) -> None:
"""Removing a transition for one user leaves another user."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 200, "Hello", "world."),
],
)
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") == "world."
async def test_start_word_removal_scoped_to_word(self, db: Database) -> None:
"""Removing one start word leaves a different start word."""
await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "World")],
[],
)
await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) == "World"
async def test_transition_removal_scoped_to_word(self, db: Database) -> None:
"""Removing one word's transition leaves another word's."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Goodbye", "world."),
],
)
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") is None
assert await db.get_random_next_word(1, "Goodbye") == "world."
async def test_transition_removal_scoped_to_next_word(self, db: Database) -> None:
"""Removing one next_word leaves a different next_word."""
await db.add_markov_data(
[],
[
Transition(1, 100, "Hello", "world."),
Transition(1, 100, "Hello", "friend."),
],
)
await db.remove_markov_data([], [Transition(1, 100, "Hello", "world.")])
assert await db.get_random_next_word(1, "Hello") == "friend."
async def test_empty_lists_is_noop(self, db: Database) -> None:
"""Empty lists do not error."""
await db.remove_markov_data([], [])
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_images([ChannelImage(1, 100, "https://example.com/cat.png")])
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
async def test_image_scoped_to_channel(self, db: Database) -> None:
"""An image in one channel is not returned for another channel."""
await db.add_images([ChannelImage(1, 100, "https://example.com/cat.png")])
assert await db.get_random_image(2) is None
async def test_empty_list_is_noop(self, db: Database) -> None:
"""Empty list does not store any data."""
await db.add_images([])
assert await db.get_random_image(1) is None
async def test_image_pooled_across_users(self, db: Database) -> None:
"""Images from different users are visible in the same channel query."""
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 200, "https://example.com/b.png"),
],
)
assert await db.get_random_image(1) in {
"https://example.com/a.png",
"https://example.com/b.png",
}
class TestRemoveImages:
"""Image removal via remove_images."""
async def test_removes_one_image(self, db: Database) -> None:
"""Removes exactly one matching image row."""
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 100, "https://example.com/a.png"),
],
)
await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
# One copy should remain.
assert await db.get_random_image(1) == "https://example.com/a.png"
async def test_removes_last_image(self, db: Database) -> None:
"""Removing the only image leaves none for that channel."""
await db.add_images([ChannelImage(1, 100, "https://example.com/a.png")])
await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
assert await db.get_random_image(1) is None
async def test_image_removal_scoped_to_channel(self, db: Database) -> None:
"""Removing an image in one channel leaves another channel intact."""
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(2, 200, "https://example.com/a.png"),
],
)
await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
assert await db.get_random_image(1) is None
assert await db.get_random_image(2) == "https://example.com/a.png"
async def test_image_removal_scoped_to_user(self, db: Database) -> None:
"""Removing an image for one user leaves another user."""
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 200, "https://example.com/a.png"),
],
)
await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
assert await db.get_random_image(1) == "https://example.com/a.png"
async def test_image_removal_scoped_to_url(self, db: Database) -> None:
"""Removing one URL leaves a different URL for the same user."""
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 100, "https://example.com/b.png"),
],
)
await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
assert await db.get_random_image(1) == "https://example.com/b.png"
async def test_no_match_is_noop(self, db: Database) -> None:
"""Removing a non-existent image does not raise."""
await db.remove_images([ChannelImage(1, 100, "https://example.com/nope.png")])
async def test_empty_list_is_noop(self, db: Database) -> None:
"""Empty list does not error."""
await db.remove_images([])
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(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_REPLY) 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(EntityType.CHANNEL, "123", Flag.NO_REPLY) is False
async def test_clear_flag(self, db: Database) -> None:
"""A cleared flag is no longer reported as set."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.clear_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_REPLY) is False
async def test_set_idempotent(self, db: Database) -> None:
"""Setting the same flag twice does not raise."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_REPLY) is True
async def test_scoped_to_entity_id(self, db: Database) -> None:
"""A flag set on one entity is not visible on another entity."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "456", Flag.NO_REPLY) is False
async def test_scoped_to_entity_type(self, db: Database) -> None:
"""A flag set on one entity type is not visible on another."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.USER, "123", Flag.NO_REPLY) is False
async def test_scoped_to_flag_name(self, db: Database) -> None:
"""A flag set under one name is not visible under a different name."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_INGEST) is False
async def test_clear_scoped_to_flag_name(self, db: Database) -> None:
"""Clearing one flag leaves other flags on the same entity intact."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_INGEST)
await db.clear_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_INGEST) is True
async def test_clear_scoped_to_entity_id(self, db: Database) -> None:
"""Clearing a flag on one entity leaves another entity."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.set_flag(EntityType.CHANNEL, "456", Flag.NO_REPLY)
await db.clear_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.CHANNEL, "456", Flag.NO_REPLY) is True
async def test_clear_scoped_to_entity_type(self, db: Database) -> None:
"""Clearing a flag on one entity type leaves another type."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.set_flag(EntityType.USER, "123", Flag.NO_REPLY)
await db.clear_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
assert await db.is_flag_set(EntityType.USER, "123", Flag.NO_REPLY) is True
async def test_clear_unset_flag_is_noop(self, db: Database) -> None:
"""Clearing a flag that was never set does not raise or affect other flags."""
await db.set_flag(EntityType.CHANNEL, "123", Flag.NO_REPLY)
await db.clear_flag(EntityType.CHANNEL, "123", Flag.NO_INGEST)
assert await db.is_flag_set(EntityType.CHANNEL, "123", Flag.NO_REPLY) 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)
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)
assert await db.is_channel_ingested(42) is True
async def test_scoped_to_channel(self, db: Database) -> None:
"""Marking one channel as ingested does not affect another channel."""
await db.mark_channel_ingested(42)
assert await db.is_channel_ingested(99) is False
class TestTransactionRollback:
"""Transaction rolls back all changes on error."""
async def test_error_rolls_back_insert(self, db: Database) -> None:
"""An error during a transaction prevents partial data from persisting."""
async def insert_and_fail() -> None:
async with db._transaction():
await db._connection.execute(
"INSERT INTO markov_start_words"
" (channel_id, user_id, word)"
" VALUES (?, ?, ?)",
(1, 100, "should_not_persist"),
)
msg = "simulated failure"
raise RuntimeError(msg)
with pytest.raises(RuntimeError, match="simulated"):
await insert_and_fail()
assert await db.get_random_start_word(1) is None
class TestForgetUser:
"""Atomic forget-user transaction across all tables."""
async def test_deletes_data_and_sets_no_ingest(self, db: Database) -> None:
"""All user data is removed and noIngest flag is set."""
await db.add_markov_data(
[StartWord(1, 100, "Hello")],
[Transition(1, 100, "Hello", "world.")],
)
await db.add_images([ChannelImage(1, 100, "https://example.com/a.png")])
await db.set_flag(EntityType.USER, "100", Flag.ALLOW_PINGS)
await db.forget_user(100, Flag.NO_INGEST)
assert await db.get_random_start_word(1) is None
assert await db.get_random_next_word(1, "Hello") is None
assert await db.get_random_image(1) is None
assert await db.is_flag_set(EntityType.USER, "100", Flag.ALLOW_PINGS) is False
assert await db.is_flag_set(EntityType.USER, "100", Flag.NO_INGEST) is True
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")],
[
Transition(1, 100, "Gone", "away."),
Transition(1, 200, "Keep", "this."),
],
)
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/gone.png"),
ChannelImage(1, 200, "https://example.com/stay.png"),
],
)
await db.set_flag(EntityType.USER, "200", Flag.ALLOW_PINGS)
await db.forget_user(100, Flag.NO_INGEST)
assert await db.get_random_start_word(1) == "Keep"
assert await db.get_random_next_word(1, "Keep") == "this."
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
async def test_preserves_other_entity_type_flags(self, db: Database) -> None:
"""Flags on channels with the same entity ID are not affected."""
await db.set_flag(EntityType.CHANNEL, "100", Flag.NO_REPLY)
await db.set_flag(EntityType.USER, "100", Flag.NO_REPLY)
await db.forget_user(100, Flag.NO_INGEST)
assert await db.is_flag_set(EntityType.CHANNEL, "100", Flag.NO_REPLY) is True
assert await db.is_flag_set(EntityType.USER, "100", Flag.NO_REPLY) is False
async def test_clears_existing_flags_except_no_ingest(self, db: Database) -> None:
"""Existing user flags are cleared but noIngest remains."""
await db.set_flag(EntityType.USER, "100", Flag.NO_REPLY)
await db.set_flag(EntityType.USER, "100", Flag.ALLOW_PINGS)
await db.forget_user(100, Flag.NO_INGEST)
assert await db.is_flag_set(EntityType.USER, "100", Flag.NO_REPLY) is False
assert await db.is_flag_set(EntityType.USER, "100", Flag.ALLOW_PINGS) is False
assert await db.is_flag_set(EntityType.USER, "100", Flag.NO_INGEST) is True
async def test_deletes_across_channels(self, db: Database) -> None:
"""All user data is removed from every channel."""
await db.add_markov_data(
[StartWord(1, 100, "One"), StartWord(2, 100, "Two")],
[
Transition(1, 100, "One", "fish."),
Transition(2, 100, "Two", "fish."),
],
)
await db.add_images(
[
ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(2, 100, "https://example.com/b.png"),
],
)
await db.forget_user(100, Flag.NO_INGEST)
assert await db.get_random_start_word(1) is None
assert await db.get_random_start_word(2) is None
assert await db.get_random_next_word(1, "One") is None
assert await db.get_random_next_word(2, "Two") is None
assert await db.get_random_image(1) is None
assert await db.get_random_image(2) is None
async def test_noop_for_nonexistent_user(self, db: Database) -> None:
"""Forgetting a user with no data does not raise."""
await db.forget_user(999, Flag.NO_INGEST)
assert await db.is_flag_set(EntityType.USER, "999", Flag.NO_INGEST) is True
class TestWriteDurability:
"""Writes persist across close and reopen."""
async def test_markov_data_survives_reopen(self, tmp_path: Path) -> None:
"""Data written via add_markov_data is durable after close/reopen."""
db_path = str(tmp_path / "durability.db")
db = await Database.connect(db_path)
try:
await db.add_markov_data(
[StartWord(1, 100, "Hello")],
[Transition(1, 100, "Hello", "world.")],
)
finally:
await db.close()
db2 = await Database.connect(db_path)
try:
assert await db2.get_random_start_word(1) == "Hello"
assert await db2.get_random_next_word(1, "Hello") == "world."
finally:
await db2.close()
@@ -14,8 +14,10 @@
"""Integration tests for the full ingest → uningest cycle."""
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING
import aiosqlite
import pytest
from crabstero.cache import CachedMessage, IngestCache
@@ -24,18 +26,31 @@ from crabstero.markov import ingest, uningest
from crabstero.messages import uningest_message
if TYPE_CHECKING:
import aiosqlite
from crabstero.database import Database
type SnapshotRows = Callable[["Database"], Awaitable[dict[str, list[aiosqlite.Row]]]]
async def _snapshot(db: Database) -> dict[str, list[aiosqlite.Row]]:
"""Return sorted rows from all Markov-related tables."""
tables: dict[str, list[aiosqlite.Row]] = {}
for table in ("markov_start_words", "markov_transitions", "channel_images"):
async with db._connection.execute(f"SELECT * FROM {table}") as cursor: # noqa: S608
tables[table] = sorted(await cursor.fetchall())
return tables
@pytest.fixture
def snapshot_rows() -> SnapshotRows:
"""Return a snapshot reader for Markov and image tables."""
async def read(db: Database) -> dict[str, list[aiosqlite.Row]]:
tables: dict[str, list[aiosqlite.Row]] = {}
for table in ("markov_start_words", "markov_transitions", "channel_images"):
async with db._connection.execute(
f"SELECT * FROM {table}", # noqa: S608
) as cursor:
tables[table] = sorted(await cursor.fetchall())
return tables
return read
@pytest.fixture
def ingest_cache() -> IngestCache:
"""Return an empty ingest cache for uningest-message tests."""
return IngestCache()
class TestIngestUningestCycle:
@@ -114,12 +129,17 @@ class TestUningestRestoresState:
),
],
)
async def test_uningest_restores_empty_db(self, db: Database, text: str) -> None:
async def test_uningest_restores_empty_db(
self,
db: Database,
text: str,
snapshot_rows: SnapshotRows,
) -> None:
"""Ingest then uningest on an empty database leaves all tables empty."""
before = await _snapshot(db)
before = await snapshot_rows(db)
await ingest(db, 1, 100, text)
await uningest(db, 1, 100, text)
after = await _snapshot(db)
after = await snapshot_rows(db)
assert after == before
@pytest.mark.parametrize(
@@ -136,26 +156,26 @@ class TestUningestRestoresState:
self,
db: Database,
text: str,
snapshot_rows: SnapshotRows,
) -> None:
"""Ingest then uningest preserves unrelated pre-existing data exactly."""
await ingest(db, 99, 200, "Pre-existing data stays safe.")
await db.add_images([ChannelImage(99, 200, "https://example.com/existing.png")])
before = await _snapshot(db)
before = await snapshot_rows(db)
await ingest(db, 1, 100, text)
await uningest(db, 1, 100, text)
after = await _snapshot(db)
after = await snapshot_rows(db)
assert after == before
class TestUningestMessage:
"""Orchestrated uningest via cache lookup and database reversal."""
async def test_content_only(self, db: Database) -> None:
async def test_content_only(self, db: Database, ingest_cache: IngestCache) -> None:
"""Uningest reverses a content-only message via the cache."""
await ingest(db, 1, 100, "Hello beautiful world.")
cache = IngestCache()
cache.put(
ingest_cache.put(
555,
CachedMessage(
channel_id=1,
@@ -166,17 +186,16 @@ class TestUningestMessage:
),
)
await uningest_message(db, cache, 555)
await uningest_message(db, ingest_cache, 555)
assert await db.get_random_start_word(1) is None
assert await db.get_random_next_word(1, "Hello") is None
async def test_embeds_only(self, db: Database) -> None:
async def test_embeds_only(self, db: Database, ingest_cache: IngestCache) -> None:
"""Uningest reverses embed text ingestion."""
await ingest(db, 1, 100, "Embed title here.")
await ingest(db, 1, 100, "Embed description here.")
cache = IngestCache()
cache.put(
ingest_cache.put(
556,
CachedMessage(
channel_id=1,
@@ -187,17 +206,20 @@ class TestUningestMessage:
),
)
await uningest_message(db, cache, 556)
await uningest_message(db, ingest_cache, 556)
assert await db.get_random_start_word(1) is None
async def test_content_with_embeds_and_images(self, db: Database) -> None:
async def test_content_with_embeds_and_images(
self,
db: Database,
ingest_cache: IngestCache,
) -> None:
"""Uningest reverses content, embed text, and image data together."""
await ingest(db, 1, 100, "Body text here.")
await ingest(db, 1, 100, "Embed title.")
await db.add_images([ChannelImage(1, 100, "https://example.com/img.png")])
cache = IngestCache()
cache.put(
ingest_cache.put(
557,
CachedMessage(
channel_id=1,
@@ -208,28 +230,35 @@ class TestUningestMessage:
),
)
await uningest_message(db, cache, 557)
await uningest_message(db, ingest_cache, 557)
assert await db.get_random_start_word(1) is None
assert await db.get_random_image(1) is None
async def test_cache_miss_is_noop(self, db: Database) -> None:
async def test_cache_miss_is_noop(
self,
db: Database,
ingest_cache: IngestCache,
snapshot_rows: SnapshotRows,
) -> None:
"""A message not in the cache leaves the database unchanged."""
await ingest(db, 1, 100, "Keep this data.")
cache = IngestCache()
before = await _snapshot(db)
before = await snapshot_rows(db)
await uningest_message(db, cache, 999)
await uningest_message(db, ingest_cache, 999)
after = await _snapshot(db)
after = await snapshot_rows(db)
assert after == before
async def test_preserves_other_messages(self, db: Database) -> None:
async def test_preserves_other_messages(
self,
db: Database,
ingest_cache: IngestCache,
) -> None:
"""Uningesting one message leaves another message's data intact."""
await ingest(db, 1, 100, "First message.")
await ingest(db, 1, 100, "Second message.")
cache = IngestCache()
cache.put(
ingest_cache.put(
601,
CachedMessage(
channel_id=1,
@@ -240,16 +269,19 @@ class TestUningestMessage:
),
)
await uningest_message(db, cache, 601)
await uningest_message(db, ingest_cache, 601)
assert await db.get_random_start_word(1) == "Second"
assert await db.get_random_next_word(1, "Second") == "message."
async def test_pops_entry_from_cache(self, db: Database) -> None:
async def test_pops_entry_from_cache(
self,
db: Database,
ingest_cache: IngestCache,
) -> None:
"""The cache entry is consumed after uningest."""
await ingest(db, 1, 100, "Hello world.")
cache = IngestCache()
cache.put(
ingest_cache.put(
602,
CachedMessage(
channel_id=1,
@@ -260,6 +292,6 @@ class TestUningestMessage:
),
)
await uningest_message(db, cache, 602)
await uningest_message(db, ingest_cache, 602)
assert cache.pop(602) is None
assert ingest_cache.pop(602) is None
+15
View File
@@ -0,0 +1,15 @@
# 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.
"""Discord-boundary integration tests."""
+129
View File
@@ -0,0 +1,129 @@
# 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.
"""Fixtures used only by Simcord-backed Discord integration tests."""
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING
import pytest
from crabstero.bot import Crabstero
from crabstero.database import Database
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from pathlib import Path
from simcord import ChannelHandle, Env, GuildHandle, MemberActor
type StartWordsForChannel = Callable[[Database, int], Awaitable[list[str]]]
type MakeSimcordTextChannel = Callable[["Env"], Awaitable["SimcordTextChannel"]]
type MakeSimcordMemberChannel = Callable[["Env"], Awaitable["SimcordMemberChannel"]]
@dataclass(frozen=True, slots=True)
class SimcordTextChannel:
"""Guild and text channel created inside a running Simcord environment."""
guild: GuildHandle
channel: ChannelHandle
@dataclass(frozen=True, slots=True)
class SimcordMemberChannel:
"""Guild, text channel, and human member for Discord flow tests."""
guild: GuildHandle
member: MemberActor
channel: ChannelHandle
@pytest.fixture
async def crabstero_bot(tmp_path: Path) -> AsyncGenerator[Crabstero]:
"""Yield the Crabstero bot instance inspected by Discord integration tests."""
bot = Crabstero(str(tmp_path / "crabstero.db"))
try:
yield bot
finally:
bot.ws = None # type: ignore[assignment]
await bot.close()
@pytest.fixture
def simcord_bot(crabstero_bot: Crabstero) -> Crabstero:
"""Expose Crabstero under the fixture name required by Simcord."""
return crabstero_bot
@pytest.fixture
def make_simcord_text_channel() -> MakeSimcordTextChannel:
"""Return a factory for guild text channels in any Simcord environment."""
async def make(env: Env) -> SimcordTextChannel:
guild = env.create_guild()
await env.settle()
channel = guild.create_text_channel("general")
await env.settle()
return SimcordTextChannel(guild, channel)
return make
@pytest.fixture
async def simcord_text_channel(
simcord_env: Env,
make_simcord_text_channel: MakeSimcordTextChannel,
) -> SimcordTextChannel:
"""Create one guild text channel in the default Simcord environment."""
return await make_simcord_text_channel(simcord_env)
@pytest.fixture
def make_simcord_member_channel() -> MakeSimcordMemberChannel:
"""Return a factory for guild/member/channel triples in any Simcord env."""
async def make(env: Env) -> SimcordMemberChannel:
guild = env.create_guild()
await env.settle()
channel = guild.create_text_channel("general")
member = guild.add_member(env.create_user("Ada"))
await env.settle()
return SimcordMemberChannel(guild, member, channel)
return make
@pytest.fixture
async def simcord_member_channel(
simcord_env: Env,
make_simcord_member_channel: MakeSimcordMemberChannel,
) -> SimcordMemberChannel:
"""Create one guild text channel and human member in the default Simcord env."""
return await make_simcord_member_channel(simcord_env)
@pytest.fixture
def start_words_for_channel() -> StartWordsForChannel:
"""Return a reader for persisted Markov start words in one Discord channel."""
async def read(db: Database, channel_id: int) -> list[str]:
async with db._connection.execute(
"SELECT word FROM markov_start_words WHERE channel_id = ? ORDER BY word",
(channel_id,),
) as cursor:
return [str(row[0]) for row in await cursor.fetchall()]
return read
@@ -0,0 +1,251 @@
# 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.
"""Simcord integration tests for Crabstero bot lifecycle behavior."""
from typing import TYPE_CHECKING, Any, cast
import discord
import pytest
from discord import app_commands
from simcord import run
from crabstero import metrics
from crabstero.bot import Crabstero, TrackedModal, TrackedView
if TYPE_CHECKING:
from pathlib import Path
from simcord import Env
from tests.integration.discord.conftest import (
MakeSimcordMemberChannel,
SimcordMemberChannel,
)
def _counter_value(counter: Any, **labels: str) -> float:
"""Return the current value for a labelled Prometheus counter."""
return float(counter.labels(**labels)._value.get())
class TestSetupHook:
"""Bot startup wires Discord cogs and slash commands under Simcord."""
def test_db_before_setup_raises(self, tmp_path: Path) -> None:
"""The database property is unavailable before setup_hook runs."""
bot = Crabstero(str(tmp_path / "not-started.db"))
with pytest.raises(RuntimeError, match="Database is not initialized"):
_ = bot.db
async def test_loads_cogs_and_syncs_commands(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
) -> None:
"""setup_hook loads expected cogs and syncs app commands."""
assert set(crabstero_bot.cogs) == {
"InteractionCog",
"MessageCog",
"ServerEventsCog",
}
commands = simcord_env.backend.commands[None]
assert {name for name, _ in commands} == {"forgetme", "pingme"}
application_id = simcord_env.backend.application_id
http_routes = [f"{method} {path}" for method, path, _ in simcord_env.http_log]
assert f"GET /applications/{application_id}/commands" in http_routes
assert f"PUT /applications/{application_id}/commands" in http_routes
async def test_metrics_server_lifecycle_starts_and_stops(
self,
tmp_path: Path,
) -> None:
"""A configured metrics server is started and stopped with the bot."""
bot = Crabstero(
str(tmp_path / "metrics.db"),
metrics_address=metrics.TcpMetricsAddress("127.0.0.1", 0),
)
try:
async with run(bot):
assert bot._metrics_server is not None
assert bot._metrics_server.port > 0
finally:
bot.ws = None # type: ignore[assignment]
await bot.close()
class TestMetrics:
"""Bot-level metrics are incremented by lifecycle error paths."""
def test_dispatch_increments_discord_event_metric(
self,
crabstero_bot: Crabstero,
) -> None:
"""Dispatch increments the labelled Discord event counter."""
event = "codex_lifecycle_metric"
before = _counter_value(metrics.DISCORD_EVENTS, event=event)
crabstero_bot.dispatch(event)
assert _counter_value(metrics.DISCORD_EVENTS, event=event) == before + 1
async def test_on_error_increments_event_error_metric(
self,
crabstero_bot: Crabstero,
) -> None:
"""Unhandled event listener failures are counted by event name."""
source = "on_message"
before = _counter_value(metrics.ERRORS, source=source)
await crabstero_bot.on_error(source)
assert _counter_value(metrics.ERRORS, source=source) == before + 1
async def test_tracked_view_and_modal_errors_increment_metrics(self) -> None:
"""Tracked UI error handlers increment their error counters."""
view_before = _counter_value(metrics.ERRORS, source="view")
modal_before = _counter_value(metrics.ERRORS, source="modal")
interaction = cast("discord.Interaction", object())
button = cast("discord.ui.Item[TrackedView]", discord.ui.Button(label="Run"))
await TrackedView().on_error(interaction, RuntimeError("view failed"), button)
await TrackedModal(title="Tracked").on_error(
interaction,
RuntimeError("modal failed"),
)
assert _counter_value(metrics.ERRORS, source="view") == view_before + 1
assert _counter_value(metrics.ERRORS, source="modal") == modal_before + 1
class TestCommandErrors:
"""Unhandled app command errors produce an ephemeral fallback response."""
async def test_app_command_error_sends_ephemeral_fallback(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A failing slash command is captured and answered by tree.on_error."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
async def raise_on_set_flag(
_entity_type: str,
_entity_id: str,
_flag_name: str,
) -> None:
raise RuntimeError("database unavailable")
monkeypatch.setattr(crabstero_bot.db, "set_flag", raise_on_set_flag)
errors_before = _counter_value(metrics.ERRORS, source="command")
result = await member.slash(channel, "pingme")
assert result.response is not None
assert result.response.ephemeral is True
assert result.response.content == (
"I encountered an error while processing this command."
" Please try again later."
)
assert _counter_value(metrics.ERRORS, source="command") == errors_before + 1
async def test_app_command_error_mentions_developer_when_metrics_are_enabled(
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
make_simcord_member_channel: MakeSimcordMemberChannel,
) -> None:
"""Configured metrics switch command errors to the notified-developer copy."""
bot = Crabstero(
str(tmp_path / "metrics-command-errors.db"),
metrics_address=metrics.TcpMetricsAddress("127.0.0.1", 0),
)
try:
async with run(bot) as env:
context = await make_simcord_member_channel(env)
member = context.member
channel = context.channel
async def raise_on_set_flag(
_entity_type: str,
_entity_id: str,
_flag_name: str,
) -> None:
raise RuntimeError("database unavailable")
monkeypatch.setattr(bot.db, "set_flag", raise_on_set_flag)
errors_before = _counter_value(metrics.ERRORS, source="command")
result = await member.slash(channel, "pingme")
assert result.response is not None
assert result.response.ephemeral is True
assert result.response.content == (
"I encountered an error while processing this command."
" The developer has been notified,"
" please try again later."
)
assert (
_counter_value(metrics.ERRORS, source="command")
== errors_before + 1
)
finally:
bot.ws = None # type: ignore[assignment]
await bot.close()
async def test_app_command_error_after_defer_sends_ephemeral_followup(
self,
tmp_path: Path,
make_simcord_member_channel: MakeSimcordMemberChannel,
) -> None:
"""A command that has already acknowledged uses a followup fallback."""
bot = Crabstero(str(tmp_path / "deferred-command-errors.db"))
@app_commands.command(
name="deferboom",
description="Fail after acknowledging the interaction.",
)
async def deferboom(interaction: discord.Interaction) -> None:
await interaction.response.defer(ephemeral=True)
raise RuntimeError("deferred command failed")
bot.tree.add_command(deferboom)
try:
async with run(bot) as env:
context = await make_simcord_member_channel(env)
errors_before = _counter_value(metrics.ERRORS, source="command")
result = await context.member.slash(context.channel, "deferboom")
assert result.deferred is True
assert result.response is None
assert len(result.followups) == 1
followup = result.followups[0]
assert followup.ephemeral is True
assert followup.content == (
"I encountered an error while processing this command."
" Please try again later."
)
assert (
_counter_value(metrics.ERRORS, source="command")
== errors_before + 1
)
finally:
bot.ws = None # type: ignore[assignment]
await bot.close()
@@ -0,0 +1,276 @@
# 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.
"""Simcord integration tests for Crabstero slash commands."""
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING
import pytest
from crabstero.database import ChannelImage, StartWord, Transition
from crabstero.flags import Flag
if TYPE_CHECKING:
from simcord import Env
from crabstero.bot import Crabstero
from crabstero.database import Database
from tests.integration.discord.conftest import SimcordMemberChannel
type SeedForgetmeData = Callable[["Database", int], Awaitable[None]]
type ForgetmeRowsForUser = Callable[["Database", int], Awaitable[tuple[int, int, int]]]
@pytest.fixture
def seed_forgetme_data() -> SeedForgetmeData:
"""Return a seeder for user-owned and unrelated /forgetme database rows."""
async def seed(db: Database, user_id: int) -> None:
await db.add_markov_data(
[
StartWord(10, user_id, "delete"),
StartWord(10, 999, "keep"),
],
[
Transition(10, user_id, "delete", "me."),
Transition(10, 999, "keep", "me."),
],
)
await db.add_images(
[
ChannelImage(10, user_id, "https://example.com/delete.png"),
ChannelImage(10, 999, "https://example.com/keep.png"),
],
)
return seed
@pytest.fixture
def forgetme_rows_for_user() -> ForgetmeRowsForUser:
"""Return a reader for Markov/image row counts owned by one user."""
async def read(db: Database, user_id: int) -> tuple[int, int, int]:
async with db._connection.execute(
"SELECT COUNT(*) FROM markov_start_words WHERE user_id = ?",
(user_id,),
) as cursor:
row = await cursor.fetchone()
assert row is not None
start_words = row[0]
async with db._connection.execute(
"SELECT COUNT(*) FROM markov_transitions WHERE user_id = ?",
(user_id,),
) as cursor:
row = await cursor.fetchone()
assert row is not None
transitions = row[0]
async with db._connection.execute(
"SELECT COUNT(*) FROM channel_images WHERE user_id = ?",
(user_id,),
) as cursor:
row = await cursor.fetchone()
assert row is not None
images = row[0]
return int(start_words), int(transitions), int(images)
return read
class TestPingMe:
"""The /pingme command toggles persisted user opt-in state."""
async def test_first_call_opts_in(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
) -> None:
"""The first /pingme call sets allowPings and responds ephemerally."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
result = await member.slash(channel, "pingme")
assert result.response is not None
assert result.response.ephemeral is True
assert "I will now ping you" in result.response.content
assert await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.ALLOW_PINGS,
)
async def test_second_call_opts_out(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
) -> None:
"""The second /pingme call clears allowPings and responds ephemerally."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await member.slash(channel, "pingme")
result = await member.slash(channel, "pingme")
assert result.response is not None
assert result.response.ephemeral is True
assert "I will no longer ping you" in result.response.content
assert not await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.ALLOW_PINGS,
)
class TestForgetMe:
"""The /forgetme command confirms, cancels, and times out hermetically."""
async def test_initial_response_has_confirmation_buttons(
self,
simcord_member_channel: SimcordMemberChannel,
) -> None:
"""The first /forgetme response is ephemeral and asks for confirmation."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
result = await member.slash(channel, "forgetme")
assert result.response is not None
assert result.response.ephemeral is True
assert "Would you like to proceed?" in result.response.content
labels = [
component["label"]
for row in result.response.components
for component in row["components"]
]
assert labels == ["Confirm", "Cancel"]
async def test_confirm_deletes_user_data_and_sets_no_ingest(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
seed_forgetme_data: SeedForgetmeData,
forgetme_rows_for_user: ForgetmeRowsForUser,
) -> None:
"""Confirming /forgetme deletes user data and persists noIngest."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_forgetme_data(crabstero_bot.db, member.id)
await crabstero_bot.db.set_flag("user", str(member.id), Flag.ALLOW_PINGS)
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_REPLY)
prompt = (await member.slash(channel, "forgetme")).response
assert prompt is not None
result = await member.click(prompt, label="Confirm")
assert result.response is not None
assert result.response.ephemeral is True
assert "I have deleted your data" in result.response.content
assert result.response.components == []
assert await forgetme_rows_for_user(crabstero_bot.db, member.id) == (0, 0, 0)
assert await forgetme_rows_for_user(crabstero_bot.db, 999) == (1, 1, 1)
assert await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.NO_INGEST,
)
assert not await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.ALLOW_PINGS,
)
assert not await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.NO_REPLY,
)
async def test_cancel_leaves_user_data_and_flags(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
seed_forgetme_data: SeedForgetmeData,
forgetme_rows_for_user: ForgetmeRowsForUser,
) -> None:
"""Cancelling /forgetme leaves data and user flags unchanged."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_forgetme_data(crabstero_bot.db, member.id)
await crabstero_bot.db.set_flag("user", str(member.id), Flag.ALLOW_PINGS)
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_REPLY)
prompt = (await member.slash(channel, "forgetme")).response
assert prompt is not None
result = await member.click(prompt, label="Cancel")
assert result.response is not None
assert result.response.ephemeral is True
assert result.response.content == (
"Action cancelled. I have not modified your data."
)
assert result.response.components == []
assert await forgetme_rows_for_user(crabstero_bot.db, member.id) == (1, 1, 1)
assert await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.ALLOW_PINGS,
)
assert await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.NO_REPLY,
)
assert not await crabstero_bot.db.is_flag_set(
"user",
str(member.id),
Flag.NO_INGEST,
)
async def test_already_forgotten_user_gets_terminal_response(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
) -> None:
"""A noIngest user does not get another confirmation view."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_INGEST)
result = await member.slash(channel, "forgetme")
assert result.response is not None
assert result.response.ephemeral is True
assert result.response.content == (
"I have already removed your data and I am not using your messages."
)
assert result.response.components == []
async def test_confirmation_timeout_removes_view(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
) -> None:
"""The /forgetme view timeout edits the original response without sleeping."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
prompt = (await member.slash(channel, "forgetme")).response
assert prompt is not None
await simcord_env.advance_time(181)
assert prompt.content == (
"This timed out. Run `/forgetme` again if you still want to."
)
assert prompt.components == []
+470
View File
@@ -0,0 +1,470 @@
# 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.
"""Simcord integration tests for Discord message events."""
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING
import discord
import pytest
from simcord import run
from crabstero.bot import Crabstero
from crabstero.database import ChannelImage, StartWord, Transition
from crabstero.flags import EntityType, Flag
from crabstero.messages import ingest_message, reply_to_message
if TYPE_CHECKING:
from pathlib import Path
from simcord import Env
from crabstero.database import Database
from tests.integration.discord.conftest import (
MakeSimcordMemberChannel,
SimcordMemberChannel,
StartWordsForChannel,
)
type SeedReply = Callable[["Database", int], Awaitable[None]]
@pytest.fixture
def seed_reply() -> SeedReply:
"""Return a seeder for deterministic generated replies in one channel."""
async def seed(db: Database, channel_id: int) -> None:
await db.add_markov_data(
[StartWord(channel_id, 123, "Generated")],
[Transition(channel_id, 123, "Generated", "reply.")],
)
return seed
class TestMessageIngestion:
"""Normal Discord message flow populates the real database."""
async def test_normal_guild_user_message_is_ingested(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""A default guild text message adds Markov data."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await member.send(channel, "Alpha beta.")
assert await start_words_for_channel(crabstero_bot.db, channel.id) == ["Alpha"]
assert (
await crabstero_bot.db.get_random_next_word(channel.id, "Alpha") == "beta."
)
async def test_bot_messages_are_ignored(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""A bot-authored gateway message is not ingested."""
channel = simcord_member_channel.channel
simcord_env.backend.create_message(
channel.id,
simcord_env.backend.bot_user.id,
"Ignore bot.",
)
await simcord_env.settle()
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
async def test_dm_messages_are_ignored(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""DM messages are outside the guild channel types Crabstero handles."""
user = simcord_env.create_user("Ada")
await user.send_dm("Direct message.")
assert (
await start_words_for_channel(
crabstero_bot.db,
user.dm_channel.id,
)
== []
)
async def test_thread_messages_do_not_create_separate_chain(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""Messages inside threads are not ingested under the thread channel id."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
simcord_env.backend.create_thread(channel.id, "thread", member.id)
await simcord_env.settle()
thread = channel.threads[0]
await member.send(thread, "Thread only.")
assert await start_words_for_channel(crabstero_bot.db, thread.id) == []
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
async def test_embed_message_text_and_image_are_ingested(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""Embed titles, descriptions, and image URLs are ingested."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
simcord_env.backend.create_message(
channel.id,
member.id,
embeds=[
{
"title": "Title words.",
"description": "Description words.",
"image": {"url": "https://example.com/embed.png"},
},
],
)
await simcord_env.settle()
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
"Description",
"Title",
]
assert (
await crabstero_bot.db.get_random_image(channel.id)
== "https://example.com/embed.png"
)
async def test_empty_messages_are_not_ingested(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""A message with no content or embeds is ignored."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await member.send(channel)
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
async def test_ingest_only_mode_ingests_without_replying(
self,
tmp_path: Path,
make_simcord_member_channel: MakeSimcordMemberChannel,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""An ingest-only bot still ingests eligible messages and never replies."""
bot = Crabstero(str(tmp_path / "ingest-only.db"), ingest_only=True)
try:
async with run(bot) as env:
context = await make_simcord_member_channel(env)
member = context.member
channel = context.channel
await member.send(channel, "Ingest only.")
assert await start_words_for_channel(bot.db, channel.id) == ["Ingest"]
assert [message.content for message in channel.history()] == [
"Ingest only.",
]
finally:
bot.ws = None # type: ignore[assignment]
await bot.close()
class TestReplies:
"""Mentions produce replies through the real Discord message path."""
async def test_mentioning_bot_sends_seeded_reply(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
seed_reply: SeedReply,
) -> None:
"""A mention causes a deterministic Markov reply to be posted."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_reply(crabstero_bot.db, channel.id)
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
history = channel.history()
assert [message.content for message in history] == [
f"<@{simcord_env.backend.bot_user.id}> please reply",
"Generated reply.",
]
assert history[-1].reference is not None
assert history[-1].author == simcord_env.bot.user
async def test_mention_reply_can_include_embed_and_image(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
monkeypatch: pytest.MonkeyPatch,
seed_reply: SeedReply,
) -> None:
"""The optional reply embed path uses generated text and stored images."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_reply(crabstero_bot.db, channel.id)
await crabstero_bot.db.add_images(
[ChannelImage(channel.id, 123, "https://example.com/reply.png")],
)
monkeypatch.setattr("crabstero.messages.secrets.randbelow", lambda _upper: 95)
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
reply = channel.history()[-1]
assert reply.content == "Generated reply."
assert len(reply.embeds) == 1
assert reply.embeds[0].title == "Generated reply."
assert reply.embeds[0].description == "Generated reply."
assert reply.embeds[0].image.url == "https://example.com/reply.png"
async def test_mention_in_thread_uses_parent_channel_chain(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
seed_reply: SeedReply,
) -> None:
"""A thread mention generates from the parent channel's Markov chain."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_reply(crabstero_bot.db, channel.id)
simcord_env.backend.create_thread(channel.id, "thread", member.id)
await simcord_env.settle()
thread = channel.threads[0]
await member.send(thread, f"<@{simcord_env.backend.bot_user.id}> thread reply")
assert [message.content for message in thread.history()] == [
f"<@{simcord_env.backend.bot_user.id}> thread reply",
"Generated reply.",
]
async def test_allow_pings_flag_allows_generated_user_mentions(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
) -> None:
"""Generated mentions are allowed for users who opted in."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
target = simcord_env.create_user("Mentioned")
await crabstero_bot.db.add_markov_data(
[StartWord(channel.id, 123, "Hello")],
[Transition(channel.id, 123, "Hello", f"{target.mention}.")],
)
await crabstero_bot.db.set_flag(
EntityType.USER,
str(target.id),
Flag.ALLOW_PINGS,
)
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
assert channel.history()[-1].content == f"Hello {target.mention}."
async def test_no_reply_flag_suppresses_reply(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
seed_reply: SeedReply,
) -> None:
"""The noReply flag prevents a mention response."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await seed_reply(crabstero_bot.db, channel.id)
await crabstero_bot.db.set_flag(
EntityType.CHANNEL,
str(channel.id),
Flag.NO_REPLY,
)
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
assert [message.content for message in channel.history()] == [
f"<@{simcord_env.backend.bot_user.id}> please reply",
]
async def test_missing_send_permission_suppresses_reply(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
seed_reply: SeedReply,
) -> None:
"""A bot without send_messages permission does not reply."""
guild = simcord_env.create_guild()
bot_role = guild.roles[simcord_env.backend.bot_user.name]
channel = guild.create_text_channel(
"readonly",
overwrites={
bot_role: discord.PermissionOverwrite(send_messages=False),
},
)
member = guild.add_member(simcord_env.create_user("Ada"))
await simcord_env.settle()
await seed_reply(crabstero_bot.db, channel.id)
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
assert [message.content for message in channel.history()] == [
f"<@{simcord_env.backend.bot_user.id}> please reply",
]
async def test_direct_dm_reply_is_ignored(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
) -> None:
"""The reply helper ignores messages outside guilds."""
user = simcord_env.create_user("Ada")
message = await user.send_dm(f"<@{simcord_env.backend.bot_user.id}> hi")
await reply_to_message(crabstero_bot.db, message)
assert [message.content for message in user.dm_channel.history()] == [
f"<@{simcord_env.backend.bot_user.id}> hi",
]
class TestDeletes:
"""Raw delete events reverse recent message ingestion."""
async def test_delete_reverses_recent_ingest(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""Deleting a cached message removes its Markov rows."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
message = await member.send(channel, "Delete me.")
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
"Delete",
]
await member.delete(message)
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
async def test_bulk_delete_reverses_recent_ingests(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""Bulk-deleting cached messages removes their Markov rows."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
first = await member.send(channel, "Bulk one.")
second = await member.send(channel, "Bulk two.")
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
"Bulk",
"Bulk",
]
simcord_env.backend.bulk_delete_messages(channel.id, [first.id, second.id])
await simcord_env.settle()
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
class TestFlagSuppression:
"""Message behavior respects persisted noIngest and noReply flags."""
async def test_no_ingest_user_flag_suppresses_ingestion(
self,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""The noIngest flag prevents storing a user's message."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await crabstero_bot.db.set_flag(
EntityType.USER,
str(member.id),
Flag.NO_INGEST,
)
await member.send(channel, "Do not learn.")
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
async def test_direct_dm_ingest_is_ignored(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""The ingest helper ignores messages outside guilds."""
user = simcord_env.create_user("Ada")
message = await user.send_dm("Direct helper call.")
await ingest_message(crabstero_bot.db, message)
assert (
await start_words_for_channel(
crabstero_bot.db,
user.dm_channel.id,
)
== []
)
async def test_direct_bot_mention_ingest_is_ignored(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""The ingest helper ignores messages that mention the bot."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
message = await member.send(
channel,
f"<@{simcord_env.backend.bot_user.id}> do not learn",
)
await ingest_message(crabstero_bot.db, message)
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
@@ -0,0 +1,277 @@
# 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.
"""Simcord integration tests for server event ingestion triggers."""
from typing import TYPE_CHECKING
import discord
import pytest
from simcord.backend.models import Overwrite
from simcord.enums import OverwriteType
from crabstero.listeners.server_events import ServerEventsCog
from crabstero.tasks.ingestion import ingest_channel
if TYPE_CHECKING:
from simcord import Env
from crabstero.bot import Crabstero
from tests.integration.discord.conftest import (
SimcordMemberChannel,
SimcordTextChannel,
StartWordsForChannel,
)
@pytest.fixture
def server_events_cog(crabstero_bot: Crabstero) -> ServerEventsCog:
"""Return the loaded server-events cog from the Simcord-backed bot."""
cog = crabstero_bot.get_cog("ServerEventsCog")
assert isinstance(cog, ServerEventsCog)
return cog
class TestGuildEvents:
"""Guild availability and joins queue local channel history ingestion."""
async def test_guild_available_queues_text_and_voice_channels(
self,
simcord_env: Env,
server_events_cog: ServerEventsCog,
) -> None:
"""Available guilds enqueue all textable channels."""
guild = simcord_env.create_guild()
await simcord_env.settle()
text = guild.create_text_channel("general")
voice = guild.create_voice_channel("voice")
await simcord_env.settle()
cached_guild = simcord_env.bot.get_guild(guild.id)
assert cached_guild is not None
await server_events_cog.on_guild_available(cached_guild)
await simcord_env.settle()
assert await server_events_cog.bot.db.is_channel_ingested(text.id)
assert await server_events_cog.bot.db.is_channel_ingested(voice.id)
async def test_guild_join_queues_ingestion_and_attempts_owner_notification(
self,
simcord_env: Env,
simcord_text_channel: SimcordTextChannel,
server_events_cog: ServerEventsCog,
) -> None:
"""Joining a guild enqueues ingestion and uses the fake owner DM path."""
guild = simcord_text_channel.guild
channel = simcord_text_channel.channel
cached_guild = simcord_env.bot.get_guild(guild.id)
assert cached_guild is not None
await server_events_cog.on_guild_join(cached_guild)
await simcord_env.settle()
assert await server_events_cog.bot.db.is_channel_ingested(channel.id)
http_routes = [f"{method} {path}" for method, path, _ in simcord_env.http_log]
assert "GET /oauth2/applications/@me" in http_routes
assert "POST /users/@me/channels" in http_routes
assert any(route.endswith("/messages") for route in http_routes)
class TestPermissionUpdateEvents:
"""Permission-changing events trigger ingestion only when relevant."""
async def test_role_update_queues_only_when_permissions_change_for_bot_role(
self,
simcord_env: Env,
simcord_text_channel: SimcordTextChannel,
crabstero_bot: Crabstero,
) -> None:
"""Role updates require changed permissions and bot membership."""
guild = simcord_text_channel.guild
channel = simcord_text_channel.channel
role = guild.create_role("reader", permissions=discord.Permissions.none())
await simcord_env.settle()
simcord_env.backend.edit_role(
guild.id,
role.id,
{"permissions": discord.Permissions(view_channel=True).value},
)
await simcord_env.settle()
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
bot_role = guild.roles[simcord_env.backend.bot_user.name]
simcord_env.backend.edit_role(
guild.id,
bot_role.id,
{
"permissions": discord.Permissions(
view_channel=True,
read_message_history=True,
).value,
},
)
await simcord_env.settle()
assert await crabstero_bot.db.is_channel_ingested(channel.id)
async def test_role_update_same_permissions_does_not_queue(
self,
simcord_env: Env,
simcord_text_channel: SimcordTextChannel,
crabstero_bot: Crabstero,
) -> None:
"""A role update without permission changes is ignored."""
guild = simcord_text_channel.guild
channel = simcord_text_channel.channel
role = guild.create_role(
"reader",
permissions=discord.Permissions(read_message_history=True),
)
simcord_env.backend.add_member_role(
guild.id,
simcord_env.backend.bot_user.id,
role.id,
)
await simcord_env.settle()
simcord_env.backend.edit_role(
guild.id,
role.id,
{"permissions": discord.Permissions(read_message_history=True).value},
)
await simcord_env.settle()
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
async def test_channel_update_queues_only_when_overwrites_change(
self,
simcord_env: Env,
simcord_text_channel: SimcordTextChannel,
crabstero_bot: Crabstero,
) -> None:
"""Channel updates without overwrite changes do not enqueue ingestion."""
guild = simcord_text_channel.guild
channel = simcord_text_channel.channel
simcord_env.backend.edit_channel(channel.id, {"topic": "no permission change"})
await simcord_env.settle()
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
simcord_env.backend.set_overwrite(
channel.id,
Overwrite(
target_id=guild.default_role.id,
type=OverwriteType.ROLE,
allow=discord.Permissions(read_message_history=True).value,
deny=0,
),
)
await simcord_env.settle()
assert await crabstero_bot.db.is_channel_ingested(channel.id)
class TestIngestionTasks:
"""Server-triggered ingestion task behavior stays local and deterministic."""
async def test_duplicate_channel_queue_requests_share_one_task(
self,
simcord_env: Env,
simcord_text_channel: SimcordTextChannel,
crabstero_bot: Crabstero,
) -> None:
"""Queuing the same channel twice before the loop runs creates one task."""
channel = simcord_text_channel.channel
cached_channel = simcord_env.bot.get_channel(channel.id)
assert isinstance(cached_channel, discord.TextChannel)
crabstero_bot.queue_channel_for_ingestion(cached_channel)
crabstero_bot.queue_channel_for_ingestion(cached_channel)
assert list(crabstero_bot._ingestion_tasks) == [channel.id]
await simcord_env.settle()
async def test_channel_history_ingestion_requires_read_history_permission(
self,
simcord_env: Env,
crabstero_bot: Crabstero,
) -> None:
"""A channel missing read history permission is skipped."""
guild = simcord_env.create_guild()
channel = guild.create_text_channel(
"hidden-history",
overwrites={
guild.default_role: discord.PermissionOverwrite(
read_message_history=False,
),
},
)
await simcord_env.settle()
cached_channel = simcord_env.bot.get_channel(channel.id)
assert isinstance(cached_channel, discord.TextChannel)
await ingest_channel(cached_channel, crabstero_bot.db)
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
async def test_channel_history_ingestion_reads_existing_messages(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""Bulk channel ingestion reads historical messages through Discord."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
cached_channel = simcord_env.bot.get_channel(channel.id)
assert isinstance(cached_channel, discord.TextChannel)
simcord_env.backend.create_message(
channel.id,
member.id,
"Historical message.",
broadcast=False,
)
await ingest_channel(cached_channel, crabstero_bot.db)
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
"Historical",
]
assert await crabstero_bot.db.is_channel_ingested(channel.id)
async def test_already_ingested_channel_is_skipped(
self,
simcord_env: Env,
simcord_member_channel: SimcordMemberChannel,
crabstero_bot: Crabstero,
start_words_for_channel: StartWordsForChannel,
) -> None:
"""A channel marked ingested is not read again."""
member = simcord_member_channel.member
channel = simcord_member_channel.channel
await crabstero_bot.db.mark_channel_ingested(channel.id)
cached_channel = simcord_env.bot.get_channel(channel.id)
assert isinstance(cached_channel, discord.TextChannel)
simcord_env.backend.create_message(
channel.id,
member.id,
"Historical message.",
broadcast=False,
)
await ingest_channel(cached_channel, crabstero_bot.db)
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
+15
View File
@@ -0,0 +1,15 @@
# 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.
"""Metrics integration tests for Crabstero."""
+196
View File
@@ -0,0 +1,196 @@
# 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.
"""Integration tests for the Prometheus metrics HTTP server."""
import errno
import os
import socket
from dataclasses import dataclass
from http import HTTPStatus
from typing import TYPE_CHECKING, Literal, cast
import aiohttp
import pytest
from crabstero.metrics import MetricsServer, TcpMetricsAddress, UnixMetricsAddress
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from pathlib import Path
type MetricsTransport = Literal["tcp", "unix-socket"]
requires_unix_socket = pytest.mark.skipif(
os.name != "posix" or not hasattr(socket, "AF_UNIX"),
reason="Unix-socket metrics transports require Unix-domain socket support",
)
@dataclass(frozen=True, slots=True)
class MetricsEndpoint:
"""Client details for one running metrics transport."""
base_url: str
unix_socket_path: str | None = None
def client_session(self) -> aiohttp.ClientSession:
"""Create an aiohttp client session for this metrics transport."""
connector = (
aiohttp.UnixConnector(path=self.unix_socket_path)
if self.unix_socket_path is not None
else None
)
return aiohttp.ClientSession(connector=connector)
@pytest.fixture(
params=[
pytest.param("tcp", id="tcp"),
pytest.param("unix-socket", marks=requires_unix_socket, id="unix-socket"),
],
)
async def metrics_endpoint(
request: pytest.FixtureRequest,
tmp_path: Path,
) -> AsyncGenerator[MetricsEndpoint]:
"""Start a MetricsServer on each supported transport."""
transport = cast("MetricsTransport", request.param)
match transport:
case "tcp":
server = MetricsServer(TcpMetricsAddress("127.0.0.1", 0))
await server.start()
endpoint = MetricsEndpoint(f"http://127.0.0.1:{server.port}")
case "unix-socket":
socket_path = tmp_path / "metrics.sock"
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
await server.start()
endpoint = MetricsEndpoint(
"http://crabstero",
unix_socket_path=str(socket_path),
)
try:
yield endpoint
finally:
await server.stop()
class TestMetricsServer:
"""HTTP server serves Prometheus metrics on /metrics."""
@requires_unix_socket
def test_unix_socket_server_has_no_tcp_port(self, tmp_path: Path) -> None:
"""Unix-socket metrics servers do not expose a TCP port."""
server = MetricsServer(UnixMetricsAddress(str(tmp_path / "metrics.sock")))
with pytest.raises(RuntimeError, match="does not have a TCP port"):
_ = server.port
async def test_serves_metrics_endpoint(
self,
metrics_endpoint: MetricsEndpoint,
) -> None:
"""GET /metrics returns 200 with metric output containing our metrics."""
async with (
metrics_endpoint.client_session() as session,
session.get(f"{metrics_endpoint.base_url}/metrics") as resp,
):
assert resp.status == HTTPStatus.OK
body = await resp.text()
assert "crabstero_build_info" in body
async def test_non_metrics_path_returns_404(
self,
metrics_endpoint: MetricsEndpoint,
) -> None:
"""GET on an unknown path returns 404."""
async with (
metrics_endpoint.client_session() as session,
session.get(f"{metrics_endpoint.base_url}/notfound") as resp,
):
assert resp.status == HTTPStatus.NOT_FOUND
@requires_unix_socket
async def test_unix_socket_mode_is_applied(self, tmp_path: Path) -> None:
"""Configured Unix socket mode is applied after startup."""
socket_path = tmp_path / "metrics.sock"
server = MetricsServer(UnixMetricsAddress(str(socket_path), mode=0o666))
await server.start()
try:
assert socket_path.stat().st_mode & 0o777 == 0o666
finally:
await server.stop()
@requires_unix_socket
async def test_stale_unix_socket_path_is_recovered(
self,
tmp_path: Path,
) -> None:
"""Startup removes an abandoned Unix socket file left by a crash."""
socket_path = tmp_path / "metrics.sock"
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as stale_socket:
stale_socket.bind(str(socket_path))
server = MetricsServer(UnixMetricsAddress(str(socket_path), mode=0o666))
await server.start()
try:
assert socket_path.stat().st_mode & 0o777 == 0o666
connector = aiohttp.UnixConnector(path=str(socket_path))
async with (
aiohttp.ClientSession(connector=connector) as session,
session.get("http://crabstero/metrics") as resp,
):
assert resp.status == HTTPStatus.OK
body = await resp.text()
assert "crabstero_build_info" in body
finally:
await server.stop()
@requires_unix_socket
async def test_active_unix_socket_path_fails(
self,
tmp_path: Path,
) -> None:
"""Startup refuses to replace a Unix socket path that is still in use."""
socket_path = tmp_path / "metrics.sock"
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as active_socket:
active_socket.bind(str(socket_path))
active_socket.listen(1)
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
with pytest.raises(OSError, match="already in use") as exc_info:
await server.start()
assert exc_info.value.errno == errno.EADDRINUSE
@requires_unix_socket
async def test_non_socket_unix_path_fails(self, tmp_path: Path) -> None:
"""Startup refuses to replace a non-socket path."""
socket_path = tmp_path / "metrics.sock"
socket_path.write_text("")
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
with pytest.raises(FileExistsError):
await server.start()
async def test_starting_server_twice_fails(self) -> None:
"""A running metrics server cannot be started twice."""
server = MetricsServer(TcpMetricsAddress("127.0.0.1", 0))
await server.start()
try:
with pytest.raises(RuntimeError, match="already running"):
await server.start()
finally:
await server.stop()