From 0a6db1dc9cef0cd63b2075e67e5304de76c833ea Mon Sep 17 00:00:00 2001 From: Zomatree Date: Tue, 11 Apr 2023 19:05:56 +0100 Subject: [PATCH] inital permissions calculations --- revolt/channel.py | 22 ++++++- revolt/ext/commands/checks.py | 5 ++ revolt/member.py | 15 +++++ revolt/permissions.py | 57 ++++++++++++++++++- revolt/state.py | 8 ++- revolt/types/user.py | 1 + revolt/user.py | 21 ++++++- typings/msgpack/{__init__.py => __init__.pyi} | 0 8 files changed, 123 insertions(+), 6 deletions(-) rename typings/msgpack/{__init__.py => __init__.pyi} (100%) diff --git a/revolt/channel.py b/revolt/channel.py index 2cfb0b9..24383e0 100755 --- a/revolt/channel.py +++ b/revolt/channel.py @@ -2,6 +2,8 @@ from __future__ import annotations from typing import TYPE_CHECKING, Any, Optional, Union +from revolt.user import User + from .utils import Missing, Ulid from .asset import Asset from .enums import ChannelType @@ -49,7 +51,7 @@ class EditableChannel: nsfw: bool Sets whether the channel is nsfw or not """ - remove = [] + remove: list[str] = [] if kwargs.get("icon", Missing) == None: remove.append("Icon") @@ -123,12 +125,28 @@ class DMChannel(Channel, Messageable): The id of the last message in this channel, if any """ - __slots__ = ("last_message_id",) + __slots__ = ("last_message_id", "recipients") def __init__(self, data: DMChannelPayload, state: State): super().__init__(data, state) + self.recipient_ids: tuple[str, str] = tuple(data["recipients"]) self.last_message_id = data.get("last_message_id") + @property + def recipients(self) -> tuple[User, User]: + a, b = self.recipient_ids + + return (self.state.get_user(a), self.state.get_user(b)) + + @property + def recipient(self) -> User: + if self.recipient_ids[0] != self.state.user_id: + user_id = self.recipient_ids[0] + else: + user_id = self.recipient_ids[1] + + return self.state.get_user(user_id) + @property def last_message(self) -> Message: """Gets the last message from the channel, shorthand for `client.get_message(channel.last_message_id)` diff --git a/revolt/ext/commands/checks.py b/revolt/ext/commands/checks.py index 454db3c..bb5d5fe 100755 --- a/revolt/ext/commands/checks.py +++ b/revolt/ext/commands/checks.py @@ -59,3 +59,8 @@ def is_server_owner(): raise NotServerOwner return inner + +def has_permissions(**permissions: bool): + @check + def inner(context: Context[ClientT]): + ... \ No newline at end of file diff --git a/revolt/member.py b/revolt/member.py index f4af55e..a5988ec 100755 --- a/revolt/member.py +++ b/revolt/member.py @@ -1,7 +1,11 @@ from __future__ import annotations +import this from typing import TYPE_CHECKING, Optional import datetime +from revolt.channel import Channel + +from revolt.permissions import Permissions from .asset import Asset from .user import User @@ -114,3 +118,14 @@ class Member(User): ends_at = datetime.datetime.utcnow() + length await self.state.http.edit_member(self.server.id, self.id, None, {"timeout": ends_at.isoformat()}) + + def get_permissions(self) -> Permissions: + return calculate_permissions(self, self.server) + + def get_channel_permissions(self, channel: Channel): + return calculate_permissions(self, channel) + + def has_permissions(self, **kwargs: bool) -> bool: + calculated_perms = self.get_permissions() + + return all([getattr(calculated_perms, key) == value for key, value in kwargs.items()]) diff --git a/revolt/permissions.py b/revolt/permissions.py index cc65ff9..224f314 100755 --- a/revolt/permissions.py +++ b/revolt/permissions.py @@ -1,11 +1,38 @@ from __future__ import annotations +from datetime import datetime from typing import TYPE_CHECKING, Any, Optional from typing_extensions import Self +from revolt.enums import ChannelType + +from .channel import Channel, DMChannel +from .member import Member +from .server import Server from .types.permissions import Overwrite from .flags import Flags, Flag -__all__ = ("Permissions", "PermissionsOverwrite") +__all__ = ("Permissions", "PermissionsOverwrite", "UserPermissions") + +class UserPermissions(Flags): + @Flag + def access() -> int: + return 1 << 0 + + @Flag + def view_profile() -> int: + return 1 << 1 + + @Flag + def send_message() -> int: + return 1 << 2 + + @Flag + def invite() -> int: + return 1 << 3 + + @classmethod + def all(cls) -> Self: + return cls(access=True, view_profile=True, send_message=True, invite=True) class Permissions(Flags): @Flag @@ -198,3 +225,31 @@ class PermissionsOverwrite: deny = Permissions(overwrite["d"]) return cls(allow, deny) + +def calculate_permissions(member: Member, target: Server | Channel) -> Permissions: + if member.privileged: + return Permissions.all() + + if isinstance(target, Server): + if target.owner_id == member.id: + return Permissions.all() + + permissions = target.default_permissions + + for role in member.roles: + permissions = (permissions | role.permissions._allow) & (~role.permissions._deny) + + if member.current_timeout and member.current_timeout > datetime.now(): + permissions = permissions & Permissions.default_view_only() + + return permissions + + else: + channel_type = target.channel_type + + if channel_type is ChannelType.saved_messages: + return Permissions.all() + elif channel_type is ChannelType.direct_message: + assert isinstance(target, DMChannel) + + user_permissions = target.recipient.permissions \ No newline at end of file diff --git a/revolt/state.py b/revolt/state.py index 19abd7c..4883e7a 100755 --- a/revolt/state.py +++ b/revolt/state.py @@ -23,18 +23,19 @@ if TYPE_CHECKING: __all__ = ("State",) class State: - __slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages") + __slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages", "global_emojis", "user_id") def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int): self.http = http self.api_info = api_info self.max_messages = max_messages + self.user_id = "" self.users: dict[str, User] = {} self.channels: dict[str, Channel] = {} self.servers: dict[str, Server] = {} self.messages: deque[Message] = deque() - self.global_emojis: list[Emoji] + self.global_emojis: list[Emoji] = [] def get_user(self, id: str) -> User: try: @@ -59,6 +60,9 @@ class State: raise LookupError from None def add_user(self, payload: UserPayload) -> User: + if payload["relationship"] == "User": + self.user_id = payload["_id"] + user = User(payload, self) self.users[user.id] = user return user diff --git a/revolt/types/user.py b/revolt/types/user.py index df03691..508a599 100755 --- a/revolt/types/user.py +++ b/revolt/types/user.py @@ -40,6 +40,7 @@ class User(TypedDict): online: NotRequired[bool] flags: NotRequired[int] bot: NotRequired[UserBot] + privileged: NotRequired[bool] class UserProfile(TypedDict, total=False): content: str diff --git a/revolt/user.py b/revolt/user.py index 2c1ae08..31de325 100755 --- a/revolt/user.py +++ b/revolt/user.py @@ -3,6 +3,7 @@ from __future__ import annotations from typing import TYPE_CHECKING, NamedTuple, Optional, Union from weakref import WeakSet +from .permissions import UserPermissions from .asset import Asset, PartialAsset from .channel import DMChannel from .enums import PresenceType, RelationshipType @@ -60,8 +61,10 @@ class User(Messageable, Ulid): The users status dm_channel: Optional[:class:`DMChannel`] The dm channel between the client and the user, this will only be set if the client has dm'ed the user or :meth:`User.open_dm` was run + privileged: :class:`bool` + Whether the user is privileged """ - __flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel") + __flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel", "privileged") __slots__ = (*__flattern_attributes__, "state", "_members") def __init__(self, data: UserPayload, state: State): @@ -82,6 +85,7 @@ class User(Messageable, Ulid): self.badges = UserBadges._from_value(data.get("badges", 0)) self.online = data.get("online", False) self.flags = data.get("flags", 0) + self.privileged = data.get("privileged", False) avatar = data.get("avatar") self.original_avatar = Asset(avatar, state) if avatar else None @@ -109,6 +113,21 @@ class User(Messageable, Ulid): self.masquerade_avatar: Optional[PartialAsset] = None self.masquerade_name: Optional[str] = None + @property + def permissions(self) -> UserPermissions: + permissions = UserPermissions() + + if self.relationship in [RelationshipType.friend, RelationshipType.user]: + return UserPermissions.all() + elif self.relationship in [RelationshipType.blocked, RelationshipType.blocked_other]: + return UserPermissions(access=True) + elif self.relationship in [RelationshipType.incoming_friend_request, RelationshipType.outgoing_friend_request]: + permissions.access = True + + + + return permissions + async def _get_channel_id(self): if not self.dm_channel: payload = await self.state.http.open_dm(self.id) diff --git a/typings/msgpack/__init__.py b/typings/msgpack/__init__.pyi similarity index 100% rename from typings/msgpack/__init__.py rename to typings/msgpack/__init__.pyi