From 9bf443e36b73a2a0cff4f574307400560000a2bc Mon Sep 17 00:00:00 2001 From: Zomatree Date: Fri, 29 Dec 2023 20:09:38 +0000 Subject: [PATCH] Fix channel.history not fetching members to go along with the messages --- revolt/channel.py | 2 +- revolt/http.py | 4 ++-- revolt/message.py | 48 ++++++++++++++++++++++++++----------------- revolt/messageable.py | 25 ++++++++++++++++++---- revolt/server.py | 7 +++++++ revolt/state.py | 4 +--- revolt/types/http.py | 3 ++- 7 files changed, 63 insertions(+), 30 deletions(-) diff --git a/revolt/channel.py b/revolt/channel.py index c32e91c..06459ca 100755 --- a/revolt/channel.py +++ b/revolt/channel.py @@ -266,7 +266,7 @@ class ServerChannel(Channel): def __init__(self, data: ServerChannelPayload, state: State): super().__init__(data, state) - self.server_id: str = data["server"] + self.server_id: Optional[str] = data["server"] self.name: str = data["name"] self.description: Optional[str] = data.get("description") self.nsfw: bool = data.get("nsfw", False) diff --git a/revolt/http.py b/revolt/http.py index 6f6f3f3..56c5631 100755 --- a/revolt/http.py +++ b/revolt/http.py @@ -140,7 +140,7 @@ class HttpClient: return await self.request("POST", f"/channels/{channel}/messages", json=json) def edit_message(self, channel: str, message: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]] = None) -> Request[None]: - json = {} + json: dict[str, Any] = {} if content is not None: json["content"] = content @@ -399,7 +399,7 @@ class HttpClient: return self.request("PUT", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}") def remove_reaction(self, channel_id: str, message_id: str, emoji: str, user_id: Optional[str], remove_all: bool) -> Request[None]: - parameters = {} + parameters: dict[str, str] = {} if user_id: parameters["user_id"] = user_id diff --git a/revolt/message.py b/revolt/message.py index ca6f3d3..30d904a 100755 --- a/revolt/message.py +++ b/revolt/message.py @@ -48,8 +48,6 @@ class Message(Ulid): The time at which the message was edited, will be None if the message has not been edited raw_mentions: list[:class:`str`] A list of ids of the mentions in this message - mentions: list[Union[:class:`Member`, :class:`User`]] - The users or members that where mentioned in the message replies: list[:class:`Message`] The message's this message has replied to, this may not contain all the messages if they are outside the cache reply_ids: list[:class:`str`] @@ -59,7 +57,7 @@ class Message(Ulid): interactions: Optional[:class:`MessageInteractions`] The interactions on the message, if any """ - __slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "mentions", "replies", "reply_ids", "reactions", "interactions") + __slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "replies", "reply_ids", "reactions", "interactions") def __init__(self, data: MessagePayload, state: State): self.state: State = state @@ -79,7 +77,6 @@ class Message(Ulid): self.server_id: str | None = self.channel.server_id self.raw_mentions: list[str] = data.get("mentions", []) - self.mentions: list[Member | User] = [] if self.system_content: author_id: str = self.system_content.get("id", data["author"]) @@ -89,21 +86,9 @@ class Message(Ulid): if self.server_id: author = state.get_member(self.server_id, author_id) - for mention in self.raw_mentions: - try: - self.mentions.append(self.server.get_member(mention)) - except LookupError: - pass - else: author = state.get_user(author_id) - for mention in self.raw_mentions: - try: - self.mentions.append(state.get_user(mention)) - except LookupError: - pass - self.author: Member | User = author if masquerade := data.get("masquerade"): @@ -152,6 +137,31 @@ class Message(Ulid): if edited is not None: self.edited_at = parse_timestamp(edited) + @property + def mentions(self) -> list[User | Member]: + """The users or members that where mentioned in the message + + Returns: list[Union[:class:`Member`, :class:`User`]] + """ + + mentions: list[User | Member] = [] + + if self.server_id: + for mention in self.raw_mentions: + try: + self.mentions.append(self.server.get_member(mention)) + except LookupError: + pass + + else: + for mention in self.raw_mentions: + try: + self.mentions.append(self.state.get_user(mention)) + except LookupError: + pass + + return mentions + async def edit(self, *, content: Optional[str] = None, embeds: Optional[list[SendableEmbed]] = None) -> None: """Edits the message. The bot can only edit its own message @@ -232,12 +242,12 @@ class MessageReply: """ __slots__ = ("message", "mention") - def __init__(self, message: Message, mention: bool = False): - self.message: Message = message + def __init__(self, message: Ulid, mention: bool = False): + self.message: Ulid = message self.mention: bool = mention def to_dict(self) -> MessageReplyPayload: - return { "id": self.message.id, "mention": self.mention } + return {"id": self.message.id, "mention": self.mention} class Masquerade: """represents a message's masquerade. diff --git a/revolt/messageable.py b/revolt/messageable.py index d66f0ed..c16e10c 100755 --- a/revolt/messageable.py +++ b/revolt/messageable.py @@ -9,6 +9,7 @@ if TYPE_CHECKING: from .file import File from .message import Masquerade, Message, MessageInteractions, MessageReply from .state import State + from .types.http import MessageWithUserData __all__ = ("Messageable",) @@ -86,6 +87,18 @@ class Messageable: payload = await self.state.http.fetch_message(await self._get_channel_id(), message_id) return Message(payload, self.state) + def _add_missing_users(self, payload: MessageWithUserData): + for user in payload["users"]: + if user["_id"] not in self.state.users: + self.state.add_user(user) + + if members := payload.get("members", []): + server = self.state.get_server(members[0]["_id"]["server"]) + + for member in members: + if member["_id"]["user"] not in server._members: + server._add_member(member) + async def history(self, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None, nearby: Optional[str] = None) -> list[Message]: """Fetches multiple messages from the channel's history @@ -109,8 +122,10 @@ class Messageable: """ from .message import Message - payloads = await self.state.http.fetch_messages(await self._get_channel_id(), sort=sort, limit=limit, before=before, after=after, nearby=nearby) - return [Message(payload, self.state) for payload in payloads] + payload = await self.state.http.fetch_messages(await self._get_channel_id(), sort=sort, limit=limit, before=before, after=after, nearby=nearby, include_users=True) + self._add_missing_users(payload) + + return [Message(msg, self.state) for msg in payload["messages"]] async def search(self, query: str, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None) -> list[Message]: """searches the channel for a query @@ -135,8 +150,10 @@ class Messageable: """ from .message import Message - payloads = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after) - return [Message(payload, self.state) for payload in payloads] + payload = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after, include_users=True) + self._add_missing_users(payload) + + return [Message(msg, self.state) for msg in payload["messages"]] async def delete_messages(self, messages: list[Message]) -> None: """Bulk deletes messages from the channel diff --git a/revolt/server.py b/revolt/server.py index e992689..d32368e 100755 --- a/revolt/server.py +++ b/revolt/server.py @@ -20,6 +20,7 @@ if TYPE_CHECKING: from .types import File as FilePayload from .types import Server as ServerPayload from .types import SystemMessagesConfig + from .types import Member as MemberPayload __all__ = ("Server", "SystemMessages", "ServerBan") @@ -184,6 +185,12 @@ class Server(Ulid): if channels is not None: self._channels = {channel_id: self.state.get_channel(channel_id) for channel_id in channels} + def _add_member(self, payload: MemberPayload) -> Member: + member = Member(payload, self, self.state) + self._members[member.id] = member + + return member + @property def roles(self) -> list[Role]: """list[:class:`Role`] Gets all roles in the server in decending order""" diff --git a/revolt/state.py b/revolt/state.py index f091161..fc29a16 100755 --- a/revolt/state.py +++ b/revolt/state.py @@ -73,10 +73,8 @@ class State: def add_member(self, server_id: str, payload: MemberPayload) -> Member: server = self.get_server(server_id) - member = Member(payload, server, self) - server._members[member.id] = member - return member + return server._add_member(payload) def add_channel(self, payload: ChannelPayload) -> Channel: channel = channel_factory(payload, self) diff --git a/revolt/types/http.py b/revolt/types/http.py index 9ff45a4..5be3b9d 100755 --- a/revolt/types/http.py +++ b/revolt/types/http.py @@ -1,6 +1,7 @@ from __future__ import annotations from typing import TYPE_CHECKING, TypedDict +from typing_extensions import NotRequired if TYPE_CHECKING: from .member import Member @@ -48,5 +49,5 @@ class GetServerMembers(TypedDict): class MessageWithUserData(TypedDict): messages: list[Message] - members: list[Member] + members: NotRequired[list[Member]] users: list[User]