mirror of
https://github.com/stoatchat/python-client-sdk.git
synced 2026-07-21 01:55:23 -04:00
fix issues stemming from bugs in the revolt api
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user