diff --git a/changelog/1522.feature.rst b/changelog/1522.feature.rst new file mode 100644 index 0000000000..18b54c793d --- /dev/null +++ b/changelog/1522.feature.rst @@ -0,0 +1 @@ +Add :meth:`Guild.search_messages`, which allows searching for messages in a guild or channels using several different query parameters. diff --git a/disnake/enums.py b/disnake/enums.py index f5527bcd35..087fd2ca90 100644 --- a/disnake/enums.py +++ b/disnake/enums.py @@ -75,6 +75,7 @@ "MessageReferenceType", "SeparatorSpacing", "NameplatePalette", + "MessageSearchSortMode", ) EnumMetaT = TypeVar("EnumMetaT", bound="EnumMeta") @@ -2499,6 +2500,40 @@ class NameplatePalette(Enum): """White color palette.""" +class MessageSearchSortMode(Enum): + """Represents the sorting algorithm/direction used for :meth:`Guild.search_messages`. + + .. versionadded:: |vnext| + """ + + timestamp_desc = "timestamp_desc" + """Sort by message creation time, descending.""" + timestamp_asc = "timestamp_asc" + """Sort by message creation time, ascending.""" + relevance = "relevance" + """Sort by relevance of the message to the search query.""" + + @property + def sort_key(self) -> str: + match self: + case MessageSearchSortMode.timestamp_desc: + return "timestamp" + case MessageSearchSortMode.timestamp_asc: + return "timestamp" + case MessageSearchSortMode.relevance: + return "relevance" + + @property + def sort_order(self) -> str | None: + match self: + case MessageSearchSortMode.timestamp_desc: + return "desc" + case MessageSearchSortMode.timestamp_asc: + return "asc" + case MessageSearchSortMode.relevance: + return None + + T = TypeVar("T", bound="Enum") diff --git a/disnake/errors.py b/disnake/errors.py index 69bd9d5c16..40e5b5ec4f 100644 --- a/disnake/errors.py +++ b/disnake/errors.py @@ -10,6 +10,7 @@ from requests import Response from .client import SessionStartLimit + from .guild import Guild from .interactions import Interaction, ModalInteraction _ResponseType: TypeAlias = ClientResponse | Response @@ -36,6 +37,7 @@ "ModalChainNotSupported", "InteractionNotEditable", "LocalizationKeyError", + "MessageSearchIndexUnavailableError", ) @@ -433,3 +435,20 @@ class LocalizationKeyError(DiscordException): def __init__(self, key: str) -> None: self.key: str = key super().__init__(f"No localizations were found for the key '{key}'.") + + +class MessageSearchIndexUnavailableError(DiscordException): + """Exception that's raised when the message search target is not yet + indexed, and all retries (if configured) have been exhausted. + + .. versionadded:: |vnext| + + Attributes + ---------- + guild: :class:`Guild` + The guild whose messages are not yet indexed. + """ + + def __init__(self, guild: Guild) -> None: + self.guild: Guild = guild + super().__init__(f"Message search indexing for guild ID {guild.id} is still in progress.") diff --git a/disnake/guild.py b/disnake/guild.py index bc5faf324f..e4a5eabb71 100644 --- a/disnake/guild.py +++ b/disnake/guild.py @@ -43,6 +43,7 @@ GuildScheduledEventEntityType, GuildScheduledEventPrivacyLevel, Locale, + MessageSearchSortMode, NotificationLevel, NSFWLevel, ThreadLayout, @@ -59,7 +60,7 @@ from .guild_scheduled_event import GuildScheduledEvent, GuildScheduledEventMetadata from .integrations import Integration, _integration_factory from .invite import Invite -from .iterators import AuditLogIterator, BanIterator, MemberIterator +from .iterators import AuditLogIterator, BanIterator, MemberIterator, MessageSearchIterator from .member import Member, VoiceState from .mixins import Hashable from .object import Object @@ -77,6 +78,9 @@ from .widget import Widget, WidgetSettings __all__ = ( + "MessageSearchAuthorType", + "MessageSearchHasType", + "MessageSearchEmbedType", "IncidentsData", "Guild", ) @@ -102,6 +106,7 @@ MFALevel, ) from .types.integration import Integration as IntegrationPayload, IntegrationType + from .types.message import MessageSearchQuery from .types.role import CreateRole as CreateRolePayload from .types.sticker import CreateGuildSticker as CreateStickerPayload from .types.threads import Thread as ThreadPayload, ThreadArchiveDurationLiteral @@ -116,6 +121,30 @@ ByCategoryItem: TypeAlias = tuple[CategoryChannel | None, list[GuildChannel]] +# These literals are here such that they can (in theory) be used at runtime; +# disnake.types isn't necessarily runtime-importable due to cycles + +# fmt: off +MessageSearchAuthorType = Literal[ + "user", "-user", + "bot", "-bot", + "webhook", "-webhook" +] +MessageSearchHasType = Literal[ + "image", "-image", + "sound", "-sound", + "video", "-video", + "file", "-file", + "sticker", "-sticker", + "embed", "-embed", + "link", "-link", + "poll", "-poll", + "snapshot", "-snapshot", +] +# fmt: on +MessageSearchEmbedType = Literal["image", "video", "gif", "sound", "article"] + + class _GuildLimit(NamedTuple): emoji: int stickers: int @@ -5344,3 +5373,180 @@ async def fetch_soundboard_sounds(self) -> list[GuildSoundboardSound]: return [ GuildSoundboardSound(data=d, state=self._state, guild_id=self.id) for d in data["items"] ] + + def search_messages( + self, + *, + # common iterator params + limit: int | None = 25, + before: SnowflakeTime | None = None, + after: SnowflakeTime | None = None, + sort: MessageSearchSortMode = MessageSearchSortMode.timestamp_desc, + # search filters + content: str | None = None, + slop: int | None = None, + channel: Sequence[Snowflake] | Snowflake | None = None, + author: Sequence[Snowflake] | Snowflake | None = None, + author_type: Sequence[MessageSearchAuthorType] | MessageSearchAuthorType | None = None, + mentions: Sequence[Snowflake] | Snowflake | None = None, + mentions_role: Sequence[Snowflake] | Snowflake | None = None, + mentions_everyone: bool | None = None, + replied_to_user: Sequence[Snowflake] | Snowflake | None = None, + replied_to_message: Sequence[Snowflake] | Snowflake | None = None, + pinned: bool | None = None, + has: Sequence[MessageSearchHasType] | MessageSearchHasType | None = None, + embed_type: Sequence[MessageSearchEmbedType] | MessageSearchEmbedType | None = None, + embed_provider: Sequence[str] | str | None = None, + link_hostname: Sequence[str] | str | None = None, + attachment_filename: Sequence[str] | str | None = None, + attachment_extension: Sequence[str] | str | None = None, + include_nsfw: bool = False, + # for handling indexing errors + retries: int = 3, + ) -> MessageSearchIterator: + r"""Returns an :class:`.AsyncIterator` representing the messages matching the query parameters. + + Results are returned from newest to oldest by default; this is configurable using + the ``sort`` parameter. + + You must have :attr:`~Permissions.read_message_history` permissions to do this, + and the :attr:`~Intents.message_content` intent must be enabled for this bot. + + .. versionadded:: |vnext| + + Parameters + ---------- + limit: :class:`int` | :data:`None` + The number of messages to retrieve, up to 10000. + If :data:`None`, retrieves the maximum number of matching messages. + Note, however, that this would make it a slow operation. + Defaults to ``25``. + before: :class:`.abc.Snowflake` | :class:`datetime.datetime` | :data:`None` + Retrieves messages created before this date or object. + If a datetime is provided, it is recommended to use a UTC aware datetime. + If the datetime is naive, it is assumed to be local time. + after: :class:`.abc.Snowflake` | :class:`datetime.datetime` | :data:`None` + Retrieve messages created after this date or object. + If a datetime is provided, it is recommended to use a UTC aware datetime. + If the datetime is naive, it is assumed to be local time. + sort: :class:`MessageSearchSortMode` + The sorting algorithm/direction to use for retrieving search results. + Defaults to :attr:`~MessageSearchSortMode.timestamp_desc`. + content: :class:`str` | :data:`None` + Filter messages by content (up to 1024 characters). + channel: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages by channels (up to 500). + author: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages by authors (up to 100). + author_type: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by author types. + + Can be any subset of ``["user", "bot", "webhook"]``. Types can also be negated with a + ``-`` prefix to exclude that type, e.g. ``["bot", "-webhook"]`` would be a valid value. + mentions: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages that mention these users (up to 100). + mentions_role: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages that mention these roles (up to 100). + mentions_everyone: :class:`bool` | :data:`None` + Filter messages that do/don't mention ``@everyone``. + replied_to_user: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages that reply to these users (up to 100). + replied_to_message: :class:`~collections.abc.Sequence`\[:class:`.abc.Snowflake`] | :class:`.abc.Snowflake` | :data:`None` + Filter messages that reply to these messages (up to 100). + pinned: :class:`bool` | :data:`None` + Filter messages that are/aren't pinned. + has: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by whether or not they have specific things. + + Can be any subset of ``["image", "sound", "video", "file", "sticker", "embed", "link", "poll", "snapshot"]``. + Types can also be negated with a ``-`` prefix to exclude that type, + e.g. ``["image", "-link"]`` would be a valid value. + embed_type: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by embed type. + + Can be any subset of ``["image", "video", "gif", "sound", "article"]``. + embed_provider: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by embed provider (up to 100, with up to 256 characters each). + link_hostname: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by link hostname, e.g. ``discordapp.com`` (up to 100, with up to 256 characters each). + attachment_filename: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by attachment filename (up too 100, with up to 1024 characters each). + attachment_extension: :class:`~collections.abc.Sequence`\[:class:`str`] | :class:`str` | :data:`None` + Filter messages by attachment extension, e.g. ``txt`` (up too 100, with up to 256 characters each). + include_nsfw: :class:`bool` + Whether to include results from age-restricted channels. Defaults to ``False``. + retries: :class:`int` + The number of times to wait and retry fetching results in case the guild is still being indexed. + Can be set to 0 to disable retries and raise an error immediately instead of retrying. + Defaults to 3. + + Raises + ------ + Forbidden + You do not have permission to search messages, + or the :attr:`~Intents.message_content` intent is not enabled. + HTTPException + Retrieving the search results failed. + MessageSearchIndexUnavailableError + Exceeded maximum number of retries while waiting for messages to finish indexing. + + Yields + ------ + :class:`.Message` + The message matching the given query parameters. + """ + + def listify_snowflakes(arg: abc.Snowflake | Sequence[abc.Snowflake]) -> Sequence[int]: + if isinstance(arg, abc.Snowflake): + return [arg.id] + return [item.id for item in arg] + + def listify_strs(arg: str | Sequence[str]) -> Sequence[str]: + if isinstance(arg, str): + return [arg] + return arg + + query: MessageSearchQuery = {"include_nsfw": include_nsfw} + + query["sort_by"] = sort.sort_key + if sort_order := sort.sort_order: + query["sort_order"] = sort_order + + if content is not None: + query["content"] = content + if slop is not None: + query["slop"] = slop + if channel is not None: + query["channel_id"] = listify_snowflakes(channel) + if author is not None: + query["author_id"] = listify_snowflakes(author) + if author_type is not None: + query["author_type"] = listify_strs(author_type) + if mentions is not None: + query["mentions"] = listify_snowflakes(mentions) + if mentions_role is not None: + query["mentions_role"] = listify_snowflakes(mentions_role) + if mentions_everyone is not None: + query["mentions_everyone"] = mentions_everyone + if replied_to_user is not None: + query["replied_to_user_id"] = listify_snowflakes(replied_to_user) + if replied_to_message is not None: + query["replied_to_message_id"] = listify_snowflakes(replied_to_message) + if pinned is not None: + query["pinned"] = pinned + if has is not None: + query["has"] = listify_strs(has) + if embed_type is not None: + query["embed_type"] = listify_strs(embed_type) + if embed_provider is not None: + query["embed_provider"] = listify_strs(embed_provider) + if link_hostname is not None: + query["link_hostname"] = listify_strs(link_hostname) + if attachment_filename is not None: + query["attachment_filename"] = listify_strs(attachment_filename) + if attachment_extension is not None: + query["attachment_extension"] = listify_strs(attachment_extension) + + return MessageSearchIterator( + self, query, retries=retries, limit=limit, before=before, after=after + ) diff --git a/disnake/http.py b/disnake/http.py index dad432fdc2..d9a6a4318e 100644 --- a/disnake/http.py +++ b/disnake/http.py @@ -7,7 +7,7 @@ import re import sys import weakref -from collections.abc import Coroutine, Iterable, Sequence +from collections.abc import Coroutine, Iterable, Mapping, Sequence from errno import ECONNRESET from typing import ( TYPE_CHECKING, @@ -909,6 +909,15 @@ def get_pins( return self.request(r, params=params) + def search_guild_messages( + self, guild_id: Snowflake, params: Mapping[str, Any] + ) -> Response[message.MessageSearchResult | message.MessageSearchNotIndexedResult]: + # turn bools into 0/1 + params = {k: (int(v) if isinstance(v, bool) else v) for k, v in params.items()} + + r = Route("GET", "/guilds/{guild_id}/messages/search", guild_id=guild_id) + return self.request(r, params=params) + # Member management def search_guild_members( diff --git a/disnake/iterators.py b/disnake/iterators.py index 0203c86f9a..d60e4bfb2b 100644 --- a/disnake/iterators.py +++ b/disnake/iterators.py @@ -4,6 +4,7 @@ import asyncio import datetime +import logging from collections.abc import AsyncIterator, Awaitable, Callable, Generator from typing import ( TYPE_CHECKING, @@ -20,7 +21,7 @@ from .automod import AutoModRule from .bans import BanEntry from .entitlement import Entitlement -from .errors import NoMoreItems +from .errors import MessageSearchIndexUnavailableError, NoMoreItems from .guild_scheduled_event import GuildScheduledEvent from .integrations import PartialIntegration from .object import Object @@ -39,6 +40,7 @@ "EntitlementIterator", "SubscriptionIterator", "PollAnswerIterator", + "MessageSearchIterator", ) if TYPE_CHECKING: @@ -60,7 +62,7 @@ GuildScheduledEventUser as GuildScheduledEventUserPayload, ) from .types.member import MemberWithUser as MemberWithUserPayload - from .types.message import Message as MessagePayload + from .types.message import Message as MessagePayload, MessageSearchQuery, MessageSearchResult from .types.subscription import Subscription as SubscriptionPayload from .types.threads import Thread as ThreadPayload from .types.user import PartialUser as PartialUserPayload @@ -72,6 +74,8 @@ OLDEST_OBJECT = Object(id=0) +_log = logging.getLogger(__name__) + class _AsyncIterator(AsyncIterator[T]): __slots__ = () @@ -1402,3 +1406,126 @@ async def fill_messages(self) -> None: message = self._state.create_message(channel=self.channel, data=element["message"]) message._pinned_at = parse_time(element["pinned_at"]) await self.messages.put(message) + + +class MessageSearchIterator(_AsyncIterator["Message"]): + def __init__( + self, + guild: Guild, + query: MessageSearchQuery, + *, + retries: int, + limit: int | None, + before: Snowflake | datetime.datetime | None = None, + after: Snowflake | datetime.datetime | None = None, + ) -> None: + if isinstance(before, datetime.datetime): + before = Object(id=time_snowflake(before, high=False)) + if isinstance(after, datetime.datetime): + after = Object(id=time_snowflake(after, high=True)) + + self.guild = guild + self._state = guild._state + + self.limit = limit + self.offset: int = 0 + + self.query = query + # since the given `query` dict should always be ephemeral and only created by the lib, + # we can just mutate it directly + if before is not None: + self.query["max_id"] = before.id + if after is not None: + self.query["min_id"] = after.id + + self.max_retries = retries + self.getter = self._state.http.search_guild_messages + self.messages: asyncio.Queue[Message] = asyncio.Queue() + + async def next(self) -> Message: + # note: unlike other endpoints, this one can return empty pages, + # especially with higher offsets. therefore, continue iterating empty pages + # until we either get some results or reach the definitive end + while self.messages.empty() and self.limit != 0: + await self.fill_messages() + + try: + return self.messages.get_nowait() + except asyncio.QueueEmpty: + raise NoMoreItems from None + + def _get_retrieve(self) -> bool: + self.retrieve = min(self.limit, 25) if self.limit is not None else 25 + return self.retrieve > 0 + + # this endpoint is somewhat special in that the guild/channel may still be in the + # process of being indexed, in which case we receive a `202 Accepted` with a `retry_after` field. + async def _try_fetch(self) -> MessageSearchResult: + retries = 0 + while True: + data = await self.getter( + guild_id=self.guild.id, + params=self.query, + ) + + if "code" not in data: + # success + return data + + if data["code"] != 110000: + msg = f"Received unexpected error code {data['code']}" + raise RuntimeError(msg) + + # if we have `"code": 110000`, this was a 202 response and message indexing is likely still in progress + if retries >= self.max_retries: + raise MessageSearchIndexUnavailableError(self.guild) + + retry_after = data["retry_after"] + # "If the retry_after field is 0, you should retry the request after a short delay." + retry_after = max(retry_after, 1) + _log.info( + "Message search index for guild ID %d is not yet available. Retrying in %.2fs.", + self.guild.id, + retry_after, + ) + + await asyncio.sleep(retry_after) + retries += 1 + + async def fill_messages(self) -> None: + if not self._get_retrieve(): + return + + self.query["limit"] = self.retrieve + self.query["offset"] = self.offset + data = await self._try_fetch() + + threads = { + int(t["id"]): Thread(guild=self.guild, state=self._state, data=t) + for t in data.get("threads") or [] + } + message_data = [m for ms in data["messages"] for m in ms] + + if self.limit is not None: + self.limit -= len(message_data) + self.offset += self.retrieve + + # if the next offset would exceed the total number of results or maximum allowed offset, stop + if self.offset >= data["total_results"] or self.offset > 9975: + self.limit = 0 # terminate loop + + for element in message_data: + if message := self.create_message(element, threads): + await self.messages.put(message) + + def create_message(self, data: MessagePayload, threads: dict[int, Thread]) -> Message | None: + from .abc import Messageable + + channel_id = int(data["channel_id"]) + channel = self.guild.get_channel_or_thread(channel_id) or threads.get(channel_id) + if not isinstance(channel, Messageable): + # this should never happen. we're here either because the channel resolved to a + # non-messageable guild channel, or because we can't find the channel/thread + return None + + return self._state.create_message(channel=channel, data=data) diff --git a/disnake/types/message.py b/disnake/types/message.py index e1dd6dcbc0..692003b03d 100644 --- a/disnake/types/message.py +++ b/disnake/types/message.py @@ -2,6 +2,7 @@ from __future__ import annotations +from collections.abc import Sequence from typing import Literal, TypedDict from typing_extensions import NotRequired @@ -15,7 +16,7 @@ from .poll import Poll from .snowflake import Snowflake, SnowflakeList from .sticker import StickerItem -from .threads import Thread +from .threads import Thread, ThreadMember from .user import User @@ -170,3 +171,48 @@ class MessagePin(TypedDict): class MessageCall(TypedDict): participants: SnowflakeList ended_timestamp: NotRequired[str | None] + + +class MessageSearchQuery(TypedDict, total=False): + # pagination + limit: int + offset: int + max_id: Snowflake + min_id: Snowflake + # query + slop: int + content: str + channel_id: Sequence[Snowflake] + author_type: Sequence[str] + author_id: Sequence[Snowflake] + mentions: Sequence[Snowflake] + mentions_role: Sequence[Snowflake] + mentions_everyone: bool + replied_to_user_id: Sequence[Snowflake] + replied_to_message_id: Sequence[Snowflake] + pinned: bool + has: Sequence[str] + embed_type: Sequence[str] + embed_provider: Sequence[str] + link_hostname: Sequence[str] + attachment_filename: Sequence[str] + attachment_extension: Sequence[str] + sort_by: str + sort_order: str + include_nsfw: bool + + +class MessageSearchResult(TypedDict): + doing_deep_historical_index: bool + documents_indexed: NotRequired[int] + total_results: int + messages: list[list[Message]] + threads: NotRequired[list[Thread]] + members: NotRequired[list[ThreadMember]] + + +class MessageSearchNotIndexedResult(TypedDict): + message: str + code: int + documents_indexed: int + retry_after: int diff --git a/docs/api/exceptions.rst b/docs/api/exceptions.rst index be2e773ae9..b9f2a723de 100644 --- a/docs/api/exceptions.rst +++ b/docs/api/exceptions.rst @@ -111,6 +111,11 @@ LocalizationKeyError .. autoexception:: LocalizationKeyError +MessageSearchIndexUnavailableError +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +.. autoexception:: MessageSearchIndexUnavailableError + OpusError ~~~~~~~~~ @@ -146,6 +151,7 @@ Exception Hierarchy - :exc:`NotFound` - :exc:`DiscordServerError` - :exc:`LocalizationKeyError` + - :exc:`MessageSearchIndexUnavailableError` - :exc:`WebhookTokenMissing` diff --git a/docs/api/guilds.rst b/docs/api/guilds.rst index 9aacb75b8e..a24f1a23c9 100644 --- a/docs/api/guilds.rst +++ b/docs/api/guilds.rst @@ -19,7 +19,7 @@ Guild .. autoclass:: Guild() :members: - :exclude-members: fetch_members, audit_logs + :exclude-members: fetch_members, audit_logs, search_messages .. automethod:: fetch_members :async-for: @@ -27,6 +27,9 @@ Guild .. automethod:: audit_logs :async-for: + .. automethod:: search_messages + :async-for: + GuildPreview ~~~~~~~~~~~~ @@ -172,6 +175,12 @@ OnboardingPromptType .. autoclass:: OnboardingPromptType() :members: +MessageSearchSortMode +~~~~~~~~~~~~~~~~~~~~~ + +.. autoclass:: MessageSearchSortMode() + :members: + Events ------