34 Commits

Author SHA1 Message Date
Zomatree 4e70b8076c bump version 2023-06-13 03:34:42 +01:00
Zomatree 828b428f48 Add display_name and discriminator 2023-06-13 03:10:45 +01:00
Zomatree 4efaac40f6 fix FLAG_NAMES not being set 2023-06-04 15:22:21 +01:00
Zomatree 78f14fa9a0 Revert: test CI 2023-05-20 13:09:59 +01:00
Zomatree 300169d71a Add back type checking to CI 2023-05-20 13:07:21 +01:00
Zomatree 983907e0ad test CI 2023-05-20 13:04:20 +01:00
Zomatree 4c8553d5ef ignore external types 2023-05-20 03:18:19 +01:00
Zomatree f717163e17 include py.typed 2023-05-20 03:14:10 +01:00
Zomatree 60474fdeef fix verify-types parameter 2023-05-20 03:08:24 +01:00
Zomatree ccfae66e16 dont cache deps 2023-05-20 03:06:04 +01:00
Zomatree 60ff8d81c9 bring type completeness to 100% 2023-05-20 03:04:52 +01:00
TheBobBobs dfb45494ba Fix Message._update (#50)
* fix Message._update

* handle edited being int in MessageUpdate event
2023-05-19 23:16:33 +01:00
Zomatree 52c2be4e91 Fix using .server instead of .server_id
Fixes #51
2023-05-19 23:13:55 +01:00
Zomatree d5a15c0f44 fix typing error related to bad aiohttp types 2023-05-19 22:23:47 +01:00
Zomatree 4bf0530f8e clean up code 2023-05-11 07:02:39 +01:00
Zomatree 77ff484ad8 small bug fixes 2023-04-12 20:20:18 +01:00
Zomatree d909b0eb3e Merge branch 'master' of github.com:revoltchat/revolt.py 2023-04-12 16:45:11 +01:00
Zomatree 32d6b1d6e8 Add missing checks to __all__ 2023-04-12 16:41:37 +01:00
Zomatree f9ba869e75 fix fetch_emoji not returning an Emoji instance 2023-04-12 16:37:05 +01:00
Angelo Kontaxis 0e694db586 Merge pull request #49 from revoltchat/permissions
Permissions
2023-04-11 22:34:30 +01:00
Zomatree ba3f74dcd0 Add docs 2023-04-11 22:23:12 +01:00
Zomatree 625dd4eac2 finish implementation of permissions 2023-04-11 22:07:56 +01:00
Zomatree 0a6db1dc9c inital permissions calculations 2023-04-11 19:05:56 +01:00
Zomatree a672f949a1 fix servermemberjoin, closes #48 2023-03-27 00:43:12 +01:00
Zomatree df9893aaff add bulk message delete 2023-03-26 21:44:04 +01:00
Zomatree 35a4614b61 fix incorrect routes 2023-03-26 18:42:16 +01:00
Zomatree 263f99f281 use is not None in _update 2023-03-26 18:09:35 +01:00
Zomatree e99b6edee3 correctly handle role creation 2023-03-26 18:08:30 +01:00
Zomatree 720272d3cb switch from deprecated config 2023-03-26 18:08:06 +01:00
Zomatree ab5a2751df fix docs 2023-03-15 19:41:31 +00:00
Zomatree 6059e4e4dc remove bugged dep 2023-03-15 19:30:41 +00:00
Zomatree 63491ec76e update docs 2023-03-15 19:00:49 +00:00
Zomatree 76364572f1 timeout and upload_file 2023-02-22 18:49:55 +00:00
Zomatree 75f980e0f9 use our own msgpack types 2023-02-22 18:40:06 +00:00
50 changed files with 1321 additions and 577 deletions
+26 -4
View File
@@ -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
View File
@@ -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"
+3 -3
View File
@@ -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
+37
View File
@@ -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
View File
@@ -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
View File
@@ -18,4 +18,4 @@ from .role import *
from .server import *
from .user import *
__version__ = "0.1.9"
__version__ = "0.1.11"
+20 -21
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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"
+40 -10
View File
@@ -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
+10 -11
View File
@@ -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
+6 -6
View File
@@ -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]
+29 -27
View File
@@ -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
-----------
+36 -15
View File
@@ -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)
+16 -16
View File
@@ -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)
+13 -1
View File
@@ -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"""
+9 -13
View File
@@ -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
View File
@@ -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)
+5 -2
View File
@@ -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] = []
+5 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+80
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 -1
View File
@@ -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
View File
@@ -1,6 +1,5 @@
from typing import TypedDict
__all__ = ("Category",)
class Category(TypedDict):
+2 -2
View File
@@ -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]
+2
View File
@@ -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
View File
@@ -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]
+3 -1
View File
@@ -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]]
+1
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from typing import TypedDict
class Overwrite(TypedDict):
a: int
d: int
+1
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
+3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+40
View File
@@ -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: ...
+1
View File
@@ -0,0 +1 @@
def get_html_theme_path() -> str: ...