18 Commits

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

* handle edited being int in MessageUpdate event
2023-05-19 23:16:33 +01:00
Zomatree 52c2be4e91 Fix using .server instead of .server_id
Fixes #51
2023-05-19 23:13:55 +01:00
Zomatree d5a15c0f44 fix typing error related to bad aiohttp types 2023-05-19 22:23:47 +01:00
Zomatree 4bf0530f8e clean up code 2023-05-11 07:02:39 +01:00
Zomatree 77ff484ad8 small bug fixes 2023-04-12 20:20:18 +01:00
Zomatree d909b0eb3e Merge branch 'master' of github.com:revoltchat/revolt.py 2023-04-12 16:45:11 +01:00
Angelo Kontaxis 0e694db586 Merge pull request #49 from revoltchat/permissions
Permissions
2023-04-11 22:34:30 +01:00
40 changed files with 668 additions and 427 deletions
+26 -4
View File
@@ -1,16 +1,38 @@
on: [push, pull_request]
name: pyright
jobs:
pyright:
pyright-type-checking:
strategy:
matrix:
version: ["3.9", "3.10", "3.11"]
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- uses: actions/setup-python@v2
with:
python-version: '3.9'
python-version: ${{ matrix.version }}
- run: pip install .[speedups,docs]
- uses: jakebailey/pyright-action@v1
with:
lib: true
python-version: 3.9
python-version: ${{ matrix.version }}
working-directory: revolt
pyright-type-completeness:
strategy:
matrix:
version: ["3.9", "3.10", "3.11"]
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- uses: actions/setup-python@v2
with:
python-version: ${{ matrix.version }}
- run: pip install .[speedups,docs]
- uses: jakebailey/pyright-action@v1
with:
python-version: ${{ matrix.version }}
working-directory: revolt
verify-types: revolt
ignore-external: true
+2 -2
View File
@@ -11,10 +11,10 @@ upload:
python -m twine upload dist/*
lint:
pyright --lib
pyright .
coverage:
pyright --lib --ignoreexternal --verifytypes revolt
pyright --ignoreexternal --verifytypes revolt
docs:
cd docs && make html
+9
View File
@@ -49,11 +49,20 @@ path = "revolt/__init__.py"
[tool.hatch.build]
only-packages = true
include = ["revolt/**/*"]
[tool.pyright]
reportPrivateUsage = false
reportImportCycles = false
reportIncompatibleMethodOverride = false
typeCheckingMode = "strict"
[tool.hatch.build.targets.sdist]
strict-naming = false
[tool.hatch.build.targets.wheel]
strict-naming = false
[build-system]
requires = ["hatchling"]
+1 -1
View File
@@ -18,4 +18,4 @@ from .role import *
from .server import *
from .user import *
__version__ = "0.1.9"
__version__ = "0.1.11"
+20 -21
View File
@@ -42,14 +42,16 @@ class Asset(Ulid):
__slots__ = ("state", "id", "tag", "size", "filename", "content_type", "width", "height", "type", "url")
def __init__(self, data: FilePayload, state: State):
self.state = state
self.state: State = state
self.id = data['_id']
self.tag = data['tag']
self.size = data['size']
self.filename = data['filename']
self.id: str = data['_id']
self.tag: str = data['tag']
self.size: int = data['size']
self.filename: str = data['filename']
metadata = data['metadata']
self.height: int | None
self.width: int | None
if metadata["type"] == "Image" or metadata["type"] == "Video": # cannot use `in` because type narrowing will not happen
self.height = metadata["height"]
@@ -58,17 +60,17 @@ class Asset(Ulid):
self.height = None
self.width = None
self.content_type = data["content_type"]
self.type = AssetType(metadata["type"])
self.content_type: str | None = data["content_type"]
self.type: AssetType = AssetType(metadata["type"])
base_url = self.state.api_info["features"]["autumn"]["url"]
self.url = f"{base_url}/{self.tag}/{self.id}"
self.url: str = f"{base_url}/{self.tag}/{self.id}"
async def read(self) -> bytes:
"""Reads the files content into bytes"""
return await self.state.http.request_file(self.url)
async def save(self, fp: IOBase):
async def save(self, fp: IOBase) -> None:
"""Reads the files content and saves it to a file
Parameters
@@ -85,8 +87,6 @@ class PartialAsset(Asset):
-----------
id: :class:`str`
The id of the asset, this will always be ``"0"``
tag: Optional[:class:`str`]
The tag of the asset, this corrasponds to where the asset is used, this will always be ``None``
size: :class:`int`
Amount of bytes in the file, this will always be ``0``
filename: :class:`str`
@@ -102,13 +102,12 @@ class PartialAsset(Asset):
"""
def __init__(self, url: str, state: State):
self.state = state
self.id = "0"
self.tag = None
self.size = 0
self.filename = ""
self.height = None
self.width = None
self.content_type = mimetypes.guess_extension(url)
self.type = AssetType.file
self.url = url
self.state: State = state
self.id: str = "0"
self.size: int = 0
self.filename: str = ""
self.height: int | None = None
self.width: int | None = None
self.content_type: str | None = mimetypes.guess_extension(url)
self.type: AssetType = AssetType.file
self.url: str = url
+4 -4
View File
@@ -25,10 +25,10 @@ class Category(Ulid):
"""
def __init__(self, data: CategoryPayload, state: State):
self.state = state
self.name = data["title"]
self.id = data["id"]
self.channel_ids = data["channels"]
self.state: State = state
self.name: str = data["title"]
self.id: str = data["id"]
self.channel_ids: list[str] = data["channels"]
@property
def channels(self) -> list[Channel]:
+33 -22
View File
@@ -31,7 +31,7 @@ class EditableChannel:
state: State
id: str
async def edit(self, **kwargs: Any):
async def edit(self, **kwargs: Any) -> None:
"""Edits the channel
Passing ``None`` to the parameters that accept it will remove them.
@@ -80,24 +80,30 @@ class Channel(Ulid):
__slots__ = ("state", "id", "channel_type", "server_id")
def __init__(self, data: ChannelPayload, state: State):
self.state = state
self.id = data["_id"]
self.channel_type = ChannelType(data["channel_type"])
self.state: State = state
self.id: str = data["_id"]
self.channel_type: ChannelType = ChannelType(data["channel_type"])
self.server_id: Optional[str] = None
async def _get_channel_id(self) -> str:
return self.id
def _update(self, **_: Any):
def _update(self, **_: Any) -> None:
pass
async def delete(self):
async def delete(self) -> None:
"""Deletes or closes the channel"""
await self.state.http.close_channel(self.id)
@property
def server(self) -> Server:
""":class:`Server` The server this voice channel belongs too"""
""":class:`Server` The server this voice channel belongs too
Raises
-------
:class:`LookupError`
Raises if the channel is not part of a server
"""
if not self.server_id:
raise LookupError
@@ -128,7 +134,7 @@ class DMChannel(Channel, Messageable):
def __init__(self, data: DMChannelPayload, state: State):
super().__init__(data, state)
self.recipient_ids: tuple[str, str] = tuple(data["recipients"])
self.last_message_id = data.get("last_message_id")
self.last_message_id: str | None = data.get("last_message_id")
@property
def recipients(self) -> tuple[User, User]:
@@ -184,20 +190,22 @@ class GroupDMChannel(Channel, Messageable, EditableChannel):
def __init__(self, data: GroupDMChannelPayload, state: State):
super().__init__(data, state)
self.recipient_ids = data["recipients"]
self.name = data["name"]
self.owner_id = data["owner"]
self.description: Optional[str] = data.get("description")
self.last_message_id = data.get("last_message_id")
self.recipient_ids: list[str] = data["recipients"]
self.name: str = data["name"]
self.owner_id: str = data["owner"]
self.description: str | None = data.get("description")
self.last_message_id: str | None = data.get("last_message_id")
self.icon: Asset | None
if icon := data.get("icon"):
self.icon = Asset(icon, state)
else:
self.icon = None
self.permissions = Permissions(data.get("permissions", 0))
self.permissions: Permissions = Permissions(data.get("permissions", 0))
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None, description: Optional[str] = None):
def _update(self, *, name: Optional[str] = None, recipients: Optional[list[str]] = None, description: Optional[str] = None) -> None:
if name is not None:
self.name = name
@@ -257,12 +265,12 @@ class ServerChannel(Channel):
def __init__(self, data: ServerChannelPayload, state: State):
super().__init__(data, state)
self.server_id = data["server"]
self.name = data["name"]
self.server_id: str = data["server"]
self.name: str = data["name"]
self.description: Optional[str] = data.get("description")
self.nsfw = data.get("nsfw", False)
self.active = False
self.default_permissions = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0}))
self.nsfw: bool = data.get("nsfw", False)
self.active: bool = False
self.default_permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0}))
permissions: dict[str, PermissionsOverwrite] = {}
@@ -270,7 +278,10 @@ class ServerChannel(Channel):
overwrite = PermissionsOverwrite._from_overwrite(overwrite_data)
permissions[role_name] = overwrite
self.permissions = permissions
self.permissions: dict[str, PermissionsOverwrite] = permissions
self.icon: Asset | None
if icon := data.get("icon"):
self.icon = Asset(icon, state)
else:
@@ -353,7 +364,7 @@ class TextChannel(ServerChannel, Messageable, EditableChannel):
def __init__(self, data: TextChannelPayload, state: State):
super().__init__(data, state)
self.last_message_id = data.get("last_message_id")
self.last_message_id: str | None = data.get("last_message_id")
async def _get_channel_id(self) -> str:
return self.id
+13 -13
View File
@@ -30,7 +30,7 @@ if TYPE_CHECKING:
__all__ = ("Client",)
logger = logging.getLogger("revolt")
logger: logging.Logger = logging.getLogger("revolt")
class Client:
"""The client for interacting with revolt
@@ -48,11 +48,11 @@ class Client:
"""
def __init__(self, session: aiohttp.ClientSession, token: str, *, api_url: str = "https://api.revolt.chat", max_messages: int = 5000, bot: bool = True):
self.session = session
self.token = token
self.api_url = api_url
self.max_messages = max_messages
self.bot = bot
self.session: aiohttp.ClientSession = session
self.token: str = token
self.api_url: str = api_url
self.max_messages: int = max_messages
self.bot: bool = bot
self.api_info: ApiInfo
self.http: HttpClient
@@ -63,7 +63,7 @@ class Client:
super().__init__()
def dispatch(self, event: str, *args: Any):
def dispatch(self, event: str, *args: Any) -> None:
"""Dispatch an event, this is typically used for testing and internals.
Parameters
@@ -88,7 +88,7 @@ class Client:
async with self.session.get(self.api_url) as resp:
return json.loads(await resp.text())
async def start(self):
async def start(self) -> None:
"""Starts the client"""
api_info = await self.get_api_info()
@@ -98,7 +98,7 @@ class Client:
self.websocket = WebsocketHandler(self.session, self.token, api_info["ws"], self.dispatch, self.state)
await self.websocket.start()
async def stop(self):
async def stop(self) -> None:
await self.websocket.websocket.close()
def get_user(self, id: str) -> User:
@@ -300,7 +300,7 @@ class Client:
raise LookupError
async def edit_self(self, **kwargs: Any):
async def edit_self(self, **kwargs: Any) -> None:
"""Edits the client's own user
Parameters
@@ -316,7 +316,7 @@ class Client:
await self.state.http.edit_self(remove, kwargs)
async def edit_status(self, **kwargs: Any):
async def edit_status(self, **kwargs: Any) -> None:
"""Edits the client's own status
Parameters
@@ -337,7 +337,7 @@ class Client:
await self.state.http.edit_self(remove, {"status": kwargs})
async def edit_profile(self, **kwargs: Any):
async def edit_profile(self, **kwargs: Any) -> None:
"""Edits the client's own profile
Parameters
@@ -347,7 +347,7 @@ class Client:
background: Optional[:class:`File`]
The new background for the profile, passing in ``None`` will remove the profile background
"""
remove = []
remove: list[str] = []
if kwargs.get("content", Missing) is None:
del kwargs["content"]
+26 -21
View File
@@ -4,6 +4,8 @@ from typing import TYPE_CHECKING, Optional, TypedDict, Union
from typing_extensions import NotRequired, Unpack
from revolt.types.embed import WebsiteSpecial
from .asset import Asset
from .enums import EmbedType
@@ -14,6 +16,7 @@ if TYPE_CHECKING:
from .types import SendableEmbed as SendableEmbedPayload
from .types import TextEmbed as TextEmbedPayload
from .types import WebsiteEmbed as WebsiteEmbedPayload
from .types import JanuaryImage, JanuaryVideo
__all__ = ("Embed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "to_embed", "SendableEmbed")
@@ -21,43 +24,45 @@ class WebsiteEmbed:
type = EmbedType.website
def __init__(self, embed: WebsiteEmbedPayload):
self.url = embed.get("url")
self.special = embed.get("special")
self.title = embed.get("title")
self.description = embed.get("description")
self.image = embed.get("image")
self.video = embed.get("video")
self.site_name = embed.get("site_name")
self.icon_url = embed.get("icon_url")
self.colour = embed.get("colour")
self.url: str | None = embed.get("url")
self.special: WebsiteSpecial | None = embed.get("special")
self.title: str | None = embed.get("title")
self.description: str | None = embed.get("description")
self.image: JanuaryImage | None = embed.get("image")
self.video: JanuaryVideo | None = embed.get("video")
self.site_name: str | None = embed.get("site_name")
self.icon_url: str | None = embed.get("icon_url")
self.colour: str | None = embed.get("colour")
class ImageEmbed:
type = EmbedType.image
type: EmbedType = EmbedType.image
def __init__(self, image: ImageEmbedPayload):
self.url = image.get("url")
self.width = image.get("width")
self.height = image.get("height")
self.size = image.get("size")
self.url: str = image.get("url")
self.width: int = image.get("width")
self.height: int = image.get("height")
self.size: str = image.get("size")
class TextEmbed:
type = EmbedType.text
type: EmbedType = EmbedType.text
def __init__(self, embed: TextEmbedPayload, state: State):
self.icon_url = embed.get("icon_url")
self.url = embed.get("url")
self.title = embed.get("title")
self.description = embed.get("description")
self.icon_url: str | None = embed.get("icon_url")
self.url: str | None = embed.get("url")
self.title: str | None = embed.get("title")
self.description: str | None = embed.get("description")
self.media: Asset | None
if media := embed.get("media"):
self.media = Asset(media, state)
else:
self.media = None
self.colour = embed.get("colour")
self.colour: str | None = embed.get("colour")
class NoneEmbed:
type = EmbedType.none
type: EmbedType = EmbedType.none
Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed]
+9 -9
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
from .utils import Ulid
@@ -30,16 +30,16 @@ class Emoji(Ulid):
The server id this emoji belongs to, if any
"""
def __init__(self, payload: EmojiPayload, state: State):
self.state = state
self.state: State = state
self.id = payload["_id"]
self.author_id = payload["creator_id"]
self.name = payload["name"]
self.animated = payload.get("animated", False)
self.nsfw = payload.get("nsfw", False)
self.server_id: Optional[str] = payload["parent"].get("id")
self.id: str = payload["_id"]
self.author_id: str = payload["creator_id"]
self.name: str = payload["name"]
self.animated: bool = payload.get("animated", False)
self.nsfw: bool = payload.get("nsfw", False)
self.server_id: str | None = payload["parent"].get("id")
async def delete(self):
async def delete(self) -> None:
"""Deletes the emoji."""
await self.state.http.delete_emoji(self.id)
+1
View File
@@ -4,6 +4,7 @@ __all__ = (
"ServerError",
"FeatureDisabled",
"AutumnDisabled",
"Forbidden",
)
class RevoltError(Exception):
+20 -15
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from typing import Any, Callable, Coroutine, TypeVar, Union, cast
from typing import Any, Callable, Coroutine, Union, cast
from typing_extensions import TypeVar
import revolt
@@ -12,11 +13,11 @@ from .utils import ClientT
__all__ = ("check", "Check", "is_bot_owner", "is_server_owner", "has_permissions", "has_channel_permissions")
T = TypeVar("T", Callable[..., Any], Command)
T = TypeVar("T", Callable[..., Any], Command, default=Command)
Check = Callable[[Context[ClientT]], Union[Any, Coroutine[Any, Any, Any]]]
def check(check: Check[ClientT]):
def check(check: Check[ClientT]) -> Callable[[T], T]:
"""A decorator for adding command checks
Parameters
@@ -37,7 +38,7 @@ def check(check: Check[ClientT]):
return inner
def is_bot_owner():
def is_bot_owner() -> Callable[[T], T]:
"""A command check for limiting the command to only the bot's owner"""
@check
def inner(context: Context[ClientT]):
@@ -48,11 +49,11 @@ def is_bot_owner():
return inner
def is_server_owner():
def is_server_owner() -> Callable[[T], T]:
"""A command check for limiting the command to only a server's owner"""
@check
def inner(context: Context[ClientT]):
if not context.server:
def inner(context: Context[ClientT]) -> bool:
if not context.server_id:
raise ServerOnly
if context.author.id == context.server.owner_id:
@@ -62,25 +63,29 @@ def is_server_owner():
return inner
def has_permissions(**permissions: bool):
def has_permissions(**permissions: bool) -> Callable[[T], T]:
@check
def inner(context: Context[ClientT]):
def inner(context: Context[ClientT]) -> bool:
author = context.author
if not author.has_permissions(**permissions):
raise MissingPermissionsError
raise MissingPermissionsError(permissions)
return True
return inner
def has_channel_permissions(**permissions: bool):
def has_channel_permissions(**permissions: bool) -> Callable[[T], T]:
@check
def inner(context: Context[ClientT]):
def inner(context: Context[ClientT]) -> bool:
author = context.author
if isinstance(author, revolt.User):
raise MissingPermissionsError
if not isinstance(author, revolt.Member):
raise ServerOnly
if not author.has_channel_permissions(context.channel, **permissions):
raise MissingPermissionsError
raise MissingPermissionsError(permissions)
return True
return inner
+9 -10
View File
@@ -36,7 +36,7 @@ class ExtensionProtocol(Protocol):
class CommandsMeta(type):
_commands: list[Command[Any]]
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any]):
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any]) -> Self:
commands: list[Command[Any]] = []
self = super().__new__(cls, name, bases, attrs)
for base in reversed(self.__mro__):
@@ -95,6 +95,8 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
for alias in command.aliases:
self.all_commands[alias] = command
self.help_command: HelpCommand[Self] | None
if help_command is not None:
self.help_command = help_command or DefaultHelpCommand[Self]()
self.add_command(HelpCommandImpl(self))
@@ -137,7 +139,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""
return self.all_commands[name]
def add_command(self, command: Command[Self]):
def add_command(self, command: Command[Self]) -> None:
"""Adds a command, this is typically only used for dynamic commands, you should use the `commands.command` decorator for most usecases.
Parameters
@@ -194,9 +196,6 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""
content = message.content
if not isinstance(content, str):
return
prefixes = await self.get_prefix(message)
if isinstance(prefixes, str):
@@ -246,7 +245,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
await command._error_handler(command.cog or self, context, e)
self.dispatch("command_error", context, e)
async def on_command_error(self, ctx: Context[Self], error: Exception, /):
async def on_command_error(self, ctx: Context[Self], error: Exception, /) -> None:
traceback.print_exception(type(error), error, error.__traceback__)
on_message = process_commands
@@ -266,7 +265,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return True
def add_cog(self, cog: Cog[Self]):
def add_cog(self, cog: Cog[Self]) -> None:
"""Adds a cog to the bot, this cog must subclass `Cog`.
Parameters
@@ -294,7 +293,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return cog
def load_extension(self, name: str):
def load_extension(self, name: str) -> None:
"""Loads an extension, this takes a module name and runs the setup function inside of it.
Parameters
@@ -310,7 +309,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
self.extensions[name] = extension
extension.setup(self)
def unload_extension(self, name: str):
def unload_extension(self, name: str) -> None:
"""Unloads an extension, this takes a module name and runs the teardown function inside of it.
Parameters
@@ -325,7 +324,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
if teardown := getattr(extension, "teardown", None):
teardown(self)
def reload_extension(self, name: str):
def reload_extension(self, name: str) -> None:
"""Reloads an extension, this will unload and reload the extension.
Parameters
+6 -5
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from typing import Any, Generic, Optional, cast
from typing_extensions import Self
from .command import Command
from .utils import ClientT
@@ -11,7 +12,7 @@ class CogMeta(type, Generic[ClientT]):
_commands: list[Command[ClientT]]
qualified_name: str
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any], *, qualified_name: Optional[str] = None):
def __new__(cls, name: str, bases: tuple[type, ...], attrs: dict[str, Any], *, qualified_name: Optional[str] = None) -> Self:
commands: list[Command[ClientT]] = []
self = super().__new__(cls, name, bases, attrs)
@@ -29,15 +30,15 @@ class Cog(Generic[ClientT], metaclass=CogMeta):
_commands: list[Command[ClientT]]
qualified_name: str
def cog_load(self):
def cog_load(self) -> None:
"""A special method that is called when the cog gets loaded."""
pass
def cog_unload(self):
def cog_unload(self) -> None:
"""A special method that is called when the cog gets removed."""
pass
def _inject(self, client: ClientT):
def _inject(self, client: ClientT) -> None:
client.cogs[self.qualified_name] = self
for command in self._commands:
@@ -46,7 +47,7 @@ class Cog(Generic[ClientT], metaclass=CogMeta):
self.cog_load()
def _uninject(self, client: ClientT):
def _uninject(self, client: ClientT) -> None:
for name, command in client.all_commands.copy().items():
if command in self._commands:
del client.all_commands[name]
+27 -25
View File
@@ -5,11 +5,12 @@ import traceback
from contextlib import suppress
from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine,
Generic, Literal, Optional, Union, get_args, get_origin)
from typing_extensions import ParamSpec
from revolt.utils import copy_doc, maybe_coroutine
from .errors import InvalidLiteralArgument, UnionConverterError
from .utils import ClientT, evaluate_parameters
from .utils import ClientCoT, evaluate_parameters
if TYPE_CHECKING:
from .checks import Check
@@ -17,15 +18,15 @@ if TYPE_CHECKING:
from .context import Context
from .group import Group
__all__ = (
__all__: tuple[str, ...] = (
"Command",
"command"
)
NoneType = type(None)
NoneType: type[None] = type(None)
P = ParamSpec("P")
class Command(Generic[ClientT]):
class Command(Generic[ClientCoT]):
"""Class for holding info about a command.
Parameters
@@ -46,19 +47,19 @@ class Command(Generic[ClientT]):
__slots__ = ("callback", "name", "aliases", "signature", "checks", "parent", "_error_handler", "cog", "description", "usage", "parameters")
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str], usage: Optional[str] = None):
self.callback = callback
self.name = name
self.aliases = aliases
self.usage = usage
self.signature = inspect.signature(self.callback)
self.parameters = evaluate_parameters(self.signature.parameters.values(), getattr(callback, "__globals__", {}))
self.checks: list[Check[ClientT]] = getattr(callback, "_checks", [])
self.parent: Optional[Group[ClientT]] = None
self.cog: Optional[Cog[ClientT]] = None
self._error_handler: Callable[[Any, Context[ClientT], Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
self.description = callback.__doc__
self.callback: Callable[..., Coroutine[Any, Any, Any]] = callback
self.name: str = name
self.aliases: list[str] = aliases
self.usage: str | None = usage
self.signature: inspect.Signature = inspect.signature(self.callback)
self.parameters: list[inspect.Parameter] = evaluate_parameters(self.signature.parameters.values(), getattr(callback, "__globals__", {}))
self.checks: list[Check[ClientCoT]] = getattr(callback, "_checks", [])
self.parent: Optional[Group[ClientCoT]] = None
self.cog: Optional[Cog[ClientCoT]] = None
self._error_handler: Callable[[Any, Context[ClientCoT], Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
self.description: str | None = callback.__doc__
async def invoke(self, context: Context[ClientT], *args: Any, **kwargs: Any) -> Any:
async def invoke(self, context: Context[ClientCoT], *args: Any, **kwargs: Any) -> Any:
"""Runs the command and calls the error handler if the command errors.
Parameters
@@ -74,10 +75,10 @@ class Command(Generic[ClientT]):
return await self._error_handler(self.cog or context.client, context, err)
@copy_doc(invoke)
def __call__(self, context: Context[ClientT], *args: Any, **kwargs: Any) -> Any:
def __call__(self, context: Context[ClientCoT], *args: Any, **kwargs: Any) -> Any:
return self.invoke(context, *args, **kwargs)
def error(self, func: Callable[..., Coroutine[Any, Any, Any]]):
def error(self, func: Callable[..., Coroutine[Any, Any, Any]]) -> Callable[..., Coroutine[Any, Any, Any]]:
"""Sets the error handler for the command.
Parameters
@@ -97,11 +98,11 @@ class Command(Generic[ClientT]):
self._error_handler = func
return func
async def _default_error_handler(self, ctx: Context[ClientT], error: Exception):
async def _default_error_handler(self, ctx: Context[ClientCoT], error: Exception):
traceback.print_exception(type(error), error, error.__traceback__)
@classmethod
async def handle_origin(cls, context: Context[ClientT], origin: Any, annotation: Any, arg: str) -> Any:
async def handle_origin(cls, context: Context[ClientCoT], origin: Any, annotation: Any, arg: str) -> Any:
if origin is Union:
for converter in get_args(annotation):
try:
@@ -128,11 +129,12 @@ class Command(Generic[ClientT]):
raise InvalidLiteralArgument(arg)
@classmethod
async def convert_argument(cls, arg: str, annotation: Any, context: Context[ClientT]) -> Any:
async def convert_argument(cls, arg: str, annotation: Any, context: Context[ClientCoT]) -> Any:
if annotation is not inspect.Signature.empty:
if annotation is str: # no converting is needed - its already a string
return arg
origin: Any
if origin := get_origin(annotation):
return await cls.handle_origin(context, origin, annotation, arg)
else:
@@ -140,7 +142,7 @@ class Command(Generic[ClientT]):
else:
return arg
async def parse_arguments(self, context: Context[ClientT]):
async def parse_arguments(self, context: Context[ClientCoT]) -> None:
# please pr if you can think of a better way to do this
for parameter in self.parameters[2:]:
@@ -213,8 +215,8 @@ class Command(Generic[ClientT]):
return f"{' '.join(parents[::-1])} {self.name} {' '.join(parameters)}"
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command, usage: Optional[str] = None):
"""A decorator that turns a function into a :class:`Command`.
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientCoT]] = Command, usage: Optional[str] = None) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientCoT]]:
"""A decorator that turns a function into a :class:`Command`.n
Parameters
-----------
+36 -15
View File
@@ -7,16 +7,17 @@ from revolt.utils import maybe_coroutine
from .command import Command
from .group import Group
from .utils import ClientT
from .utils import ClientCoT
if TYPE_CHECKING:
from .view import StringView
from revolt.state import State
__all__ = (
"Context",
)
class Context(revolt.Messageable, Generic[ClientT]):
class Context(revolt.Messageable, Generic[ClientCoT]):
"""Stores metadata the commands execution.
Attributes
@@ -29,7 +30,7 @@ class Context(revolt.Messageable, Generic[ClientT]):
The message that was sent to invoke the command
channel: :class:`Messageable`
The channel the command was invoked in
server: :class:`Server`
server_id: Optional[:class:`Server`]
The server the command was invoked in
author: Union[:class:`Member`, :class:`User`]
The user or member that invoked the commad, will be :class:`User` in DMs
@@ -40,23 +41,37 @@ class Context(revolt.Messageable, Generic[ClientT]):
client: :class:`CommandsClient`
The revolt client
"""
__slots__ = ("command", "invoked_with", "args", "message", "server", "channel", "author", "view", "kwargs", "state", "client")
__slots__ = ("command", "invoked_with", "args", "message", "channel", "author", "view", "kwargs", "state", "client", "server_id")
async def _get_channel_id(self) -> str:
return self.channel.id
def __init__(self, command: Optional[Command[ClientT]], invoked_with: str, view: StringView, message: revolt.Message, client: ClientT):
self.command = command
self.invoked_with = invoked_with
self.view = view
self.message = message
self.client = client
def __init__(self, command: Optional[Command[ClientCoT]], invoked_with: str, view: StringView, message: revolt.Message, client: ClientCoT):
self.command: Command[ClientCoT] | None = command
self.invoked_with: str = invoked_with
self.view: StringView = view
self.message: revolt.Message = message
self.client: ClientCoT = client
self.args: list[Any] = []
self.kwargs: dict[str, Any] = {}
self.server = message.server
self.channel = message.channel
self.author = message.author
self.state = message.state
self.server_id: str | None = message.server_id
self.channel: revolt.TextChannel | revolt.GroupDMChannel | revolt.DMChannel = message.channel
self.author: revolt.Member | revolt.User = message.author
self.state: State = message.state
@property
def server(self) -> revolt.Server:
""":class:`Server` The server this voice channel belongs too
Raises
-------
:class:`LookupError`
Raises if the channel is not part of a server
"""
if not self.server_id:
raise LookupError
return self.state.get_server(self.server_id)
async def invoke(self) -> Any:
"""Invokes the command.
@@ -85,8 +100,14 @@ class Context(revolt.Messageable, Generic[ClientT]):
await command.parse_arguments(self)
return await command.invoke(self, *self.args, **self.kwargs)
async def can_run(self, command: Optional[Command[ClientT]] = None) -> bool:
async def can_run(self, command: Optional[Command[ClientCoT]] = None) -> bool:
"""Runs all of the commands checks, and returns true if all of them pass"""
command = command or self.command
return all([await maybe_coroutine(check, self) for check in (command.checks if command else [])])
async def send_help(self, argument: Command[Any] | Group[Any] | ClientCoT | None = None) -> None:
argument = argument or self.client
command = self.client.get_command("help")
await command.invoke(self, argument)
+15 -15
View File
@@ -13,46 +13,46 @@ from .errors import (BadBoolArgument, CategoryConverterError,
if TYPE_CHECKING:
from .client import CommandsClient
__all__ = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter")
__all__: tuple[str, ...] = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter")
channel_regex = re.compile("<#([A-z0-9]{26})>")
user_regex = re.compile("<@([A-z0-9]{26})>")
channel_regex: re.Pattern[str] = re.compile("<#([A-z0-9]{26})>")
user_regex: re.Pattern[str] = re.compile("<@([A-z0-9]{26})>")
ClientT = TypeVar("ClientT", bound="CommandsClient")
def bool_converter(arg: str, _):
def bool_converter(arg: str, _: Context[ClientT]) -> bool:
lowered = arg.lower()
if lowered in ["yes", "true", "ye", "y", "1", "on", "enable"]:
if lowered in ("yes", "true", "ye", "y", "1", "on", "enable"):
return True
elif lowered in ('no', 'n', 'false', 'f', '0', 'disable', 'off'):
elif lowered in ("no", "false", "n", "f", "0", "off", "disabled"):
return False
else:
raise BadBoolArgument(lowered)
def category_converter(arg: str, context: Context[ClientT]) -> Category:
if not (server := context.server):
if not context.server_id:
raise ServerOnly
try:
return server.get_category(arg)
return context.server.get_category(arg)
except KeyError:
try:
return utils.get(server.categories, name=arg)
return utils.get(context.server.categories, name=arg)
except LookupError:
raise CategoryConverterError(arg)
def channel_converter(arg: str, context: Context[ClientT]) -> Channel:
if not (server := context.server):
if not context.server_id:
raise ServerOnly
if (match := channel_regex.match(arg)):
arg = match.group(1)
try:
return server.get_channel(arg)
return context.server.get_channel(arg)
except KeyError:
try:
return utils.get(server.channels, name=arg)
return utils.get(context.server.channels, name=arg)
except LookupError:
raise ChannelConverterError(arg)
@@ -69,17 +69,17 @@ def user_converter(arg: str, context: Context[ClientT]) -> User:
raise UserConverterError(arg)
def member_converter(arg: str, context: Context[ClientT]) -> Member:
if not (server := context.server):
if not context.server_id:
raise ServerOnly
if (match := user_regex.match(arg)):
arg = match.group(1)
try:
return server.get_member(arg)
return context.server.get_member(arg)
except KeyError:
try:
return utils.get(server.members, name=arg)
return utils.get(context.server.members, name=arg)
except LookupError:
raise MemberConverterError(arg)
+11 -2
View File
@@ -32,7 +32,7 @@ class CommandNotFound(CommandError):
__slots__ = ("command_name",)
def __init__(self, command_name: str):
self.command_name = command_name
self.command_name: str = command_name
class NoClosingQuote(CommandError):
"""Raised when there is no closing quote for a command argument"""
@@ -50,7 +50,16 @@ class ServerOnly(CheckError):
"""Raised when a check requires the command to be ran in a server"""
class MissingPermissionsError(CheckError):
"""Raised when a check requires permissions the user does not have"""
"""Raised when a check requires permissions the user does not have
Attributes
-----------
permissions: :class:`dict[str, bool]`
The permissions which the user did not have
"""
def __init__(self, permissions: dict[str, bool]):
self.permissions = permissions
class ConverterError(CommandError):
"""Base class for all converter errors"""
+9 -13
View File
@@ -1,21 +1,17 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Callable, Coroutine, Optional, TypeVar
from typing import Any, Callable, Coroutine, Optional
from .command import Command
from .utils import ClientCoT, ClientT
if TYPE_CHECKING:
from .client import CommandsClient
__all__ = (
"Group",
"group"
)
ClientT = TypeVar("ClientT", bound="CommandsClient")
class Group(Command[ClientT]):
class Group(Command[ClientCoT]):
"""Class for holding info about a group command.
Parameters
@@ -30,13 +26,13 @@ class Group(Command[ClientT]):
The group's subcommands.
"""
__slots__ = ("subcommands",)
__slots__: tuple[str, ...] = ("subcommands",)
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]):
self.subcommands: dict[str, Command[ClientT]] = {}
self.subcommands: dict[str, Command[ClientCoT]] = {}
super().__init__(callback, name, aliases)
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT]] = Command[ClientT]):
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientCoT]] = Command[ClientCoT]) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientCoT]]:
"""A decorator that turns a function into a :class:`Command` and registers the command as a subcommand.
Parameters
@@ -61,7 +57,7 @@ class Group(Command[ClientT]):
return inner
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type[Group[ClientT]]] = None):
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type[Group[ClientCoT]]] = None) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientCoT]]:
"""A decorator that turns a function into a :class:`Group` and registers the command as a subcommand
Parameters
@@ -92,10 +88,10 @@ class Group(Command[ClientT]):
return f"<Group name=\"{self.name}\">"
@property
def commands(self) -> list[Command[ClientT]]:
def commands(self) -> list[Command[ClientCoT]]:
return list(self.subcommands.values())
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group[ClientT]] = Group):
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group[ClientT]] = Group) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientT]]:
"""A decorator that turns a function into a :class:`Group`
Parameters
+30 -28
View File
@@ -9,7 +9,7 @@ from .cog import Cog
from .command import Command
from .context import Context
from .group import Group
from .utils import ClientT
from .utils import ClientCoT, ClientT
if TYPE_CHECKING:
from revolt import File, Message, Messageable, MessageReply, SendableEmbed
@@ -18,6 +18,7 @@ if TYPE_CHECKING:
__all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl")
class MessagePayload(TypedDict):
content: str
embed: NotRequired[SendableEmbed]
@@ -25,28 +26,28 @@ class MessagePayload(TypedDict):
attachments: NotRequired[list[File]]
replies: NotRequired[list[MessageReply]]
class HelpCommand(ABC, Generic[ClientT]):
class HelpCommand(ABC, Generic[ClientCoT]):
@abstractmethod
async def create_bot_help(self, context: Context[ClientT], commands: dict[Optional[Cog[ClientT]], list[Command[ClientT]]]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_command_help(self, context: Context[ClientT], command: Command[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_command_help(self, context: Context[ClientCoT], command: Command[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_group_help(self, context: Context[ClientT], group: Group[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_group_help(self, context: Context[ClientCoT], group: Group[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
@abstractmethod
async def create_cog_help(self, context: Context[ClientT], cog: Cog[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_cog_help(self, context: Context[ClientCoT], cog: Cog[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError
async def send_help_command(self, context: Context[ClientT], message_payload: MessagePayload) -> Message:
async def send_help_command(self, context: Context[ClientCoT], message_payload: MessagePayload) -> Message:
return await context.send(**message_payload)
async def filter_commands(self, context: Context[ClientT], commands: list[Command[ClientT]]) -> list[Command[ClientT]]:
filtered: list[Command[ClientT]] = []
async def filter_commands(self, context: Context[ClientCoT], commands: list[Command[ClientCoT]]) -> list[Command[ClientCoT]]:
filtered: list[Command[ClientCoT]] = []
for command in commands:
try:
@@ -57,34 +58,34 @@ class HelpCommand(ABC, Generic[ClientT]):
return filtered
async def group_commands(self, context: Context[ClientT], commands: list[Command[ClientT]]) -> dict[Optional[Cog[ClientT]], list[Command[ClientT]]]:
cogs: dict[Optional[Cog[ClientT]], list[Command[ClientT]]] = {}
async def group_commands(self, context: Context[ClientCoT], commands: list[Command[ClientCoT]]) -> dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]:
cogs: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]] = {}
for command in commands:
cogs.setdefault(command.cog, []).append(command)
return cogs
async def handle_message(self, context: Context[ClientT], message: Message):
async def handle_message(self, context: Context[ClientCoT], message: Message) -> None:
pass
async def get_channel(self, context: Context[ClientT]) -> Messageable:
async def get_channel(self, context: Context) -> Messageable:
return context
@abstractmethod
async def handle_no_command_found(self, context: Context[ClientT], name: str) -> Any:
async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> Any:
raise NotImplementedError
@abstractmethod
async def handle_no_cog_found(self, context: Context[ClientT], name: str) -> Any:
async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> Any:
raise NotImplementedError
class DefaultHelpCommand(HelpCommand[ClientT]):
class DefaultHelpCommand(HelpCommand[ClientCoT]):
def __init__(self, default_cog_name: str = "No Cog"):
self.default_cog_name = default_cog_name
async def create_bot_help(self, context: Context[ClientT], commands: dict[Optional[Cog[ClientT]], list[Command[ClientT]]]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
for cog, cog_commands in commands.items():
@@ -99,7 +100,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
lines.append("```")
return "\n".join(lines)
async def create_cog_help(self, context: Context[ClientT], cog: Cog[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_cog_help(self, context: Context[ClientCoT], cog: Cog[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{cog.qualified_name}:")
@@ -110,7 +111,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
lines.append("```")
return "\n".join(lines)
async def create_command_help(self, context: Context[ClientT], command: Command[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_command_help(self, context: Context[ClientCoT], command: Command[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{command.name}:")
@@ -126,7 +127,7 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
lines.append("```")
return "\n".join(lines)
async def create_group_help(self, context: Context[ClientT], group: Group[ClientT]) -> Union[str, SendableEmbed, MessagePayload]:
async def create_group_help(self, context: Context[ClientCoT], group: Group[ClientCoT]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"]
lines.append(f"{group.name}:")
@@ -144,27 +145,27 @@ class DefaultHelpCommand(HelpCommand[ClientT]):
lines.append("```")
return "\n".join(lines)
async def handle_no_command_found(self, context: Context[ClientT], name: str):
async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> None:
channel = await self.get_channel(context)
await channel.send(f"Command `{name}` not found.")
async def handle_no_cog_found(self, context: Context[ClientT], name: str):
async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> None:
channel = await self.get_channel(context)
await channel.send(f"Cog `{name}` not found.")
class HelpCommandImpl(Command[ClientT]):
def __init__(self, client: ClientT):
class HelpCommandImpl(Command[ClientCoT]):
def __init__(self, client: ClientCoT):
self.client = client
async def callback(_: Union[ClientT, Cog[ClientT]], context: Context[ClientT], *args: str):
async def callback(_: Union[ClientCoT, Cog[ClientCoT]], context: Context[ClientCoT], *args: str) -> None:
await help_command_impl(context.client, context, *args)
super().__init__(callback=callback, name="help", aliases=[])
self.description = "Shows help for a command, cog or the entire bot"
self.description: str | None = "Shows help for a command, cog or the entire bot"
async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments: str):
async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments: str) -> None:
help_command = self.help_command
if not help_command:
@@ -201,4 +202,5 @@ async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments
else:
msg_payload = payload
await help_command.send_help_command(context, msg_payload)
message = await help_command.send_help_command(context, msg_payload)
await help_command.handle_message(context, message)
+3 -1
View File
@@ -7,12 +7,14 @@ from typing_extensions import TypeVar
if TYPE_CHECKING:
from .client import CommandsClient
from .context import Context
__all__ = ("evaluate_parameters",)
ClientT = TypeVar("ClientT", bound="CommandsClient", default="CommandsClient")
ClientCoT = TypeVar("ClientCoT", bound="CommandsClient", default="CommandsClient", covariant=True)
ContextT = TypeVar("ContextT", bound="Context")
def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]:
new_parameters: list[Parameter] = []
+5 -4
View File
@@ -1,13 +1,14 @@
from typing import Iterator
from .errors import NoClosingQuote
class StringView:
def __init__(self, string: str):
self.value = iter(string)
self.temp = ""
self.should_undo = False
self.value: Iterator[str] = iter(string)
self.temp: str = ""
self.should_undo: bool = False
def undo(self):
def undo(self) -> None:
self.should_undo = True
def next_char(self) -> str:
+8 -4
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
import io
from typing import Optional, Union
from typing import Optional, Union, cast
__all__ = ("File",)
@@ -18,17 +20,19 @@ class File:
__slots__ = ("f", "spoiler", "filename")
def __init__(self, file: Union[str, bytes], *, filename: Optional[str] = None, spoiler: bool = False):
self.f: io.BufferedIOBase
if isinstance(file, str):
self.f = open(file, "rb")
else:
self.f = io.BytesIO(file)
if filename is None and isinstance(file, str):
filename = self.f.name
filename = cast(Optional[str], self.f.name)
self.spoiler = spoiler or (filename and filename.startswith("SPOILER_"))
self.spoiler: bool = spoiler or (bool(filename) and filename.startswith("SPOILER_"))
if self.spoiler and (filename and not filename.startswith("SPOILER_")):
filename = f"SPOILER_{filename}"
self.filename = filename
self.filename: str | None = filename
+11 -5
View File
@@ -11,8 +11,8 @@ class Flag:
__slots__ = ("flag", "__doc__")
def __init__(self, func: Callable[[], int]):
self.flag = func()
self.__doc__ = func.__doc__
self.flag: int = func()
self.__doc__: str | None = func.__doc__
@overload
def __get__(self: Self, instance: None, owner: type[Flags]) -> Self:
@@ -28,7 +28,7 @@ class Flag:
return instance._check_flag(self.flag)
def __set__(self, instance: Flags, value: bool):
def __set__(self, instance: Flags, value: bool) -> None:
instance._set_flag(self.flag, value)
class Flags:
@@ -37,6 +37,12 @@ class Flags:
def __init_subclass__(cls) -> None:
cls.FLAG_NAMES = []
for name in dir(cls):
value = getattr(cls, name)
if isinstance(value, Flag):
cls.FLAG_NAMES.append(name)
def __init__(self, value: int = 0, **flags: bool):
self.value = value
@@ -52,7 +58,7 @@ class Flags:
def _check_flag(self, flag: int) -> bool:
return (self.value & flag) == flag
def _set_flag(self, flag: int, value: bool):
def _set_flag(self, flag: int, value: bool) -> None:
if value:
self.value |= flag
else:
@@ -85,7 +91,7 @@ class Flags:
def __gt__(self, other: Self) -> bool:
return self.value > other.value
def __repr__(self):
def __repr__(self) -> str:
return f"<{self.__class__.__name__} value={self.value}>"
def __iter__(self) -> Iterator[tuple[str, bool]]:
+17 -18
View File
@@ -8,7 +8,6 @@ import ulid
from .errors import Forbidden, HTTPError, ServerError
from .file import File
from .utils import Missing
try:
import ujson as _json
@@ -47,11 +46,11 @@ class HttpClient:
__slots__ = ("session", "token", "api_url", "api_info", "auth_header")
def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo, bot: bool = True):
self.session = session
self.token = token
self.api_url = api_url
self.api_info = api_info
self.auth_header = "x-bot-token" if bot else "x-session-token"
self.session: aiohttp.ClientSession = session
self.token: str = token
self.api_url: str = api_url
self.api_info: ApiInfo = api_info
self.auth_header: str = "x-bot-token" if bot else "x-session-token"
async def request(self, method: Literal["GET", "POST", "PUT", "DELETE", "PATCH"], route: str, *, json: Optional[dict[str, Any]] = None, nonce: bool = True, params: Optional[dict[str, Any]] = None) -> Any:
url = f"{self.api_url}{route}"
@@ -203,7 +202,7 @@ class HttpClient:
include_users: bool = False
) -> Request[Union[list[MessagePayload], MessageWithUserData]]:
json = {"sort": sort.value, "include_users": str(include_users)}
json: dict[str, Any] = {"sort": sort.value, "include_users": str(include_users)}
if limit:
json["limit"] = limit
@@ -360,19 +359,19 @@ class HttpClient:
def delete_invite(self, code: str) -> Request[None]:
return self.request("DELETE", f"/invites/{code}")
def edit_channel(self, channel_id: str, remove: list[str] | None, values: dict[str, Any]):
def edit_channel(self, channel_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove:
values["remove"] = remove
return self.request("PATCH", f"/channels/{channel_id}", json=values)
def edit_role(self, server_id: str, role_id: str, remove: list[str] | None, values: dict[str, Any]):
def edit_role(self, server_id: str, role_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove:
values["remove"] = remove
return self.request("PATCH", f"/servers/{server_id}/roles/{role_id}", json=values)
async def edit_self(self, remove: list[str] | None, values: dict[str, Any]):
async def edit_self(self, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove:
values["remove"] = remove
@@ -393,25 +392,25 @@ class HttpClient:
def set_guild_channel_role_permissions(self, channel_id: str, role_id: str, allow: int, deny: int) -> Request[None]:
return self.request("PUT", f"/channels/{channel_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}})
def set_group_channel_default_permissions(self, channel_id: str, value: int):
def set_group_channel_default_permissions(self, channel_id: str, value: int) -> Request[None]:
return self.request("PUT", f"/channels/{channel_id}/permissions/default", json={"permissions": value})
def set_server_role_permissions(self, server_id: str, role_id: str, allow: int, deny: int):
def set_server_role_permissions(self, server_id: str, role_id: str, allow: int, deny: int) -> Request[None]:
return self.request("PUT", f"/servers/{server_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}})
def set_server_default_permissions(self, server_id: str, value: int):
def set_server_default_permissions(self, server_id: str, value: int) -> Request[None]:
return self.request("PUT", f"/servers/{server_id}/permissions/default", json={"permissions": value})
def add_reaction(self, channel_id: str, message_id: str, emoji: str):
def add_reaction(self, channel_id: str, message_id: str, emoji: str) -> Request[None]:
return self.request("PUT", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}")
def remove_reaction(self, channel_id: str, message_id: str, emoji: str, user_id: Optional[str], remove_all: bool):
def remove_reaction(self, channel_id: str, message_id: str, emoji: str, user_id: Optional[str], remove_all: bool) -> Request[None]:
return self.request("PUT", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}")
def remove_all_reactions(self, channel_id: str, message_id: str):
def remove_all_reactions(self, channel_id: str, message_id: str) -> Request[None]:
return self.request("DELETE", f"/channels/{channel_id}/messages/{message_id}/reactions")
def delete_emoji(self, emoji_id: str):
def delete_emoji(self, emoji_id: str) -> Request[None]:
return self.request("DELETE", f"/custom/emoji/{emoji_id}")
def fetch_emoji(self, emoji_id: str) -> Request[EmojiPayload]:
@@ -425,5 +424,5 @@ class HttpClient:
def edit_member(self, server_id: str, member_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[MemberPayload]:
return self.request("PATCH", f"/servers/{server_id}/members/{member_id}", json={"remove": remove, **values})
def delete_messages(self, channel_id: str, messages: list[str]):
def delete_messages(self, channel_id: str, messages: list[str]) -> Request[None]:
return self.request("DELETE", f"/channels/{channel_id}/messages/bulk", json={"ids": messages})
+16 -9
View File
@@ -2,12 +2,17 @@ from __future__ import annotations
from typing import TYPE_CHECKING
from .asset import Asset
from .utils import Ulid
if TYPE_CHECKING:
from .state import State
from .channel import Channel
from .server import Server
from .types import Invite as InvitePayload
from .user import User
__all__ = ("Invite",)
@@ -37,22 +42,24 @@ class Invite(Ulid):
__slots__ = ("state", "code", "id", "server", "channel", "user_name", "user_avatar", "user", "member_count")
def __init__(self, data: InvitePayload, code: str, state: State):
self.state = state
self.state: State = state
self.code = code
self.id = code
self.server = state.get_server(data["server_id"])
self.channel = self.server.get_channel(data["channel_id"])
self.code: str = code
self.id: str = code
self.server: Server = state.get_server(data["server_id"])
self.channel: Channel = self.server.get_channel(data["channel_id"])
self.user_name = data["user_name"]
self.user = None
self.user_name: str = data["user_name"]
self.user: User | None = None
self.user_avatar: Asset | None
if avatar := data.get("user_avatar"):
self.user_avatar = Asset(avatar, state)
else:
self.user_avatar = None
self.member_count = data["member_count"]
self.member_count: int = data["member_count"]
@staticmethod
def _from_partial(code: str, server: str, creator: str, channel: str, state: State) -> Invite:
@@ -69,6 +76,6 @@ class Invite(Ulid):
return invite
async def delete(self):
async def delete(self) -> None:
"""Deletes the invite"""
await self.state.http.delete_invite(self.code)
+69 -17
View File
@@ -1,23 +1,28 @@
from __future__ import annotations
import datetime
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Any, Optional
from .utils import _Missing, Missing
from .asset import Asset
from .permissions import Permissions
from .permissions_calculator import calculate_permissions
from .user import User
from .file import File
if TYPE_CHECKING:
from .channel import Channel
from .server import Server
from .state import State
from .types import File
from .types import File as FilePayload
from .types import Member as MemberPayload
from .role import Role
__all__ = ("Member",)
def flattern_user(member: Member, user: User):
def flattern_user(member: Member, user: User) -> None:
for attr in user.__flattern_attributes__:
setattr(member, attr, getattr(user, attr))
@@ -42,9 +47,11 @@ class Member(User):
# due to not having a user payload and only a user object we have to manually add all the attributes instead of calling User.__init__
flattern_user(self, user)
user._members.add(self)
user._members[server.id] = self
self.state = state
self.state: State = state
self.guild_avatar: Asset | None
if avatar := data.get("avatar"):
self.guild_avatar = Asset(avatar, state)
@@ -52,32 +59,40 @@ class Member(User):
self.guild_avatar = None
roles = [server.get_role(role_id) for role_id in data.get("roles", [])]
self.roles = sorted(roles, key=lambda role: role.rank, reverse=True)
self.roles: list[Role] = sorted(roles, key=lambda role: role.rank, reverse=True)
self.server = server
self.nickname = data.get("nickname")
self.server: Server = server
self.nickname: str | None = data.get("nickname")
joined_at = data["joined_at"]
if isinstance(joined_at, int):
self.joined_at = datetime.datetime.fromtimestamp(joined_at / 1000)
self.joined_at: datetime.datetime = datetime.datetime.fromtimestamp(joined_at / 1000)
else:
self.joined_at = datetime.datetime.strptime(joined_at, "%Y-%m-%dT%H:%M:%S.%f%z")
self.current_timeout = None
self.joined_at: datetime.datetime = datetime.datetime.strptime(joined_at, "%Y-%m-%dT%H:%M:%S.%f%z")
self.current_timeout: datetime.datetime | None
if current_timeout := data.get("timeout"):
self.current_timeout = datetime.datetime.strptime(current_timeout, "%Y-%m-%dT%H:%M:%S.%f%z")
else:
self.current_timeout = None
@property
def avatar(self) -> Optional[Asset]:
"""Optional[:class:`Asset`] The avatar the member is displaying, this includes guild avatars and masqueraded avatar"""
return self.masquerade_avatar or self.guild_avatar or self.original_avatar
@property
def name(self) -> str:
""":class:`str` The name the user is displaying, this includes (in order) their masqueraded name, display name and orginal name"""
return self.nickname or self.display_name or self.masquerade_name or self.original_name
@property
def mention(self) -> str:
""":class:`str`: Returns a string that allows you to mention the given member."""
return f"<@{self.id}>"
def _update(self, *, nickname: Optional[str] = None, avatar: Optional[File] = None, roles: Optional[list[str]] = None):
def _update(self, *, nickname: Optional[str] = None, avatar: Optional[FilePayload] = None, roles: Optional[list[str]] = None):
if nickname is not None:
self.nickname = nickname
@@ -88,11 +103,11 @@ class Member(User):
member_roles = [self.server.get_role(role_id) for role_id in roles]
self.roles = sorted(member_roles, key=lambda role: role.rank, reverse=True)
async def kick(self):
async def kick(self) -> None:
"""Kicks the member from the server"""
await self.state.http.kick_member(self.server.id, self.id)
async def ban(self, *, reason: Optional[str] = None):
async def ban(self, *, reason: Optional[str] = None) -> None:
"""Bans the member from the server
Parameters
@@ -102,11 +117,48 @@ class Member(User):
"""
await self.state.http.ban_member(self.server.id, self.id, reason)
async def unban(self):
async def unban(self) -> None:
"""Unbans the member from the server"""
await self.state.http.unban_member(self.server.id, self.id)
async def timeout(self, length: datetime.timedelta):
async def edit(
self,
*,
nickname: str | None | _Missing = Missing,
roles: list[Role] | None | _Missing = Missing,
avatar: File | None | _Missing = Missing,
timeout: datetime.timedelta | None | _Missing = Missing
) -> None:
remove: list[str] = []
data: dict[str, Any] = {}
if nickname is None:
remove.append("Nickname")
elif nickname is not Missing:
data["nickname"] = nickname
if roles is None:
remove.append("Roles")
elif roles is not Missing:
data["roles"] = roles
if avatar is None:
remove.append("Avatar")
elif avatar is not Missing:
# pyright cant understand custom singletons - it doesnt know this will never be an instance of _Missing here because Missing is the only instance
assert not isinstance(avatar, _Missing)
data["avatar"] = (await self.state.http.upload_file(avatar, "avatars"))["id"]
if timeout is None:
remove.append("Timeout")
elif timeout is not Missing:
assert not isinstance(timeout, _Missing)
data["timeout"] = (datetime.datetime.now(datetime.timezone.utc) + timeout).isoformat()
await self.state.http.edit_member(self.server.id, self.id, remove, data)
async def timeout(self, length: datetime.timedelta) -> None:
"""Timeouts the member
Parameters
@@ -128,7 +180,7 @@ class Member(User):
"""
return calculate_permissions(self, self.server)
def get_channel_permissions(self, channel: Channel):
def get_channel_permissions(self, channel: Channel) -> Permissions:
"""Gets the permissions for the member in the server taking into account the channel as well
Parameters
+58 -34
View File
@@ -1,11 +1,13 @@
from __future__ import annotations
import datetime
from typing import TYPE_CHECKING, Any, Optional
from typing import TYPE_CHECKING, Any, Coroutine, Optional, Union
from revolt.types.message import SystemMessageContent
from .asset import Asset, PartialAsset
from .channel import Messageable
from .embed import SendableEmbed, to_embed
from .channel import DMChannel, GroupDMChannel, TextChannel
from .embed import Embed, SendableEmbed, to_embed
from .utils import Ulid
if TYPE_CHECKING:
@@ -17,6 +19,7 @@ if TYPE_CHECKING:
from .types import Message as MessagePayload
from .types import MessageReplyPayload
from .user import User
from .member import Member
__all__ = (
"Message",
@@ -58,23 +61,37 @@ class Message(Ulid):
__slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "mentions", "replies", "reply_ids", "reactions", "interactions")
def __init__(self, data: MessagePayload, state: State):
self.state = state
self.state: State = state
self.id = data["_id"]
self.content = data.get("content", "")
self.attachments = [Asset(attachment, state) for attachment in data.get("attachments", [])]
self.embeds = [to_embed(embed, state) for embed in data.get("embeds", [])]
self.id: str = data["_id"]
content = data.get("content", "")
if not isinstance(content, str):
self.system_content: SystemMessageContent = content
self.content: str = ""
else:
self.content = content
self.attachments: list[Asset] = [Asset(attachment, state) for attachment in data.get("attachments", [])]
self.embeds: list[Embed] = [to_embed(embed, state) for embed in data.get("embeds", [])]
channel = state.get_channel(data["channel"])
assert isinstance(channel, Messageable)
self.channel = channel
assert isinstance(channel, Union[TextChannel, GroupDMChannel, DMChannel])
self.channel: TextChannel | GroupDMChannel | DMChannel = channel
if server_id := self.channel.server_id:
author = state.get_member(server_id, data["author"])
self.server_id: str | None = self.channel.server_id
self.mentions: list[Member | User]
if self.server_id:
author = state.get_member(self.server_id, data["author"])
self.mentions = [self.server.get_member(member_id) for member_id in data.get("mentions", [])]
else:
author = state.get_user(data["author"])
self.mentions = [state.get_user(member_id) for member_id in data.get("mentions", [])]
self.author = author
self.author: Member | User = author
if masquerade := data.get("masquerade"):
if name := masquerade.get("name"):
@@ -86,11 +103,6 @@ class Message(Ulid):
if edited_at := data.get("edited"):
self.edited_at: Optional[datetime.datetime] = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z")
if self.server:
self.mentions = [self.server.get_member(member_id) for member_id in data.get("mentions", [])]
else:
self.mentions = [state.get_user(member_id) for member_id in data.get("mentions", [])]
self.replies: list[Message] = []
self.reply_ids: list[str] = []
@@ -110,20 +122,26 @@ class Message(Ulid):
for emoji, users in reactions.items():
self.reactions[emoji] = [self.state.get_user(user_id) for user_id in users]
self.interactions: MessageInteractions | None
if interactions := data.get("interactions"):
self.interactions = MessageInteractions(reactions=interactions.get("reactions"), restrict_reactions=interactions.get("restrict_reactions", False))
else:
self.interactions = None
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited: int):
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited: Optional[Union[str, int]] = None):
if content is not None:
self.content = content
self.edited = datetime.datetime.fromtimestamp(edited / 1000)
if embeds is not None:
self.embeds = [to_embed(embed, self.state) for embed in embeds]
if edited is not None:
if isinstance(edited, int):
self.edited_at = datetime.datetime.fromtimestamp(edited / 1000, tz=datetime.timezone.utc)
else:
self.edited_at = datetime.datetime.strptime(edited, "%Y-%m-%dT%H:%M:%S.%f%z")
async def edit(self, *, content: Optional[str] = None, embeds: Optional[list[SendableEmbed]] = None) -> None:
"""Edits the message. The bot can only edit its own message
Parameters
@@ -140,7 +158,7 @@ class Message(Ulid):
"""Deletes the message. The bot can only delete its own messages and messages it has permission to delete """
await self.state.http.delete_message(self.channel.id, self.id)
def reply(self, *args: Any, mention: bool = False, **kwargs: Any):
def reply(self, *args: Any, mention: bool = False, **kwargs: Any) -> Coroutine[Any, Any, Message]:
"""Replies to this message, equivilant to:
.. code-block:: python
@@ -150,19 +168,25 @@ class Message(Ulid):
"""
return self.channel.send(*args, **kwargs, replies=[MessageReply(self, mention)])
async def add_reaction(self, emoji: str):
async def add_reaction(self, emoji: str) -> None:
await self.state.http.add_reaction(self.channel.id, self.id, emoji)
async def remove_reaction(self, emoji: str, user: Optional[User] = None, remove_all: bool = False):
async def remove_reaction(self, emoji: str, user: Optional[User] = None, remove_all: bool = False) -> None:
await self.state.http.remove_reaction(self.channel.id, self.id, emoji, user.id if user else None, remove_all)
async def remove_all_reactions(self):
async def remove_all_reactions(self) -> None:
await self.state.http.remove_all_reactions(self.channel.id, self.id)
@property
def server(self) -> Server:
""":class:`Server` The server this voice channel belongs too"""
""":class:`Server` The server this voice channel belongs too
Raises
-------
:class:`LookupError`
Raises if the channel is not part of a server
"""
return self.channel.server
class MessageReply:
@@ -178,8 +202,8 @@ class MessageReply:
__slots__ = ("message", "mention")
def __init__(self, message: Message, mention: bool = False):
self.message = message
self.mention = mention
self.message: Message = message
self.mention: bool = mention
def to_dict(self) -> MessageReplyPayload:
return { "id": self.message.id, "mention": self.mention }
@@ -199,9 +223,9 @@ class Masquerade:
__slots__ = ("name", "avatar", "colour")
def __init__(self, name: Optional[str] = None, avatar: Optional[str] = None, colour: Optional[str] = None):
self.name = name
self.avatar = avatar
self.colour = colour
self.name: str | None = name
self.avatar: str | None = avatar
self.colour: str | None = colour
def to_dict(self) -> MasqueradePayload:
output: MasqueradePayload = {}
@@ -230,10 +254,10 @@ class MessageInteractions:
__slots__ = ("reactions", "restrict_reactions")
def __init__(self, *, reactions: Optional[list[str]] = None, restrict_reactions: bool = False):
self.reactions = reactions
self.restrict_reactions = restrict_reactions
self.reactions: list[str] | None = reactions
self.restrict_reactions: bool = restrict_reactions
def to_dict(self):
def to_dict(self) -> InteractionsPayload:
output: InteractionsPayload = {}
if reactions := self.reactions:
+1 -1
View File
@@ -138,7 +138,7 @@ class Messageable:
payloads = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after)
return [Message(payload, self.state) for payload in payloads]
async def delete_messages(self, messages: list[Message]):
async def delete_messages(self, messages: list[Message]) -> None:
"""Bulk deletes messages from the channel
.. note:: The messages must have been sent in the last 7 days.
+2 -2
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional, cast
from typing import TYPE_CHECKING, Any, Optional
from typing_extensions import Self
@@ -176,7 +176,7 @@ class PermissionsOverwrite:
super().__setattr__(perm, value)
def __setattr__(self, key: str, value: Any):
def __setattr__(self, key: str, value: Any) -> None:
if key in Permissions.FLAG_NAMES:
if key is True:
setattr(self._allow, key, True)
+13 -13
View File
@@ -35,20 +35,20 @@ class Role(Ulid):
channel_permissions: :class:`ChannelPermissions`
The channel permissions for the role
"""
__slots__ = ("id", "name", "colour", "hoist", "rank", "state", "server", "permissions")
__slots__: tuple[str, ...] = ("id", "name", "colour", "hoist", "rank", "state", "server", "permissions")
def __init__(self, data: RolePayload, role_id: str, server: Server, state: State):
self.state = state
self.id = role_id
self.name = data["name"]
self.colour = data.get("colour", None)
self.hoist = False
self.rank = 0
self.server = server
self.permissions = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
self.state: State = state
self.id: str = role_id
self.name: str = data["name"]
self.colour: str | None = data.get("colour", None)
self.hoist: bool = data.get("hoist", False)
self.rank: int = data["rank"]
self.server: Server = server
self.permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
@property
def color(self):
def color(self) -> str | None:
return self.colour
async def set_permissions_overwrite(self, *, permissions: PermissionsOverwrite) -> None:
@@ -63,7 +63,7 @@ class Role(Ulid):
allow, deny = permissions.to_pair()
await self.state.http.set_server_role_permissions(self.server.id, self.id, allow.value, deny.value)
def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None, permissions: Optional[Overwrite] = None):
def _update(self, *, name: Optional[str] = None, colour: Optional[str] = None, hoist: Optional[bool] = None, rank: Optional[int] = None, permissions: Optional[Overwrite] = None) -> None:
if name is not None:
self.name = name
@@ -79,11 +79,11 @@ class Role(Ulid):
if permissions is not None:
self.permissions = PermissionsOverwrite._from_overwrite(permissions)
async def delete(self):
async def delete(self) -> None:
"""Deletes the role"""
await self.state.http.delete_role(self.server.id, self.id)
async def edit(self, **kwargs: Any):
async def edit(self, **kwargs: Any) -> None:
"""Edits the role
Parameters
+26 -22
View File
@@ -25,11 +25,11 @@ __all__ = ("Server", "SystemMessages", "ServerBan")
class SystemMessages:
def __init__(self, data: SystemMessagesConfig, state: State):
self.state = state
self.user_joined_id = data.get("user_joined")
self.user_left_id = data.get("user_left")
self.user_kicked_id = data.get("user_kicked")
self.user_banned_id = data.get("user_banned")
self.state: State = state
self.user_joined_id: str | None = data.get("user_joined")
self.user_left_id: str | None = data.get("user_left")
self.user_kicked_id: str | None = data.get("user_kicked")
self.user_banned_id: str | None = data.get("user_banned")
@property
def user_joined(self) -> Optional[TextChannel]:
@@ -94,21 +94,25 @@ class Server(Ulid):
__slots__ = ("state", "id", "name", "owner_id", "default_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "_categories", "_emojis")
def __init__(self, data: ServerPayload, state: State):
self.state = state
self.id = data["_id"]
self.name = data["name"]
self.owner_id = data["owner"]
self.description = data.get("description") or None
self.nsfw = data.get("nsfw", False)
self.system_messages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state)
self._categories = {data["id"]: Category(data, state) for data in data.get("categories", [])}
self.default_permissions = Permissions(data["default_permissions"])
self.state: State = state
self.id: str = data["_id"]
self.name: str = data["name"]
self.owner_id: str = data["owner"]
self.description: str | None = data.get("description") or None
self.nsfw: bool = data.get("nsfw", False)
self.system_messages: SystemMessages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state)
self._categories: dict[str, Category] = {data["id"]: Category(data, state) for data in data.get("categories", [])}
self.default_permissions: Permissions = Permissions(data["default_permissions"])
self.icon: Asset | None
if icon := data.get("icon"):
self.icon = Asset(icon, state)
else:
self.icon = None
self.banner: Asset | None
if banner := data.get("banner"):
self.banner = Asset(banner, state)
else:
@@ -279,11 +283,11 @@ class Server(Ulid):
await self.state.http.set_server_default_permissions(self.id, permissions.value)
async def leave_server(self):
async def leave_server(self) -> None:
"""Leaves or deletes the server"""
await self.state.http.delete_leave_server(self.id)
async def delete_server(self):
async def delete_server(self) -> None:
"""Leaves or deletes a server, alias to :meth`Server.leave_server`"""
await self.leave_server()
@@ -388,7 +392,7 @@ class Server(Ulid):
return Role(payload, name, self, self.state)
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False):
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False) -> Emoji:
"""Creates an emoji
Parameters
@@ -421,11 +425,11 @@ class ServerBan:
__slots__ = ("reason", "server", "user_id", "state")
def __init__(self, ban: Ban, state: State):
self.reason = ban.get("reason")
self.server = state.get_server(ban["_id"]["server"])
self.user_id = ban["_id"]["user"]
self.state = state
self.reason: str | None = ban.get("reason")
self.server: Server = state.get_server(ban["_id"]["server"])
self.user_id: str = ban["_id"]["user"]
self.state: State = state
async def unban(self):
async def unban(self) -> None:
"""Unbans the user"""
await self.state.http.unban_member(self.server.id, self.user_id)
+5 -5
View File
@@ -26,9 +26,9 @@ class State:
__slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages", "global_emojis", "user_id", "me")
def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int):
self.http = http
self.api_info = api_info
self.max_messages = max_messages
self.http: HttpClient = http
self.api_info: ApiInfo = api_info
self.max_messages: int = max_messages
self.me: User
@@ -114,7 +114,7 @@ class State:
raise LookupError
async def fetch_server_members(self, server_id: str):
async def fetch_server_members(self, server_id: str) -> None:
data = await self.http.fetch_members(server_id)
for user in data["users"]:
@@ -123,6 +123,6 @@ class State:
for member in data["members"]:
self.add_member(server_id, member)
async def fetch_all_server_members(self):
async def fetch_all_server_members(self) -> None:
for server_id in self.servers:
await self.fetch_server_members(server_id)
+3 -1
View File
@@ -11,6 +11,7 @@ from .permissions import Overwrite
if TYPE_CHECKING:
from .category import Category
from .embed import Embed
from .emoji import Emoji
from .file import File
from .member import Member, MemberID
@@ -65,7 +66,8 @@ class MessageEventPayload(BasePayload, Message):
class MessageUpdateData(TypedDict):
content: str
edited: int
embeds: list[Embed]
edited: Union[str, int]
class MessageUpdateEventPayload(BasePayload):
channel: str
+3 -1
View File
@@ -65,11 +65,13 @@ class Interactions(TypedDict):
reactions: NotRequired[list[str]]
restrict_reactions: NotRequired[bool]
SystemMessageContent = Union[UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent]
class Message(TypedDict):
_id: str
channel: str
author: str
content: Union[str, UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent]
content: Union[str, SystemMessageContent]
attachments: NotRequired[list[File]]
embeds: NotRequired[list[Embed]]
mentions: NotRequired[list[str]]
+2
View File
@@ -32,6 +32,8 @@ class UserRelation(TypedDict):
class User(TypedDict):
_id: str
username: str
discriminator: str
display_name: NotRequired[str]
avatar: NotRequired[File]
relations: NotRequired[list[UserRelation]]
badges: NotRequired[int]
+61 -18
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from typing import TYPE_CHECKING, NamedTuple, Optional, Union
from weakref import WeakSet
from weakref import WeakValueDictionary
from .asset import Asset, PartialAsset
from .channel import DMChannel, GroupDMChannel
@@ -18,6 +18,7 @@ if TYPE_CHECKING:
from .types import Status as StatusPayload
from .types import User as UserPayload
from .types import UserProfile as UserProfileData
from .server import Server
__all__ = ("User", "Status", "Relation", "UserProfile")
@@ -42,7 +43,11 @@ class User(Messageable, Ulid):
Attributes
-----------
id: :class:`str`
The users id
The user's id
discriminator: :class:`str`
The user's discriminator
display_name: Optional[:class:`str`]
The user's display name if they have one
bot: :class:`bool`
Whether or not the user is a bot
owner_id: Optional[:class:`str`]
@@ -64,17 +69,23 @@ class User(Messageable, Ulid):
privileged: :class:`bool`
Whether the user is privileged
"""
__flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel", "privileged")
__slots__ = (*__flattern_attributes__, "state", "_members")
__flattern_attributes__: tuple[str, ...] = ("id", "discriminator", "display_name", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel", "privileged")
__slots__: tuple[str, ...] = (*__flattern_attributes__, "state", "_members")
def __init__(self, data: UserPayload, state: State):
self.state = state
self._members: WeakSet[Member] = WeakSet() # we store all member versions of this user to avoid having to check every guild when needing to update.
self.id = data["_id"]
self.original_name = data["username"]
self.dm_channel = None
self._members: WeakValueDictionary[str, Member] = WeakValueDictionary() # we store all member versions of this user to avoid having to check every guild when needing to update.
self.id: str = data["_id"]
self.discriminator = data["discriminator"]
self.display_name = data.get("display_name")
self.original_name: str = data["username"]
self.dm_channel: DMChannel | None = None
bot = data.get("bot")
self.bot: bool
self.owner_id: str | None
if bot:
self.bot = True
self.owner_id = bot["owner"]
@@ -82,13 +93,13 @@ class User(Messageable, Ulid):
self.bot = False
self.owner_id = None
self.badges = UserBadges._from_value(data.get("badges", 0))
self.online = data.get("online", False)
self.flags = data.get("flags", 0)
self.privileged = data.get("privileged", False)
self.badges: UserBadges = UserBadges._from_value(data.get("badges", 0))
self.online: bool = data.get("online", False)
self.flags: int = data.get("flags", 0)
self.privileged: bool = data.get("privileged", False)
avatar = data.get("avatar")
self.original_avatar = Asset(avatar, state) if avatar else None
self.original_avatar: Asset | None = Asset(avatar, state) if avatar else None
relations: list[Relation] = []
@@ -96,12 +107,14 @@ class User(Messageable, Ulid):
user = state.get_user(relation["_id"])
if user:
relations.append(Relation(RelationshipType(relation["status"]), user))
self.relations = relations
self.relations: list[Relation] = relations
relationship = data.get("relationship")
self.relationship = RelationshipType(relationship) if relationship else None
self.relationship: RelationshipType | None = RelationshipType(relationship) if relationship else None
status = data.get("status")
self.status: Status | None
if status:
presence = status.get("presence")
self.status = Status(status.get("text"), PresenceType(presence) if presence else None) if status else None
@@ -177,8 +190,8 @@ class User(Messageable, Ulid):
@property
def name(self) -> str:
""":class:`str` The name the user is displaying, this includes there orginal name and masqueraded name"""
return self.masquerade_name or self.original_name
""":class:`str` The name the user is displaying, this includes (in order) their masqueraded name, display name and orginal name"""
return self.display_name or self.masquerade_name or self.original_name
@property
def avatar(self) -> Union[Asset, PartialAsset, None]:
@@ -212,7 +225,7 @@ class User(Messageable, Ulid):
# update user infomation for all members
if self.__class__ is User:
for member in self._members:
for member in self._members.values():
User._update(member, status=status, profile=profile, avatar=avatar, online=online)
async def default_avatar(self) -> bytes:
@@ -245,3 +258,33 @@ class User(Messageable, Ulid):
self.profile = UserProfile(payload.get("content"), background)
return self.profile
def to_member(self, server: Server) -> Member:
"""Gets the member instance for this user for a specific server.
Roughly equivelent to:
.. code-block:: python
member = server.get_member(user.id)
Parameters
-----------
server: :class:`Server`
The server to get the member for
Returns
--------
:class:`Member`
The member
Raises
-------
:class:`LookupError`
"""
try:
return self._members[server.id]
except IndexError:
raise LookupError from None
+2 -2
View File
@@ -11,13 +11,13 @@ from typing_extensions import ParamSpec
__all__ = ("Missing", "copy_doc", "maybe_coroutine", "get", "client_session")
class _Missing:
def __repr__(self):
def __repr__(self) -> str:
return "<Missing>"
def __bool__(self) -> Literal[False]:
return False
Missing = _Missing()
Missing: _Missing = _Missing()
T = TypeVar("T")
+55 -45
View File
@@ -4,7 +4,7 @@ import asyncio
import logging
import time
from copy import copy
from typing import TYPE_CHECKING, Callable, cast
from typing import TYPE_CHECKING, Callable, NamedTuple, cast
from . import utils
from .channel import GroupDMChannel, TextChannel, VoiceChannel
@@ -27,13 +27,17 @@ from .types import (ServerCreateEventPayload, ServerDeleteEventPayload,
ServerRoleDeleteEventPayload, ServerRoleUpdateEventPayload,
ServerUpdateEventPayload, UserRelationshipEventPayload,
UserUpdateEventPayload)
from .user import Status, UserProfile
from .user import Status, User, UserProfile
import aiohttp
try:
import ujson as json
except ImportError:
import json
use_msgpack: bool
try:
import msgpack
use_msgpack = True
@@ -46,44 +50,48 @@ if TYPE_CHECKING:
from .state import State
from .types import (AuthenticatePayload, BasePayload, MessageEventPayload,
ReadyEventPayload)
from .message import Message
class WSMessage(NamedTuple):
type: aiohttp.WSMsgType
data: str | bytes | aiohttp.WSCloseCode
__all__ = ("WebsocketHandler",)
__all__: tuple[str, ...] = ("WebsocketHandler",)
logger = logging.getLogger("revolt")
logger: logging.Logger = logging.getLogger("revolt")
class WebsocketHandler:
__slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready", "server_events")
def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State):
self.session = session
self.token = token
self.ws_url = ws_url
self.dispatch = dispatch
self.state = state
self.session: aiohttp.ClientSession = session
self.token: str = token
self.ws_url: str = ws_url
self.dispatch: Callable[..., None] = dispatch
self.state: State = state
self.websocket: aiohttp.ClientWebSocketResponse
self.loop = asyncio.get_running_loop()
self.user = None
self.ready = asyncio.Event()
self.loop: asyncio.AbstractEventLoop = asyncio.get_running_loop()
self.user: User | None = None
self.ready: asyncio.Event = asyncio.Event()
self.server_events: dict[str, asyncio.Event] = {}
async def _wait_for_server_ready(self, server_id: str):
async def _wait_for_server_ready(self, server_id: str) -> None:
if event := self.server_events.get(server_id):
await event.wait()
async def send_payload(self, payload: BasePayload):
async def send_payload(self, payload: BasePayload) -> None:
if use_msgpack:
await self.websocket.send_bytes(msgpack.packb(payload))
else:
await self.websocket.send_str(json.dumps(payload))
async def heartbeat(self):
async def heartbeat(self) -> None:
while not self.websocket.closed:
logger.info("Sending hearbeat")
await self.websocket.ping()
await asyncio.sleep(15)
async def send_authenticate(self):
async def send_authenticate(self) -> None:
payload: AuthenticatePayload = {
"type": "Authenticate",
"token": self.token
@@ -91,7 +99,7 @@ class WebsocketHandler:
await self.send_payload(payload)
async def handle_event(self, payload: BasePayload):
async def handle_event(self, payload: BasePayload) -> None:
event_type = payload["type"].lower()
logger.debug("Recieved event %s %s", event_type, payload)
try:
@@ -104,10 +112,10 @@ class WebsocketHandler:
await func(payload)
async def handle_authenticated(self, _):
async def handle_authenticated(self, _: BasePayload) -> None:
logger.info("Successfully authenticated")
async def handle_ready(self, payload: ReadyEventPayload):
async def handle_ready(self, payload: ReadyEventPayload) -> None:
for user_payload in payload["users"]:
user = self.state.add_user(user_payload)
@@ -132,7 +140,7 @@ class WebsocketHandler:
self.ready.set()
self.dispatch("ready")
async def handle_message(self, payload: MessageEventPayload):
async def handle_message(self, payload: MessageEventPayload) -> None:
if server := self.state.get_channel(payload["channel"]).server_id:
await self._wait_for_server_ready(server)
@@ -141,7 +149,7 @@ class WebsocketHandler:
self.dispatch("message", message)
async def handle_messageupdate(self, payload: MessageUpdateEventPayload):
async def handle_messageupdate(self, payload: MessageUpdateEventPayload) -> None:
self.dispatch("raw_message_update", payload)
try:
@@ -156,7 +164,7 @@ class WebsocketHandler:
self.dispatch("message_update", message)
async def handle_messagedelete(self, payload: MessageDeleteEventPayload):
async def handle_messagedelete(self, payload: MessageDeleteEventPayload) -> None:
self.dispatch("raw_message_delete", payload)
try:
@@ -172,7 +180,7 @@ class WebsocketHandler:
self.dispatch("message_delete", message)
async def handle_channelcreate(self, payload: ChannelCreateEventPayload):
async def handle_channelcreate(self, payload: ChannelCreateEventPayload) -> None:
channel = self.state.add_channel(payload)
if server_id := channel.server_id:
@@ -180,7 +188,7 @@ class WebsocketHandler:
self.dispatch("channel_create", channel)
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload):
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload) -> None:
# Revolt sends channel updates for channels we dont have permissions to see, a bug, but still can cause issues as its not in the cache
if not (channel := self.state.channels.get(payload["id"], None)):
@@ -205,7 +213,7 @@ class WebsocketHandler:
self.dispatch("channel_update", old_channel, channel)
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload):
async def handle_channeldelete(self, payload: ChannelDeleteEventPayload) -> None:
channel = self.state.channels.pop(payload["id"])
if server_id := channel.server_id:
@@ -213,7 +221,7 @@ class WebsocketHandler:
self.dispatch("channel_delete", channel)
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload):
async def handle_channelstarttyping(self, payload: ChannelStartTypingEventPayload) -> None:
channel = self.state.get_channel(payload["id"])
if server_id := channel.server_id:
@@ -223,7 +231,7 @@ class WebsocketHandler:
self.dispatch("typing_start", channel, user)
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload):
async def handle_channelstoptyping(self, payload: ChannelDeleteTypingEventPayload) -> None:
channel = self.state.get_channel(payload["id"])
if server_id := channel.server_id:
@@ -233,7 +241,7 @@ class WebsocketHandler:
self.dispatch("typing_stop", channel, user)
async def handle_serverupdate(self, payload: ServerUpdateEventPayload):
async def handle_serverupdate(self, payload: ServerUpdateEventPayload) -> None:
await self._wait_for_server_ready(payload["id"])
server = self.state.get_server(payload["id"])
@@ -255,7 +263,7 @@ class WebsocketHandler:
self.dispatch("server_update", old_server, server)
async def handle_serverdelete(self, payload: ServerDeleteEventPayload):
async def handle_serverdelete(self, payload: ServerDeleteEventPayload) -> None:
server = self.state.servers.pop(payload["id"])
for channel in server.channels:
@@ -265,7 +273,7 @@ class WebsocketHandler:
self.dispatch("server_delete", server)
async def handle_servercreate(self, payload: ServerCreateEventPayload):
async def handle_servercreate(self, payload: ServerCreateEventPayload) -> None:
for channel in payload["channels"]:
self.state.add_channel(channel)
@@ -278,7 +286,7 @@ class WebsocketHandler:
self.dispatch("server_join", server)
async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload):
async def handle_servermemberupdate(self, payload: ServerMemberUpdateEventPayload) -> None:
await self._wait_for_server_ready(payload["id"]["server"])
member = self.state.get_member(payload["id"]["server"], payload["id"]["user"])
@@ -294,7 +302,7 @@ class WebsocketHandler:
self.dispatch("member_update", old_member, member)
async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload):
async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload) -> None:
# avoid an api request if possible
if payload["user"] not in self.state.users:
user = await self.state.http.fetch_user(payload["user"])
@@ -304,7 +312,7 @@ class WebsocketHandler:
self.dispatch("member_join", member)
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload):
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload) -> None:
await self._wait_for_server_ready(payload["id"])
server = self.state.get_server(payload["id"])
@@ -313,11 +321,11 @@ class WebsocketHandler:
# remove the member from the user
user = self.state.get_user(payload["user"])
user._members.remove(member)
user._members.pop(server.id)
self.dispatch("member_leave", member)
async def handle_serverroleupdate(self, payload: ServerRoleUpdateEventPayload):
async def handle_serverroleupdate(self, payload: ServerRoleUpdateEventPayload) -> None:
server = self.state.get_server(payload["id"])
await self._wait_for_server_ready(server.id)
@@ -340,7 +348,7 @@ class WebsocketHandler:
self.dispatch("role_update", old_role, role)
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload):
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload) -> None:
server = self.state.get_server(payload["id"])
role = server._roles.pop(payload["role_id"])
@@ -348,7 +356,7 @@ class WebsocketHandler:
self.dispatch("role_delete", role)
async def handle_userupdate(self, payload: UserUpdateEventPayload):
async def handle_userupdate(self, payload: UserUpdateEventPayload) -> None:
user = self.state.get_user(payload["id"])
old_user = copy(user)
@@ -371,14 +379,14 @@ class WebsocketHandler:
self.dispatch("user_update", old_user, user)
async def handle_userrelationship(self, payload: UserRelationshipEventPayload):
async def handle_userrelationship(self, payload: UserRelationshipEventPayload) -> None:
user = self.state.get_user(payload["user"])
old_relationship = user.relationship
user.relationship = RelationshipType(payload["status"])
self.dispatch("user_relationship_update", user, old_relationship, user.relationship)
async def handle_messagereact(self, payload: MessageReactEventPayload):
async def handle_messagereact(self, payload: MessageReactEventPayload) -> None:
if server := self.state.get_channel(payload["channel_id"]).server_id:
await self._wait_for_server_ready(server)
@@ -395,7 +403,7 @@ class WebsocketHandler:
self.dispatch("reaction_add", message, user, emoji_id)
async def handle_messageunreact(self, payload: MessageUnreactEventPayload):
async def handle_messageunreact(self, payload: MessageUnreactEventPayload) -> None:
if server := self.state.get_channel(payload["channel_id"]).server_id:
await self._wait_for_server_ready(server)
@@ -411,7 +419,7 @@ class WebsocketHandler:
self.dispatch("reaction_remove", message, user, payload["emoji_id"])
async def handle_messageremovereaction(self, payload: MessageRemoveReactionEventPayload):
async def handle_messageremovereaction(self, payload: MessageRemoveReactionEventPayload) -> None:
if server := self.state.get_channel(payload["channel_id"]).server_id:
await self._wait_for_server_ready(server)
@@ -426,12 +434,12 @@ class WebsocketHandler:
self.dispatch("reaction_clear", message, users, payload["emoji_id"])
async def handle_bulkmessagedelete(self, payload: BulkMessageDeleteEventPayload):
async def handle_bulkmessagedelete(self, payload: BulkMessageDeleteEventPayload) -> None:
channel = self.state.get_channel(payload["channel"])
self.dispatch("raw_bulk_message_delete", payload)
messages = []
messages: list[Message] = []
for message_id in payload["ids"]:
if server_id := channel.server_id:
@@ -451,17 +459,19 @@ class WebsocketHandler:
self.dispatch("bulk_message_delete", messages)
async def start(self):
async def start(self) -> None:
if use_msgpack:
url = f"{self.ws_url}?format=msgpack"
else:
url = f"{self.ws_url}?format=json"
self.websocket = await self.session.ws_connect(url)
self.websocket = await self.session.ws_connect(url) # type: ignore
await self.send_authenticate()
asyncio.create_task(self.heartbeat())
async for msg in self.websocket:
msg = cast(WSMessage, msg) # aiohttp doesnt use NamedTuple so the type info is missing
if use_msgpack:
data = cast(bytes, msg.data)
+1
View File
@@ -0,0 +1 @@
def get_html_theme_path() -> str: ...