Replaced regex with native str methods for whitespace normalization and mention extraction.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s

This commit is contained in:
2026-03-23 14:39:38 -04:00
parent 0fa51d6199
commit 70095f3184
2 changed files with 6 additions and 9 deletions
+4 -5
View File
@@ -34,17 +34,16 @@ if TYPE_CHECKING:
DEFAULT_SENTENCE_END = "\u00a7" DEFAULT_SENTENCE_END = "\u00a7"
_TERMINATORS = frozenset({DEFAULT_SENTENCE_END, ".", "!", "?"}) _TERMINATORS = frozenset({DEFAULT_SENTENCE_END, ".", "!", "?"})
_MULTI_SPACE = re.compile(r" +")
_SENTENCE_SPLIT = re.compile(r"(?<=[.!?]) ") _SENTENCE_SPLIT = re.compile(r"(?<=[.!?]) ")
def _normalize_whitespace(text: str) -> str: def _normalize_whitespace(text: str) -> str:
"""Collapse runs of spaces into single spaces and strip edges. """Collapse runs of whitespace into single spaces and strip edges.
:param text: The raw text to normalize. :param text: The raw text to normalize.
:return: The normalized text. :return: The normalized text.
""" """
return _MULTI_SPACE.sub(" ", text.strip()) return " ".join(text.split())
def is_complete_sentence(sentence: str) -> bool: def is_complete_sentence(sentence: str) -> bool:
@@ -70,7 +69,7 @@ def _split_sentences(paragraph: str) -> list[str]:
""" """
if not is_complete_sentence(paragraph): if not is_complete_sentence(paragraph):
paragraph += DEFAULT_SENTENCE_END paragraph += DEFAULT_SENTENCE_END
normalized = _normalize_whitespace(paragraph.replace("\n", " ")) normalized = _normalize_whitespace(paragraph)
return _SENTENCE_SPLIT.split(normalized) return _SENTENCE_SPLIT.split(normalized)
@@ -84,7 +83,7 @@ def _tokenize_sentence(
""" """
if not is_complete_sentence(sentence): if not is_complete_sentence(sentence):
sentence += DEFAULT_SENTENCE_END sentence += DEFAULT_SENTENCE_END
words = _normalize_whitespace(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) >= 2 else []
transitions: list[tuple[str, str]] = list(itertools.pairwise(words)) transitions: list[tuple[str, str]] = list(itertools.pairwise(words))
+2 -4
View File
@@ -35,7 +35,7 @@ if TYPE_CHECKING:
# Copied from discordjs/discord-api-types: # Copied from discordjs/discord-api-types:
# https://github.com/discordjs/discord-api-types/blob/662cb0cb0ac9c6f9ad93e180849476714bfceb0c/globals.ts#L39 # https://github.com/discordjs/discord-api-types/blob/662cb0cb0ac9c6f9ad93e180849476714bfceb0c/globals.ts#L39
_MENTION_PATTERN = re.compile(r"<@!?(?P<id>\d{17,20})>") _MENTION_PATTERN = re.compile(r"<@!?(\d{17,20})>")
_EMBED_CHANCE_THRESHOLD = 95 # Out of 100; sends an embed ~5% of the time. _EMBED_CHANCE_THRESHOLD = 95 # Out of 100; sends an embed ~5% of the time.
@@ -90,9 +90,7 @@ async def reply_to_message(db: Database, message: discord.Message) -> None:
# Suppress all mentions by default; only ping users who opted in. # Suppress all mentions by default; only ping users who opted in.
allowed_user_ids: set[int] = set() allowed_user_ids: set[int] = set()
for match in _MENTION_PATTERN.finditer(body): for user_id in {int(uid) for uid in _MENTION_PATTERN.findall(body)}:
user_id = int(match.group("id"))
if await flags.is_flag_set(db, user_id, EntityType.USER, Flag.ALLOW_PINGS): if await flags.is_flag_set(db, user_id, EntityType.USER, Flag.ALLOW_PINGS):
allowed_user_ids.add(user_id) allowed_user_ids.add(user_id)