Expanded integration coverage and enforced test categories.
This commit is contained in:
@@ -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"
|
||||
@@ -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."""
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
+72
-40
@@ -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
|
||||
@@ -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."""
|
||||
@@ -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 == []
|
||||
@@ -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) == []
|
||||
@@ -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."""
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user