Enabled all Ruff linter rules and fixed resulting violations.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 28s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s

This commit is contained in:
2026-03-29 20:08:30 -04:00
parent 603a927f12
commit abe1444e4c
16 changed files with 188 additions and 125 deletions
+2 -2
View File
@@ -106,7 +106,7 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
host, sep, port_str = args.listen_metrics.rpartition(":") host, sep, port_str = args.listen_metrics.rpartition(":")
if not sep or not host: if not sep or not host:
parser.error( parser.error(
"--listen-metrics must be in HOST:PORT format (e.g. 127.0.0.1:9090)" "--listen-metrics must be in HOST:PORT format (e.g. 127.0.0.1:9090)",
) )
try: try:
args.listen_metrics = (host, int(port_str)) args.listen_metrics = (host, int(port_str))
@@ -117,7 +117,7 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser.error( parser.error(
"a Discord bot token is required via --token," "a Discord bot token is required via --token,"
" the TOKEN environment variable," " the TOKEN environment variable,"
" or a systemd credential named 'token'" " or a systemd credential named 'token'",
) )
return args return args
+18 -6
View File
@@ -58,6 +58,7 @@ class Crabstero(commands.Bot):
self, self,
token: str, token: str,
database_path: str, database_path: str,
*,
ingest_only: bool = False, ingest_only: bool = False,
metrics_address: tuple[str, int] | None = None, metrics_address: tuple[str, int] | None = None,
) -> None: ) -> None:
@@ -104,7 +105,8 @@ class Crabstero(commands.Bot):
:raises RuntimeError: If accessed before :meth:`setup_hook` has run. :raises RuntimeError: If accessed before :meth:`setup_hook` has run.
""" """
if self._db is None: if self._db is None:
raise RuntimeError("Database is not initialized") msg = "Database is not initialized"
raise RuntimeError(msg)
return self._db return self._db
@property @property
@@ -123,7 +125,11 @@ class Crabstero(commands.Bot):
await server.start() await server.start()
self._metrics_server = server self._metrics_server = server
from crabstero.listeners import interaction, message, server_events from crabstero.listeners import ( # noqa: PLC0415
interaction,
message,
server_events,
)
if not self._ingest_only: if not self._ingest_only:
await interaction.setup(self) await interaction.setup(self)
@@ -133,7 +139,8 @@ class Crabstero(commands.Bot):
@self.tree.error @self.tree.error
async def on_app_command_error( async def on_app_command_error(
interaction: discord.Interaction, error: AppCommandError interaction: discord.Interaction,
error: AppCommandError,
) -> None: ) -> None:
original = ( original = (
error.original if isinstance(error, CommandInvokeError) else error error.original if isinstance(error, CommandInvokeError) else error
@@ -203,14 +210,18 @@ class Crabstero(commands.Bot):
@override @override
async def on_error( async def on_error(
self, event_method: str, /, *args: object, **kwargs: object self,
event_method: str,
/,
*args: object,
**kwargs: object,
) -> None: ) -> None:
"""Increment the global error counter for event listener exceptions. """Increment the global error counter for event listener exceptions.
:param event_method: The name of the event that raised the exception. :param event_method: The name of the event that raised the exception.
""" """
ERRORS.labels(source=event_method).inc() ERRORS.labels(source=event_method).inc()
logger.error("Unhandled exception in %s.", event_method, exc_info=True) logger.error("Unhandled exception in %s.", event_method, exc_info=True) # noqa: LOG014
async def on_ready(self) -> None: async def on_ready(self) -> None:
"""Log that the bot has started successfully.""" """Log that the bot has started successfully."""
@@ -250,7 +261,8 @@ class Crabstero(commands.Bot):
if channel.id in self._ingestion_tasks: if channel.id in self._ingestion_tasks:
return return
task = asyncio.create_task( task = asyncio.create_task(
self._ingest_one(channel), name=f"ingest-{channel.id}" self._ingest_one(channel),
name=f"ingest-{channel.id}",
) )
self._ingestion_tasks[channel.id] = task self._ingestion_tasks[channel.id] = task
task.add_done_callback(lambda t: self._on_ingestion_done(channel.id, t)) task.add_done_callback(lambda t: self._on_ingestion_done(channel.id, t))
+2 -1
View File
@@ -103,7 +103,8 @@ class IngestCache:
if self._task is not None: if self._task is not None:
return return
self._task = asyncio.create_task( self._task = asyncio.create_task(
self._cleanup_loop(), name="ingest-cache-cleanup" self._cleanup_loop(),
name="ingest-cache-cleanup",
) )
async def stop(self) -> None: async def stop(self) -> None:
+11 -3
View File
@@ -261,7 +261,9 @@ class Database:
return row[0] if row else None return row[0] if row else None
async def get_random_completing_next_word( async def get_random_completing_next_word(
self, channel_id: int, word: str self,
channel_id: int,
word: str,
) -> str | None: ) -> str | None:
"""Return a random sentence-ending next word for a given word in a channel. """Return a random sentence-ending next word for a given word in a channel.
@@ -352,7 +354,10 @@ class Database:
) )
async def clear_flag( async def clear_flag(
self, entity_type: str, entity_id: str, flag_name: str self,
entity_type: str,
entity_id: str,
flag_name: str,
) -> None: ) -> None:
"""Clear a flag on a given entity. If the flag is not set, this is a no-op. """Clear a flag on a given entity. If the flag is not set, this is a no-op.
@@ -370,7 +375,10 @@ class Database:
) )
async def is_flag_set( async def is_flag_set(
self, entity_type: str, entity_id: str, flag_name: str self,
entity_type: str,
entity_id: str,
flag_name: str,
) -> bool: ) -> bool:
"""Check whether a flag is set on a given entity. """Check whether a flag is set on a given entity.
+28 -9
View File
@@ -39,7 +39,10 @@ class ForgetMeView(TrackedView):
"""Confirmation view with Confirm/Cancel buttons for /forgetme.""" """Confirmation view with Confirm/Cancel buttons for /forgetme."""
def __init__( def __init__(
self, user_id: int, db: Database, interaction: discord.Interaction self,
user_id: int,
db: Database,
interaction: discord.Interaction,
) -> None: ) -> None:
"""Create a new ForgetMeView. """Create a new ForgetMeView.
@@ -64,12 +67,14 @@ class ForgetMeView(TrackedView):
@discord.ui.button(label="Confirm", style=discord.ButtonStyle.danger) @discord.ui.button(label="Confirm", style=discord.ButtonStyle.danger)
async def confirm( async def confirm(
self, interaction: discord.Interaction, button: discord.ui.Button[ForgetMeView] self,
interaction: discord.Interaction,
_button: discord.ui.Button[ForgetMeView],
) -> None: ) -> None:
"""Delete all user data, clear flags, and set noIngest. """Delete all user data, clear flags, and set noIngest.
:param interaction: The interaction event. :param interaction: The interaction event.
:param button: The button that was pressed. :param _button: The button that was pressed.
""" """
self._responded = True self._responded = True
await self._db.forget_user(self._user_id, Flag.NO_INGEST) await self._db.forget_user(self._user_id, Flag.NO_INGEST)
@@ -85,12 +90,14 @@ class ForgetMeView(TrackedView):
@discord.ui.button(label="Cancel", style=discord.ButtonStyle.secondary) @discord.ui.button(label="Cancel", style=discord.ButtonStyle.secondary)
async def cancel( async def cancel(
self, interaction: discord.Interaction, button: discord.ui.Button[ForgetMeView] self,
interaction: discord.Interaction,
_button: discord.ui.Button[ForgetMeView],
) -> None: ) -> None:
"""Cancel the /forgetme action. """Cancel the /forgetme action.
:param interaction: The interaction event. :param interaction: The interaction event.
:param button: The button that was pressed. :param _button: The button that was pressed.
""" """
self._responded = True self._responded = True
metrics.FORGETME.labels(outcome="cancelled").inc() metrics.FORGETME.labels(outcome="cancelled").inc()
@@ -145,10 +152,16 @@ class InteractionCog(commands.Cog):
:param interaction: The interaction event. :param interaction: The interaction event.
""" """
if await flags.is_flag_set( if await flags.is_flag_set(
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS self.bot.db,
interaction.user,
EntityType.USER,
Flag.ALLOW_PINGS,
): ):
await flags.clear_flag( await flags.clear_flag(
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS self.bot.db,
interaction.user,
EntityType.USER,
Flag.ALLOW_PINGS,
) )
metrics.PINGME.labels(outcome="opted_out").inc() metrics.PINGME.labels(outcome="opted_out").inc()
await interaction.response.send_message( await interaction.response.send_message(
@@ -159,7 +172,10 @@ class InteractionCog(commands.Cog):
) )
else: else:
await flags.set_flag( await flags.set_flag(
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS self.bot.db,
interaction.user,
EntityType.USER,
Flag.ALLOW_PINGS,
) )
metrics.PINGME.labels(outcome="opted_in").inc() metrics.PINGME.labels(outcome="opted_in").inc()
await interaction.response.send_message( await interaction.response.send_message(
@@ -179,7 +195,10 @@ class InteractionCog(commands.Cog):
:param interaction: The interaction event. :param interaction: The interaction event.
""" """
if await flags.is_flag_set( if await flags.is_flag_set(
self.bot.db, interaction.user, EntityType.USER, Flag.NO_INGEST self.bot.db,
interaction.user,
EntityType.USER,
Flag.NO_INGEST,
): ):
metrics.FORGETME.labels(outcome="already_forgotten").inc() metrics.FORGETME.labels(outcome="already_forgotten").inc()
await interaction.response.send_message( await interaction.response.send_message(
+10 -5
View File
@@ -49,7 +49,8 @@ class MessageCog(commands.Cog):
channel = message.channel channel = message.channel
if not isinstance( if not isinstance(
channel, (discord.TextChannel, discord.Thread, discord.VoiceChannel) channel,
(discord.TextChannel, discord.Thread, discord.VoiceChannel),
): ):
return return
@@ -61,7 +62,8 @@ class MessageCog(commands.Cog):
if self.bot.ingest_only: if self.bot.ingest_only:
# In ingest-only mode, never reply — only ingest eligible messages. # In ingest-only mode, never reply — only ingest eligible messages.
if message.type == discord.MessageType.default and not isinstance( if message.type == discord.MessageType.default and not isinstance(
channel, discord.Thread channel,
discord.Thread,
): ):
await ingest_message(self.bot.db, message, self.bot.ingest_cache) await ingest_message(self.bot.db, message, self.bot.ingest_cache)
return return
@@ -72,13 +74,15 @@ class MessageCog(commands.Cog):
# Only ingest non-thread messages; threads share their parent channel's chain. # Only ingest non-thread messages; threads share their parent channel's chain.
if message.type == discord.MessageType.default and not isinstance( if message.type == discord.MessageType.default and not isinstance(
channel, discord.Thread channel,
discord.Thread,
): ):
await ingest_message(self.bot.db, message, self.bot.ingest_cache) await ingest_message(self.bot.db, message, self.bot.ingest_cache)
@commands.Cog.listener() @commands.Cog.listener()
async def on_raw_message_delete( async def on_raw_message_delete(
self, payload: discord.RawMessageDeleteEvent self,
payload: discord.RawMessageDeleteEvent,
) -> None: ) -> None:
"""Uningest a recently ingested message when it is deleted. """Uningest a recently ingested message when it is deleted.
@@ -88,7 +92,8 @@ class MessageCog(commands.Cog):
@commands.Cog.listener() @commands.Cog.listener()
async def on_raw_bulk_message_delete( async def on_raw_bulk_message_delete(
self, payload: discord.RawBulkMessageDeleteEvent self,
payload: discord.RawBulkMessageDeleteEvent,
) -> None: ) -> None:
"""Uningest recently ingested messages when they are bulk-deleted. """Uningest recently ingested messages when they are bulk-deleted.
+8 -3
View File
@@ -55,7 +55,8 @@ class ServerEventsCog(commands.Cog):
logger.info('Joined new server! (Name: "%s" | ID: %s)', guild.name, guild.id) logger.info('Joined new server! (Name: "%s" | ID: %s)', guild.name, guild.id)
embed = discord.Embed( embed = discord.Embed(
title="Joined New Server", color=discord.Color.from_rgb(255, 165, 0) title="Joined New Server",
color=discord.Color.from_rgb(255, 165, 0),
) )
embed.add_field( embed.add_field(
name=f"{guild.name} ({guild.id})", name=f"{guild.name} ({guild.id})",
@@ -84,7 +85,9 @@ class ServerEventsCog(commands.Cog):
@commands.Cog.listener() @commands.Cog.listener()
async def on_guild_role_update( async def on_guild_role_update(
self, before: discord.Role, after: discord.Role self,
before: discord.Role,
after: discord.Role,
) -> None: ) -> None:
"""Queue all text channels for ingestion when role permissions change. """Queue all text channels for ingestion when role permissions change.
@@ -101,7 +104,9 @@ class ServerEventsCog(commands.Cog):
@commands.Cog.listener() @commands.Cog.listener()
async def on_guild_channel_update( async def on_guild_channel_update(
self, before: discord.abc.GuildChannel, after: discord.abc.GuildChannel self,
before: discord.abc.GuildChannel,
after: discord.abc.GuildChannel,
) -> None: ) -> None:
"""Queue all text channels for ingestion on permission change. """Queue all text channels for ingestion on permission change.
+16 -4
View File
@@ -73,6 +73,9 @@ def _split_sentences(paragraph: str) -> list[str]:
return _SENTENCE_SPLIT.split(normalized) return _SENTENCE_SPLIT.split(normalized)
_MIN_WORDS_FOR_START = 2
def _tokenize_sentence( def _tokenize_sentence(
sentence: str, sentence: str,
) -> tuple[list[str], list[tuple[str, str]]]: ) -> tuple[list[str], list[tuple[str, str]]]:
@@ -85,7 +88,7 @@ def _tokenize_sentence(
sentence += DEFAULT_SENTENCE_END sentence += DEFAULT_SENTENCE_END
words = sentence.split() words = sentence.split()
start_words: list[str] = [words[0]] if len(words) >= 2 else [] start_words: list[str] = [words[0]] if len(words) >= _MIN_WORDS_FOR_START else []
transitions: list[tuple[str, str]] = list(itertools.pairwise(words)) transitions: list[tuple[str, str]] = list(itertools.pairwise(words))
return start_words, transitions return start_words, transitions
@@ -106,7 +109,10 @@ async def ingest(db: Database, channel_id: int, user_id: int, paragraph: str) ->
async def _ingest_sentence( async def _ingest_sentence(
db: Database, channel_id: int, user_id: int, sentence: str db: Database,
channel_id: int,
user_id: int,
sentence: str,
) -> None: ) -> None:
"""Ingest a single sentence into the Markov chain for a given channel. """Ingest a single sentence into the Markov chain for a given channel.
@@ -137,7 +143,10 @@ async def uningest(db: Database, channel_id: int, user_id: int, paragraph: str)
async def _uningest_sentence( async def _uningest_sentence(
db: Database, channel_id: int, user_id: int, sentence: str db: Database,
channel_id: int,
user_id: int,
sentence: str,
) -> None: ) -> None:
"""Remove a single sentence's Markov data from the chain. """Remove a single sentence's Markov data from the chain.
@@ -153,7 +162,10 @@ async def _uningest_sentence(
async def generate( async def generate(
db: Database, channel_id: int, soft_limit: int = 750, hard_limit: int = 1000 db: Database,
channel_id: int,
soft_limit: int = 750,
hard_limit: int = 1000,
) -> str: ) -> str:
"""Generate a new sentence from previously ingested words for a channel. """Generate a new sentence from previously ingested words for a channel.
+32 -15
View File
@@ -79,7 +79,10 @@ async def reply_to_message(db: Database, message: discord.Message) -> None:
embed = discord.Embed( embed = discord.Embed(
title=await markov.generate(db, channel_id, soft_limit=200, hard_limit=300), title=await markov.generate(db, channel_id, soft_limit=200, hard_limit=300),
description=await markov.generate( description=await markov.generate(
db, channel_id, soft_limit=300, hard_limit=500 db,
channel_id,
soft_limit=300,
hard_limit=500,
), ),
) )
@@ -117,8 +120,30 @@ async def reply_to_message(db: Database, message: discord.Message) -> None:
metrics.REPLIES_SENT.inc() metrics.REPLIES_SENT.inc()
def _extract_embed_data(
embeds: list[discord.Embed],
) -> tuple[list[str], list[str]]:
"""Extract text and image URLs from message embeds.
:param embeds: The embeds to process.
:return: A tuple of (texts, image_urls).
"""
texts: list[str] = []
image_urls: list[str] = []
for embed in embeds:
if embed.title:
texts.append(embed.title)
if embed.description:
texts.append(embed.description)
if embed.image and embed.image.url:
image_urls.append(embed.image.url)
return texts, image_urls
async def ingest_message( async def ingest_message(
db: Database, message: discord.Message, cache: IngestCache | None = None db: Database,
message: discord.Message,
cache: IngestCache | None = None,
) -> None: ) -> None:
"""Ingest a given message into its channel's Markov chain. """Ingest a given message into its channel's Markov chain.
@@ -150,26 +175,18 @@ async def ingest_message(
if not message.content and not message.embeds: if not message.content and not message.embeds:
return return
embed_texts: list[str] = [] embed_texts, image_urls = _extract_embed_data(message.embeds)
image_urls: list[str] = []
with metrics.MESSAGE_INGESTION_DURATION.time(): with metrics.MESSAGE_INGESTION_DURATION.time():
if message.content: if message.content:
await markov.ingest(db, channel_id, user_id, message.content) await markov.ingest(db, channel_id, user_id, message.content)
for embed in message.embeds: for text in embed_texts:
if embed.title: await markov.ingest(db, channel_id, user_id, text)
embed_texts.append(embed.title)
await markov.ingest(db, channel_id, user_id, embed.title)
if embed.description:
embed_texts.append(embed.description)
await markov.ingest(db, channel_id, user_id, embed.description)
if embed.image and embed.image.url:
image_urls.append(embed.image.url)
if image_urls: if image_urls:
await db.add_images( await db.add_images(
[ChannelImage(channel_id, user_id, url) for url in image_urls] [ChannelImage(channel_id, user_id, url) for url in image_urls],
) )
if cache is not None: if cache is not None:
@@ -210,7 +227,7 @@ async def uningest_message(db: Database, cache: IngestCache, message_id: int) ->
[ [
ChannelImage(entry.channel_id, entry.user_id, url) ChannelImage(entry.channel_id, entry.user_id, url)
for url in entry.image_urls for url in entry.image_urls
] ],
) )
metrics.MESSAGES_UNINGESTED.inc() metrics.MESSAGES_UNINGESTED.inc()
+2 -1
View File
@@ -148,7 +148,8 @@ class MetricsServer:
async def start(self) -> None: async def start(self) -> None:
"""Create the aiohttp application and start listening.""" """Create the aiohttp application and start listening."""
if self._runner is not None: if self._runner is not None:
raise RuntimeError("Metrics server is already running") msg = "Metrics server is already running"
raise RuntimeError(msg)
app = web.Application() app = web.Application()
app.router.add_get("/metrics", make_aiohttp_handler()) app.router.add_get("/metrics", make_aiohttp_handler())
self._runner = web.AppRunner(app) self._runner = web.AppRunner(app)
+8 -48
View File
@@ -59,61 +59,21 @@ target-version = "py314"
extend-exclude = ["crabstero/_version.py"] # auto-generated by hatch-vcs extend-exclude = ["crabstero/_version.py"] # auto-generated by hatch-vcs
[tool.ruff.lint] [tool.ruff.lint]
select = [ select = ["ALL"]
# Core
"F", # Pyflakes
"E", # pycodestyle errors
"W", # pycodestyle warnings
"N", # pep8-naming
"D", # pydocstyle
"I", # isort
"ICN", # flake8-import-conventions
# Correctness & bugs
"B", # flake8-bugbear
"ASYNC", # flake8-async
"DTZ", # flake8-datetimez
"RSE", # flake8-raise
"RET", # flake8-return
"A", # flake8-builtins
"PIE", # flake8-pie
# Modernization & simplification
"UP", # pyupgrade
"SIM", # flake8-simplify
"C4", # flake8-comprehensions
"FLY", # flynt (f-string conversion)
"PTH", # flake8-use-pathlib
# Performance
"PERF", # Perflint
# Security
"S", # flake8-bandit
# Code hygiene
"T10", # flake8-debugger
"T20", # flake8-print
"ERA", # eradicate
"PGH", # pygrep-hooks
"TC", # flake8-type-checking
# Testing
"PT", # flake8-pytest-style
# Ruff-specific
"RUF", # Ruff-specific rules
]
ignore = [ ignore = [
"D203", # incompatible with D211 (no blank line before class docstring) "COM812", # handled by the formatter
"D213", # incompatible with D212 (summary on first line) "D203", # incompatible with D211 (no blank line before class docstring)
"D213", # incompatible with D212 (summary on first line)
] ]
[tool.pytest.ini_options] [tool.pytest.ini_options]
asyncio_mode = "auto" asyncio_mode = "auto"
[tool.ruff.lint.per-file-ignores] [tool.ruff.lint.per-file-ignores]
"tests/**" = ["S101"] # assert is standard for pytest "tests/**" = [
"S101", # assert is standard for pytest
"SLF001", # tests legitimately access private members for verification
]
[tool.coverage.run] [tool.coverage.run]
source = ["crabstero"] source = ["crabstero"]
+6 -2
View File
@@ -95,7 +95,9 @@ class TestReadCredential:
"""Systemd credential file reading.""" """Systemd credential file reading."""
def test_reads_credential_file( def test_reads_credential_file(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
"""Reads and strips the credential value from the file.""" """Reads and strips the credential value from the file."""
(tmp_path / "mytoken").write_text(" secret123 \n") (tmp_path / "mytoken").write_text(" secret123 \n")
@@ -108,7 +110,9 @@ class TestReadCredential:
assert _read_credential("anything") is None assert _read_credential("anything") is None
def test_returns_none_for_missing_file( def test_returns_none_for_missing_file(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
"""Returns None when the credential file does not exist.""" """Returns None when the credential file does not exist."""
monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path)) monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path))
+23 -15
View File
@@ -42,7 +42,7 @@ class TestConnect:
async def test_schema_creates_tables(self, db: Database) -> None: async def test_schema_creates_tables(self, db: Database) -> None:
"""All expected tables exist after connect.""" """All expected tables exist after connect."""
async with db._connection.execute( async with db._connection.execute(
"SELECT name FROM sqlite_master WHERE type = 'table' ORDER BY name" "SELECT name FROM sqlite_master WHERE type = 'table' ORDER BY name",
) as cursor: ) as cursor:
tables = [row[0] for row in await cursor.fetchall()] tables = [row[0] for row in await cursor.fetchall()]
assert tables == [ assert tables == [
@@ -128,7 +128,9 @@ class TestMarkovReadMethods:
], ],
) )
async def test_completing_word_filters_punctuation( async def test_completing_word_filters_punctuation(
self, db: Database, completing_word: str self,
db: Database,
completing_word: str,
) -> None: ) -> None:
"""get_random_completing_next_word only returns sentence-ending words.""" """get_random_completing_next_word only returns sentence-ending words."""
await db.add_markov_data( await db.add_markov_data(
@@ -151,7 +153,8 @@ class TestMarkovReadMethods:
async def test_start_word_pooled_across_users(self, db: Database) -> 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.""" """Start words from different users are visible in the same channel query."""
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 200, "Goodbye")], [] [StartWord(1, 100, "Hello"), StartWord(1, 200, "Goodbye")],
[],
) )
assert await db.get_random_start_word(1) in {"Hello", "Goodbye"} assert await db.get_random_start_word(1) in {"Hello", "Goodbye"}
@@ -187,7 +190,8 @@ class TestRemoveMarkovData:
async def test_removes_one_start_word(self, db: Database) -> None: async def test_removes_one_start_word(self, db: Database) -> None:
"""Removes exactly one matching start word row.""" """Removes exactly one matching start word row."""
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "Hello")], [] [StartWord(1, 100, "Hello"), StartWord(1, 100, "Hello")],
[],
) )
await db.remove_markov_data([StartWord(1, 100, "Hello")], []) await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
# One copy should remain. # One copy should remain.
@@ -227,7 +231,8 @@ class TestRemoveMarkovData:
async def test_start_word_removal_scoped_to_channel(self, db: Database) -> None: async def test_start_word_removal_scoped_to_channel(self, db: Database) -> None:
"""Removing a start word in one channel leaves another channel intact.""" """Removing a start word in one channel leaves another channel intact."""
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(2, 200, "Hello")], [] [StartWord(1, 100, "Hello"), StartWord(2, 200, "Hello")],
[],
) )
await db.remove_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 assert await db.get_random_start_word(1) is None
@@ -249,7 +254,8 @@ class TestRemoveMarkovData:
async def test_start_word_removal_scoped_to_user(self, db: Database) -> None: async def test_start_word_removal_scoped_to_user(self, db: Database) -> None:
"""Removing a start word for one user leaves another user.""" """Removing a start word for one user leaves another user."""
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 200, "Hello")], [] [StartWord(1, 100, "Hello"), StartWord(1, 200, "Hello")],
[],
) )
await db.remove_markov_data([StartWord(1, 100, "Hello")], []) await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) == "Hello" assert await db.get_random_start_word(1) == "Hello"
@@ -269,7 +275,8 @@ class TestRemoveMarkovData:
async def test_start_word_removal_scoped_to_word(self, db: Database) -> None: async def test_start_word_removal_scoped_to_word(self, db: Database) -> None:
"""Removing one start word leaves a different start word.""" """Removing one start word leaves a different start word."""
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "Hello"), StartWord(1, 100, "World")], [] [StartWord(1, 100, "Hello"), StartWord(1, 100, "World")],
[],
) )
await db.remove_markov_data([StartWord(1, 100, "Hello")], []) await db.remove_markov_data([StartWord(1, 100, "Hello")], [])
assert await db.get_random_start_word(1) == "World" assert await db.get_random_start_word(1) == "World"
@@ -334,7 +341,7 @@ class TestImages:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 200, "https://example.com/b.png"), ChannelImage(1, 200, "https://example.com/b.png"),
] ],
) )
assert await db.get_random_image(1) in { assert await db.get_random_image(1) in {
"https://example.com/a.png", "https://example.com/a.png",
@@ -351,7 +358,7 @@ class TestRemoveImages:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
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")]) await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
# One copy should remain. # One copy should remain.
@@ -369,7 +376,7 @@ class TestRemoveImages:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(2, 200, "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")]) 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(1) is None
@@ -381,7 +388,7 @@ class TestRemoveImages:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 200, "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")]) await db.remove_images([ChannelImage(1, 100, "https://example.com/a.png")])
assert await db.get_random_image(1) == "https://example.com/a.png" assert await db.get_random_image(1) == "https://example.com/a.png"
@@ -392,7 +399,7 @@ class TestRemoveImages:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(1, 100, "https://example.com/b.png"), ChannelImage(1, 100, "https://example.com/b.png"),
] ],
) )
await db.remove_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) == "https://example.com/b.png" assert await db.get_random_image(1) == "https://example.com/b.png"
@@ -515,7 +522,8 @@ class TestTransactionRollback:
" VALUES (?, ?, ?)", " VALUES (?, ?, ?)",
(1, 100, "should_not_persist"), (1, 100, "should_not_persist"),
) )
raise RuntimeError("simulated failure") msg = "simulated failure"
raise RuntimeError(msg)
class TestForgetUser: class TestForgetUser:
@@ -551,7 +559,7 @@ class TestForgetUser:
[ [
ChannelImage(1, 100, "https://example.com/gone.png"), ChannelImage(1, 100, "https://example.com/gone.png"),
ChannelImage(1, 200, "https://example.com/stay.png"), ChannelImage(1, 200, "https://example.com/stay.png"),
] ],
) )
await db.set_flag(EntityType.USER, "200", Flag.ALLOW_PINGS) await db.set_flag(EntityType.USER, "200", Flag.ALLOW_PINGS)
@@ -596,7 +604,7 @@ class TestForgetUser:
[ [
ChannelImage(1, 100, "https://example.com/a.png"), ChannelImage(1, 100, "https://example.com/a.png"),
ChannelImage(2, 100, "https://example.com/b.png"), ChannelImage(2, 100, "https://example.com/b.png"),
] ],
) )
await db.forget_user(100, Flag.NO_INGEST) await db.forget_user(100, Flag.NO_INGEST)
+8 -4
View File
@@ -53,7 +53,7 @@ class TestIsCompleteSentence:
pytest.param("Hello,", False, id="comma"), pytest.param("Hello,", False, id="comma"),
], ],
) )
def test_detection(self, sentence: str, expected: bool) -> None: def test_detection(self, sentence: str, expected: bool) -> None: # noqa: FBT001
"""Correctly identifies sentence completeness.""" """Correctly identifies sentence completeness."""
assert is_complete_sentence(sentence) is expected assert is_complete_sentence(sentence) is expected
@@ -147,7 +147,9 @@ class TestIngestParagraph:
], ],
) )
async def test_empty_input_stores_nothing( async def test_empty_input_stores_nothing(
self, db: Database, paragraph: str self,
db: Database,
paragraph: str,
) -> None: ) -> None:
"""Empty or whitespace-only input does not store any data.""" """Empty or whitespace-only input does not store any data."""
await ingest(db, 1, 100, paragraph) await ingest(db, 1, 100, paragraph)
@@ -260,13 +262,15 @@ class TestGenerate:
assert result == "A B C" assert result == "A B C"
async def test_soft_limit_falls_back_to_regular_next_word( async def test_soft_limit_falls_back_to_regular_next_word(
self, db: Database self,
db: Database,
) -> None: ) -> None:
"""After soft_limit, falls back when no completing word exists.""" """After soft_limit, falls back when no completing word exists."""
# "A" -> "B" (no completing transition). Past soft_limit, # "A" -> "B" (no completing transition). Past soft_limit,
# get_random_completing_next_word returns None, falls back to "B". # get_random_completing_next_word returns None, falls back to "B".
await db.add_markov_data( await db.add_markov_data(
[StartWord(1, 100, "A")], [Transition(1, 100, "A", "B")] [StartWord(1, 100, "A")],
[Transition(1, 100, "A", "B")],
) )
result = await generate(db, 1, soft_limit=1, hard_limit=1000) result = await generate(db, 1, soft_limit=1, hard_limit=1000)
+11 -6
View File
@@ -17,6 +17,7 @@
Tests cover metric object registration and the MetricsServer HTTP endpoint. Tests cover metric object registration and the MetricsServer HTTP endpoint.
""" """
from http import HTTPStatus
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import aiohttp import aiohttp
@@ -54,14 +55,17 @@ class TestMetricObjects:
id="channel-ingestion-duration", id="channel-ingestion-duration",
), ),
pytest.param( pytest.param(
"crabstero_generation_duration_seconds", id="generation-duration" "crabstero_generation_duration_seconds",
id="generation-duration",
), ),
pytest.param( pytest.param(
"crabstero_channel_ingestion_messages", id="channel-ingestion-messages" "crabstero_channel_ingestion_messages",
id="channel-ingestion-messages",
), ),
pytest.param("crabstero_discord_events_total", id="discord-events"), pytest.param("crabstero_discord_events_total", id="discord-events"),
pytest.param( pytest.param(
"crabstero_messages_uningested_total", id="messages-uningested" "crabstero_messages_uningested_total",
id="messages-uningested",
), ),
], ],
) )
@@ -89,16 +93,17 @@ class TestMetricsServer:
aiohttp.ClientSession() as session, aiohttp.ClientSession() as session,
session.get(f"http://127.0.0.1:{metrics_server.port}/metrics") as resp, session.get(f"http://127.0.0.1:{metrics_server.port}/metrics") as resp,
): ):
assert resp.status == 200 assert resp.status == HTTPStatus.OK
body = await resp.text() body = await resp.text()
assert "crabstero_build_info" in body assert "crabstero_build_info" in body
async def test_non_metrics_path_returns_404( async def test_non_metrics_path_returns_404(
self, metrics_server: MetricsServer self,
metrics_server: MetricsServer,
) -> None: ) -> None:
"""GET on an unknown path returns 404.""" """GET on an unknown path returns 404."""
async with ( async with (
aiohttp.ClientSession() as session, aiohttp.ClientSession() as session,
session.get(f"http://127.0.0.1:{metrics_server.port}/notfound") as resp, session.get(f"http://127.0.0.1:{metrics_server.port}/notfound") as resp,
): ):
assert resp.status == 404 assert resp.status == HTTPStatus.NOT_FOUND
+3 -1
View File
@@ -133,7 +133,9 @@ class TestUningestRestoresState:
], ],
) )
async def test_uningest_restores_preexisting_data( async def test_uningest_restores_preexisting_data(
self, db: Database, text: str self,
db: Database,
text: str,
) -> None: ) -> None:
"""Ingest then uningest preserves unrelated pre-existing data exactly.""" """Ingest then uningest preserves unrelated pre-existing data exactly."""
await ingest(db, 99, 200, "Pre-existing data stays safe.") await ingest(db, 99, 200, "Pre-existing data stays safe.")