Fix channel.history not fetching members to go along with the messages

This commit is contained in:
Zomatree
2023-12-29 20:09:38 +00:00
parent 5d3250bcce
commit 9bf443e36b
7 changed files with 63 additions and 30 deletions
+1 -1
View File
@@ -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)
+2 -2
View File
@@ -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
+29 -19
View File
@@ -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.
+21 -4
View File
@@ -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
+7
View File
@@ -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"""
+1 -3
View File
@@ -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)
+2 -1
View File
@@ -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]