mirror of
https://github.com/stoatchat/python-client-sdk.git
synced 2026-07-25 16:35:33 -04:00
Compare commits
34 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4e70b8076c | |||
| 828b428f48 | |||
| 4efaac40f6 | |||
| 78f14fa9a0 | |||
| 300169d71a | |||
| 983907e0ad | |||
| 4c8553d5ef | |||
| f717163e17 | |||
| 60474fdeef | |||
| ccfae66e16 | |||
| 60ff8d81c9 | |||
| dfb45494ba | |||
| 52c2be4e91 | |||
| d5a15c0f44 | |||
| 4bf0530f8e | |||
| 77ff484ad8 | |||
| d909b0eb3e | |||
| 32d6b1d6e8 | |||
| f9ba869e75 | |||
| 0e694db586 | |||
| ba3f74dcd0 | |||
| 625dd4eac2 | |||
| 0a6db1dc9c | |||
| a672f949a1 | |||
| df9893aaff | |||
| 35a4614b61 | |||
| 263f99f281 | |||
| e99b6edee3 | |||
| 720272d3cb | |||
| ab5a2751df | |||
| 6059e4e4dc | |||
| 63491ec76e | |||
| 76364572f1 | |||
| 75f980e0f9 |
@@ -1,16 +1,38 @@
|
||||
on: [push, pull_request]
|
||||
name: pyright
|
||||
jobs:
|
||||
pyright:
|
||||
pyright-type-checking:
|
||||
strategy:
|
||||
matrix:
|
||||
version: ["3.9", "3.10", "3.11"]
|
||||
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
- uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: '3.9'
|
||||
python-version: ${{ matrix.version }}
|
||||
- run: pip install .[speedups,docs]
|
||||
- uses: jakebailey/pyright-action@v1
|
||||
with:
|
||||
lib: true
|
||||
python-version: 3.9
|
||||
python-version: ${{ matrix.version }}
|
||||
working-directory: revolt
|
||||
|
||||
pyright-type-completeness:
|
||||
strategy:
|
||||
matrix:
|
||||
version: ["3.9", "3.10", "3.11"]
|
||||
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
- uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: ${{ matrix.version }}
|
||||
- run: pip install .[speedups,docs]
|
||||
- uses: jakebailey/pyright-action@v1
|
||||
with:
|
||||
python-version: ${{ matrix.version }}
|
||||
working-directory: revolt
|
||||
verify-types: revolt
|
||||
ignore-external: true
|
||||
|
||||
+3
-2
@@ -4,7 +4,6 @@ sphinx:
|
||||
configuration: docs/conf.py
|
||||
|
||||
python:
|
||||
version: "3.9"
|
||||
install:
|
||||
- method: pip
|
||||
path: .
|
||||
@@ -12,4 +11,6 @@ python:
|
||||
- docs
|
||||
|
||||
build:
|
||||
image: testing
|
||||
tools:
|
||||
python: "3.9"
|
||||
os: "ubuntu-22.04"
|
||||
|
||||
@@ -8,13 +8,13 @@ build:
|
||||
python -m build
|
||||
|
||||
upload:
|
||||
python -m twine upload dist/* -u $PYPI_USERNAME -p $PYPI_PASSWORD
|
||||
python -m twine upload dist/*
|
||||
|
||||
lint:
|
||||
pyright . --venv-path .venv
|
||||
pyright .
|
||||
|
||||
coverage:
|
||||
pyright --lib --ignoreexternal --verifytypes revolt
|
||||
pyright --ignoreexternal --verifytypes revolt
|
||||
|
||||
docs:
|
||||
cd docs && make html
|
||||
|
||||
@@ -6,90 +6,127 @@ API Reference
|
||||
|
||||
.. autoclass:: Client
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Asset
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: PartialAsset
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Channel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: ServerChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: SavedMessageChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: DMChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: GroupDMChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: TextChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: VoiceChannel
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Embed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: WebsiteEmbed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: ImageEmbed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: TextEmbed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: NoneEmbed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: SendableEmbed
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: File
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Member
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Message
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: MessageReply
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Masquerade
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Messageable
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Permissions
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: UserPermissions
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: PermissionsOverwrite
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Role
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Server
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: ServerBan
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: Category
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: SystemMessages
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autoclass:: User
|
||||
:members:
|
||||
:inherited-members:
|
||||
|
||||
.. autonamedtuple:: Relation
|
||||
|
||||
|
||||
+11
-6
@@ -25,7 +25,6 @@ dependencies = [
|
||||
[project.optional-dependencies]
|
||||
speedups = [
|
||||
"ujson==5.1.*",
|
||||
"aiohttp[speedups]==3.8.*",
|
||||
"msgpack==1.0.*"
|
||||
]
|
||||
docs = [
|
||||
@@ -34,9 +33,6 @@ docs = [
|
||||
"sphinx-toolbox==3.2.*",
|
||||
"setuptools==65.4.*"
|
||||
]
|
||||
dev = [
|
||||
"msgpack-types @ git+https://github.com/Zomatree/msgpack-types"
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/revoltchat/revolt.py"
|
||||
@@ -51,13 +47,22 @@ email = "me@zomatree.live"
|
||||
[tool.hatch.version]
|
||||
path = "revolt/__init__.py"
|
||||
|
||||
[tool.hatch.metadata]
|
||||
allow-direct-references = true
|
||||
[tool.hatch.build]
|
||||
only-packages = true
|
||||
include = ["revolt/**/*"]
|
||||
|
||||
|
||||
[tool.pyright]
|
||||
reportPrivateUsage = false
|
||||
reportImportCycles = false
|
||||
reportIncompatibleMethodOverride = false
|
||||
typeCheckingMode = "strict"
|
||||
|
||||
[tool.hatch.build.targets.sdist]
|
||||
strict-naming = false
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
strict-naming = false
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
|
||||
+1
-1
@@ -18,4 +18,4 @@ from .role import *
|
||||
from .server import *
|
||||
from .user import *
|
||||
|
||||
__version__ = "0.1.9"
|
||||
__version__ = "0.1.11"
|
||||
|
||||
+20
-21
@@ -42,14 +42,16 @@ class Asset(Ulid):
|
||||
__slots__ = ("state", "id", "tag", "size", "filename", "content_type", "width", "height", "type", "url")
|
||||
|
||||
def __init__(self, data: FilePayload, state: State):
|
||||
self.state = state
|
||||
self.state: State = state
|
||||
|
||||
self.id = data['_id']
|
||||
self.tag = data['tag']
|
||||
self.size = data['size']
|
||||
self.filename = data['filename']
|
||||
self.id: str = data['_id']
|
||||
self.tag: str = data['tag']
|
||||
self.size: int = data['size']
|
||||
self.filename: str = data['filename']
|
||||
|
||||
metadata = data['metadata']
|
||||
self.height: int | None
|
||||
self.width: int | None
|
||||
|
||||
if metadata["type"] == "Image" or metadata["type"] == "Video": # cannot use `in` because type narrowing will not happen
|
||||
self.height = metadata["height"]
|
||||
@@ -58,17 +60,17 @@ class Asset(Ulid):
|
||||
self.height = None
|
||||
self.width = None
|
||||
|
||||
self.content_type = data["content_type"]
|
||||
self.type = AssetType(metadata["type"])
|
||||
self.content_type: str | None = data["content_type"]
|
||||
self.type: AssetType = AssetType(metadata["type"])
|
||||
|
||||
base_url = self.state.api_info["features"]["autumn"]["url"]
|
||||
self.url = f"{base_url}/{self.tag}/{self.id}"
|
||||
self.url: str = f"{base_url}/{self.tag}/{self.id}"
|
||||
|
||||
async def read(self) -> bytes:
|
||||
"""Reads the files content into bytes"""
|
||||
return await self.state.http.request_file(self.url)
|
||||
|
||||
async def save(self, fp: IOBase):
|
||||
async def save(self, fp: IOBase) -> None:
|
||||
"""Reads the files content and saves it to a file
|
||||
|
||||
Parameters
|
||||
@@ -85,8 +87,6 @@ class PartialAsset(Asset):
|
||||
-----------
|
||||
id: :class:`str`
|
||||
The id of the asset, this will always be ``"0"``
|
||||
tag: Optional[:class:`str`]
|
||||
The tag of the asset, this corrasponds to where the asset is used, this will always be ``None``
|
||||
size: :class:`int`
|
||||
Amount of bytes in the file, this will always be ``0``
|
||||
filename: :class:`str`
|
||||
@@ -102,13 +102,12 @@ class PartialAsset(Asset):
|
||||
"""
|
||||
|
||||
def __init__(self, url: str, state: State):
|
||||
self.state = state
|
||||
self.id = "0"
|
||||
self.tag = None
|
||||
self.size = 0
|
||||
self.filename = ""
|
||||
self.height = None
|
||||
self.width = None
|
||||
self.content_type = mimetypes.guess_extension(url)
|
||||
self.type = AssetType.file
|
||||
self.url = url
|
||||
self.state: State = state
|
||||
self.id: str = "0"
|
||||
self.size: int = 0
|
||||
self.filename: str = ""
|
||||
self.height: int | None = None
|
||||
self.width: int | None = None
|
||||
self.content_type: str | None = mimetypes.guess_extension(url)
|
||||
self.type: AssetType = AssetType.file
|
||||
self.url: str = url
|
||||
|
||||
+4
-4
@@ -25,10 +25,10 @@ class Category(Ulid):
|
||||
"""
|
||||
|
||||
def __init__(self, data: CategoryPayload, state: State):
|
||||
self.state = state
|
||||
self.name = data["title"]
|
||||
self.id = data["id"]
|
||||
self.channel_ids = data["channels"]
|
||||
self.state: State = state
|
||||
self.name: str = data["title"]
|
||||
self.id: str = data["id"]
|
||||
self.channel_ids: list[str] = data["channels"]
|
||||
|
||||
@property
|
||||
def channels(self) -> list[Channel]:
|
||||
|
||||
+100
-46
@@ -2,12 +2,11 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||
|
||||
from .utils import Missing, Ulid
|
||||
from .asset import Asset
|
||||
from .enums import ChannelType
|
||||
from .messageable import Messageable
|
||||
from .permissions import Permissions, PermissionsOverwrite
|
||||
from .utils import Missing
|
||||
from .utils import Missing, Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .message import Message
|
||||
@@ -16,14 +15,15 @@ if TYPE_CHECKING:
|
||||
from .state import State
|
||||
from .types import Channel as ChannelPayload
|
||||
from .types import DMChannel as DMChannelPayload
|
||||
from .types import GroupDMChannel as GroupDMChannelPayload
|
||||
from .types import SavedMessages as SavedMessagesPayload
|
||||
from .types import TextChannel as TextChannelPayload
|
||||
from .types import GuildChannel as GuildChannelPayload
|
||||
from .types import File as FilePayload
|
||||
from .types import GroupDMChannel as GroupDMChannelPayload
|
||||
from .types import Overwrite as OverwritePayload
|
||||
from .types import SavedMessages as SavedMessagesPayload
|
||||
from .types import ServerChannel as ServerChannelPayload
|
||||
from .types import TextChannel as TextChannelPayload
|
||||
from .user import User
|
||||
|
||||
__all__ = ("DMChannel", "GroupDMChannel", "SavedMessageChannel", "TextChannel", "VoiceChannel", "Channel")
|
||||
__all__ = ("DMChannel", "GroupDMChannel", "SavedMessageChannel", "TextChannel", "VoiceChannel", "Channel", "ServerChannel")
|
||||
|
||||
class EditableChannel:
|
||||
__slots__ = ()
|
||||
@@ -31,7 +31,7 @@ class EditableChannel:
|
||||
state: State
|
||||
id: str
|
||||
|
||||
async def edit(self, **kwargs: Any):
|
||||
async def edit(self, **kwargs: Any) -> None:
|
||||
"""Edits the channel
|
||||
|
||||
Passing ``None`` to the parameters that accept it will remove them.
|
||||
@@ -49,12 +49,12 @@ class EditableChannel:
|
||||
nsfw: bool
|
||||
Sets whether the channel is nsfw or not
|
||||
"""
|
||||
remove: list[str] = []
|
||||
|
||||
if kwargs.get("icon", Missing) == None:
|
||||
remove = "Icon"
|
||||
remove.append("Icon")
|
||||
elif kwargs.get("description", Missing) == None:
|
||||
remove = "Description"
|
||||
else:
|
||||
remove = None
|
||||
remove.append("Description")
|
||||
|
||||
if icon := kwargs.get("icon"):
|
||||
asset = await self.state.http.upload_file(icon, "icons")
|
||||
@@ -80,26 +80,32 @@ class Channel(Ulid):
|
||||
__slots__ = ("state", "id", "channel_type", "server_id")
|
||||
|
||||
def __init__(self, data: ChannelPayload, state: State):
|
||||
self.state = state
|
||||
self.id = data["_id"]
|
||||
self.channel_type = ChannelType(data["channel_type"])
|
||||
self.state: State = state
|
||||
self.id: str = data["_id"]
|
||||
self.channel_type: ChannelType = ChannelType(data["channel_type"])
|
||||
self.server_id: Optional[str] = None
|
||||
|
||||
async def _get_channel_id(self) -> str:
|
||||
return self.id
|
||||
|
||||
def _update(self, **_: Any):
|
||||
def _update(self, **_: Any) -> None:
|
||||
pass
|
||||
|
||||
async def delete(self):
|
||||
async def delete(self) -> None:
|
||||
"""Deletes or closes the channel"""
|
||||
await self.state.http.close_channel(self.id)
|
||||
|
||||
@property
|
||||
def server(self) -> Server:
|
||||
""":class:`Server` The server this voice channel belongs too"""
|
||||
""":class:`Server` The server this voice channel belongs too
|
||||
|
||||
Raises
|
||||
-------
|
||||
:class:`LookupError`
|
||||
Raises if the channel is not part of a server
|
||||
"""
|
||||
if not self.server_id:
|
||||
raise IndexError
|
||||
raise LookupError
|
||||
|
||||
return self.state.get_server(self.server_id)
|
||||
|
||||
@@ -123,11 +129,27 @@ class DMChannel(Channel, Messageable):
|
||||
The id of the last message in this channel, if any
|
||||
"""
|
||||
|
||||
__slots__ = ("last_message_id",)
|
||||
__slots__ = ("last_message_id", "recipient_ids")
|
||||
|
||||
def __init__(self, data: DMChannelPayload, state: State):
|
||||
super().__init__(data, state)
|
||||
self.last_message_id = data.get("last_message_id")
|
||||
self.recipient_ids: tuple[str, str] = tuple(data["recipients"])
|
||||
self.last_message_id: str | None = 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:
|
||||
@@ -164,33 +186,43 @@ class GroupDMChannel(Channel, Messageable, EditableChannel):
|
||||
The id of the last message in this channel, if any
|
||||
"""
|
||||
|
||||
__slots__ = ("recipients", "name", "owner", "permissions", "icon", "description", "last_message_id")
|
||||
__slots__ = ("recipient_ids", "name", "owner_id", "permissions", "icon", "description", "last_message_id")
|
||||
|
||||
def __init__(self, data: GroupDMChannelPayload, state: State):
|
||||
super().__init__(data, state)
|
||||
self.recipients = [state.get_user(user_id) for user_id in data["recipients"]]
|
||||
self.name = data["name"]
|
||||
self.owner = state.get_user(data["owner"])
|
||||
self.description: Optional[str] = data.get("description")
|
||||
self.last_message_id = data.get("last_message_id")
|
||||
self.recipient_ids: list[str] = data["recipients"]
|
||||
self.name: str = data["name"]
|
||||
self.owner_id: str = data["owner"]
|
||||
self.description: str | None = data.get("description")
|
||||
self.last_message_id: str | None = data.get("last_message_id")
|
||||
|
||||
self.icon: Asset | None
|
||||
|
||||
if icon := data.get("icon"):
|
||||
self.icon = Asset(icon, state)
|
||||
else:
|
||||
self.icon = None
|
||||
|
||||
self.permissions = Permissions(data.get("permissions", 0))
|
||||
self.permissions: Permissions = Permissions(data.get("permissions", 0))
|
||||
|
||||
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None, description: Optional[str] = None):
|
||||
if name:
|
||||
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None, description: Optional[str] = None) -> None:
|
||||
if name is not None:
|
||||
self.name = name
|
||||
|
||||
if recipients:
|
||||
self.recipients = [self.state.get_user(user_id) for user_id in recipients]
|
||||
if recipients is not None:
|
||||
self.recipient_ids = recipients
|
||||
|
||||
if description:
|
||||
if description is not None:
|
||||
self.description = description
|
||||
|
||||
@property
|
||||
def recipients(self) -> list[User]:
|
||||
return [self.state.get_user(user_id) for user_id in self.recipient_ids]
|
||||
|
||||
@property
|
||||
def owner(self) -> User:
|
||||
return self.state.get_user(self.owner_id)
|
||||
|
||||
async def set_default_permissions(self, permissions: Permissions) -> None:
|
||||
"""Sets the default permissions for a group.
|
||||
Parameters
|
||||
@@ -214,16 +246,31 @@ class GroupDMChannel(Channel, Messageable, EditableChannel):
|
||||
|
||||
return self.state.get_message(self.last_message_id)
|
||||
|
||||
class GuildChannel(Channel):
|
||||
def __init__(self, data: GuildChannelPayload, state: State):
|
||||
class ServerChannel(Channel):
|
||||
"""Base class for all guild channels
|
||||
|
||||
Attributes
|
||||
-----------
|
||||
server_id: :class:`str`
|
||||
The id of the server this text channel belongs to
|
||||
name: :class:`str`
|
||||
The name of the text channel
|
||||
description: Optional[:class:`str`]
|
||||
The description of the channel, if any
|
||||
nsfw: bool
|
||||
Sets whether the channel is nsfw or not
|
||||
default_permissions: :class:`ChannelPermissions`
|
||||
The default permissions for all users in the text channel
|
||||
"""
|
||||
def __init__(self, data: ServerChannelPayload, state: State):
|
||||
super().__init__(data, state)
|
||||
|
||||
self.server_id = data["server"]
|
||||
self.name = data["name"]
|
||||
self.server_id: str = data["server"]
|
||||
self.name: str = data["name"]
|
||||
self.description: Optional[str] = data.get("description")
|
||||
self.nsfw = data.get("nsfw", False)
|
||||
self.active = False
|
||||
self.default_permissions = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0}))
|
||||
self.nsfw: bool = data.get("nsfw", False)
|
||||
self.active: bool = False
|
||||
self.default_permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0}))
|
||||
|
||||
permissions: dict[str, PermissionsOverwrite] = {}
|
||||
|
||||
@@ -231,7 +278,10 @@ class GuildChannel(Channel):
|
||||
overwrite = PermissionsOverwrite._from_overwrite(overwrite_data)
|
||||
permissions[role_name] = overwrite
|
||||
|
||||
self.permissions = permissions
|
||||
self.permissions: dict[str, PermissionsOverwrite] = permissions
|
||||
|
||||
self.icon: Asset | None
|
||||
|
||||
if icon := data.get("icon"):
|
||||
self.icon = Asset(icon, state)
|
||||
else:
|
||||
@@ -265,7 +315,7 @@ class GuildChannel(Channel):
|
||||
if description is not None:
|
||||
self.description = description
|
||||
|
||||
if icon:
|
||||
if icon is not None:
|
||||
self.icon = Asset(icon, self.state)
|
||||
|
||||
if nsfw is not None:
|
||||
@@ -284,11 +334,13 @@ class GuildChannel(Channel):
|
||||
self.permissions = permissions
|
||||
|
||||
if default_permissions is not None:
|
||||
self.default_permissions = default_permissions
|
||||
self.default_permissions = PermissionsOverwrite._from_overwrite(default_permissions)
|
||||
|
||||
class TextChannel(GuildChannel, Messageable, EditableChannel):
|
||||
class TextChannel(ServerChannel, Messageable, EditableChannel):
|
||||
"""A text channel
|
||||
|
||||
Subclasses :class:`ServerChannel` and :class:`Messageable`
|
||||
|
||||
Attributes
|
||||
-----------
|
||||
name: :class:`str`
|
||||
@@ -312,7 +364,7 @@ class TextChannel(GuildChannel, Messageable, EditableChannel):
|
||||
def __init__(self, data: TextChannelPayload, state: State):
|
||||
super().__init__(data, state)
|
||||
|
||||
self.last_message_id = data.get("last_message_id")
|
||||
self.last_message_id: str | None = data.get("last_message_id")
|
||||
|
||||
async def _get_channel_id(self) -> str:
|
||||
return self.id
|
||||
@@ -331,9 +383,11 @@ class TextChannel(GuildChannel, Messageable, EditableChannel):
|
||||
|
||||
return self.state.get_message(self.last_message_id)
|
||||
|
||||
class VoiceChannel(GuildChannel, EditableChannel):
|
||||
class VoiceChannel(ServerChannel, EditableChannel):
|
||||
"""A voice channel
|
||||
|
||||
Subclasses :class:`ServerChannel`
|
||||
|
||||
Attributes
|
||||
-----------
|
||||
name: :class:`str`
|
||||
|
||||
+41
-33
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any, Callable, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Callable, Literal, Optional, Union, cast
|
||||
|
||||
import aiohttp
|
||||
|
||||
@@ -12,7 +12,7 @@ from .http import HttpClient
|
||||
from .invite import Invite
|
||||
from .message import Message
|
||||
from .state import State
|
||||
from .utils import Missing
|
||||
from .utils import Missing, Ulid
|
||||
from .websocket import WebsocketHandler
|
||||
|
||||
try:
|
||||
@@ -22,15 +22,15 @@ except ImportError:
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .channel import Channel
|
||||
from .emoji import Emoji
|
||||
from .file import File
|
||||
from .server import Server
|
||||
from .types import ApiInfo
|
||||
from .user import User
|
||||
from .emoji import Emoji
|
||||
from .file import File
|
||||
|
||||
__all__ = ("Client",)
|
||||
|
||||
logger = logging.getLogger("revolt")
|
||||
logger: logging.Logger = logging.getLogger("revolt")
|
||||
|
||||
class Client:
|
||||
"""The client for interacting with revolt
|
||||
@@ -48,11 +48,11 @@ class Client:
|
||||
"""
|
||||
|
||||
def __init__(self, session: aiohttp.ClientSession, token: str, *, api_url: str = "https://api.revolt.chat", max_messages: int = 5000, bot: bool = True):
|
||||
self.session = session
|
||||
self.token = token
|
||||
self.api_url = api_url
|
||||
self.max_messages = max_messages
|
||||
self.bot = bot
|
||||
self.session: aiohttp.ClientSession = session
|
||||
self.token: str = token
|
||||
self.api_url: str = api_url
|
||||
self.max_messages: int = max_messages
|
||||
self.bot: bool = bot
|
||||
|
||||
self.api_info: ApiInfo
|
||||
self.http: HttpClient
|
||||
@@ -63,7 +63,7 @@ class Client:
|
||||
|
||||
super().__init__()
|
||||
|
||||
def dispatch(self, event: str, *args: Any):
|
||||
def dispatch(self, event: str, *args: Any) -> None:
|
||||
"""Dispatch an event, this is typically used for testing and internals.
|
||||
|
||||
Parameters
|
||||
@@ -88,7 +88,7 @@ class Client:
|
||||
async with self.session.get(self.api_url) as resp:
|
||||
return json.loads(await resp.text())
|
||||
|
||||
async def start(self):
|
||||
async def start(self) -> None:
|
||||
"""Starts the client"""
|
||||
api_info = await self.get_api_info()
|
||||
|
||||
@@ -98,7 +98,7 @@ class Client:
|
||||
self.websocket = WebsocketHandler(self.session, self.token, api_info["ws"], self.dispatch, self.state)
|
||||
await self.websocket.start()
|
||||
|
||||
async def stop(self):
|
||||
async def stop(self) -> None:
|
||||
await self.websocket.websocket.close()
|
||||
|
||||
def get_user(self, id: str) -> User:
|
||||
@@ -300,7 +300,7 @@ class Client:
|
||||
|
||||
raise LookupError
|
||||
|
||||
async def edit_self(self, **kwargs: Any):
|
||||
async def edit_self(self, **kwargs: Any) -> None:
|
||||
"""Edits the client's own user
|
||||
|
||||
Parameters
|
||||
@@ -310,13 +310,13 @@ class Client:
|
||||
"""
|
||||
if kwargs.get("avatar", Missing) is None:
|
||||
del kwargs["avatar"]
|
||||
remove = "Avatar"
|
||||
remove = ["Avatar"]
|
||||
else:
|
||||
remove = None
|
||||
|
||||
await self.state.http.edit_self(remove, kwargs)
|
||||
|
||||
async def edit_status(self, **kwargs: Any):
|
||||
async def edit_status(self, **kwargs: Any) -> None:
|
||||
"""Edits the client's own status
|
||||
|
||||
Parameters
|
||||
@@ -328,7 +328,7 @@ class Client:
|
||||
"""
|
||||
if kwargs.get("text", Missing) is None:
|
||||
del kwargs["text"]
|
||||
remove = "StatusText"
|
||||
remove = ["StatusText"]
|
||||
else:
|
||||
remove = None
|
||||
|
||||
@@ -337,7 +337,7 @@ class Client:
|
||||
|
||||
await self.state.http.edit_self(remove, {"status": kwargs})
|
||||
|
||||
async def edit_profile(self, **kwargs: Any):
|
||||
async def edit_profile(self, **kwargs: Any) -> None:
|
||||
"""Edits the client's own profile
|
||||
|
||||
Parameters
|
||||
@@ -347,14 +347,15 @@ class Client:
|
||||
background: Optional[:class:`File`]
|
||||
The new background for the profile, passing in ``None`` will remove the profile background
|
||||
"""
|
||||
remove: list[str] = []
|
||||
|
||||
if kwargs.get("content", Missing) is None:
|
||||
del kwargs["content"]
|
||||
remove = "ProfileContent"
|
||||
elif kwargs.get("background", Missing) is None:
|
||||
remove.append("ProfileContent")
|
||||
|
||||
if kwargs.get("background", Missing) is None:
|
||||
del kwargs["background"]
|
||||
remove = "ProfileBackground"
|
||||
else:
|
||||
remove = None
|
||||
remove.append("ProfileBackground")
|
||||
|
||||
await self.state.http.edit_self(remove, {"profile": kwargs})
|
||||
|
||||
@@ -372,20 +373,27 @@ class Client:
|
||||
The emoji with the corrasponding id
|
||||
"""
|
||||
|
||||
return await self.state.http.fetch_emoji(emoji_id)
|
||||
emoji = await self.state.http.fetch_emoji(emoji_id)
|
||||
|
||||
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False):
|
||||
"""Creates an emoji
|
||||
return Emoji(emoji, self.state)
|
||||
|
||||
async def upload_file(self, file: File, tag: Literal['attachments', 'avatars', 'backgrounds', 'icons', 'banners', 'emojis']) -> Ulid:
|
||||
"""Uploads a file to revolt
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
name: :class:`str`
|
||||
The name for the emoji
|
||||
file: :class:`File`
|
||||
The image for the emoji
|
||||
nsfw: :class:`bool`
|
||||
Whether or not the emoji is nsfw
|
||||
The file to upload
|
||||
tag: :class:`str`
|
||||
The type of file to upload, this should a string of either `'attachments'`, `'avatars'`, `'backgrounds'`, `'icons'`, `'banners'` or `'emojis'`
|
||||
Returns
|
||||
--------
|
||||
:class:`Ulid`
|
||||
The id of the file that was uploaded
|
||||
"""
|
||||
payload = await self.http.create_emoji(name, file, nsfw, {"type": "Detached"})
|
||||
asset = await self.http.upload_file(file, tag)
|
||||
|
||||
return self.state.add_emoji(payload)
|
||||
ulid = Ulid()
|
||||
ulid.id = asset["id"]
|
||||
|
||||
return ulid
|
||||
|
||||
+57
-22
@@ -2,7 +2,9 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional, TypedDict, Union
|
||||
|
||||
from typing_extensions import Unpack, NotRequired
|
||||
from typing_extensions import NotRequired, Unpack
|
||||
|
||||
from revolt.types.embed import WebsiteSpecial
|
||||
|
||||
from .asset import Asset
|
||||
from .enums import EmbedType
|
||||
@@ -14,6 +16,7 @@ if TYPE_CHECKING:
|
||||
from .types import SendableEmbed as SendableEmbedPayload
|
||||
from .types import TextEmbed as TextEmbedPayload
|
||||
from .types import WebsiteEmbed as WebsiteEmbedPayload
|
||||
from .types import JanuaryImage, JanuaryVideo
|
||||
|
||||
__all__ = ("Embed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "to_embed", "SendableEmbed")
|
||||
|
||||
@@ -21,43 +24,45 @@ class WebsiteEmbed:
|
||||
type = EmbedType.website
|
||||
|
||||
def __init__(self, embed: WebsiteEmbedPayload):
|
||||
self.url = embed.get("url")
|
||||
self.special = embed.get("special")
|
||||
self.title = embed.get("title")
|
||||
self.description = embed.get("description")
|
||||
self.image = embed.get("image")
|
||||
self.video = embed.get("video")
|
||||
self.site_name = embed.get("site_name")
|
||||
self.icon_url = embed.get("icon_url")
|
||||
self.colour = embed.get("colour")
|
||||
self.url: str | None = embed.get("url")
|
||||
self.special: WebsiteSpecial | None = embed.get("special")
|
||||
self.title: str | None = embed.get("title")
|
||||
self.description: str | None = embed.get("description")
|
||||
self.image: JanuaryImage | None = embed.get("image")
|
||||
self.video: JanuaryVideo | None = embed.get("video")
|
||||
self.site_name: str | None = embed.get("site_name")
|
||||
self.icon_url: str | None = embed.get("icon_url")
|
||||
self.colour: str | None = embed.get("colour")
|
||||
|
||||
class ImageEmbed:
|
||||
type = EmbedType.image
|
||||
type: EmbedType = EmbedType.image
|
||||
|
||||
def __init__(self, image: ImageEmbedPayload):
|
||||
self.url = image.get("url")
|
||||
self.width = image.get("width")
|
||||
self.height = image.get("height")
|
||||
self.size = image.get("size")
|
||||
self.url: str = image.get("url")
|
||||
self.width: int = image.get("width")
|
||||
self.height: int = image.get("height")
|
||||
self.size: str = image.get("size")
|
||||
|
||||
class TextEmbed:
|
||||
type = EmbedType.text
|
||||
type: EmbedType = EmbedType.text
|
||||
|
||||
def __init__(self, embed: TextEmbedPayload, state: State):
|
||||
self.icon_url = embed.get("icon_url")
|
||||
self.url = embed.get("url")
|
||||
self.title = embed.get("title")
|
||||
self.description = embed.get("description")
|
||||
self.icon_url: str | None = embed.get("icon_url")
|
||||
self.url: str | None = embed.get("url")
|
||||
self.title: str | None = embed.get("title")
|
||||
self.description: str | None = embed.get("description")
|
||||
|
||||
self.media: Asset | None
|
||||
|
||||
if media := embed.get("media"):
|
||||
self.media = Asset(media, state)
|
||||
else:
|
||||
self.media = None
|
||||
|
||||
self.colour = embed.get("colour")
|
||||
self.colour: str | None = embed.get("colour")
|
||||
|
||||
class NoneEmbed:
|
||||
type = EmbedType.none
|
||||
type: EmbedType = EmbedType.none
|
||||
|
||||
Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed]
|
||||
|
||||
@@ -80,6 +85,29 @@ class EmbedParameters(TypedDict):
|
||||
url: NotRequired[str]
|
||||
|
||||
class SendableEmbed:
|
||||
"""
|
||||
Represents an embed that can be sent in a message, you will never receive this, you will receive :class:`Embed`.
|
||||
|
||||
Attributes
|
||||
-----------
|
||||
title: Optional[:class:`str`]
|
||||
The title of the embed
|
||||
|
||||
description: Optional[:class:`str`]
|
||||
The description of the embed
|
||||
|
||||
media: Optional[:class:`str`]
|
||||
The file inside the embed, this is the ID of the file, you can use :meth:`Client.upload_file` to get an ID.
|
||||
|
||||
icon_url: Optional[:class:`str`]
|
||||
The url of the icon url
|
||||
|
||||
colour: Optional[:class:`str`]
|
||||
The embed's accent colour, this is any valid `CSS color <https://developer.mozilla.org/en-US/docs/Web/CSS/color_value>`_
|
||||
|
||||
url: Optional[:class:`str`]
|
||||
URL for hyperlinking the embed's title
|
||||
"""
|
||||
def __init__(self, **attrs: Unpack[EmbedParameters]):
|
||||
self.title: Optional[str] = None
|
||||
self.description: Optional[str] = None
|
||||
@@ -92,6 +120,13 @@ class SendableEmbed:
|
||||
setattr(self, key, value)
|
||||
|
||||
def to_dict(self) -> SendableEmbedPayload:
|
||||
"""Converts the embed to a dictionary which Revolt accepts
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`dict[str, Any]`
|
||||
The embed
|
||||
"""
|
||||
output: SendableEmbedPayload = {"type": "Text"}
|
||||
|
||||
if title := self.title:
|
||||
|
||||
+10
-10
@@ -1,13 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from .utils import Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .server import Server
|
||||
from .state import State
|
||||
from .types import Emoji as EmojiPayload
|
||||
from .server import Server
|
||||
|
||||
__all__ = ("Emoji",)
|
||||
|
||||
@@ -30,16 +30,16 @@ class Emoji(Ulid):
|
||||
The server id this emoji belongs to, if any
|
||||
"""
|
||||
def __init__(self, payload: EmojiPayload, state: State):
|
||||
self.state = state
|
||||
self.state: State = state
|
||||
|
||||
self.id = payload["_id"]
|
||||
self.author_id = payload["creator_id"]
|
||||
self.name = payload["name"]
|
||||
self.animated = payload.get("animated", False)
|
||||
self.nsfw = payload.get("nsfw", False)
|
||||
self.server_id: Optional[str] = payload["parent"].get("id")
|
||||
self.id: str = payload["_id"]
|
||||
self.author_id: str = payload["creator_id"]
|
||||
self.name: str = payload["name"]
|
||||
self.animated: bool = payload.get("animated", False)
|
||||
self.nsfw: bool = payload.get("nsfw", False)
|
||||
self.server_id: str | None = payload["parent"].get("id")
|
||||
|
||||
async def delete(self):
|
||||
async def delete(self) -> None:
|
||||
"""Deletes the emoji."""
|
||||
await self.state.http.delete_emoji(self.id)
|
||||
|
||||
|
||||
+6
-2
@@ -4,6 +4,7 @@ __all__ = (
|
||||
"ServerError",
|
||||
"FeatureDisabled",
|
||||
"AutumnDisabled",
|
||||
"Forbidden",
|
||||
)
|
||||
|
||||
class RevoltError(Exception):
|
||||
@@ -16,7 +17,10 @@ class ServerError(RevoltError):
|
||||
"Internal server error"
|
||||
|
||||
class FeatureDisabled(RevoltError):
|
||||
"""Base class for any feature disabled errors"""
|
||||
"Base class for any feature disabled errors"
|
||||
|
||||
class AutumnDisabled(FeatureDisabled):
|
||||
"""The autumn feature is disabled"""
|
||||
"The autumn feature is disabled"
|
||||
|
||||
class Forbidden(HTTPError):
|
||||
"Missing permissions"
|
||||
|
||||
@@ -1,20 +1,23 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable, Coroutine, TypeVar, Union, cast
|
||||
from typing import Any, Callable, Coroutine, Union, cast
|
||||
from typing_extensions import TypeVar
|
||||
|
||||
import revolt
|
||||
|
||||
from .command import Command
|
||||
from .context import Context
|
||||
from .errors import NotBotOwner, NotServerOwner, ServerOnly
|
||||
from .errors import (MissingPermissionsError, NotBotOwner, NotServerOwner,
|
||||
ServerOnly)
|
||||
from .utils import ClientT
|
||||
|
||||
__all__ = ("check", "Check", "is_bot_owner", "is_server_owner", "has_permissions", "has_channel_permissions")
|
||||
|
||||
__all__ = ("check", "Check", "is_bot_owner", "is_server_owner")
|
||||
|
||||
T = TypeVar("T", Callable[..., Any], Command)
|
||||
T = TypeVar("T", Callable[..., Any], Command, default=Command)
|
||||
|
||||
Check = Callable[[Context[ClientT]], Union[Any, Coroutine[Any, Any, Any]]]
|
||||
|
||||
def check(check: Check[ClientT]):
|
||||
def check(check: Check[ClientT]) -> Callable[[T], T]:
|
||||
"""A decorator for adding command checks
|
||||
|
||||
Parameters
|
||||
@@ -35,7 +38,7 @@ def check(check: Check[ClientT]):
|
||||
|
||||
return inner
|
||||
|
||||
def is_bot_owner():
|
||||
def is_bot_owner() -> Callable[[T], T]:
|
||||
"""A command check for limiting the command to only the bot's owner"""
|
||||
@check
|
||||
def inner(context: Context[ClientT]):
|
||||
@@ -46,11 +49,11 @@ def is_bot_owner():
|
||||
|
||||
return inner
|
||||
|
||||
def is_server_owner():
|
||||
def is_server_owner() -> Callable[[T], T]:
|
||||
"""A command check for limiting the command to only a server's owner"""
|
||||
@check
|
||||
def inner(context: Context[ClientT]):
|
||||
if not context.server:
|
||||
def inner(context: Context[ClientT]) -> bool:
|
||||
if not context.server_id:
|
||||
raise ServerOnly
|
||||
|
||||
if context.author.id == context.server.owner_id:
|
||||
@@ -59,3 +62,30 @@ def is_server_owner():
|
||||
raise NotServerOwner
|
||||
|
||||
return inner
|
||||
|
||||
def has_permissions(**permissions: bool) -> Callable[[T], T]:
|
||||
@check
|
||||
def inner(context: Context[ClientT]) -> bool:
|
||||
author = context.author
|
||||
|
||||
if not author.has_permissions(**permissions):
|
||||
raise MissingPermissionsError(permissions)
|
||||
|
||||
return True
|
||||
|
||||
return inner
|
||||
|
||||
def has_channel_permissions(**permissions: bool) -> Callable[[T], T]:
|
||||
@check
|
||||
def inner(context: Context[ClientT]) -> bool:
|
||||
author = context.author
|
||||
|
||||
if not isinstance(author, revolt.Member):
|
||||
raise ServerOnly
|
||||
|
||||
if not author.has_channel_permissions(context.channel, **permissions):
|
||||
raise MissingPermissionsError(permissions)
|
||||
|
||||
return True
|
||||
|
||||
return inner
|
||||
|
||||
@@ -4,7 +4,7 @@ import sys
|
||||
import traceback
|
||||
from importlib import import_module
|
||||
from typing import (TYPE_CHECKING, Any, Optional, Protocol, TypeVar, Union,
|
||||
runtime_checkable, overload)
|
||||
overload, runtime_checkable)
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
@@ -36,7 +36,7 @@ class ExtensionProtocol(Protocol):
|
||||
class CommandsMeta(type):
|
||||
_commands: list[Command[Any]]
|
||||
|
||||
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any]):
|
||||
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any]) -> Self:
|
||||
commands: list[Command[Any]] = []
|
||||
self = super().__new__(cls, name, bases, attrs)
|
||||
for base in reversed(self.__mro__):
|
||||
@@ -95,6 +95,8 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
for alias in command.aliases:
|
||||
self.all_commands[alias] = command
|
||||
|
||||
self.help_command: HelpCommand[Self] | None
|
||||
|
||||
if help_command is not None:
|
||||
self.help_command = help_command or DefaultHelpCommand[Self]()
|
||||
self.add_command(HelpCommandImpl(self))
|
||||
@@ -137,7 +139,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
"""
|
||||
return self.all_commands[name]
|
||||
|
||||
def add_command(self, command: Command[Self]):
|
||||
def add_command(self, command: Command[Self]) -> None:
|
||||
"""Adds a command, this is typically only used for dynamic commands, you should use the `commands.command` decorator for most usecases.
|
||||
|
||||
Parameters
|
||||
@@ -194,9 +196,6 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
"""
|
||||
content = message.content
|
||||
|
||||
if not isinstance(content, str):
|
||||
return
|
||||
|
||||
prefixes = await self.get_prefix(message)
|
||||
|
||||
if isinstance(prefixes, str):
|
||||
@@ -246,7 +245,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
await command._error_handler(command.cog or self, context, e)
|
||||
self.dispatch("command_error", context, e)
|
||||
|
||||
async def on_command_error(self, ctx: Context[Self], error: Exception, /):
|
||||
async def on_command_error(self, ctx: Context[Self], error: Exception, /) -> None:
|
||||
traceback.print_exception(type(error), error, error.__traceback__)
|
||||
|
||||
on_message = process_commands
|
||||
@@ -266,7 +265,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
|
||||
return True
|
||||
|
||||
def add_cog(self, cog: Cog[Self]):
|
||||
def add_cog(self, cog: Cog[Self]) -> None:
|
||||
"""Adds a cog to the bot, this cog must subclass `Cog`.
|
||||
|
||||
Parameters
|
||||
@@ -294,7 +293,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
|
||||
return cog
|
||||
|
||||
def load_extension(self, name: str):
|
||||
def load_extension(self, name: str) -> None:
|
||||
"""Loads an extension, this takes a module name and runs the setup function inside of it.
|
||||
|
||||
Parameters
|
||||
@@ -310,7 +309,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
self.extensions[name] = extension
|
||||
extension.setup(self)
|
||||
|
||||
def unload_extension(self, name: str):
|
||||
def unload_extension(self, name: str) -> None:
|
||||
"""Unloads an extension, this takes a module name and runs the teardown function inside of it.
|
||||
|
||||
Parameters
|
||||
@@ -325,7 +324,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
|
||||
if teardown := getattr(extension, "teardown", None):
|
||||
teardown(self)
|
||||
|
||||
def reload_extension(self, name: str):
|
||||
def reload_extension(self, name: str) -> None:
|
||||
"""Reloads an extension, this will unload and reload the extension.
|
||||
|
||||
Parameters
|
||||
|
||||
@@ -1,18 +1,18 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Generic, Optional, cast
|
||||
from typing_extensions import Self
|
||||
|
||||
from .command import Command
|
||||
from .utils import ClientT
|
||||
|
||||
|
||||
__all__ = ("Cog", "CogMeta")
|
||||
|
||||
class CogMeta(type, Generic[ClientT]):
|
||||
_commands: list[Command[ClientT]]
|
||||
qualified_name: str
|
||||
|
||||
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any], *, qualified_name: Optional[str] = None):
|
||||
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any], *, qualified_name: Optional[str] = None) -> Self:
|
||||
commands: list[Command[ClientT]] = []
|
||||
self = super().__new__(cls, name, bases, attrs)
|
||||
|
||||
@@ -30,15 +30,15 @@ class Cog(Generic[ClientT], metaclass=CogMeta):
|
||||
_commands: list[Command[ClientT]]
|
||||
qualified_name: str
|
||||
|
||||
def cog_load(self):
|
||||
def cog_load(self) -> None:
|
||||
"""A special method that is called when the cog gets loaded."""
|
||||
pass
|
||||
|
||||
def cog_unload(self):
|
||||
def cog_unload(self) -> None:
|
||||
"""A special method that is called when the cog gets removed."""
|
||||
pass
|
||||
|
||||
def _inject(self, client: ClientT):
|
||||
def _inject(self, client: ClientT) -> None:
|
||||
client.cogs[self.qualified_name] = self
|
||||
|
||||
for command in self._commands:
|
||||
@@ -47,7 +47,7 @@ class Cog(Generic[ClientT], metaclass=CogMeta):
|
||||
|
||||
self.cog_load()
|
||||
|
||||
def _uninject(self, client: ClientT):
|
||||
def _uninject(self, client: ClientT) -> None:
|
||||
for name, command in client.all_commands.copy().items():
|
||||
if command in self._commands:
|
||||
del client.all_commands[name]
|
||||
|
||||
@@ -3,13 +3,14 @@ from __future__ import annotations
|
||||
import inspect
|
||||
import traceback
|
||||
from contextlib import suppress
|
||||
from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine, Generic,
|
||||
Literal, Optional, Union, get_args, get_origin)
|
||||
from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine,
|
||||
Generic, Literal, Optional, Union, get_args, get_origin)
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from revolt.utils import copy_doc, maybe_coroutine
|
||||
|
||||
from .errors import InvalidLiteralArgument, UnionConverterError
|
||||
from .utils import evaluate_parameters, ClientT
|
||||
from .utils import ClientCoT, evaluate_parameters
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .checks import Check
|
||||
@@ -17,15 +18,15 @@ if TYPE_CHECKING:
|
||||
from .context import Context
|
||||
from .group import Group
|
||||
|
||||
__all__ = (
|
||||
__all__: tuple[str, ...] = (
|
||||
"Command",
|
||||
"command"
|
||||
)
|
||||
|
||||
NoneType = type(None)
|
||||
NoneType: type[None] = type(None)
|
||||
P = ParamSpec("P")
|
||||
|
||||
|
||||
class Command(Generic[ClientT]):
|
||||
class Command(Generic[ClientCoT]):
|
||||
"""Class for holding info about a command.
|
||||
|
||||
Parameters
|
||||
@@ -46,19 +47,19 @@ class Command(Generic[ClientT]):
|
||||
__slots__ = ("callback", "name", "aliases", "signature", "checks", "parent", "_error_handler", "cog", "description", "usage", "parameters")
|
||||
|
||||
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str], usage: Optional[str] = None):
|
||||
self.callback = callback
|
||||
self.name = name
|
||||
self.aliases = aliases
|
||||
self.usage = usage
|
||||
self.signature = inspect.signature(self.callback)
|
||||
self.parameters = evaluate_parameters(self.signature.parameters.values(), getattr(callback, "__globals__", {}))
|
||||
self.checks: list[Check[ClientT]] = getattr(callback, "_checks", [])
|
||||
self.parent: Optional[Group[ClientT]] = None
|
||||
self.cog: Optional[Cog[ClientT]] = None
|
||||
self._error_handler: Callable[[Any, Context[ClientT], Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
|
||||
self.description = callback.__doc__
|
||||
self.callback: Callable[..., Coroutine[Any, Any, Any]] = callback
|
||||
self.name: str = name
|
||||
self.aliases: list[str] = aliases
|
||||
self.usage: str | None = usage
|
||||
self.signature: inspect.Signature = inspect.signature(self.callback)
|
||||
self.parameters: list[inspect.Parameter] = evaluate_parameters(self.signature.parameters.values(), getattr(callback, "__globals__", {}))
|
||||
self.checks: list[Check[ClientCoT]] = getattr(callback, "_checks", [])
|
||||
self.parent: Optional[Group[ClientCoT]] = None
|
||||
self.cog: Optional[Cog[ClientCoT]] = None
|
||||
self._error_handler: Callable[[Any, Context[ClientCoT], Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
|
||||
self.description: str | None = callback.__doc__
|
||||
|
||||
async def invoke(self, context: Context[ClientT], *args: Any, **kwargs: Any) -> Any:
|
||||
async def invoke(self, context: Context[ClientCoT], *args: Any, **kwargs: Any) -> Any:
|
||||
"""Runs the command and calls the error handler if the command errors.
|
||||
|
||||
Parameters
|
||||
@@ -74,10 +75,10 @@ class Command(Generic[ClientT]):
|
||||
return await self._error_handler(self.cog or context.client, context, err)
|
||||
|
||||
@copy_doc(invoke)
|
||||
def __call__(self, context: Context[ClientT], *args: Any, **kwargs: Any) -> Any:
|
||||
def __call__(self, context: Context[ClientCoT], *args: Any, **kwargs: Any) -> Any:
|
||||
return self.invoke(context, *args, **kwargs)
|
||||
|
||||
def error(self, func: Callable[..., Coroutine[Any, Any, Any]]):
|
||||
def error(self, func: Callable[..., Coroutine[Any, Any, Any]]) -> Callable[..., Coroutine[Any, Any, Any]]:
|
||||
"""Sets the error handler for the command.
|
||||
|
||||
Parameters
|
||||
@@ -97,11 +98,11 @@ class Command(Generic[ClientT]):
|
||||
self._error_handler = func
|
||||
return func
|
||||
|
||||
async def _default_error_handler(self, ctx: Context[ClientT], error: Exception):
|
||||
async def _default_error_handler(self, ctx: Context[ClientCoT], error: Exception):
|
||||
traceback.print_exception(type(error), error, error.__traceback__)
|
||||
|
||||
@classmethod
|
||||
async def handle_origin(cls, context: Context[ClientT], origin: Any, annotation: Any, arg: str) -> Any:
|
||||
async def handle_origin(cls, context: Context[ClientCoT], origin: Any, annotation: Any, arg: str) -> Any:
|
||||
if origin is Union:
|
||||
for converter in get_args(annotation):
|
||||
try:
|
||||
@@ -128,11 +129,12 @@ class Command(Generic[ClientT]):
|
||||
raise InvalidLiteralArgument(arg)
|
||||
|
||||
@classmethod
|
||||
async def convert_argument(cls, arg: str, annotation: Any, context: Context[ClientT]) -> Any:
|
||||
async def convert_argument(cls, arg: str, annotation: Any, context: Context[ClientCoT]) -> Any:
|
||||
if annotation is not inspect.Signature.empty:
|
||||
if annotation is str: # no converting is needed - its already a string
|
||||
return arg
|
||||
|
||||
origin: Any
|
||||
if origin := get_origin(annotation):
|
||||
return await cls.handle_origin(context, origin, annotation, arg)
|
||||
else:
|
||||
@@ -140,7 +142,7 @@ class Command(Generic[ClientT]):
|
||||
else:
|
||||
return arg
|
||||
|
||||
async def parse_arguments(self, context: Context[ClientT]):
|
||||
async def parse_arguments(self, context: Context[ClientCoT]) -> None:
|
||||
# please pr if you can think of a better way to do this
|
||||
|
||||
for parameter in self.parameters[2:]:
|
||||
@@ -213,8 +215,8 @@ class Command(Generic[ClientT]):
|
||||
|
||||
return f"{' '.join(parents[::-1])} {self.name} {' '.join(parameters)}"
|
||||
|
||||
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command, usage: Optional[str] = None):
|
||||
"""A decorator that turns a function into a :class:`Command`.
|
||||
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientCoT]] = Command, usage: Optional[str] = None) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientCoT]]:
|
||||
"""A decorator that turns a function into a :class:`Command`.n
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
|
||||
@@ -7,16 +7,17 @@ from revolt.utils import maybe_coroutine
|
||||
|
||||
from .command import Command
|
||||
from .group import Group
|
||||
from .utils import ClientT
|
||||
from .utils import ClientCoT
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .view import StringView
|
||||
from revolt.state import State
|
||||
|
||||
__all__ = (
|
||||
"Context",
|
||||
)
|
||||
|
||||
class Context(revolt.Messageable, Generic[ClientT]):
|
||||
class Context(revolt.Messageable, Generic[ClientCoT]):
|
||||
"""Stores metadata the commands execution.
|
||||
|
||||
Attributes
|
||||
@@ -29,7 +30,7 @@ class Context(revolt.Messageable, Generic[ClientT]):
|
||||
The message that was sent to invoke the command
|
||||
channel: :class:`Messageable`
|
||||
The channel the command was invoked in
|
||||
server: :class:`Server`
|
||||
server_id: Optional[:class:`Server`]
|
||||
The server the command was invoked in
|
||||
author: Union[:class:`Member`, :class:`User`]
|
||||
The user or member that invoked the commad, will be :class:`User` in DMs
|
||||
@@ -40,23 +41,37 @@ class Context(revolt.Messageable, Generic[ClientT]):
|
||||
client: :class:`CommandsClient`
|
||||
The revolt client
|
||||
"""
|
||||
__slots__ = ("command", "invoked_with", "args", "message", "server", "channel", "author", "view", "kwargs", "state", "client")
|
||||
__slots__ = ("command", "invoked_with", "args", "message", "channel", "author", "view", "kwargs", "state", "client", "server_id")
|
||||
|
||||
async def _get_channel_id(self) -> str:
|
||||
return self.channel.id
|
||||
|
||||
def __init__(self, command: Optional[Command[ClientT]], invoked_with: str, view: StringView, message: revolt.Message, client: ClientT):
|
||||
self.command = command
|
||||
self.invoked_with = invoked_with
|
||||
self.view = view
|
||||
self.message = message
|
||||
self.client = client
|
||||
def __init__(self, command: Optional[Command[ClientCoT]], invoked_with: str, view: StringView, message: revolt.Message, client: ClientCoT):
|
||||
self.command: Command[ClientCoT] | None = command
|
||||
self.invoked_with: str = invoked_with
|
||||
self.view: StringView = view
|
||||
self.message: revolt.Message = message
|
||||
self.client: ClientCoT = client
|
||||
self.args: list[Any] = []
|
||||
self.kwargs: dict[str, Any] = {}
|
||||
self.server = message.server
|
||||
self.channel = message.channel
|
||||
self.author = message.author
|
||||
self.state = message.state
|
||||
self.server_id: str | None = message.server_id
|
||||
self.channel: revolt.TextChannel | revolt.GroupDMChannel | revolt.DMChannel = message.channel
|
||||
self.author: revolt.Member | revolt.User = message.author
|
||||
self.state: State = message.state
|
||||
|
||||
@property
|
||||
def server(self) -> revolt.Server:
|
||||
""":class:`Server` The server this voice channel belongs too
|
||||
|
||||
Raises
|
||||
-------
|
||||
:class:`LookupError`
|
||||
Raises if the channel is not part of a server
|
||||
"""
|
||||
if not self.server_id:
|
||||
raise LookupError
|
||||
|
||||
return self.state.get_server(self.server_id)
|
||||
|
||||
async def invoke(self) -> Any:
|
||||
"""Invokes the command.
|
||||
@@ -85,8 +100,14 @@ class Context(revolt.Messageable, Generic[ClientT]):
|
||||
await command.parse_arguments(self)
|
||||
return await command.invoke(self, *self.args, **self.kwargs)
|
||||
|
||||
async def can_run(self, command: Optional[Command[ClientT]] = None) -> bool:
|
||||
async def can_run(self, command: Optional[Command[ClientCoT]] = None) -> bool:
|
||||
"""Runs all of the commands checks, and returns true if all of them pass"""
|
||||
command = command or self.command
|
||||
|
||||
return all([await maybe_coroutine(check, self) for check in (command.checks if command else [])])
|
||||
|
||||
async def send_help(self, argument: Command[Any] | Group[Any] | ClientCoT | None = None) -> None:
|
||||
argument = argument or self.client
|
||||
|
||||
command = self.client.get_command("help")
|
||||
await command.invoke(self, argument)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Annotated, TypeVar, TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Annotated, TypeVar
|
||||
|
||||
from revolt import Category, Channel, Member, User, utils
|
||||
|
||||
@@ -13,46 +13,46 @@ from .errors import (BadBoolArgument, CategoryConverterError,
|
||||
if TYPE_CHECKING:
|
||||
from .client import CommandsClient
|
||||
|
||||
__all__ = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter")
|
||||
__all__: tuple[str, ...] = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter")
|
||||
|
||||
channel_regex = re.compile("<#([A-z0-9]{26})>")
|
||||
user_regex = re.compile("<@([A-z0-9]{26})>")
|
||||
channel_regex: re.Pattern[str] = re.compile("<#([A-z0-9]{26})>")
|
||||
user_regex: re.Pattern[str] = re.compile("<@([A-z0-9]{26})>")
|
||||
|
||||
ClientT = TypeVar("ClientT", bound="CommandsClient")
|
||||
|
||||
def bool_converter(arg: str, _):
|
||||
def bool_converter(arg: str, _: Context[ClientT]) -> bool:
|
||||
lowered = arg.lower()
|
||||
if lowered in ["yes", "true", "ye", "y", "1", "on", "enable"]:
|
||||
if lowered in ("yes", "true", "ye", "y", "1", "on", "enable"):
|
||||
return True
|
||||
elif lowered in ('no', 'n', 'false', 'f', '0', 'disable', 'off'):
|
||||
elif lowered in ("no", "false", "n", "f", "0", "off", "disabled"):
|
||||
return False
|
||||
else:
|
||||
raise BadBoolArgument(lowered)
|
||||
|
||||
def category_converter(arg: str, context: Context[ClientT]) -> Category:
|
||||
if not (server := context.server):
|
||||
if not context.server_id:
|
||||
raise ServerOnly
|
||||
|
||||
try:
|
||||
return server.get_category(arg)
|
||||
return context.server.get_category(arg)
|
||||
except KeyError:
|
||||
try:
|
||||
return utils.get(server.categories, name=arg)
|
||||
return utils.get(context.server.categories, name=arg)
|
||||
except LookupError:
|
||||
raise CategoryConverterError(arg)
|
||||
|
||||
def channel_converter(arg: str, context: Context[ClientT]) -> Channel:
|
||||
if not (server := context.server):
|
||||
if not context.server_id:
|
||||
raise ServerOnly
|
||||
|
||||
if (match := channel_regex.match(arg)):
|
||||
arg = match.group(1)
|
||||
|
||||
try:
|
||||
return server.get_channel(arg)
|
||||
return context.server.get_channel(arg)
|
||||
except KeyError:
|
||||
try:
|
||||
return utils.get(server.channels, name=arg)
|
||||
return utils.get(context.server.channels, name=arg)
|
||||
except LookupError:
|
||||
raise ChannelConverterError(arg)
|
||||
|
||||
@@ -69,17 +69,17 @@ def user_converter(arg: str, context: Context[ClientT]) -> User:
|
||||
raise UserConverterError(arg)
|
||||
|
||||
def member_converter(arg: str, context: Context[ClientT]) -> Member:
|
||||
if not (server := context.server):
|
||||
if not context.server_id:
|
||||
raise ServerOnly
|
||||
|
||||
if (match := user_regex.match(arg)):
|
||||
arg = match.group(1)
|
||||
|
||||
try:
|
||||
return server.get_member(arg)
|
||||
return context.server.get_member(arg)
|
||||
except KeyError:
|
||||
try:
|
||||
return utils.get(server.members, name=arg)
|
||||
return utils.get(context.server.members, name=arg)
|
||||
except LookupError:
|
||||
raise MemberConverterError(arg)
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ class CommandNotFound(CommandError):
|
||||
__slots__ = ("command_name",)
|
||||
|
||||
def __init__(self, command_name: str):
|
||||
self.command_name = command_name
|
||||
self.command_name: str = command_name
|
||||
|
||||
class NoClosingQuote(CommandError):
|
||||
"""Raised when there is no closing quote for a command argument"""
|
||||
@@ -49,6 +49,18 @@ class NotServerOwner(CheckError):
|
||||
class ServerOnly(CheckError):
|
||||
"""Raised when a check requires the command to be ran in a server"""
|
||||
|
||||
class MissingPermissionsError(CheckError):
|
||||
"""Raised when a check requires permissions the user does not have
|
||||
|
||||
Attributes
|
||||
-----------
|
||||
permissions: :class:`dict[str, bool]`
|
||||
The permissions which the user did not have
|
||||
"""
|
||||
|
||||
def __init__(self, permissions: dict[str, bool]):
|
||||
self.permissions = permissions
|
||||
|
||||
class ConverterError(CommandError):
|
||||
"""Base class for all converter errors"""
|
||||
|
||||
|
||||
@@ -1,21 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Callable, Coroutine, Optional, TypeVar
|
||||
from typing import Any, Callable, Coroutine, Optional
|
||||
|
||||
from .command import Command
|
||||
from .utils import ClientCoT, ClientT
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .client import CommandsClient
|
||||
|
||||
__all__ = (
|
||||
"Group",
|
||||
"group"
|
||||
)
|
||||
|
||||
ClientT = TypeVar("ClientT", bound="CommandsClient")
|
||||
|
||||
|
||||
class Group(Command[ClientT]):
|
||||
class Group(Command[ClientCoT]):
|
||||
"""Class for holding info about a group command.
|
||||
|
||||
Parameters
|
||||
@@ -30,13 +26,13 @@ class Group(Command[ClientT]):
|
||||
The group's subcommands.
|
||||
"""
|
||||
|
||||
__slots__ = ("subcommands",)
|
||||
__slots__: tuple[str, ...] = ("subcommands",)
|
||||
|
||||
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]):
|
||||
self.subcommands: dict[str, Command[ClientT]] = {}
|
||||
self.subcommands: dict[str, Command[ClientCoT]] = {}
|
||||
super().__init__(callback, name, aliases)
|
||||
|
||||
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command[ClientT]):
|
||||
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientCoT]] = Command[ClientCoT]) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientCoT]]:
|
||||
"""A decorator that turns a function into a :class:`Command` and registers the command as a subcommand.
|
||||
|
||||
Parameters
|
||||
@@ -61,7 +57,7 @@ class Group(Command[ClientT]):
|
||||
|
||||
return inner
|
||||
|
||||
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type[Group[ClientT]]] = None):
|
||||
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type[Group[ClientCoT]]] = None) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientCoT]]:
|
||||
"""A decorator that turns a function into a :class:`Group` and registers the command as a subcommand
|
||||
|
||||
Parameters
|
||||
@@ -92,10 +88,10 @@ class Group(Command[ClientT]):
|
||||
return f"<Group name=\"{self.name}\">"
|
||||
|
||||
@property
|
||||
def commands(self) -> list[Command[ClientT]]:
|
||||
def commands(self) -> list[Command[ClientCoT]]:
|
||||
return list(self.subcommands.values())
|
||||
|
||||
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group[ClientT]] = Group):
|
||||
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group[ClientT]] = Group) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientT]]:
|
||||
"""A decorator that turns a function into a :class:`Group`
|
||||
|
||||
Parameters
|
||||
|
||||
+32
-30
@@ -1,15 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Optional, TypedDict, Union, Generic
|
||||
from typing import TYPE_CHECKING, Any, Generic, Optional, TypedDict, Union
|
||||
|
||||
from typing_extensions import NotRequired
|
||||
|
||||
from .cog import Cog
|
||||
from .command import Command
|
||||
from .context import Context
|
||||
from .group import Group
|
||||
from .utils import ClientT
|
||||
from .cog import Cog
|
||||
from .utils import ClientCoT, ClientT
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from revolt import File, Message, Messageable, MessageReply, SendableEmbed
|
||||
@@ -18,6 +18,7 @@ if TYPE_CHECKING:
|
||||
|
||||
__all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl")
|
||||
|
||||
|
||||
class MessagePayload(TypedDict):
|
||||
content: str
|
||||
embed: NotRequired[SendableEmbed]
|
||||
@@ -25,28 +26,28 @@ class MessagePayload(TypedDict):
|
||||
attachments: NotRequired[list[File]]
|
||||
replies: NotRequired[list[MessageReply]]
|
||||
|
||||
class HelpCommand(ABC, Generic[ClientT]):
|
||||
class HelpCommand(ABC, Generic[ClientCoT]):
|
||||
@abstractmethod
|
||||
async def create_bot_help(self, context: Context[ClientT], commands: dict[Optional[Cog[ClientT]], list[Command[ClientT]]]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def create_command_help(self, context: Context[ClientT], command: Command[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_command_help(self, context: Context[ClientCoT], command: Command[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def create_group_help(self, context: Context[ClientT], group: Group[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_group_help(self, context: Context[ClientCoT], group: Group[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def create_cog_help(self, context: Context[ClientT], cog: Cog[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_cog_help(self, context: Context[ClientCoT], cog: Cog[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
raise NotImplementedError
|
||||
|
||||
async def send_help_command(self, context: Context[ClientT], message_payload: MessagePayload) -> Message:
|
||||
async def send_help_command(self, context: Context[ClientCoT], message_payload: MessagePayload) -> Message:
|
||||
return await context.send(**message_payload)
|
||||
|
||||
async def filter_commands(self, context: Context[ClientT], commands: list[Command[ClientT]]) -> list[Command[ClientT]]:
|
||||
filtered: list[Command[ClientT]] = []
|
||||
async def filter_commands(self, context: Context[ClientCoT], commands: list[Command[ClientCoT]]) -> list[Command[ClientCoT]]:
|
||||
filtered: list[Command[ClientCoT]] = []
|
||||
|
||||
for command in commands:
|
||||
try:
|
||||
@@ -57,34 +58,34 @@ class HelpCommand(ABC, Generic[ClientT]):
|
||||
|
||||
return filtered
|
||||
|
||||
async def group_commands(self, context: Context[ClientT], commands: list[Command[ClientT]]) -> dict[Optional[Cog[ClientT]], list[Command[ClientT]]]:
|
||||
cogs: dict[Optional[Cog[ClientT]], list[Command[ClientT]]] = {}
|
||||
async def group_commands(self, context: Context[ClientCoT], commands: list[Command[ClientCoT]]) -> dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]:
|
||||
cogs: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]] = {}
|
||||
|
||||
for command in commands:
|
||||
cogs.setdefault(command.cog, []).append(command)
|
||||
|
||||
return cogs
|
||||
|
||||
async def handle_message(self, context: Context[ClientT], message: Message):
|
||||
async def handle_message(self, context: Context[ClientCoT], message: Message) -> None:
|
||||
pass
|
||||
|
||||
async def get_channel(self, context: Context[ClientT]) -> Messageable:
|
||||
async def get_channel(self, context: Context) -> Messageable:
|
||||
return context
|
||||
|
||||
@abstractmethod
|
||||
async def handle_no_command_found(self, context: Context[ClientT], name: str) -> Any:
|
||||
async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> Any:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def handle_no_cog_found(self, context: Context[ClientT], name: str) -> Any:
|
||||
async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> Any:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DefaultHelpCommand(HelpCommand[ClientT]):
|
||||
class DefaultHelpCommand(HelpCommand[ClientCoT]):
|
||||
def __init__(self, default_cog_name: str = "No Cog"):
|
||||
self.default_cog_name = default_cog_name
|
||||
|
||||
async def create_bot_help(self, context: Context[ClientT], commands: dict[Optional[Cog[ClientT]], list[Command[ClientT]]]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
lines = ["```"]
|
||||
|
||||
for cog, cog_commands in commands.items():
|
||||
@@ -99,7 +100,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
|
||||
lines.append("```")
|
||||
return "\n".join(lines)
|
||||
|
||||
async def create_cog_help(self, context: Context[ClientT], cog: Cog[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_cog_help(self, context: Context[ClientCoT], cog: Cog[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
lines = ["```"]
|
||||
|
||||
lines.append(f"{cog.qualified_name}:")
|
||||
@@ -110,7 +111,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
|
||||
lines.append("```")
|
||||
return "\n".join(lines)
|
||||
|
||||
async def create_command_help(self, context: Context[ClientT], command: Command[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_command_help(self, context: Context[ClientCoT], command: Command[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
lines = ["```"]
|
||||
|
||||
lines.append(f"{command.name}:")
|
||||
@@ -126,7 +127,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
|
||||
lines.append("```")
|
||||
return "\n".join(lines)
|
||||
|
||||
async def create_group_help(self, context: Context[ClientT], group: Group[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
async def create_group_help(self, context: Context[ClientCoT], group: Group[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
|
||||
lines = ["```"]
|
||||
|
||||
lines.append(f"{group.name}:")
|
||||
@@ -144,27 +145,27 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
|
||||
lines.append("```")
|
||||
return "\n".join(lines)
|
||||
|
||||
async def handle_no_command_found(self, context: Context[ClientT], name: str):
|
||||
async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> None:
|
||||
channel = await self.get_channel(context)
|
||||
await channel.send(f"Command `{name}` not found.")
|
||||
|
||||
async def handle_no_cog_found(self, context: Context[ClientT], name: str):
|
||||
async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> None:
|
||||
channel = await self.get_channel(context)
|
||||
await channel.send(f"Cog `{name}` not found.")
|
||||
|
||||
|
||||
class HelpCommandImpl(Command[ClientT]):
|
||||
def __init__(self, client: ClientT):
|
||||
class HelpCommandImpl(Command[ClientCoT]):
|
||||
def __init__(self, client: ClientCoT):
|
||||
self.client = client
|
||||
|
||||
async def callback(_: Union[ClientT, Cog[ClientT]], context: Context[ClientT], *args: str):
|
||||
async def callback(_: Union[ClientCoT, Cog[ClientCoT]], context: Context[ClientCoT], *args: str) -> None:
|
||||
await help_command_impl(context.client, context, *args)
|
||||
|
||||
super().__init__(callback=callback, name="help", aliases=[])
|
||||
self.description = "Shows help for a command, cog or the entire bot"
|
||||
self.description: str | None = "Shows help for a command, cog or the entire bot"
|
||||
|
||||
|
||||
async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments: str):
|
||||
async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments: str) -> None:
|
||||
help_command = self.help_command
|
||||
|
||||
if not help_command:
|
||||
@@ -201,4 +202,5 @@ async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments
|
||||
else:
|
||||
msg_payload = payload
|
||||
|
||||
await help_command.send_help_command(context, msg_payload)
|
||||
message = await help_command.send_help_command(context, msg_payload)
|
||||
await help_command.handle_message(context, message)
|
||||
|
||||
@@ -1,17 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from inspect import Parameter
|
||||
from typing import Any, Iterable, TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any, Iterable
|
||||
|
||||
from typing_extensions import TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .client import CommandsClient
|
||||
from .context import Context
|
||||
|
||||
|
||||
__all__ = ("evaluate_parameters",)
|
||||
|
||||
ClientT = TypeVar("ClientT", bound="CommandsClient", default="CommandsClient")
|
||||
|
||||
ClientCoT = TypeVar("ClientCoT", bound="CommandsClient", default="CommandsClient", covariant=True)
|
||||
ContextT = TypeVar("ContextT", bound="Context")
|
||||
|
||||
def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]:
|
||||
new_parameters: list[Parameter] = []
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from typing import Iterator
|
||||
from .errors import NoClosingQuote
|
||||
|
||||
|
||||
class StringView:
|
||||
def __init__(self, string: str):
|
||||
self.value = iter(string)
|
||||
self.temp = ""
|
||||
self.should_undo = False
|
||||
self.value: Iterator[str] = iter(string)
|
||||
self.temp: str = ""
|
||||
self.should_undo: bool = False
|
||||
|
||||
def undo(self):
|
||||
def undo(self) -> None:
|
||||
self.should_undo = True
|
||||
|
||||
def next_char(self) -> str:
|
||||
|
||||
+8
-4
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
from typing import Optional, Union
|
||||
from typing import Optional, Union, cast
|
||||
|
||||
__all__ = ("File",)
|
||||
|
||||
@@ -18,17 +20,19 @@ class File:
|
||||
__slots__ = ("f", "spoiler", "filename")
|
||||
|
||||
def __init__(self, file: Union[str, bytes], *, filename: Optional[str] = None, spoiler: bool = False):
|
||||
self.f: io.BufferedIOBase
|
||||
|
||||
if isinstance(file, str):
|
||||
self.f = open(file, "rb")
|
||||
else:
|
||||
self.f = io.BytesIO(file)
|
||||
|
||||
if filename is None and isinstance(file, str):
|
||||
filename = self.f.name
|
||||
filename = cast(Optional[str], self.f.name)
|
||||
|
||||
self.spoiler = spoiler or (filename and filename.startswith("SPOILER_"))
|
||||
self.spoiler: bool = spoiler or (bool(filename) and filename.startswith("SPOILER_"))
|
||||
|
||||
if self.spoiler and (filename and not filename.startswith("SPOILER_")):
|
||||
filename = f"SPOILER_{filename}"
|
||||
|
||||
self.filename = filename
|
||||
self.filename: str | None = filename
|
||||
|
||||
+12
-5
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Callable, Iterator, Optional, Union, overload
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
__all__ = ("Flag", "Flags", "UserBadges")
|
||||
@@ -10,8 +11,8 @@ class Flag:
|
||||
__slots__ = ("flag", "__doc__")
|
||||
|
||||
def __init__(self, func: Callable[[], int]):
|
||||
self.flag = func()
|
||||
self.__doc__ = func.__doc__
|
||||
self.flag: int = func()
|
||||
self.__doc__: str | None = func.__doc__
|
||||
|
||||
@overload
|
||||
def __get__(self: Self, instance: None, owner: type[Flags]) -> Self:
|
||||
@@ -27,7 +28,7 @@ class Flag:
|
||||
|
||||
return instance._check_flag(self.flag)
|
||||
|
||||
def __set__(self, instance: Flags, value: bool):
|
||||
def __set__(self, instance: Flags, value: bool) -> None:
|
||||
instance._set_flag(self.flag, value)
|
||||
|
||||
class Flags:
|
||||
@@ -36,6 +37,12 @@ class Flags:
|
||||
def __init_subclass__(cls) -> None:
|
||||
cls.FLAG_NAMES = []
|
||||
|
||||
for name in dir(cls):
|
||||
value = getattr(cls, name)
|
||||
|
||||
if isinstance(value, Flag):
|
||||
cls.FLAG_NAMES.append(name)
|
||||
|
||||
def __init__(self, value: int = 0, **flags: bool):
|
||||
self.value = value
|
||||
|
||||
@@ -51,7 +58,7 @@ class Flags:
|
||||
def _check_flag(self, flag: int) -> bool:
|
||||
return (self.value & flag) == flag
|
||||
|
||||
def _set_flag(self, flag: int, value: bool):
|
||||
def _set_flag(self, flag: int, value: bool) -> None:
|
||||
if value:
|
||||
self.value |= flag
|
||||
else:
|
||||
@@ -84,7 +91,7 @@ class Flags:
|
||||
def __gt__(self, other: Self) -> bool:
|
||||
return self.value > other.value
|
||||
|
||||
def __repr__(self):
|
||||
def __repr__(self) -> str:
|
||||
return f"<{self.__class__.__name__} value={self.value}>"
|
||||
|
||||
def __iter__(self) -> Iterator[tuple[str, bool]]:
|
||||
|
||||
+40
-36
@@ -6,9 +6,7 @@ from typing import (TYPE_CHECKING, Any, Coroutine, Literal, Optional, TypeVar,
|
||||
import aiohttp
|
||||
import ulid
|
||||
|
||||
from revolt.utils import Missing
|
||||
|
||||
from .errors import HTTPError, ServerError
|
||||
from .errors import Forbidden, HTTPError, ServerError
|
||||
from .file import File
|
||||
|
||||
try:
|
||||
@@ -21,19 +19,23 @@ if TYPE_CHECKING:
|
||||
|
||||
from .enums import SortType
|
||||
from .file import File
|
||||
from .types import (MessageReplyPayload, MessageWithUserData,
|
||||
PartialInvite, Role, EmojiParent, Member, ApiInfo)
|
||||
from .types import ApiInfo
|
||||
from .types import Autumn as AutumnPayload
|
||||
from .types import Channel, DMChannel
|
||||
from .types import GetServerMembers, GroupDMChannel, Invite
|
||||
from .types import Emoji as EmojiPayload
|
||||
from .types import EmojiParent, GetServerMembers, GroupDMChannel
|
||||
from .types import Interactions as InteractionsPayload
|
||||
from .types import Invite
|
||||
from .types import Masquerade as MasqueradePayload
|
||||
from .types import Member
|
||||
from .types import Member as MemberPayload
|
||||
from .types import Message as MessagePayload
|
||||
from .types import (MessageReplyPayload, MessageWithUserData,
|
||||
PartialInvite, Role)
|
||||
from .types import SendableEmbed as SendableEmbedPayload
|
||||
from .types import Server, ServerBans, TextChannel
|
||||
from .types import User as UserPayload
|
||||
from .types import UserProfile, VoiceChannel
|
||||
from .types import Interactions as InteractionsPayload
|
||||
from .types import Emoji as EmojiPayload
|
||||
|
||||
__all__ = ("HttpClient",)
|
||||
|
||||
@@ -44,11 +46,11 @@ class HttpClient:
|
||||
__slots__ = ("session", "token", "api_url", "api_info", "auth_header")
|
||||
|
||||
def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo, bot: bool = True):
|
||||
self.session = session
|
||||
self.token = token
|
||||
self.api_url = api_url
|
||||
self.api_info = api_info
|
||||
self.auth_header = "x-bot-token" if bot else "x-session-token"
|
||||
self.session: aiohttp.ClientSession = session
|
||||
self.token: str = token
|
||||
self.api_url: str = api_url
|
||||
self.api_info: ApiInfo = api_info
|
||||
self.auth_header: str = "x-bot-token" if bot else "x-session-token"
|
||||
|
||||
async def request(self, method: Literal["GET", "POST", "PUT", "DELETE", "PATCH"], route: str, *, json: Optional[dict[str, Any]] = None, nonce: bool = True, params: Optional[dict[str, Any]] = None) -> Any:
|
||||
url = f"{self.api_url}{route}"
|
||||
@@ -87,6 +89,8 @@ class HttpClient:
|
||||
|
||||
if 200 <= resp_code <= 300:
|
||||
return response
|
||||
elif resp_code == 401:
|
||||
raise Forbidden("401: Missing Permissions")
|
||||
else:
|
||||
raise HTTPError(resp_code)
|
||||
|
||||
@@ -198,7 +202,7 @@ class HttpClient:
|
||||
include_users: bool = False
|
||||
) -> Request[Union[list[MessagePayload], MessageWithUserData]]:
|
||||
|
||||
json = {"sort": sort.value, "include_users": str(include_users)}
|
||||
json: dict[str, Any] = {"sort": sort.value, "include_users": str(include_users)}
|
||||
|
||||
if limit:
|
||||
json["limit"] = limit
|
||||
@@ -355,19 +359,19 @@ class HttpClient:
|
||||
def delete_invite(self, code: str) -> Request[None]:
|
||||
return self.request("DELETE", f"/invites/{code}")
|
||||
|
||||
def edit_channel(self, channel_id: str, remove: Optional[str], values: dict[str, Any]):
|
||||
def edit_channel(self, channel_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
|
||||
if remove:
|
||||
values["remove"] = remove
|
||||
|
||||
return self.request("PATCH", f"/channels/{channel_id}", json=values)
|
||||
|
||||
def edit_role(self, server_id: str, role_id: str, remove: Optional[str], values: dict[str, Any]):
|
||||
def edit_role(self, server_id: str, role_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
|
||||
if remove:
|
||||
values["remove"] = remove
|
||||
|
||||
return self.request("PATCH", f"/servers/{server_id}/roles/{role_id}", json=values)
|
||||
|
||||
async def edit_self(self, remove: Optional[str], values: dict[str, Any]):
|
||||
async def edit_self(self, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
|
||||
if remove:
|
||||
values["remove"] = remove
|
||||
|
||||
@@ -380,12 +384,6 @@ class HttpClient:
|
||||
asset = await self.upload_file(background, "backgrounds")
|
||||
profile["background"] = asset["id"]
|
||||
|
||||
if not values.get("profile", Missing):
|
||||
del values["profile"]
|
||||
|
||||
if not values.get("status", Missing):
|
||||
del values["status"]
|
||||
|
||||
return await self.request("PATCH", "/users/@me", json=values)
|
||||
|
||||
def set_guild_channel_default_permissions(self, channel_id: str, allow: int, deny: int) -> Request[None]:
|
||||
@@ -394,31 +392,37 @@ class HttpClient:
|
||||
def set_guild_channel_role_permissions(self, channel_id: str, role_id: str, allow: int, deny: int) -> Request[None]:
|
||||
return self.request("PUT", f"/channels/{channel_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}})
|
||||
|
||||
def set_group_channel_default_permissions(self, channel_id: str, value: int):
|
||||
def set_group_channel_default_permissions(self, channel_id: str, value: int) -> Request[None]:
|
||||
return self.request("PUT", f"/channels/{channel_id}/permissions/default", json={"permissions": value})
|
||||
|
||||
def set_server_role_permissions(self, server_id: str, role_id: str, allow: int, deny: int):
|
||||
return self.request("PUT", f"/server/{server_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}})
|
||||
def set_server_role_permissions(self, server_id: str, role_id: str, allow: int, deny: int) -> Request[None]:
|
||||
return self.request("PUT", f"/servers/{server_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}})
|
||||
|
||||
def set_server_default_permissions(self, server_id: str, value: int):
|
||||
return self.request("PUT", f"/server/{server_id}/permissions/default", json={"permissions": value})
|
||||
def set_server_default_permissions(self, server_id: str, value: int) -> Request[None]:
|
||||
return self.request("PUT", f"/servers/{server_id}/permissions/default", json={"permissions": value})
|
||||
|
||||
def add_reaction(self, channel_id: str, message_id: str, emoji: str):
|
||||
return self.request("PUT", f"/channel/{channel_id}/message/{message_id}/reactions/{emoji}")
|
||||
def add_reaction(self, channel_id: str, message_id: str, emoji: str) -> Request[None]:
|
||||
return self.request("PUT", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}")
|
||||
|
||||
def remove_reaction(self, channel_id: str, message_id: str, emoji: str, user_id: Optional[str], remove_all: bool):
|
||||
return self.request("PUT", f"/channel/{channel_id}/message/{message_id}/reactions/{emoji}")
|
||||
def remove_reaction(self, channel_id: str, message_id: str, emoji: str, user_id: Optional[str], remove_all: bool) -> Request[None]:
|
||||
return self.request("PUT", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}")
|
||||
|
||||
def remove_all_reactions(self, channel_id: str, message_id: str):
|
||||
return self.request("DELETE", f"/channel/{channel_id}/message/{message_id}/reactions")
|
||||
def remove_all_reactions(self, channel_id: str, message_id: str) -> Request[None]:
|
||||
return self.request("DELETE", f"/channels/{channel_id}/messages/{message_id}/reactions")
|
||||
|
||||
def delete_emoji(self, emoji_id: str):
|
||||
def delete_emoji(self, emoji_id: str) -> Request[None]:
|
||||
return self.request("DELETE", f"/custom/emoji/{emoji_id}")
|
||||
|
||||
def fetch_emoji(self, emoji_id: str):
|
||||
def fetch_emoji(self, emoji_id: str) -> Request[EmojiPayload]:
|
||||
return self.request("GET", f"/custom/emoji/{emoji_id}")
|
||||
|
||||
async def create_emoji(self, name: str, file: File, nsfw: bool, parent: EmojiParent) -> EmojiPayload:
|
||||
asset = await self.upload_file(file, "emojis")
|
||||
|
||||
return await self.request("PUT", f"/custom/emoji/{asset['id']}", json={"name": name, "parent": parent, "nsfw": nsfw})
|
||||
|
||||
def edit_member(self, server_id: str, member_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[MemberPayload]:
|
||||
return self.request("PATCH", f"/servers/{server_id}/members/{member_id}", json={"remove": remove, **values})
|
||||
|
||||
def delete_messages(self, channel_id: str, messages: list[str]) -> Request[None]:
|
||||
return self.request("DELETE", f"/channels/{channel_id}/messages/bulk", json={"ids": messages})
|
||||
|
||||
+16
-9
@@ -2,12 +2,17 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
|
||||
from .asset import Asset
|
||||
from .utils import Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .state import State
|
||||
from .channel import Channel
|
||||
from .server import Server
|
||||
from .types import Invite as InvitePayload
|
||||
from .user import User
|
||||
|
||||
|
||||
__all__ = ("Invite",)
|
||||
|
||||
@@ -37,22 +42,24 @@ class Invite(Ulid):
|
||||
__slots__ = ("state", "code", "id", "server", "channel", "user_name", "user_avatar", "user", "member_count")
|
||||
|
||||
def __init__(self, data: InvitePayload, code: str, state: State):
|
||||
self.state = state
|
||||
self.state: State = state
|
||||
|
||||
self.code = code
|
||||
self.id = code
|
||||
self.server = state.get_server(data["server_id"])
|
||||
self.channel = self.server.get_channel(data["channel_id"])
|
||||
self.code: str = code
|
||||
self.id: str = code
|
||||
self.server: Server = state.get_server(data["server_id"])
|
||||
self.channel: Channel = self.server.get_channel(data["channel_id"])
|
||||
|
||||
self.user_name = data["user_name"]
|
||||
self.user = None
|
||||
self.user_name: str = data["user_name"]
|
||||
self.user: User | None = None
|
||||
|
||||
self.user_avatar: Asset | None
|
||||
|
||||
if avatar := data.get("user_avatar"):
|
||||
self.user_avatar = Asset(avatar, state)
|
||||
else:
|
||||
self.user_avatar = None
|
||||
|
||||
self.member_count = data["member_count"]
|
||||
self.member_count: int = data["member_count"]
|
||||
|
||||
@staticmethod
|
||||
def _from_partial(code: str, server: str, creator: str, channel: str, state: State) -> Invite:
|
||||
@@ -69,6 +76,6 @@ class Invite(Ulid):
|
||||
|
||||
return invite
|
||||
|
||||
async def delete(self):
|
||||
async def delete(self) -> None:
|
||||
"""Deletes the invite"""
|
||||
await self.state.http.delete_invite(self.code)
|
||||
|
||||
+148
-20
@@ -1,20 +1,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
import datetime
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
|
||||
from .utils import _Missing, Missing
|
||||
|
||||
from .asset import Asset
|
||||
from .permissions import Permissions
|
||||
from .permissions_calculator import calculate_permissions
|
||||
from .user import User
|
||||
from .file import File
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .channel import Channel
|
||||
from .server import Server
|
||||
from .state import State
|
||||
from .types import File
|
||||
from .types import File as FilePayload
|
||||
from .types import Member as MemberPayload
|
||||
from .role import Role
|
||||
|
||||
__all__ = ("Member",)
|
||||
|
||||
def flattern_user(member: Member, user: User):
|
||||
def flattern_user(member: Member, user: User) -> None:
|
||||
for attr in user.__flattern_attributes__:
|
||||
setattr(member, attr, getattr(user, attr))
|
||||
|
||||
@@ -32,16 +40,18 @@ class Member(User):
|
||||
guild_avatar: Optional[:class:`Asset`]
|
||||
The member's guild avatar if any
|
||||
"""
|
||||
__slots__ = ("state", "nickname", "roles", "server", "guild_avatar", "joined_at", "timeout")
|
||||
__slots__ = ("state", "nickname", "roles", "server", "guild_avatar", "joined_at", "current_timeout")
|
||||
|
||||
def __init__(self, data: MemberPayload, server: Server, state: State):
|
||||
user = state.get_user(data["_id"]["user"])
|
||||
|
||||
# due to not having a user payload and only a user object we have to manually add all the attributes instead of calling User.__init__
|
||||
flattern_user(self, user)
|
||||
user._members.add(self)
|
||||
user._members[server.id] = self
|
||||
|
||||
self.state = state
|
||||
self.state: State = state
|
||||
|
||||
self.guild_avatar: Asset | None
|
||||
|
||||
if avatar := data.get("avatar"):
|
||||
self.guild_avatar = Asset(avatar, state)
|
||||
@@ -49,47 +59,55 @@ class Member(User):
|
||||
self.guild_avatar = None
|
||||
|
||||
roles = [server.get_role(role_id) for role_id in data.get("roles", [])]
|
||||
self.roles = sorted(roles, key=lambda role: role.rank, reverse=True)
|
||||
self.roles: list[Role] = sorted(roles, key=lambda role: role.rank, reverse=True)
|
||||
|
||||
self.server = server
|
||||
self.nickname = data.get("nickname")
|
||||
self.server: Server = server
|
||||
self.nickname: str | None = data.get("nickname")
|
||||
joined_at = data["joined_at"]
|
||||
|
||||
if isinstance(joined_at, int):
|
||||
self.joined_at = datetime.datetime.fromtimestamp(joined_at / 1000)
|
||||
self.joined_at: datetime.datetime = 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
|
||||
self.joined_at: datetime.datetime = datetime.datetime.strptime(joined_at, "%Y-%m-%dT%H:%M:%S.%f%z")
|
||||
|
||||
if timeout := data.get("timeout"):
|
||||
self.timeout = datetime.datetime.strptime(timeout, "%Y-%m-%dT%H:%M:%S.%f%z")
|
||||
self.current_timeout: datetime.datetime | None
|
||||
|
||||
if current_timeout := data.get("timeout"):
|
||||
self.current_timeout = datetime.datetime.strptime(current_timeout, "%Y-%m-%dT%H:%M:%S.%f%z")
|
||||
else:
|
||||
self.current_timeout = None
|
||||
|
||||
@property
|
||||
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
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
""":class:`str` The name the user is displaying, this includes (in order) their masqueraded name, display name and orginal name"""
|
||||
return self.nickname or self.display_name or self.masquerade_name or self.original_name
|
||||
|
||||
@property
|
||||
def mention(self) -> str:
|
||||
""":class:`str`: Returns a string that allows you to mention the given member."""
|
||||
return f"<@{self.id}>"
|
||||
|
||||
def _update(self, *, nickname: Optional[str] = None, avatar: Optional[File] = None, roles: Optional[list[str]] = None):
|
||||
if nickname:
|
||||
def _update(self, *, nickname: Optional[str] = None, avatar: Optional[FilePayload] = None, roles: Optional[list[str]] = None):
|
||||
if nickname is not None:
|
||||
self.nickname = nickname
|
||||
|
||||
if avatar:
|
||||
if avatar is not None:
|
||||
self.guild_avatar = Asset(avatar, self.state)
|
||||
|
||||
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)
|
||||
|
||||
async def kick(self):
|
||||
async def kick(self) -> None:
|
||||
"""Kicks the member from the server"""
|
||||
await self.state.http.kick_member(self.server.id, self.id)
|
||||
|
||||
async def ban(self, *, reason: Optional[str] = None):
|
||||
async def ban(self, *, reason: Optional[str] = None) -> None:
|
||||
"""Bans the member from the server
|
||||
|
||||
Parameters
|
||||
@@ -99,6 +117,116 @@ class Member(User):
|
||||
"""
|
||||
await self.state.http.ban_member(self.server.id, self.id, reason)
|
||||
|
||||
async def unban(self):
|
||||
async def unban(self) -> None:
|
||||
"""Unbans the member from the server"""
|
||||
await self.state.http.unban_member(self.server.id, self.id)
|
||||
|
||||
async def edit(
|
||||
self,
|
||||
*,
|
||||
nickname: str | None | _Missing = Missing,
|
||||
roles: list[Role] | None | _Missing = Missing,
|
||||
avatar: File | None | _Missing = Missing,
|
||||
timeout: datetime.timedelta | None | _Missing = Missing
|
||||
) -> None:
|
||||
remove: list[str] = []
|
||||
data: dict[str, Any] = {}
|
||||
|
||||
if nickname is None:
|
||||
remove.append("Nickname")
|
||||
elif nickname is not Missing:
|
||||
data["nickname"] = nickname
|
||||
|
||||
if roles is None:
|
||||
remove.append("Roles")
|
||||
elif roles is not Missing:
|
||||
data["roles"] = roles
|
||||
|
||||
if avatar is None:
|
||||
remove.append("Avatar")
|
||||
elif avatar is not Missing:
|
||||
# pyright cant understand custom singletons - it doesnt know this will never be an instance of _Missing here because Missing is the only instance
|
||||
assert not isinstance(avatar, _Missing)
|
||||
|
||||
data["avatar"] = (await self.state.http.upload_file(avatar, "avatars"))["id"]
|
||||
|
||||
if timeout is None:
|
||||
remove.append("Timeout")
|
||||
elif timeout is not Missing:
|
||||
assert not isinstance(timeout, _Missing)
|
||||
data["timeout"] = (datetime.datetime.now(datetime.timezone.utc) + timeout).isoformat()
|
||||
|
||||
await self.state.http.edit_member(self.server.id, self.id, remove, data)
|
||||
|
||||
async def timeout(self, length: datetime.timedelta) -> None:
|
||||
"""Timeouts the member
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
length: :class:`datetime.timedelta`
|
||||
The length of the timeout
|
||||
"""
|
||||
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:
|
||||
"""Gets the permissions for the member in the server
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`Permissions`
|
||||
The members permissions
|
||||
"""
|
||||
return calculate_permissions(self, self.server)
|
||||
|
||||
def get_channel_permissions(self, channel: Channel) -> Permissions:
|
||||
"""Gets the permissions for the member in the server taking into account the channel as well
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
channel: :class:`Channel`
|
||||
The channel to calculate permissions with
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`Permissions`
|
||||
The members permissions
|
||||
"""
|
||||
return calculate_permissions(self, channel)
|
||||
|
||||
def has_permissions(self, **permissions: bool) -> bool:
|
||||
"""Computes if the member has the specified permissions
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
permissions: :class:`bool`
|
||||
The permissions to check, this also accepted `False` if you need to check if the member does not have the permission
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`bool`
|
||||
Whether or not they have the permissions
|
||||
"""
|
||||
calculated_perms = self.get_permissions()
|
||||
|
||||
return all([getattr(calculated_perms, key, False) == value for key, value in permissions.items()])
|
||||
|
||||
def has_channel_permissions(self, channel: Channel, **permissions: bool) -> bool:
|
||||
"""Computes if the member has the specified permissions, taking into account the channel as well
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
channel: :class:`Channel`
|
||||
The channel to calculate permissions with
|
||||
permissions: :class:`bool`
|
||||
The permissions to check, this also accepted `False` if you need to check if the member does not have the permission
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`bool`
|
||||
Whether or not they have the permissions
|
||||
"""
|
||||
calculated_perms = self.get_channel_permissions(channel)
|
||||
|
||||
return all([getattr(calculated_perms, key, False) == value for key, value in permissions.items()])
|
||||
|
||||
+62
-38
@@ -1,22 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
from typing import TYPE_CHECKING, Any, Coroutine, Optional, Union
|
||||
|
||||
from revolt.types.message import SystemMessageContent
|
||||
|
||||
from .asset import Asset, PartialAsset
|
||||
from .channel import Messageable
|
||||
from .embed import SendableEmbed, to_embed
|
||||
from .channel import DMChannel, GroupDMChannel, TextChannel
|
||||
from .embed import Embed, SendableEmbed, to_embed
|
||||
from .utils import Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .server import Server
|
||||
from .state import State
|
||||
from .types import Embed as EmbedPayload
|
||||
from .types import Interactions as InteractionsPayload
|
||||
from .types import Masquerade as MasqueradePayload
|
||||
from .types import Message as MessagePayload
|
||||
from .types import Interactions as InteractionsPayload
|
||||
from .types import MessageReplyPayload
|
||||
from .server import Server
|
||||
from .user import User
|
||||
from .member import Member
|
||||
|
||||
__all__ = (
|
||||
"Message",
|
||||
@@ -58,23 +61,37 @@ class Message(Ulid):
|
||||
__slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "mentions", "replies", "reply_ids", "reactions", "interactions")
|
||||
|
||||
def __init__(self, data: MessagePayload, state: State):
|
||||
self.state = state
|
||||
self.state: State = state
|
||||
|
||||
self.id = data["_id"]
|
||||
self.content = data.get("content", "")
|
||||
self.attachments = [Asset(attachment, state) for attachment in data.get("attachments", [])]
|
||||
self.embeds = [to_embed(embed, state) for embed in data.get("embeds", [])]
|
||||
self.id: str = data["_id"]
|
||||
|
||||
content = data.get("content", "")
|
||||
|
||||
if not isinstance(content, str):
|
||||
self.system_content: SystemMessageContent = content
|
||||
self.content: str = ""
|
||||
else:
|
||||
self.content = content
|
||||
|
||||
self.attachments: list[Asset] = [Asset(attachment, state) for attachment in data.get("attachments", [])]
|
||||
self.embeds: list[Embed] = [to_embed(embed, state) for embed in data.get("embeds", [])]
|
||||
|
||||
channel = state.get_channel(data["channel"])
|
||||
assert isinstance(channel, Messageable)
|
||||
self.channel = channel
|
||||
assert isinstance(channel, Union[TextChannel, GroupDMChannel, DMChannel])
|
||||
self.channel: TextChannel | GroupDMChannel | DMChannel = channel
|
||||
|
||||
if server_id := self.channel.server_id:
|
||||
author = state.get_member(server_id, data["author"])
|
||||
self.server_id: str | None = self.channel.server_id
|
||||
|
||||
self.mentions: list[Member | User]
|
||||
|
||||
if self.server_id:
|
||||
author = state.get_member(self.server_id, data["author"])
|
||||
self.mentions = [self.server.get_member(member_id) for member_id in data.get("mentions", [])]
|
||||
else:
|
||||
author = state.get_user(data["author"])
|
||||
self.mentions = [state.get_user(member_id) for member_id in data.get("mentions", [])]
|
||||
|
||||
self.author = author
|
||||
self.author: Member | User = author
|
||||
|
||||
if masquerade := data.get("masquerade"):
|
||||
if name := masquerade.get("name"):
|
||||
@@ -86,11 +103,6 @@ class Message(Ulid):
|
||||
if edited_at := data.get("edited"):
|
||||
self.edited_at: Optional[datetime.datetime] = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z")
|
||||
|
||||
if self.server:
|
||||
self.mentions = [self.server.get_member(member_id) for member_id in data.get("mentions", [])]
|
||||
else:
|
||||
self.mentions = [state.get_user(member_id) for member_id in data.get("mentions", [])]
|
||||
|
||||
self.replies: list[Message] = []
|
||||
self.reply_ids: list[str] = []
|
||||
|
||||
@@ -110,20 +122,26 @@ class Message(Ulid):
|
||||
for emoji, users in reactions.items():
|
||||
self.reactions[emoji] = [self.state.get_user(user_id) for user_id in users]
|
||||
|
||||
self.interactions: MessageInteractions | None
|
||||
|
||||
if interactions := data.get("interactions"):
|
||||
self.interactions = MessageInteractions(reactions=interactions.get("reactions"), restrict_reactions=interactions.get("restrict_reactions", False))
|
||||
else:
|
||||
self.interactions = None
|
||||
|
||||
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited: int):
|
||||
if content:
|
||||
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited: Optional[Union[str, int]] = None):
|
||||
if content is not None:
|
||||
self.content = content
|
||||
|
||||
self.edited = datetime.datetime.fromtimestamp(edited / 1000)
|
||||
|
||||
if embeds:
|
||||
if embeds is not None:
|
||||
self.embeds = [to_embed(embed, self.state) for embed in embeds]
|
||||
|
||||
if edited is not None:
|
||||
if isinstance(edited, int):
|
||||
self.edited_at = datetime.datetime.fromtimestamp(edited / 1000, tz=datetime.timezone.utc)
|
||||
else:
|
||||
self.edited_at = datetime.datetime.strptime(edited, "%Y-%m-%dT%H:%M:%S.%f%z")
|
||||
|
||||
async def edit(self, *, content: Optional[str] = None, embeds: Optional[list[SendableEmbed]] = None) -> None:
|
||||
"""Edits the message. The bot can only edit its own message
|
||||
Parameters
|
||||
@@ -140,7 +158,7 @@ class Message(Ulid):
|
||||
"""Deletes the message. The bot can only delete its own messages and messages it has permission to delete """
|
||||
await self.state.http.delete_message(self.channel.id, self.id)
|
||||
|
||||
def reply(self, *args: Any, mention: bool = False, **kwargs: Any):
|
||||
def reply(self, *args: Any, mention: bool = False, **kwargs: Any) -> Coroutine[Any, Any, Message]:
|
||||
"""Replies to this message, equivilant to:
|
||||
|
||||
.. code-block:: python
|
||||
@@ -150,19 +168,25 @@ class Message(Ulid):
|
||||
"""
|
||||
return self.channel.send(*args, **kwargs, replies=[MessageReply(self, mention)])
|
||||
|
||||
async def add_reaction(self, emoji: str):
|
||||
async def add_reaction(self, emoji: str) -> None:
|
||||
await self.state.http.add_reaction(self.channel.id, self.id, emoji)
|
||||
|
||||
async def remove_reaction(self, emoji: str, user: Optional[User] = None, remove_all: bool = False):
|
||||
async def remove_reaction(self, emoji: str, user: Optional[User] = None, remove_all: bool = False) -> None:
|
||||
await self.state.http.remove_reaction(self.channel.id, self.id, emoji, user.id if user else None, remove_all)
|
||||
|
||||
async def remove_all_reactions(self):
|
||||
async def remove_all_reactions(self) -> None:
|
||||
await self.state.http.remove_all_reactions(self.channel.id, self.id)
|
||||
|
||||
|
||||
@property
|
||||
def server(self) -> Server:
|
||||
""":class:`Server` The server this voice channel belongs too"""
|
||||
""":class:`Server` The server this voice channel belongs too
|
||||
|
||||
Raises
|
||||
-------
|
||||
:class:`LookupError`
|
||||
Raises if the channel is not part of a server
|
||||
"""
|
||||
return self.channel.server
|
||||
|
||||
class MessageReply:
|
||||
@@ -178,8 +202,8 @@ class MessageReply:
|
||||
__slots__ = ("message", "mention")
|
||||
|
||||
def __init__(self, message: Message, mention: bool = False):
|
||||
self.message = message
|
||||
self.mention = mention
|
||||
self.message: Message = message
|
||||
self.mention: bool = mention
|
||||
|
||||
def to_dict(self) -> MessageReplyPayload:
|
||||
return { "id": self.message.id, "mention": self.mention }
|
||||
@@ -199,9 +223,9 @@ class Masquerade:
|
||||
__slots__ = ("name", "avatar", "colour")
|
||||
|
||||
def __init__(self, name: Optional[str] = None, avatar: Optional[str] = None, colour: Optional[str] = None):
|
||||
self.name = name
|
||||
self.avatar = avatar
|
||||
self.colour = colour
|
||||
self.name: str | None = name
|
||||
self.avatar: str | None = avatar
|
||||
self.colour: str | None = colour
|
||||
|
||||
def to_dict(self) -> MasqueradePayload:
|
||||
output: MasqueradePayload = {}
|
||||
@@ -230,10 +254,10 @@ class MessageInteractions:
|
||||
__slots__ = ("reactions", "restrict_reactions")
|
||||
|
||||
def __init__(self, *, reactions: Optional[list[str]] = None, restrict_reactions: bool = False):
|
||||
self.reactions = reactions
|
||||
self.restrict_reactions = restrict_reactions
|
||||
self.reactions: list[str] | None = reactions
|
||||
self.restrict_reactions: bool = restrict_reactions
|
||||
|
||||
def to_dict(self):
|
||||
def to_dict(self) -> InteractionsPayload:
|
||||
output: InteractionsPayload = {}
|
||||
|
||||
if reactions := self.reactions:
|
||||
|
||||
+14
-1
@@ -7,7 +7,7 @@ from .enums import SortType
|
||||
if TYPE_CHECKING:
|
||||
from .embed import SendableEmbed
|
||||
from .file import File
|
||||
from .message import Masquerade, Message, MessageReply, MessageInteractions
|
||||
from .message import Masquerade, Message, MessageInteractions, MessageReply
|
||||
from .state import State
|
||||
|
||||
|
||||
@@ -137,3 +137,16 @@ class Messageable:
|
||||
|
||||
payloads = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after)
|
||||
return [Message(payload, self.state) for payload in payloads]
|
||||
|
||||
async def delete_messages(self, messages: list[Message]) -> None:
|
||||
"""Bulk deletes messages from the channel
|
||||
|
||||
.. note:: The messages must have been sent in the last 7 days.
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
messages: list[:class:`Message`]
|
||||
The messages for deletion, this can be up to 100 messages
|
||||
"""
|
||||
|
||||
await self.state.http.delete_messages(await self._get_channel_id(), [message.id for message in messages])
|
||||
|
||||
+36
-3
@@ -1,13 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
from .flags import Flag, Flags
|
||||
from .types.permissions import Overwrite
|
||||
from .flags import Flags, Flag
|
||||
|
||||
__all__ = ("Permissions", "PermissionsOverwrite")
|
||||
__all__ = ("Permissions", "PermissionsOverwrite", "UserPermissions")
|
||||
|
||||
class UserPermissions(Flags):
|
||||
"""Permissions for users"""
|
||||
|
||||
@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):
|
||||
"""Server permissions for members and roles"""
|
||||
|
||||
@Flag
|
||||
def manage_channel() -> int:
|
||||
return 1 << 0
|
||||
@@ -128,7 +155,13 @@ class Permissions(Flags):
|
||||
def default(cls) -> Self:
|
||||
return cls.default_view_only() | cls(send_messages=True, invite_others=True, send_embeds=True, upload_files=True, connect=True, speak=True)
|
||||
|
||||
@classmethod
|
||||
def default_direct_message(cls) -> Self:
|
||||
return cls.default_view_only() | cls(react=True, manage_channel=True)
|
||||
|
||||
class PermissionsOverwrite:
|
||||
"""A permissions overwrite in a channel"""
|
||||
|
||||
def __init__(self, allow: Permissions, deny: Permissions):
|
||||
self._allow = allow
|
||||
self._deny = deny
|
||||
@@ -143,7 +176,7 @@ class PermissionsOverwrite:
|
||||
|
||||
super().__setattr__(perm, value)
|
||||
|
||||
def __setattr__(self, key: str, value: Any):
|
||||
def __setattr__(self, key: str, value: Any) -> None:
|
||||
if key in Permissions.FLAG_NAMES:
|
||||
if key is True:
|
||||
setattr(self._allow, key, True)
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, cast
|
||||
|
||||
from revolt.enums import ChannelType
|
||||
|
||||
from .permissions import Permissions
|
||||
from .server import Server
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .channel import Channel, DMChannel, GroupDMChannel, ServerChannel
|
||||
from .member import Member
|
||||
|
||||
|
||||
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:
|
||||
target = cast("DMChannel", target)
|
||||
|
||||
user_permissions = target.recipient.get_permissions()
|
||||
|
||||
if user_permissions.send_message:
|
||||
return Permissions.default_direct_message()
|
||||
|
||||
else:
|
||||
return Permissions.default_view_only()
|
||||
|
||||
elif channel_type is ChannelType.group:
|
||||
target = cast("GroupDMChannel", target)
|
||||
|
||||
if target.owner.id != member.id:
|
||||
return Permissions.default_direct_message()
|
||||
else:
|
||||
if target.permissions.value == 0:
|
||||
return Permissions.default_direct_message()
|
||||
else:
|
||||
return target.permissions
|
||||
|
||||
else:
|
||||
target = cast("ServerChannel", target)
|
||||
server = target.server
|
||||
|
||||
if server.owner_id == member.id:
|
||||
return Permissions.all()
|
||||
|
||||
else:
|
||||
perms = calculate_permissions(member, server)
|
||||
perms = (perms | target.default_permissions._allow) & (~target.default_permissions._deny)
|
||||
|
||||
for role in server.roles[::-1]:
|
||||
if overwrite :=target.permissions.get(role.id):
|
||||
perms = (perms | overwrite._allow) & (~overwrite._deny)
|
||||
|
||||
if member.current_timeout and member.current_timeout > datetime.now():
|
||||
perms = perms & Permissions(view_channel=True, read_message_history=True)
|
||||
|
||||
return perms
|
||||
+22
-19
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from .permissions import PermissionsOverwrite
|
||||
from .permissions import Overwrite, PermissionsOverwrite
|
||||
from .utils import Missing, Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -35,20 +35,20 @@ class Role(Ulid):
|
||||
channel_permissions: :class:`ChannelPermissions`
|
||||
The channel permissions for the role
|
||||
"""
|
||||
__slots__ = ("id", "name", "colour", "hoist", "rank", "state", "server", "permissions")
|
||||
__slots__: tuple[str, ...] = ("id", "name", "colour", "hoist", "rank", "state", "server", "permissions")
|
||||
|
||||
def __init__(self, data: RolePayload, role_id: str, server: Server, state: State):
|
||||
self.state = state
|
||||
self.id = role_id
|
||||
self.name = data["name"]
|
||||
self.colour = data.get("colour", None)
|
||||
self.hoist = False
|
||||
self.rank = 0
|
||||
self.server = server
|
||||
self.permissions = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
|
||||
self.state: State = state
|
||||
self.id: str = role_id
|
||||
self.name: str = data["name"]
|
||||
self.colour: str | None = data.get("colour", None)
|
||||
self.hoist: bool = data.get("hoist", False)
|
||||
self.rank: int = data["rank"]
|
||||
self.server: Server = server
|
||||
self.permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
|
||||
|
||||
@property
|
||||
def color(self):
|
||||
def color(self) -> str | None:
|
||||
return self.colour
|
||||
|
||||
async def set_permissions_overwrite(self, *, permissions: PermissionsOverwrite) -> None:
|
||||
@@ -63,24 +63,27 @@ class Role(Ulid):
|
||||
allow, deny = permissions.to_pair()
|
||||
await self.state.http.set_server_role_permissions(self.server.id, self.id, allow.value, deny.value)
|
||||
|
||||
def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None):
|
||||
if name:
|
||||
def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None, permissions: Optional[Overwrite] = None) -> None:
|
||||
if name is not None:
|
||||
self.name = name
|
||||
|
||||
if colour:
|
||||
if colour is not None:
|
||||
self.colour = colour
|
||||
|
||||
if hoist:
|
||||
if hoist is not None:
|
||||
self.hoist = hoist
|
||||
|
||||
if rank:
|
||||
if rank is not None:
|
||||
self.rank = rank
|
||||
|
||||
async def delete(self):
|
||||
if permissions is not None:
|
||||
self.permissions = PermissionsOverwrite._from_overwrite(permissions)
|
||||
|
||||
async def delete(self) -> None:
|
||||
"""Deletes the role"""
|
||||
await self.state.http.delete_role(self.server.id, self.id)
|
||||
|
||||
async def edit(self, **kwargs: Any):
|
||||
async def edit(self, **kwargs: Any) -> None:
|
||||
"""Edits the role
|
||||
|
||||
Parameters
|
||||
@@ -95,7 +98,7 @@ class Role(Ulid):
|
||||
The position of the role
|
||||
"""
|
||||
if kwargs.get("colour", Missing) is None:
|
||||
remove = "Colour"
|
||||
remove = ["Colour"]
|
||||
else:
|
||||
remove = None
|
||||
|
||||
|
||||
+35
-32
@@ -4,14 +4,15 @@ from typing import TYPE_CHECKING, Optional, cast
|
||||
|
||||
from .asset import Asset
|
||||
from .category import Category
|
||||
from .channel import Channel, VoiceChannel
|
||||
from .invite import Invite
|
||||
from .permissions import Permissions
|
||||
from .role import Role
|
||||
from .utils import Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .channel import TextChannel
|
||||
from .channel import Channel, TextChannel, VoiceChannel
|
||||
from .emoji import Emoji
|
||||
from .file import File
|
||||
from .member import Member
|
||||
from .state import State
|
||||
from .types import Ban
|
||||
@@ -19,18 +20,16 @@ if TYPE_CHECKING:
|
||||
from .types import File as FilePayload
|
||||
from .types import Server as ServerPayload
|
||||
from .types import SystemMessagesConfig
|
||||
from .emoji import Emoji
|
||||
from .file import File
|
||||
|
||||
__all__ = ("Server", "SystemMessages", "ServerBan")
|
||||
|
||||
class SystemMessages:
|
||||
def __init__(self, data: SystemMessagesConfig, state: State):
|
||||
self.state = state
|
||||
self.user_joined_id = data.get("user_joined")
|
||||
self.user_left_id = data.get("user_left")
|
||||
self.user_kicked_id = data.get("user_kicked")
|
||||
self.user_banned_id = data.get("user_banned")
|
||||
self.state: State = state
|
||||
self.user_joined_id: str | None = data.get("user_joined")
|
||||
self.user_left_id: str | None = data.get("user_left")
|
||||
self.user_kicked_id: str | None = data.get("user_kicked")
|
||||
self.user_banned_id: str | None = data.get("user_banned")
|
||||
|
||||
@property
|
||||
def user_joined(self) -> Optional[TextChannel]:
|
||||
@@ -95,21 +94,25 @@ class Server(Ulid):
|
||||
__slots__ = ("state", "id", "name", "owner_id", "default_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "_categories", "_emojis")
|
||||
|
||||
def __init__(self, data: ServerPayload, state: State):
|
||||
self.state = state
|
||||
self.id = data["_id"]
|
||||
self.name = data["name"]
|
||||
self.owner_id = data["owner"]
|
||||
self.description = data.get("description") or None
|
||||
self.nsfw = data.get("nsfw", False)
|
||||
self.system_messages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state)
|
||||
self._categories = {data["id"]: Category(data, state) for data in data.get("categories", [])}
|
||||
self.default_permissions = Permissions(data["default_permissions"])
|
||||
self.state: State = state
|
||||
self.id: str = data["_id"]
|
||||
self.name: str = data["name"]
|
||||
self.owner_id: str = data["owner"]
|
||||
self.description: str | None = data.get("description") or None
|
||||
self.nsfw: bool = data.get("nsfw", False)
|
||||
self.system_messages: SystemMessages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state)
|
||||
self._categories: dict[str, Category] = {data["id"]: Category(data, state) for data in data.get("categories", [])}
|
||||
self.default_permissions: Permissions = Permissions(data["default_permissions"])
|
||||
|
||||
self.icon: Asset | None
|
||||
|
||||
if icon := data.get("icon"):
|
||||
self.icon = Asset(icon, state)
|
||||
else:
|
||||
self.icon = None
|
||||
|
||||
self.banner: Asset | None
|
||||
|
||||
if banner := data.get("banner"):
|
||||
self.banner = Asset(banner, state)
|
||||
else:
|
||||
@@ -130,15 +133,15 @@ class Server(Ulid):
|
||||
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):
|
||||
if owner:
|
||||
if owner is not None:
|
||||
self.owner_id = owner
|
||||
if name:
|
||||
if name is not None:
|
||||
self.name = name
|
||||
if description is not None:
|
||||
self.description = description or None
|
||||
if icon:
|
||||
if icon is not None:
|
||||
self.icon = Asset(icon, self.state)
|
||||
if banner:
|
||||
if banner is not None:
|
||||
self.banner = Asset(banner, self.state)
|
||||
if default_permissions is not None:
|
||||
self.default_permissions = Permissions(default_permissions)
|
||||
@@ -280,11 +283,11 @@ class Server(Ulid):
|
||||
|
||||
await self.state.http.set_server_default_permissions(self.id, permissions.value)
|
||||
|
||||
async def leave_server(self):
|
||||
async def leave_server(self) -> None:
|
||||
"""Leaves or deletes the server"""
|
||||
await self.state.http.delete_leave_server(self.id)
|
||||
|
||||
async def delete_server(self):
|
||||
async def delete_server(self) -> None:
|
||||
"""Leaves or deletes a server, alias to :meth`Server.leave_server`"""
|
||||
await self.leave_server()
|
||||
|
||||
@@ -327,10 +330,10 @@ class Server(Ulid):
|
||||
"""
|
||||
payload = await self.state.http.create_channel(self.id, "Voice", name, description)
|
||||
|
||||
channel = VoiceChannel(payload, self.state)
|
||||
channel = self.state.add_channel(payload)
|
||||
self._channels[channel.id] = channel
|
||||
|
||||
return channel
|
||||
return cast("VoiceChannel", channel)
|
||||
|
||||
async def fetch_invites(self) -> list[Invite]:
|
||||
"""Fetches all invites in the server
|
||||
@@ -389,7 +392,7 @@ class Server(Ulid):
|
||||
|
||||
return Role(payload, name, self, self.state)
|
||||
|
||||
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False):
|
||||
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False) -> Emoji:
|
||||
"""Creates an emoji
|
||||
|
||||
Parameters
|
||||
@@ -422,11 +425,11 @@ class ServerBan:
|
||||
__slots__ = ("reason", "server", "user_id", "state")
|
||||
|
||||
def __init__(self, ban: Ban, state: State):
|
||||
self.reason = ban.get("reason")
|
||||
self.server = state.get_server(ban["_id"]["server"])
|
||||
self.user_id = ban["_id"]["user"]
|
||||
self.state = state
|
||||
self.reason: str | None = ban.get("reason")
|
||||
self.server: Server = state.get_server(ban["_id"]["server"])
|
||||
self.user_id: str = ban["_id"]["user"]
|
||||
self.state: State = state
|
||||
|
||||
async def unban(self):
|
||||
async def unban(self) -> None:
|
||||
"""Unbans the user"""
|
||||
await self.state.http.unban_member(self.server.id, self.user_id)
|
||||
|
||||
+17
-9
@@ -4,37 +4,39 @@ from collections import deque
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from .channel import Channel, channel_factory
|
||||
from .emoji import Emoji
|
||||
from .member import Member
|
||||
from .message import Message
|
||||
from .server import Server
|
||||
from .user import User
|
||||
from .emoji import Emoji
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .http import HttpClient
|
||||
from .types import ApiInfo
|
||||
from .types import Channel as ChannelPayload
|
||||
from .types import Emoji as EmojiPayload
|
||||
from .types import Member as MemberPayload
|
||||
from .types import Message as MessagePayload
|
||||
from .types import Server as ServerPayload
|
||||
from .types import User as UserPayload
|
||||
from .types import Emoji as EmojiPayload
|
||||
|
||||
__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", "me")
|
||||
|
||||
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.http: HttpClient = http
|
||||
self.api_info: ApiInfo = api_info
|
||||
self.max_messages: int = max_messages
|
||||
|
||||
self.me: User
|
||||
|
||||
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,7 +61,13 @@ class State:
|
||||
raise LookupError from None
|
||||
|
||||
def add_user(self, payload: UserPayload) -> User:
|
||||
|
||||
|
||||
user = User(payload, self)
|
||||
|
||||
if payload.get("relationship") == "User":
|
||||
self.me = user
|
||||
|
||||
self.users[user.id] = user
|
||||
return user
|
||||
|
||||
@@ -106,7 +114,7 @@ class State:
|
||||
|
||||
raise LookupError
|
||||
|
||||
async def fetch_server_members(self, server_id: str):
|
||||
async def fetch_server_members(self, server_id: str) -> None:
|
||||
data = await self.http.fetch_members(server_id)
|
||||
|
||||
for user in data["users"]:
|
||||
@@ -115,6 +123,6 @@ class State:
|
||||
for member in data["members"]:
|
||||
self.add_member(server_id, member)
|
||||
|
||||
async def fetch_all_server_members(self):
|
||||
async def fetch_all_server_members(self) -> None:
|
||||
for server_id in self.servers:
|
||||
await self.fetch_server_members(server_id)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from .category import *
|
||||
from .channel import *
|
||||
from .embed import *
|
||||
from .emoji import *
|
||||
from .file import *
|
||||
from .gateway import *
|
||||
from .http import *
|
||||
@@ -11,4 +12,3 @@ from .permissions import *
|
||||
from .role import *
|
||||
from .server import *
|
||||
from .user import *
|
||||
from .emoji import *
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from typing import TypedDict
|
||||
|
||||
|
||||
__all__ = ("Category",)
|
||||
|
||||
class Category(TypedDict):
|
||||
|
||||
@@ -14,7 +14,7 @@ __all__ = (
|
||||
"GroupDMChannel",
|
||||
"TextChannel",
|
||||
"VoiceChannel",
|
||||
"GuildChannel",
|
||||
"ServerChannel",
|
||||
"Channel",
|
||||
)
|
||||
|
||||
@@ -64,5 +64,5 @@ class VoiceChannel(BaseChannel):
|
||||
role_permissions: NotRequired[dict[str, Overwrite]]
|
||||
nsfw: NotRequired[bool]
|
||||
|
||||
GuildChannel = Union[TextChannel, VoiceChannel]
|
||||
ServerChannel = Union[TextChannel, VoiceChannel]
|
||||
Channel = Union[SavedMessages, DMChannel, GroupDMChannel, TextChannel, VoiceChannel]
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from typing import Literal, TypedDict, Union
|
||||
|
||||
from typing_extensions import NotRequired
|
||||
|
||||
|
||||
class EmojiParentServer(TypedDict):
|
||||
type: Literal["Server"]
|
||||
id: str
|
||||
|
||||
+14
-6
@@ -1,20 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Literal, TypedDict, Union
|
||||
|
||||
from typing_extensions import NotRequired
|
||||
|
||||
from .permissions import Overwrite
|
||||
from .channel import (Channel, DMChannel, GroupDMChannel, SavedMessages,
|
||||
TextChannel, VoiceChannel)
|
||||
from .message import Message
|
||||
from .permissions import Overwrite
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .category import Category
|
||||
from .member import Member, MemberID
|
||||
from .server import Server, SystemMessagesConfig
|
||||
from .user import User, UserProfile, Status
|
||||
from .embed import Embed
|
||||
from .emoji import Emoji
|
||||
from .file import File
|
||||
from .member import Member, MemberID
|
||||
from .server import Server, SystemMessagesConfig
|
||||
from .user import Status, User, UserProfile
|
||||
|
||||
|
||||
__all__ = (
|
||||
@@ -42,7 +44,8 @@ __all__ = (
|
||||
"ServerCreateEventPayload",
|
||||
"MessageReactEventPayload",
|
||||
"MessageUnreactEventPayload",
|
||||
"MessageRemoveReactionEventPayload"
|
||||
"MessageRemoveReactionEventPayload",
|
||||
"BulkMessageDeleteEventPayload"
|
||||
)
|
||||
|
||||
class BasePayload(TypedDict):
|
||||
@@ -63,7 +66,8 @@ class MessageEventPayload(BasePayload, Message):
|
||||
|
||||
class MessageUpdateData(TypedDict):
|
||||
content: str
|
||||
edited: int
|
||||
embeds: list[Embed]
|
||||
edited: Union[str, int]
|
||||
|
||||
class MessageUpdateEventPayload(BasePayload):
|
||||
channel: str
|
||||
@@ -198,3 +202,7 @@ class MessageRemoveReactionEventPayload(BasePayload):
|
||||
id: str
|
||||
channel_id: str
|
||||
emoji_id: str
|
||||
|
||||
class BulkMessageDeleteEventPayload(BasePayload):
|
||||
channel: str
|
||||
ids: list[str]
|
||||
|
||||
@@ -65,11 +65,13 @@ class Interactions(TypedDict):
|
||||
reactions: NotRequired[list[str]]
|
||||
restrict_reactions: NotRequired[bool]
|
||||
|
||||
SystemMessageContent = Union[UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent]
|
||||
|
||||
class Message(TypedDict):
|
||||
_id: str
|
||||
channel: str
|
||||
author: str
|
||||
content: Union[str, UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent]
|
||||
content: Union[str, SystemMessageContent]
|
||||
attachments: NotRequired[list[File]]
|
||||
embeds: NotRequired[list[Embed]]
|
||||
mentions: NotRequired[list[str]]
|
||||
|
||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import TypedDict
|
||||
|
||||
|
||||
class Overwrite(TypedDict):
|
||||
a: int
|
||||
d: int
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, TypedDict
|
||||
|
||||
from typing_extensions import NotRequired
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -32,6 +32,8 @@ class UserRelation(TypedDict):
|
||||
class User(TypedDict):
|
||||
_id: str
|
||||
username: str
|
||||
discriminator: str
|
||||
display_name: NotRequired[str]
|
||||
avatar: NotRequired[File]
|
||||
relations: NotRequired[list[UserRelation]]
|
||||
badges: NotRequired[int]
|
||||
@@ -40,6 +42,7 @@ class User(TypedDict):
|
||||
online: NotRequired[bool]
|
||||
flags: NotRequired[int]
|
||||
bot: NotRequired[UserBot]
|
||||
privileged: NotRequired[bool]
|
||||
|
||||
class UserProfile(TypedDict, total=False):
|
||||
content: str
|
||||
|
||||
+120
-27
@@ -1,22 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional, Union
|
||||
from weakref import WeakSet
|
||||
from weakref import WeakValueDictionary
|
||||
|
||||
from .asset import Asset, PartialAsset
|
||||
from .channel import DMChannel
|
||||
from .channel import DMChannel, GroupDMChannel
|
||||
from .enums import PresenceType, RelationshipType
|
||||
from .flags import UserBadges
|
||||
from .messageable import Messageable
|
||||
from .permissions import UserPermissions
|
||||
from .utils import Ulid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .member import Member
|
||||
from .state import State
|
||||
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
|
||||
from .server import Server
|
||||
|
||||
__all__ = ("User", "Status", "Relation", "UserProfile")
|
||||
|
||||
@@ -41,11 +43,15 @@ class User(Messageable, Ulid):
|
||||
Attributes
|
||||
-----------
|
||||
id: :class:`str`
|
||||
The users id
|
||||
The user's id
|
||||
discriminator: :class:`str`
|
||||
The user's discriminator
|
||||
display_name: Optional[:class:`str`]
|
||||
The user's display name if they have one
|
||||
bot: :class:`bool`
|
||||
Whether or not the user is a bot
|
||||
owner: Optional[:class:`User`]
|
||||
The bot's owner if the user is a bot
|
||||
owner_id: Optional[:class:`str`]
|
||||
The bot's owner id if the user is a bot
|
||||
badges: :class:`UserBadges`
|
||||
The users badges
|
||||
online: :class:`bool`
|
||||
@@ -60,18 +66,26 @@ 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")
|
||||
__slots__ = (*__flattern_attributes__, "state", "_members")
|
||||
__flattern_attributes__: tuple[str, ...] = ("id", "discriminator", "display_name", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel", "privileged")
|
||||
__slots__: tuple[str, ...] = (*__flattern_attributes__, "state", "_members")
|
||||
|
||||
def __init__(self, data: UserPayload, state: State):
|
||||
self.state = state
|
||||
self._members: WeakSet[Member] = WeakSet() # we store all member versions of this user to avoid having to check every guild when needing to update.
|
||||
self.id = data["_id"]
|
||||
self.original_name = data["username"]
|
||||
self.dm_channel = None
|
||||
self._members: WeakValueDictionary[str, Member] = WeakValueDictionary() # we store all member versions of this user to avoid having to check every guild when needing to update.
|
||||
self.id: str = data["_id"]
|
||||
self.discriminator = data["discriminator"]
|
||||
self.display_name = data.get("display_name")
|
||||
self.original_name: str = data["username"]
|
||||
self.dm_channel: DMChannel | None = None
|
||||
|
||||
bot = data.get("bot")
|
||||
|
||||
self.bot: bool
|
||||
self.owner_id: str | None
|
||||
|
||||
if bot:
|
||||
self.bot = True
|
||||
self.owner_id = bot["owner"]
|
||||
@@ -79,12 +93,13 @@ class User(Messageable, Ulid):
|
||||
self.bot = False
|
||||
self.owner_id = None
|
||||
|
||||
self.badges = UserBadges._from_value(data.get("badges", 0))
|
||||
self.online = data.get("online", False)
|
||||
self.flags = data.get("flags", 0)
|
||||
self.badges: UserBadges = UserBadges._from_value(data.get("badges", 0))
|
||||
self.online: bool = data.get("online", False)
|
||||
self.flags: int = data.get("flags", 0)
|
||||
self.privileged: bool = data.get("privileged", False)
|
||||
|
||||
avatar = data.get("avatar")
|
||||
self.original_avatar = Asset(avatar, state) if avatar else None
|
||||
self.original_avatar: Asset | None = Asset(avatar, state) if avatar else None
|
||||
|
||||
relations: list[Relation] = []
|
||||
|
||||
@@ -92,12 +107,14 @@ class User(Messageable, Ulid):
|
||||
user = state.get_user(relation["_id"])
|
||||
if user:
|
||||
relations.append(Relation(RelationshipType(relation["status"]), user))
|
||||
self.relations = relations
|
||||
self.relations: list[Relation] = relations
|
||||
|
||||
relationship = data.get("relationship")
|
||||
self.relationship = RelationshipType(relationship) if relationship else None
|
||||
self.relationship: RelationshipType | None = RelationshipType(relationship) if relationship else None
|
||||
|
||||
status = data.get("status")
|
||||
self.status: Status | None
|
||||
|
||||
if status:
|
||||
presence = status.get("presence")
|
||||
self.status = Status(status.get("text"), PresenceType(presence) if presence else None) if status else None
|
||||
@@ -109,6 +126,52 @@ class User(Messageable, Ulid):
|
||||
self.masquerade_avatar: Optional[PartialAsset] = None
|
||||
self.masquerade_name: Optional[str] = None
|
||||
|
||||
def get_permissions(self) -> UserPermissions:
|
||||
"""Gets the permissions for the user
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`UserPermissions`
|
||||
The users permissions
|
||||
"""
|
||||
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
|
||||
|
||||
for channel in self.state.channels.values():
|
||||
if (isinstance(channel, (GroupDMChannel, DMChannel)) and self.id in channel.recipient_ids) or any(self.id in (m.id for m in server.members) for server in self.state.servers.values()):
|
||||
if self.state.me.bot or self.bot:
|
||||
permissions.send_message = True
|
||||
|
||||
permissions.access = True
|
||||
permissions.view_profile = True
|
||||
|
||||
return permissions
|
||||
|
||||
def has_permissions(self, **permissions: bool) -> bool:
|
||||
"""Computes if the user has the specified permissions
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
permissions: :class:`bool`
|
||||
The permissions to check, this also accepted `False` if you need to check if the user does not have the permission
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`bool`
|
||||
Whether or not they have the permissions
|
||||
"""
|
||||
perms = self.get_permissions()
|
||||
|
||||
return all([getattr(perms, key, False) == value for key, value in permissions.items()])
|
||||
|
||||
async def _get_channel_id(self):
|
||||
if not self.dm_channel:
|
||||
payload = await self.state.http.open_dm(self.id)
|
||||
@@ -117,18 +180,18 @@ class User(Messageable, Ulid):
|
||||
return self.id
|
||||
|
||||
@property
|
||||
def owner(self) -> Optional[User]:
|
||||
owner_id = self.owner_id
|
||||
def owner(self) -> User:
|
||||
""":class:`User` the owner of the bot account"""
|
||||
|
||||
if not owner_id:
|
||||
return
|
||||
if not self.owner_id:
|
||||
raise LookupError
|
||||
|
||||
return self.state.get_user(owner_id)
|
||||
return self.state.get_user(self.owner_id)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
""":class:`str` The name the user is displaying, this includes there orginal name and masqueraded name"""
|
||||
return self.masquerade_name or self.original_name
|
||||
""":class:`str` The name the user is displaying, this includes (in order) their masqueraded name, display name and orginal name"""
|
||||
return self.display_name or self.masquerade_name or self.original_name
|
||||
|
||||
@property
|
||||
def avatar(self) -> Union[Asset, PartialAsset, None]:
|
||||
@@ -153,7 +216,7 @@ class User(Messageable, Ulid):
|
||||
|
||||
self.profile = UserProfile(profile.get("content"), background)
|
||||
|
||||
if avatar:
|
||||
if avatar is not None:
|
||||
self.original_avatar = Asset(avatar, self.state)
|
||||
|
||||
if online is not None:
|
||||
@@ -162,7 +225,7 @@ class User(Messageable, Ulid):
|
||||
# update user infomation for all members
|
||||
|
||||
if self.__class__ is User:
|
||||
for member in self._members:
|
||||
for member in self._members.values():
|
||||
User._update(member, status=status, profile=profile, avatar=avatar, online=online)
|
||||
|
||||
async def default_avatar(self) -> bytes:
|
||||
@@ -195,3 +258,33 @@ class User(Messageable, Ulid):
|
||||
|
||||
self.profile = UserProfile(payload.get("content"), background)
|
||||
return self.profile
|
||||
|
||||
def to_member(self, server: Server) -> Member:
|
||||
"""Gets the member instance for this user for a specific server.
|
||||
|
||||
Roughly equivelent to:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
member = server.get_member(user.id)
|
||||
|
||||
|
||||
Parameters
|
||||
-----------
|
||||
server: :class:`Server`
|
||||
The server to get the member for
|
||||
|
||||
Returns
|
||||
--------
|
||||
:class:`Member`
|
||||
The member
|
||||
|
||||
Raises
|
||||
-------
|
||||
:class:`LookupError`
|
||||
|
||||
"""
|
||||
try:
|
||||
return self._members[server.id]
|
||||
except IndexError:
|
||||
raise LookupError from None
|
||||
|
||||
+4
-4
@@ -1,23 +1,23 @@
|
||||
import inspect
|
||||
import datetime
|
||||
import inspect
|
||||
from contextlib import asynccontextmanager
|
||||
from operator import attrgetter
|
||||
from typing import Any, Callable, Coroutine, Iterable, Literal, TypeVar, Union
|
||||
import ulid
|
||||
|
||||
import ulid
|
||||
from aiohttp import ClientSession
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
__all__ = ("Missing", "copy_doc", "maybe_coroutine", "get", "client_session")
|
||||
|
||||
class _Missing:
|
||||
def __repr__(self):
|
||||
def __repr__(self) -> str:
|
||||
return "<Missing>"
|
||||
|
||||
def __bool__(self) -> Literal[False]:
|
||||
return False
|
||||
|
||||
Missing = _Missing()
|
||||
Missing: _Missing = _Missing()
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
+119
-70
@@ -1,34 +1,43 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
import logging
|
||||
import time
|
||||
from copy import copy
|
||||
from typing import TYPE_CHECKING, Callable, cast
|
||||
from typing import TYPE_CHECKING, Callable, NamedTuple, cast
|
||||
|
||||
from . import utils
|
||||
from .channel import GroupDMChannel, TextChannel, VoiceChannel
|
||||
from .enums import RelationshipType
|
||||
from .role import Role
|
||||
from .types import (BulkMessageDeleteEventPayload, ChannelCreateEventPayload,
|
||||
ChannelDeleteEventPayload, ChannelDeleteTypingEventPayload,
|
||||
ChannelStartTypingEventPayload, ChannelUpdateEventPayload)
|
||||
from .types import Member as MemberPayload
|
||||
from .types import Message as MessagePayload
|
||||
from .types import MemberID as MemberIDPayload
|
||||
from .types import (MessageDeleteEventPayload, MessageUpdateEventPayload,
|
||||
ServerDeleteEventPayload, ServerMemberJoinEventPayload,
|
||||
from .types import Message as MessagePayload
|
||||
from .types import (MessageDeleteEventPayload, MessageReactEventPayload,
|
||||
MessageRemoveReactionEventPayload,
|
||||
MessageUnreactEventPayload, MessageUpdateEventPayload)
|
||||
from .types import Role as RolePayload
|
||||
from .types import (ServerCreateEventPayload, ServerDeleteEventPayload,
|
||||
ServerMemberJoinEventPayload,
|
||||
ServerMemberLeaveEventPayload,
|
||||
ServerCreateEventPayload,
|
||||
ServerMemberUpdateEventPayload,
|
||||
ServerRoleDeleteEventPayload, ServerRoleUpdateEventPayload,
|
||||
ServerUpdateEventPayload, UserRelationshipEventPayload,
|
||||
UserUpdateEventPayload, MessageReactEventPayload, MessageUnreactEventPayload, MessageRemoveReactionEventPayload, ChannelCreateEventPayload, ChannelDeleteEventPayload,
|
||||
ChannelDeleteTypingEventPayload,
|
||||
ChannelStartTypingEventPayload, ChannelUpdateEventPayload)
|
||||
from .user import Status, UserProfile
|
||||
from . import utils
|
||||
UserUpdateEventPayload)
|
||||
from .user import Status, User, UserProfile
|
||||
|
||||
import aiohttp
|
||||
|
||||
try:
|
||||
import ujson as json
|
||||
except ImportError:
|
||||
import json
|
||||
|
||||
use_msgpack: bool
|
||||
|
||||
try:
|
||||
import msgpack
|
||||
use_msgpack = True
|
||||
@@ -39,46 +48,50 @@ if TYPE_CHECKING:
|
||||
import aiohttp
|
||||
|
||||
from .state import State
|
||||
from .types import AuthenticatePayload, BasePayload
|
||||
from .types import MessageEventPayload, ReadyEventPayload
|
||||
from .types import (AuthenticatePayload, BasePayload, MessageEventPayload,
|
||||
ReadyEventPayload)
|
||||
from .message import Message
|
||||
|
||||
class WSMessage(NamedTuple):
|
||||
type: aiohttp.WSMsgType
|
||||
data: str | bytes | aiohttp.WSCloseCode
|
||||
|
||||
__all__ = ("WebsocketHandler",)
|
||||
__all__: tuple[str, ...] = ("WebsocketHandler",)
|
||||
|
||||
logger = logging.getLogger("revolt")
|
||||
logger: logging.Logger = logging.getLogger("revolt")
|
||||
|
||||
class WebsocketHandler:
|
||||
__slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready", "server_events")
|
||||
|
||||
def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State):
|
||||
self.session = session
|
||||
self.token = token
|
||||
self.ws_url = ws_url
|
||||
self.dispatch = dispatch
|
||||
self.state = state
|
||||
self.session: aiohttp.ClientSession = session
|
||||
self.token: str = token
|
||||
self.ws_url: str = ws_url
|
||||
self.dispatch: Callable[..., None] = dispatch
|
||||
self.state: State = state
|
||||
self.websocket: aiohttp.ClientWebSocketResponse
|
||||
self.loop = asyncio.get_running_loop()
|
||||
self.user = None
|
||||
self.ready = asyncio.Event()
|
||||
self.loop: asyncio.AbstractEventLoop = asyncio.get_running_loop()
|
||||
self.user: User | None = None
|
||||
self.ready: asyncio.Event = asyncio.Event()
|
||||
self.server_events: dict[str, asyncio.Event] = {}
|
||||
|
||||
async def _wait_for_server_ready(self, server_id: str):
|
||||
async def _wait_for_server_ready(self, server_id: str) -> None:
|
||||
if event := self.server_events.get(server_id):
|
||||
await event.wait()
|
||||
|
||||
async def send_payload(self, payload: BasePayload):
|
||||
async def send_payload(self, payload: BasePayload) -> None:
|
||||
if use_msgpack:
|
||||
await self.websocket.send_bytes(msgpack.packb(payload)) # type: ignore
|
||||
await self.websocket.send_bytes(msgpack.packb(payload))
|
||||
else:
|
||||
await self.websocket.send_str(json.dumps(payload))
|
||||
|
||||
async def heartbeat(self):
|
||||
async def heartbeat(self) -> None:
|
||||
while not self.websocket.closed:
|
||||
logger.info("Sending hearbeat")
|
||||
await self.websocket.ping()
|
||||
await asyncio.sleep(15)
|
||||
|
||||
async def send_authenticate(self):
|
||||
async def send_authenticate(self) -> None:
|
||||
payload: AuthenticatePayload = {
|
||||
"type": "Authenticate",
|
||||
"token": self.token
|
||||
@@ -86,7 +99,7 @@ class WebsocketHandler:
|
||||
|
||||
await self.send_payload(payload)
|
||||
|
||||
async def handle_event(self, payload: BasePayload):
|
||||
async def handle_event(self, payload: BasePayload) -> None:
|
||||
event_type = payload["type"].lower()
|
||||
logger.debug("Recieved event %s %s", event_type, payload)
|
||||
try:
|
||||
@@ -95,15 +108,14 @@ class WebsocketHandler:
|
||||
|
||||
func = getattr(self, f"handle_{event_type}")
|
||||
except AttributeError:
|
||||
logger.debug("Unknown event '%s'", event_type)
|
||||
return
|
||||
return logger.debug("Unknown event '%s'", event_type)
|
||||
|
||||
await func(payload)
|
||||
|
||||
async def handle_authenticated(self, _):
|
||||
async def handle_authenticated(self, _: BasePayload) -> None:
|
||||
logger.info("Successfully authenticated")
|
||||
|
||||
async def handle_ready(self, payload: ReadyEventPayload):
|
||||
async def handle_ready(self, payload: ReadyEventPayload) -> None:
|
||||
for user_payload in payload["users"]:
|
||||
user = self.state.add_user(user_payload)
|
||||
|
||||
@@ -128,7 +140,7 @@ class WebsocketHandler:
|
||||
self.ready.set()
|
||||
self.dispatch("ready")
|
||||
|
||||
async def handle_message(self, payload: MessageEventPayload):
|
||||
async def handle_message(self, payload: MessageEventPayload) -> None:
|
||||
if server := self.state.get_channel(payload["channel"]).server_id:
|
||||
await self._wait_for_server_ready(server)
|
||||
|
||||
@@ -137,7 +149,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("message", message)
|
||||
|
||||
async def handle_messageupdate(self, payload: MessageUpdateEventPayload):
|
||||
async def handle_messageupdate(self, payload: MessageUpdateEventPayload) -> None:
|
||||
self.dispatch("raw_message_update", payload)
|
||||
|
||||
try:
|
||||
@@ -152,7 +164,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("message_update", message)
|
||||
|
||||
async def handle_messagedelete(self, payload: MessageDeleteEventPayload):
|
||||
async def handle_messagedelete(self, payload: MessageDeleteEventPayload) -> None:
|
||||
self.dispatch("raw_message_delete", payload)
|
||||
|
||||
try:
|
||||
@@ -168,7 +180,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("message_delete", message)
|
||||
|
||||
async def handle_channelcreate(self, payload: ChannelCreateEventPayload):
|
||||
async def handle_channelcreate(self, payload: ChannelCreateEventPayload) -> None:
|
||||
channel = self.state.add_channel(payload)
|
||||
|
||||
if server_id := channel.server_id:
|
||||
@@ -176,7 +188,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("channel_create", channel)
|
||||
|
||||
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload):
|
||||
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload) -> None:
|
||||
# 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 not (channel := self.state.channels.get(payload["id"], None)):
|
||||
@@ -201,7 +213,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("channel_update", old_channel, channel)
|
||||
|
||||
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload):
|
||||
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload) -> None:
|
||||
channel = self.state.channels.pop(payload["id"])
|
||||
|
||||
if server_id := channel.server_id:
|
||||
@@ -209,7 +221,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("channel_delete", channel)
|
||||
|
||||
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload):
|
||||
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload) -> None:
|
||||
channel = self.state.get_channel(payload["id"])
|
||||
|
||||
if server_id := channel.server_id:
|
||||
@@ -219,7 +231,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("typing_start", channel, user)
|
||||
|
||||
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload):
|
||||
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload) -> None:
|
||||
channel = self.state.get_channel(payload["id"])
|
||||
|
||||
if server_id := channel.server_id:
|
||||
@@ -229,7 +241,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("typing_stop", channel, user)
|
||||
|
||||
async def handle_serverupdate(self, payload: ServerUpdateEventPayload):
|
||||
async def handle_serverupdate(self, payload: ServerUpdateEventPayload) -> None:
|
||||
await self._wait_for_server_ready(payload["id"])
|
||||
|
||||
server = self.state.get_server(payload["id"])
|
||||
@@ -251,7 +263,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("server_update", old_server, server)
|
||||
|
||||
async def handle_serverdelete(self, payload: ServerDeleteEventPayload):
|
||||
async def handle_serverdelete(self, payload: ServerDeleteEventPayload) -> None:
|
||||
server = self.state.servers.pop(payload["id"])
|
||||
|
||||
for channel in server.channels:
|
||||
@@ -261,7 +273,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("server_delete", server)
|
||||
|
||||
async def handle_servercreate(self, payload: ServerCreateEventPayload):
|
||||
async def handle_servercreate(self, payload: ServerCreateEventPayload) -> None:
|
||||
for channel in payload["channels"]:
|
||||
self.state.add_channel(channel)
|
||||
|
||||
@@ -274,7 +286,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("server_join", server)
|
||||
|
||||
async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload):
|
||||
async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload) -> None:
|
||||
await self._wait_for_server_ready(payload["id"]["server"])
|
||||
|
||||
member = self.state.get_member(payload["id"]["server"], payload["id"]["user"])
|
||||
@@ -290,15 +302,17 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("member_update", old_member, member)
|
||||
|
||||
async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload):
|
||||
member = self.state.add_member(payload["id"], MemberPayload(_id=MemberIDPayload(server=payload["id"], user=payload["user"]), joined_at=int(time.time()))) # revolt doesnt give us the joined at time
|
||||
async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload) -> None:
|
||||
# avoid an api request if possible
|
||||
if payload["user"] not in self.state.users:
|
||||
user = await self.state.http.fetch_user(payload["user"])
|
||||
self.state.add_user(user)
|
||||
|
||||
user = await self.state.http.fetch_user(member.id)
|
||||
self.state.add_user(user)
|
||||
member = self.state.add_member(payload["id"], MemberPayload(_id=MemberIDPayload(server=payload["id"], user=payload["user"]), joined_at=int(time.time()))) # revolt doesnt give us the joined at time
|
||||
|
||||
self.dispatch("member_join", member)
|
||||
|
||||
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload):
|
||||
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload) -> None:
|
||||
await self._wait_for_server_ready(payload["id"])
|
||||
|
||||
server = self.state.get_server(payload["id"])
|
||||
@@ -307,26 +321,34 @@ class WebsocketHandler:
|
||||
# remove the member from the user
|
||||
|
||||
user = self.state.get_user(payload["user"])
|
||||
user._members.remove(member)
|
||||
user._members.pop(server.id)
|
||||
|
||||
self.dispatch("member_leave", member)
|
||||
|
||||
async def handle_serveroleupdate(self, payload: ServerRoleUpdateEventPayload):
|
||||
async def handle_serverroleupdate(self, payload: ServerRoleUpdateEventPayload) -> None:
|
||||
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"])
|
||||
|
||||
await self._wait_for_server_ready(server.id)
|
||||
|
||||
self.dispatch("role_update", old_role, role)
|
||||
try:
|
||||
role = server.get_role(payload["role_id"])
|
||||
except LookupError:
|
||||
# the role wasnt found meaning it was just created
|
||||
|
||||
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload):
|
||||
role = Role(cast(RolePayload, payload["data"]), payload["role_id"], server, self.state)
|
||||
server._roles[role.id] = role
|
||||
self.dispatch("role_create", role)
|
||||
else:
|
||||
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) -> None:
|
||||
server = self.state.get_server(payload["id"])
|
||||
role = server._roles.pop(payload["role_id"])
|
||||
|
||||
@@ -334,7 +356,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("role_delete", role)
|
||||
|
||||
async def handle_userupdate(self, payload: UserUpdateEventPayload):
|
||||
async def handle_userupdate(self, payload: UserUpdateEventPayload) -> None:
|
||||
user = self.state.get_user(payload["id"])
|
||||
old_user = copy(user)
|
||||
|
||||
@@ -357,14 +379,14 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("user_update", old_user, user)
|
||||
|
||||
async def handle_userrelationship(self, payload: UserRelationshipEventPayload):
|
||||
async def handle_userrelationship(self, payload: UserRelationshipEventPayload) -> None:
|
||||
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 handle_messagereact(self, payload: MessageReactEventPayload):
|
||||
async def handle_messagereact(self, payload: MessageReactEventPayload) -> None:
|
||||
if server := self.state.get_channel(payload["channel_id"]).server_id:
|
||||
await self._wait_for_server_ready(server)
|
||||
|
||||
@@ -381,7 +403,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("reaction_add", message, user, emoji_id)
|
||||
|
||||
async def handle_messageunreact(self, payload: MessageUnreactEventPayload):
|
||||
async def handle_messageunreact(self, payload: MessageUnreactEventPayload) -> None:
|
||||
if server := self.state.get_channel(payload["channel_id"]).server_id:
|
||||
await self._wait_for_server_ready(server)
|
||||
|
||||
@@ -397,7 +419,7 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("reaction_remove", message, user, payload["emoji_id"])
|
||||
|
||||
async def handle_messageremovereaction(self, payload: MessageRemoveReactionEventPayload):
|
||||
async def handle_messageremovereaction(self, payload: MessageRemoveReactionEventPayload) -> None:
|
||||
if server := self.state.get_channel(payload["channel_id"]).server_id:
|
||||
await self._wait_for_server_ready(server)
|
||||
|
||||
@@ -412,17 +434,44 @@ class WebsocketHandler:
|
||||
|
||||
self.dispatch("reaction_clear", message, users, payload["emoji_id"])
|
||||
|
||||
async def start(self):
|
||||
async def handle_bulkmessagedelete(self, payload: BulkMessageDeleteEventPayload) -> None:
|
||||
channel = self.state.get_channel(payload["channel"])
|
||||
|
||||
self.dispatch("raw_bulk_message_delete", payload)
|
||||
|
||||
messages: list[Message] = []
|
||||
|
||||
for message_id in payload["ids"]:
|
||||
if server_id := channel.server_id:
|
||||
await self._wait_for_server_ready(server_id)
|
||||
|
||||
self.dispatch("raw_message_delete", MessageDeleteEventPayload(type="messagedelete", channel=payload["channel"], id=message_id))
|
||||
|
||||
try:
|
||||
message = self.state.get_message(message_id)
|
||||
except LookupError:
|
||||
pass
|
||||
else:
|
||||
self.state.messages.remove(message)
|
||||
self.dispatch("message_delete", message)
|
||||
|
||||
messages.append(message)
|
||||
|
||||
self.dispatch("bulk_message_delete", messages)
|
||||
|
||||
async def start(self) -> None:
|
||||
if use_msgpack:
|
||||
url = f"{self.ws_url}?format=msgpack"
|
||||
else:
|
||||
url = f"{self.ws_url}?format=json"
|
||||
|
||||
self.websocket = await self.session.ws_connect(url)
|
||||
self.websocket = await self.session.ws_connect(url) # type: ignore
|
||||
await self.send_authenticate()
|
||||
asyncio.create_task(self.heartbeat())
|
||||
|
||||
async for msg in self.websocket:
|
||||
msg = cast(WSMessage, msg) # aiohttp doesnt use NamedTuple so the type info is missing
|
||||
|
||||
if use_msgpack:
|
||||
data = cast(bytes, msg.data)
|
||||
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||
|
||||
from typing_extensions import Protocol
|
||||
|
||||
class _FileLike(Protocol):
|
||||
def read(self, n: int) -> bytes: ...
|
||||
|
||||
def unpackb(
|
||||
packed: bytes,
|
||||
file_like: Optional[_FileLike] = ...,
|
||||
read_size: int = ...,
|
||||
use_list: bool = ...,
|
||||
raw: bool = ...,
|
||||
timestamp: int = ...,
|
||||
strict_map_key: bool = ...,
|
||||
object_hook: Optional[Callable[[Dict[Any, Any]], Any]] = ...,
|
||||
object_pairs_hook: Optional[Callable[[List[Tuple[Any, Any]]], Any]] = ...,
|
||||
list_hook: Optional[Callable[[List[Any]], Any]] = ...,
|
||||
unicode_errors: Optional[str] = ...,
|
||||
max_buffer_size: int = ...,
|
||||
ext_hook: Callable[[int, bytes], Any] = ...,
|
||||
max_str_len: int = ...,
|
||||
max_bin_len: int = ...,
|
||||
max_array_len: int = ...,
|
||||
max_map_len: int = ...,
|
||||
max_ext_len: int = ...,
|
||||
) -> Any: ...
|
||||
|
||||
def packb(
|
||||
o: Any,
|
||||
default: Optional[Callable[[Any], Any]] = ...,
|
||||
use_single_float: bool = ...,
|
||||
autoreset: bool = ...,
|
||||
use_bin_type: bool = ...,
|
||||
strict_types: bool = ...,
|
||||
datetime: bool = ...,
|
||||
unicode_errors: Optional[str] = ...,
|
||||
) -> bytes: ...
|
||||
@@ -0,0 +1 @@
|
||||
def get_html_theme_path() -> str: ...
|
||||
Reference in New Issue
Block a user