Experimental Discord bot written in Python
You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276
  1. """
  2. General utility functions.
  3. """
  4. import re
  5. import sys
  6. import traceback
  7. from datetime import datetime, timedelta, timezone
  8. from typing import Any
  9. import discord
  10. from discord import Guild, Interaction, Permissions
  11. from discord.app_commands import Transformer
  12. from discord.ext.commands import BadArgument, Cog
  13. def dump_stacktrace(e: BaseException) -> None:
  14. print(e, file=sys.stderr)
  15. traceback.print_exception(type(e), e, e.__traceback__)
  16. def timedelta_from_str(s: str) -> timedelta:
  17. """
  18. Parses a timespan.
  19. Format examples:
  20. "30m"
  21. "10s"
  22. "90d"
  23. "1h30m"
  24. "73d18h22m52s"
  25. Parameters
  26. ----------
  27. s : str
  28. string to parse
  29. Returns
  30. -------
  31. timedelta
  32. Raises
  33. ------
  34. ValueError
  35. if parsing fails
  36. """
  37. p: re.Pattern = re.compile('^(?:[0-9]+[a-zA-Z])+$')
  38. if p.match(s) is None:
  39. raise ValueError(f'Illegal timespan value "{s}". Examples: 30s, 5m, 1h30m, 30d')
  40. p = re.compile('([0-9]+)([dhms])')
  41. days: int = 0
  42. hours: int = 0
  43. minutes: int = 0
  44. seconds: int = 0
  45. for m in p.finditer(s):
  46. scalar = int(m.group(1))
  47. unit = m.group(2).lower()
  48. if unit == 'd':
  49. days = scalar
  50. elif unit == 'h':
  51. hours = scalar
  52. elif unit == 'm':
  53. minutes = scalar
  54. elif unit == 's':
  55. seconds = scalar
  56. else:
  57. raise ValueError(f'Invalid unit "{unit}". Valid units: "s"=seconds, "m"=minutes, "h"=hours, "d"=days')
  58. return timedelta(days=days, hours=hours, minutes=minutes, seconds=seconds)
  59. def str_from_timedelta(td: timedelta) -> str:
  60. """
  61. Encodes a timedelta as a str. E.g. "3d2h"
  62. """
  63. d: int = td.days
  64. h: int = td.seconds // 3600
  65. m: int = (td.seconds // 60) % 60
  66. s: int = td.seconds % 60
  67. components: list[str] = []
  68. if d != 0:
  69. components.append(f'{d}d')
  70. if h != 0:
  71. components.append(f'{h}h')
  72. if m != 0:
  73. components.append(f'{m}m')
  74. if s != 0 or len(components) == 0:
  75. components.append(f'{s}s')
  76. return ''.join(components)
  77. def describe_timedelta(td: timedelta, max_components: int = 2) -> str:
  78. """
  79. Formats a human-readable description of a time span. E.g. "3 days 2 hours".
  80. """
  81. d: int = td.days
  82. h: int = td.seconds // 3600
  83. m: int = (td.seconds // 60) % 60
  84. s: int = td.seconds % 60
  85. components: list[str] = []
  86. if d != 0:
  87. components.append('1 day' if d == 1 else f'{d} days')
  88. if h != 0:
  89. components.append('1 hour' if h == 1 else f'{h} hours')
  90. if m != 0:
  91. components.append('1 minute' if m == 1 else f'{m} minutes')
  92. if s != 0 or len(components) == 0:
  93. components.append('1 second' if s == 1 else f'{s} seconds')
  94. if len(components) > max_components:
  95. components = components[0:max_components]
  96. return ' '.join(components)
  97. def _old_first_command_group(cog: Cog) -> discord.ext.commands.Group | None:
  98. """Returns the first command Group found in a cog."""
  99. for member_name in dir(cog):
  100. member = getattr(cog, member_name)
  101. if isinstance(member, discord.ext.commands.Group):
  102. return member
  103. return None
  104. def first_command_group(cog: Cog) -> discord.app_commands.Group | None:
  105. """Returns the first slash command Group found in a cog."""
  106. for member_name in dir(cog):
  107. member = getattr(cog, member_name)
  108. if isinstance(member, discord.app_commands.Group):
  109. return member
  110. return None
  111. def bot_log(guild: Guild | None, cog_class: type | None, message: Any) -> None:
  112. """Logs a message to stdout with time, cog, and guild info."""
  113. now: datetime = datetime.now(tz=None) # noqa: DTZ005
  114. s = f'[{now.strftime("%Y-%m-%dT%H:%M:%S")}|'
  115. s += f'{cog_class.__name__}|' if cog_class else '-|'
  116. s += f'{guild.name}] ' if guild else '-] '
  117. s += str(message)
  118. print(s)
  119. __QUOTE_CHARS: str = '\'"'
  120. __ID_REGEX: re.Pattern = re.compile('^[0-9]{17,20}$')
  121. __MENTION_REGEX: re.Pattern = re.compile('^<@[!&]([0-9]{17,20})>$')
  122. __USER_MENTION_REGEX: re.Pattern = re.compile('^<@!([0-9]{17,20})>$')
  123. __ROLE_MENTION_REGEX: re.Pattern = re.compile('^<@&([0-9]{17,20})>$')
  124. __EMAIL_REGEX: re.Pattern = re.compile(r'^(?:(?:[^<>()\[\]\\.,;:\s@"]+(?:\.[^<>()\[\]\\.,;:\s@"]+)*)|(?:".+"))@(?:(?:\[[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}])|(?:(?:[a-zA-Z\-0-9]+\.)+[a-zA-Z]{2,}))$')
  125. __USERNAME_REGEX: re.Pattern = re.compile(r'[a-z0-9\._]{2,32}')
  126. def is_user_id(val: str) -> bool:
  127. """Tests if a string is in user/role ID format."""
  128. return __ID_REGEX.match(val) is not None
  129. def is_mention(val: str) -> bool:
  130. """Tests if a string is a user or role mention."""
  131. return __MENTION_REGEX.match(val) is not None
  132. def is_role_mention(val: str) -> bool:
  133. """Tests if a string is a role mention."""
  134. return __ROLE_MENTION_REGEX.match(val) is not None
  135. def is_user_mention(val: str) -> bool:
  136. """Tests if a string is a user mention."""
  137. return __USER_MENTION_REGEX.match(val) is not None
  138. def is_email_address(val: str) -> bool:
  139. """Tests if a string is a well-formed email address."""
  140. return __EMAIL_REGEX.match(val) is not None
  141. def is_discord_username(val: str) -> bool:
  142. """Tests if a string is a properly formatted Discord username."""
  143. return __USERNAME_REGEX.match(val.lower())
  144. def user_id_from_mention(mention: str) -> str:
  145. """Extracts the user ID from a mention. Raises a ValueError if malformed."""
  146. m = __USER_MENTION_REGEX.match(mention)
  147. if m:
  148. return m.group(1)
  149. raise ValueError(f'"{mention}" is not an @ user mention')
  150. def mention_from_user_id(user_id: str | int) -> str:
  151. """Returns a Markdown user mention from a user id."""
  152. return f'<@!{user_id}>'
  153. def mention_from_role_id(role_id: str | int) -> str:
  154. """Returns a Markdown role mention from a role id."""
  155. return f'<@&{role_id}>'
  156. def str_from_quoted_str(val: str) -> str:
  157. """Removes the leading and trailing quotes from a string."""
  158. if len(val) < 2 or val[0:1] not in __QUOTE_CHARS or val[-1:] not in __QUOTE_CHARS:
  159. raise ValueError(f'Not a quoted string: {val}')
  160. return val[1:-1]
  161. def blockquote_markdown(markdown: str) -> str:
  162. """Encloses some Markdown in a blockquote."""
  163. return '> ' + (markdown.replace('\n', '\n> '))
  164. def indent_markdown(markdown: str) -> str:
  165. """Indents a block of Markdown by one level."""
  166. return ' ' + (markdown.replace('\n', '\n '))
  167. def suppress_markdown_url_previews(markdown: str) -> str:
  168. """Finds URLs in markdown and encloses them in <...> to suppress the preview."""
  169. return re.sub(r'(?<!<)(https?://\S+)(?!>)', '<\\1>', markdown)
  170. def truncate_markdown(markdown: str, max_length: int) -> str:
  171. """Truncates markdown in a way that attempts to minimize formatting disruption."""
  172. if len(markdown) <= max_length:
  173. return markdown
  174. markdown = markdown[:max_length]
  175. # Try to cut at a newline if it's in the latter 20% of the max length
  176. last_newline_index = markdown.rfind('\n')
  177. if last_newline_index >= 0 and last_newline_index < (max_length * 8 / 10):
  178. return markdown[:last_newline_index] + "\n\u2026"
  179. # Cut at the last space
  180. last_space_index = markdown.rfind(' ')
  181. if last_space_index >= 0:
  182. return markdown[:last_space_index] + " \u2026"
  183. # Last resort, do a blind substring
  184. return markdown[:max_length - 1] + "\u2026"
  185. def format_bytes(size: int) -> str:
  186. """Formats s size in bytes to a human readable description (e.g. "3.2 KiB")"""
  187. size = max(size, 0)
  188. kib = 1024
  189. mib = kib * kib
  190. gib = mib * kib
  191. if size < kib:
  192. return f"{size:,} bytes"
  193. if size < 10 * kib:
  194. return f"{size/kib:,.1f} KiB"
  195. if size < mib:
  196. return f"{size/kib:,.0f} KiB"
  197. if size < 10 * mib:
  198. return f"{size/mib:,.1f} MiB"
  199. if size < gib:
  200. return f"{size/mib:,.0f} MiB"
  201. if size < 10 * gib:
  202. return f"{size/gib:,.1f} GiB"
  203. return f"{size/gib:,.0f} GiB"
  204. def norm_datetime(dt: datetime) -> datetime:
  205. """Converts a datetime to UTC for consistent comparison."""
  206. # "Naive" datetimes (without a time zone) are assumed as system local time zone
  207. if dt is None:
  208. return dt
  209. return datetime.fromtimestamp(dt.timestamp(), timezone.utc)
  210. def levenshtein(a: str, b: str) -> int:
  211. """Returns the Levenshtein distance between two strings."""
  212. # Based on https://en.wikipedia.org/wiki/Levenshtein_distance#Iterative_with_two_matrix_rows
  213. m: int = len(a)
  214. n: int = len(b)
  215. v0: list[int] = [ i for i in range(n + 1) ]
  216. v1: list[int] = [ 0 for i in range(n + 1) ]
  217. for i in range(m):
  218. v1[0] = i + 1
  219. for j in range(n):
  220. deletion_cost = v0[j + 1] + 1
  221. insertion_cost = v1[j] + 1
  222. substitution_cost = v0[j] + (0 if a[i] == b[j] else 1)
  223. v1[j + 1] = min(deletion_cost, insertion_cost, substitution_cost)
  224. h = v0
  225. v0 = v1
  226. v1 = h
  227. return v0[n]
  228. MOD_PERMISSIONS: Permissions = Permissions(Permissions.manage_messages.flag)
  229. ADMIN_PERMISSIONS: Permissions = Permissions(Permissions.administrator.flag)
  230. class TimeDeltaTransformer(Transformer):
  231. async def transform(self, interaction: Interaction, value: Any) -> timedelta:
  232. try:
  233. return timedelta_from_str(str(value))
  234. except ValueError as e:
  235. print("Invalid time delta:", e)
  236. raise BadArgument(str(e))
  237. @property
  238. def _error_display_name(self) -> str:
  239. return 'timedelta'