mirror of
https://github.com/stoatchat/python-client-sdk.git
synced 2026-07-21 18:15:28 -04:00
inital permissions calculations
This commit is contained in:
+20
-2
@@ -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)`
|
||||
|
||||
@@ -59,3 +59,8 @@ def is_server_owner():
|
||||
raise NotServerOwner
|
||||
|
||||
return inner
|
||||
|
||||
def has_permissions(**permissions: bool):
|
||||
@check
|
||||
def inner(context: Context[ClientT]):
|
||||
...
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user