diff --git a/revolt/client.py b/revolt/client.py index 59fd16e..93121c5 100755 --- a/revolt/client.py +++ b/revolt/client.py @@ -165,3 +165,10 @@ class Client: self.listeners.setdefault(event, []).append((check, future)) return await asyncio.wait_for(future, timeout) + + @property + def user(self) -> User: + user = self.websocket.user + + assert user + return user diff --git a/revolt/member.py b/revolt/member.py index 2c55194..3475ce8 100755 --- a/revolt/member.py +++ b/revolt/member.py @@ -54,3 +54,15 @@ class Member(User): def avatar(self) -> Optional[Asset]: """Optional[:class:`Asset`] The avatar the member is displaying, this includes guild avatars and masqueraded avatar""" return self.masquerade_avatar or self.guild_avatar or self.original_avatar + + def _update(self, *, nickname: Optional[str] = None, avatar: Optional[File] = None, roles: Optional[list[str]] = None): + if nickname: + self.nickname = nickname + + if avatar: + self.guild_avatar = Asset(avatar, self.state) + + if roles: + 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) + diff --git a/revolt/role.py b/revolt/role.py index 9bfbe4b..5d15a26 100755 --- a/revolt/role.py +++ b/revolt/role.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Optional from .permissions import ServerPermissions @@ -56,3 +56,16 @@ class Role: The new permissions for the role """ await self.state.http.set_role_permissions(self.server.id, self.id, *permissions.value) + + def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None): + if name: + self.name = name + + if colour: + self.colour = colour + + if hoist: + self.hoist = hoist + + if rank: + self.rank = rank diff --git a/revolt/server.py b/revolt/server.py index 6ff5aea..e095444 100755 --- a/revolt/server.py +++ b/revolt/server.py @@ -111,7 +111,7 @@ class Server: self._members: dict[str, Member] = {} self._roles: dict[str, Role] = {role_id: Role(role, role_id, state, self) for role_id, role in data.get("roles", {}).items()} - channels = cast(list[Channel], list(filter(bool, [state.get_channel(channel_id) for channel_id in data["channels"]]))) + channels = [state.get_channel(channel_id) for channel_id in data["channels"]] self._channels: dict[str, Channel] = {channel.id: channel for channel in channels} 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[PermissionPayload] = None, nsfw: Optional[bool] = None, system_messages: Optional[SystemMessagesConfig] = None, categories: Optional[list[CategoryPayload]] = None): diff --git a/revolt/types/gateway.py b/revolt/types/gateway.py index 75a5150..e91f884 100755 --- a/revolt/types/gateway.py +++ b/revolt/types/gateway.py @@ -7,10 +7,9 @@ from .channel import (Channel, DMChannel, Group, SavedMessages, TextChannel, from .message import Message if TYPE_CHECKING: - from .member import Member + from .member import Member, MemberID from .server import Server - from .user import User - + from .user import User, Status __all__ = ( "BasePayload", @@ -26,7 +25,15 @@ __all__ = ( "ChannelDeleteEventPayload", "ChannelStartTypingEventPayload", "ChannelDeleteTypingEventPayload", - "ServerUpdateEventPayload" + "ServerUpdateEventPayload", + "ServerDeleteEventPayload", + "ServerMemberUpdateEventPayload", + "ServerMemberJoinEventPayload", + "ServerMemberLeaveEventPayload", + "ServerRoleUpdateEventPayload", + "ServerRoleDeleteEventPayload", + "UserUpdateEventPayload", + "UserRelationshipEventPayload" ) class BasePayload(TypedDict): @@ -94,3 +101,37 @@ class ServerUpdateEventPayload(BasePayload): id: str data: dict clear: Literal["Icon", "Banner", "Description"] + +class ServerDeleteEventPayload(BasePayload): + id: str + +class ServerMemberUpdateEventPayload(BasePayload): + id: MemberID + data: dict + clear: Literal["Nickname", "Avatar"] + +class ServerMemberJoinEventPayload(BasePayload): + id: str + user: str + +ServerMemberLeaveEventPayload = ServerMemberJoinEventPayload + +class ServerRoleUpdateEventPayload(BasePayload): + id: str + role_id: str + data: dict + clear: Literal["Color"] + +class ServerRoleDeleteEventPayload(BasePayload): + id: str + role_id: str + +class UserUpdateEventPayload(BasePayload): + id: str + data: dict + clear: Literal["ProfileContent", "ProfileBackground", "StatusText", "Avatar"] + +class UserRelationshipEventPayload(BasePayload): + id: str + user: str + status: Status diff --git a/revolt/user.py b/revolt/user.py index fbefc71..d47e23b 100755 --- a/revolt/user.py +++ b/revolt/user.py @@ -7,8 +7,7 @@ from .enums import PresenceType, RelationshipType if TYPE_CHECKING: from .state import State - from .types import User as UserPayload - from .types import UserRelation + from .types import User as UserPayload, Status as StatusPayload, File __all__ = ("User",) @@ -23,6 +22,11 @@ class Status(NamedTuple): text: Optional[str] presence: Optional[PresenceType] +class UserProfile(NamedTuple): + """A namedtuple representing a users profile""" + content: Optional[str] + background: Optional[Asset] + class User: """Represents a user @@ -47,7 +51,7 @@ class User: status: Optional[:class:`Status`] The users status """ - __flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar") + __flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile") __slots__ = (*__flattern_attributes__, "state") def __init__(self, data: UserPayload, state: State): @@ -88,6 +92,8 @@ class User: else: self.status = None + self.profile: Optional[UserProfile] = None + self.masquerade_avatar: Optional[PartialAsset] = None self.masquerade_name: Optional[str] = None @@ -109,3 +115,20 @@ class User: def avatar(self) -> Union[Asset, PartialAsset, None]: """Optional[:class:`Asset`] The avatar the member is displaying, this includes there orginal avatar and masqueraded avatar""" return self.masquerade_avatar or self.original_avatar + + 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: + 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_content: + self.profile = UserProfile(profile_content, self.profile.background if self.profile else None) + + if avatar: + self.original_avatar = Asset(avatar, self.state) + + if online: + self.online = online diff --git a/revolt/websocket.py b/revolt/websocket.py index 8252a07..c5406f0 100755 --- a/revolt/websocket.py +++ b/revolt/websocket.py @@ -5,9 +5,12 @@ import logging from copy import copy from typing import TYPE_CHECKING, Callable, cast +from .enums import RelationshipType +from .user import Status + from .types import (ChannelCreateEventPayload, ChannelDeleteEventPayload, ChannelDeleteTypingEventPayload, - ChannelStartTypingEventPayload, ChannelUpdateEventPayload, ServerUpdateEventPayload) + ChannelStartTypingEventPayload, ChannelUpdateEventPayload, ServerUpdateEventPayload, UserRelationshipEventPayload, ServerRoleDeleteEventPayload, UserUpdateEventPayload, ServerDeleteEventPayload, ServerMemberUpdateEventPayload, ServerMemberJoinEventPayload, ServerMemberLeaveEventPayload, ServerRoleUpdateEventPayload) from .types import Message as MessagePayload from .types import MessageDeleteEventPayload, MessageUpdateEventPayload @@ -36,7 +39,7 @@ __all__ = ("WebsocketHandler",) logger = logging.getLogger("revolt") class WebsocketHandler: - __slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop") + __slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready") def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State): self.session = session @@ -46,6 +49,8 @@ class WebsocketHandler: self.state = state self.websocket: aiohttp.ClientWebSocketResponse self.loop = asyncio.get_running_loop() + self.user = None + self.ready = asyncio.Event() async def send_payload(self, payload: BasePayload): if use_msgpack: @@ -71,6 +76,9 @@ class WebsocketHandler: event_type = payload["type"].lower() logger.debug("Recieved event %s %s", event_type, payload) try: + if event_type != "ready": + await self.ready.wait() + func = getattr(self, f"handle_{event_type}") except: logger.debug("Unknown event '%s'", event_type) @@ -82,8 +90,11 @@ class WebsocketHandler: logger.info("Successfully authenticated") async def handle_ready(self, payload: ReadyEventPayload): - for user in payload["users"]: - self.state.add_user(user) + for user_payload in payload["users"]: + user = self.state.add_user(user_payload) + + if user.relationship == RelationshipType.user: + self.user = user for channel in payload["channels"]: self.state.add_channel(channel) @@ -97,6 +108,7 @@ class WebsocketHandler: await self.state.fetch_all_server_members() + self.ready.set() self.dispatch("ready") async def handle_message(self, payload: MessageEventPayload): @@ -189,6 +201,93 @@ class WebsocketHandler: self.dispatch("server_update", old_server, server) + async def handle_serverdelete(self, payload: ServerDeleteEventPayload): + server = self.state.servers.pop(payload["id"]) + + for channel in server.channels: + del self.state.channels[channel.id] + + self.dispatch("server_delete", server) + + async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload): + member = self.state.get_member(payload["id"]["server"], payload["id"]["user"]) + old_member = copy(member) + + if clear := payload.get("clear"): + if clear == "Nickname": + member.nickname = None + elif clear == "Avatar": + member.guild_avatar = None + + member._update(**payload["data"]) + + self.dispatch("member_update", old_member, member) + + async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload): + member = self.state.add_member(payload["id"], {"_id": {"server": payload["id"], "user": payload["user"]}}) + self.dispatch("member_join", member) + + async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload): + server = self.state.get_server(payload["id"]) + member = server._members.pop(payload["user"]) + + self.dispatch("member_leave", member) + + async def handle_serveroleupdate(self, payload: ServerRoleUpdateEventPayload): + server = self.state.get_server(payload["id"]) + role = server.get_role(payload["role_id"]) + old_role = copy(role) + + if clear := payload.get("clear"): + if clear == "Colour": + role.colour = None + + role._update(**payload["data"]) + + self.dispatch("role_update", old_role, role) + + async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload): + server = self.state.get_server(payload["id"]) + role = server._roles.pop(payload["role_id"]) + + self.dispatch("role_delete", role) + + async def handle_userupdate(self, payload: UserUpdateEventPayload): + user = self.state.get_user(payload["id"]) + old_user = copy(user) + + if clear := payload.get("clear"): + if clear == "ProfileContent": + ... + elif clear == "ProfileBackground": + ... + elif clear == "StatusText": + # user.status will never be none because they are trying to remove the text + if user.status.presence is None: # type: ignore + user.status = None + else: + user.status = Status(None, user.status.presence) # type: ignore + + elif clear == "Avatar": + user.original_avatar = None + + # the keys have . in them so i need to replace with _ + + data = payload["data"] + data["profile_content"] = data.get("profile.content", None) + data["profile_background"] = data.get("profile.background", None) + + user._update(**data) + + self.dispatch("user_update", old_user, user) + + async def handle_userrelationship(self, payload: UserRelationshipEventPayload): + user = self.state.get_user(payload["user"]) + old_relationship = user.relationship + user.relationship = RelationshipType(payload["status"]) + + self.dispatch("user_relationship_update", user, old_relationship, user.relationship) + async def start(self): if use_msgpack: url = f"{self.ws_url}?format=msgpack"