inital permissions calculations

This commit is contained in:
Zomatree
2023-04-11 19:05:56 +01:00
parent a672f949a1
commit 0a6db1dc9c
8 changed files with 123 additions and 6 deletions
+20 -2
View File
@@ -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)`
+5
View File
@@ -59,3 +59,8 @@ def is_server_owner():
raise NotServerOwner
return inner
def has_permissions(**permissions: bool):
@check
def inner(context: Context[ClientT]):
...
+15
View File
@@ -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()])
+56 -1
View File
@@ -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
+6 -2
View File
@@ -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
+1
View File
@@ -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
+20 -1
View File
@@ -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)