From 96922d021e567a620d948148e58df5f337da8b60 Mon Sep 17 00:00:00 2001 From: Zomatree Date: Sun, 29 Aug 2021 17:05:28 +0100 Subject: [PATCH] fetch all members at startup --- revolt/channel.py | 9 +++++++++ revolt/http.py | 8 +++++++- revolt/member.py | 15 ++++++++++++--- revolt/message.py | 16 +++++++++------- revolt/state.py | 10 ++++++++++ revolt/user.py | 2 ++ 6 files changed, 49 insertions(+), 11 deletions(-) diff --git a/revolt/channel.py b/revolt/channel.py index 4937f59..10b8ba2 100644 --- a/revolt/channel.py +++ b/revolt/channel.py @@ -32,6 +32,15 @@ class TextChannel(Channel, Messageable): super().__init__(data, state) Messageable.__init__(self, state) +class PartialTextChannel(Messageable): + def __init__(self, channel_id: str, state: State): + super().__init__(state) + + self.id = channel_id + self.name = "Unknown Channel" + self.channel_type = "TextChannel" + self.server = None + class VoiceChannel(Channel): def __init__(self, data: ChannelPayload, state: State): super().__init__(data, state) diff --git a/revolt/http.py b/revolt/http.py index 1753bbf..4e15a4e 100644 --- a/revolt/http.py +++ b/revolt/http.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Optional, TYPE_CHECKING, Literal +from typing import Any, Coroutine, Optional, TYPE_CHECKING, Literal, TypeVar import aiohttp import ulid @@ -17,6 +17,9 @@ if TYPE_CHECKING: from .types import ApiInfo, Autumn as AutumnPayload, Message as MessagePayload, Embed as EmbedPayload, GetServerMembers from .file import File +T = TypeVar("T") +Request = Coroutine[Any, Any, T] + class HttpClient: def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo): self.session = session @@ -95,3 +98,6 @@ class HttpClient: json["attachments"] = attachment_ids return await self.request("POST", f"/channels/{channel}/messages", json=json) + + def get_server_members(self, server_id: str) -> Request[GetServerMembers]: + return self.request("GET", f"/servers/{server_id}/members") diff --git a/revolt/member.py b/revolt/member.py index dc1f656..521145b 100644 --- a/revolt/member.py +++ b/revolt/member.py @@ -2,16 +2,25 @@ from __future__ import annotations from typing import TYPE_CHECKING +from .user import User + if TYPE_CHECKING: from .state import State from .types import Member as MemberPayload from .server import Server -class Member: +def flattern_user(member: Member, user: User): + for attr in user.__flattern_attributes__: + setattr(member, attr, getattr(user, attr)) + +class Member(User): def __init__(self, data: MemberPayload, server: Server, state: State): + user = state.get_user(data["_id"]["user"]) + assert user + flattern_user(self, user) + self._state = state - self.id = data["_id"] + self.id = data["_id"]["user"] self.nickname = data.get("nickname") self.roles = [server.get_role(role_id) for role_id in data.get("roles", [])] - self.server = server diff --git a/revolt/message.py b/revolt/message.py index a4349e0..a9c7ec8 100644 --- a/revolt/message.py +++ b/revolt/message.py @@ -1,11 +1,10 @@ from __future__ import annotations - from typing import TYPE_CHECKING from .asset import Asset from .embed import Embed -from .channel import TextChannel +from .channel import TextChannel, PartialTextChannel, Messageable if TYPE_CHECKING: from .state import State @@ -16,14 +15,17 @@ class Message: self.state = state self.id = data["_id"] - self.content = data['content'] - self.attachments = [Asset(attachment, state) for attachment in data.get('attachments', [])] + self.content = data["content"] + self.attachments = [Asset(attachment, state) for attachment in data.get("attachments", [])] self.embeds = [Embed.from_dict(embed) for embed in data.get("embeds", [])] - self.channel = state.get_channel(data['channel']) + channel = state.get_channel(data["channel"]) or PartialTextChannel(data["channel"], state) + assert isinstance(channel, Messageable) + self.channel = channel + self.server = self.channel and self.channel.server if isinstance(self.channel, TextChannel) and self.server: - self.author = state.get_member(self.server.id, data['author']) + self.author = state.get_member(self.server.id, data["author"]) else: - self.author = state.get_user(data['author']) + self.author = state.get_user(data["author"]) diff --git a/revolt/state.py b/revolt/state.py index 6353d66..f722aa1 100644 --- a/revolt/state.py +++ b/revolt/state.py @@ -75,3 +75,13 @@ class State: self.messages.appendleft(message) return message + + async def fetch_all_server_members(self): + for server_id in self.servers.keys(): + data = await self.http.get_server_members(server_id) + + for user in data["users"]: + self.add_user(user) + + for member in data["members"]: + self.add_member(server_id, member) diff --git a/revolt/user.py b/revolt/user.py index 04514e6..1b21df8 100644 --- a/revolt/user.py +++ b/revolt/user.py @@ -7,6 +7,8 @@ if TYPE_CHECKING: from .types import User as UserPayload class User: + __flattern_attributes__ = ("name", "bot", "owner", "badges", "online", "flags") + def __init__(self, data: UserPayload, state: State): self.state = state self.id = data["_id"]