1 Commits

Author SHA1 Message Date
Zomatree b61bb49949 inital voice support 2021-12-31 04:52:56 +00:00
62 changed files with 1243 additions and 4220 deletions
+4 -2
View File
@@ -8,9 +8,11 @@ jobs:
- uses: actions/setup-python@v2
with:
python-version: '3.9'
- run: pip install .[speedups]
- uses: BSFishy/pip-action@v1
with:
requirements: pyright_requirements.txt
- uses: jakebailey/pyright-action@v1
with:
lib: true
python-version: 3.9
python-version: 3.9.7
working-directory: revolt
-1
View File
@@ -7,4 +7,3 @@ docs/_build
.vscode
.env
.mypy_cache
build
+5 -8
View File
@@ -1,15 +1,12 @@
version: 2
sphinx:
configuration: docs/conf.py
configuration: docs/conf.py
python:
version: "3.9"
install:
- method: pip
path: .
extra_requirements:
- docs
version: "3.9"
install:
- requirements: docs_requirements.txt
build:
image: testing
image: testing
+11 -3
View File
@@ -1,13 +1,21 @@
set dotenv-load := true
activate:
source .venv/bin/activate
venv:
python -m venv .venv
just activate
python -m pip install -r requirements.txt
python -m pip install -r docs_requirements.txt
test:
python test.py
build:
rm -rf dist/*
build: venv
python -m build
upload:
upload: venv
python -m twine upload dist/* -u $PYPI_USERNAME -p $PYPI_PASSWORD
lint:
+5 -15
View File
@@ -1,26 +1,16 @@
# Revolt.py
An async library to interact with the https://revolt.chat API.
An async library to interact with the https://revolt.chat api.
You can join the support server [here](https://rvlt.gg/FDXER6hr) and find the library's documentation [here](https://revoltpy.readthedocs.io/en/latest/).
This library will be focused on making bots and i will not implement anything only for user accounts.
## Installing
Support server: https://app.revolt.chat/invite/FDXER6hr
You can use `pip` to install revolt.py. It differs slightly depending on what OS/Distro you use.
On Windows
```
py -m pip install -U revolt.py # -U to update
```
On macOS and Linux
```
python3 -m pip install -U revolt.py
```
Documentation is [here](https://revoltpy.readthedocs.io/en/latest/)
## Example
More examples can be found in the [examples folder](https://github.com/revoltchat/revolt.py/blob/master/examples).
More examples in the [examples folder](https://github.com/zomatree/revolt.py/blob/master/examples)
```py
import revolt
+9 -137
View File
@@ -16,13 +16,6 @@ Asset
.. autoclass:: Asset
:members:
PartialAsset
~~~~~~~~~~~~~
.. autoclass:: PartialAsset
:members:
Channel
~~~~~~~~
@@ -54,7 +47,7 @@ TextChannel
:members:
VoiceChannel
~~~~~~~~~~~~~
~~~~~~~~~~~~
.. autoclass:: VoiceChannel
:members:
@@ -66,35 +59,6 @@ Embed
.. autoclass:: Embed
:members:
WebsiteEmbed
~~~~~~~~~~~~~
.. autoclass:: WebsiteEmbed
:members:
ImageEmbed
~~~~~~~~~~~
.. autoclass:: ImageEmbed
:members:
TextEmbed
~~~~~~~~~~
.. autoclass:: TextEmbed
:members:
NoneEmbed
~~~~~~~~~~
.. autoclass:: NoneEmbed
:members:
SendableEmbed
~~~~~~~~~~~~~~
.. autoclass:: SendableEmbed
:members:
File
~~~~~
@@ -109,21 +73,14 @@ Member
:members:
Message
~~~~~~~~
~~~~~~~
.. autoclass:: Message
:members:
MessageReply
~~~~~~~~~~~~~
.. autoclass:: MessageReply
:members:
Masquerade
~~~~~~~~~~~~~
.. autoclass:: Masquerade
:members:
Messageable
~~~~~~~~~~~~
@@ -133,12 +90,10 @@ Messageable
Permissions
~~~~~~~~~~~~
.. autoclass:: Permissions
.. autoclass:: ChannelPermissions
:members:
PermissionsOverwrite
~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: PermissionsOverwrite
.. autoclass:: ServerPermissions
:members:
Role
@@ -153,12 +108,6 @@ Server
.. autoclass:: Server
:members:
ServerBan
~~~~~~~~~~
.. autoclass:: ServerBan
:members:
Category
~~~~~~~~~
@@ -196,23 +145,11 @@ UserBadges
.. autoclass:: UserBadges
:members:
UserProfile
~~~~~~~~~~~~
.. autoclass:: UserProfile
:members:
Invite
~~~~~~~
.. autoclass:: Invite
:members:
Enums
======
The api uses enums to say what variant of something is,
these represent those enums
The api uses enums to say what variant of something is, these represent those enums
All enums subclass `aenum.Enum`.
@@ -224,10 +161,10 @@ All enums subclass `aenum.Enum`.
A private channel only you can access.
.. attribute:: direct_message
A private direct message channel between you and another user
.. attribute:: group
A private group channel for messages between a group of users
.. attribute:: text_channel
@@ -274,77 +211,12 @@ All enums subclass `aenum.Enum`.
They are sending you a friend request
.. attribute:: none
You have no relationship with them
.. attribute:: outgoing_friend_request
You are sending them a friend request
.. attribute:: user
That user is yourself
.. class:: AssetType
Specifies the type of asset
.. attribute:: image
The asset is an image
.. attribute:: video
The asset is a video
.. attribute:: text
The asset is a text file
.. attribute:: audio
The asset is an audio file
.. attribute:: file
The asset is a generic file
.. class:: SortType
The sort type for a message search
.. attribute:: latest
Sort by the latest message
.. attribute:: oldest
Sort by the oldest message
.. attribute:: relevance
Sort by the relevance of the message
.. class:: EmbedType
The type of embed
.. attribute:: website
The embed is a website
.. attribute:: image
The embed is an image
.. attribute:: text
The embed is text
.. attribute:: video
The embed is a video
.. attribute:: unknown
The embed is unknown
Utils
======
.. currentmodule:: revolt.utils
A collection a utility functions and classes to aid in making your bot
.. autofunction:: get
.. autofunction:: client_session
-1
View File
@@ -17,7 +17,6 @@ import sphinx_nameko_theme
sys.path.insert(0, os.path.abspath('..'))
import revolt
# -- Project information -----------------------------------------------------
-20
View File
@@ -19,11 +19,6 @@ Command
.. autoclass:: revolt.ext.commands.Command
:members:
Cog
~~~~
.. autoclass:: revolt.ext.commands.Cog
:members:
command
~~~~~~~~
.. autodecorator:: revolt.ext.commands.command
@@ -93,18 +88,3 @@ BadBoolArgument
~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.BadBoolArgument
:members:
CategoryConverterError
~~~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.CategoryConverterError
:members:
UserConverterError
~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.UserConverterError
:members:
MemberConverterError
~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.MemberConverterError
:members:
+6
View File
@@ -0,0 +1,6 @@
Sphinx==3.5.4
sphinx-nameko-theme==0.0.3
aiohttp==3.7.4.post0
ulid-py==1.1.0
sphinx-toolbox==2.13.0
aenum==3.1.0
Generated
-1298
View File
File diff suppressed because it is too large Load Diff
-46
View File
@@ -1,46 +0,0 @@
[tool.poetry]
name = "revolt.py"
version = "0.1.9"
description = "Python wrapper for the revolt.chat API"
authors = ["Zomatee <me@zomatree.live>"]
license = "MIT"
readme = "README.md"
homepage = "https://github.com/revoltchat/revolt.py"
documentation = "https://revoltpy.readthedocs.io/en/latest/"
classifiers = [
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"License :: OSI Approved :: MIT License",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3 :: Only",
]
keywords = ["wrapper", "async", "api", "websockets", "http"]
packages = [
{ include = "revolt" }
]
[tool.poetry.dependencies]
python = "^3.9"
aiohttp = "3.7.4"
ulid-py = "1.1.0"
aenum = "3.1.8"
typing-extensions = "4.1.1"
ujson = { version = "5.1.0", optional = true }
msgpack = { version = "", optional = true }
Sphinx = { version = "4.3.2", optional = true }
sphinx-nameko-theme = { version = "0.0.3", optional = true }
sphinx-toolbox = { version = "2.15.2", optional = true }
[tool.poetry.extras]
speedups = ["ujson", "aiohttp[speedups]", "msgpack"]
docs = ["Sphinx", "sphinx-nameko-theme", "sphinx-toolbox"]
[build-system]
requires = ["poetry-core>=1.0.0"]
build-backend = "poetry.core.masonry.api"
[tool.poetry.urls]
"Bug Tracker" = "https://github.com/revoltchat/revolt.py/issues"
Source = "https://github.com/revoltchat/revolt.py/"
+6
View File
@@ -0,0 +1,6 @@
aiohttp==3.7.4.post0
ulid-py==1.1.0
aenum==3.1.0
ujson
msgpack==1.0.2
.
+5
View File
@@ -0,0 +1,5 @@
aiohttp==3.7.4.post0
ulid-py==1.1.0
aenum==3.1.0
typing_extensions
pynacl==1.4.0
+19 -19
View File
@@ -1,20 +1,20 @@
from . import utils
from .asset import *
from .category import *
from .channel import *
from .client import *
from .embed import *
from .enums import *
from .errors import *
from .file import *
from .flags import *
from .invite import *
from .member import *
from .message import *
from .messageable import *
from .permissions import *
from .role import *
from .server import *
from .user import *
from .asset import Asset
from .category import Category
from .channel import (Channel, DMChannel, GroupDMChannel, SavedMessageChannel,
TextChannel, VoiceChannel)
from .client import Client
from .embed import Embed
from .enums import (AssetType, ChannelType, PresenceType, RelationshipType,
SortType)
from .errors import HTTPError, RevoltError, ServerError
from .file import File
from .flags import UserBadges
from .member import Member
from .message import Masquerade, Message, MessageReply
from .messageable import Messageable
from .permissions import ChannelPermissions, ServerPermissions
from .role import Role
from .server import Server, SystemMessages
from .user import Relation, Status, User
__version__ = (0, 1, 9)
__version__ = (0, 1, 1)
+6 -6
View File
@@ -13,11 +13,11 @@ if TYPE_CHECKING:
from .types import File as FilePayload
__all__ = ("Asset", "PartialAsset")
__all__ = ("Asset",)
class Asset:
"""Represents a file on revolt
Attributes
-----------
id: :class:`str`
@@ -40,7 +40,7 @@ class Asset:
The assets url
"""
__slots__ = ("state", "id", "tag", "size", "filename", "content_type", "width", "height", "type", "url")
def __init__(self, data: FilePayload, state: State):
self.state = state
@@ -48,10 +48,10 @@ class Asset:
self.tag = data['tag']
self.size = data['size']
self.filename = data['filename']
metadata = data['metadata']
if metadata["type"] == "Image" or metadata["type"] == "Video": # cannot use `in` because type narrowing will not happen
if metadata["type"] == "Image" or metadata["type"] == "Video": # cant use `in` because type narrowing wont happen
self.height = metadata["height"]
self.width = metadata["width"]
else:
@@ -70,7 +70,7 @@ class Asset:
async def save(self, fp: IOBase):
"""Reads the files content and saves it to a file
Parameters
-----------
fp: IOBase
+105 -259
View File
@@ -1,14 +1,13 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal, Optional, Union
from typing import TYPE_CHECKING, Optional, cast
import asyncio
from revolt.utils import Missing
from .asset import Asset
from .enums import ChannelType
from .messageable import Messageable
from .permissions import Permissions, PermissionsOverwrite
from .utils import Missing
from .permissions import ChannelPermissions
from .voice_client import VoiceClient
from .errors import FeatureDisabled
if TYPE_CHECKING:
from .message import Message
@@ -17,67 +16,25 @@ 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 Group as GroupDMChannelPayload
from .types import SavedMessages as SavedMessagesPayload
from .types import TextChannel as TextChannelPayload
from .types import VoiceChannel as VoiceChannelPayload
from .types import GuildChannel as GuildChannelPayload
from .types import File as FilePayload
from .types import Overwrite as OverwritePayload
from .user import User
__all__ = ("DMChannel", "GroupDMChannel", "SavedMessageChannel", "TextChannel", "VoiceChannel", "Channel")
class EditableChannel:
__slots__ = ()
state: State
id: str
async def edit(self, **kwargs):
"""Edits the channel
Passing ``None`` to the parameters that accept it will remove them.
Parameters
-----------
name: str
The new name for the channel
description: Optional[str]
The new description for the channel
owner: User
The new owner for the group dm channel
icon: Optional[File]
The new icon for the channel
nsfw: bool
Sets whether the channel is nsfw or not
"""
if kwargs.get("icon", Missing) == None:
remove = "Icon"
elif kwargs.get("description", Missing) == None:
remove = "Description"
else:
remove = None
if icon := kwargs.get("icon"):
asset = await self.state.http.upload_file(icon, "icons")
kwargs["icon"] = asset["id"]
if owner := kwargs.get("owner"):
kwargs["owner"] = owner.id
await self.state.http.edit_channel(self.id, remove, kwargs)
__all__ = ("Channel",)
class Channel:
"""Base class for all channels
Attributes
-----------
id: :class:`str`
The id of the channel
channel_type: ChannelType
The type of the channel
server_id: Optional[:class:`str`]
The server id of the chanel, if any
server: Optional[:class:`Server`]
The server the channel is part of
"""
__slots__ = ("state", "id", "channel_type", "server_id")
@@ -85,277 +42,166 @@ class Channel:
self.state = state
self.id = data["_id"]
self.channel_type = ChannelType(data["channel_type"])
self.server_id: Optional[str] = None
self.server_id = ""
async def _get_channel_id(self) -> str:
return self.id
@property
def server(self) -> Optional[Server]:
return self.state.get_server(self.server_id) if self.server_id else None
def _update(self, **_):
def _update(self):
pass
async def delete(self):
"""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"""
if not self.server_id:
raise IndexError
return self.state.get_server(self.server_id)
@property
def mention(self) -> str:
""":class:`str`: Returns a string that allows you to mention the given channel."""
return f"<#{self.id}>"
class SavedMessageChannel(Channel, Messageable):
"""The Saved Message Channel"""
def __init__(self, data: SavedMessagesPayload, state: State):
super().__init__(data, state)
class DMChannel(Channel, Messageable):
"""A DM channel
Attributes
-----------
last_message_id: Optional[:class:`str`]
The id of the last message in this channel, if any
"""
__slots__ = ("last_message_id",)
"""A DM channel"""
def __init__(self, data: DMChannelPayload, state: State):
super().__init__(data, state)
self.last_message_id = data.get("last_message_id")
@property
def last_message(self) -> Message:
"""Gets the last message from the channel, shorthand for `client.get_message(channel.last_message_id)`
Returns
--------
:class:`Message` the last message in the channel
"""
if not self.last_message_id:
raise LookupError
return self.state.get_message(self.last_message_id)
class GroupDMChannel(Channel, Messageable, EditableChannel):
"""A group DM channel
Attributes
-----------
recipients: list[:class:`User`]
The recipients of the group dm channel
name: :class:`str`
The name of the group dm channel
owner: :class:`User`
The user who created the group dm channel
icon: Optional[:class:`Asset`]
The icon of the group dm channel
permissions: :class:`ChannelPermissions`
The permissions of the users inside the group dm channel
description: Optional[:class:`str`]
The description of the channel, if any
last_message_id: Optional[:class:`str`]
The id of the last message in this channel, if any
"""
__slots__ = ("recipients", "name", "owner", "permissions", "icon", "description", "last_message_id")
class GroupDMChannel(Channel, Messageable):
__slots__ = ("recipients", "name", "owner", "permissions")
"""A group DM channel"""
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")
if icon := data.get("icon"):
self.icon = Asset(icon, state)
else:
self.icon = None
if perms := data.get("permissions"):
self.permissions = ChannelPermissions._from_value(perms)
self.permissions = Permissions(data.get("permissions", 0))
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None, description: Optional[str] = None):
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None):
if name:
self.name = name
if recipients:
self.recipients = [self.state.get_user(user_id) for user_id in recipients]
if description:
self.description = description
async def set_default_permissions(self, permissions: Permissions) -> None:
async def set_default_permissions(self, permissions: ChannelPermissions) -> None:
"""Sets the default permissions for a group.
Parameters
-----------
permissions: :class:`ChannelPermissions`
The new default group permissions
"""
await self.state.http.set_group_channel_default_permissions(self.id, permissions.value)
await self.state.http.set_channel_default_permissions(self.id, permissions.value)
@property
def last_message(self) -> Message:
"""Gets the last message from the channel, shorthand for `client.get_message(channel.last_message_id)`
class TextChannel(Channel, Messageable):
__slots__ = ("name", "description", "last_message_id", "server_id", "default_permissions", "role_permissions")
Returns
--------
:class:`Message` the last message in the channel
"""
if not self.last_message_id:
raise LookupError
return self.state.get_message(self.last_message_id)
class GuildChannel(Channel):
def __init__(self, data: GuildChannelPayload, state: State):
"""A text channel"""
def __init__(self, data: TextChannelPayload, state: State):
super().__init__(data, state)
self.server_id = data["server"]
self.name = 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.description = data.get("description")
permissions: dict[str, PermissionsOverwrite] = {}
last_message_id = data.get("last_message")
self.last_message_id = last_message_id
for role_name, overwrite_data in data.get("role_permissions", {}).items():
overwrite = PermissionsOverwrite._from_overwrite(overwrite_data)
permissions[role_name] = overwrite
if perms := data.get("default_permissions"):
self.default_permissions = ChannelPermissions._from_value(perms)
self.permissions = permissions
if icon := data.get("icon"):
self.icon = Asset(icon, state)
else:
self.icon = None
if role_perms := data.get("role_permissions"):
self.role_permissions = {role_id: ChannelPermissions._from_value(perms) for role_id, perms in role_perms.items()}
async def set_default_permissions(self, permissions: PermissionsOverwrite) -> None:
"""Sets the default permissions for the channel.
def _get_channel_id(self) -> str:
return self.id
@property
def last_message(self) -> Message:
return self.state.get_message(self.last_message_id)
def _update(self, *, name: Optional[str] = None, description: Optional[str] = None):
if name:
self.name = name
if description:
self.description = description
async def set_default_permissions(self, permissions: ChannelPermissions) -> None:
"""Sets the default permissions for a channel.
Parameters
-----------
permissions: :class:`ChannelPermissions`
The new default channel permissions
"""
allow, deny = permissions.to_pair()
await self.state.http.set_guild_channel_default_permissions(self.id, allow.value, deny.value)
await self.state.http.set_channel_default_permissions(self.id, permissions.value)
async def set_role_permissions(self, role: Role, permissions: PermissionsOverwrite) -> None:
"""Sets the permissions for a role in the channel.
async def set_role_permissions(self, role: Role, permissions: ChannelPermissions) -> None:
"""Sets the permissions for a role in a channel.
Parameters
-----------
permissions: :class:`ChannelPermissions`
The new channel permissions
"""
allow, deny = permissions.to_pair()
await self.state.http.set_channel_role_permissions(self.id, role.id, permissions.value)
await self.state.http.set_guild_channel_role_permissions(self.id, role.id, allow.value, deny.value)
def _update(self, *, name: Optional[str] = None, description: Optional[str] = None, icon: Optional[FilePayload] = None, nsfw: Optional[bool] = None, active: Optional[bool] = None, role_permissions: Optional[dict[str, OverwritePayload]] = None, default_permissions: Optional[OverwritePayload] = None):
if name is not None:
self.name = name
if description is not None:
self.description = description
if icon:
self.icon = Asset(icon, self.state)
if nsfw is not None:
self.nsfw = nsfw
if active is not None:
self.active = active
if role_permissions is not None:
permissions = {}
for role_name, overwrite_data in role_permissions.items():
overwrite = PermissionsOverwrite._from_overwrite(overwrite_data)
permissions[role_name] = overwrite
self.permissions = permissions
if default_permissions is not None:
self.default_permissions = default_permissions
class TextChannel(GuildChannel, Messageable, EditableChannel):
"""A text channel
Attributes
-----------
name: :class:`str`
The name of the text channel
server_id: :class:`str`
The id of the server this text channel belongs to
last_message_id: Optional[:class:`str`]
The id of the last message in this channel, if any
default_permissions: :class:`ChannelPermissions`
The default permissions for all users in the text channel
role_permissions: dict[:class:`str`, :class:`ChannelPermissions`]
A dictionary of role id's to the permissions of that role in the text channel
icon: Optional[:class:`Asset`]
The icon of the text channel, if any
description: Optional[:class:`str`]
The description of the channel, if any
"""
__slots__ = ("name", "description", "last_message_id", "default_permissions", "icon", "overwrites")
def __init__(self, data: TextChannelPayload, state: State):
class VoiceChannel(Channel):
"""A voice channel"""
def __init__(self, data: VoiceChannelPayload, state: State):
super().__init__(data, state)
self.last_message_id = data.get("last_message_id")
self.server_id = data["server"]
self.name = data["name"]
self.description = data.get("description")
async def _get_channel_id(self) -> str:
return self.id
if perms := data.get("default_permissions"):
self.default_permissions = ChannelPermissions._from_value(perms)
@property
def last_message(self) -> Message:
"""Gets the last message from the channel, shorthand for `client.get_message(channel.last_message_id)`
if role_perms := data.get("role_permissions"):
self.role_permissions = {role_id: ChannelPermissions._from_value(perms) for role_id, perms in role_perms.items()}
Returns
--------
:class:`Message` the last message in the channel
def _update(self, *, name: Optional[str] = None, description: Optional[str] = None):
if name:
self.name = name
if description:
self.description = description
async def set_default_permissions(self, permissions: ChannelPermissions) -> None:
"""Sets the default permissions for a voice channel.
Parameters
-----------
permissions: :class:`ChannelPermissions`
The new default channel permissions
"""
await self.state.http.set_channel_default_permissions(self.id, permissions.value)
if not self.last_message_id:
raise LookupError
async def set_role_permissions(self, role: Role, permissions: ChannelPermissions) -> None:
"""Sets the permissions for a role in a voice channel
Parameters
-----------
permissions: :class:`ChannelPermissions`
The new channel permissions
"""
await self.state.http.set_channel_role_permissions(self.id, role.id, permissions.value)
return self.state.get_message(self.last_message_id)
async def connect(self):
token = (await self.state.http.connect_to_voice(self.id))["token"]
voso = self.state.api_info["features"]["voso"]
class VoiceChannel(GuildChannel, EditableChannel):
"""A voice channel
if not voso["enabled"]:
raise FeatureDisabled("vortex is disabled.")
Attributes
-----------
name: :class:`str`
The name of the voice channel
server_id: :class:`str`
The id of the server this voice channel belongs to
last_message_id: Optional[:class:`str`]
The id of the last message in this channel, if any
default_permissions: :class:`ChannelPermissions`
The default permissions for all users in the voice channel
role_permissions: dict[:class:`str`, :class:`ChannelPermissions`]
A dictionary of role id's to the permissions of that role in the voice channel
icon: Optional[:class:`Asset`]
The icon of the voice channel, if any
description: Optional[:class:`str`]
The description of the channel, if any
"""
voice_client = VoiceClient(self.state.http.session, voso["ws"], self.id, token, self.state)
def channel_factory(data: ChannelPayload, state: State) -> Union[DMChannel, GroupDMChannel, SavedMessageChannel, TextChannel, VoiceChannel]:
if data["channel_type"] == "SavedMessages":
await voice_client.start()
await voice_client.send_authenticate()
print("auth")
await voice_client.send_initialize_transport()
print("init")
await voice_client.send_connect_transport()
print("connect")
return voice_client
def channel_factory(data: ChannelPayload, state: State) -> Channel:
if data["channel_type"] == "SavedMessage":
return SavedMessageChannel(data, state)
elif data["channel_type"] == "DirectMessage":
return DMChannel(data, state)
+24 -201
View File
@@ -2,17 +2,12 @@ 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, Coroutine, Optional, TypeVar
import aiohttp
from .channel import (DMChannel, GroupDMChannel, SavedMessageChannel,
TextChannel, VoiceChannel, channel_factory)
from .http import HttpClient
from .invite import Invite
from .message import Message
from .state import State
from .utils import Missing
from .websocket import WebsocketHandler
try:
@@ -33,7 +28,7 @@ logger = logging.getLogger("revolt")
class Client:
"""The client for interacting with revolt
Parameters
-----------
session: :class:`aiohttp.ClientSession`
@@ -45,14 +40,13 @@ class Client:
max_messages: :class:`int`
The max amount of messages stored in the cache, by default this is 5k
"""
def __init__(self, session: aiohttp.ClientSession, token: str, *, api_url: str = "https://api.revolt.chat", max_messages: int = 5000, bot: bool = True):
def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str = "https://api.revolt.chat", max_messages: int = 5000):
self.session = session
self.token = token
self.api_url = api_url
self.max_messages = max_messages
self.bot = bot
self.api_info: ApiInfo
self.http: HttpClient
self.state: State
@@ -64,7 +58,7 @@ class Client:
def dispatch(self, event: str, *args: Any):
"""Dispatch an event, this is typically used for testing and internals.
Parameters
----------
event: class:`str`
@@ -92,59 +86,59 @@ class Client:
api_info = await self.get_api_info()
self.api_info = api_info
self.http = HttpClient(self.session, self.token, self.api_url, self.api_info, self.bot)
self.state = State(self.http, api_info, self.max_messages)
self.http = HttpClient(self.session, self.token, self.api_url, self.api_info)
self.state = State(self.http, api_info, self.max_messages, self.dispatch)
self.websocket = WebsocketHandler(self.session, self.token, api_info["ws"], self.dispatch, self.state)
await self.websocket.start()
def get_user(self, id: str) -> User:
def get_user(self, id: str) -> Optional[User]:
"""Gets a user from the cache
Parameters
-----------
id: :class:`str`
The id of the user
Returns
--------
:class:`User`
The user
Optional[:class:`User`]
The user if found
"""
return self.state.get_user(id)
def get_channel(self, id: str) -> Channel:
def get_channel(self, id: str) -> Optional[Channel]:
"""Gets a channel from the cache
Parameters
-----------
id: :class:`str`
The id of the channel
Returns
--------
:class:`Channel`
The channel
Optional[:class:`Channel`]
The channel if found
"""
return self.state.get_channel(id)
def get_server(self, id: str) -> Server:
def get_server(self, id: str) -> Optional[Server]:
"""Gets a server from the cache
Parameters
-----------
id: :class:`str`
The id of the server
Returns
--------
:class:`Server`
The server
Optional[:class:`Server`]
The server if found
"""
return self.state.get_server(id)
async def wait_for(self, event: str, *, check: Optional[Callable[..., bool]] = None, timeout: Optional[float] = None) -> Any:
"""Waits for an event
Parameters
-----------
event: :class:`str`
@@ -174,178 +168,7 @@ class Client:
@property
def user(self) -> User:
""":class:`User` the user corrasponding to the client"""
user = self.websocket.user
assert user
return user
@property
def users(self) -> list[User]:
"""list[:class:`User`] All users the client can see"""
return list(self.state.users.values())
@property
def servers(self) -> list[Server]:
"""list[:class:'Server'] All servers the client can see"""
return list(self.state.servers.values())
async def fetch_user(self, user_id: str) -> User:
"""Fetchs a user
Parameters
-----------
user_id: :class:`str`
The id of the user you are fetching
Returns
--------
:class:`User`
The user with the matching id
"""
payload = await self.http.fetch_user(user_id)
return User(payload, self.state)
async def fetch_dm_channels(self) -> list[Union[DMChannel, GroupDMChannel]]:
"""Fetchs all dm channels the client has made
Returns
--------
list[Union[:class:`DMChanel`, :class:`GroupDMChannel`]]
A list of :class:`DMChannel` or :class`GroupDMChannel`
"""
channel_payloads = await self.http.fetch_dm_channels()
return cast(list[Union[DMChannel, GroupDMChannel]], [channel_factory(payload, self.state) for payload in channel_payloads])
async def fetch_channel(self, channel_id: str) -> Union[DMChannel, GroupDMChannel, SavedMessageChannel, TextChannel, VoiceChannel]:
"""Fetches a channel
Parameters
-----------
channel_id: :class:`str`
The id of the channel
Returns
--------
Union[:class:`DMChannel`, :class:`GroupDMChannel`, :class:`SavedMessageChannel`, :class:`TextChannel`, :class:`VoiceChannel`]
The channel with the matching id
"""
payload = await self.http.fetch_channel(channel_id)
return channel_factory(payload, self.state)
async def fetch_server(self, server_id: str) -> Server:
"""Fetchs a server
Parameters
-----------
server_id: :class:`str`
The id of the server you are fetching
Returns
--------
:class:`Server`
The server with the matching id
"""
payload = await self.http.fetch_server(server_id)
return Server(payload, self.state)
async def fetch_invite(self, code: str) -> Invite:
"""Fetchs an invite
Parameters
-----------
code: :class:`str`
The code of the invite you are fetching
Returns
--------
:class:`Invite`
The invite with the matching code
"""
payload = await self.http.fetch_invite(code)
return Invite(payload, code, self.state)
def get_message(self, message_id: str) -> Message:
"""Gets a message from the cache
Parameters
-----------
message_id: :class:`str`
The id of the message you are getting
Returns
--------
:class:`Message`
The message with the matching id
Raises
-------
LookupError
This raises if the message is not found in the cache
"""
for message in self.state.messages:
if message.id == message_id:
return message
raise LookupError
async def edit_self(self, **kwargs):
"""Edits the client's own user
Parameters
-----------
avatar: Optional[:class:`File`]
The avatar to change to, passing in ``None`` will remove the avatar
"""
if kwargs.get("avatar", Missing) is None:
del kwargs["avatar"]
remove = "Avatar"
else:
remove = None
await self.state.http.edit_self(remove, kwargs)
async def edit_status(self, **kwargs):
"""Edits the client's own status
Parameters
-----------
presence: :class:`PresenceType`
The presence to change to
text: Optional[:class:`str`]
The text to change the status to, passing in ``None`` will remove the status
"""
if kwargs.get("text", Missing) is None:
del kwargs["text"]
remove = "StatusText"
else:
remove = None
if presence := kwargs.get("presence"):
kwargs["presence"] = presence.value
await self.state.http.edit_self(remove, {"status": kwargs})
async def edit_profile(self, **kwargs):
"""Edits the client's own profile
Parameters
-----------
content: Optional[:class:`str`]
The new content for the profile, passing in ``None`` will remove the profile content
background: Optional[:class:`File`]
The new background for the profile, passing in ``None`` will remove the profile background
"""
if kwargs.get("content", Missing) is None:
del kwargs["content"]
remove = "ProfileContent"
elif kwargs.get("background", Missing) is None:
del kwargs["background"]
remove = "ProfileBackground"
else:
remove = None
await self.state.http.edit_self(remove, {"profile": kwargs})
+8 -97
View File
@@ -1,106 +1,17 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional, Union
from .asset import Asset
from .enums import EmbedType
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .state import State
from .types import Embed as EmbedPayload
from .types import ImageEmbed as ImageEmbedPayload
from .types import NoneEmbed as NoneEmbedPayload
from .types import SendableEmbed as SendableEmbedPayload
from .types import TextEmbed as TextEmbedPayload
from .types import WebsiteEmbed as WebsiteEmbedPayload
__all__ = ("Embed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "to_embed", "SendableEmbed")
class WebsiteEmbed:
type = EmbedType.website
__all__ = ("Embed",)
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")
class Embed:
@classmethod
def from_dict(cls, data: EmbedPayload):
...
class ImageEmbed:
type = 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")
class TextEmbed:
type = 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")
if media := embed.get("media"):
self.media = Asset(media, state)
else:
self.media = None
self.colour = embed.get("colour")
class NoneEmbed:
type = EmbedType.none
Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed]
def to_embed(payload: EmbedPayload, state: State) -> Embed:
if payload["type"] == "Website":
return WebsiteEmbed(payload)
elif payload["type"] == "Image":
return ImageEmbed(payload)
elif payload["type"] == "Text":
return TextEmbed(payload, state)
else:
return NoneEmbed()
class SendableEmbed:
def __init__(self, **attrs):
self.title: Optional[str] = None
self.description: Optional[str] = None
self.media: Optional[str] = None
self.icon_url: Optional[str] = None
self.colour: Optional[str] = None
self.url: Optional[str] = None
for key, value in attrs.items():
setattr(self, key, value)
def to_dict(self) -> SendableEmbedPayload:
output: SendableEmbedPayload = {"type": "Text"}
if title := self.title:
output["title"] = title
if description := self.description:
output["description"] = description
if media := self.media:
output["media"] = media
if icon_url := self.icon_url:
output["icon_url"] = icon_url
if colour := self.colour:
output["colour"] = colour
if url := self.url:
output["url"] = url
return output
def to_dict(self) -> EmbedPayload:
...
+2 -9
View File
@@ -1,6 +1,6 @@
from typing import TYPE_CHECKING
# typing does not understand aenum so I am pretending its stdlib enum while type checking
# typing doesnt understand aenum so im pretending its stdlib enum while type checking
if TYPE_CHECKING:
import enum
@@ -14,11 +14,10 @@ __all__ = (
"RelationshipType",
"AssetType",
"SortType",
"EmbedType"
)
class ChannelType(enum.Enum):
saved_messages = "SavedMessages"
saved_message = "SavedMessage"
direct_message = "DirectMessage"
group = "Group"
text_channel = "TextChannel"
@@ -50,9 +49,3 @@ class SortType(enum.Enum):
latest = "Latest"
oldest = "Oldest"
relevance = "Relevance"
class EmbedType(enum.Enum):
website = "Website"
image = "Image"
text = "Text"
none = "None"
+1 -1
View File
@@ -16,7 +16,7 @@ class ServerError(RevoltError):
"Internal server error"
class FeatureDisabled(RevoltError):
"""Base class for any feature disabled errors"""
"""Base class for any any feature disabled"""
class AutumnDisabled(FeatureDisabled):
"""The autumn feature is disabled"""
-4
View File
@@ -1,9 +1,5 @@
from .checks import *
from .client import *
from .cog import *
from .command import *
from .context import *
from .converters import *
from .errors import *
from .group import *
from .help import *
+1 -1
View File
@@ -29,7 +29,7 @@ def check(check: Check):
else:
checks = getattr(func, "_checks", [])
checks.append(check)
func._checks = checks # type: ignore
func._checks = checks
return func
+15 -173
View File
@@ -1,22 +1,14 @@
from __future__ import annotations
import sys
import re
import traceback
from importlib import import_module
from typing import (TYPE_CHECKING, Any, Optional, Protocol, Union,
runtime_checkable)
from typing_extensions import Self
from typing import Any, Callable, Union
import revolt
if TYPE_CHECKING:
from .help import HelpCommand
from .cog import Cog
from .command import Command
from .context import Context
from .errors import CheckError, CommandNotFound, MissingSetup
from .errors import CheckError, CommandNotFound
from .view import StringView
__all__ = (
@@ -24,11 +16,9 @@ __all__ = (
"CommandsClient"
)
@runtime_checkable
class ExtensionProtocol(Protocol):
@staticmethod
def setup(client: CommandsClient) -> None:
raise NotImplementedError
quote_regex = re.compile(r"[\"']")
chunk_regex = re.compile(r"\S+")
class CommandsMeta(type):
_commands: list[Command]
@@ -36,6 +26,7 @@ class CommandsMeta(type):
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any]):
commands: list[Command] = []
self = super().__new__(cls, name, bases, attrs)
for base in reversed(self.__mro__):
for value in base.__dict__.values():
if isinstance(value, Command):
@@ -43,37 +34,18 @@ class CommandsMeta(type):
self._commands = commands
for command in commands:
command._client = self # type: ignore
return self
class CaseInsensitiveDict(dict):
def __setitem__(self, key: str, value: Any) -> None:
super().__setitem__(key.casefold(), value)
def __getitem__(self, key: str) -> Any:
return super().__getitem__(key.casefold())
def __contains__(self, key: str) -> bool:
return super().__contains__(key.casefold())
def get(self, key: str, default: Any = None) -> Any:
return super().get(key.casefold(), default)
def __delitem__(self, key: str) -> None:
super().__delitem__(key.casefold())
class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""Main class that adds commands, this class should be subclassed along with `revolt.Client`."""
_commands: list[Command]
def __init__(self, *args, help_command: Optional[HelpCommand] = None, case_insensitive: bool = False, **kwargs):
from .help import DefaultHelpCommand, HelpCommandImpl
self.all_commands: dict[str, Command] = {} if not case_insensitive else CaseInsensitiveDict()
self.cogs: dict[str, Cog] = {}
self.extensions: dict[str, ExtensionProtocol] = {}
def __init__(self, *args, **kwargs):
self.all_commands: dict[str, Command] = {}
for command in self._commands:
self.all_commands[command.name] = command
@@ -81,11 +53,6 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
for alias in command.aliases:
self.all_commands[alias] = command
if help_command is None:
help_command = DefaultHelpCommand()
self.help_command = DefaultHelpCommand()
self.add_command(HelpCommandImpl(self))
super().__init__(*args, **kwargs)
@property
@@ -122,7 +89,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""
return self.all_commands[name]
def add_command(self, command: Command):
def add_command(self, name: str, command: Command):
"""Adds a command, this is typically only used for dynamic commands, you should use the `commands.command` decorator for most usecases.
Parameters
@@ -132,36 +99,12 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
command: :class:`Command`
The command to be added
"""
self.all_commands[command.name] = command
for alias in command.aliases:
self.all_commands[alias] = command
def remove_command(self, name: str) -> Optional[Command]:
"""Removes a command.
Parameters
-----------
name: :class:`str`
The name or alias of the command
Returns
--------
Optional[:class:`Command`]
The command that was removed
"""
command = self.all_commands.pop(name, None)
if command is not None:
for alias in command.aliases:
self.all_commands.pop(alias, None)
return command
self.all_commands[name] = command
def get_view(self, message: revolt.Message) -> type[StringView]:
return StringView
def get_context(self, message: revolt.Message) -> type[Context[Self]]:
def get_context(self, message: revolt.Message) -> type[Context]:
return Context
async def process_commands(self, message: revolt.Message) -> Any:
@@ -228,7 +171,6 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return output
except Exception as e:
await command._error_handler(command.cog or self, context, e)
self.dispatch("command_error", context, e)
@staticmethod
@@ -251,103 +193,3 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""
return True
def add_cog(self, cog: Cog):
"""Adds a cog to the bot, this cog must subclass `Cog`.
Parameters
-----------
cog: :class:`Cog`
The cog to be added
"""
cog._inject(self)
def remove_cog(self, cog_name: str) -> Cog:
"""Removes a cog from the bot.
Parameters
-----------
cog_name: :class:`str`
The name of the cog to be removed
Returns
--------
:class:`Cog`
The cog that was removed
"""
cog = self.cogs.pop(cog_name)
cog._uninject(self)
return cog
def load_extension(self, name: str):
"""Loads an extension, this takes a module name and runs the setup function inside of it.
Parameters
-----------
name: :class:`str`
The name of the extension to be loaded
"""
extension = import_module(name)
if not isinstance(extension, ExtensionProtocol):
raise MissingSetup(f"'{extension}' is missing a setup function")
self.extensions[name] = extension
extension.setup(self)
def unload_extension(self, name: str):
"""Unloads an extension, this takes a module name and runs the teardown function inside of it.
Parameters
-----------
name: :class:`str`
The name of the extension to be unloaded
"""
extension = self.extensions.pop(name)
del sys.modules[name]
if teardown := getattr(extension, "teardown", None):
teardown(self)
def reload_extension(self, name: str):
"""Reloads an extension, this will unload and reload the extension.
Parameters
-----------
name: :class:`str`
The name of the extension to be reloaded
"""
self.unload_extension(name)
self.load_extension(name)
def get_cog(self, name: str) -> Cog:
"""Gets a cog from the bot.
Parameters
-----------
name: :class:`str`
The name of the cog to get
Returns
--------
:class:`Cog`
The cog that was requested
"""
return self.cogs[name]
def get_extension(self, name: str) -> ExtensionProtocol:
"""Gets an extension from the bot.
Parameters
-----------
name: :class:`str`
The name of the extension to get
Returns
--------
:class:`ExtensionProtocol`
The extension that was requested
"""
return self.extensions[name]
-61
View File
@@ -1,61 +0,0 @@
from __future__ import annotations
from distutils import command
from typing import TYPE_CHECKING, Any, Optional
from .command import Command
if TYPE_CHECKING:
from .client import CommandsClient
__all__ = ("Cog", "CogMeta")
class CogMeta(type):
_commands: list[Command]
qualified_name: str
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any], *, qualified_name: Optional[str] = None):
commands: list[Command] = []
self = super().__new__(cls, name, bases, attrs)
for base in reversed(self.__mro__):
for value in base.__dict__.values():
if isinstance(value, Command):
commands.append(value)
self._commands = commands
self.qualified_name = qualified_name or name
return self
class Cog(metaclass=CogMeta):
_commands: list[Command]
qualified_name: str
def cog_load(self):
"""A special method that is called when the cog gets loaded."""
pass
def cog_unload(self):
"""A special method that is called when the cog gets removed."""
pass
def _inject(self, client: CommandsClient):
client.cogs[self.qualified_name] = self
for command in self._commands:
command.cog = self
client.add_command(command)
self.cog_load()
def _uninject(self, client: CommandsClient):
for name, command in client.all_commands.copy().items():
if command in self._commands:
del client.all_commands[name]
self.cog_unload()
@property
def commands(self) -> list[Command]:
return self._commands
+49 -121
View File
@@ -4,28 +4,22 @@ import inspect
import traceback
from contextlib import suppress
from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine,
Literal, Optional, Union, cast, get_args, get_origin)
Literal, Optional, Union, get_args, get_origin)
import revolt
from revolt.utils import copy_doc, maybe_coroutine
from .errors import InvalidLiteralArgument, UnionConverterError
from .utils import evaluate_parameters
from .errors import InvalidLiteralArgument
if TYPE_CHECKING:
from .checks import Check
from .cog import Cog
from .context import Context
from .group import Group
__all__ = (
"Command",
"command"
)
NoneType = type(None)
class Command:
"""Class for holding info about a command.
@@ -37,27 +31,17 @@ class Command:
The name of the command
aliases: list[:class:`str`]
The aliases of the command
parent: Optional[:class:`Group`]
The parent of the command if this command is a subcommand
cog: Optional[:class:`Cog`]
The cog the command is apart of.
usage: Optional[:class:`str`]
The usage string for the command
"""
__slots__ = ("callback", "name", "aliases", "signature", "checks", "parent", "_error_handler", "cog", "description", "usage", "parameters")
__slots__ = ("callback", "name", "aliases", "_client", "signature", "checks")
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str], usage: Optional[str] = None):
_client: revolt.Client
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]):
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] = getattr(callback, "_checks", [])
self.parent: Optional[Group] = None
self.cog: Optional[Cog] = None
self._error_handler: Callable[[Any, Context, Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
self.description = callback.__doc__
async def invoke(self, context: Context, *args, **kwargs) -> Any:
"""Runs the command and calls the error handler if the command errors.
@@ -70,9 +54,9 @@ class Command:
The arguments for the command
"""
try:
return await self.callback(self.cog or context.client, context, *args, **kwargs)
return await self.callback(self._client, context, *args, **kwargs)
except Exception as err:
return await self._error_handler(self.cog or context.client, context, err)
return await self._error_handler(context, err)
@copy_doc(invoke)
def __call__(self, context: Context, *args, **kwargs) -> Any:
@@ -95,126 +79,70 @@ class Command:
await ctx.send(str(error))
"""
self._error_handler = func
self._error_handler = func # type: ignore
return func
async def _default_error_handler(self, ctx: Context, error: Exception):
@staticmethod
async def _error_handler(ctx: Context, error: Exception):
traceback.print_exception(type(error), error, error.__traceback__)
@classmethod
async def handle_origin(cls, context: Context, origin: Any, annotation: Any, arg: str) -> Any:
if origin is Union:
for converter in get_args(annotation):
try:
return await cls.convert_argument(arg, converter, context)
except:
if converter is NoneType:
context.view.undo()
return None
@staticmethod
def extract_type(t):
if origin := get_origin(t):
if origin is Annotated:
return get_args(t)[1]
raise UnionConverterError(arg)
elif origin is Annotated:
annotated_args = get_args(annotation)
if origin := get_origin(annotated_args[0]):
return await cls.handle_origin(context, origin, annotated_args[1], arg)
else:
return await cls.convert_argument(arg, annotated_args[1], context)
elif origin is Literal:
if arg in get_args(annotation):
return arg
else:
raise InvalidLiteralArgument(arg)
return t
@classmethod
async def convert_argument(cls, arg: str, annotation: Any, context: Context) -> Any:
if annotation is not inspect.Signature.empty:
if annotation is str: # no converting is needed - its already a string
async def convert_argument(cls, arg: str, parameter: inspect.Parameter, context: Context):
if annot := parameter.annotation:
if annot is str: # no converting is needed - its already a string
return arg
if origin := get_origin(annotation):
return await cls.handle_origin(context, origin, annotation, arg)
if origin := get_origin(annot):
if origin is Union:
for converter in get_args(annot):
try:
return await maybe_coroutine(converter, arg, context)
except:
if converter is None:
context.view.undo()
return None
elif origin is Annotated:
converter: Callable[[str, Context], Any] = get_args(annot)[1] # the typehint affects the other if statement somehow
return await maybe_coroutine(converter, arg, context)
elif origin is Literal:
if arg in get_args(annot):
return arg
else:
raise InvalidLiteralArgument(arg)
else:
return await maybe_coroutine(annotation, arg, context)
annot: Callable[..., Any]
return await maybe_coroutine(annot, arg, context)
else:
return arg
async def parse_arguments(self, context: Context):
# please pr if you can think of a better way to do this
for parameter in self.parameters[2:]:
for name, parameter in list(self.signature.parameters.items())[2:]:
if parameter.kind == parameter.KEYWORD_ONLY:
try:
arg = await self.convert_argument(context.view.get_rest(), parameter.annotation, context)
except StopIteration:
if parameter.default is not parameter.empty:
arg = parameter.default
else:
raise
context.kwargs[parameter.name] = arg
context.kwargs[name] = await self.convert_argument(context.view.get_rest(), parameter, context)
elif parameter.kind == parameter.VAR_POSITIONAL:
with suppress(StopIteration):
while True:
context.args.append(await self.convert_argument(context.view.get_next_word(), parameter.annotation, context))
context.args.append(await self.convert_argument(context.view.get_next_word(), parameter, context))
elif parameter.kind == parameter.POSITIONAL_OR_KEYWORD:
try:
rest = context.view.get_next_word()
arg = await self.convert_argument(rest, parameter.annotation, context)
except StopIteration:
if parameter.default is not parameter.empty:
arg = parameter.default
else:
raise
context.args.append(arg)
context.args.append(await self.convert_argument(context.view.get_next_word(), parameter, context))
def __repr__(self) -> str:
return f"<{self.__class__.__name__} name=\"{self.name}\">"
return f"<Command name=\"{self.name}\">"
@property
def short_description(self) -> Optional[str]:
"""Returns the first line of the description or None if there is no description."""
if self.description:
return self.description.split("\n")[0]
def get_usage(self) -> str:
"""Returns the usage string for the command."""
if self.usage:
return self.usage
parents = []
if self.parent:
parent = self.parent
while parent:
parents.append(parent.name)
parent = parent.parent
parameters = []
for parameter in self.parameters[2:]:
if parameter.kind == parameter.POSITIONAL_OR_KEYWORD:
if parameter.default is not parameter.empty:
parameters.append(f"[{parameter.name}]")
else:
parameters.append(f"<{parameter.name}>")
elif parameter.kind == parameter.KEYWORD_ONLY:
if parameter.default is not parameter.empty:
parameters.append(f"[{parameter.name}]")
else:
parameters.append(f"<{parameter.name}...>")
elif parameter.kind == parameter.VAR_POSITIONAL:
parameters.append(f"[{parameter.name}...]")
return f"{' '.join(parents[::-1])} {self.name} {' '.join(parameters)}"
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command] = Command, usage: Optional[str] = None):
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command] = Command):
"""A decorator that turns a function into a :class:`Command`.
Parameters
@@ -232,6 +160,6 @@ def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None,
A function that takes the command callback and returns a :class:`Command`
"""
def inner(func: Callable[..., Coroutine[Any, Any, Any]]):
return cls(func, name or func.__name__, aliases or [], usage)
return cls(func, name or func.__name__, aliases or [])
return inner
+8 -23
View File
@@ -1,24 +1,21 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Generic, Optional, TypeVar
from typing import TYPE_CHECKING, Any, Optional
import revolt
from revolt.utils import maybe_coroutine
from .command import Command
from .group import Group
if TYPE_CHECKING:
from .client import CommandsClient
from .view import StringView
ClientT = TypeVar("ClientT", bound="CommandsClient")
__all__ = (
"Context",
)
class Context(revolt.Messageable, Generic[ClientT]):
class Context(revolt.Messageable):
"""Stores metadata the commands execution.
Attributes
@@ -44,10 +41,10 @@ class Context(revolt.Messageable, Generic[ClientT]):
"""
__slots__ = ("command", "invoked_with", "args", "message", "server", "channel", "author", "view", "kwargs", "state", "client")
async def _get_channel_id(self) -> str:
def _get_channel_id(self) -> str:
return self.channel.id
def __init__(self, command: Optional[Command], invoked_with: str, view: StringView, message: revolt.Message, client: ClientT):
def __init__(self, command: Optional[Command], invoked_with: str, view: StringView, message: revolt.Message, client: CommandsClient):
self.command = command
self.invoked_with = invoked_with
self.view = view
@@ -61,7 +58,7 @@ class Context(revolt.Messageable, Generic[ClientT]):
self.state = message.state
async def invoke(self) -> Any:
"""Invokes the command.
"""Invokes the command, this is equal to `await command.invoke(context, context.args)`.
.. note:: If the command is `None`, this function will do nothing.
@@ -72,23 +69,11 @@ class Context(revolt.Messageable, Generic[ClientT]):
"""
if command := self.command:
if isinstance(command, Group):
try:
subcommand_name = self.view.get_next_word()
except StopIteration:
pass
else:
if subcommand := command.subcommands.get(subcommand_name):
self.command = command = subcommand
return await self.invoke()
self.view.undo()
await command.parse_arguments(self)
return await command.invoke(self, *self.args, **self.kwargs)
async def can_run(self, command: Optional[Command] = None) -> bool:
async def can_run(self) -> 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 (self.command.checks if self.command else [])])
return all([await maybe_coroutine(check, self) for check in (command.checks if command else [])])
+2 -70
View File
@@ -1,17 +1,8 @@
import re
from typing import Annotated
from revolt import Category, Channel, Member, User, utils
from .errors import BadBoolArgument
from .context import Context
from .errors import (BadBoolArgument, CategoryConverterError,
ChannelConverterError, MemberConverterError, ServerOnly,
UserConverterError)
__all__ = ("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})>")
IntConverter = Annotated[int, lambda arg, _: int(arg)]
def bool_converter(arg: str, _):
lowered = arg.lower()
@@ -22,63 +13,4 @@ def bool_converter(arg: str, _):
else:
raise BadBoolArgument(lowered)
def category_converter(arg: str, context: Context) -> Category:
if not (server := context.server):
raise ServerOnly
try:
return server.get_category(arg)
except KeyError:
try:
return utils.get(server.categories, name=arg)
except LookupError:
raise CategoryConverterError(arg)
def channel_converter(arg: str, context: Context) -> Channel:
if not (server := context.server):
raise ServerOnly
if (match := channel_regex.match(arg)):
arg = match.group(1)
try:
return server.get_channel(arg)
except KeyError:
try:
return utils.get(server.channels, name=arg)
except LookupError:
raise ChannelConverterError(arg)
def user_converter(arg: str, context: Context) -> User:
if (match := user_regex.match(arg)):
arg = match.group(1)
try:
return context.client.get_user(arg)
except KeyError:
try:
return utils.get(context.client.users, name=arg)
except LookupError:
raise UserConverterError(arg)
def member_converter(arg: str, context: Context) -> Member:
if not (server := context.server):
raise ServerOnly
if (match := user_regex.match(arg)):
arg = match.group(1)
try:
return server.get_member(arg)
except KeyError:
try:
return utils.get(server.members, name=arg)
except LookupError:
raise MemberConverterError(arg)
IntConverter = Annotated[int, lambda arg, _: int(arg)]
BoolConverter = Annotated[bool, bool_converter]
CategoryConverter = Annotated[Category, category_converter]
UserConverter = Annotated[User, user_converter]
MemberConverter = Annotated[Member, member_converter]
ChannelConverter = Annotated[Channel, channel_converter]
+1 -34
View File
@@ -10,12 +10,7 @@ __all__ = (
"ServerOnly",
"ConverterError",
"InvalidLiteralArgument",
"BadBoolArgument",
"CategoryConverterError",
"ChannelConverterError",
"UserConverterError",
"MemberConverterError",
"MissingSetup",
"BadBoolArgument"
)
class CommandError(RevoltError):
@@ -57,31 +52,3 @@ class InvalidLiteralArgument(ConverterError):
class BadBoolArgument(ConverterError):
"""Raised when the bool converter fails"""
class CategoryConverterError(ConverterError):
"""Raised when the Category conveter fails"""
def __init__(self, argument: str):
self.argument = argument
class ChannelConverterError(ConverterError):
"""Raised when the Channel conveter fails"""
def __init__(self, argument: str):
self.argument = argument
class UserConverterError(ConverterError):
"""Raised when the Category conveter fails"""
def __init__(self, argument: str):
self.argument = argument
class MemberConverterError(ConverterError):
"""Raised when the Category conveter fails"""
def __init__(self, argument: str):
self.argument = argument
class UnionConverterError(ConverterError):
"""Raised when all converters in a union fails"""
def __init__(self, argument: str):
self.argument = argument
class MissingSetup(CommandError):
"""Raised when an extension is missing the `setup` function"""
-117
View File
@@ -1,117 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Callable, Coroutine, Optional
from .command import Command
if TYPE_CHECKING:
from .context import Context
__all__ = (
"Group",
"group"
)
class Group(Command):
"""Class for holding info about a group command.
Parameters
-----------
callback: Callable[..., Coroutine[Any, Any, Any]]
The callback for the group command
name: :class:`str`
The name of the command
aliases: list[:class:`str`]
The aliases of the group command
subcommands: dict[:class:`str`, :class:`Command`]
The group's subcommands.
"""
__slots__ = ("subcommands",)
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]):
self.subcommands: dict[str, Command] = {}
super().__init__(callback, name, aliases)
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command] = Command):
"""A decorator that turns a function into a :class:`Command` and registers the command as a subcommand.
Parameters
-----------
name: Optional[:class:`str`]
The name of the command, this defaults to the functions name
aliases: Optional[list[:class:`str`]]
The aliases of the command, defaults to no aliases
cls: type[:class:`Command`]
The class used for creating the command, this defaults to :class:`Command` but can be used to use a custom command subclass
Returns
--------
Callable[Callable[..., Coroutine], :class:`Command`]
A function that takes the command callback and returns a :class:`Command`
"""
def inner(func: Callable[..., Coroutine[Any, Any, Any]]):
command = cls(func, name or func.__name__, aliases or [])
command.parent = self
self.subcommands[command.name] = command
return command
return inner
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type["Group"]] = None):
"""A decorator that turns a function into a :class:`Group` and registers the command as a subcommand
Parameters
-----------
name: Optional[:class:`str`]
The name of the group command, this defaults to the functions name
aliases: Optional[list[:class:`str`]]
The aliases of the group command, defaults to no aliases
cls: type[:class:`Group`]
The class used for creating the command, this defaults to :class:`Group` but can be used to use a custom group subclass
Returns
--------
Callable[Callable[..., Coroutine], :class:`Group`]
A function that takes the command callback and returns a :class:`Group`
"""
cls = cls or type(self)
def inner(func: Callable[..., Coroutine[Any, Any, Any]]):
command = cls(func, name or func.__name__, aliases or [])
command.parent = self
self.subcommands[command.name] = command
return command
return inner
def __repr__(self) -> str:
return f"<Group name=\"{self.name}\">"
@property
def commands(self) -> list[Command]:
return list(self.subcommands.values())
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group] = Group):
"""A decorator that turns a function into a :class:`Group`
Parameters
-----------
name: Optional[:class:`str`]
The name of the group command, this defaults to the functions name
aliases: Optional[list[:class:`str`]]
The aliases of the group command, defaults to no aliases
cls: type[:class:`Group`]
The class used for creating the command, this defaults to :class:`Group` but can be used to use a custom group subclass
Returns
--------
Callable[Callable[..., Coroutine], :class:`Group`]
A function that takes the command callback and returns a :class:`Group`
"""
def inner(func: Callable[..., Coroutine[Any, Any, Any]]):
return cls(func, name or func.__name__, aliases or [])
return inner
-197
View File
@@ -1,197 +0,0 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from itertools import groupby
from typing import TYPE_CHECKING, Optional, TypedDict, Union
from typing_extensions import NotRequired
from .client import CommandsClient
from .command import Command, command
from .context import Context
from .group import Group
from .utils import evaluate_parameters
if TYPE_CHECKING:
from revolt import File, Message, Messageable, MessageReply, SendableEmbed
from .cog import Cog
__all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl")
class MessagePayload(TypedDict):
content: str
embed: NotRequired[SendableEmbed]
embeds: NotRequired[list[SendableEmbed]]
attachments: NotRequired[list[File]]
replies: NotRequired[list[MessageReply]]
class HelpCommand(ABC):
@abstractmethod
async def create_bot_help(self, context: Context, commands: dict[Optional[Cog], list[Command]]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_command_help(self, context: Context, command: Command) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_group_help(self, context: Context, group: Group) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_cog_help(self, context: Context, cog: Cog) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
async def send_help_command(self, context: Context, message_payload: MessagePayload) -> Message:
return await context.send(**message_payload)
async def filter_commands(self, context: Context, commands: list[Command]) -> list[Command]:
filtered: list[Command] = []
for command in commands:
try:
if await context.can_run(command):
filtered.append(command)
except Exception:
pass
return filtered
async def group_commands(self, context: Context, commands: list[Command]) -> dict[Optional[Cog], list[Command]]:
cogs = {}
for command in commands:
cogs.setdefault(command.cog, []).append(command)
return cogs
async def handle_message(self, context: Context, message: Message):
pass
async def get_channel(self, context: Context) -> Messageable:
return context
@abstractmethod
async def handle_no_command_found(self, context: Context, name: str):
raise NotImplementedError
@abstractmethod
async def handle_no_cog_found(self, context: Context, name: str):
raise NotImplementedError
class DefaultHelpCommand(HelpCommand):
def __init__(self, default_cog_name: str = "No Cog"):
self.default_cog_name = default_cog_name
async def create_bot_help(self, context: Context, commands: dict[Optional[Cog], list[Command]]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
for cog, cog_commands in commands.items():
cog_lines = []
cog_lines.append(f"{cog.qualified_name if cog else self.default_cog_name}:")
for command in cog_commands:
cog_lines.append(f" {command.name} - {command.short_description or 'No description'}")
lines.append("\n".join(cog_lines))
lines.append("```")
return "\n".join(lines)
async def create_cog_help(self, context: Context, cog: Cog) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{cog.qualified_name}:")
for command in cog.commands:
lines.append(f" {command.name} - {command.short_description or 'No description'}")
lines.append("```")
return "\n".join(lines)
async def create_command_help(self, context: Context, command: Command) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{command.name}:")
lines.append(f" Usage: {command.get_usage()}")
if command.aliases:
lines.append(f" Aliases: {', '.join(command.aliases)}")
if command.description:
lines.append(command.description)
lines.append("```")
return "\n".join(lines)
async def create_group_help(self, context: Context, group: Group) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{group.name}:")
lines.append(f" Usage: {group.get_usage()}")
if group.aliases:
lines.append(f" Aliases: {', '.join(group.aliases)}")
if group.description:
lines.append(group.description)
for command in group.commands:
lines.append(f" {command.name} - {command.short_description or 'No description'}")
lines.append("```")
return "\n".join(lines)
async def handle_no_command_found(self, context: Context, name: str):
channel = await self.get_channel(context)
await channel.send(f"Command `{name}` not found.")
async def handle_no_cog_found(self, context: Context, name: str):
channel = await self.get_channel(context)
await channel.send(f"Cog `{name}` not found.")
class HelpCommandImpl(Command):
def __init__(self, client: CommandsClient):
self.client = client
super().__init__(callback=lambda _, context, *args: help_command_impl(self.client, context, *args), name="help", aliases=[])
self.description = "Shows help for a command, cog or the entire bot"
async def help_command_impl(self: CommandsClient, context: Context, *arguments: str):
filtered_commands = await context.client.help_command.filter_commands(context, self.commands)
commands = await self.help_command.group_commands(context, filtered_commands)
if not arguments:
payload = await self.help_command.create_bot_help(context, commands)
else:
command_name = arguments[0]
try:
command = self.get_command(command_name)
except KeyError:
cog = self.cogs.get(command_name)
if cog:
payload = await self.help_command.create_cog_help(context, cog)
else:
return await self.help_command.handle_no_command_found(context, command_name)
else:
if isinstance(command, Group):
payload = await self.help_command.create_group_help(context, command)
else:
payload = await self.help_command.create_command_help(context, command)
msg_payload: MessagePayload
if isinstance(payload, str):
msg_payload = {"content": payload}
elif isinstance(payload, SendableEmbed):
msg_payload = {"embed": payload, "content": " "}
else:
msg_payload = payload
await self.help_command.send_help_command(context, msg_payload)
-16
View File
@@ -1,16 +0,0 @@
from inspect import Parameter
from typing import Any, Iterable
__all__ = ("evaluate_parameters",)
def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]:
new_parameters = []
for parameter in parameters:
if parameter.annotation is not parameter.empty:
if isinstance(parameter.annotation, str):
parameter = parameter.replace(annotation=eval(parameter.annotation, globals))
new_parameters.append(parameter)
return new_parameters
+2 -5
View File
@@ -14,9 +14,6 @@ class StringView:
return next(self.value)
def get_rest(self) -> str:
if self.should_undo:
return f"{self.temp} {''.join(self.value)}"
return "".join(self.value)
def get_next_word(self) -> str:
@@ -25,7 +22,7 @@ class StringView:
return self.temp
char = self.next_char()
temp: list[str] = []
temp = []
while char == " ":
char = self.next_char()
@@ -41,7 +38,7 @@ class StringView:
else:
temp.append(char)
try:
while (char := self.next_char()) not in " \n":
while (char := self.next_char()) != " ":
temp.append(char)
except StopIteration:
pass
+4 -4
View File
@@ -6,7 +6,7 @@ __all__ = ("File",)
class File:
"""Respresents a file about to be uploaded to revolt
Parameters
-----------
file: Union[str, bytes]
@@ -17,7 +17,7 @@ class File:
Determines if the file will be a spoiler, this prefexes the filename with `SPOILER_`
"""
__slots__ = ("f", "spoiler", "filename")
def __init__(self, file: Union[str, bytes], *, filename: Optional[str] = None, spoiler: bool = False):
if isinstance(file, str):
self.f = open(file, "rb")
@@ -26,10 +26,10 @@ class File:
if filename is None and isinstance(file, str):
filename = self.f.name
self.spoiler = spoiler or (filename and filename.startswith("SPOILER_"))
if self.spoiler and (filename and not filename.startswith("SPOILER_")):
filename = f"SPOILER_{filename}"
self.filename = filename
self.filename = filename
+2 -12
View File
@@ -33,14 +33,8 @@ class flag_value:
instance._set_flag(self.flag, value)
class Flags:
FLAG_NAMES: list[str]
def __init_subclass__(cls) -> None:
flags = cls._flags()
cls.FLAG_NAMES = list(flags.keys())
def __init__(self, value: int = 0, **kwargs: bool):
self.value = value
def __init__(self, **kwargs: bool):
self.value = 0
for k, v in kwargs.items():
setattr(self, k, v)
@@ -98,10 +92,6 @@ class Flags:
def __hash__(self) -> int:
return hash(self.value)
@classmethod
def _flags(cls) -> dict[str, flag_value]:
return {name: value for name, value in cls.__dict__.items() if isinstance(value, flag_value)}
class UserBadges(Flags):
"""Contains all user badges"""
+86 -131
View File
@@ -6,8 +6,6 @@ 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 .file import File
@@ -25,40 +23,38 @@ if TYPE_CHECKING:
from .types import Autumn as AutumnPayload
from .types import Channel, DMChannel
from .types import Embed as EmbedPayload
from .types import GetServerMembers, GroupDMChannel, Invite
from .types import GetServerMembers
from .types import Masquerade as MasqueradePayload
from .types import Member
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 (MessageReplyPayload, MessageWithUserData, Role, Server,
ServerBans, ServerInvite, TextChannel, JoinCallResponse)
from .types import User as UserPayload
from .types import UserProfile, VoiceChannel
__all__ = ("HttpClient",)
T = TypeVar("T")
Request = Coroutine[Any, Any, T]
class HttpClient:
__slots__ = ("session", "token", "api_url", "api_info", "auth_header")
__slots__ = ("session", "token", "api_url", "api_info")
def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo, bot: bool = True):
def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo):
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"
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}"
kwargs = {}
headers = {
"User-Agent": "Revolt.py (https://github.com/revoltchat/revolt.py)",
self.auth_header: self.token
"User-Agent": "Revolt.py https://github.com/Zomatree/revolt.py",
"x-bot-token": self.token
}
if json:
@@ -77,10 +73,7 @@ class HttpClient:
async with self.session.request(method, url, **kwargs) as resp:
text = await resp.text()
if text:
try:
response = _json.loads(await resp.text())
except ValueError:
raise HTTPError(f"Invalid json response:\n{text}") from None
response = _json.loads(await resp.text())
else:
response = text
@@ -95,7 +88,7 @@ class HttpClient:
url = f"{self.api_info['features']['autumn']['url']}/{tag}"
headers = {
"User-Agent": "Revolt.py (https://github.com/revoltchat/revolt.py)"
"User-Agent": "Revolt.py https://github.com/Zomatree/revolt.py"
}
form = aiohttp.FormData()
@@ -113,7 +106,7 @@ class HttpClient:
else:
return response
async def send_message(self, channel: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]], attachments: Optional[list[File]], replies: Optional[list[MessageReplyPayload]], masquerade: Optional[MasqueradePayload]) -> MessagePayload:
async def send_message(self, channel: str, content: Optional[str], embeds: Optional[list[EmbedPayload]], attachments: Optional[list[File]], replies: Optional[list[MessageReplyPayload]], masquerade: Optional[MasqueradePayload]) -> MessagePayload:
json: dict[str, Any] = {}
if content:
@@ -139,15 +132,8 @@ class HttpClient:
return await self.request("POST", f"/channels/{channel}/messages", json=json)
def edit_message(self, channel: str, message: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]] = None) -> Request[None]:
json = {}
if content is not None:
json["content"] = content
if embeds is not None:
json["embeds"] = embeds
def edit_message(self, channel: str, message: str, content: str) -> Request[None]:
json = {"content": content}
return self.request("PATCH", f"/channels/{channel}/messages/{message}", json=json)
def delete_message(self, channel: str, message: str) -> Request[None]:
@@ -155,44 +141,44 @@ class HttpClient:
def fetch_message(self, channel: str, message: str) -> Request[MessagePayload]:
return self.request("GET", f"/channels/{channel}/messages/{message}")
@overload
def fetch_messages(
self,
channel: str,
self,
channel: str,
sort: SortType,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
nearby: Optional[str] = ...,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
nearby: Optional[str] = ...,
include_users: Literal[False] = ...
) -> Request[list[MessagePayload]]:
...
@overload
def fetch_messages(
self,
channel: str,
self,
channel: str,
sort: SortType,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
nearby: Optional[str] = ...,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
nearby: Optional[str] = ...,
include_users: Literal[True] = ...
) -> Request[MessageWithUserData]:
...
def fetch_messages(
self,
channel: str,
self,
channel: str,
sort: SortType,
*,
limit: Optional[int] = None,
before: Optional[str] = None,
after: Optional[str] = None,
nearby: Optional[str] = None,
*,
limit: Optional[int] = None,
before: Optional[str] = None,
after: Optional[str] = None,
nearby: Optional[str] = None,
include_users: bool = False
) -> Request[Union[list[MessagePayload], MessageWithUserData]]:
@@ -214,12 +200,12 @@ class HttpClient:
@overload
def search_messages(
self,
channel: str,
self,
channel: str,
query: str,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
sort: Optional[SortType] = ...,
include_users: Literal[False] = ...
@@ -228,12 +214,12 @@ class HttpClient:
@overload
def search_messages(
self,
channel: str,
self,
channel: str,
query: str,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
*,
limit: Optional[int] = ...,
before: Optional[str] = ...,
after: Optional[str] = ...,
sort: Optional[SortType] = ...,
include_users: Literal[True] = ...
@@ -241,12 +227,12 @@ class HttpClient:
...
def search_messages(
self,
channel: str,
self,
channel: str,
query: str,
*,
limit: Optional[int] = None,
before: Optional[str] = None,
*,
limit: Optional[int] = None,
before: Optional[str] = None,
after: Optional[str] = None,
sort: Optional[SortType] = None,
include_users: bool = False
@@ -280,8 +266,8 @@ class HttpClient:
def fetch_default_avatar(self, user_id: str) -> Request[bytes]:
return self.request_file(f"{self.api_url}/users/{user_id}/default_avatar")
def fetch_dm_channels(self) -> Request[list[Union[DMChannel, GroupDMChannel]]]:
def fetch_dm_channels(self) -> Request[list[Channel]]:
return self.request("GET", "/users/dms")
def open_dm(self, user_id: str) -> Request[DMChannel]:
@@ -293,20 +279,20 @@ class HttpClient:
def close_channel(self, channel_id: str) -> Request[None]:
return self.request("DELETE", f"/channels/{channel_id}")
def set_channel_role_permissions(self, channel_id: str, role_id: str, channel_permissions: int) -> Request[None]:
payload = {"permissions": channel_permissions}
return self.request("PUT", f"/channels/{channel_id}/permissions/{role_id}", json=payload)
def set_channel_default_permissions(self, channel_id: str, channel_permissions: int) -> Request[None]:
payload = {"permissions": channel_permissions}
return self.request("PUT", f"/channels/{channel_id}/permissions/default", json=payload)
def fetch_server(self, server_id: str) -> Request[Server]:
return self.request("GET", f"/servers/{server_id}")
def delete_leave_server(self, server_id: str) -> Request[None]:
return self.request("DELETE", f"/servers/{server_id}")
@overload
def create_channel(self, server_id: str, channel_type: Literal["Text"], name: str, description: Optional[str]) -> Request[TextChannel]:
...
@overload
def create_channel(self, server_id: str, channel_type: Literal["Voice"], name: str, description: Optional[str]) -> Request[VoiceChannel]:
...
def create_channel(self, server_id: str, channel_type: Literal["Text", "Voice"], name: str, description: Optional[str]) -> Request[Union[TextChannel, VoiceChannel]]:
payload = {
"type": channel_type,
@@ -318,7 +304,7 @@ class HttpClient:
return self.request("POST", f"/servers/{server_id}/channels", json=payload)
def fetch_server_invites(self, server_id: str) -> Request[list[PartialInvite]]:
def fetch_server_invites(self, server_id: str) -> Request[list[ServerInvite]]:
return self.request("GET", f"/servers/{server_id}/invites")
def fetch_member(self, server_id: str, member_id: str) -> Request[Member]:
@@ -337,66 +323,35 @@ class HttpClient:
def unban_member(self, server_id: str, member_id: str) -> Request[None]:
return self.request("DELETE", f"/servers/{server_id}/bans/{member_id}")
def fetch_bans(self, server_id: str) -> Request[ServerBans]:
return self.request("GET", f"/servers/{server_id}/bans")
def set_role_permissions(self, server_id: str, role_id: str, server_permissions: int, channel_permissions: int) -> Request[None]:
payload = {
"permissions": {
"server": server_permissions,
"channel": channel_permissions
}
}
return self.request("PUT", f"/servers/{server_id}/permissions/{role_id}", json=payload, nonce=False)
def set_default_permissions(self, server_id: str, server_permissions: int, channel_permissions: int) -> Request[None]:
payload = {
"permissions": {
"server": server_permissions,
"channel": channel_permissions
}
}
return self.request("PUT", f"/servers/{server_id}/permissions/default", json=payload, nonce=False)
def create_role(self, server_id: str, name: str) -> Request[Role]:
return self.request("POST", f"/servers/{server_id}/roles", json={"name": name}, nonce=False)
def delete_role(self, server_id: str, role_id: str) -> Request[None]:
return self.request("DELETE", f"/servers/{server_id}/roles/{role_id}")
def fetch_invite(self, code: str) -> Request[Invite]:
return self.request("GET", f"/invites/{code}")
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]):
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]):
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]):
if remove:
values["remove"] = remove
if avatar := values.get("avatar"):
asset = await self.upload_file(avatar, "avatars")
values["avatar"] = asset["id"]
if profile := values.get("profile"):
if background := profile.background():
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:
return self.request("PUT", f"/channels/{channel_id}/permissions/default", json={"permissions": {"allow": allow, "deny": deny}})
def set_guild_channel_role_permissions(self, channel_id: str, role_id: str, allow: int, deny: int) -> Request:
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):
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_default_permissions(self, server_id: str, value: int):
return self.request("PUT", f"/server/{server_id}/permissions/default", json={"permissions": value})
def connect_to_voice(self, channel_id) -> Request[JoinCallResponse]:
return self.request("POST", f"/channels/{channel_id}/join_call")
-67
View File
@@ -1,67 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from .asset import Asset
if TYPE_CHECKING:
from .state import State
from .types import Invite as InvitePayload
__all__ = ("Invite",)
class Invite:
"""Represents a server invite.
Attributes
-----------
code: :class:`str`
The code for the invite
server: :class:`Server`
The server this invite is for
channel: :class:`Channel`
The channel this invite is for
user_name: :class:`str`
The name of the user who made the invite
user: Optional[:class:`User`]
The user who made the invite, this is only set if this was fetched via :meth:`Server.fetch_invites`
user_avatar: Optional[:class:`Asset`]
The invite creator's avatar, if any
member_count: :class:`int`
The member count of the server this invite is for
"""
def __init__(self, data: InvitePayload, code: str, state: State):
self.state = state
self.code = code
self.server = state.get_server(data["server_id"])
self.channel = self.server.get_channel(data["channel_id"])
self.user_name = data["user_name"]
self.user = None
if avatar := data.get("user_avatar"):
self.user_avatar = Asset(avatar, state)
else:
self.user_avatar = None
self.member_count = data["member_count"]
@staticmethod
def _from_partial(code: str, server: str, creator: str, channel: str, state: State) -> Invite:
invite = Invite.__new__(Invite)
invite.state = state
invite.code = code
invite.server = state.get_server(server)
invite.channel = state.get_channel(channel)
invite.user = state.get_user(creator)
invite.user_name = invite.user.name
invite.user_avatar = invite.user.avatar
invite.member_count = len(invite.server.members)
return invite
async def delete(self):
"""Deletes the invite"""
await self.state.http.delete_invite(self.code)
+3 -30
View File
@@ -19,7 +19,7 @@ def flattern_user(member: Member, user: User):
class Member(User):
"""Represents a member of a server, subclasses :class:`User`
Attributes
-----------
nickname: Optional[:class:`str`]
@@ -32,16 +32,13 @@ class Member(User):
The member's guild avatar if any
"""
__slots__ = ("_state", "nickname", "roles", "server", "guild_avatar")
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.append(self)
self._state = state
self.nickname = data.get("nickname")
if avatar := data.get("avatar"):
self.guild_avatar = Asset(avatar, state)
@@ -52,18 +49,12 @@ class Member(User):
self.roles = sorted(roles, key=lambda role: role.rank, reverse=True)
self.server = server
self.nickname = data.get("nickname")
@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 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:
self.nickname = nickname
@@ -74,21 +65,3 @@ class Member(User):
if roles:
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):
"""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):
"""Bans the member from the server
Parameters
-----------
reason: Optional[:class:`str`]
The reason for the ban
"""
await self.state.http.ban_member(self.server.id, self.id, reason)
async def unban(self):
"""Unbans the member from the server"""
await self.state.http.unban_member(self.server.id, self.id)
+18 -39
View File
@@ -5,15 +5,14 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
from .asset import Asset, PartialAsset
from .channel import Messageable
from .embed import SendableEmbed, to_embed
from .embed import Embed
if TYPE_CHECKING:
from .state import State
from .types import Embed as EmbedPayload
from .types import Masquerade as MasqueradePayload
from .types import Message as MessagePayload
from .types import MessageReplyPayload
from .server import Server
__all__ = (
"Message",
@@ -32,10 +31,12 @@ class Message:
The content of the message, this will not include system message's content
attachments: list[:class:`Asset`]
The attachments of the message
embeds: list[Union[:class:`WebsiteEmbed`, :class:`ImageEmbed`, :class:`TextEmbed`, :class:`NoneEmbed`]]
embeds: list[:class:`Embed`]
The embeds of the message
channel: :class:`Messageable`
The channel the message was sent in
server: :class:`Server`
The server the message was sent in
author: Union[:class:`Member`, :class:`User`]
The author of the message, will be :class:`User` in DMs
edited_at: Optional[:class:`datetime.datetime`]
@@ -47,7 +48,7 @@ class Message:
reply_ids: list[:class:`str`]
The message's ids this message has replies to
"""
__slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "mentions", "replies", "reply_ids")
__slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "server", "author", "edited_at", "mentions", "replies", "reply_ids")
def __init__(self, data: MessagePayload, state: State):
self.state = state
@@ -55,14 +56,16 @@ class Message:
self.id = data["_id"]
self.content = data["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.embeds = [Embed.from_dict(embed) for embed in data.get("embeds", [])]
channel = state.get_channel(data["channel"])
assert isinstance(channel, Messageable)
self.channel = channel
if server_id := self.channel.server_id:
author = state.get_member(server_id, data["author"])
self.server = self.channel and self.channel.server
if self.server:
author = state.get_member(self.server.id, data["author"])
else:
author = state.get_user(data["author"])
@@ -95,47 +98,29 @@ class Message:
self.reply_ids.append(reply)
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited_at: str):
def _update(self, *, content: Optional[str] = None, edited_at: Optional[str] = None) -> Message:
if content:
self.content = content
self.edited_at = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z")
# strptime is used here instead of fromisoformat because of its inability to parse `Z` (Zulu or UTC time) in the RFCC 3339 format provided by API
if edited_at:
self.edited_at = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z")
# strptime is used here instead of fromisoformat because of its inability to parse `Z` (Zulu or UTC time) in the RFCC 3339 format provided by API
if embeds:
self.embeds = [to_embed(embed, self.state) for embed in embeds]
return self
async def edit(self, *, content: Optional[str] = None, embeds: Optional[list[SendableEmbed]] = None) -> None:
async def edit(self, *, content: str) -> None:
"""Edits the message. The bot can only edit its own message
Parameters
-----------
content: :class:`str`
The new content of the message
"""
new_embeds = [embed.to_dict() for embed in embeds] if embeds else None
await self.state.http.edit_message(self.channel.id, self.id, content, new_embeds)
await self.state.http.edit_message(self.channel.id, self.id, content)
async def delete(self) -> None:
"""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, mention: bool = False, **kwargs):
"""Replies to this message, equivilant to:
.. code-block:: python
await channel.send(..., replies=[MessageReply(message, mention)])
"""
return self.channel.send(*args, **kwargs, replies=[MessageReply(self, mention)])
@property
def server(self) -> Server:
""":class:`Server` The server this voice channel belongs too"""
return self.channel.server
class MessageReply(NamedTuple):
"""A namedtuple which represents a reply to a message.
@@ -161,12 +146,9 @@ class Masquerade(NamedTuple):
The name to display for the message
avatar: Optional[:class:`str`]
The avatar's url to display for the message
colour: Optional[:class:`str`]
The colour of the name, similar to role colours
"""
name: Optional[str] = None
avatar: Optional[str] = None
colour: Optional[str] = None
def to_dict(self) -> MasqueradePayload:
output: MasqueradePayload = {}
@@ -177,7 +159,4 @@ class Masquerade(NamedTuple):
if avatar := self.avatar:
output["avatar"] = avatar
if colour := self.colour:
output["colour"] = colour
return output
+6 -75
View File
@@ -2,10 +2,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from .enums import SortType
if TYPE_CHECKING:
from .embed import Embed, SendableEmbed
from .embed import Embed
from .file import File
from .message import Masquerade, Message, MessageReply
from .state import State
@@ -25,10 +23,10 @@ class Messageable:
__slots__ = ()
async def _get_channel_id(self) -> str:
def _get_channel_id(self) -> str:
raise NotImplementedError
async def send(self, content: Optional[str] = None, *, embeds: Optional[list[SendableEmbed]] = None, embed: Optional[SendableEmbed] = None, attachments: Optional[list[File]] = None, replies: Optional[list[MessageReply]] = None, reply: Optional[MessageReply] = None, masquerade: Optional[Masquerade] = None) -> Message:
async def send(self, content: Optional[str] = None, *, embeds: Optional[list[Embed]] = None, embed: Optional[Embed] = None, attachments: Optional[list[File]] = None, replies: Optional[list[MessageReply]] = None, reply: Optional[MessageReply] = None, masquerade: Optional[Masquerade] = None) -> Message:
"""Sends a message in a channel, you must send at least one of either `content`, `embeds` or `attachments`
Parameters
@@ -37,10 +35,8 @@ class Messageable:
The content of the message, this will not include system message's content
attachments: Optional[list[:class:`File`]]
The attachments of the message
embed: Optional[:class:`SendableEmbed`]
The embed to send with the message
embeds: Optional[list[:class:`SendableEmbed`]]
The embeds to send with the message
embeds: Optional[list[:class:`Embed`]]
The embeds of the message
replies: Optional[list[:class:`MessageReply`]]
The list of messages to reply to.
@@ -59,70 +55,5 @@ class Messageable:
reply_payload = [reply.to_dict() for reply in replies] if replies else None
masquerade_payload = masquerade.to_dict() if masquerade else None
message = await self.state.http.send_message(await self._get_channel_id(), content, embed_payload, attachments, reply_payload, masquerade_payload)
message = await self.state.http.send_message(self._get_channel_id(), content, embed_payload, attachments, reply_payload, masquerade_payload)
return self.state.add_message(message)
async def fetch_message(self, message_id: str) -> Message:
"""Fetches a message from the channel
Parameters
-----------
message_id: :class:`str`
The id of the message you want to fetch
Returns
--------
:class:`Message`
The message with the matching id
"""
payload = await self.state.http.fetch_message(await self._get_channel_id(), message_id)
return Message(payload, self.state)
async def history(self, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None, nearby: Optional[str] = None) -> list[Message]:
"""Fetches multiple messages from the channel's history
Parameters
-----------
sort: :class:`SortType`
The order to sort the messages in
limit: :class:`int`
How many messages to fetch
before: Optional[:class:`str`]
The id of the message which should come *before* all the messages to be fetched
after: Optional[:class:`str`]
The id of the message which should come *after* all the messages to be fetched
nearby: Optional[:class:`str`]
The id of the message which should be nearby all the messages to be fetched
Returns
--------
list[:class:`Message`]
The messages found in order of the sort parameter
"""
payloads = await self.state.http.fetch_messages(await self._get_channel_id(), sort=sort, limit=limit, before=before, after=after, nearby=nearby)
return [Message(payload, self.state) for payload in payloads]
async def search(self, query: str, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None) -> list[Message]:
"""searches the channel for a query
Parameters
-----------
query: :class:`str`
The query to search for in the channel
sort: :class:`SortType`
The order to sort the messages in
limit: :class:`int`
How many messages to fetch
before: Optional[:class:`str`]
The id of the message which should come *before* all the messages to be fetched
after: Optional[:class:`str`]
The id of the message which should come *after* all the messages to be fetched
Returns
--------
list[:class:`Message`]
The messages found in order of the sort parameter
"""
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]
+121 -168
View File
@@ -1,200 +1,153 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional
from typing_extensions import Self
from .types.permissions import Overwrite
from .flags import Flags, flag_value
__all__ = ("Permissions", "PermissionsOverwrite")
class Permissions(Flags):
__all__ = (
"ChannelPermissions",
"ServerPermissions"
)
# Channel permissions
#
# View = 0b00000000000000000000000000000001 // 1
# SendMessage = 0b00000000000000000000000000000010 // 2
# ManageMessages = 0b00000000000000000000000000000100 // 4
# ManageChannel = 0b00000000000000000000000000001000 // 8
# VoiceCall = 0b00000000000000000000000000010000 // 16
# InviteOthers = 0b00000000000000000000000000100000 // 32
# EmbedLinks = 0b00000000000000000000000001000000 // 64
# UploadFiles = 0b00000000000000000000000010000000 // 128
# Server permissions
#
# View = 0b00000000000000000000000000000001 // 1
# ManageRoles = 0b00000000000000000000000000000010 // 2
# ManageChannels = 0b00000000000000000000000000000100 // 4
# ManageServer = 0b00000000000000000000000000001000 // 8
# KickMembers = 0b00000000000000000000000000010000 // 16
# BanMembers = 0b00000000000000000000000000100000 // 32
# ChangeNickname = 0b00000000000000000001000000000000 // 4096
# ManageNicknames = 0b00000000000000000010000000000000 // 8192
# ChangeAvatar = 0b00000000000000000100000000000000 // 16382
# RemoveAvatars = 0b00000000000000001000000000000000 // 32768
class ChannelPermissions(Flags):
"""Represents the channel permissions for a role as seen in channel settings."""
@classmethod
def none(cls) -> ChannelPermissions:
return cls._from_value(0)
@classmethod
def all(cls) -> ChannelPermissions:
return cls._from_value(0b11111011)
@classmethod
def view(cls) -> ChannelPermissions:
return cls._from_value(0b1)
@classmethod
def send_message(cls) -> ChannelPermissions:
return cls._from_value(0b11)
@classmethod
def manage_channel(cls) -> ChannelPermissions:
return cls._from_value(0b1001)
@classmethod
def voice_call(cls) -> ChannelPermissions:
return cls._from_value(0b10001)
@classmethod
def invite_others(cls) -> ChannelPermissions:
return cls._from_value(0b100001)
@classmethod
def embed_links(cls) -> ChannelPermissions:
return cls._from_value(0b1000001)
@classmethod
def upload_files(cls) -> ChannelPermissions:
return cls._from_value(0b10000001)
@flag_value
def manage_channel() -> int:
def can_view() -> int:
return 1 << 0
@flag_value
def manage_server() -> int:
def can_send_message() -> int:
return 1 << 1
@flag_value
def manage_permissions() -> int:
def can_manage_channel() -> int:
return 1 << 3
@flag_value
def can_voice_call() -> int:
return 1 << 4
@flag_value
def can_invite_others() -> int:
return 1 << 5
@flag_value
def can_embed_links() -> int:
return 1 << 6
@flag_value
def can_upload_files() -> int:
return 1 << 7
class ServerPermissions(Flags):
"""Represents the server permissions for a role as seen in server settings."""
@classmethod
def none(cls) -> ServerPermissions:
return cls._from_value(0)
@classmethod
def all(cls) -> ServerPermissions:
return cls._from_value(0b1111000000111111)
@flag_value
def view_server() -> int:
return 1 << 0
@flag_value
def manage_roles() -> int:
return 1 << 1
@flag_value
def manage_channels() -> int:
return 1 << 2
@flag_value
def manage_role() -> int:
def manage_server() -> int:
return 1 << 3
@flag_value
def kick_members() -> int:
return 1 << 6
return 1 << 4
@flag_value
def ban_members() -> int:
return 1 << 7
return 1 << 5
@flag_value
def timeout_members() -> int:
return 1 << 8
@flag_value
def asign_roles() -> int:
return 1 << 9
@flag_value
def change_nickname() -> int:
return 1 << 10
@flag_value
def manage_nicknames() -> int:
return 1 << 11
@flag_value
def change_avatars() -> int:
def change_nicknames() -> int:
return 1 << 12
@flag_value
def remove_avatars() -> int:
def manage_nicknames() -> int:
return 1 << 13
@flag_value
def view_channel() -> int:
return 1 << 20
def change_avatar() -> int:
return 1 << 14
@flag_value
def read_message_history() -> int:
return 1 << 21
@flag_value
def send_messages() -> int:
return 1 << 22
@flag_value
def manage_messages() -> int:
return 1 << 23
@flag_value
def manage_webhooks() -> int:
return 1 << 24
@flag_value
def invite_others() -> int:
return 1 << 25
@flag_value
def send_embeds() -> int:
return 1 << 26
@flag_value
def upload_files() -> int:
return 1 << 27
@flag_value
def masquerade() -> int:
return 1 << 28
@flag_value
def connect() -> int:
return 1 << 30
@flag_value
def speak() -> int:
return 1 << 31
@flag_value
def video() -> int:
return 1 << 32
@flag_value
def mute_members() -> int:
return 1 << 33
@flag_value
def deafen_members() -> int:
return 1 << 34
@flag_value
def move_members() -> int:
return 1 << 35
@classmethod
def all(cls) -> Self:
return cls(0x000F_FFFF_FFFF_FFFF)
@classmethod
def default_view_only(cls) -> Self:
return cls(view_channel=True, read_message_history=True)
@classmethod
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)
class PermissionsOverwrite:
def __init__(self, allow: Permissions, deny: Permissions):
self._allow = allow
self._deny = deny
for perm in Permissions.FLAG_NAMES:
if getattr(allow, perm):
value = True
elif getattr(deny, perm):
value = False
else:
value = None
super().__setattr__(perm, value)
def __setattr__(self, key: str, value: Any):
if key in Permissions.FLAG_NAMES:
if key is True:
setattr(self._allow, key, True)
super().__setattr__(key, True)
elif key is False:
setattr(self._deny, key, True)
super().__setattr__(key, False)
else:
setattr(self._allow, key, False)
setattr(self._deny, key, False)
super().__setattr__(key, None)
else:
super().__setattr__(key, value)
if TYPE_CHECKING:
manage_channel: Optional[bool]
manage_server: Optional[bool]
manage_permissions: Optional[bool]
manage_role: Optional[bool]
kick_members: Optional[bool]
ban_members: Optional[bool]
timeout_members: Optional[bool]
asign_roles: Optional[bool]
change_nickname: Optional[bool]
manage_nicknames: Optional[bool]
change_avatars: Optional[bool]
remove_avatars: Optional[bool]
view_channel: Optional[bool]
read_message_history: Optional[bool]
send_messages: Optional[bool]
manage_messages: Optional[bool]
manage_webhooks: Optional[bool]
invite_others: Optional[bool]
send_embeds: Optional[bool]
upload_files: Optional[bool]
masquerade: Optional[bool]
connect: Optional[bool]
speak: Optional[bool]
video: Optional[bool]
mute_members: Optional[bool]
deafen_members: Optional[bool]
move_members: Optional[bool]
def to_pair(self) -> tuple[Permissions, Permissions]:
return self._allow, self._deny
@classmethod
def _from_overwrite(cls, overwrite: Overwrite) -> Self:
allow = Permissions(overwrite["a"])
deny = Permissions(overwrite["d"])
return cls(allow, deny)
def remove_avatars() -> int:
return 1 << 15
+20 -29
View File
@@ -2,8 +2,9 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from .permissions import Permissions, PermissionsOverwrite
from .utils import Missing
from revolt.types import server
from .permissions import ChannelPermissions, ServerPermissions
if TYPE_CHECKING:
from .server import Server
@@ -22,7 +23,7 @@ class Role:
The id of the role
name: :class:`str`
The name of the role
colour: Optional[:class:`str`]
colour: :class:`str`
The colour of the role
hoist: :class:`bool`
Whether members with the role will display seperate from everyone else
@@ -35,23 +36,24 @@ class Role:
channel_permissions: :class:`ChannelPermissions`
The channel permissions for the role
"""
__slots__ = ("id", "name", "colour", "hoist", "rank", "state", "server", "permissions")
__slots__ = ("id", "name", "colour", "hoist", "rank", "state", "server", "server_permissions", "channel_permissions")
def __init__(self, data: RolePayload, role_id: str, server: Server, state: State):
def __init__(self, data: RolePayload, role_id: str, state: State, server: Server):
self.state = state
self.id = role_id
self.name = data["name"]
self.colour = None
self.hoist = False
self.rank = 0
self.colour = data.get("colour")
self.hoist = data.get("hoist", False)
self.rank = data.get("rank", 0)
self.server = server
self.permissions = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
self.server_permissions = ServerPermissions._from_value(data["permissions"][0])
self.channel_permissions = ChannelPermissions._from_value(data["permissions"][1])
@property
def color(self):
return self.colour
async def set_permissions_overwrite(self, *, permissions: PermissionsOverwrite) -> None:
async def set_permissions(self, *, server_permissions: Optional[ServerPermissions] = None, channel_permissions: Optional[ChannelPermissions] = None) -> None:
"""Sets the permissions for a role in a server.
Parameters
-----------
@@ -60,8 +62,14 @@ class Role:
channel_permissions: Optional[:class:`ChannelPermissions`]
The new channel permissions for the role
"""
allow, deny = permissions.to_pair()
await self.state.http.set_server_role_permissions(self.server.id, self.id, allow.value, deny.value)
if not server_permissions and not channel_permissions:
return
server_value = (server_permissions or self.server_permissions).value
channel_value = (channel_permissions or self.channel_permissions).value
await self.state.http.set_role_permissions(self.server.id, self.id, server_value, channel_value)
def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None):
if name:
@@ -75,20 +83,3 @@ class Role:
if rank:
self.rank = rank
async def delete(self):
"""Deletes the role"""
await self.state.http.delete_role(self.server.id, self.id)
async def edit(self, **kwargs):
"""Edits the role
Parameters
-----------
"""
if kwargs.get("colour", Missing) is None:
remove = "Colour"
else:
remove = None
await self.state.http.edit_role(self.server.id, self.id, remove, kwargs)
+33 -194
View File
@@ -1,26 +1,25 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional, cast
from typing import TYPE_CHECKING, Optional
from .asset import Asset
from .category import Category
from .channel import Channel, VoiceChannel
from .invite import Invite
from .permissions import Permissions
from .channel import Channel
from .permissions import ChannelPermissions, ServerPermissions
from .role import Role
if TYPE_CHECKING:
from .channel import TextChannel
from .member import Member
from .state import State
from .types import Ban
from .types import Category as CategoryPayload
from .types import File as FilePayload
from .types import Permission as PermissionPayload
from .types import Server as ServerPayload
from .types import SystemMessagesConfig
__all__ = ("Server", "SystemMessages", "ServerBan")
__all__ = ("Server", "SystemMessages")
class SystemMessages:
def __init__(self, data: SystemMessagesConfig, state: State):
@@ -83,25 +82,26 @@ class Server:
Whether the server is nsfw or not
system_messages: :class:`SystemMessages`
The system message config for the server
categories: list[:class:`Category`]
The categories in the server
icon: Optional[:class:`Asset`]
The servers icon
banner: Optional[:class:`Asset`]
The servers banner
default_permissions: :class:`Permissions`
The permissions for the default role
"""
__slots__ = ("state", "id", "name", "owner_id", "default_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "_categories")
__slots__ = ("state", "id", "name", "owner_id", "default_server_permissions", "default_channel_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "categories")
def __init__(self, data: ServerPayload, state: State):
self.state = state
self.id = data["_id"]
self.name = data["name"]
self.owner_id = data["owner"]
self.default_server_permissions = ServerPermissions._from_value(data["default_permissions"][0])
self.default_channel_permissions = ChannelPermissions._from_value(data["default_permissions"][1])
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.system_messages = SystemMessages(data.get("system_messages", {}), state)
self.categories = [Category(data, state) for data in data.get("categories", [])]
if icon := data.get("icon"):
self.icon = Asset(icon, state)
@@ -114,11 +114,12 @@ class Server:
self.banner = None
self._members: dict[str, Member] = {}
self._roles: dict[str, Role] = {role_id: Role(role, role_id, self, state) for role_id, role in data.get("roles", {}).items()}
self._roles: dict[str, Role] = {role_id: Role(role, role_id, state, self) for role_id, role in data.get("roles", {}).items()}
self._channels: dict[str, Channel] = {channel_id: state.get_channel(channel_id) for channel_id in data.get("channels", [])}
channels = [state.get_channel(channel_id) for channel_id in data["channels"]]
self._channels: dict[str, Channel] = {channel.id: channel for channel in channels}
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):
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[PermissionPayload] = None, nsfw: Optional[bool] = None, system_messages: Optional[SystemMessagesConfig] = None, categories: Optional[list[CategoryPayload]] = None):
if owner:
self.owner_id = owner
if name:
@@ -129,22 +130,21 @@ class Server:
self.icon = Asset(icon, self.state)
if banner:
self.banner = Asset(banner, self.state)
if default_permissions is not None:
self.default_permissions = Permissions(default_permissions)
if default_permissions:
self.default_server_permissions = ServerPermissions._from_value(default_permissions[0])
self.default_channel_permissions = ChannelPermissions._from_value(default_permissions[1])
if nsfw is not None:
self.nsfw = nsfw
if system_messages is not None:
self.system_messages = SystemMessages(system_messages, self.state)
if categories is not None:
self._categories = {data["id"]: Category(data, self.state) for data in categories}
if channels is not None:
self._channels = {channel_id: self.state.get_channel(channel_id) for channel_id in channels}
self.categories = [Category(data, self.state) for data in categories]
@property
def roles(self) -> list[Role]:
"""list[:class:`Role`] Gets all roles in the server in decending order"""
return list(self._roles.values())
@property
def members(self) -> list[Member]:
"""list[:class:`Member`] Gets all members in the server"""
@@ -155,19 +155,14 @@ class Server:
"""list[:class:`Member`] Gets all channels in the server"""
return list(self._channels.values())
@property
def categories(self) -> list[Category]:
"""list[:class:`Category`] Gets all categories in the server"""
return list(self._categories.values())
def get_role(self, role_id: str) -> Role:
"""Gets a role from the cache
Parameters
-----------
id: :class:`str`
The id of the role
Returns
--------
:class:`Role`
@@ -177,64 +172,40 @@ class Server:
def get_member(self, member_id: str) -> Member:
"""Gets a member from the cache
Parameters
-----------
id: :class:`str`
The id of the member
Returns
--------
:class:`Member`
The member
"""
try:
return self._members[member_id]
except KeyError as e:
raise LookupError from e
return self._members[member_id]
def get_channel(self, channel_id: str) -> Channel:
"""Gets a channel from the cache
Parameters
-----------
id: :class:`str`
The id of the channel
Returns
--------
:class:`Channel`
The channel
"""
try:
return self._channels[channel_id]
except KeyError as e:
raise LookupError from e
def get_category(self, category_id: str) -> Category:
"""Gets a category from the cache
Parameters
-----------
id: :class:`str`
The id of the category
Returns
--------
:class:`Category`
The category
"""
try:
return self._categories[category_id]
except KeyError as e:
raise LookupError from e
return self._channels[channel_id]
@property
def owner(self) -> Member:
""":class:`Member` The owner of the server"""
return self.get_member(self.owner_id)
async def set_default_permissions(self, permissions: Permissions) -> None:
async def set_default_permissions(self, *, server_permissions: Optional[ServerPermissions] = None, channel_permissions: Optional[ChannelPermissions] = None) -> None:
"""Sets the default server permissions.
Parameters
-----------
@@ -243,139 +214,7 @@ class Server:
channel_permissions: Optional[:class:`ChannelPermissions`]
the new default channel permissions
"""
server_value = (server_permissions or self.default_server_permissions).value
channel_value = (channel_permissions or self.default_channel_permissions).value
await self.state.http.set_server_default_permissions(self.id, permissions.value)
async def leave_server(self):
"""Leaves or deletes the server"""
await self.state.http.delete_leave_server(self.id)
async def delete_server(self):
"""Leaves or deletes a server, alias to :meth`Server.leave_server`"""
await self.leave_server()
async def create_text_channel(self, *, name: str, description: Optional[str] = None) -> TextChannel:
"""Creates a text channel in the server
Parameters
-----------
name: :class:`str`
The name of the channel
description: Optional[:class:`str`]
The channel's description
Returns
--------
:class:`TextChannel`
The text channel that was just created
"""
payload = await self.state.http.create_channel(self.id, "Text", name, description)
channel = TextChannel(payload, self.state)
self._channels[channel.id] = channel
return channel
async def create_voice_channel(self, *, name: str, description: Optional[str] = None) -> VoiceChannel:
"""Creates a voice channel in the server
Parameters
-----------
name: :class:`str`
The name of the channel
description: Optional[:class:`str`]
The channel's description
Returns
--------
:class:`VoiceChannel`
The voice channel that was just created
"""
payload = await self.state.http.create_channel(self.id, "Voice", name, description)
channel = VoiceChannel(payload, self.state)
self._channels[channel.id] = channel
return channel
async def fetch_invites(self) -> list[Invite]:
"""Fetches all invites in the server
Returns
--------
list[:class:`Invite`]
"""
invite_payloads = await self.state.http.fetch_server_invites(self.id)
return [Invite._from_partial(payload["_id"], payload["server"], payload["creator"], payload["channel"], self.state) for payload in invite_payloads]
async def fetch_member(self, member_id: str) -> Member:
"""Fetches a member from this server
Parameters
-----------
member_id: :class:`str`
The id of the member you are fetching
Returns
--------
:class:`Member`
The member with the matching id
"""
payload = await self.state.http.fetch_member(self.id, member_id)
return Member(payload, self, self.state)
async def fetch_bans(self) -> list[ServerBan]:
"""Fetches all invites in the server
Returns
--------
list[:class:`Invite`]
"""
payload = await self.state.http.fetch_bans(self.id)
return [ServerBan(ban, self.state) for ban in payload["bans"]]
async def create_role(self, name: str) -> Role:
"""Creates a role in the server
Parameters
-----------
name: :class:`str`
The name of the role
Returns
--------
:class:`Role`
The role that was just created
"""
payload = await self.state.http.create_role(self.id, name)
return Role(payload, name, self, self.state)
class ServerBan:
"""Represents a server ban
Attributes
-----------
reason: Optional[:class:str`]
The reason the user was banned
server: :class:`Server`
The server the user was banned in
user_id: :class:`str`
The id of the user who was banned
"""
__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
async def unban(self):
"""Unbans the user"""
await self.state.http.unban_member(self.server.id, self.user_id)
await self.state.http.set_default_permissions(self.id, server_value, channel_value)
Executable
+36
View File
@@ -0,0 +1,36 @@
import time
import random
import bitarray
class SecureRTP:
def __init__(self, salt: str, hash: str):
self.salt = salt
self.hash = hash
self.roc = 0
self.seq = 0
self.packet_counter = 0
self.ssrc = random.randbytes(4) # 32 bytes
self.version = [1, 0] # binary for 2
# documentation of header can be found at https://datatracker.ietf.org/doc/html/rfc3550#section-5.1
def create_packet(self, data: bytes) -> bytes:
arr = bitarray.bitarray()
arr.extend(self.version)
arr.extend([0, 0])
arr.extend([0, 0, 0, 0])
arr.append(0)
arr.extend([0, 0, 0, 0, 0, 0, 0])
seq = list(map(int, bin(self.seq)[2:]))
arr.extend(([0] * (16 - len(seq))) + seq)
timestamp = list(map(int, bin(int(time.monotonic()))[2:]))
arr.extend(([0] * (32 - len(timestamp))) + timestamp)
arr.extend(self.ssrc)
arr.extend(data)
return arr.tobytes()
+18 -29
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from collections import deque
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Callable, Optional
from .channel import Channel, channel_factory
from .member import Member
@@ -22,12 +22,13 @@ if TYPE_CHECKING:
__all__ = ("State",)
class State:
__slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages")
__slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages", "dispatch")
def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int):
def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int, dispatch: Callable[..., None]):
self.http = http
self.api_info = api_info
self.max_messages = max_messages
self.dispatch = dispatch
self.users: dict[str, User] = {}
self.channels: dict[str, Channel] = {}
@@ -35,26 +36,17 @@ class State:
self.messages: deque[Message] = deque()
def get_user(self, id: str) -> User:
try:
return self.users[id]
except KeyError as e:
raise LookupError from e
return self.users[id]
def get_member(self, server_id: str, member_id: str) -> Member:
server = self.servers[server_id]
return server.get_member(member_id)
def get_channel(self, id: str) -> Channel:
try:
return self.channels[id]
except KeyError as e:
raise LookupError from e
return self.channels[id]
def get_server(self, id: str) -> Server:
try:
return self.servers[id]
except KeyError as e:
raise LookupError from e
return self.servers[id]
def add_user(self, payload: UserPayload) -> User:
user = User(payload, self)
@@ -82,7 +74,7 @@ class State:
message = Message(payload, self)
if len(self.messages) >= self.max_messages:
self.messages.pop()
self.messages.appendleft(message)
return message
@@ -90,18 +82,15 @@ class State:
for msg in self.messages:
if msg.id == message_id:
return msg
raise LookupError
async def fetch_server_members(self, server_id: str):
data = await self.http.fetch_members(server_id)
for user in data["users"]:
self.add_user(user)
for member in data["members"]:
self.add_member(server_id, member)
raise KeyError
async def fetch_all_server_members(self):
for server_id in self.servers:
await self.fetch_server_members(server_id)
for server_id in self.servers.keys():
data = await self.http.fetch_members(server_id)
for user in data["users"]:
self.add_user(user)
for member in data["members"]:
self.add_member(server_id, member)
+1 -1
View File
@@ -7,7 +7,7 @@ from .http import *
from .invite import *
from .member import *
from .message import *
from .permissions import Overwrite
from .role import *
from .server import *
from .user import *
from .voice import *
+32 -32
View File
@@ -1,69 +1,69 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal, Text, TypedDict, Union
from typing_extensions import NotRequired
from typing import TYPE_CHECKING, Literal, TypedDict, Union
if TYPE_CHECKING:
from .file import File
from .message import Message
from .permissions import Overwrite
__all__ = (
"SavedMessages",
"DMChannel",
"GroupDMChannel",
"Group",
"TextChannel",
"VoiceChannel",
"GuildChannel",
"Channel",
)
class BaseChannel(TypedDict):
_id: str
class _NonceChannel(TypedDict, total=False):
nonce: str
class SavedMessages(BaseChannel):
user: str
channel_type: Literal["SavedMessages"]
class BaseChannel(TypedDict):
_id: str
class DMChannel(BaseChannel):
class SavedMessages(_NonceChannel, BaseChannel):
user: str
channel_type: Literal["SavedMessage"]
class DMChannel(_NonceChannel, BaseChannel):
active: bool
recipients: list[str]
last_message_id: NotRequired[str]
last_message: Message
channel_type: Literal["DirectMessage"]
class GroupDMChannel(BaseChannel):
class _GroupOptional(TypedDict):
icon: File
permissions: int
description: str
class Group(_NonceChannel, _GroupOptional, BaseChannel):
recipients: list[str]
name: str
owner: str
channel_type: Literal["Group"]
icon: NotRequired[File]
permissions: NotRequired[int]
description: NotRequired[str]
nsfw: NotRequired[bool]
last_message_id: NotRequired[str]
class TextChannel(BaseChannel):
class _TextChannelOptional(TypedDict, total=False):
icon: File
default_permissions: int
role_permissions: dict[str, int]
class TextChannel(_NonceChannel, _TextChannelOptional, BaseChannel):
server: str
name: str
description: str
last_message: str
channel_type: Literal["TextChannel"]
icon: NotRequired[File]
default_permissions: NotRequired[Overwrite]
role_permissions: NotRequired[dict[str, Overwrite]]
nsfw: NotRequired[bool]
last_message_id: NotRequired[str]
class VoiceChannel(BaseChannel):
class _VoiceChannelOptional(TypedDict, total=False):
icon: File
default_permissions: int
role_permissions: int
class VoiceChannel(_NonceChannel, _TextChannelOptional, BaseChannel):
server: str
name: str
description: str
channel_type: Literal["VoiceChannel"]
icon: NotRequired[File]
default_permissions: NotRequired[Overwrite]
role_permissions: NotRequired[dict[str, Overwrite]]
nsfw: NotRequired[bool]
GuildChannel = Union[TextChannel, VoiceChannel]
Channel = Union[SavedMessages, DMChannel, GroupDMChannel, TextChannel, VoiceChannel]
Channel = Union[SavedMessages, DMChannel, Group, TextChannel, VoiceChannel]
+4 -81
View File
@@ -1,85 +1,8 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict, Union
from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .file import File
__all__ = ("Embed", "SendableEmbed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "YoutubeSpecial", "TwitchSpecial", "SpotifySpecial", "SoundcloudSpecial", "BandcampSpecial", "WebsiteSpecial", "JanuaryImage", "JanuaryVideo")
class YoutubeSpecial(TypedDict):
type: Literal["Youtube"]
id: str
timestamp: NotRequired[str]
class TwitchSpecial(TypedDict):
type: Literal["Twitch"]
content_type: Literal["Channel", "Video", "Clip"]
id: str
class SpotifySpecial(TypedDict):
type: Literal["Spotify"]
content_type: str
id: str
class SoundcloudSpecial(TypedDict):
type: Literal["Soundcloud"]
class BandcampSpecial(TypedDict):
type: Literal["Bandcamp"]
content_type: Literal["Album", "Track"]
id: str
WebsiteSpecial = Union[YoutubeSpecial, TwitchSpecial, SpotifySpecial, SoundcloudSpecial, BandcampSpecial]
class JanuaryImage(TypedDict):
url: str
width: int
height: int
size: Literal["Large", "Preview"]
class JanuaryVideo(TypedDict):
url: str
width: int
height: int
class WebsiteEmbed(TypedDict):
type: Literal["Website"]
url: NotRequired[str]
special: NotRequired[WebsiteSpecial]
title: NotRequired[str]
description: NotRequired[str]
image: NotRequired[JanuaryImage]
video: NotRequired[JanuaryVideo]
site_name: NotRequired[str]
icon_url: NotRequired[str]
colour: NotRequired[str]
class ImageEmbed(JanuaryImage):
type: Literal["Image"]
class TextEmbed(TypedDict):
type: Literal["Text"]
icon_url: NotRequired[str]
url: NotRequired[str]
title: NotRequired[str]
description: NotRequired[str]
media: NotRequired[File]
colour: NotRequired[str]
class NoneEmbed(TypedDict):
type: Literal["None"]
Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed]
class SendableEmbed(TypedDict):
type: Literal["Text"]
icon_url: NotRequired[str]
url: NotRequired[str]
title: NotRequired[str]
description: NotRequired[str]
media: NotRequired[str]
colour: NotRequired[str]
__all__ = ("Embed",)
class Embed(TypedDict):
pass # TODO
+11 -62
View File
@@ -2,19 +2,14 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict, Union
from revolt.types.permissions import Overwrite
from .channel import (Channel, DMChannel, GroupDMChannel, SavedMessages,
TextChannel, VoiceChannel)
from .file import File
from .channel import (Channel, DMChannel, Group, SavedMessages, TextChannel,
VoiceChannel)
from .message import Message
from .user import Status
if TYPE_CHECKING:
from .category import Category
from .member import Member, MemberID
from .server import Server, SystemMessagesConfig
from .user import User
from .server import Server
from .user import Status, User
__all__ = (
"BasePayload",
@@ -38,8 +33,7 @@ __all__ = (
"ServerRoleUpdateEventPayload",
"ServerRoleDeleteEventPayload",
"UserUpdateEventPayload",
"UserRelationshipEventPayload",
"ServerCreateEventPayload"
"UserRelationshipEventPayload"
)
class BasePayload(TypedDict):
@@ -75,7 +69,7 @@ class MessageDeleteEventPayload(BasePayload):
class ChannelCreateEventPayload_SavedMessages(BasePayload, SavedMessages):
pass
class ChannelCreateEventPayload_Group(BasePayload, GroupDMChannel):
class ChannelCreateEventPayload_Group(BasePayload, Group):
pass
class ChannelCreateEventPayload_TextChannel(BasePayload, TextChannel):
@@ -89,18 +83,9 @@ class ChannelCreateEventPayload_DMChannel(BasePayload, DMChannel):
ChannelCreateEventPayload = Union[ChannelCreateEventPayload_Group, ChannelCreateEventPayload_Group, ChannelCreateEventPayload_TextChannel, ChannelCreateEventPayload_VoiceChannel, ChannelCreateEventPayload_DMChannel]
class ChannelUpdateEventPayloadData(TypedDict, total=False):
name: str
description: str
icon: File
nsfw: bool
active: bool
role_permissions: dict[str, Overwrite]
default_permissions: Overwrite
class ChannelUpdateEventPayload(BasePayload):
id: str
data: ChannelUpdateEventPayloadData
data: ...
clear: Literal["Icon", "Description"]
class ChannelDeleteEventPayload(BasePayload):
@@ -112,38 +97,17 @@ class ChannelStartTypingEventPayload(BasePayload):
ChannelDeleteTypingEventPayload = ChannelStartTypingEventPayload
class ServerUpdateEventPayloadData(TypedDict, total=False):
owner: str
name: str
description: str
icon: File
banner: File
default_permissions: int
nsfw: bool
system_messages: SystemMessagesConfig
categories: list[Category]
class ServerUpdateEventPayload(BasePayload):
id: str
data: ServerUpdateEventPayloadData
data: dict
clear: Literal["Icon", "Banner", "Description"]
class ServerDeleteEventPayload(BasePayload):
id: str
class ServerCreateEventPayload(BasePayload):
id: str
server: Server
channels: list[Channel]
class ServerMemberUpdateEventPayloadData(TypedDict, total=False):
nickname: str
avatar: File
roles: list[str]
class ServerMemberUpdateEventPayload(BasePayload):
id: MemberID
data: ServerMemberUpdateEventPayloadData
data: dict
clear: Literal["Nickname", "Avatar"]
class ServerMemberJoinEventPayload(BasePayload):
@@ -152,34 +116,19 @@ class ServerMemberJoinEventPayload(BasePayload):
ServerMemberLeaveEventPayload = ServerMemberJoinEventPayload
class ServerRoleUpdateEventPayloadData(TypedDict, total=False):
name: str
colour: str
hoist: bool
rank: int
class ServerRoleUpdateEventPayload(BasePayload):
id: str
role_id: str
data: ServerRoleUpdateEventPayloadData
data: dict
clear: Literal["Color"]
class ServerRoleDeleteEventPayload(BasePayload):
id: str
role_id: str
UserUpdateEventPayloadData = TypedDict("UserUpdateEventPayloadData", {
"status": Status,
"profile.background": File,
"profile.content": str,
"avatar": File,
"online": bool
}, total=False)
class UserUpdateEventPayload(BasePayload):
id: str
data: UserUpdateEventPayloadData
data: dict
clear: Literal["ProfileContent", "ProfileBackground", "StatusText", "Avatar"]
class UserRelationshipEventPayload(BasePayload):
+4
View File
@@ -14,6 +14,7 @@ __all__ = (
"Autumn",
"GetServerMembers",
"MessageWithUserData",
"JoinCallResponse"
)
@@ -50,3 +51,6 @@ class MessageWithUserData(TypedDict):
messages: list[Message]
members: list[Member]
users: list[User]
class JoinCallResponse(TypedDict):
token: str
+15 -23
View File
@@ -1,30 +1,22 @@
from __future__ import annotations
from typing import Literal, TypedDict, Union
from typing import TYPE_CHECKING, Literal, TypedDict
__all__ = (
"GroupInvite",
"ServerInvite",
"Invite",
)
from typing_extensions import NotRequired
class GroupInvite(TypedDict):
type: Literal["Group"]
_id: str
creator: str
channel: str
if TYPE_CHECKING:
from .file import File
__all__ = ("Invite", "PartialInvite")
class Invite(TypedDict):
class ServerInvite(TypedDict):
type: Literal["Server"]
server_id: str
server_name: str
server_icon: NotRequired[str]
server_banner: NotRequired[str]
channel_id: str
channel_name: str
channel_description: NotRequired[str]
user_name: str
user_avatar: NotRequired[File]
member_count: int
class PartialInvite(TypedDict):
_id: str
server: str
channel: str
creator: str
channel: str
Invite = Union[ServerInvite, GroupInvite]
+6 -6
View File
@@ -2,20 +2,20 @@ from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .file import File
__all__ = ("Member",)
class _MemberOptional(TypedDict, total=False):
nickname: str
avatar: File
roles: list[str]
class MemberID(TypedDict):
server: str
user: str
class Member(TypedDict):
class Member(_MemberOptional):
_id: MemberID
nickname: NotRequired[str]
avatar: NotRequired[File]
roles: NotRequired[list[str]]
+9 -10
View File
@@ -2,8 +2,6 @@ from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict, Union
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .embed import Embed
from .file import File
@@ -51,19 +49,20 @@ MessageEdited = TypedDict("MessageEdited", {"$date": str})
class Masquerade(TypedDict, total=False):
name: str
avatar: str
colour: str
class Message(TypedDict):
class _OptionalMessage(TypedDict):
attachments: list[File]
embeds: list[Embed]
mentions: list[str]
replies: list[str]
edited: MessageEdited
masquerade: Masquerade
class Message(_OptionalMessage):
_id: str
channel: str
author: str
content: Union[str, UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent]
attachments: NotRequired[list[File]]
embeds: NotRequired[list[Embed]]
mentions: NotRequired[list[str]]
replies: NotRequired[list[str]]
edited: NotRequired[MessageEdited]
masquerade: NotRequired[Masquerade]
class MessageReplyPayload(TypedDict):
id: str
-7
View File
@@ -1,7 +0,0 @@
from __future__ import annotations
from typing import TypedDict
class Overwrite(TypedDict):
a: int
d: int
+11 -10
View File
@@ -1,18 +1,19 @@
from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .permissions import Overwrite
from typing import TypedDict
__all__ = (
"Permission",
"Role",
)
class Role(TypedDict):
name: str
permissions: Overwrite
colour: NotRequired[str]
hoist: NotRequired[bool]
Permission = tuple[int, int]
class _RoleOptional(TypedDict, total=False):
colour: str
hoist: bool
rank: int
class Role(_RoleOptional):
name: str
permissions: Permission
+20 -21
View File
@@ -2,13 +2,11 @@ from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .category import Category
from .channel import Channel
from .file import File
from .role import Role
from .role import Permission, Role
__all__ = (
"Server",
@@ -24,34 +22,35 @@ class SystemMessagesConfig(TypedDict, total=False):
user_kicked: str
user_banned: str
class _ServerOptional(TypedDict, total=False):
nonce: str
description: str
categories: list[Category]
system_messages: SystemMessagesConfig
roles: list[Role]
icon: File
banner: File
nsfw: bool
class Server(TypedDict):
class Server(_ServerOptional):
_id: str
owner: str
name: str
channels: list[str]
default_permissions: int
nonce: NotRequired[str]
description: NotRequired[str]
categories: NotRequired[list[Category]]
system_messages: NotRequired[SystemMessagesConfig]
roles: NotRequired[dict[str, Role]]
icon: NotRequired[File]
banner: NotRequired[File]
nsfw: NotRequired[bool]
default_permissions: Permission
class BannedUser(TypedDict):
class _OptionalBannedUser(TypedDict, total=False):
avatar: File
class BannedUser(_OptionalBannedUser):
_id: str
username: str
avatar: NotRequired[File]
class BanId(TypedDict):
server: str
user: str
class _OptionalBan(TypedDict, total=False):
reason: str
class Ban(TypedDict):
_id: BanId
reason: NotRequired[str]
class Ban(_OptionalBan):
_id: str
class ServerBans(TypedDict):
users: list[BannedUser]
+11 -11
View File
@@ -2,8 +2,6 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING:
from .file import File
@@ -29,17 +27,19 @@ class UserRelation(TypedDict):
status: Relation
_id: str
class User(TypedDict):
class _OptionalUser(TypedDict, total=False):
avatar: File
relations: list[UserRelation]
badges: int
status: Status
relationship: Relation
online: bool
flags: int
bot: UserBot
class User(_OptionalUser):
_id: str
username: str
avatar: NotRequired[File]
relations: NotRequired[list[UserRelation]]
badges: NotRequired[int]
status: NotRequired[Status]
relationship: NotRequired[Relation]
online: NotRequired[bool]
flags: NotRequired[int]
bot: NotRequired[UserBot]
class UserProfile(TypedDict, total=False):
content: str
+192
View File
@@ -0,0 +1,192 @@
from typing import Optional, TypedDict, Union, Literal
from typing_extensions import NotRequired
TransportProtocol = Literal["TCP", "UDP"]
class BaseWS(TypedDict):
id: NotRequired[int]
class AuthenticationCommandData(TypedDict):
token: str
roomId: str
class AuthenticateCommand(BaseWS):
type: Literal["Authenticate"]
data: AuthenticationCommandData
class RTPCodecFeedback(TypedDict):
parameter: str
type: str
class RTPCodec(TypedDict):
channels: int
clockRate: int
kind: str
mimeType: str
parameters: dict
preferredPayloadType: int
rtcpFeedback: list[RTPCodecFeedback]
class RTPHeaderExtension(TypedDict):
direction: str
kind: str
preferredEncrypt: bool
preferredId: int
uri: str
class RTPCapabilities(TypedDict):
codecs: list[RTPCodec]
headerExtensions: list[RTPHeaderExtension]
class AuthenticationResponseData(TypedDict):
roomId: str
userId: str
version: str
rtpCapabilities: RTPCapabilities
class AuthenticateResponse(BaseWS):
type: Literal["Authenticate"]
data: AuthenticationResponseData
class RoomInfoCommand(BaseWS):
type: Literal["RoomInfo"]
class VoiceUser(TypedDict):
audio: bool
class RoomInfoResponseData(TypedDict):
id: str
users: dict[str, VoiceUser]
videoAllowed: bool
class RoomInfoResponse(BaseWS):
type: Literal["RoomInfo"]
data: RoomInfoResponseData
class InitializeTransportCommandData(TypedDict):
mode: Literal["SplitWebRTC", "CombinedWebRTC", "CombinedRTP"]
rtpCapabilities: RTPCapabilities
class InitializeTransportCommand(BaseWS):
type: Literal["InitializeTransports"]
data: InitializeTransportCommandData
class IceParameters(TypedDict):
usernameFragment: str
password: str
iceLite: NotRequired[bool]
class IceCandidate(TypedDict):
foundation: str
priority: str
ip: str
protocol: TransportProtocol
tcp_type: NotRequired[Literal["passive"]]
class DTLSFingerprints(TypedDict):
algorithm: Literal["Sha1", "Sha224", "Sha256", "Sha384", "Sha512"]
value: str
class DtlsParameters(TypedDict):
role: Literal["auto", "client", "server"]
fingerprints: list[DTLSFingerprints]
class SctpParameters(TypedDict):
port: int
OS: int
MIS: int
max_message_size: int
class TransportInfo(TypedDict):
id: str
iceParameters: IceParameters
iceCandidates: list[IceCandidate]
dtlsParameters: DtlsParameters
sctpParameters: Optional[SctpParameters]
class InitializeTransportResponseDataSplitWebRTC(TypedDict):
recvTransport: TransportInfo
sendTransport: TransportInfo
class InitializeTransportResponseDataCombinedWebRTC(TypedDict):
transport: TransportInfo
class InitializeTransportResponseDataCombinedRtp(TypedDict):
ip: bytes
port: int
protocol: TransportProtocol
id: str
srtpCryptoSuite: Literal["AES_CM_128_HMAC_SHA1_80", "AES_CM_128_HMAC_SHA1_32"]
InitializeTransportResponseData = Union[InitializeTransportResponseDataSplitWebRTC, InitializeTransportResponseDataCombinedWebRTC, InitializeTransportResponseDataCombinedRtp]
class InitializeTransportResponse(BaseWS):
type: Literal["InitializeTransports"]
data: InitializeTransportResponseData
class SrtpParameters(TypedDict):
cryptoSuite: Literal["AES_CM_128_HMAC_SHA1_80", "AES_CM_128_HMAC_SHA1_32"]
keyBase64: str
class ConnectTransportData(TypedDict):
id: str
dtlsParameters: NotRequired[DtlsParameters]
srtpParameters: NotRequired[SrtpParameters]
class ConnectTransportCommand(BaseWS):
type: Literal["ConnectTransport"]
data: ConnectTransportData
class ConnectTransportResponse(BaseWS):
type: Literal["ConnectTransport"]
class RTPEncodingRtx(TypedDict):
ssrc: int
class RTPEncoding(TypedDict, total=False):
ssrc: int
rid: str
codec_payload_type: int
rtx: RTPEncodingRtx
dtx: bool
scalability_mode: str
scale_resolution_down_by: float
max_bitrate: int
class RTPHeaderExtensionParameters(TypedDict):
uri: str
id: int
encrypt: bool
class RtcpParameters(TypedDict):
cname: NotRequired[str]
reduceSize: bool
mux: NotRequired[bool]
class RTPParameters(TypedDict):
mid: NotRequired[str]
codecs: list[RTPCodec]
headerExtensions: list[RTPHeaderExtensionParameters]
encodings: list[RTPEncoding]
rtcp: RtcpParameters
class StartProduceCommandData(TypedDict):
type: Literal["audio", "video", "saudio", "svideo"]
rtpParameters: RTPParameters
class StartProduceCommand(BaseWS):
type: Literal["StartProduce"]
data: StartProduceCommandData
class StartProduceResponseData(TypedDict):
producerId: str
class StartProduceResponse(BaseWS):
type: Literal["StartProduce"]
data: StartProduceResponseData
WSCommand = Union[AuthenticateCommand, RoomInfoCommand, InitializeTransportCommand, ConnectTransportCommand, StartProduceCommand]
WSResponse = Union[AuthenticateResponse, RoomInfoResponse, InitializeTransportResponse, ConnectTransportResponse, StartProduceResponse]
__all__ = ("BaseWS", "WSCommand", "WSResponse", "AuthenticateCommand", "RoomInfoCommand", "InitializeTransportCommand", "ConnectTransportCommand", "StartProduceCommand", "AuthenticateResponse", "RoomInfoResponse", "InitializeTransportResponse", "ConnectTransportResponse", "StartProduceResponse")
+8 -62
View File
@@ -1,21 +1,19 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union
from typing import TYPE_CHECKING, NamedTuple, Optional, Union
from .asset import Asset, PartialAsset
from .channel import DMChannel
from .enums import PresenceType, RelationshipType
from .flags import UserBadges
from .messageable import Messageable
if TYPE_CHECKING:
from .state import State
from .types import File
from .types import Status as StatusPayload
from .types import User as UserPayload
from .member import Member
__all__ = ("User", "Status", "Relation", "UserProfile")
__all__ = ("User",)
class Relation(NamedTuple):
"""A namedtuple representing a relation between the bot and a user"""
@@ -32,9 +30,9 @@ class UserProfile(NamedTuple):
content: Optional[str]
background: Optional[Asset]
class User(Messageable):
class User:
"""Represents a user
Attributes
-----------
id: :class:`str`
@@ -55,18 +53,14 @@ class User(Messageable):
The relationship between the user and the bot
status: Optional[:class:`Status`]
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
"""
__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__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile")
__slots__ = (*__flattern_attributes__, "state")
def __init__(self, data: UserPayload, state: State):
self.state = state
self._members: list[Member] = [] # 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
bot = data.get("bot")
if bot:
@@ -83,7 +77,7 @@ class User(Messageable):
avatar = data.get("avatar")
self.original_avatar = Asset(avatar, state) if avatar else None
relations: list[Relation] = []
relations = []
for relation in data.get("relations", []):
user = state.get_user(relation["_id"])
@@ -106,13 +100,6 @@ class User(Messageable):
self.masquerade_avatar: Optional[PartialAsset] = None
self.masquerade_name: Optional[str] = None
async def _get_channel_id(self):
if not self.dm_channel:
payload = await self.state.http.open_dm(self.id)
self.dm_channel = DMChannel(payload, self.state)
return self.id
@property
def owner(self) -> Optional[User]:
owner_id = self.owner_id
@@ -132,11 +119,6 @@ class User(Messageable):
"""Optional[:class:`Asset`] The avatar the member is displaying, this includes there orginal avatar and masqueraded avatar"""
return self.masquerade_avatar or self.original_avatar
@property
def mention(self) -> str:
""":class:`str`: Returns a string that allows you to mention the given user."""
return f"<@{self.id}>"
def _update(self, *, status: Optional[StatusPayload] = None, profile_content: Optional[str] = None, profile_background: Optional[File] = None, avatar: Optional[File] = None, online: Optional[bool] = None):
if status:
presence = status.get("presence")
@@ -153,39 +135,3 @@ class User(Messageable):
if online:
self.online = online
# update user infomation for all members
for member in self._members:
User._update(member, status=status, profile_content=profile_content, profile_background=profile_background, avatar=avatar, online=online)
async def default_avatar(self) -> bytes:
"""Returns the default avatar for this user
Returns
--------
:class:`bytes`
The bytes of the image
"""
return await self.state.http.fetch_default_avatar(self.id)
async def fetch_profile(self) -> UserProfile:
"""Fetches the user's profile
Returns
--------
:class:`UserProfile`
The user's profile
"""
if profile := self.profile:
return profile
payload = await self.state.http.fetch_profile(self.id)
if file := payload.get("background"):
background = Asset(file, self.state)
else:
background = None
self.profile = UserProfile(payload.get("content"), background)
return self.profile
+6 -75
View File
@@ -1,12 +1,9 @@
import inspect
from contextlib import asynccontextmanager
from operator import attrgetter
from typing import Any, Callable, Coroutine, Iterable, TypeVar, Union
from typing import Any, Callable, Coroutine, TypeVar, Union
from aiohttp import ClientSession
from typing_extensions import ParamSpec
__all__ = ("Missing", "copy_doc", "maybe_coroutine", "get", "client_session")
__all__ = ("Missing", "copy_doc", "maybe_coroutine")
class _Missing:
def __repr__(self):
@@ -20,15 +17,15 @@ def copy_doc(from_t: T) -> Callable[[T], T]:
def inner(to_t: T) -> T:
to_t.__doc__ = from_t.__doc__
return to_t
return inner
R_T = TypeVar("R_T")
P = ParamSpec("P")
# it is impossible to type this function correctly for a couple reasons:
# 1. isawaitable does not narrow while keeping typevars - there is an open PR for this (typeshed#5658) but it cannot be merged because mypy does not support the feature fully
# 2. typeguard does not narrow for the negative case which is dumb in my opinion, so `value` would stay being a union even after the if statement (PEP 647 - "The type is not narrowed in the negative case")
# its impossible to type this function correctly for a couple reasons:
# 1. isawaitable doesnt narrow while keeping typevars - there is an open pr for this (typeshed#5658) but it cant be merged because mypy doesnt support the feature fully
# 2. typeguard doesnt narrow for the negative case which is dumb imo, so `value` would stay being a union even after the if statement (PEP 647 - "The type is not narrowed in the negative case")
async def maybe_coroutine(func: Callable[P, Union[R_T, Coroutine[Any, Any, R_T]]], *args: P.args, **kwargs: P.kwargs) -> R_T:
value = func(*args, **kwargs)
@@ -37,69 +34,3 @@ async def maybe_coroutine(func: Callable[P, Union[R_T, Coroutine[Any, Any, R_T]]
value = await value
return value # type: ignore
def get(iterable: Iterable[T], **attrs: Any) -> T:
"""A convenience function to help get a value from an iterable with a specific attribute
Examples
---------
.. code-block:: python
:emphasize-lines: 3
from revolt import utils
channel = utils.get(server.channels, name="General")
await channel.send("Hello general chat.")
Parameters
-----------
iterable: Iterable
The values to search though
**attrs: Any
The attributes to check
Returns
--------
Any
The value from the iterable with the met attributes
Raises
-------
LookupError
Raises when none of the values in the iterable matches the attributes
"""
converted = [(attrgetter(attr.replace('__', '.')), value) for attr, value in attrs.items()]
for elem in iterable:
if all(pred(elem) == value for pred, value in converted):
return elem
raise LookupError
@asynccontextmanager
async def client_session():
"""A context manager that creates a new aiohttp.ClientSession() and closes it when exiting the context.
Examples
---------
.. code-block:: python
:emphasize-lines: 3
async def main():
async with client_session() as session:
client = revolt.Client(session, "TOKEN")
await client.start()
asyncio.run(main())
"""
session = ClientSession()
try:
yield session
finally:
await session.close()
+204
View File
@@ -0,0 +1,204 @@
from __future__ import annotations
import json
import logging
import asyncio
import aiohttp
import random
import string
import nacl
import nacl.public
import nacl.encoding
import nacl.utils
import base64
from typing import TYPE_CHECKING, Literal, BinaryIO
from revolt.types.voice import AuthenticationResponseData, InitializeTransportResponseDataCombinedRtp, RTPCapabilities, RTPHeaderExtensionParameters, StartProduceResponseData
if TYPE_CHECKING:
from .state import State
from .types.voice import *
__all__ = ("VoiceClient",)
logger = logging.getLogger("revolt.voice_client")
class VoiceClient:
__slots__ = ("session", "token", "websocket", "dispatch", "state", "ws_url", "channel_id", "loop", "id", "futures", "capabilities", "rtp_info", "ssrc", "cname", "producer_id", "key", "salt")
def __init__(self, session: aiohttp.ClientSession, ws_url: str, channel_id: str, token: str, state: State):
self.id = 0
self.session = session
self.ws_url = ws_url
self.channel_id = channel_id
self.websocket: aiohttp.ClientWebSocketResponse
self.token = token
self.dispatch = state.dispatch
self.state = state
self.loop = asyncio.get_running_loop()
self.futures = dict[int, asyncio.Future]()
self.ssrc = random.randint(0, 1 << 31)
self.cname = "".join(random.choices(string.ascii_letters, k=15))
self.capabilities: RTPCapabilities
self.rtp_info: InitializeTransportResponseDataCombinedRtp
self.producer_id: str
self.key = nacl.utils.random(16)
self.salt = nacl.utils.random(14)
async def send_payload(self, payload) -> int:
self.id += 1
payload["id"] = self.id
logging.info(payload)
await self.websocket.send_str(json.dumps(payload))
return self.id
async def heartbeat(self):
while not self.websocket.closed:
# logger.info("Sending voice hearbeat")
await self.websocket.ping()
await asyncio.sleep(15)
async def send_authenticate(self):
payload: AuthenticateCommand = {
"type": "Authenticate",
"data": {
"token": self.token,
"roomId": self.channel_id
}
}
id = await self.send_payload(payload)
response: AuthenticationResponseData = await self.set_future(id)
self.capabilities = response["rtpCapabilities"]
async def send_initialize_transport(self):
payload: InitializeTransportCommand = {
"type": "InitializeTransports",
"data": {
"mode": "CombinedRTP",
"rtpCapabilities": self.capabilities
}
}
id = await self.send_payload(payload)
response: InitializeTransportResponseDataCombinedRtp = await self.set_future(id)
self.rtp_info = response
async def send_connect_transport(self):
payload: ConnectTransportCommand = {
"type": "ConnectTransport",
"data": {
"srtpParameters": {
"cryptoSuite": self.rtp_info["srtpCryptoSuite"],
"keyBase64": base64.b64encode(self.salt + self.key).decode(),
},
"id": self.rtp_info["id"],
}
}
await self.send_payload(payload)
async def send_start_produce(self, type: Literal["audio", "video", "saudio", "svideo"]):
header_extensions: list[RTPHeaderExtensionParameters] = [{"uri": d["uri"], "id": d["preferredId"], "encrypt": d["preferredEncrypt"]} for d in self.capabilities["headerExtensions"]]
payload: StartProduceCommand = {
"type": "StartProduce",
"data": {
"type": type,
"rtpParameters": {
"codecs": self.capabilities["codecs"],
"headerExtensions": header_extensions,
"encodings": [{
"ssrc": self.ssrc,
}],
"mid": "0",
"rtcp": {
"cname": self.cname,
"reduceSize": True
}
}
}
}
id = await self.send_payload(payload)
response: StartProduceResponseData = await self.set_future(id)
producer_id = response["producerId"]
self.producer_id = producer_id
async def handle_event(self, payload: WSResponse):
event_type = payload["type"].lower()
logger.debug("Recieved event %s %s", event_type, payload)
try:
func = getattr(self, f"handle_{event_type}")
except:
logger.debug("Unknown event '%s'", event_type)
return
await func(payload)
async def handle_authenticate(self, _):
logger.info("Successfully authenticated")
async def set_future(self, id: int):
future = asyncio.Future()
self.futures[id] = future
return await future
async def websocket_loop(self):
async for msg in self.websocket:
payload: WSResponse = json.loads(msg.data)
print(payload)
if future := self.futures.get(payload.get("id", "")):
future.set_result(payload.get("data"))
self.loop.create_task(self.handle_event(payload))
async def start(self):
self.websocket = await self.session.ws_connect(self.ws_url)
asyncio.create_task(self.heartbeat())
asyncio.create_task(self.websocket_loop())
async def send_audio(self, audio: BinaryIO):
await self.send_start_produce("audio")
"""
{
'type': 'StartProduce',
'data': {
'type': 'audio',
'rtpParameters': {
'codecs': [
{'kind': 'audio', 'mimeType': 'audio/opus', 'preferredPayloadType': 100, 'clockRate': 48000, 'channels': 2, 'parameters': {}, 'rtcpFeedback': [{'type': 'transport-cc', 'parameter': ''}]}
],
'headerExtensions': [{'uri': 'urn:ietf:params:rtp-hdrext:sdes:mid', 'id': 1, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:sdes:mid', 'id': 1, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:sdes:rtp-stream-id', 'id': 2, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:sdes:repaired-rtp-stream-id', 'id': 3, 'encrypt': False}, {'uri': 'http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time', 'id': 4, 'encrypt': False}, {'uri': 'http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time', 'id': 4, 'encrypt': False}, {'uri': 'http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01', 'id': 5, 'encrypt': False}, {'uri': 'http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01', 'id': 5, 'encrypt': False}, {'uri': 'http://tools.ietf.org/html/draft-ietf-avtext-framemarking-07', 'id': 6, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:framemarking', 'id': 7, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:ssrc-audio-level', 'id': 10, 'encrypt': False}, {'uri': 'urn:3gpp:video-orientation', 'id': 11, 'encrypt': False}, {'uri': 'urn:ietf:params:rtp-hdrext:toffset', 'id': 12, 'encrypt': False}],
'encodings': [{'ssrc': 896646738}],
'mid': '0',
'rtcp': {'cname': 'iUaoUqyQjShruzl', 'reduceSize': True}}
},
'id': 4}
"""
"""
{
"type": "StartProduce",
"data": {
"type":"audio",
"rtpParameters": {
"codecs": [
{"mimeType":"audio/opus","payloadType":111,"clockRate":48000,"channels":2,"parameters":{"minptime":10,"useinbandfec":1},"rtcpFeedback":[{"type":"transport-cc","parameter":""}]}
],
"headerExtensions":[{"uri":"urn:ietf:params:rtp-hdrext:sdes:mid","id":4,"encrypt":false,"parameters":{}},{"uri":"http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time","id":2,"encrypt":false,"parameters":{}},{"uri":"http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01","id":3,"encrypt":false,"parameters":{}},{"uri":"urn:ietf:params:rtp-hdrext:ssrc-audio-level","id":1,"encrypt":false,"parameters":{}}],
"encodings":[{"ssrc":2276570677,"dtx":false}],
"mid":"0"}
"rtcp":{"cname":"KAhFB58ELM5Tf+9a","reducedSize":true},
}
}
"id":7,
"""
+22 -96
View File
@@ -3,10 +3,8 @@ from __future__ import annotations
import asyncio
import logging
from copy import copy
from traceback import print_exception
from typing import TYPE_CHECKING, Callable, cast
from .channel import GroupDMChannel, TextChannel, VoiceChannel
from .enums import RelationshipType
from .types import (ChannelCreateEventPayload, ChannelDeleteEventPayload,
ChannelDeleteTypingEventPayload,
@@ -15,12 +13,11 @@ from .types import Message as MessagePayload
from .types import (MessageDeleteEventPayload, MessageUpdateEventPayload,
ServerDeleteEventPayload, ServerMemberJoinEventPayload,
ServerMemberLeaveEventPayload,
ServerCreateEventPayload,
ServerMemberUpdateEventPayload,
ServerRoleDeleteEventPayload, ServerRoleUpdateEventPayload,
ServerUpdateEventPayload, UserRelationshipEventPayload,
UserUpdateEventPayload)
from .user import Status, UserProfile
from .user import Status
try:
import ujson as json
@@ -47,7 +44,7 @@ __all__ = ("WebsocketHandler",)
logger = logging.getLogger("revolt")
class WebsocketHandler:
__slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready", "server_events")
__slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready")
def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State):
self.session = session
@@ -59,11 +56,6 @@ class WebsocketHandler:
self.loop = asyncio.get_running_loop()
self.user = None
self.ready = asyncio.Event()
self.server_events: dict[str, asyncio.Event] = {}
async def _wait_for_server_ready(self, server_id: str):
if event := self.server_events.get(server_id):
await event.wait()
async def send_payload(self, payload: BasePayload):
if use_msgpack:
@@ -93,7 +85,7 @@ class WebsocketHandler:
await self.ready.wait()
func = getattr(self, f"handle_{event_type}")
except AttributeError:
except:
logger.debug("Unknown event '%s'", event_type)
return
@@ -126,10 +118,6 @@ class WebsocketHandler:
async def handle_message(self, payload: MessageEventPayload):
message = self.state.add_message(cast(MessagePayload, payload))
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message", message)
async def handle_messageupdate(self, payload: MessageUpdateEventPayload):
@@ -140,19 +128,14 @@ class WebsocketHandler:
data = payload["data"]
kwargs = {}
if content := data.get("content"):
kwargs["content"] = content
if data["content"]:
kwargs["content"] = data["content"]
kwargs["edited_at"] = data["edited"]["$date"]
if embeds := data.get("embeds"):
kwargs["embeds"] = embeds
if data["edited"]["$date"]:
kwargs["edited_at"] = data["edited"]["$date"]
message._update(**kwargs)
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message_update", message)
async def handle_messagedelete(self, payload: MessageDeleteEventPayload):
@@ -160,22 +143,15 @@ class WebsocketHandler:
try:
message = self.state.get_message(payload["id"])
except KeyError:
except:
return
self.state.messages.remove(message)
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message_delete", message)
async def handle_channelcreate(self, payload: ChannelCreateEventPayload):
channel = self.state.add_channel(payload)
if server := channel.server:
await self._wait_for_server_ready(server.id)
self.dispatch("channel_create", channel)
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload):
@@ -187,42 +163,28 @@ class WebsocketHandler:
if clear := payload.get("clear"):
if clear == "Icon":
if isinstance(channel, (TextChannel, VoiceChannel, GroupDMChannel)):
channel.icon = None
pass # TODO
elif clear == "Description":
if isinstance(channel, (TextChannel, VoiceChannel, GroupDMChannel)):
channel.description = None
if server := channel.server:
await self._wait_for_server_ready(server.id)
channel.description = None # type: ignore
self.dispatch("channel_update", old_channel, channel)
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload):
channel = self.state.channels.pop(payload["id"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
self.dispatch("channel_delete", channel)
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload):
channel = self.state.get_channel(payload["id"])
user = self.state.get_user(payload["user"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
self.dispatch("typing_start", channel, user)
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload):
channel = self.state.get_channel(payload["id"])
user = self.state.get_user(payload["user"])
if server := channel.server:
await self._wait_for_server_ready(server.id)
self.dispatch("typing_stop", channel, user)
async def handle_serverupdate(self, payload: ServerUpdateEventPayload):
@@ -242,8 +204,6 @@ class WebsocketHandler:
elif clear == "Description":
server.description = None
await self._wait_for_server_ready(server.id)
self.dispatch("server_update", old_server, server)
async def handle_serverdelete(self, payload: ServerDeleteEventPayload):
@@ -252,26 +212,9 @@ class WebsocketHandler:
for channel in server.channels:
del self.state.channels[channel.id]
await self._wait_for_server_ready(server.id)
self.dispatch("server_delete", server)
async def handle_servercreate(self, payload: ServerCreateEventPayload):
for channel in payload["channels"]:
self.state.add_channel(channel)
server = self.state.add_server(payload["server"])
# lock all server events until we fetch all the members, otherwise the cache will be incomplete
self.server_events[server.id] = asyncio.Event()
await self.state.fetch_server_members(server.id)
self.server_events.pop(server.id).set()
self.dispatch("server_create", server)
async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload):
await self._wait_for_server_ready(payload["id"]["server"])
member = self.state.get_member(payload["id"]["server"], payload["id"]["user"])
old_member = copy(member)
@@ -290,16 +233,9 @@ class WebsocketHandler:
self.dispatch("member_join", member)
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload):
await self._wait_for_server_ready(payload["id"])
server = self.state.get_server(payload["id"])
member = server._members.pop(payload["user"])
# remove the member from the user
user = self.state.get_user(payload["user"])
user._members.remove(member)
self.dispatch("member_leave", member)
async def handle_serveroleupdate(self, payload: ServerRoleUpdateEventPayload):
@@ -313,16 +249,12 @@ class WebsocketHandler:
role._update(**payload["data"])
await self._wait_for_server_ready(server.id)
self.dispatch("role_update", old_role, role)
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload):
server = self.state.get_server(payload["id"])
role = server._roles.pop(payload["role_id"])
await self._wait_for_server_ready(server.id)
self.dispatch("role_delete", role)
async def handle_userupdate(self, payload: UserUpdateEventPayload):
@@ -331,27 +263,26 @@ class WebsocketHandler:
if clear := payload.get("clear"):
if clear == "ProfileContent":
if profile := user.profile:
user.profile = UserProfile(None, profile.background)
...
elif clear == "ProfileBackground":
if profile := user.profile:
user.profile = UserProfile(profile.content, None)
...
elif clear == "StatusText":
user.status = Status(None, user.status.presence if user.status else None)
# user.status will never be none because they are trying to remove the text
if user.status.presence is None: # type: ignore
user.status = None
else:
user.status = Status(None, user.status.presence) # type: ignore
elif clear == "Avatar":
user.original_avatar = None
# the keys have . in them so I need to replace with _
# type: ignore is for it to stop complaining about the keys not existing in the typeddict
# the keys have . in them so i need to replace with _
data = payload["data"]
data["profile_content"] = data.pop("profile.content", None) # type: ignore
data["profile_background"] = data.pop("profile.background", None) # type: ignore
data["profile_content"] = data.get("profile.content", None)
data["profile_background"] = data.get("profile.background", None)
user._update(**data) # type: ignore
user._update(**data)
self.dispatch("user_update", old_user, user)
@@ -378,9 +309,4 @@ class WebsocketHandler:
else:
payload = json.loads(msg.data)
task = self.loop.create_task(self.handle_event(payload))
# task.add_done_callback(task_done)
def task_done(task: asyncio.Task[None]):
if exception := task.exception():
print_exception(type(exception), exception, exception.__traceback__)
self.loop.create_task(self.handle_event(payload))
Executable
+46
View File
@@ -0,0 +1,46 @@
import pathlib
from setuptools import find_packages, setup
here = pathlib.Path(__file__).parent.resolve()
long_description = (here / "README.md").read_text(encoding="utf-8")
requirements = []
with open('requirements.txt') as f:
requirements = f.read().splitlines()
setup(
name="revolt.py",
version="0.1.1",
description="Python wrapper around revolt.chat",
long_description=long_description,
long_description_content_type="text/markdown",
url="https://github.com/Zomatree/revolt.py",
author="Zomatree",
classifiers=[
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"License :: OSI Approved :: MIT License",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3 :: Only",
],
keywords="wrapper, async, api, websockets, http",
packages=find_packages(),
python_requires=">=3.9",
extras_require={
"speedups": ["ujson", "aiohttp[speedups]==3.7.4.post0", "msgpack==1.0.2"],
},
install_requires=requirements,
project_urls={
"Bug Reports": "https://github.com/Zomatree/revolt.py/issues",
"Source": "https://github.com/Zomatree/revolt.py/",
},
)