fix issues stemming from bugs in the revolt api

This commit is contained in:
Zomatree
2022-11-16 23:06:31 +00:00
parent 9504c05e7b
commit a92324149e
10 changed files with 85 additions and 82 deletions
+1 -1
View File
@@ -36,7 +36,7 @@ class Group(Command[ClientT]):
self.subcommands: dict[str, Command[ClientT]] = {}
super().__init__(callback, name, aliases)
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command):
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command[ClientT]):
"""A decorator that turns a function into a :class:`Command` and registers the command as a subcommand.
Parameters
+3 -2
View File
@@ -1,7 +1,8 @@
from __future__ import annotations
from inspect import Parameter
from typing import Any, Iterable, TYPE_CHECKING, TypeVar
from typing import Any, Iterable, TYPE_CHECKING
from typing_extensions import TypeVar
if TYPE_CHECKING:
from .client import CommandsClient
@@ -9,7 +10,7 @@ if TYPE_CHECKING:
__all__ = ("evaluate_parameters",)
ClientT = TypeVar("ClientT", bound="CommandsClient")
ClientT = TypeVar("ClientT", bound="CommandsClient", default="CommandsClient")
def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]:
+13 -2
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import datetime
from .asset import Asset
from .user import User
@@ -31,7 +32,7 @@ class Member(User):
guild_avatar: Optional[:class:`Asset`]
The member's guild avatar if any
"""
__slots__ = ("state", "nickname", "roles", "server", "guild_avatar")
__slots__ = ("state", "nickname", "roles", "server", "guild_avatar", "joined_at", "timeout")
def __init__(self, data: MemberPayload, server: Server, state: State):
user = state.get_user(data["_id"]["user"])
@@ -53,6 +54,16 @@ class Member(User):
self.server = server
self.nickname = data.get("nickname")
joined_at = data["joined_at"]
if isinstance(joined_at, int):
self.joined_at = datetime.datetime.fromtimestamp(joined_at / 1000)
else:
self.joined_at = datetime.datetime.strptime(joined_at, "%Y-%m-%dT%H:%M:%S.%f%z")
self.timeout = None
if timeout := data.get("timeout"):
self.timeout = datetime.datetime.strptime(timeout, "%Y-%m-%dT%H:%M:%S.%f%z")
@property
def avatar(self) -> Optional[Asset]:
@@ -71,7 +82,7 @@ class Member(User):
if avatar:
self.guild_avatar = Asset(avatar, self.state)
if roles:
if roles is not None:
member_roles = [self.server.get_role(role_id) for role_id in roles]
self.roles = sorted(member_roles, key=lambda role: role.rank, reverse=True)
+3 -4
View File
@@ -98,7 +98,7 @@ class Message(Ulid):
try:
message = state.get_message(reply)
self.replies.append(message)
except KeyError:
except LookupError:
pass
self.reply_ids.append(reply)
@@ -115,12 +115,11 @@ class Message(Ulid):
else:
self.interactions = None
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited_at: str):
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited: int):
if content:
self.content = content
self.edited_at = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z")
# strptime is used here instead of fromisoformat because of its inability to parse `Z` (Zulu or UTC time) in the RFCC 3339 format provided by API
self.edited = datetime.datetime.fromtimestamp(edited / 1000)
if embeds:
self.embeds = [to_embed(embed, self.state) for embed in embeds]
+15 -7
View File
@@ -118,7 +118,15 @@ class Server(Ulid):
self._members: dict[str, Member] = {}
self._roles: dict[str, Role] = {role_id: Role(role, role_id, self, state) for role_id, role in data.get("roles", {}).items()}
self._channels: dict[str, Channel] = {channel_id: state.get_channel(channel_id) for channel_id in data.get("channels", [])}
self._channels: dict[str, Channel] = {}
# The api doesnt send us all the channels but sends us all the ids, this is because channels we dont have permissions to see are not sent
# this causes get_channel to error so we have to first check ourself if its in the cache.
for channel_id in data["channels"]:
if channel := state.channels.get(channel_id):
self._channels[channel_id] = channel
self._emojis: dict[str, Emoji] = {}
def _update(self, *, owner: Optional[str] = None, name: Optional[str] = None, description: Optional[str] = None, icon: Optional[FilePayload] = None, banner: Optional[FilePayload] = None, default_permissions: Optional[int] = None, nsfw: Optional[bool] = None, system_messages: Optional[SystemMessagesConfig] = None, categories: Optional[list[CategoryPayload]] = None, channels: Optional[list[str]] = None):
@@ -198,8 +206,8 @@ class Server(Ulid):
"""
try:
return self._members[member_id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def get_channel(self, channel_id: str) -> Channel:
"""Gets a channel from the cache
@@ -216,8 +224,8 @@ class Server(Ulid):
"""
try:
return self._channels[channel_id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def get_category(self, category_id: str) -> Category:
"""Gets a category from the cache
@@ -234,8 +242,8 @@ class Server(Ulid):
"""
try:
return self._categories[category_id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def get_emoji(self, emoji_id: str) -> Emoji:
"""Gets a emoji from the cache
+6 -6
View File
@@ -39,8 +39,8 @@ class State:
def get_user(self, id: str) -> User:
try:
return self.users[id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def get_member(self, server_id: str, member_id: str) -> Member:
server = self.servers[server_id]
@@ -49,14 +49,14 @@ class State:
def get_channel(self, id: str) -> Channel:
try:
return self.channels[id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def get_server(self, id: str) -> Server:
try:
return self.servers[id]
except KeyError as e:
raise LookupError from e
except KeyError:
raise LookupError from None
def add_user(self, payload: UserPayload) -> User:
user = User(payload, self)
+11 -17
View File
@@ -1,28 +1,27 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict, Union
from typing_extensions import NotRequired
from revolt.types.permissions import Overwrite
from .permissions import Overwrite
from .channel import (Channel, DMChannel, GroupDMChannel, SavedMessages,
TextChannel, VoiceChannel)
from .file import File
from .message import Message
from .user import Status
if TYPE_CHECKING:
from .category import Category
from .member import Member, MemberID
from .server import Server, SystemMessagesConfig
from .user import User
from .user import User, UserProfile, Status
from .emoji import Emoji
from .file import File
__all__ = (
"BasePayload",
"AuthenticatePayload",
"ReadyEventPayload",
"MessageEventPayload",
"MessageUpdateEditedData",
"MessageUpdateData",
"MessageUpdateEventPayload",
"MessageDeleteEventPayload",
@@ -62,11 +61,9 @@ class ReadyEventPayload(BasePayload):
class MessageEventPayload(BasePayload, Message):
pass
MessageUpdateEditedData = TypedDict("MessageUpdateEditedData", {"$date": str})
class MessageUpdateData(TypedDict):
content: str
edited: MessageUpdateEditedData
edited: int
class MessageUpdateEventPayload(BasePayload):
channel: str
@@ -173,14 +170,11 @@ class ServerRoleDeleteEventPayload(BasePayload):
id: str
role_id: str
UserUpdateEventPayloadData = TypedDict("UserUpdateEventPayloadData", {
"status": Status,
"profile.background": File,
"profile.content": str,
"avatar": File,
"online": bool
}, total=False)
class UserUpdateEventPayloadData(TypedDict):
status: NotRequired[Status]
avatar: NotRequired[File]
online: NotRequired[bool]
profile: NotRequired[UserProfile]
class UserUpdateEventPayload(BasePayload):
id: str
+2
View File
@@ -19,3 +19,5 @@ class Member(TypedDict):
nickname: NotRequired[str]
avatar: NotRequired[File]
roles: NotRequired[list[str]]
joined_at: int | str
timeout: NotRequired[str]
+11 -8
View File
@@ -14,6 +14,7 @@ if TYPE_CHECKING:
from .types import File
from .types import Status as StatusPayload
from .types import User as UserPayload
from .types import UserProfile as UserProfileData
from .member import Member
__all__ = ("User", "Status", "Relation", "UserProfile")
@@ -138,27 +139,29 @@ class User(Messageable, Ulid):
""":class:`str`: Returns a string that allows you to mention the given user."""
return f"<@{self.id}>"
def _update(self, *, status: Optional[StatusPayload] = None, profile_content: Optional[str] = None, profile_background: Optional[File] = None, avatar: Optional[File] = None, online: Optional[bool] = None):
if status:
def _update(self, *, status: Optional[StatusPayload] = None, profile: Optional[UserProfileData] = None, avatar: Optional[File] = None, online: Optional[bool] = None):
if status is not None:
presence = status.get("presence")
self.status = Status(status.get("text"), PresenceType(presence) if presence else None)
if profile_background:
self.profile = UserProfile(self.profile.content if self.profile else None, Asset(profile_background, self.state))
if profile is not None:
if background_file := profile.get("background"):
background = Asset(background_file, self.state)
else:
background = None
if profile_content:
self.profile = UserProfile(profile_content, self.profile.background if self.profile else None)
self.profile = UserProfile(profile.get("content"), background)
if avatar:
self.original_avatar = Asset(avatar, self.state)
if online:
if online is not None:
self.online = online
# update user infomation for all members
for member in self._members:
User._update(member, status=status, profile_content=profile_content, profile_background=profile_background, avatar=avatar, online=online)
User._update(member, status=status, profile=profile, avatar=avatar, online=online)
async def default_avatar(self) -> bytes:
"""Returns the default avatar for this user
+20 -35
View File
@@ -144,21 +144,10 @@ class WebsocketHandler:
except LookupError:
return
if server := message.server:
await self._wait_for_server_ready(server.id)
if server_id := message.channel.server_id:
await self._wait_for_server_ready(server_id)
data = payload["data"]
kwargs = {}
if content := data.get("content"):
kwargs["content"] = content
kwargs["edited_at"] = data["edited"]["$date"]
if embeds := data.get("embeds"):
kwargs["embeds"] = embeds
message._update(**kwargs)
message._update(**payload["data"])
self.dispatch("message_update", message)
@@ -170,8 +159,8 @@ class WebsocketHandler:
except LookupError:
return
if server := message.server:
await self._wait_for_server_ready(server.id)
if server_id := message.channel.server_id:
await self._wait_for_server_ready(server_id)
self.state.messages.remove(message)
@@ -181,16 +170,19 @@ class WebsocketHandler:
async def handle_channelcreate(self, payload: ChannelCreateEventPayload):
channel = self.state.add_channel(payload)
if server := channel.server:
await self._wait_for_server_ready(server.id)
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
self.dispatch("channel_create", channel)
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload):
channel = self.state.get_channel(payload["id"])
# Revolt sends channel updates for channels we dont have permissions to see, a bug, but still can cause issues as its not in the cache
if server := channel.server:
await self._wait_for_server_ready(server.id)
if not (channel := self.state.channels.get(payload["id"], None)):
return
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
old_channel = copy(channel)
@@ -211,16 +203,16 @@ class WebsocketHandler:
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload):
channel = self.state.channels.pop(payload["id"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
self.dispatch("channel_delete", channel)
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload):
channel = self.state.get_channel(payload["id"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
user = self.state.get_user(payload["user"])
@@ -229,8 +221,8 @@ class WebsocketHandler:
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload):
channel = self.state.get_channel(payload["id"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
user = self.state.get_user(payload["user"])
@@ -360,14 +352,7 @@ class WebsocketHandler:
elif clear == "Avatar":
user.original_avatar = None
# the keys have . in them so I need to replace with _
# type: ignore is for it to stop complaining about the keys not existing in the typeddict
data = payload["data"]
data["profile_content"] = data.pop("profile.content", None) # type: ignore
data["profile_background"] = data.pop("profile.background", None) # type: ignore
user._update(**data) # type: ignore
user._update(**payload["data"])
self.dispatch("user_update", old_user, user)