Enabled all Ruff linter rules and fixed resulting violations.
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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.")
|
||||||
|
|||||||
Reference in New Issue
Block a user