106 Commits

Author SHA1 Message Date
Zomatree efce2aedac add archived notice 2025-12-11 01:22:41 +00:00
green. ca37846df7 fix: dont compare offset-naive and offset-aware datetimes (#89) 2024-11-29 21:54:44 +00:00
Zomatree 86c07a6d6b fix: cleanup docs to include missing things 2024-10-29 17:02:40 +00:00
Zomatree a8be358339 fix: various bugs
Closes #85
Closes #82
2024-10-29 15:57:26 +00:00
Zomatree e34ca0afda Greedy support, closes #77 2024-06-11 21:33:11 +01:00
IAmTomahawkx 060e6ea32a fix: prevent new space from appearing if buffer is depleted 2024-05-29 21:00:20 +01:00
Zomatree 2dabb0dbb8 fix patch for older python versions 2024-05-29 19:11:56 +01:00
Zomatree 48e0f0fb2e better handling of optional args 2024-05-29 00:18:00 +01:00
Zomatree bc8d650659 bump aiohttp version, fixes #73 2024-05-13 13:36:12 +01:00
Zomatree d8fde3d0a6 fix converters not catching errors 2024-05-13 13:32:21 +01:00
Zomatree d4568c02ea fix create_role 2024-05-13 13:27:11 +01:00
Zomatree ff49f8e3e5 fix pyright changing generics in metaclasses 2024-05-13 13:23:27 +01:00
Zomatree eed960b9de chore: bump version 2024-02-20 17:15:01 +00:00
Zomatree 185bd00eb9 fix: typing coverage 2024-02-15 20:30:34 +00:00
Zomatree 41be16477d feat: Command cooldowns 2024-02-15 20:25:37 +00:00
Zomatree 5a22d6063d Fix typing bugs 2023-12-29 20:21:48 +00:00
Zomatree 7b95ec1b6b Simple reconnection logic 2023-12-29 20:11:55 +00:00
Zomatree dffbc1f665 Add an Object class to make mocking data with ids easier 2023-12-29 20:11:27 +00:00
Zomatree 1cfc21d4ff Undo on failed args when the default exists 2023-12-29 20:10:57 +00:00
Zomatree 9bf443e36b Fix channel.history not fetching members to go along with the messages 2023-12-29 20:09:38 +00:00
Zomatree 5d3250bcce keep waiting when wait_for check fails 2023-09-13 21:30:49 +01:00
Zomatree 1a697457b7 Fix member.edit 2023-09-11 15:53:54 +01:00
Zomatree 51c7e45821 provide autocomplete for events 2023-09-11 15:53:36 +01:00
Zomatree 1c041d1e9f re-export types 2023-09-11 15:53:05 +01:00
Zomatree 04b2bfbc40 Merge branch 'master' of github.com:revoltchat/revolt.py 2023-08-28 18:57:10 +01:00
Zomatree 02cefeab46 Fix member.edit 2023-08-28 18:57:05 +01:00
William M 5aa3c2c00a Fix wrong query param types for remove_reaction (#65) 2023-07-27 15:51:27 +01:00
Zomatree 11663d493a Fix circular imports 2023-07-24 20:00:13 +01:00
Zomatree 411a0f8203 Merge branch 'master' of github.com:revoltchat/revolt.py 2023-07-20 00:30:30 +01:00
William M 7ae2a56610 Fix the 'user' and 'remove_all' parameters not working in Message#remove_reaction() (#63)
* Fix parameters not being passed in remove_reaction

* Update revolt/http.py

---------

Co-authored-by: Angelo Kontaxis <angelokontaxis@hotmail.com>
2023-07-20 00:30:07 +01:00
MysticMia 04e63e9468 Edit function and class docstrings to fix copy-paste error and improve english (#62)
* Update context.py

* Update asset.py

spelling / grammar errors

* Fix missing 'not' in docstring in Context.server
2023-07-20 00:19:10 +01:00
Zomatree c8c11b5394 Fix nameerror
fixes #61
2023-07-20 00:15:37 +01:00
Zomatree 2749822782 Add missing type-hint 2023-07-02 19:58:21 +01:00
Zomatree 71a33c16b4 Fix mentions for users who dont share a server
fixes #60
2023-07-02 19:56:59 +01:00
Zomatree 5ee3c60ba5 Add listeners 2023-07-01 01:41:17 +01:00
Zomatree 8524ff4036 rename typevars to be more informative 2023-06-30 16:20:46 +01:00
Zomatree c97085cd55 provide original message in message_update event 2023-06-30 16:17:28 +01:00
Zomatree 85910a703a fix subcommands in cogs 2023-06-26 17:19:37 +01:00
Zomatree 042a7da85e fix help command argument parsing 2023-06-24 01:19:34 +01:00
Zomatree 1841d21bee fix subcommands being shown as base commands 2023-06-24 01:18:55 +01:00
Zomatree 69602fc812 make created_at a property 2023-06-24 01:18:01 +01:00
Zomatree a52febcb4e Fix incorrect docs, closes #59 2023-06-22 02:54:03 +01:00
Zomatree da6fb64b89 Fix trying to send a DM to yourself 2023-06-22 02:53:21 +01:00
Zomatree 072a4b1f2d fix for older python versions 2023-06-21 00:45:12 +01:00
Zomatree 826e0f8c85 Fix varies bugs, fixes #55, fixes #56 2023-06-21 00:41:15 +01:00
Zomatree a23615a5ae Fix bug for older python versions 2023-06-15 01:58:46 +01:00
Zomatree 2c30502b50 Raise an error if the token is invalid 2023-06-13 04:17:14 +01:00
Zomatree e23effd605 add missing annotations 2023-06-13 04:14:49 +01:00
Zomatree 4e70b8076c bump version 2023-06-13 03:34:42 +01:00
Zomatree 828b428f48 Add display_name and discriminator 2023-06-13 03:10:45 +01:00
Zomatree 4efaac40f6 fix FLAG_NAMES not being set 2023-06-04 15:22:21 +01:00
Zomatree 78f14fa9a0 Revert: test CI 2023-05-20 13:09:59 +01:00
Zomatree 300169d71a Add back type checking to CI 2023-05-20 13:07:21 +01:00
Zomatree 983907e0ad test CI 2023-05-20 13:04:20 +01:00
Zomatree 4c8553d5ef ignore external types 2023-05-20 03:18:19 +01:00
Zomatree f717163e17 include py.typed 2023-05-20 03:14:10 +01:00
Zomatree 60474fdeef fix verify-types parameter 2023-05-20 03:08:24 +01:00
Zomatree ccfae66e16 dont cache deps 2023-05-20 03:06:04 +01:00
Zomatree 60ff8d81c9 bring type completeness to 100% 2023-05-20 03:04:52 +01:00
TheBobBobs dfb45494ba Fix Message._update (#50)
* fix Message._update

* handle edited being int in MessageUpdate event
2023-05-19 23:16:33 +01:00
Zomatree 52c2be4e91 Fix using .server instead of .server_id
Fixes #51
2023-05-19 23:13:55 +01:00
Zomatree d5a15c0f44 fix typing error related to bad aiohttp types 2023-05-19 22:23:47 +01:00
Zomatree 4bf0530f8e clean up code 2023-05-11 07:02:39 +01:00
Zomatree 77ff484ad8 small bug fixes 2023-04-12 20:20:18 +01:00
Zomatree d909b0eb3e Merge branch 'master' of github.com:revoltchat/revolt.py 2023-04-12 16:45:11 +01:00
Zomatree 32d6b1d6e8 Add missing checks to __all__ 2023-04-12 16:41:37 +01:00
Zomatree f9ba869e75 fix fetch_emoji not returning an Emoji instance 2023-04-12 16:37:05 +01:00
Angelo Kontaxis 0e694db586 Merge pull request #49 from revoltchat/permissions
Permissions
2023-04-11 22:34:30 +01:00
Zomatree ba3f74dcd0 Add docs 2023-04-11 22:23:12 +01:00
Zomatree 625dd4eac2 finish implementation of permissions 2023-04-11 22:07:56 +01:00
Zomatree 0a6db1dc9c inital permissions calculations 2023-04-11 19:05:56 +01:00
Zomatree a672f949a1 fix servermemberjoin, closes #48 2023-03-27 00:43:12 +01:00
Zomatree df9893aaff add bulk message delete 2023-03-26 21:44:04 +01:00
Zomatree 35a4614b61 fix incorrect routes 2023-03-26 18:42:16 +01:00
Zomatree 263f99f281 use is not None in _update 2023-03-26 18:09:35 +01:00
Zomatree e99b6edee3 correctly handle role creation 2023-03-26 18:08:30 +01:00
Zomatree 720272d3cb switch from deprecated config 2023-03-26 18:08:06 +01:00
Zomatree ab5a2751df fix docs 2023-03-15 19:41:31 +00:00
Zomatree 6059e4e4dc remove bugged dep 2023-03-15 19:30:41 +00:00
Zomatree 63491ec76e update docs 2023-03-15 19:00:49 +00:00
Zomatree 76364572f1 timeout and upload_file 2023-02-22 18:49:55 +00:00
Zomatree 75f980e0f9 use our own msgpack types 2023-02-22 18:40:06 +00:00
Zomatree 8db036e131 fix CI 2023-01-02 22:10:07 +00:00
Zomatree 4d35de3ac8 qol updates 2023-01-02 22:06:25 +00:00
Zomatree ce5bf29f59 rework flags 2023-01-02 22:05:51 +00:00
Zomatree 656b9e6483 convert to hatch from poetry 2023-01-02 22:04:37 +00:00
Zomatree a92324149e fix issues stemming from bugs in the revolt api 2022-11-16 23:06:31 +00:00
Zomatree 9504c05e7b bump aiohttp dep to fix python 3.11 2022-11-16 23:05:45 +00:00
Zomatree d3f7275d1b add created_at to all models with ids 2022-11-16 23:05:06 +00:00
Zomatree 588d62e093 Use revolt.utils.client_session in example (#46) 2022-10-10 09:38:06 +01:00
Zomatree 8534458e28 fix type narrowing 2022-10-10 00:29:29 +01:00
Zomatree 95100a5b77 fix typing issues 2022-10-10 00:20:35 +01:00
Zomatree 81de5a6539 allow passing None to help_command to disable it 2022-10-10 00:19:37 +01:00
Zomatree b3e1ffa3bb reword typeguard comment to be updated 2022-10-10 00:18:24 +01:00
Zomatree c531aa649e bump doc deps 2022-10-08 14:38:39 +01:00
Zomatree 23354d030b make metaclass not generic 2022-09-26 01:26:03 +01:00
Zomatree bf8ded7588 make client generic 2022-09-26 01:16:33 +01:00
Zomatree 7ddfa2635c Fix colour always being none 2022-09-11 19:17:41 +01:00
Zomatree 7c412a056a fix race condition when joining a server 2022-08-21 13:19:45 +01:00
Zomatree 19cfa0a96c implement interactions 2022-08-18 22:51:47 +01:00
Zomatree 88af87a4e1 rename server_create event to server_join 2022-08-18 21:58:20 +01:00
Zomatree 6436901e6f move _members to a instance attr due being read only 2022-08-18 21:50:12 +01:00
Zomatree bdafa57437 fix incorrect __all__ 2022-08-18 21:04:18 +01:00
Zomatree 46c3cf893e fix varies bugs 2022-08-18 20:16:29 +01:00
Zomatree 737515f9b5 remove unnessary channel parameter 2022-08-18 03:25:51 +01:00
Zomatree 285974feef add reactions and interactions 2022-08-18 02:19:54 +01:00
62 changed files with 3123 additions and 2341 deletions
+27 -5
View File
@@ -1,16 +1,38 @@
on: [push, pull_request] on: [push, pull_request]
name: pyright name: pyright
jobs: jobs:
pyright: pyright-type-checking:
strategy:
matrix:
version: ["3.9", "3.10", "3.11"]
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v2 - uses: actions/checkout@v2
- uses: actions/setup-python@v2 - uses: actions/setup-python@v2
with: with:
python-version: '3.9' python-version: ${{ matrix.version }}
- run: pip install .[speedups] - run: pip install .[speedups,docs]
- uses: jakebailey/pyright-action@v1 - uses: jakebailey/pyright-action@v1
with: with:
lib: true python-version: ${{ matrix.version }}
python-version: 3.9
working-directory: revolt working-directory: revolt
pyright-type-completeness:
strategy:
matrix:
version: ["3.9", "3.10", "3.11"]
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- uses: actions/setup-python@v2
with:
python-version: ${{ matrix.version }}
- run: pip install .[speedups,docs]
- uses: jakebailey/pyright-action@v1
with:
python-version: ${{ matrix.version }}
working-directory: revolt
verify-types: revolt
ignore-external: true
+3 -2
View File
@@ -4,7 +4,6 @@ sphinx:
configuration: docs/conf.py configuration: docs/conf.py
python: python:
version: "3.9"
install: install:
- method: pip - method: pip
path: . path: .
@@ -12,4 +11,6 @@ python:
- docs - docs
build: build:
image: testing tools:
python: "3.9"
os: "ubuntu-22.04"
+3 -3
View File
@@ -8,13 +8,13 @@ build:
python -m build python -m build
upload: upload:
python -m twine upload dist/* -u $PYPI_USERNAME -p $PYPI_PASSWORD python -m twine upload dist/*
lint: lint:
pyright . --venv-path .venv pyright .
coverage: coverage:
pyright --lib --ignoreexternal --verifytypes revolt pyright --ignoreexternal --verifytypes revolt
docs: docs:
cd docs && make html cd docs && make html
+3 -2
View File
@@ -1,5 +1,7 @@
# Revolt.py # Revolt.py
> # This project is archived and is no longer receiving updates.
An async library to interact with the https://revolt.chat API. An async library to interact with the https://revolt.chat API.
You can join the support server [here](https://rvlt.gg/FDXER6hr) and find the library's documentation [here](https://revoltpy.readthedocs.io/en/latest/). You can join the support server [here](https://rvlt.gg/FDXER6hr) and find the library's documentation [here](https://revoltpy.readthedocs.io/en/latest/).
@@ -25,7 +27,6 @@ More examples can be found in the [examples folder](https://github.com/revoltcha
```py ```py
import revolt import revolt
import asyncio import asyncio
import aiohttp
class Client(revolt.Client): class Client(revolt.Client):
async def on_message(self, message: revolt.Message): async def on_message(self, message: revolt.Message):
@@ -33,7 +34,7 @@ class Client(revolt.Client):
await message.channel.send("hi how are you") await message.channel.send("hi how are you")
async def main(): async def main():
async with aiohttp.ClientSession() as session: async with revolt.utils.client_session() as session:
client = Client(session, "BOT TOKEN HERE") client = Client(session, "BOT TOKEN HERE")
await client.start() await client.start()
View File
+51 -106
View File
@@ -4,212 +4,151 @@ API Reference
=============== ===============
Client
~~~~~~~
.. autoclass:: Client .. autoclass:: Client
:members: :members:
:inherited-members:
Asset
~~~~~~
.. autoclass:: Asset .. autoclass:: Asset
:members: :members:
:inherited-members:
PartialAsset
~~~~~~~~~~~~~
.. autoclass:: PartialAsset .. autoclass:: PartialAsset
:members: :members:
:inherited-members:
Channel
~~~~~~~~
.. autoclass:: Channel .. autoclass:: Channel
:members: :members:
:inherited-members:
SavedMessageChannel .. autoclass:: ServerChannel
~~~~~~~~~~~~~~~~~~~~ :members:
:inherited-members:
.. autoclass:: SavedMessageChannel .. autoclass:: SavedMessageChannel
:members: :members:
:inherited-members:
DMChannel
~~~~~~~~~~
.. autoclass:: DMChannel .. autoclass:: DMChannel
:members: :members:
:inherited-members:
GroupDMChannel
~~~~~~~~~~~~~~~
.. autoclass:: GroupDMChannel .. autoclass:: GroupDMChannel
:members: :members:
:inherited-members:
TextChannel
~~~~~~~~~~~~
.. autoclass:: TextChannel .. autoclass:: TextChannel
:members: :members:
:inherited-members:
VoiceChannel
~~~~~~~~~~~~~
.. autoclass:: VoiceChannel .. autoclass:: VoiceChannel
:members: :members:
:inherited-members:
Embed
~~~~~~
.. autoclass:: Embed .. autoclass:: Embed
:members: :members:
:inherited-members:
WebsiteEmbed
~~~~~~~~~~~~~
.. autoclass:: WebsiteEmbed .. autoclass:: WebsiteEmbed
:members: :members:
:inherited-members:
ImageEmbed
~~~~~~~~~~~
.. autoclass:: ImageEmbed .. autoclass:: ImageEmbed
:members: :members:
:inherited-members:
TextEmbed
~~~~~~~~~~
.. autoclass:: TextEmbed .. autoclass:: TextEmbed
:members: :members:
:inherited-members:
NoneEmbed
~~~~~~~~~~
.. autoclass:: NoneEmbed .. autoclass:: NoneEmbed
:members: :members:
:inherited-members:
SendableEmbed
~~~~~~~~~~~~~~
.. autoclass:: SendableEmbed .. autoclass:: SendableEmbed
:members: :members:
:inherited-members:
File
~~~~~
.. autoclass:: File .. autoclass:: File
:members: :members:
:inherited-members:
Member
~~~~~~~
.. autoclass:: Member .. autoclass:: Member
:members: :members:
:inherited-members:
Message
~~~~~~~~
.. autoclass:: Message .. autoclass:: Message
:members: :members:
:inherited-members:
MessageReply
~~~~~~~~~~~~~
.. autoclass:: MessageReply .. autoclass:: MessageReply
:members: :members:
:inherited-members:
Masquerade
~~~~~~~~~~~~~
.. autoclass:: Masquerade .. autoclass:: Masquerade
:members: :members:
:inherited-members:
Messageable
~~~~~~~~~~~~
.. autoclass:: Messageable .. autoclass:: Messageable
:members: :members:
:inherited-members:
Permissions
~~~~~~~~~~~~
.. autoclass:: Permissions .. autoclass:: Permissions
:members: :members:
:inherited-members:
.. autoclass:: UserPermissions
:members:
:inherited-members:
PermissionsOverwrite
~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: PermissionsOverwrite .. autoclass:: PermissionsOverwrite
:members: :members:
:inherited-members:
Role
~~~~~
.. autoclass:: Role .. autoclass:: Role
:members: :members:
:inherited-members:
Server
~~~~~~~
.. autoclass:: Server .. autoclass:: Server
:members: :members:
:inherited-members:
ServerBan
~~~~~~~~~~
.. autoclass:: ServerBan .. autoclass:: ServerBan
:members: :members:
:inherited-members:
Category
~~~~~~~~~
.. autoclass:: Category .. autoclass:: Category
:members: :members:
:inherited-members:
SystemMessages
~~~~~~~~~~~~~~~
.. autoclass:: SystemMessages .. autoclass:: SystemMessages
:members: :members:
:inherited-members:
User
~~~~~
.. autoclass:: User .. autoclass:: User
:members: :members:
:inherited-members:
Relation
~~~~~~~~~
.. autonamedtuple:: Relation .. autonamedtuple:: Relation
Status
~~~~~~~
.. autonamedtuple:: Status .. autonamedtuple:: Status
UserBadges
~~~~~~~~~~~
.. autoclass:: UserBadges .. autoclass:: UserBadges
:members: :members:
UserProfile
~~~~~~~~~~~~
.. autoclass:: UserProfile .. autoclass:: UserProfile
:members: :members:
Invite
~~~~~~~
.. autoclass:: Invite .. autoclass:: Invite
:members: :members:
.. autoclass:: Emoji
:members:
.. autoclass:: MessageInteractions
:members:
Enums Enums
====== ------
The api uses enums to say what variant of something is, The api uses enums to say what variant of something is,
these represent those enums these represent those enums
@@ -257,7 +196,7 @@ All enums subclass `aenum.Enum`.
The user is offline or invisible The user is offline or invisible
.. attribute:: RelationshipType .. class:: RelationshipType
Specifies the relationship between two users Specifies the relationship between two users
@@ -339,7 +278,7 @@ All enums subclass `aenum.Enum`.
The embed is unknown The embed is unknown
Utils Utils
====== ------
.. currentmodule:: revolt.utils .. currentmodule:: revolt.utils
@@ -348,3 +287,9 @@ A collection a utility functions and classes to aid in making your bot
.. autofunction:: get .. autofunction:: get
.. autofunction:: client_session .. autofunction:: client_session
.. autoclass:: Ulid
:members:
.. autoclass:: Object
:members:
+1 -1
View File
@@ -24,7 +24,7 @@ import revolt
project = 'Revolt.py' project = 'Revolt.py'
copyright = '2021-present, Zomatree' copyright = '2021-present, Zomatree'
author = 'Zomatree' author = 'Zomatree'
version = "0.0.1" version = ".".join(map(str, revolt.__version__))
# -- General configuration --------------------------------------------------- # -- General configuration ---------------------------------------------------
+91 -1
View File
@@ -19,6 +19,11 @@ Command
.. autoclass:: revolt.ext.commands.Command .. autoclass:: revolt.ext.commands.Command
:members: :members:
Group
~~~~~~~~
.. autoclass:: revolt.ext.commands.Group
:members:
Cog Cog
~~~~ ~~~~
.. autoclass:: revolt.ext.commands.Cog .. autoclass:: revolt.ext.commands.Cog
@@ -28,6 +33,13 @@ command
~~~~~~~~ ~~~~~~~~
.. autodecorator:: revolt.ext.commands.command .. autodecorator:: revolt.ext.commands.command
group
~~~~~~~~
.. autodecorator:: revolt.ext.commands.group
Checks
-------
check check
~~~~~~ ~~~~~~
.. autodecorator:: revolt.ext.commands.check .. autodecorator:: revolt.ext.commands.check
@@ -40,9 +52,62 @@ is_server_owner
~~~~~~~~~~~~~~~~ ~~~~~~~~~~~~~~~~
.. autodecorator:: revolt.ext.commands.is_server_owner .. autodecorator:: revolt.ext.commands.is_server_owner
has_permissions
~~~~~~~~~~~~~~~~
.. autodecorator:: revolt.ext.commands.has_permissions
has_channel_permissions
~~~~~~~~~~~~~~~~~~~~~~~~
.. autodecorator:: revolt.ext.commands.has_channel_permissions
Converters
-----------
IntConverter
~~~~~~~~~~~~~
Converts the parameter to an int
BoolConverter
~~~~~~~~~~~~~~
Converts the parameter to a bool
CategoryConverter
~~~~~~~~~~~~~~~~~~
Converts the parameter to a category
UserConverter
~~~~~~~~~~~~~~
Converts the parameter to a category
MemberConverter
~~~~~~~~~~~~~~~~
Converts the parameter to a category
ChannelConverter
~~~~~~~~~~~~~~~~~
Converts the parameter to a category
Greedy
~~~~~~
Converts the parameter to a greedy parameter which will take as many arguments which convert successfully.
Allows you to have var-args in the middle of a signature.
Help Commands
--------------
HelpCommand
~~~~~~~~~~~~
.. autoclass:: revolt.ext.commands.HelpCommand
:members:
DefaultHelpCommand
~~~~~~~~~~~~~~~~~~~
.. autoclass:: revolt.ext.commands.DefaultHelpCommand
:members:
Exceptions Exceptions
=========== -----------
CommandError CommandError
~~~~~~~~~~~~~ ~~~~~~~~~~~~~
@@ -79,6 +144,11 @@ ServerOnly
.. autoexception:: revolt.ext.commands.ServerOnly .. autoexception:: revolt.ext.commands.ServerOnly
:members: :members:
MissingPermissionsError
~~~~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.MissingPermissionsError
:members:
ConverterError ConverterError
~~~~~~~~~~~~~~~ ~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.ConverterError .. autoexception:: revolt.ext.commands.ConverterError
@@ -99,6 +169,11 @@ CategoryConverterError
.. autoexception:: revolt.ext.commands.CategoryConverterError .. autoexception:: revolt.ext.commands.CategoryConverterError
:members: :members:
ChannelConverterError
~~~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.ChannelConverterError
:members:
UserConverterError UserConverterError
~~~~~~~~~~~~~~~~~~~ ~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.UserConverterError .. autoexception:: revolt.ext.commands.UserConverterError
@@ -108,3 +183,18 @@ MemberConverterError
~~~~~~~~~~~~~~~~~~~~~ ~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.MemberConverterError .. autoexception:: revolt.ext.commands.MemberConverterError
:members: :members:
UnionConverterError
~~~~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.UnionConverterError
:members:
MissingSetup
~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.MissingSetup
:members:
CommandOnCooldown
~~~~~~~~~~~~~~~~~~
.. autoexception:: revolt.ext.commands.CommandOnCooldown
:members:
Generated
-1298
View File
File diff suppressed because it is too large Load Diff
+51 -28
View File
@@ -1,12 +1,11 @@
[tool.poetry] [project]
name = "revolt.py" name = "revolt.py"
version = "0.1.9" dynamic = ["version"]
description = "Python wrapper for the revolt.chat API" description = "Python wrapper for the revolt.chat API"
authors = ["Zomatee <me@zomatree.live>"] requires-python = ">=3.9"
license = "MIT" license = "MIT"
readme = "README.md" readme = "README.md"
homepage = "https://github.com/revoltchat/revolt.py" keywords = ["wrapper", "async", "api", "websockets", "http"]
documentation = "https://revoltpy.readthedocs.io/en/latest/"
classifiers = [ classifiers = [
"Development Status :: 4 - Beta", "Development Status :: 4 - Beta",
"Intended Audience :: Developers", "Intended Audience :: Developers",
@@ -16,31 +15,55 @@ classifiers = [
"Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3 :: Only", "Programming Language :: Python :: 3 :: Only",
] ]
keywords = ["wrapper", "async", "api", "websockets", "http"] dependencies = [
packages = [ "aiohttp==3.10.*",
{ include = "revolt" } "ulid-py==1.1.*",
"aenum==3.1.*",
"typing_extensions>=4.4.0"
] ]
[tool.poetry.dependencies] [project.optional-dependencies]
python = "^3.9" speedups = [
aiohttp = "3.7.4" "ujson==5.1.*",
ulid-py = "1.1.0" "msgpack==1.0.*"
aenum = "3.1.8" ]
typing-extensions = "4.1.1" docs = [
ujson = { version = "5.1.0", optional = true } "Sphinx==5.2.*",
msgpack = { version = "", optional = true } "sphinx-nameko-theme==0.0.*",
Sphinx = { version = "4.3.2", optional = true } "sphinx-toolbox==3.2.*",
sphinx-nameko-theme = { version = "0.0.3", optional = true } "setuptools==65.4.*"
sphinx-toolbox = { version = "2.15.2", optional = true } ]
[tool.poetry.extras] [project.urls]
speedups = ["ujson", "aiohttp[speedups]", "msgpack"] Homepage = "https://github.com/revoltchat/revolt.py"
docs = ["Sphinx", "sphinx-nameko-theme", "sphinx-toolbox"] Documentation = "https://revoltpy.readthedocs.io/en/latest/"
"Source Code" = "https://github.com/revoltchat/revolt.py"
"Bug Tracker" = "https://github.com/revoltchat/revolt.py/issues"
[[project.authors]]
name = "Zomatree"
email = "me@zomatree.live"
[tool.hatch.version]
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] [build-system]
requires = ["poetry-core>=1.0.0"] requires = ["hatchling"]
build-backend = "poetry.core.masonry.api" build-backend = "hatchling.build"
[tool.poetry.urls]
"Bug Tracker" = "https://github.com/revoltchat/revolt.py/issues"
Source = "https://github.com/revoltchat/revolt.py/"
+4 -2
View File
@@ -1,9 +1,11 @@
from . import utils from . import utils as utils
from . import types as types
from .asset import * from .asset import *
from .category import * from .category import *
from .channel import * from .channel import *
from .client import * from .client import *
from .embed import * from .embed import *
from .emoji import *
from .enums import * from .enums import *
from .errors import * from .errors import *
from .file import * from .file import *
@@ -17,4 +19,4 @@ from .role import *
from .server import * from .server import *
from .user import * from .user import *
__version__ = (0, 1, 9) __version__ = "0.2.0"
+25 -26
View File
@@ -4,7 +4,7 @@ import mimetypes
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .enums import AssetType from .enums import AssetType
from .errors import AutumnDisabled from .utils import Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from io import IOBase from io import IOBase
@@ -15,7 +15,7 @@ if TYPE_CHECKING:
__all__ = ("Asset", "PartialAsset") __all__ = ("Asset", "PartialAsset")
class Asset: class Asset(Ulid):
"""Represents a file on revolt """Represents a file on revolt
Attributes Attributes
@@ -23,7 +23,7 @@ class Asset:
id: :class:`str` id: :class:`str`
The id of the asset The id of the asset
tag: :class:`str` tag: :class:`str`
The tag of the asset, this corrasponds to where the asset is used The tag of the asset, this corresponds to where the asset is used
size: :class:`int` size: :class:`int`
Amount of bytes in the file Amount of bytes in the file
filename: :class:`str` filename: :class:`str`
@@ -37,19 +37,21 @@ class Asset:
type: :class:`AssetType` type: :class:`AssetType`
The type of asset it is The type of asset it is
url: :class:`str` url: :class:`str`
The assets url The asset's url
""" """
__slots__ = ("state", "id", "tag", "size", "filename", "content_type", "width", "height", "type", "url") __slots__ = ("state", "id", "tag", "size", "filename", "content_type", "width", "height", "type", "url")
def __init__(self, data: FilePayload, state: State): def __init__(self, data: FilePayload, state: State):
self.state = state self.state: State = state
self.id = data['_id'] self.id: str = data['_id']
self.tag = data['tag'] self.tag: str = data['tag']
self.size = data['size'] self.size: int = data['size']
self.filename = data['filename'] self.filename: str = data['filename']
metadata = data['metadata'] 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 if metadata["type"] == "Image" or metadata["type"] == "Video": # cannot use `in` because type narrowing will not happen
self.height = metadata["height"] self.height = metadata["height"]
@@ -58,23 +60,23 @@ class Asset:
self.height = None self.height = None
self.width = None self.width = None
self.content_type = data["content_type"] self.content_type: str | None = data["content_type"]
self.type = AssetType(metadata["type"]) self.type: AssetType = AssetType(metadata["type"])
base_url = self.state.api_info["features"]["autumn"]["url"] 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: async def read(self) -> bytes:
"""Reads the files content into bytes""" """Reads the files content into bytes"""
return await self.state.http.request_file(self.url) 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 """Reads the files content and saves it to a file
Parameters Parameters
----------- -----------
fp: IOBase fp: IOBase
The file to write too. The file to write to
""" """
fp.write(await self.read()) fp.write(await self.read())
@@ -85,8 +87,6 @@ class PartialAsset(Asset):
----------- -----------
id: :class:`str` id: :class:`str`
The id of the asset, this will always be ``"0"`` 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` size: :class:`int`
Amount of bytes in the file, this will always be ``0`` Amount of bytes in the file, this will always be ``0``
filename: :class:`str` filename: :class:`str`
@@ -102,13 +102,12 @@ class PartialAsset(Asset):
""" """
def __init__(self, url: str, state: State): def __init__(self, url: str, state: State):
self.state = state self.state: State = state
self.id = "0" self.id: str = "0"
self.tag = None self.size: int = 0
self.size = 0 self.filename: str = ""
self.filename = "" self.height: int | None = None
self.height = None self.width: int | None = None
self.width = None self.content_type: str | None = mimetypes.guess_extension(url)
self.content_type = mimetypes.guess_extension(url) self.type: AssetType = AssetType.file
self.type = AssetType.file self.url: str = url
self.url = url
+7 -5
View File
@@ -2,6 +2,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .utils import Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from .channel import Channel from .channel import Channel
from .state import State from .state import State
@@ -9,7 +11,7 @@ if TYPE_CHECKING:
__all__ = ("Category",) __all__ = ("Category",)
class Category: class Category(Ulid):
"""Represents a category in a server that stores channels. """Represents a category in a server that stores channels.
Attributes Attributes
@@ -23,10 +25,10 @@ class Category:
""" """
def __init__(self, data: CategoryPayload, state: State): def __init__(self, data: CategoryPayload, state: State):
self.state = state self.state: State = state
self.name = data["title"] self.name: str = data["title"]
self.id = data["id"] self.id: str = data["id"]
self.channel_ids = data["channels"] self.channel_ids: list[str] = data["channels"]
@property @property
def channels(self) -> list[Channel]: def channels(self) -> list[Channel]:
+110 -52
View File
@@ -1,14 +1,12 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Literal, Optional, Union from typing import TYPE_CHECKING, Any, Optional, Union
from revolt.utils import Missing
from .asset import Asset from .asset import Asset
from .enums import ChannelType from .enums import ChannelType
from .messageable import Messageable from .messageable import Messageable
from .permissions import Permissions, PermissionsOverwrite from .permissions import Permissions, PermissionsOverwrite
from .utils import Missing from .utils import Missing, Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from .message import Message from .message import Message
@@ -17,15 +15,15 @@ if TYPE_CHECKING:
from .state import State from .state import State
from .types import Channel as ChannelPayload from .types import Channel as ChannelPayload
from .types import DMChannel as DMChannelPayload from .types import DMChannel as DMChannelPayload
from .types import GroupDMChannel as GroupDMChannelPayload
from .types import SavedMessages as SavedMessagesPayload
from .types import TextChannel as TextChannelPayload
from .types import VoiceChannel as VoiceChannelPayload
from .types import GuildChannel as GuildChannelPayload
from .types import File as FilePayload from .types import File as FilePayload
from .types import GroupDMChannel as GroupDMChannelPayload
from .types import Overwrite as OverwritePayload from .types import Overwrite as OverwritePayload
from .types import SavedMessages as SavedMessagesPayload
from .types import ServerChannel as ServerChannelPayload
from .types import TextChannel as TextChannelPayload
from .user import User
__all__ = ("DMChannel", "GroupDMChannel", "SavedMessageChannel", "TextChannel", "VoiceChannel", "Channel") __all__ = ("DMChannel", "GroupDMChannel", "SavedMessageChannel", "TextChannel", "VoiceChannel", "Channel", "ServerChannel")
class EditableChannel: class EditableChannel:
__slots__ = () __slots__ = ()
@@ -33,7 +31,7 @@ class EditableChannel:
state: State state: State
id: str id: str
async def edit(self, **kwargs): async def edit(self, **kwargs: Any) -> None:
"""Edits the channel """Edits the channel
Passing ``None`` to the parameters that accept it will remove them. Passing ``None`` to the parameters that accept it will remove them.
@@ -51,12 +49,13 @@ class EditableChannel:
nsfw: bool nsfw: bool
Sets whether the channel is nsfw or not Sets whether the channel is nsfw or not
""" """
remove: list[str] = []
if kwargs.get("icon", Missing) == None: if kwargs.get("icon", Missing) == None:
remove = "Icon" remove.append("Icon")
elif kwargs.get("description", Missing) == None:
remove = "Description" if kwargs.get("description", Missing) == None:
else: remove.append("Description")
remove = None
if icon := kwargs.get("icon"): if icon := kwargs.get("icon"):
asset = await self.state.http.upload_file(icon, "icons") asset = await self.state.http.upload_file(icon, "icons")
@@ -67,7 +66,7 @@ class EditableChannel:
await self.state.http.edit_channel(self.id, remove, kwargs) await self.state.http.edit_channel(self.id, remove, kwargs)
class Channel: class Channel(Ulid):
"""Base class for all channels """Base class for all channels
Attributes Attributes
@@ -82,26 +81,32 @@ class Channel:
__slots__ = ("state", "id", "channel_type", "server_id") __slots__ = ("state", "id", "channel_type", "server_id")
def __init__(self, data: ChannelPayload, state: State): def __init__(self, data: ChannelPayload, state: State):
self.state = state self.state: State = state
self.id = data["_id"] self.id: str = data["_id"]
self.channel_type = ChannelType(data["channel_type"]) self.channel_type: ChannelType = ChannelType(data["channel_type"])
self.server_id: Optional[str] = None self.server_id: Optional[str] = None
async def _get_channel_id(self) -> str: async def _get_channel_id(self) -> str:
return self.id return self.id
def _update(self, **_): def _update(self, **_: Any) -> None:
pass pass
async def delete(self): async def delete(self) -> None:
"""Deletes or closes the channel""" """Deletes or closes the channel"""
await self.state.http.close_channel(self.id) await self.state.http.close_channel(self.id)
@property @property
def server(self) -> Server: 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: if not self.server_id:
raise IndexError raise LookupError
return self.state.get_server(self.server_id) return self.state.get_server(self.server_id)
@@ -125,11 +130,28 @@ class DMChannel(Channel, Messageable):
The id of the last message in this channel, if any The id of the last message in this channel, if any
""" """
__slots__ = ("last_message_id",) __slots__ = ("last_message_id", "recipient_ids")
def __init__(self, data: DMChannelPayload, state: State): def __init__(self, data: DMChannelPayload, state: State):
super().__init__(data, state) super().__init__(data, state)
self.last_message_id = data.get("last_message_id")
self.recipient_ids: list[str] = data["recipients"]
self.last_message_id: str | None = data.get("last_message_id")
@property
def recipients(self) -> tuple[User, User]:
a, b = self.recipient_ids
return (self.state.get_user(a), self.state.get_user(b))
@property
def recipient(self) -> User:
if self.recipient_ids[0] != self.state.user_id:
user_id = self.recipient_ids[0]
else:
user_id = self.recipient_ids[1]
return self.state.get_user(user_id)
@property @property
def last_message(self) -> Message: def last_message(self) -> Message:
@@ -166,33 +188,43 @@ class GroupDMChannel(Channel, Messageable, EditableChannel):
The id of the last message in this channel, if any The id of the last message in this channel, if any
""" """
__slots__ = ("recipients", "name", "owner", "permissions", "icon", "description", "last_message_id") __slots__ = ("recipient_ids", "name", "owner_id", "permissions", "icon", "description", "last_message_id")
def __init__(self, data: GroupDMChannelPayload, state: State): def __init__(self, data: GroupDMChannelPayload, state: State):
super().__init__(data, state) super().__init__(data, state)
self.recipients = [state.get_user(user_id) for user_id in data["recipients"]] self.recipient_ids: list[str] = data["recipients"]
self.name = data["name"] self.name: str = data["name"]
self.owner = state.get_user(data["owner"]) self.owner_id: str = data["owner"]
self.description: Optional[str] = data.get("description") self.description: str | None = data.get("description")
self.last_message_id = data.get("last_message_id") self.last_message_id: str | None = data.get("last_message_id")
self.icon: Asset | None
if icon := data.get("icon"): if icon := data.get("icon"):
self.icon = Asset(icon, state) self.icon = Asset(icon, state)
else: else:
self.icon = None 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: if name is not None:
self.name = name self.name = name
if recipients: if recipients is not None:
self.recipients = [self.state.get_user(user_id) for user_id in recipients] self.recipient_ids = recipients
if description: if description is not None:
self.description = description self.description = description
@property
def recipients(self) -> list[User]:
return [self.state.get_user(user_id) for user_id in self.recipient_ids]
@property
def owner(self) -> User:
return self.state.get_user(self.owner_id)
async def set_default_permissions(self, permissions: Permissions) -> None: async def set_default_permissions(self, permissions: Permissions) -> None:
"""Sets the default permissions for a group. """Sets the default permissions for a group.
Parameters Parameters
@@ -216,16 +248,31 @@ class GroupDMChannel(Channel, Messageable, EditableChannel):
return self.state.get_message(self.last_message_id) return self.state.get_message(self.last_message_id)
class GuildChannel(Channel): class ServerChannel(Channel):
def __init__(self, data: GuildChannelPayload, state: State): """Base class for all guild channels
Attributes
-----------
server_id: :class:`str`
The id of the server this text channel belongs to
name: :class:`str`
The name of the text channel
description: Optional[:class:`str`]
The description of the channel, if any
nsfw: bool
Sets whether the channel is nsfw or not
default_permissions: :class:`ChannelPermissions`
The default permissions for all users in the text channel
"""
def __init__(self, data: ServerChannelPayload, state: State):
super().__init__(data, state) super().__init__(data, state)
self.server_id = data["server"] self.server_id: Optional[str] = data["server"]
self.name = data["name"] self.name: str = data["name"]
self.description: Optional[str] = data.get("description") self.description: Optional[str] = data.get("description")
self.nsfw = data.get("nsfw", False) self.nsfw: bool = data.get("nsfw", False)
self.active = False self.active: bool = False
self.default_permissions = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0})) self.default_permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("default_permissions", {"a": 0, "d": 0}))
permissions: dict[str, PermissionsOverwrite] = {} permissions: dict[str, PermissionsOverwrite] = {}
@@ -233,7 +280,10 @@ class GuildChannel(Channel):
overwrite = PermissionsOverwrite._from_overwrite(overwrite_data) overwrite = PermissionsOverwrite._from_overwrite(overwrite_data)
permissions[role_name] = overwrite permissions[role_name] = overwrite
self.permissions = permissions self.permissions: dict[str, PermissionsOverwrite] = permissions
self.icon: Asset | None
if icon := data.get("icon"): if icon := data.get("icon"):
self.icon = Asset(icon, state) self.icon = Asset(icon, state)
else: else:
@@ -241,9 +291,10 @@ class GuildChannel(Channel):
async def set_default_permissions(self, permissions: PermissionsOverwrite) -> None: async def set_default_permissions(self, permissions: PermissionsOverwrite) -> None:
"""Sets the default permissions for the channel. """Sets the default permissions for the channel.
Parameters Parameters
----------- -----------
permissions: :class:`ChannelPermissions` permissions: :class:`PermissionsOverwrite`
The new default channel permissions The new default channel permissions
""" """
allow, deny = permissions.to_pair() allow, deny = permissions.to_pair()
@@ -251,8 +302,11 @@ class GuildChannel(Channel):
async def set_role_permissions(self, role: Role, permissions: PermissionsOverwrite) -> None: async def set_role_permissions(self, role: Role, permissions: PermissionsOverwrite) -> None:
"""Sets the permissions for a role in the channel. """Sets the permissions for a role in the channel.
Parameters Parameters
----------- -----------
role: :class:`Role`
The role to set permissions for
permissions: :class:`ChannelPermissions` permissions: :class:`ChannelPermissions`
The new channel permissions The new channel permissions
""" """
@@ -267,7 +321,7 @@ class GuildChannel(Channel):
if description is not None: if description is not None:
self.description = description self.description = description
if icon: if icon is not None:
self.icon = Asset(icon, self.state) self.icon = Asset(icon, self.state)
if nsfw is not None: if nsfw is not None:
@@ -286,11 +340,13 @@ class GuildChannel(Channel):
self.permissions = permissions self.permissions = permissions
if default_permissions is not None: if default_permissions is not None:
self.default_permissions = default_permissions self.default_permissions = PermissionsOverwrite._from_overwrite(default_permissions)
class TextChannel(GuildChannel, Messageable, EditableChannel): class TextChannel(ServerChannel, Messageable, EditableChannel):
"""A text channel """A text channel
Subclasses :class:`ServerChannel` and :class:`Messageable`
Attributes Attributes
----------- -----------
name: :class:`str` name: :class:`str`
@@ -314,7 +370,7 @@ class TextChannel(GuildChannel, Messageable, EditableChannel):
def __init__(self, data: TextChannelPayload, state: State): def __init__(self, data: TextChannelPayload, state: State):
super().__init__(data, 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: async def _get_channel_id(self) -> str:
return self.id return self.id
@@ -333,9 +389,11 @@ class TextChannel(GuildChannel, Messageable, EditableChannel):
return self.state.get_message(self.last_message_id) return self.state.get_message(self.last_message_id)
class VoiceChannel(GuildChannel, EditableChannel): class VoiceChannel(ServerChannel, EditableChannel):
"""A voice channel """A voice channel
Subclasses :class:`ServerChannel`
Attributes Attributes
----------- -----------
name: :class:`str` name: :class:`str`
+250 -34
View File
@@ -2,18 +2,23 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
from typing import TYPE_CHECKING, Any, Callable, Optional, Union, cast from typing import TYPE_CHECKING, Any, Callable, Coroutine, Literal, Optional, TypeVar, Union, cast, overload
from typing_extensions import ParamSpec
import aiohttp import aiohttp
from .errors import RevoltError
from .channel import (DMChannel, GroupDMChannel, SavedMessageChannel, from .channel import (DMChannel, GroupDMChannel, SavedMessageChannel,
TextChannel, VoiceChannel, channel_factory) TextChannel, VoiceChannel, channel_factory)
from .http import HttpClient from .http import HttpClient
from .invite import Invite from .invite import Invite
from .message import Message from .message import Message
from .state import State from .state import State
from .utils import Missing from .utils import Missing, Ulid
from .websocket import WebsocketHandler from .websocket import WebsocketHandler
from .emoji import Emoji
from .server import Server
from .user import User
try: try:
import ujson as json import ujson as json
@@ -22,14 +27,17 @@ except ImportError:
if TYPE_CHECKING: if TYPE_CHECKING:
from .channel import Channel from .channel import Channel
from .server import Server from .file import File
from .types import ApiInfo from .types import ApiInfo
from .user import User
import revolt
__all__ = ("Client",) __all__ = ("Client",)
logger = logging.getLogger("revolt") logger: logging.Logger = logging.getLogger("revolt")
P = ParamSpec("P")
R = TypeVar("R")
class Client: class Client:
"""The client for interacting with revolt """The client for interacting with revolt
@@ -44,25 +52,28 @@ class Client:
The api url for the revolt instance you are connecting to, by default it uses the offical instance hosted at revolt.chat The api url for the revolt instance you are connecting to, by default it uses the offical instance hosted at revolt.chat
max_messages: :class:`int` max_messages: :class:`int`
The max amount of messages stored in the cache, by default this is 5k The max amount of messages stored in the cache, by default this is 5k
bot: :class:`bool`
Denotes whether the account used is a bot account or user account, by default this it assumes a bot account
""" """
def __init__(self, session: aiohttp.ClientSession, token: str, *, api_url: str = "https://api.revolt.chat", max_messages: int = 5000, bot: bool = True): def __init__(self, session: aiohttp.ClientSession, token: str, *, api_url: str = "https://api.revolt.chat", max_messages: int = 5000, bot: bool = True):
self.session = session self.session: aiohttp.ClientSession = session
self.token = token self.token: str = token
self.api_url = api_url self.api_url: str = api_url
self.max_messages = max_messages self.max_messages: int = max_messages
self.bot = bot self.bot: bool = bot
self.api_info: ApiInfo self.api_info: ApiInfo
self.http: HttpClient self.http: HttpClient
self.state: State self.state: State
self.websocket: WebsocketHandler self.websocket: WebsocketHandler
self.listeners: dict[str, list[tuple[Callable[..., bool], asyncio.Future[Any]]]] = {} self.temp_listeners: dict[str, list[tuple[Callable[..., bool], asyncio.Future[Any]]]] = {}
self.listeners: dict[str, list[Callable[..., Coroutine[Any, Any, Any]]]] = {}
super().__init__() 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. """Dispatch an event, this is typically used for testing and internals.
Parameters Parameters
@@ -72,22 +83,33 @@ class Client:
args: :class:`Any` args: :class:`Any`
The arguments passed to the event The arguments passed to the event
""" """
for check, future in self.listeners.pop(event, []):
if check(*args):
if len(args) == 1:
future.set_result(args[0])
else:
future.set_result(args)
func = getattr(self, f"on_{event}", None) if temp_listeners := self.temp_listeners.get(event, None):
if func: for check, future in temp_listeners:
if check(*args):
if len(args) == 1:
future.set_result(args[0])
else:
future.set_result(args)
self.temp_listeners[event] = [(c, f) for c, f in temp_listeners if not f.done()]
for listener in self.listeners.get(event, []):
asyncio.create_task(listener(*args))
if func := getattr(self, f"on_{event}", None):
asyncio.create_task(func(*args)) asyncio.create_task(func(*args))
async def get_api_info(self) -> ApiInfo: async def get_api_info(self) -> ApiInfo:
async with self.session.get(self.api_url) as resp: async with self.session.get(self.api_url) as resp:
return json.loads(await resp.text()) text = await resp.text()
async def start(self): try:
return json.loads(text)
except:
raise RevoltError(f"Cant fetch api info:\n{text}")
async def start(self, *, reconnect: bool = True) -> None:
"""Starts the client""" """Starts the client"""
api_info = await self.get_api_info() api_info = await self.get_api_info()
@@ -95,7 +117,11 @@ class Client:
self.http = HttpClient(self.session, self.token, self.api_url, self.api_info, self.bot) self.http = HttpClient(self.session, self.token, self.api_url, self.api_info, self.bot)
self.state = State(self.http, api_info, self.max_messages) self.state = State(self.http, api_info, self.max_messages)
self.websocket = WebsocketHandler(self.session, self.token, api_info["ws"], self.dispatch, self.state) self.websocket = WebsocketHandler(self.session, self.token, api_info["ws"], self.dispatch, self.state)
await self.websocket.start()
await self.websocket.start(reconnect)
async def stop(self) -> None:
await self.websocket.websocket.close()
def get_user(self, id: str) -> User: def get_user(self, id: str) -> User:
"""Gets a user from the cache """Gets a user from the cache
@@ -168,10 +194,64 @@ class Client:
check = lambda *_: True check = lambda *_: True
future = asyncio.get_running_loop().create_future() future = asyncio.get_running_loop().create_future()
self.listeners.setdefault(event, []).append((check, future)) self.temp_listeners.setdefault(event, []).append((check, future))
return await asyncio.wait_for(future, timeout) return await asyncio.wait_for(future, timeout)
def listen(self, name: str | None = None) -> Callable[[Callable[P, Coroutine[Any, Any, R]]], Callable[P, Coroutine[Any, Any, R]]]:
"""Registers a listener for an event, multiple listeners can be registered to the same event without conflict
Parameters
-----------
name: Optional[:class:`str`]
The name of the event to register this under, this defaults to the function's name
"""
def inner(func: Callable[P, Coroutine[Any, Any, R]]) -> Callable[P, Coroutine[Any, Any, R]]:
nonlocal name
if not name:
if not func.__name__.startswith("on_"):
raise RevoltError("listener name must begin with `on_`")
name = func.__name__[3:]
self.listeners.setdefault(name, []).append(func)
return func
return inner
@overload
def remove_listener(self, func: Callable[P, Coroutine[Any, Any, R]], *, event: str = ...) -> Callable[..., Coroutine[Any, Any, R]] | None:
...
@overload
def remove_listener(self, func: Callable[P, Coroutine[Any, Any, Any]], *, event: None = ...) -> None:
...
def remove_listener(self, func: Callable[P, Coroutine[Any, Any, R]], *, event: str | None = None) -> Callable[..., Coroutine[Any, Any, R]] | None:
"""Removes a listener registered, if the `event` parameter is passed, the listener will only be removed from that event, this can be used if the same listener is registed to multiple events at once.
Parameters
-----------
func: Callable
The function for the listener to be removed
event: Optional[:class:`str`]
The name of the event to remove this from, passing `None` will make this remove the listener from all events this is registered under
"""
if event is None:
for listeners in self.listeners.values():
try:
listeners.remove(func)
except ValueError:
pass
else:
try:
self.listeners[event].remove(func)
return func
except ValueError:
pass
@property @property
def user(self) -> User: def user(self) -> User:
""":class:`User` the user corrasponding to the client""" """:class:`User` the user corrasponding to the client"""
@@ -190,6 +270,10 @@ class Client:
"""list[:class:'Server'] All servers the client can see""" """list[:class:'Server'] All servers the client can see"""
return list(self.state.servers.values()) return list(self.state.servers.values())
@property
def global_emojis(self) -> list[Emoji]:
return self.state.global_emojis
async def fetch_user(self, user_id: str) -> User: async def fetch_user(self, user_id: str) -> User:
"""Fetchs a user """Fetchs a user
@@ -292,7 +376,7 @@ class Client:
raise LookupError raise LookupError
async def edit_self(self, **kwargs): async def edit_self(self, **kwargs: Any) -> None:
"""Edits the client's own user """Edits the client's own user
Parameters Parameters
@@ -302,13 +386,13 @@ class Client:
""" """
if kwargs.get("avatar", Missing) is None: if kwargs.get("avatar", Missing) is None:
del kwargs["avatar"] del kwargs["avatar"]
remove = "Avatar" remove = ["Avatar"]
else: else:
remove = None remove = None
await self.state.http.edit_self(remove, kwargs) await self.state.http.edit_self(remove, kwargs)
async def edit_status(self, **kwargs): async def edit_status(self, **kwargs: Any) -> None:
"""Edits the client's own status """Edits the client's own status
Parameters Parameters
@@ -320,7 +404,7 @@ class Client:
""" """
if kwargs.get("text", Missing) is None: if kwargs.get("text", Missing) is None:
del kwargs["text"] del kwargs["text"]
remove = "StatusText" remove = ["StatusText"]
else: else:
remove = None remove = None
@@ -329,7 +413,7 @@ class Client:
await self.state.http.edit_self(remove, {"status": kwargs}) await self.state.http.edit_self(remove, {"status": kwargs})
async def edit_profile(self, **kwargs): async def edit_profile(self, **kwargs: Any) -> None:
"""Edits the client's own profile """Edits the client's own profile
Parameters Parameters
@@ -339,13 +423,145 @@ class Client:
background: Optional[:class:`File`] background: Optional[:class:`File`]
The new background for the profile, passing in ``None`` will remove the profile background The new background for the profile, passing in ``None`` will remove the profile background
""" """
remove: list[str] = []
if kwargs.get("content", Missing) is None: if kwargs.get("content", Missing) is None:
del kwargs["content"] del kwargs["content"]
remove = "ProfileContent" remove.append("ProfileContent")
elif kwargs.get("background", Missing) is None:
if kwargs.get("background", Missing) is None:
del kwargs["background"] del kwargs["background"]
remove = "ProfileBackground" remove.append("ProfileBackground")
else:
remove = None
await self.state.http.edit_self(remove, {"profile": kwargs}) await self.state.http.edit_self(remove, {"profile": kwargs})
async def fetch_emoji(self, emoji_id: str) -> Emoji:
"""Fetches an emoji
Parameters
-----------
emoji_id: str
The id of the emoji
Returns
--------
:class:`Emoji`
The emoji with the corrasponding id
"""
emoji = await self.state.http.fetch_emoji(emoji_id)
return Emoji(emoji, self.state)
async def upload_file(self, file: File, tag: Literal['attachments', 'avatars', 'backgrounds', 'icons', 'banners', 'emojis']) -> Ulid:
"""Uploads a file to revolt
Parameters
-----------
file: :class:`File`
The file to upload
tag: :class:`str`
The type of file to upload, this should a string of either `'attachments'`, `'avatars'`, `'backgrounds'`, `'icons'`, `'banners'` or `'emojis'`
Returns
--------
:class:`Ulid`
The id of the file that was uploaded
"""
asset = await self.http.upload_file(file, tag)
ulid = Ulid()
ulid.id = asset["id"]
return ulid
# events
async def on_ready(self) -> None:
pass
async def on_message(self, message: revolt.Message) -> None:
pass
async def on_raw_message_update(self, payload: revolt.types.MessageUpdateEventPayload) -> None:
pass
async def on_message_update(self, before: revolt.Message, after: revolt.Message) -> None:
pass
async def on_raw_message_delete(self, payload: revolt.types.MessageDeleteEventPayload) -> None:
pass
async def on_message_delete(self, message: revolt.Message) -> None:
pass
async def on_channel_create(self, channel: revolt.Channel) -> None:
pass
async def on_channel_update(self, before: revolt.Channel, after: revolt.Channel) -> None:
pass
async def on_channel_delete(self, channel: revolt.Channel) -> None:
pass
async def on_typing_start(self, channel: revolt.Channel, user: revolt.User) -> None:
pass
async def on_typing_stop(self, channel: revolt.Channel, user: revolt.User) -> None:
pass
async def on_server_update(self, before: revolt.Server, after: revolt.Server) -> None:
pass
async def on_server_delete(self, server: revolt.Server) -> None:
pass
async def on_server_join(self, server: revolt.Server) -> None:
pass
async def on_member_update(self, before: revolt.Member, after: revolt.Member) -> None:
pass
async def on_member_join(self, member: revolt.Member) -> None:
pass
async def on_member_leave(self, member: revolt.Member) -> None:
pass
async def on_role_create(self, role: revolt.Role) -> None:
pass
async def on_role_update(self, before: revolt.Role, after: revolt.Role) -> None:
pass
async def on_role_delete(self, role: revolt.Role) -> None:
pass
async def on_user_update(self, before: revolt.User, after: revolt.User) -> None:
pass
async def on_user_relationship_update(self, user: revolt.User, before: revolt.RelationshipType, after: revolt.RelationshipType) -> None:
pass
async def on_raw_reaction_add(self, payload: revolt.types.MessageReactEventPayload) -> None:
pass
async def on_reaction_add(self, message: revolt.Message, user: revolt.User, emoji_id: str) -> None:
pass
async def on_raw_reaction_remove(self, payload: revolt.types.MessageUnreactEventPayload) -> None:
pass
async def on_reaction_remove(self, message: revolt.Message, user: revolt.User, emoji_id: str) -> None:
pass
async def on_raw_reaction_clear(self, payload: revolt.types.MessageRemoveReactionEventPayload) -> None:
pass
async def on_reaction_clear(self, message: revolt.Message, user: revolt.User, emoji_id: str) -> None:
pass
async def raw_bulk_message_delete(self, payload: revolt.types.BulkMessageDeleteEventPayload) -> None:
pass
async def bulk_message_delete(self, messages: list[revolt.Message]) -> None:
pass
+68 -24
View File
@@ -1,6 +1,10 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Optional, Union 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 .asset import Asset
from .enums import EmbedType from .enums import EmbedType
@@ -9,10 +13,10 @@ if TYPE_CHECKING:
from .state import State from .state import State
from .types import Embed as EmbedPayload from .types import Embed as EmbedPayload
from .types import ImageEmbed as ImageEmbedPayload from .types import ImageEmbed as ImageEmbedPayload
from .types import NoneEmbed as NoneEmbedPayload
from .types import SendableEmbed as SendableEmbedPayload from .types import SendableEmbed as SendableEmbedPayload
from .types import TextEmbed as TextEmbedPayload from .types import TextEmbed as TextEmbedPayload
from .types import WebsiteEmbed as WebsiteEmbedPayload from .types import WebsiteEmbed as WebsiteEmbedPayload
from .types import JanuaryImage, JanuaryVideo
__all__ = ("Embed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "to_embed", "SendableEmbed") __all__ = ("Embed", "WebsiteEmbed", "ImageEmbed", "TextEmbed", "NoneEmbed", "to_embed", "SendableEmbed")
@@ -20,43 +24,45 @@ class WebsiteEmbed:
type = EmbedType.website type = EmbedType.website
def __init__(self, embed: WebsiteEmbedPayload): def __init__(self, embed: WebsiteEmbedPayload):
self.url = embed.get("url") self.url: str | None = embed.get("url")
self.special = embed.get("special") self.special: WebsiteSpecial | None = embed.get("special")
self.title = embed.get("title") self.title: str | None = embed.get("title")
self.description = embed.get("description") self.description: str | None = embed.get("description")
self.image = embed.get("image") self.image: JanuaryImage | None = embed.get("image")
self.video = embed.get("video") self.video: JanuaryVideo | None = embed.get("video")
self.site_name = embed.get("site_name") self.site_name: str | None = embed.get("site_name")
self.icon_url = embed.get("icon_url") self.icon_url: str | None = embed.get("icon_url")
self.colour = embed.get("colour") self.colour: str | None = embed.get("colour")
class ImageEmbed: class ImageEmbed:
type = EmbedType.image type: EmbedType = EmbedType.image
def __init__(self, image: ImageEmbedPayload): def __init__(self, image: ImageEmbedPayload):
self.url = image.get("url") self.url: str = image.get("url")
self.width = image.get("width") self.width: int = image.get("width")
self.height = image.get("height") self.height: int = image.get("height")
self.size = image.get("size") self.size: str = image.get("size")
class TextEmbed: class TextEmbed:
type = EmbedType.text type: EmbedType = EmbedType.text
def __init__(self, embed: TextEmbedPayload, state: State): def __init__(self, embed: TextEmbedPayload, state: State):
self.icon_url = embed.get("icon_url") self.icon_url: str | None = embed.get("icon_url")
self.url = embed.get("url") self.url: str | None = embed.get("url")
self.title = embed.get("title") self.title: str | None = embed.get("title")
self.description = embed.get("description") self.description: str | None = embed.get("description")
self.media: Asset | None
if media := embed.get("media"): if media := embed.get("media"):
self.media = Asset(media, state) self.media = Asset(media, state)
else: else:
self.media = None self.media = None
self.colour = embed.get("colour") self.colour: str | None = embed.get("colour")
class NoneEmbed: class NoneEmbed:
type = EmbedType.none type: EmbedType = EmbedType.none
Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed] Embed = Union[WebsiteEmbed, ImageEmbed, TextEmbed, NoneEmbed]
@@ -70,8 +76,39 @@ def to_embed(payload: EmbedPayload, state: State) -> Embed:
else: else:
return NoneEmbed() return NoneEmbed()
class EmbedParameters(TypedDict):
title: NotRequired[str]
description: NotRequired[str]
media: NotRequired[str]
icon_url: NotRequired[str]
colour: NotRequired[str]
url: NotRequired[str]
class SendableEmbed: class SendableEmbed:
def __init__(self, **attrs): """
Represents an embed that can be sent in a message, you will never receive this, you will receive :class:`Embed`.
Attributes
-----------
title: Optional[:class:`str`]
The title of the embed
description: Optional[:class:`str`]
The description of the embed
media: Optional[:class:`str`]
The file inside the embed, this is the ID of the file, you can use :meth:`Client.upload_file` to get an ID.
icon_url: Optional[:class:`str`]
The url of the icon url
colour: Optional[:class:`str`]
The embed's accent colour, this is any valid `CSS color <https://developer.mozilla.org/en-US/docs/Web/CSS/color_value>`_
url: Optional[:class:`str`]
URL for hyperlinking the embed's title
"""
def __init__(self, **attrs: Unpack[EmbedParameters]):
self.title: Optional[str] = None self.title: Optional[str] = None
self.description: Optional[str] = None self.description: Optional[str] = None
self.media: Optional[str] = None self.media: Optional[str] = None
@@ -83,6 +120,13 @@ class SendableEmbed:
setattr(self, key, value) setattr(self, key, value)
def to_dict(self) -> SendableEmbedPayload: def to_dict(self) -> SendableEmbedPayload:
"""Converts the embed to a dictionary which Revolt accepts
Returns
--------
:class:`dict[str, Any]`
The embed
"""
output: SendableEmbedPayload = {"type": "Text"} output: SendableEmbedPayload = {"type": "Text"}
if title := self.title: if title := self.title:
+55
View File
@@ -0,0 +1,55 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from .utils import Ulid
if TYPE_CHECKING:
from .server import Server
from .state import State
from .types import Emoji as EmojiPayload
__all__ = ("Emoji",)
class Emoji(Ulid):
"""Represents a custom emoji.
Attributes
-----------
id: :class:`str`
The id of the emoji
author_id: :class:`str`
The id of the of user who created the emoji
name: :class:`str`
The name of the emoji
animated: :class:`bool`
Whether or not the emoji is animated
nsfw: :class:`bool`
Whether or not the emoji is nsfw
server_id: Optional[:class:`str`]
The server id this emoji belongs to, if any
"""
def __init__(self, payload: EmojiPayload, state: State):
self.state: State = state
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) -> None:
"""Deletes the emoji."""
await self.state.http.delete_emoji(self.id)
@property
def server(self) -> Server:
"""Returns the server this emoji is part of
Returns
--------
:class:`Server`
The Server this emoji is part of
"""
return self.state.get_server(self.server_id) # type: ignore
+1
View File
@@ -29,6 +29,7 @@ class PresenceType(enum.Enum):
idle = "Idle" idle = "Idle"
invisible = "Invisible" invisible = "Invisible"
online = "Online" online = "Online"
focus = "Focus"
class RelationshipType(enum.Enum): class RelationshipType(enum.Enum):
blocked = "Blocked" blocked = "Blocked"
+6 -2
View File
@@ -4,6 +4,7 @@ __all__ = (
"ServerError", "ServerError",
"FeatureDisabled", "FeatureDisabled",
"AutumnDisabled", "AutumnDisabled",
"Forbidden",
) )
class RevoltError(Exception): class RevoltError(Exception):
@@ -16,7 +17,10 @@ class ServerError(RevoltError):
"Internal server error" "Internal server error"
class FeatureDisabled(RevoltError): class FeatureDisabled(RevoltError):
"""Base class for any feature disabled errors""" "Base class for any feature disabled errors"
class AutumnDisabled(FeatureDisabled): class AutumnDisabled(FeatureDisabled):
"""The autumn feature is disabled""" "The autumn feature is disabled"
class Forbidden(HTTPError):
"Missing permissions"
+1
View File
@@ -4,6 +4,7 @@ from .cog import *
from .command import * from .command import *
from .context import * from .context import *
from .converters import * from .converters import *
from .cooldown import *
from .errors import * from .errors import *
from .group import * from .group import *
from .help import * from .help import *
+55 -21
View File
@@ -1,56 +1,63 @@
from typing import Any, Callable, Coroutine, TypeVar, Union from __future__ import annotations
from typing import Any, Callable, Coroutine, Union, cast
from typing_extensions import TypeVar
import revolt
from .command import Command from .command import Command
from .context import Context from .context import Context
from .errors import NotBotOwner, NotServerOwner, ServerOnly from .errors import (MissingPermissionsError, NotBotOwner, NotServerOwner,
ServerOnly)
from .utils import ClientT_D
__all__ = ("check", "Check", "is_bot_owner", "is_server_owner") __all__ = ("check", "Check", "is_bot_owner", "is_server_owner", "has_permissions", "has_channel_permissions")
Check = Callable[[Context], Union[Any, Coroutine[Any, Any, Any]]] T = TypeVar("T", Callable[..., Any], Command, default=Command)
T = TypeVar("T", Callable[..., Any], Command) Check = Callable[[Context[ClientT_D]], Union[Any, Coroutine[Any, Any, Any]]]
def check(check: Check): def check(check: Check[ClientT_D]) -> Callable[[T], T]:
"""A decorator for adding command checks """A decorator for adding command checks
Parameters Parameters
----------- -----------
check: Callable[[Context], Union[Any, Coroutine[Any, Any, Any]]] check: Callable[[Context], Union[Any, Coroutine[Any, Any, Any]]]
The function to be called, must take one parameter, context and optionally be a coroutine The function to be called, must take one parameter, context and optionally be a coroutine, the return value denoating whether the check should pass or fail
Returns
--------
Any
The value denoating whether the check should pass or fail
""" """
def inner(func: T) -> T: def inner(func: T) -> T:
if isinstance(func, Command): if isinstance(func, Command):
func.checks.append(check) command = cast(Command[ClientT_D], func) # cant verify generic at runtime so must cast
command.checks.append(check)
else: else:
checks = getattr(func, "_checks", []) checks = getattr(func, "_checks", [])
checks.append(check) checks.append(check)
func._checks = checks # type: ignore func._checks = checks # type: ignore
return func return func # type: ignore
return inner 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""" """A command check for limiting the command to only the bot's owner"""
@check @check
def inner(context: Context): def inner(context: Context[ClientT_D]):
if context.author.id == context.client.user.owner_id: if user_id := context.client.user.owner_id:
return True if context.author.id == user_id:
return True
else:
if context.author.id == context.client.user.id:
return True
raise NotBotOwner raise NotBotOwner
return inner 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""" """A command check for limiting the command to only a server's owner"""
@check @check
def inner(context: Context): def inner(context: Context[ClientT_D]) -> bool:
if not context.server: if not context.server_id:
raise ServerOnly raise ServerOnly
if context.author.id == context.server.owner_id: if context.author.id == context.server.owner_id:
@@ -59,3 +66,30 @@ def is_server_owner():
raise NotServerOwner raise NotServerOwner
return inner return inner
def has_permissions(**permissions: bool) -> Callable[[T], T]:
@check
def inner(context: Context[ClientT_D]) -> bool:
author = context.author
if not author.has_permissions(**permissions):
raise MissingPermissionsError(permissions)
return True
return inner
def has_channel_permissions(**permissions: bool) -> Callable[[T], T]:
@check
def inner(context: Context[ClientT_D]) -> bool:
author = context.author
if not isinstance(author, revolt.Member):
raise ServerOnly
if not author.has_channel_permissions(context.channel, **permissions):
raise MissingPermissionsError(permissions)
return True
return inner
+113 -47
View File
@@ -3,8 +3,8 @@ from __future__ import annotations
import sys import sys
import traceback import traceback
from importlib import import_module from importlib import import_module
from typing import (TYPE_CHECKING, Any, Optional, Protocol, Union, from typing import (TYPE_CHECKING, Any, Coroutine, Optional, Protocol, TypeVar, Union,
runtime_checkable) overload, runtime_checkable)
from typing_extensions import Self from typing_extensions import Self
@@ -13,6 +13,8 @@ import revolt
if TYPE_CHECKING: if TYPE_CHECKING:
from .help import HelpCommand from .help import HelpCommand
import aiohttp
from .cog import Cog from .cog import Cog
from .command import Command from .command import Command
from .context import Context from .context import Context
@@ -24,6 +26,9 @@ __all__ = (
"CommandsClient" "CommandsClient"
) )
V = TypeVar("V")
T = TypeVar("T")
@runtime_checkable @runtime_checkable
class ExtensionProtocol(Protocol): class ExtensionProtocol(Protocol):
@staticmethod @staticmethod
@@ -31,32 +36,44 @@ class ExtensionProtocol(Protocol):
raise NotImplementedError raise NotImplementedError
class CommandsMeta(type): class CommandsMeta(type):
_commands: list[Command] _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]) -> Any:
commands: list[Command] = [] commands: list[Command[Any]] = []
self = super().__new__(cls, name, bases, attrs) self = super().__new__(cls, name, bases, attrs)
for base in reversed(self.__mro__): for base in reversed(self.__mro__):
for value in base.__dict__.values(): for value in base.__dict__.values():
if isinstance(value, Command): if isinstance(value, Command) and value.parent is None: # type: ignore
commands.append(value) commands.append(value) # type: ignore
self._commands = commands self._commands = commands
return self return self
class CaseInsensitiveDict(dict): class CaseInsensitiveDict(dict[str, V]):
def __setitem__(self, key: str, value: Any) -> None: def __setitem__(self, key: str, value: V) -> None:
super().__setitem__(key.casefold(), value) super().__setitem__(key.casefold(), value)
def __getitem__(self, key: str) -> Any: def __getitem__(self, key: str) -> V:
return super().__getitem__(key.casefold()) return super().__getitem__(key.casefold())
def __contains__(self, key: str) -> bool: def __contains__(self, key: object) -> bool:
return super().__contains__(key.casefold()) if isinstance(key, str):
return super().__contains__(key.casefold())
else:
return False
def get(self, key: str, default: Any = None) -> Any: @overload
def get(self, key: str) -> V | None:
...
@overload
def get(self, key: str, default: V | T) -> V | T:
...
def get(self, key: str, default: Optional[T] = None) -> V | T | None:
return super().get(key.casefold(), default) return super().get(key.casefold(), default)
def __delitem__(self, key: str) -> None: def __delitem__(self, key: str) -> None:
@@ -64,15 +81,43 @@ class CaseInsensitiveDict(dict):
class CommandsClient(revolt.Client, metaclass=CommandsMeta): class CommandsClient(revolt.Client, metaclass=CommandsMeta):
"""Main class that adds commands, this class should be subclassed along with `revolt.Client`.""" """A subclass of :class:`~revolt.Client` which has support for commands.
_commands: list[Command] Parameters
-----------
session: :class:`~aiohttp.ClientSession`
The aiohttp session to use for http request and the websocket
token: :class:`str`
The bots token
api_url: :class:`str`
The api url for the revolt instance you are connecting to, by default it uses the offical instance hosted at revolt.chat
max_messages: :class:`int`
The max amount of messages stored in the cache, by default this is 5k
bot: :class:`bool`
Denotes whether the account used is a bot account or user account, by default this it assumes a bot account
help_command: Optional[:class:`~revolt.ext.commands.HelpCommand`]
Sets the custom help command, or remove it if passed ``None``
case_insensitive: :class:`bool`
Whether or not commands should be case insensitive
"""
def __init__(self, *args, help_command: Optional[HelpCommand] = None, case_insensitive: bool = False, **kwargs): _commands: list[Command[Self]]
def __init__(
self,
session: aiohttp.ClientSession,
token: str,
*,
api_url: str = "https://api.revolt.chat",
max_messages: int = 5000,
bot: bool = True,
help_command: Union[HelpCommand[Self], None, revolt.utils._Missing] = revolt.utils.Missing,
case_insensitive: bool = False
):
from .help import DefaultHelpCommand, HelpCommandImpl from .help import DefaultHelpCommand, HelpCommandImpl
self.all_commands: dict[str, Command] = {} if not case_insensitive else CaseInsensitiveDict() self.all_commands: dict[str, Command[Self]] | CaseInsensitiveDict[Command[Self]] = {} if not case_insensitive else CaseInsensitiveDict()
self.cogs: dict[str, Cog] = {} self.cogs: dict[str, Cog[Self]] = {}
self.extensions: dict[str, ExtensionProtocol] = {} self.extensions: dict[str, ExtensionProtocol] = {}
for command in self._commands: for command in self._commands:
@@ -81,15 +126,25 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
for alias in command.aliases: for alias in command.aliases:
self.all_commands[alias] = command self.all_commands[alias] = command
if help_command is None: self.help_command: HelpCommand[Self] | None
help_command = DefaultHelpCommand()
self.help_command = DefaultHelpCommand() if help_command is not None:
self.add_command(HelpCommandImpl(self)) self.help_command = help_command or DefaultHelpCommand[Self]()
super().__init__(*args, **kwargs) self.add_command(HelpCommandImpl(self))
else:
self.help_command = None
super().__init__(session, token, api_url=api_url, max_messages=max_messages, bot=bot)
@property @property
def commands(self) -> list[Command]: def commands(self) -> list[Command[Self]]:
"""Gets all commands registered
Returns
--------
list[:class:`Command`]
The registered commands
"""
return list(set(self.all_commands.values())) return list(set(self.all_commands.values()))
async def get_prefix(self, message: revolt.Message) -> Union[str, list[str]]: async def get_prefix(self, message: revolt.Message) -> Union[str, list[str]]:
@@ -107,7 +162,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
""" """
raise NotImplementedError raise NotImplementedError
def get_command(self, name: str) -> Command: def get_command(self, name: str) -> Command[Self]:
"""Gets a command. """Gets a command.
Parameters Parameters
@@ -122,7 +177,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
""" """
return self.all_commands[name] return self.all_commands[name]
def add_command(self, command: Command): 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. """Adds a command, this is typically only used for dynamic commands, you should use the `commands.command` decorator for most usecases.
Parameters Parameters
@@ -137,7 +192,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
for alias in command.aliases: for alias in command.aliases:
self.all_commands[alias] = command self.all_commands[alias] = command
def remove_command(self, name: str) -> Optional[Command]: def remove_command(self, name: str) -> Optional[Command[Self]]:
"""Removes a command. """Removes a command.
Parameters Parameters
@@ -159,10 +214,24 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return command return command
def get_view(self, message: revolt.Message) -> type[StringView]: def get_view(self, message: revolt.Message) -> type[StringView]:
"""Returns the StringView class to use, this can be overwritten to customize how arguments are parsed
Returns
--------
type[:class:`StringView`]
The string view class to use
"""
return StringView return StringView
def get_context(self, message: revolt.Message) -> type[Context[Self]]: def get_context(self, message: revolt.Message) -> type[Context[Self]]:
return Context """Returns the Context class to use, this can be overwritten to add extra features to context
Returns
--------
type[:class:`Context`]
The context class to use
"""
return Context[Self]
async def process_commands(self, message: revolt.Message) -> Any: async def process_commands(self, message: revolt.Message) -> Any:
"""Processes commands, if you overwrite `Client.on_message` you should manually call this function inside the event. """Processes commands, if you overwrite `Client.on_message` you should manually call this function inside the event.
@@ -179,9 +248,6 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
""" """
content = message.content content = message.content
if not isinstance(content, str):
return
prefixes = await self.get_prefix(message) prefixes = await self.get_prefix(message)
if isinstance(prefixes, str): if isinstance(prefixes, str):
@@ -217,7 +283,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
try: try:
self.dispatch("command", context) self.dispatch("command", context)
if not await self.bot_check(context): if not await self.global_check(context):
raise CheckError(f"the global check for the command failed") raise CheckError(f"the global check for the command failed")
if not await context.can_run(): if not await context.can_run():
@@ -231,14 +297,14 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
await command._error_handler(command.cog or self, context, e) await command._error_handler(command.cog or self, context, e)
self.dispatch("command_error", context, e) self.dispatch("command_error", context, e)
@staticmethod async def on_command_error(self, ctx: Context[Self], error: Exception, /) -> None:
async def on_command_error(ctx: Context, error: Exception):
traceback.print_exception(type(error), error, error.__traceback__) traceback.print_exception(type(error), error, error.__traceback__)
on_message = process_commands def on_message(self, message: revolt.Message) -> Coroutine[Any, Any, Any]:
return self.process_commands(message)
async def bot_check(self, context: Context) -> bool: async def global_check(self, context: Context[Self]) -> bool:
"""A global check for the bot that stops commands from running on certain criteria. """A global check that stops commands from running on certain criteria.
Parameters Parameters
----------- -----------
@@ -252,8 +318,8 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return True return True
def add_cog(self, cog: Cog): def add_cog(self, cog: Cog[Self]) -> None:
"""Adds a cog to the bot, this cog must subclass `Cog`. """Adds a cog, this cog must subclass `Cog`.
Parameters Parameters
----------- -----------
@@ -262,8 +328,8 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
""" """
cog._inject(self) cog._inject(self)
def remove_cog(self, cog_name: str) -> Cog: def remove_cog(self, cog_name: str) -> Cog[Self]:
"""Removes a cog from the bot. """Removes a cog.
Parameters Parameters
----------- -----------
@@ -280,7 +346,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return cog 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. """Loads an extension, this takes a module name and runs the setup function inside of it.
Parameters Parameters
@@ -296,7 +362,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
self.extensions[name] = extension self.extensions[name] = extension
extension.setup(self) 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. """Unloads an extension, this takes a module name and runs the teardown function inside of it.
Parameters Parameters
@@ -311,7 +377,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
if teardown := getattr(extension, "teardown", None): if teardown := getattr(extension, "teardown", None):
teardown(self) 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. """Reloads an extension, this will unload and reload the extension.
Parameters Parameters
@@ -322,8 +388,8 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
self.unload_extension(name) self.unload_extension(name)
self.load_extension(name) self.load_extension(name)
def get_cog(self, name: str) -> Cog: def get_cog(self, name: str) -> Cog[Self]:
"""Gets a cog from the bot. """Gets a cog.
Parameters Parameters
----------- -----------
@@ -338,7 +404,7 @@ class CommandsClient(revolt.Client, metaclass=CommandsMeta):
return self.cogs[name] return self.cogs[name]
def get_extension(self, name: str) -> ExtensionProtocol: def get_extension(self, name: str) -> ExtensionProtocol:
"""Gets an extension from the bot. """Gets an extension.
Parameters Parameters
----------- -----------
+70 -23
View File
@@ -1,61 +1,108 @@
from __future__ import annotations from __future__ import annotations
from distutils import command from typing import Any, Callable, Coroutine, Generic, Optional, TypeVar
from typing import TYPE_CHECKING, Any, Optional from typing_extensions import ParamSpec
from revolt.errors import RevoltError
from .command import Command from .command import Command
from .utils import ClientT_D
if TYPE_CHECKING: P = ParamSpec("P")
from .client import CommandsClient R = TypeVar("R")
__all__ = ("Cog", "CogMeta") __all__ = ("Cog", "CogMeta")
class CogMeta(type): class CogMeta(type):
_commands: list[Command] _cog_commands: list[Command[Any]]
_cog_listeners: dict[str, list[str]]
qualified_name: str 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, extras: dict[str, Any] | None = None) -> Any:
commands: list[Command] = [] commands: list[Command[Any]] = []
listeners: dict[str, list[str]] = {}
self = super().__new__(cls, name, bases, attrs) self = super().__new__(cls, name, bases, attrs)
extras = extras or {}
for base in reversed(self.__mro__): for base in reversed(self.__mro__):
for value in base.__dict__.values(): for key, value in base.__dict__.items():
if isinstance(value, Command): if isinstance(value, Command):
commands.append(value) for extra_key, extra_value in extras.items():
setattr(value, extra_key, extra_value) # type: ignore
commands.append(value) # type: ignore
self._commands = commands elif event_name := getattr(value, "__listener_name", None):
listeners.setdefault(event_name, []).append(key)
self._cog_commands = commands
self._cog_listeners = listeners
self.qualified_name = qualified_name or name self.qualified_name = qualified_name or name
return self return self
class Cog(metaclass=CogMeta): class Cog(Generic[ClientT_D], metaclass=CogMeta):
_commands: list[Command] _cog_commands: list[Command[ClientT_D]]
_cog_listeners: dict[str, list[str]]
qualified_name: str qualified_name: str
def cog_load(self): def cog_load(self) -> None:
"""A special method that is called when the cog gets loaded.""" """A special method that is called when the cog gets loaded."""
pass pass
def cog_unload(self): def cog_unload(self) -> None:
"""A special method that is called when the cog gets removed.""" """A special method that is called when the cog gets removed."""
pass pass
def _inject(self, client: CommandsClient): def _inject(self, client: ClientT_D) -> None:
client.cogs[self.qualified_name] = self client.cogs[self.qualified_name] = self
for command in self._commands: try:
command.cog = self for command in self._cog_commands:
client.add_command(command) command.cog = self
if command.parent is None:
client.add_command(command)
for key, listeners in self._cog_listeners.items():
for listener_name in listeners:
client.listeners.setdefault(key, []).append(getattr(self, listener_name))
except Exception as e:
self._uninject(client)
raise e
self.cog_load() self.cog_load()
def _uninject(self, client: CommandsClient): def _uninject(self, client: ClientT_D) -> None:
for name, command in client.all_commands.copy().items(): for name, command in client.all_commands.copy().items():
if command in self._commands: if command in self._cog_commands:
del client.all_commands[name] try:
del client.all_commands[name]
except KeyError:
pass
for key, listeners in self._cog_listeners.items():
for listener_name in listeners:
try:
client.listeners[key].remove(getattr(self, listener_name))
except ValueError:
pass
self.cog_unload() self.cog_unload()
@property @property
def commands(self) -> list[Command]: def commands(self) -> list[Command[ClientT_D]]:
return self._commands return self._cog_commands
@staticmethod
def listen(name: str | None = None) -> Callable[[Callable[P, Coroutine[Any, Any, R]]], Callable[P, Coroutine[Any, Any, R]]]:
def inner(func: Callable[P, Coroutine[Any, Any, R]]) -> Callable[P, Coroutine[Any, Any, R]]:
if not func.__name__.startswith("on_"):
raise RevoltError("event name must start with `on_`")
setattr(func, "__listener_name", name or func.__name__[3:])
return func
return inner
+114 -39
View File
@@ -4,13 +4,22 @@ import inspect
import traceback import traceback
from contextlib import suppress from contextlib import suppress
from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine, from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine,
Literal, Optional, Union, cast, get_args, get_origin) Generic, Literal, Optional, Union, get_args, get_origin)
from typing_extensions import ParamSpec
import sys
import revolt if sys.version_info >= (3, 10):
from revolt.utils import copy_doc, maybe_coroutine from types import UnionType
from .errors import InvalidLiteralArgument, UnionConverterError UnionTypes: tuple[Any, ...] = (Union, UnionType)
from .utils import evaluate_parameters else:
UnionTypes = (Union,)
from ...utils import maybe_coroutine
from .errors import CommandOnCooldown, InvalidLiteralArgument, UnionConverterError
from .utils import ClientT_Co_D, evaluate_parameters, ClientT_Co
from .cooldown import BucketType, CooldownMapping
if TYPE_CHECKING: if TYPE_CHECKING:
from .checks import Check from .checks import Check
@@ -18,15 +27,15 @@ if TYPE_CHECKING:
from .context import Context from .context import Context
from .group import Group from .group import Group
__all__ = ( __all__: tuple[str, ...] = (
"Command", "Command",
"command" "command"
) )
NoneType = type(None) NoneType: type[None] = type(None)
P = ParamSpec("P")
class Command(Generic[ClientT_Co_D]):
class Command:
"""Class for holding info about a command. """Class for holding info about a command.
Parameters Parameters
@@ -43,23 +52,46 @@ class Command:
The cog the command is apart of. The cog the command is apart of.
usage: Optional[:class:`str`] usage: Optional[:class:`str`]
The usage string for the command The usage string for the command
checks: Optional[list[Callable]]
The list of checks the command has
cooldown: Optional[:class:`Cooldown`]
The cooldown for the command to restrict how often the command can be used
description: Optional[:class:`str`]
The commands description if it has one
hidden: :class:`bool`
Whether or not the command should be hidden from the help command
""" """
__slots__ = ("callback", "name", "aliases", "signature", "checks", "parent", "_error_handler", "cog", "description", "usage", "parameters") __slots__ = ("callback", "name", "aliases", "signature", "checks", "parent", "_error_handler", "cog", "description", "usage", "parameters", "hidden", "cooldown", "cooldown_bucket")
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str], usage: Optional[str] = None): def __init__(
self.callback = callback self,
self.name = name callback: Callable[..., Coroutine[Any, Any, Any]],
self.aliases = aliases name: str,
self.usage = usage *,
self.signature = inspect.signature(self.callback) aliases: list[str] | None = None,
self.parameters = evaluate_parameters(self.signature.parameters.values(), getattr(callback, "__globals__", {})) usage: Optional[str] = None,
self.checks: list[Check] = getattr(callback, "_checks", []) checks: list[Check[ClientT_Co_D]] | None = None,
self.parent: Optional[Group] = None cooldown: Optional[CooldownMapping] | None = None,
self.cog: Optional[Cog] = None bucket: Optional[BucketType | Callable[[Context[ClientT_Co_D]], Coroutine[Any, Any, str]]] = None,
self._error_handler: Callable[[Any, Context, Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler description: str | None = None,
self.description = callback.__doc__ hidden: bool = False,
):
self.callback: Callable[..., Coroutine[Any, Any, Any]] = callback
self.name: str = name
self.aliases: list[str] = aliases or []
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[ClientT_Co_D]] = checks or getattr(callback, "_checks", [])
self.cooldown: CooldownMapping | None = cooldown or getattr(callback, "_cooldown", None)
self.cooldown_bucket: BucketType | Callable[[Context[ClientT_Co_D]], Coroutine[Any, Any, str]] = bucket or getattr(callback, "_bucket", BucketType.default)
self.parent: Optional[Group[ClientT_Co_D]] = None
self.cog: Optional[Cog[ClientT_Co_D]] = None
self._error_handler: Callable[[Any, Context[ClientT_Co_D], Exception], Coroutine[Any, Any, Any]] = type(self)._default_error_handler
self.description: str | None = description or callback.__doc__
self.hidden: bool = hidden
async def invoke(self, context: Context, *args, **kwargs) -> Any: async def invoke(self, context: Context[ClientT_Co_D], *args: Any, **kwargs: Any) -> Any:
"""Runs the command and calls the error handler if the command errors. """Runs the command and calls the error handler if the command errors.
Parameters Parameters
@@ -74,11 +106,10 @@ class Command:
except Exception as err: except Exception as err:
return await self._error_handler(self.cog or context.client, context, err) return await self._error_handler(self.cog or context.client, context, err)
@copy_doc(invoke) def __call__(self, context: Context[ClientT_Co_D], *args: Any, **kwargs: Any) -> Any:
def __call__(self, context: Context, *args, **kwargs) -> Any:
return self.invoke(context, *args, **kwargs) 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. """Sets the error handler for the command.
Parameters Parameters
@@ -98,12 +129,12 @@ class Command:
self._error_handler = func self._error_handler = func
return func return func
async def _default_error_handler(self, ctx: Context, error: Exception): async def _default_error_handler(self, ctx: Context[ClientT_Co_D], error: Exception):
traceback.print_exception(type(error), error, error.__traceback__) traceback.print_exception(type(error), error, error.__traceback__)
@classmethod @classmethod
async def handle_origin(cls, context: Context, origin: Any, annotation: Any, arg: str) -> Any: async def handle_origin(cls, context: Context[ClientT_Co_D], origin: Any, annotation: Any, arg: str) -> Any:
if origin is Union: if origin in UnionTypes:
for converter in get_args(annotation): for converter in get_args(annotation):
try: try:
return await cls.convert_argument(arg, converter, context) return await cls.convert_argument(arg, converter, context)
@@ -117,8 +148,20 @@ class Command:
elif origin is Annotated: elif origin is Annotated:
annotated_args = get_args(annotation) annotated_args = get_args(annotation)
if origin := get_origin(annotated_args[0]): if annotated_args[1] == "_revolt_greedy_marker":
return await cls.handle_origin(context, origin, annotated_args[1], arg) real_annotation = get_args(annotated_args[0])[0]
converted_args: list[Any] = []
converted_args.append(await cls.convert_argument(arg, real_annotation, context))
for arg in context.view:
try:
converted_args.append(await cls.convert_argument(arg, real_annotation, context))
except:
context.view.undo()
break
return converted_args
else: else:
return await cls.convert_argument(arg, annotated_args[1], context) return await cls.convert_argument(arg, annotated_args[1], context)
@@ -129,11 +172,12 @@ class Command:
raise InvalidLiteralArgument(arg) raise InvalidLiteralArgument(arg)
@classmethod @classmethod
async def convert_argument(cls, arg: str, annotation: Any, context: Context) -> Any: async def convert_argument(cls, arg: str, annotation: Any, context: Context[ClientT_Co_D]) -> Any:
if annotation is not inspect.Signature.empty: if annotation is not inspect.Signature.empty:
if annotation is str: # no converting is needed - its already a string if annotation is str: # no converting is needed - its already a string
return arg return arg
origin: Any
if origin := get_origin(annotation): if origin := get_origin(annotation):
return await cls.handle_origin(context, origin, annotation, arg) return await cls.handle_origin(context, origin, annotation, arg)
else: else:
@@ -141,7 +185,7 @@ class Command:
else: else:
return arg return arg
async def parse_arguments(self, context: Context): async def parse_arguments(self, context: Context[ClientT_Co_D]) -> None:
# please pr if you can think of a better way to do this # please pr if you can think of a better way to do this
for parameter in self.parameters[2:]: for parameter in self.parameters[2:]:
@@ -151,6 +195,10 @@ class Command:
except StopIteration: except StopIteration:
if parameter.default is not parameter.empty: if parameter.default is not parameter.empty:
arg = parameter.default arg = parameter.default
elif is_optional(parameter.annotation):
arg = None
else: else:
raise raise
@@ -168,11 +216,27 @@ class Command:
except StopIteration: except StopIteration:
if parameter.default is not parameter.empty: if parameter.default is not parameter.empty:
arg = parameter.default arg = parameter.default
elif is_optional(parameter.annotation):
arg = None
else: else:
raise raise
context.args.append(arg) context.args.append(arg)
async def run_cooldown(self, context: Context[ClientT_Co_D]) -> None:
if mapping := self.cooldown:
if isinstance(self.cooldown_bucket, BucketType):
key = self.cooldown_bucket.resolve(context)
else:
key = await self.cooldown_bucket(context)
cooldown = mapping.get_bucket(key)
if retry_after := cooldown.update_cooldown():
raise CommandOnCooldown(retry_after)
def __repr__(self) -> str: def __repr__(self) -> str:
return f"<{self.__class__.__name__} name=\"{self.name}\">" return f"<{self.__class__.__name__} name=\"{self.name}\">"
@@ -187,7 +251,7 @@ class Command:
if self.usage: if self.usage:
return self.usage return self.usage
parents = [] parents: list[str] = []
if self.parent: if self.parent:
parent = self.parent parent = self.parent
@@ -196,7 +260,7 @@ class Command:
parents.append(parent.name) parents.append(parent.name)
parent = parent.parent parent = parent.parent
parameters = [] parameters: list[str] = []
for parameter in self.parameters[2:]: for parameter in self.parameters[2:]:
if parameter.kind == parameter.POSITIONAL_OR_KEYWORD: if parameter.kind == parameter.POSITIONAL_OR_KEYWORD:
@@ -214,8 +278,17 @@ class Command:
return f"{' '.join(parents[::-1])} {self.name} {' '.join(parameters)}" return f"{' '.join(parents[::-1])} {self.name} {' '.join(parameters)}"
def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command] = Command, usage: Optional[str] = None): def is_optional(arg: Any) -> bool:
"""A decorator that turns a function into a :class:`Command`. return get_origin(arg) in UnionTypes and any(arg is NoneType for arg in get_args(arg))
def command(
*,
name: Optional[str] = None,
aliases: Optional[list[str]] = None,
cls: type[Command[ClientT_Co]] = Command,
usage: Optional[str] = None
) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientT_Co]]:
"""A decorator that turns a function into a :class:`Command`.n
Parameters Parameters
----------- -----------
@@ -225,13 +298,15 @@ def command(*, name: Optional[str] = None, aliases: Optional[list[str]] = None,
The aliases of the command, defaults to no aliases The aliases of the command, defaults to no aliases
cls: type[:class:`Command`] cls: type[:class:`Command`]
The class used for creating the command, this defaults to :class:`Command` but can be used to use a custom command subclass The class used for creating the command, this defaults to :class:`Command` but can be used to use a custom command subclass
usage: Optional[:class:`str`]
The signature for how the command should be called
Returns Returns
-------- --------
Callable[Callable[..., Coroutine], :class:`Command`] Callable[Callable[..., Coroutine], :class:`Command`]
A function that takes the command callback and returns a :class:`Command` A function that takes the command callback and returns a :class:`Command`
""" """
def inner(func: Callable[..., Coroutine[Any, Any, Any]]): def inner(func: Callable[..., Coroutine[Any, Any, Any]]) -> Command[ClientT_Co]:
return cls(func, name or func.__name__, aliases or [], usage) return cls(func, name or func.__name__, aliases=aliases or [], usage=usage)
return inner return inner
+40 -20
View File
@@ -1,24 +1,23 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Any, Generic, Optional, TypeVar from typing import TYPE_CHECKING, Any, Generic, Optional
import revolt import revolt
from revolt.utils import maybe_coroutine from revolt.utils import maybe_coroutine
from .command import Command from .command import Command
from .group import Group from .group import Group
from .utils import ClientT_Co_D
if TYPE_CHECKING: if TYPE_CHECKING:
from .client import CommandsClient
from .view import StringView from .view import StringView
from revolt.state import State
ClientT = TypeVar("ClientT", bound="CommandsClient")
__all__ = ( __all__ = (
"Context", "Context",
) )
class Context(revolt.Messageable, Generic[ClientT]): class Context(revolt.Messageable, Generic[ClientT_Co_D]):
"""Stores metadata the commands execution. """Stores metadata the commands execution.
Attributes Attributes
@@ -31,7 +30,7 @@ class Context(revolt.Messageable, Generic[ClientT]):
The message that was sent to invoke the command The message that was sent to invoke the command
channel: :class:`Messageable` channel: :class:`Messageable`
The channel the command was invoked in The channel the command was invoked in
server: :class:`Server` server_id: Optional[:class:`Server`]
The server the command was invoked in The server the command was invoked in
author: Union[:class:`Member`, :class:`User`] author: Union[:class:`Member`, :class:`User`]
The user or member that invoked the commad, will be :class:`User` in DMs The user or member that invoked the commad, will be :class:`User` in DMs
@@ -42,23 +41,37 @@ class Context(revolt.Messageable, Generic[ClientT]):
client: :class:`CommandsClient` client: :class:`CommandsClient`
The revolt client 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: async def _get_channel_id(self) -> str:
return self.channel.id return self.channel.id
def __init__(self, command: Optional[Command], invoked_with: str, view: StringView, message: revolt.Message, client: ClientT): def __init__(self, command: Optional[Command[ClientT_Co_D]], invoked_with: str, view: StringView, message: revolt.Message, client: ClientT_Co_D):
self.command = command self.command: Command[ClientT_Co_D] | None = command
self.invoked_with = invoked_with self.invoked_with: str = invoked_with
self.view = view self.view: StringView = view
self.message = message self.message: revolt.Message = message
self.client = client self.client: ClientT_Co_D = client
self.args = [] self.args: list[Any] = []
self.kwargs = {} self.kwargs: dict[str, Any] = {}
self.server = message.server self.server_id: str | None = message.server_id
self.channel = message.channel self.channel: revolt.TextChannel | revolt.GroupDMChannel | revolt.DMChannel | revolt.SavedMessageChannel = message.channel
self.author = message.author self.author: revolt.Member | revolt.User = message.author
self.state = message.state self.state: State = message.state
@property
def server(self) -> revolt.Server:
""":class:`Server` The server this context belongs too
Raises
-------
:class:`LookupError`
Raises if the context is not from a server
"""
if not self.server_id:
raise LookupError
return self.state.get_server(self.server_id)
async def invoke(self) -> Any: async def invoke(self) -> Any:
"""Invokes the command. """Invokes the command.
@@ -84,11 +97,18 @@ class Context(revolt.Messageable, Generic[ClientT]):
self.view.undo() self.view.undo()
await command.run_cooldown(self)
await command.parse_arguments(self) await command.parse_arguments(self)
return await command.invoke(self, *self.args, **self.kwargs) return await command.invoke(self, *self.args, **self.kwargs)
async def can_run(self, command: Optional[Command] = None) -> bool: async def can_run(self, command: Optional[Command[ClientT_Co_D]] = None) -> bool:
"""Runs all of the commands checks, and returns true if all of them pass""" """Runs all of the commands checks, and returns true if all of them pass"""
command = command or self.command command = command or self.command
return all([await maybe_coroutine(check, self) for check in (command.checks if command else [])]) 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] | ClientT_Co_D | None = None) -> None:
argument = argument or self.client
command = self.client.get_command("help")
await command.invoke(self, argument)
+68 -26
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
import re import re
from typing import Annotated from typing import TYPE_CHECKING, Annotated, TypeVar
from revolt import Category, Channel, Member, User, utils from revolt import Category, Channel, Member, User, utils
@@ -8,77 +10,117 @@ from .errors import (BadBoolArgument, CategoryConverterError,
ChannelConverterError, MemberConverterError, ServerOnly, ChannelConverterError, MemberConverterError, ServerOnly,
UserConverterError) UserConverterError)
__all__ = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter") if TYPE_CHECKING:
from .client import CommandsClient
channel_regex = re.compile("<#([A-z0-9]{26})>") T = TypeVar("T")
user_regex = re.compile("<@([A-z0-9]{26})>")
def bool_converter(arg: str, _): __all__: tuple[str, ...] = ("bool_converter", "category_converter", "channel_converter", "user_converter", "member_converter", "IntConverter", "BoolConverter", "CategoryConverter", "UserConverter", "MemberConverter", "ChannelConverter", "Greedy")
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, _: Context[ClientT]) -> bool:
lowered = arg.lower() 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 return True
elif lowered in ('no', 'n', 'false', 'f', '0', 'disable', 'off'): elif lowered in ("no", "false", "n", "f", "0", "off", "disabled"):
return False return False
else: else:
raise BadBoolArgument(lowered) raise BadBoolArgument(lowered)
def category_converter(arg: str, context: Context) -> Category: def category_converter(arg: str, context: Context[ClientT]) -> Category:
if not (server := context.server): if not context.server_id:
raise ServerOnly raise ServerOnly
try: try:
return server.get_category(arg) return context.server.get_category(arg)
except KeyError: except LookupError:
try: try:
return utils.get(server.categories, name=arg) return utils.get(context.server.categories, name=arg)
except LookupError: except LookupError:
raise CategoryConverterError(arg) raise CategoryConverterError(arg)
def channel_converter(arg: str, context: Context) -> Channel: def channel_converter(arg: str, context: Context[ClientT]) -> Channel:
if not (server := context.server): if not context.server_id:
raise ServerOnly raise ServerOnly
if (match := channel_regex.match(arg)): if (match := channel_regex.match(arg)):
arg = match.group(1) arg = match.group(1)
try: try:
return server.get_channel(arg) return context.server.get_channel(arg)
except KeyError: except LookupError:
try: try:
return utils.get(server.channels, name=arg) return utils.get(context.server.channels, name=arg)
except LookupError: except LookupError:
raise ChannelConverterError(arg) raise ChannelConverterError(arg)
def user_converter(arg: str, context: Context) -> User: def user_converter(arg: str, context: Context[ClientT]) -> User:
if (match := user_regex.match(arg)): if (match := user_regex.match(arg)):
arg = match.group(1) arg = match.group(1)
try: try:
return context.client.get_user(arg) return context.client.get_user(arg)
except KeyError: except LookupError:
try: try:
return utils.get(context.client.users, name=arg) parts = arg.split("#")
if len(parts) == 1:
return (
utils.get(context.client.users, original_name=arg)
or utils.get(context.client.users, display_name=arg)
)
elif len(parts) == 2:
return (
utils.get(context.client.users, original_name=parts[0], discriminator=parts[1])
or utils.get(context.client.users, display_name=parts[0], discriminator=parts[1])
)
else:
raise LookupError
except LookupError: except LookupError:
raise UserConverterError(arg) raise UserConverterError(arg)
def member_converter(arg: str, context: Context) -> Member: def member_converter(arg: str, context: Context[ClientT]) -> Member:
if not (server := context.server): if not context.server_id:
raise ServerOnly raise ServerOnly
if (match := user_regex.match(arg)): if (match := user_regex.match(arg)):
arg = match.group(1) arg = match.group(1)
try: try:
return server.get_member(arg) return context.server.get_member(arg)
except KeyError: except LookupError:
try: try:
return utils.get(server.members, name=arg) parts = arg.split("#")
if len(parts) == 1:
return (
utils.get(context.server.members, original_name=arg)
or utils.get(context.server.members, display_name=arg)
)
elif len(parts) == 2:
return (
utils.get(context.server.members, original_name=parts[0], discriminator=parts[1])
or utils.get(context.server.members, display_name=parts[0], discriminator=parts[1])
)
else:
raise LookupError
except LookupError: except LookupError:
raise MemberConverterError(arg) raise MemberConverterError(arg)
IntConverter = Annotated[int, lambda arg, _: int(arg)] def int_converter(arg: str, context: Context[ClientT]) -> int:
return int(arg)
IntConverter = Annotated[int, int_converter]
BoolConverter = Annotated[bool, bool_converter] BoolConverter = Annotated[bool, bool_converter]
CategoryConverter = Annotated[Category, category_converter] CategoryConverter = Annotated[Category, category_converter]
UserConverter = Annotated[User, user_converter] UserConverter = Annotated[User, user_converter]
MemberConverter = Annotated[Member, member_converter] MemberConverter = Annotated[Member, member_converter]
ChannelConverter = Annotated[Channel, channel_converter] ChannelConverter = Annotated[Channel, channel_converter]
Greedy = Annotated[list[T], "_revolt_greedy_marker"]
+144
View File
@@ -0,0 +1,144 @@
from __future__ import annotations
import time
from typing import TYPE_CHECKING, Any, Callable, Coroutine, TypeVar, cast
from .errors import ServerOnly
if TYPE_CHECKING:
from enum import Enum
from .context import Context
from .utils import ClientT_Co_D, ClientT_Co
else:
from aenum import Enum
__all__ = ("Cooldown", "CooldownMapping", "BucketType", "cooldown")
T = TypeVar("T")
class Cooldown:
"""Represent a single cooldown for a single key
Parameters
-----------
rate: :class:`int`
How many times it can be used
per: :class:`int`
How long the window is before the ratelimit resets
"""
def __init__(self, rate: int, per: int):
self.rate: int = rate
self.per: int = per
self.window: float = 0.0
self.tokens: int = rate
self.last: float = 0.0
def get_tokens(self, current: float | None) -> int:
current = current or time.time()
if current > (self.window + self.per):
return self.rate
else:
return self.tokens
def update_cooldown(self) -> float | None:
current = time.time()
self.last = current
self.tokens = self.get_tokens(current)
if self.tokens == 0:
return self.per - (current - self.window)
self.tokens -= 1
if self.tokens == 0:
self.window = current
return None
class CooldownMapping:
"""Holds all cooldowns for every key"""
def __init__(self, rate: int, per: int):
self.rate = rate
self.per = per
self.cache: dict[str, Cooldown] = {}
def verify_cache(self) -> None:
current = time.time()
self.cache = {k: v for k, v in self.cache.items() if current < (v.last + v.per)}
def get_bucket(self, key: str) -> Cooldown:
self.verify_cache()
if not (rl := self.cache.get(key)):
self.cache[key] = rl = Cooldown(self.rate, self.per)
return rl
class BucketType(Enum):
default = 0
user = 1
server = 2
channel = 3
member = 4
def resolve(self, context: Context[ClientT_Co_D]) -> str:
if self == BucketType.default:
return f"{context.author.id}{context.channel.id}"
elif self == BucketType.user:
return context.author.id
elif self == BucketType.server:
if id := context.server_id:
return id
raise ServerOnly
elif self == BucketType.channel:
return context.channel.id
else: # BucketType.member
if server_id := context.server_id:
return f"{context.author.id}{server_id}"
raise ServerOnly
def cooldown(rate: int, per: int, *, bucket: BucketType | Callable[[Context[ClientT_Co]], Coroutine[Any, Any, str]] = BucketType.default) -> Callable[[T], T]:
"""Adds a cooldown to a command
Parameters
-----------
rate: :class:`int`
How many times it can be used
per: :class:`int`
How long the window is before the ratelimit resets
bucket: Optional[Union[:class:`BucketType`, Callable[[Context], str]]]
Controls how the key is generated for the cooldowns
Examples
--------
.. code-block:: python
@commands.command()
@commands.cooldown(1, 5)
async def ping(self, ctx: Context):
await ctx.send("Pong")
"""
def inner(func: T) -> T:
from .command import Command
if isinstance(func, Command):
command = cast(Command[ClientT_Co], func) # cant verify generic at runtime so must cast
command.cooldown = CooldownMapping(rate, per)
command.cooldown_bucket = bucket
else:
func._cooldown = CooldownMapping(rate, per) # type: ignore
func._bucket = bucket # type: ignore
return func # type: ignore
return inner
+30 -1
View File
@@ -8,6 +8,7 @@ __all__ = (
"NotBotOwner", "NotBotOwner",
"NotServerOwner", "NotServerOwner",
"ServerOnly", "ServerOnly",
"MissingPermissionsError",
"ConverterError", "ConverterError",
"InvalidLiteralArgument", "InvalidLiteralArgument",
"BadBoolArgument", "BadBoolArgument",
@@ -15,7 +16,9 @@ __all__ = (
"ChannelConverterError", "ChannelConverterError",
"UserConverterError", "UserConverterError",
"MemberConverterError", "MemberConverterError",
"UnionConverterError",
"MissingSetup", "MissingSetup",
"CommandOnCooldown"
) )
class CommandError(RevoltError): class CommandError(RevoltError):
@@ -32,7 +35,7 @@ class CommandNotFound(CommandError):
__slots__ = ("command_name",) __slots__ = ("command_name",)
def __init__(self, command_name: str): def __init__(self, command_name: str):
self.command_name = command_name self.command_name: str = command_name
class NoClosingQuote(CommandError): class NoClosingQuote(CommandError):
"""Raised when there is no closing quote for a command argument""" """Raised when there is no closing quote for a command argument"""
@@ -49,6 +52,18 @@ class NotServerOwner(CheckError):
class ServerOnly(CheckError): class ServerOnly(CheckError):
"""Raised when a check requires the command to be ran in a server""" """Raised when a check requires the command to be ran in a server"""
class MissingPermissionsError(CheckError):
"""Raised when a check requires permissions the user does not have
Attributes
-----------
permissions: :class:`dict[str, bool]`
The permissions which the user did not have
"""
def __init__(self, permissions: dict[str, bool]):
self.permissions = permissions
class ConverterError(CommandError): class ConverterError(CommandError):
"""Base class for all converter errors""" """Base class for all converter errors"""
@@ -85,3 +100,17 @@ class UnionConverterError(ConverterError):
class MissingSetup(CommandError): class MissingSetup(CommandError):
"""Raised when an extension is missing the `setup` function""" """Raised when an extension is missing the `setup` function"""
class CommandOnCooldown(CommandError):
"""Raised when a command is on cooldown
Attributes
-----------
retry_after: :class:`float`
How long the user must wait until the cooldown resets
"""
__slots__ = ("retry_after",)
def __init__(self, retry_after: float):
self.retry_after: float = retry_after
+78 -14
View File
@@ -1,19 +1,17 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Any, Callable, Coroutine, Optional from typing import Any, Callable, Coroutine, Optional
from .command import Command from .command import Command
from .utils import ClientT_Co_D, ClientT_D
if TYPE_CHECKING:
from .context import Context
__all__ = ( __all__ = (
"Group", "Group",
"group" "group"
) )
class Group(Command[ClientT_Co_D]):
class Group(Command):
"""Class for holding info about a group command. """Class for holding info about a group command.
Parameters Parameters
@@ -28,13 +26,13 @@ class Group(Command):
The group's subcommands. The group's subcommands.
""" """
__slots__ = ("subcommands",) __slots__: tuple[str, ...] = ("subcommands",)
def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]): def __init__(self, callback: Callable[..., Coroutine[Any, Any, Any]], name: str, aliases: list[str]):
self.subcommands: dict[str, Command] = {} self.subcommands: dict[str, Command[ClientT_Co_D]] = {}
super().__init__(callback, name, aliases) super().__init__(callback, name, aliases=aliases)
def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command] = Command): def command(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Command[ClientT_Co_D]] = Command[ClientT_Co_D]) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Command[ClientT_Co_D]]:
"""A decorator that turns a function into a :class:`Command` and registers the command as a subcommand. """A decorator that turns a function into a :class:`Command` and registers the command as a subcommand.
Parameters Parameters
@@ -52,14 +50,18 @@ class Group(Command):
A function that takes the command callback and returns a :class:`Command` A function that takes the command callback and returns a :class:`Command`
""" """
def inner(func: Callable[..., Coroutine[Any, Any, Any]]): def inner(func: Callable[..., Coroutine[Any, Any, Any]]):
command = cls(func, name or func.__name__, aliases or []) command = cls(func, name or func.__name__, aliases=aliases or [])
command.parent = self command.parent = self
self.subcommands[command.name] = command self.subcommands[command.name] = command
for alias in command.aliases:
self.subcommands[alias] = command
return command return command
return inner return inner
def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type["Group"]] = None): def group(self, *, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: Optional[type[Group[ClientT_Co_D]]] = None) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientT_Co_D]]:
"""A decorator that turns a function into a :class:`Group` and registers the command as a subcommand """A decorator that turns a function into a :class:`Group` and registers the command as a subcommand
Parameters Parameters
@@ -82,6 +84,10 @@ class Group(Command):
command = cls(func, name or func.__name__, aliases or []) command = cls(func, name or func.__name__, aliases or [])
command.parent = self command.parent = self
self.subcommands[command.name] = command self.subcommands[command.name] = command
for alias in command.aliases:
self.subcommands[alias] = command
return command return command
return inner return inner
@@ -90,10 +96,68 @@ class Group(Command):
return f"<Group name=\"{self.name}\">" return f"<Group name=\"{self.name}\">"
@property @property
def commands(self) -> list[Command]: def commands(self) -> list[Command[ClientT_Co_D]]:
return list(self.subcommands.values()) """Gets all commands registered
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group] = Group): Returns
--------
list[:class:`Command`]
The registered commands
"""
return list(set(self.subcommands.values()))
def get_command(self, name: str) -> Command[ClientT_Co_D]:
"""Gets a command.
Parameters
-----------
name: :class:`str`
The name or alias of the command
Returns
--------
:class:`Command`
The command with the name
"""
return self.subcommands[name]
def add_command(self, command: Command[ClientT_Co_D]) -> None:
"""Adds a command, this is typically only used for dynamic commands, you should use the `commands.command` decorator for most usecases.
Parameters
-----------
name: :class:`str`
The name or alias of the command
command: :class:`Command`
The command to be added
"""
self.subcommands[command.name] = command
for alias in command.aliases:
self.subcommands[alias] = command
def remove_command(self, name: str) -> Optional[Command[ClientT_Co_D]]:
"""Removes a command.
Parameters
-----------
name: :class:`str`
The name or alias of the command
Returns
--------
Optional[:class:`Command`]
The command that was removed
"""
command = self.subcommands.pop(name, None)
if command is not None:
for alias in command.aliases:
self.subcommands.pop(alias, None)
return command
def group(*, name: Optional[str] = None, aliases: Optional[list[str]] = None, cls: type[Group[ClientT_D]] = Group) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientT_D]]:
"""A decorator that turns a function into a :class:`Group` """A decorator that turns a function into a :class:`Group`
Parameters Parameters
+80 -61
View File
@@ -1,24 +1,24 @@
from __future__ import annotations from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from itertools import groupby from typing import TYPE_CHECKING, Generic, Optional, TypedDict, Union, cast
from typing import TYPE_CHECKING, Optional, TypedDict, Union
from typing_extensions import NotRequired from typing_extensions import NotRequired
from .client import CommandsClient from .cog import Cog
from .command import Command, command from .command import Command
from .context import Context from .context import Context
from .group import Group from .group import Group
from .utils import evaluate_parameters from .utils import ClientT_Co_D, ClientT_D
from revolt import File, Message, Messageable, MessageReply, SendableEmbed
if TYPE_CHECKING: if TYPE_CHECKING:
from revolt import File, Message, Messageable, MessageReply, SendableEmbed
from .cog import Cog from .cog import Cog
__all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl") __all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl")
class MessagePayload(TypedDict): class MessagePayload(TypedDict):
content: str content: str
embed: NotRequired[SendableEmbed] embed: NotRequired[SendableEmbed]
@@ -26,31 +26,33 @@ class MessagePayload(TypedDict):
attachments: NotRequired[list[File]] attachments: NotRequired[list[File]]
replies: NotRequired[list[MessageReply]] replies: NotRequired[list[MessageReply]]
class HelpCommand(ABC, Generic[ClientT_Co_D]):
class HelpCommand(ABC):
@abstractmethod @abstractmethod
async def create_bot_help(self, context: Context, commands: dict[Optional[Cog], list[Command]]) -> Union[str, SendableEmbed, MessagePayload]: async def create_global_help(self, context: Context[ClientT_Co_D], commands: dict[Optional[Cog[ClientT_Co_D]], list[Command[ClientT_Co_D]]]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError raise NotImplementedError
@abstractmethod @abstractmethod
async def create_command_help(self, context: Context, command: Command) -> Union[str, SendableEmbed, MessagePayload]: async def create_command_help(self, context: Context[ClientT_Co_D], command: Command[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError raise NotImplementedError
@abstractmethod @abstractmethod
async def create_group_help(self, context: Context, group: Group) -> Union[str, SendableEmbed, MessagePayload]: async def create_group_help(self, context: Context[ClientT_Co_D], group: Group[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError raise NotImplementedError
@abstractmethod @abstractmethod
async def create_cog_help(self, context: Context, cog: Cog) -> Union[str, SendableEmbed, MessagePayload]: async def create_cog_help(self, context: Context[ClientT_Co_D], cog: Cog[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError raise NotImplementedError
async def send_help_command(self, context: Context, message_payload: MessagePayload) -> Message: async def send_help_command(self, context: Context[ClientT_Co_D], message_payload: MessagePayload) -> Message:
return await context.send(**message_payload) return await context.send(**message_payload)
async def filter_commands(self, context: Context, commands: list[Command]) -> list[Command]: async def filter_commands(self, context: Context[ClientT_Co_D], commands: list[Command[ClientT_Co_D]]) -> list[Command[ClientT_Co_D]]:
filtered: list[Command] = [] filtered: list[Command[ClientT_Co_D]] = []
for command in commands: for command in commands:
if command.hidden:
continue
try: try:
if await context.can_run(command): if await context.can_run(command):
filtered.append(command) filtered.append(command)
@@ -59,38 +61,33 @@ class HelpCommand(ABC):
return filtered return filtered
async def group_commands(self, context: Context, commands: list[Command]) -> dict[Optional[Cog], list[Command]]: async def group_commands(self, context: Context[ClientT_Co_D], commands: list[Command[ClientT_Co_D]]) -> dict[Optional[Cog[ClientT_Co_D]], list[Command[ClientT_Co_D]]]:
cogs = {} cogs: dict[Optional[Cog[ClientT_Co_D]], list[Command[ClientT_Co_D]]] = {}
for command in commands: for command in commands:
cogs.setdefault(command.cog, []).append(command) cogs.setdefault(command.cog, []).append(command)
return cogs return cogs
async def handle_message(self, context: Context, message: Message): async def handle_message(self, context: Context[ClientT_Co_D], message: Message) -> None:
pass pass
async def get_channel(self, context: Context) -> Messageable: async def get_channel(self, context: Context) -> Messageable:
return context return context
@abstractmethod @abstractmethod
async def handle_no_command_found(self, context: Context, name: str): async def handle_no_command_found(self, context: Context[ClientT_Co_D], name: str) -> Union[str, SendableEmbed, MessagePayload]:
raise NotImplementedError raise NotImplementedError
@abstractmethod class DefaultHelpCommand(HelpCommand[ClientT_Co_D]):
async def handle_no_cog_found(self, context: Context, name: str):
raise NotImplementedError
class DefaultHelpCommand(HelpCommand):
def __init__(self, default_cog_name: str = "No Cog"): def __init__(self, default_cog_name: str = "No Cog"):
self.default_cog_name = default_cog_name self.default_cog_name = default_cog_name
async def create_bot_help(self, context: Context, commands: dict[Optional[Cog], list[Command]]) -> Union[str, SendableEmbed, MessagePayload]: async def create_global_help(self, context: Context[ClientT_Co_D], commands: dict[Optional[Cog[ClientT_Co_D]], list[Command[ClientT_Co_D]]]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"] lines = ["```"]
for cog, cog_commands in commands.items(): for cog, cog_commands in commands.items():
cog_lines = [] cog_lines: list[str] = []
cog_lines.append(f"{cog.qualified_name if cog else self.default_cog_name}:") cog_lines.append(f"{cog.qualified_name if cog else self.default_cog_name}:")
for command in cog_commands: for command in cog_commands:
@@ -101,7 +98,7 @@ class DefaultHelpCommand(HelpCommand):
lines.append("```") lines.append("```")
return "\n".join(lines) return "\n".join(lines)
async def create_cog_help(self, context: Context, cog: Cog) -> Union[str, SendableEmbed, MessagePayload]: async def create_cog_help(self, context: Context[ClientT_Co_D], cog: Cog[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"] lines = ["```"]
lines.append(f"{cog.qualified_name}:") lines.append(f"{cog.qualified_name}:")
@@ -112,7 +109,7 @@ class DefaultHelpCommand(HelpCommand):
lines.append("```") lines.append("```")
return "\n".join(lines) return "\n".join(lines)
async def create_command_help(self, context: Context, command: Command) -> Union[str, SendableEmbed, MessagePayload]: async def create_command_help(self, context: Context[ClientT_Co_D], command: Command[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"] lines = ["```"]
lines.append(f"{command.name}:") lines.append(f"{command.name}:")
@@ -128,7 +125,7 @@ class DefaultHelpCommand(HelpCommand):
lines.append("```") lines.append("```")
return "\n".join(lines) return "\n".join(lines)
async def create_group_help(self, context: Context, group: Group) -> Union[str, SendableEmbed, MessagePayload]: async def create_group_help(self, context: Context[ClientT_Co_D], group: Group[ClientT_Co_D]) -> Union[str, SendableEmbed, MessagePayload]:
lines = ["```"] lines = ["```"]
lines.append(f"{group.name}:") lines.append(f"{group.name}:")
@@ -146,44 +143,65 @@ class DefaultHelpCommand(HelpCommand):
lines.append("```") lines.append("```")
return "\n".join(lines) return "\n".join(lines)
async def handle_no_command_found(self, context: Context, name: str): async def handle_no_command_found(self, context: Context[ClientT_Co_D], name: str) -> str:
channel = await self.get_channel(context) return f"Command `{name}` not found."
await channel.send(f"Command `{name}` not found.")
async def handle_no_cog_found(self, context: Context, name: str): class HelpCommandImpl(Command[ClientT_Co_D]):
channel = await self.get_channel(context) def __init__(self, client: ClientT_Co_D):
await channel.send(f"Cog `{name}` not found.")
class HelpCommandImpl(Command):
def __init__(self, client: CommandsClient):
self.client = client self.client = client
super().__init__(callback=lambda _, context, *args: help_command_impl(self.client, context, *args), name="help", aliases=[])
self.description = "Shows help for a command, cog or the entire bot" async def callback(_: Union[ClientT_Co_D, Cog[ClientT_Co_D]], context: Context[ClientT_Co_D], *args: str) -> None:
await help_command_impl(context.client, context, *args)
super().__init__(callback=callback, name="help", aliases=[])
self.description: str | None = "Shows help for a command, cog or the entire bot"
async def help_command_impl(self: CommandsClient, context: Context, *arguments: str): async def help_command_impl(client: ClientT_D, context: Context[ClientT_D], *arguments: str) -> None:
filtered_commands = await context.client.help_command.filter_commands(context, self.commands) help_command = client.help_command
commands = await self.help_command.group_commands(context, filtered_commands)
if not help_command:
return
filtered_commands = await help_command.filter_commands(context, client.commands)
commands = await help_command.group_commands(context, filtered_commands)
if not arguments: if not arguments:
payload = await self.help_command.create_bot_help(context, commands) payload = await help_command.create_global_help(context, commands)
else:
command_name = arguments[0] else:
parent: ClientT_D | Group[ClientT_D] = client
for param in arguments:
try:
command = parent.get_command(param)
except LookupError:
try:
cog = client.get_cog(param)
except LookupError:
payload = await help_command.handle_no_command_found(context, param)
else:
payload = await help_command.create_cog_help(context, cog)
finally:
break
try:
command = self.get_command(command_name)
except KeyError:
cog = self.cogs.get(command_name)
if cog:
payload = await self.help_command.create_cog_help(context, cog)
else:
return await self.help_command.handle_no_command_found(context, command_name)
else:
if isinstance(command, Group): if isinstance(command, Group):
payload = await self.help_command.create_group_help(context, command) parent = command
else: else:
payload = await self.help_command.create_command_help(context, command) payload = await help_command.create_command_help(context, command)
break
else:
if TYPE_CHECKING:
command = cast(Command[ClientT_D], ...)
if isinstance(command, Group):
payload = await help_command.create_group_help(context, command)
else:
payload = await help_command.create_command_help(context, command)
if TYPE_CHECKING:
payload = cast(MessagePayload, ...)
msg_payload: MessagePayload msg_payload: MessagePayload
@@ -194,4 +212,5 @@ async def help_command_impl(self: CommandsClient, context: Context, *arguments:
else: else:
msg_payload = payload msg_payload = payload
await self.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)
+16 -2
View File
@@ -1,10 +1,24 @@
from __future__ import annotations
from inspect import Parameter from inspect import Parameter
from typing import Any, Iterable from typing import TYPE_CHECKING, Any, Iterable
from typing_extensions import TypeVar
if TYPE_CHECKING:
from .client import CommandsClient
from .context import Context
__all__ = ("evaluate_parameters",) __all__ = ("evaluate_parameters",)
ClientT_Co = TypeVar("ClientT_Co", bound="CommandsClient", covariant=True)
ClientT_D = TypeVar("ClientT_D", bound="CommandsClient", default="CommandsClient")
ClientT_Co_D = TypeVar("ClientT_Co_D", bound="CommandsClient", default="CommandsClient", covariant=True)
ContextT = TypeVar("ContextT", bound="Context", default="Context")
def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]: def evaluate_parameters(parameters: Iterable[Parameter], globals: dict[str, Any]) -> list[Parameter]:
new_parameters = [] new_parameters: list[Parameter] = []
for parameter in parameters: for parameter in parameters:
if parameter.annotation is not parameter.empty: if parameter.annotation is not parameter.empty:
+15 -5
View File
@@ -1,13 +1,16 @@
from typing import Iterator
from typing_extensions import Self
from .errors import NoClosingQuote from .errors import NoClosingQuote
class StringView: class StringView:
def __init__(self, string: str): def __init__(self, string: str):
self.value = iter(string) self.value: Iterator[str] = iter(string)
self.temp = "" self.temp: str = ""
self.should_undo = False self.should_undo: bool = False
def undo(self): def undo(self) -> None:
self.should_undo = True self.should_undo = True
def next_char(self) -> str: def next_char(self) -> str:
@@ -15,7 +18,8 @@ class StringView:
def get_rest(self) -> str: def get_rest(self) -> str:
if self.should_undo: if self.should_undo:
return f"{self.temp} {''.join(self.value)}" return f"{self.temp} {''.join(self.value)}".rstrip()
# prevent a new space appearing at end if the buffer is depleted
return "".join(self.value) return "".join(self.value)
@@ -50,3 +54,9 @@ class StringView:
self.temp = output self.temp = output
return output return output
def __iter__(self) -> Self:
return self
def __next__(self) -> str:
return self.get_next_word()
+9 -6
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import io import io
import os from typing import Optional, Union, cast
from typing import Optional, Union
__all__ = ("File",) __all__ = ("File",)
@@ -19,17 +20,19 @@ class File:
__slots__ = ("f", "spoiler", "filename") __slots__ = ("f", "spoiler", "filename")
def __init__(self, file: Union[str, bytes], *, filename: Optional[str] = None, spoiler: bool = False): def __init__(self, file: Union[str, bytes], *, filename: Optional[str] = None, spoiler: bool = False):
self.f: io.BufferedIOBase
if isinstance(file, str): if isinstance(file, str):
self.f = open(file, "rb") self.f = open(file, "rb")
elif isinstance(file, bytes): else:
self.f = io.BytesIO(file) self.f = io.BytesIO(file)
if filename is None and isinstance(file, str): 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_")): if self.spoiler and (filename and not filename.startswith("SPOILER_")):
filename = f"SPOILER_{filename}" filename = f"SPOILER_{filename}"
self.filename = filename self.filename: str | None = filename
+43 -43
View File
@@ -1,52 +1,56 @@
from __future__ import annotations from __future__ import annotations
from typing import Callable, Iterator, Optional, TypeVar, Union, overload from typing import Callable, Iterator, Optional, Union, overload
__all__ = ("flag_value", "Flags", "UserBadges") from typing_extensions import Self
F_T = TypeVar("F_T", bound="Flags") __all__ = ("Flag", "Flags", "UserBadges")
F_V = TypeVar("F_V", bound="flag_value")
class flag_value: class Flag:
__slots__ = ("flag", "__doc__") __slots__ = ("flag", "__doc__")
def __init__(self, func: Callable[[], int]): def __init__(self, func: Callable[[], int]):
self.flag = func() self.flag: int = func()
self.__doc__ = func.__doc__ self.__doc__: str | None = func.__doc__
@overload @overload
def __get__(self: F_V, instance: None, owner: type[F_T]) -> F_V: def __get__(self: Self, instance: None, owner: type[Flags]) -> Self:
... ...
@overload @overload
def __get__(self, instance: F_T, owner: type[F_T]) -> bool: def __get__(self, instance: Flags, owner: type[Flags]) -> bool:
... ...
def __get__(self: F_V, instance: Optional[F_T], owner: type[F_T]) -> Union[F_V, bool]: def __get__(self: Self, instance: Optional[Flags], owner: type[Flags]) -> Union[Self, bool]:
if instance is None: if instance is None:
return self return self
return instance._check_flag(self.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) instance._set_flag(self.flag, value)
class Flags: class Flags:
FLAG_NAMES: list[str] FLAG_NAMES: list[str]
def __init_subclass__(cls) -> None: def __init_subclass__(cls) -> None:
flags = cls._flags() cls.FLAG_NAMES = []
cls.FLAG_NAMES = list(flags.keys())
def __init__(self, value: int = 0, **kwargs: bool): 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 self.value = value
for k, v in kwargs.items(): for k, v in flags.items():
setattr(self, k, v) setattr(self, k, v)
@classmethod @classmethod
def _from_value(cls: type[F_T], value: int) -> F_T: def _from_value(cls, value: int) -> Self:
self = cls.__new__(cls) self = cls.__new__(cls)
self.value = value self.value = value
return self return self
@@ -54,103 +58,99 @@ class Flags:
def _check_flag(self, flag: int) -> bool: def _check_flag(self, flag: int) -> bool:
return (self.value & flag) == flag return (self.value & flag) == flag
def _set_flag(self, flag: int, value: bool): def _set_flag(self, flag: int, value: bool) -> None:
if value: if value:
self.value |= flag self.value |= flag
else: else:
self.value &= ~flag self.value &= ~flag
def __eq__(self: F_T, other: F_T) -> bool: def __eq__(self, other: Self) -> bool:
return self.value == other.value return self.value == other.value
def __ne__(self: F_T, other: F_T) -> bool: def __ne__(self, other: Self) -> bool:
return not self.__eq__(other) return not self.__eq__(other)
def __or__(self: F_T, other: F_T) -> F_T: def __or__(self, other: Self) -> Self:
return self.__class__._from_value(self.value | other.value) return self.__class__._from_value(self.value | other.value)
def __and__(self: F_T, other: F_T) -> F_T: def __and__(self, other: Self) -> Self:
return self.__class__._from_value(self.value & other.value) return self.__class__._from_value(self.value & other.value)
def __invert__(self: F_T) -> F_T: def __invert__(self) -> Self:
return self.__class__._from_value(~self.value) return self.__class__._from_value(~self.value)
def __add__(self: F_T, other: F_T) -> F_T: def __add__(self, other: Self) -> Self:
return self | other return self | other
def __sub__(self: F_T, other: F_T) -> F_T: def __sub__(self, other: Self) -> Self:
return self & ~other return self & ~other
def __lt__(self: F_T, other: F_T) -> bool: def __lt__(self, other: Self) -> bool:
return self.value < other.value return self.value < other.value
def __gt__(self: F_T, other: F_T) -> bool: def __gt__(self, other: Self) -> bool:
return self.value > other.value return self.value > other.value
def __repr__(self): def __repr__(self) -> str:
return f"<{self.__class__.__name__} value={self.value}>" return f"<{self.__class__.__name__} value={self.value}>"
def __iter__(self) -> Iterator[tuple[str, bool]]: def __iter__(self) -> Iterator[tuple[str, bool]]:
for name, value in self.__class__.__dict__.items(): for name, value in self.__class__.__dict__.items():
if isinstance(value, flag_value): if isinstance(value, Flag):
yield name, value.__get__(self, self.__class__) yield name, self._check_flag(value.flag)
def __hash__(self) -> int: def __hash__(self) -> int:
return hash(self.value) return hash(self.value)
@classmethod
def _flags(cls) -> dict[str, flag_value]:
return {name: value for name, value in cls.__dict__.items() if isinstance(value, flag_value)}
class UserBadges(Flags): class UserBadges(Flags):
"""Contains all user badges""" """Contains all user badges"""
@flag_value @Flag
def developer(): def developer():
""":class:`bool` The developer badge.""" """:class:`bool` The developer badge."""
return 1 << 0 return 1 << 0
@flag_value @Flag
def translator(): def translator():
""":class:`bool` The translator badge.""" """:class:`bool` The translator badge."""
return 1 << 1 return 1 << 1
@flag_value @Flag
def supporter(): def supporter():
""":class:`bool` The supporter badge.""" """:class:`bool` The supporter badge."""
return 1 << 2 return 1 << 2
@flag_value @Flag
def responsible_disclosure(): def responsible_disclosure():
""":class:`bool` The responsible disclosure badge.""" """:class:`bool` The responsible disclosure badge."""
return 1 << 3 return 1 << 3
@flag_value @Flag
def founder(): def founder():
""":class:`bool` The founder badge.""" """:class:`bool` The founder badge."""
return 1 << 4 return 1 << 4
@flag_value @Flag
def platform_moderation(): def platform_moderation():
""":class:`bool` The platform moderation badge.""" """:class:`bool` The platform moderation badge."""
return 1 << 5 return 1 << 5
@flag_value @Flag
def active_supporter(): def active_supporter():
""":class:`bool` The active supporter badge.""" """:class:`bool` The active supporter badge."""
return 1 << 6 return 1 << 6
@flag_value @Flag
def paw(): def paw():
""":class:`bool` The paw badge.""" """:class:`bool` The paw badge."""
return 1 << 7 return 1 << 7
@flag_value @Flag
def early_adopter(): def early_adopter():
""":class:`bool` The early adopter badge.""" """:class:`bool` The early adopter badge."""
return 1 << 8 return 1 << 8
@flag_value @Flag
def reserved_relevant_joke_badge_1(): def reserved_relevant_joke_badge_1():
""":class:`bool` The reserved relevant joke badge 1 badge.""" """:class:`bool` The reserved relevant joke badge 1 badge."""
return 1 << 9 return 1 << 9
+72 -40
View File
@@ -6,9 +6,8 @@ from typing import (TYPE_CHECKING, Any, Coroutine, Literal, Optional, TypeVar,
import aiohttp import aiohttp
import ulid import ulid
from revolt.utils import Missing
from .errors import HTTPError, ServerError from .errors import Forbidden, HTTPError, ServerError
from .file import File from .file import File
try: try:
@@ -21,20 +20,18 @@ if TYPE_CHECKING:
from .enums import SortType from .enums import SortType
from .file import File from .file import File
from .types import ApiInfo
from .types import Autumn as AutumnPayload from .types import Autumn as AutumnPayload
from .types import Channel, DMChannel from .types import Emoji as EmojiPayload
from .types import Embed as EmbedPayload from .types import Interactions as InteractionsPayload
from .types import GetServerMembers, GroupDMChannel, Invite
from .types import Masquerade as MasqueradePayload from .types import Masquerade as MasqueradePayload
from .types import Member from .types import Member as MemberPayload
from .types import Message as MessagePayload from .types import Message as MessagePayload
from .types import (MessageReplyPayload, MessageWithUserData,
PartialInvite, Role)
from .types import SendableEmbed as SendableEmbedPayload from .types import SendableEmbed as SendableEmbedPayload
from .types import Server, ServerBans, TextChannel
from .types import User as UserPayload from .types import User as UserPayload
from .types import UserProfile, VoiceChannel from .types import (Server, ServerBans, TextChannel, UserProfile, VoiceChannel, Member, Invite, ApiInfo, Channel, SavedMessages,
DMChannel, EmojiParent, GetServerMembers, GroupDMChannel, MessageReplyPayload, MessageWithUserData, PartialInvite, CreateRole)
from aiohttp.client import _RequestOptions
__all__ = ("HttpClient",) __all__ = ("HttpClient",)
@@ -45,16 +42,16 @@ class HttpClient:
__slots__ = ("session", "token", "api_url", "api_info", "auth_header") __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): def __init__(self, session: aiohttp.ClientSession, token: str, api_url: str, api_info: ApiInfo, bot: bool = True):
self.session = session self.session: aiohttp.ClientSession = session
self.token = token self.token: str = token
self.api_url = api_url self.api_url: str = api_url
self.api_info = api_info self.api_info: ApiInfo = api_info
self.auth_header = "x-bot-token" if bot else "x-session-token" 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: 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}" url = f"{self.api_url}{route}"
kwargs = {} kwargs: _RequestOptions = {}
headers = { headers = {
"User-Agent": "Revolt.py (https://github.com/revoltchat/revolt.py)", "User-Agent": "Revolt.py (https://github.com/revoltchat/revolt.py)",
@@ -88,10 +85,12 @@ class HttpClient:
if 200 <= resp_code <= 300: if 200 <= resp_code <= 300:
return response return response
elif resp_code == 401:
raise Forbidden("401: Missing Permissions")
else: else:
raise HTTPError(resp_code) raise HTTPError(resp_code)
async def upload_file(self, file: File, tag: str) -> AutumnPayload: async def upload_file(self, file: File, tag: Literal["attachments", "avatars", "backgrounds", "icons", "banners", "emojis"]) -> AutumnPayload:
url = f"{self.api_info['features']['autumn']['url']}/{tag}" url = f"{self.api_info['features']['autumn']['url']}/{tag}"
headers = { headers = {
@@ -113,7 +112,7 @@ class HttpClient:
else: else:
return response return response
async def send_message(self, channel: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]], attachments: Optional[list[File]], replies: Optional[list[MessageReplyPayload]], masquerade: Optional[MasqueradePayload]) -> MessagePayload: async def send_message(self, channel: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]], attachments: Optional[list[File]], replies: Optional[list[MessageReplyPayload]], masquerade: Optional[MasqueradePayload], interactions: Optional[InteractionsPayload]) -> MessagePayload:
json: dict[str, Any] = {} json: dict[str, Any] = {}
if content: if content:
@@ -137,10 +136,13 @@ class HttpClient:
if masquerade: if masquerade:
json["masquerade"] = masquerade json["masquerade"] = masquerade
if interactions:
json["interactions"] = interactions
return await self.request("POST", f"/channels/{channel}/messages", json=json) return await self.request("POST", f"/channels/{channel}/messages", json=json)
def edit_message(self, channel: str, message: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]] = None) -> Request[None]: def edit_message(self, channel: str, message: str, content: Optional[str], embeds: Optional[list[SendableEmbedPayload]] = None) -> Request[None]:
json = {} json: dict[str, Any] = {}
if content is not None: if content is not None:
json["content"] = content json["content"] = content
@@ -196,7 +198,7 @@ class HttpClient:
include_users: bool = False include_users: bool = False
) -> Request[Union[list[MessagePayload], MessageWithUserData]]: ) -> 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: if limit:
json["limit"] = limit json["limit"] = limit
@@ -252,7 +254,7 @@ class HttpClient:
include_users: bool = False include_users: bool = False
) -> Request[Union[list[MessagePayload], MessageWithUserData]]: ) -> Request[Union[list[MessagePayload], MessageWithUserData]]:
json = {"query": query, "include_users": include_users} json: dict[str, Any] = {"query": query, "include_users": include_users}
if limit: if limit:
json["limit"] = limit json["limit"] = limit
@@ -284,7 +286,7 @@ class HttpClient:
def fetch_dm_channels(self) -> Request[list[Union[DMChannel, GroupDMChannel]]]: def fetch_dm_channels(self) -> Request[list[Union[DMChannel, GroupDMChannel]]]:
return self.request("GET", "/users/dms") return self.request("GET", "/users/dms")
def open_dm(self, user_id: str) -> Request[DMChannel]: def open_dm(self, user_id: str) -> Request[DMChannel | SavedMessages]:
return self.request("GET", f"/users/{user_id}/dm") return self.request("GET", f"/users/{user_id}/dm")
def fetch_channel(self, channel_id: str) -> Request[Channel]: def fetch_channel(self, channel_id: str) -> Request[Channel]:
@@ -341,7 +343,7 @@ class HttpClient:
def fetch_bans(self, server_id: str) -> Request[ServerBans]: def fetch_bans(self, server_id: str) -> Request[ServerBans]:
return self.request("GET", f"/servers/{server_id}/bans") return self.request("GET", f"/servers/{server_id}/bans")
def create_role(self, server_id: str, name: str) -> Request[Role]: def create_role(self, server_id: str, name: str) -> Request[CreateRole]:
return self.request("POST", f"/servers/{server_id}/roles", json={"name": name}, nonce=False) return self.request("POST", f"/servers/{server_id}/roles", json={"name": name}, nonce=False)
def delete_role(self, server_id: str, role_id: str) -> Request[None]: def delete_role(self, server_id: str, role_id: str) -> Request[None]:
@@ -353,19 +355,19 @@ class HttpClient:
def delete_invite(self, code: str) -> Request[None]: def delete_invite(self, code: str) -> Request[None]:
return self.request("DELETE", f"/invites/{code}") return self.request("DELETE", f"/invites/{code}")
def edit_channel(self, channel_id: str, remove: Optional[str], values: dict[str, Any]): def edit_channel(self, channel_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove: if remove:
values["remove"] = remove values["remove"] = remove
return self.request("PATCH", f"/channels/{channel_id}", json=values) return self.request("PATCH", f"/channels/{channel_id}", json=values)
def edit_role(self, server_id: str, role_id: str, remove: Optional[str], values: dict[str, Any]): def edit_role(self, server_id: str, role_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove: if remove:
values["remove"] = remove values["remove"] = remove
return self.request("PATCH", f"/servers/{server_id}/roles/{role_id}", json=values) return self.request("PATCH", f"/servers/{server_id}/roles/{role_id}", json=values)
async def edit_self(self, remove: Optional[str], values: dict[str, Any]): async def edit_self(self, remove: list[str] | None, values: dict[str, Any]) -> Request[None]:
if remove: if remove:
values["remove"] = remove values["remove"] = remove
@@ -378,25 +380,55 @@ class HttpClient:
asset = await self.upload_file(background, "backgrounds") asset = await self.upload_file(background, "backgrounds")
profile["background"] = asset["id"] profile["background"] = asset["id"]
if not values.get("profile", Missing):
del values["profile"]
if not values.get("status", Missing):
del values["status"]
return await self.request("PATCH", "/users/@me", json=values) return await self.request("PATCH", "/users/@me", json=values)
def set_guild_channel_default_permissions(self, channel_id: str, allow: int, deny: int) -> Request: def set_guild_channel_default_permissions(self, channel_id: str, allow: int, deny: int) -> Request[None]:
return self.request("PUT", f"/channels/{channel_id}/permissions/default", json={"permissions": {"allow": allow, "deny": deny}}) return self.request("PUT", f"/channels/{channel_id}/permissions/default", json={"permissions": {"allow": allow, "deny": deny}})
def set_guild_channel_role_permissions(self, channel_id: str, role_id: str, allow: int, deny: int) -> Request: 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}}) 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}) 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"/server/{server_id}/permissions/{role_id}", json={"permissions": {"allow": allow, "deny": deny}}) 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"/server/{server_id}/permissions/default", json={"permissions": value}) 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) -> 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) -> Request[None]:
parameters: dict[str, str] = {}
if user_id:
parameters["user_id"] = user_id
parameters["remove_all"] = "true" if remove_all else "false"
return self.request("DELETE", f"/channels/{channel_id}/messages/{message_id}/reactions/{emoji}", params=parameters)
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) -> Request[None]:
return self.request("DELETE", f"/custom/emoji/{emoji_id}")
def fetch_emoji(self, emoji_id: str) -> Request[EmojiPayload]:
return self.request("GET", f"/custom/emoji/{emoji_id}")
async def create_emoji(self, name: str, file: File, nsfw: bool, parent: EmojiParent) -> EmojiPayload:
asset = await self.upload_file(file, "emojis")
return await self.request("PUT", f"/custom/emoji/{asset['id']}", json={"name": name, "parent": parent, "nsfw": nsfw})
def edit_member(self, server_id: str, member_id: str, remove: list[str] | None, values: dict[str, Any]) -> Request[MemberPayload]:
if remove:
values["remove"] = remove
return self.request("PATCH", f"/servers/{server_id}/members/{member_id}", json=values)
def delete_messages(self, channel_id: str, messages: list[str]) -> Request[None]:
return self.request("DELETE", f"/channels/{channel_id}/messages/bulk", json={"ids": messages})
+23 -9
View File
@@ -2,21 +2,29 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .asset import Asset from .asset import Asset
from .utils import Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from .state import State from .state import State
from .channel import Channel
from .server import Server
from .types import Invite as InvitePayload from .types import Invite as InvitePayload
from .user import User
__all__ = ("Invite",) __all__ = ("Invite",)
class Invite: class Invite(Ulid):
"""Represents a server invite. """Represents a server invite.
Attributes Attributes
----------- -----------
code: :class:`str` code: :class:`str`
The code for the invite The code for the invite
id: :class:`str`
Alias for :attr:`code`
server: :class:`Server` server: :class:`Server`
The server this invite is for The server this invite is for
channel: :class:`Channel` channel: :class:`Channel`
@@ -30,22 +38,28 @@ class Invite:
member_count: :class:`int` member_count: :class:`int`
The member count of the server this invite is for The member count of the server this invite is for
""" """
__slots__ = ("state", "code", "id", "server", "channel", "user_name", "user_avatar", "user", "member_count")
def __init__(self, data: InvitePayload, code: str, state: State): def __init__(self, data: InvitePayload, code: str, state: State):
self.state = state self.state: State = state
self.code = code self.code: str = code
self.server = state.get_server(data["server_id"]) self.id: str = code
self.channel = self.server.get_channel(data["channel_id"]) 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_name: str = data["user_name"]
self.user = None self.user: User | None = None
self.user_avatar: Asset | None
if avatar := data.get("user_avatar"): if avatar := data.get("user_avatar"):
self.user_avatar = Asset(avatar, state) self.user_avatar = Asset(avatar, state)
else: else:
self.user_avatar = None self.user_avatar = None
self.member_count = data["member_count"] self.member_count: int = data["member_count"]
@staticmethod @staticmethod
def _from_partial(code: str, server: str, creator: str, channel: str, state: State) -> Invite: def _from_partial(code: str, server: str, creator: str, channel: str, state: State) -> Invite:
@@ -62,6 +76,6 @@ class Invite:
return invite return invite
async def delete(self): async def delete(self) -> None:
"""Deletes the invite""" """Deletes the invite"""
await self.state.http.delete_invite(self.code) await self.state.http.delete_invite(self.code)
+169 -17
View File
@@ -1,19 +1,28 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Optional import datetime
from typing import TYPE_CHECKING, Any, Optional
from .utils import _Missing, Missing, parse_timestamp
from .asset import Asset from .asset import Asset
from .permissions import Permissions
from .permissions_calculator import calculate_permissions
from .user import User from .user import User
from .file import File
if TYPE_CHECKING: if TYPE_CHECKING:
from .channel import Channel
from .server import Server from .server import Server
from .state import State from .state import State
from .types import File from .types import File as FilePayload
from .types import Member as MemberPayload from .types import Member as MemberPayload
from .role import Role
__all__ = ("Member",) __all__ = ("Member",)
def flattern_user(member: Member, user: User): def flattern_user(member: Member, user: User) -> None:
for attr in user.__flattern_attributes__: for attr in user.__flattern_attributes__:
setattr(member, attr, getattr(user, attr)) setattr(member, attr, getattr(user, attr))
@@ -31,17 +40,18 @@ class Member(User):
guild_avatar: Optional[:class:`Asset`] guild_avatar: Optional[:class:`Asset`]
The member's guild avatar if any The member's guild avatar if any
""" """
__slots__ = ("_state", "nickname", "roles", "server", "guild_avatar") __slots__ = ("state", "nickname", "roles", "server", "guild_avatar", "joined_at", "current_timeout")
def __init__(self, data: MemberPayload, server: Server, state: State): def __init__(self, data: MemberPayload, server: Server, state: State):
user = state.get_user(data["_id"]["user"]) user = state.get_user(data["_id"]["user"])
# due to not having a user payload and only a user object we have to manually add all the attributes instead of calling User.__init__ # 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) flattern_user(self, user)
user._members.append(self) user._members[server.id] = self
self._state = state self.state: State = state
self.guild_avatar: Asset | None
if avatar := data.get("avatar"): if avatar := data.get("avatar"):
self.guild_avatar = Asset(avatar, state) self.guild_avatar = Asset(avatar, state)
@@ -49,37 +59,60 @@ class Member(User):
self.guild_avatar = None self.guild_avatar = None
roles = [server.get_role(role_id) for role_id in data.get("roles", [])] 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.server: Server = server
self.nickname = data.get("nickname") self.nickname: str | None = data.get("nickname")
self.joined_at: datetime.datetime = parse_timestamp(data["joined_at"])
self.current_timeout: datetime.datetime | None
if current_timeout := data.get("timeout"):
self.current_timeout = parse_timestamp(current_timeout)
else:
self.current_timeout = None
@property @property
def avatar(self) -> Optional[Asset]: def avatar(self) -> Optional[Asset]:
"""Optional[:class:`Asset`] The avatar the member is displaying, this includes guild avatars and masqueraded avatar""" """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 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 @property
def mention(self) -> str: def mention(self) -> str:
""":class:`str`: Returns a string that allows you to mention the given member.""" """:class:`str`: Returns a string that allows you to mention the given member."""
return f"<@{self.id}>" return f"<@{self.id}>"
def _update(self, *, nickname: Optional[str] = None, avatar: Optional[File] = None, roles: Optional[list[str]] = None): def _update(
if nickname: self,
*,
nickname: Optional[str] = None,
avatar: Optional[FilePayload] = None,
roles: Optional[list[str]] = None,
timeout: Optional[str | int] = None
) -> None:
if nickname is not None:
self.nickname = nickname self.nickname = nickname
if avatar: if avatar is not None:
self.guild_avatar = Asset(avatar, self.state) self.guild_avatar = Asset(avatar, self.state)
if roles: if roles is not None:
member_roles = [self.server.get_role(role_id) for role_id in roles] member_roles = [self.server.get_role(role_id) for role_id in roles]
self.roles = sorted(member_roles, key=lambda role: role.rank, reverse=True) self.roles = sorted(member_roles, key=lambda role: role.rank, reverse=True)
async def kick(self): if timeout is not None:
self.current_timeout = parse_timestamp(timeout)
async def kick(self) -> None:
"""Kicks the member from the server""" """Kicks the member from the server"""
await self.state.http.kick_member(self.server.id, self.id) 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 """Bans the member from the server
Parameters Parameters
@@ -89,6 +122,125 @@ class Member(User):
""" """
await self.state.http.ban_member(self.server.id, self.id, reason) 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""" """Unbans the member from the server"""
await self.state.http.unban_member(self.server.id, self.id) await self.state.http.unban_member(self.server.id, self.id)
async def edit(
self,
*,
nickname: str | None | _Missing = Missing,
roles: list[Role] | None | _Missing = Missing,
avatar: File | None | _Missing = Missing,
timeout: datetime.timedelta | None | _Missing = Missing
) -> None:
"""Edits the member
Parameters
-----------
nickname: Union[:class:`str`, :class:`None`]
The new nickname, or :class:`None` to reset it
roles: Union[list[:class:`Role`], :class:`None`]
The new roles for the member, or :class:`None` to clear it
avatar: Union[:class:`File`, :class:`None`]
The new server avatar, or :class:`None` to reset it
timeout: Union[:class:`datetime.timedelta`, :class:`None`]
The new timeout length for the member, or :class:`None` to reset it
"""
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 not isinstance(roles, _Missing):
data["roles"] = [role.id for role in roles]
if avatar is None:
remove.append("Avatar")
elif not isinstance(avatar, _Missing):
data["avatar"] = (await self.state.http.upload_file(avatar, "avatars"))["id"]
if timeout is None:
remove.append("Timeout")
elif not isinstance(timeout, _Missing):
data["timeout"] = (datetime.datetime.now(datetime.timezone.utc) + timeout).isoformat()
await self.state.http.edit_member(self.server.id, self.id, remove, data)
async def timeout(self, length: datetime.timedelta) -> None:
"""Timeouts the member
Parameters
-----------
length: :class:`datetime.timedelta`
The length of the timeout
"""
ends_at = datetime.datetime.now(tz=datetime.timezone.utc) + length
await self.state.http.edit_member(self.server.id, self.id, None, {"timeout": ends_at.isoformat()})
def get_permissions(self) -> Permissions:
"""Gets the permissions for the member in the server
Returns
--------
:class:`Permissions`
The members permissions
"""
return calculate_permissions(self, self.server)
def get_channel_permissions(self, channel: Channel) -> Permissions:
"""Gets the permissions for the member in the server taking into account the channel as well
Parameters
-----------
channel: :class:`Channel`
The channel to calculate permissions with
Returns
--------
:class:`Permissions`
The members permissions
"""
return calculate_permissions(self, channel)
def has_permissions(self, **permissions: bool) -> bool:
"""Computes if the member has the specified permissions
Parameters
-----------
permissions: :class:`bool`
The permissions to check, this also accepted `False` if you need to check if the member does not have the permission
Returns
--------
:class:`bool`
Whether or not they have the permissions
"""
calculated_perms = self.get_permissions()
return all([getattr(calculated_perms, key, False) == value for key, value in permissions.items()])
def has_channel_permissions(self, channel: Channel, **permissions: bool) -> bool:
"""Computes if the member has the specified permissions, taking into account the channel as well
Parameters
-----------
channel: :class:`Channel`
The channel to calculate permissions with
permissions: :class:`bool`
The permissions to check, this also accepted `False` if you need to check if the member does not have the permission
Returns
--------
:class:`bool`
Whether or not they have the permissions
"""
calculated_perms = self.get_channel_permissions(channel)
return all([getattr(calculated_perms, key, False) == value for key, value in permissions.items()])
+175 -48
View File
@@ -1,27 +1,33 @@
from __future__ import annotations from __future__ import annotations
import datetime import datetime
from typing import TYPE_CHECKING, NamedTuple, Optional from typing import TYPE_CHECKING, Any, Coroutine, Optional, Union
from .asset import Asset, PartialAsset from .asset import Asset, PartialAsset
from .channel import Messageable from .channel import DMChannel, GroupDMChannel, TextChannel, SavedMessageChannel
from .embed import SendableEmbed, to_embed from .embed import Embed, SendableEmbed, to_embed
from .utils import Ulid, parse_timestamp
if TYPE_CHECKING: if TYPE_CHECKING:
from .server import Server
from .state import State from .state import State
from .types import Embed as EmbedPayload from .types import Embed as EmbedPayload
from .types import Interactions as InteractionsPayload
from .types import Masquerade as MasqueradePayload from .types import Masquerade as MasqueradePayload
from .types import Message as MessagePayload from .types import Message as MessagePayload
from .types import MessageReplyPayload from .types import MessageReplyPayload, SystemMessageContent
from .server import Server from .user import User
from .member import Member
__all__ = ( __all__ = (
"Message", "Message",
"MessageReply", "MessageReply",
"Masquerade" "Masquerade",
"MessageInteractions"
) )
class Message: class Message(Ulid):
"""Represents a message """Represents a message
Attributes Attributes
@@ -40,33 +46,50 @@ class Message:
The author of the message, will be :class:`User` in DMs The author of the message, will be :class:`User` in DMs
edited_at: Optional[:class:`datetime.datetime`] edited_at: Optional[:class:`datetime.datetime`]
The time at which the message was edited, will be None if the message has not been edited The time at which the message was edited, will be None if the message has not been edited
mentions: list[Union[:class:`Member`, :class:`User`]] raw_mentions: list[:class:`str`]
The users or members that where mentioned in the message A list of ids of the mentions in this message
replies: list[:class:`Message`] replies: list[:class:`Message`]
The message's this message has replied to, this may not contain all the messages if they are outside the cache The message's this message has replied to, this may not contain all the messages if they are outside the cache
reply_ids: list[:class:`str`] reply_ids: list[:class:`str`]
The message's ids this message has replies to The message's ids this message has replies to
reactions: dict[str, list[:class:`User`]]
The reactions on the message
interactions: Optional[:class:`MessageInteractions`]
The interactions on the message, if any
""" """
__slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "mentions", "replies", "reply_ids") __slots__ = ("state", "id", "content", "attachments", "embeds", "channel", "author", "edited_at", "replies", "reply_ids", "reactions", "interactions")
def __init__(self, data: MessagePayload, state: State): def __init__(self, data: MessagePayload, state: State):
self.state = state self.state: State = state
self.id = data["_id"] self.id: str = data["_id"]
self.content = data["content"] self.content: str = 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.system_content: SystemMessageContent | None = data.get("system")
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"]) channel = state.get_channel(data["channel"])
assert isinstance(channel, Messageable) assert isinstance(channel, (TextChannel, GroupDMChannel, DMChannel, SavedMessageChannel))
self.channel = channel self.channel: TextChannel | GroupDMChannel | DMChannel | SavedMessageChannel = channel
if server_id := self.channel.server_id: self.server_id: str | None = self.channel.server_id
author = state.get_member(server_id, data["author"])
self.raw_mentions: list[str] = data.get("mentions", [])
if self.system_content:
author_id: str = self.system_content.get("id", data["author"])
else: else:
author = state.get_user(data["author"]) author_id = data["author"]
self.author = author if self.server_id:
author = state.get_member(self.server_id, author_id)
else:
author = state.get_user(author_id)
self.author: Member | User = author
if masquerade := data.get("masquerade"): if masquerade := data.get("masquerade"):
if name := masquerade.get("name"): if name := masquerade.get("name"):
@@ -76,41 +99,78 @@ class Message:
self.author.masquerade_avatar = PartialAsset(avatar, state) self.author.masquerade_avatar = PartialAsset(avatar, state)
if edited_at := data.get("edited"): if edited_at := data.get("edited"):
self.edited_at: Optional[datetime.datetime] = datetime.datetime.strptime(edited_at["$date"], "%Y-%m-%dT%H:%M:%S.%f%z") self.edited_at: Optional[datetime.datetime] = parse_timestamp(edited_at)
if self.server: self.replies: list[Message] = []
self.mentions = [self.server.get_member(member_id) for member_id in data.get("mentions", [])] self.reply_ids: list[str] = []
else:
self.mentions = [state.get_user(member_id) for member_id in data.get("mentions", [])]
self.replies = []
self.reply_ids = []
for reply in data.get("replies", []): for reply in data.get("replies", []):
try: try:
message = state.get_message(reply) message = state.get_message(reply)
self.replies.append(message) self.replies.append(message)
except KeyError: except LookupError:
pass pass
self.reply_ids.append(reply) self.reply_ids.append(reply)
def _update(self, *, content: Optional[str] = None, embeds: Optional[list[EmbedPayload]] = None, edited_at: str): reactions = data.get("reactions", {})
if content:
self.reactions: dict[str, list[User]] = {}
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: Optional[Union[str, int]] = None):
if content is not None:
self.content = content self.content = content
self.edited_at = datetime.datetime.strptime(edited_at, "%Y-%m-%dT%H:%M:%S.%f%z") if embeds is not None:
# strptime is used here instead of fromisoformat because of its inability to parse `Z` (Zulu or UTC time) in the RFCC 3339 format provided by API
if embeds:
self.embeds = [to_embed(embed, self.state) for embed in embeds] self.embeds = [to_embed(embed, self.state) for embed in embeds]
if edited is not None:
self.edited_at = parse_timestamp(edited)
@property
def mentions(self) -> list[User | Member]:
"""The users or members that where mentioned in the message
Returns: list[Union[:class:`Member`, :class:`User`]]
"""
mentions: list[User | Member] = []
if self.server_id:
for mention in self.raw_mentions:
try:
mentions.append(self.server.get_member(mention))
except LookupError:
pass
else:
for mention in self.raw_mentions:
try:
mentions.append(self.state.get_user(mention))
except LookupError:
pass
return mentions
async def edit(self, *, content: Optional[str] = None, embeds: Optional[list[SendableEmbed]] = None) -> None: 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 """Edits the message. The bot can only edit its own message
Parameters Parameters
----------- -----------
content: :class:`str` content: :class:`str`
The new content of the message The new content of the message
embeds: list[:class:`SendableEmbed`]
The new embeds of the message
""" """
new_embeds = [embed.to_dict() for embed in embeds] if embeds else None new_embeds = [embed.to_dict() for embed in embeds] if embeds else None
@@ -121,7 +181,7 @@ class Message:
"""Deletes the message. The bot can only delete its own messages and messages it has permission to delete """ """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) await self.state.http.delete_message(self.channel.id, self.id)
def reply(self, *args, mention: bool = False, **kwargs): def reply(self, *args: Any, mention: bool = False, **kwargs: Any) -> Coroutine[Any, Any, Message]:
"""Replies to this message, equivilant to: """Replies to this message, equivilant to:
.. code-block:: python .. code-block:: python
@@ -131,13 +191,47 @@ class Message:
""" """
return self.channel.send(*args, **kwargs, replies=[MessageReply(self, mention)]) return self.channel.send(*args, **kwargs, replies=[MessageReply(self, mention)])
async def add_reaction(self, emoji: str) -> None:
"""Adds a reaction to the message
Parameters
-----------
emoji: :class:`str`
The emoji to add as a reaction
"""
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) -> None:
"""Removes a reaction from the message, this can remove either a specific users, the current users reaction or all of a specific emoji
Parameters
-----------
emoji: :class:`str`
The emoji to remove
user: Optional[:class:`User`]
The user to use for removing a reaction from
remove_all: bool
Whether or not to remove all reactions for that specific emoji
"""
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) -> None:
"""Removes all reactions from the message"""
await self.state.http.remove_all_reactions(self.channel.id, self.id)
@property @property
def server(self) -> Server: 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 return self.channel.server
class MessageReply(NamedTuple): class MessageReply:
"""A namedtuple which represents a reply to a message. """represents a reply to a message.
Parameters Parameters
----------- -----------
@@ -146,14 +240,17 @@ class MessageReply(NamedTuple):
mention: :class:`bool` mention: :class:`bool`
Whether the reply should mention the author of the message. Defaults to false. Whether the reply should mention the author of the message. Defaults to false.
""" """
message: Message __slots__ = ("message", "mention")
mention: bool = False
def __init__(self, message: Ulid, mention: bool = False):
self.message: Ulid = message
self.mention: bool = mention
def to_dict(self) -> MessageReplyPayload: def to_dict(self) -> MessageReplyPayload:
return { "id": self.message.id, "mention": self.mention } return {"id": self.message.id, "mention": self.mention}
class Masquerade(NamedTuple): class Masquerade:
"""A namedtuple which represents a message's masquerade. """represents a message's masquerade.
Parameters Parameters
----------- -----------
@@ -164,9 +261,12 @@ class Masquerade(NamedTuple):
colour: Optional[:class:`str`] colour: Optional[:class:`str`]
The colour of the name, similar to role colours The colour of the name, similar to role colours
""" """
name: Optional[str] = None __slots__ = ("name", "avatar", "colour")
avatar: Optional[str] = None
colour: Optional[str] = None def __init__(self, name: Optional[str] = None, avatar: Optional[str] = None, colour: Optional[str] = None):
self.name: str | None = name
self.avatar: str | None = avatar
self.colour: str | None = colour
def to_dict(self) -> MasqueradePayload: def to_dict(self) -> MasqueradePayload:
output: MasqueradePayload = {} output: MasqueradePayload = {}
@@ -181,3 +281,30 @@ class Masquerade(NamedTuple):
output["colour"] = colour output["colour"] = colour
return output return output
class MessageInteractions:
"""Represents a message's interactions, this is for allowing preset reactions and restricting adding reactions to only those.
Parameters
-----------
reactions: Optional[list[:class:`str`]]
The preset reactions on the message
restrict_reactions: bool
Whether or not users can only react to the interaction's reactions
"""
__slots__ = ("reactions", "restrict_reactions")
def __init__(self, *, reactions: Optional[list[str]] = None, restrict_reactions: bool = False):
self.reactions: list[str] | None = reactions
self.restrict_reactions: bool = restrict_reactions
def to_dict(self) -> InteractionsPayload:
output: InteractionsPayload = {}
if reactions := self.reactions:
output["reactions"] = reactions
if restrict_reactions := self.restrict_reactions:
output["restrict_reactions"] = restrict_reactions
return output
+49 -8
View File
@@ -5,10 +5,11 @@ from typing import TYPE_CHECKING, Optional
from .enums import SortType from .enums import SortType
if TYPE_CHECKING: if TYPE_CHECKING:
from .embed import Embed, SendableEmbed from .embed import SendableEmbed
from .file import File from .file import File
from .message import Masquerade, Message, MessageReply from .message import Masquerade, Message, MessageInteractions, MessageReply
from .state import State from .state import State
from .types.http import MessageWithUserData
__all__ = ("Messageable",) __all__ = ("Messageable",)
@@ -28,7 +29,7 @@ class Messageable:
async def _get_channel_id(self) -> str: async def _get_channel_id(self) -> str:
raise NotImplementedError raise NotImplementedError
async def send(self, content: Optional[str] = None, *, embeds: Optional[list[SendableEmbed]] = None, embed: Optional[SendableEmbed] = None, attachments: Optional[list[File]] = None, replies: Optional[list[MessageReply]] = None, reply: Optional[MessageReply] = None, masquerade: Optional[Masquerade] = None) -> Message: async def send(self, content: Optional[str] = None, *, embeds: Optional[list[SendableEmbed]] = None, embed: Optional[SendableEmbed] = None, attachments: Optional[list[File]] = None, replies: Optional[list[MessageReply]] = None, reply: Optional[MessageReply] = None, masquerade: Optional[Masquerade] = None, interactions: Optional[MessageInteractions] = None) -> Message:
"""Sends a message in a channel, you must send at least one of either `content`, `embeds` or `attachments` """Sends a message in a channel, you must send at least one of either `content`, `embeds` or `attachments`
Parameters Parameters
@@ -43,6 +44,10 @@ class Messageable:
The embeds to send with the message The embeds to send with the message
replies: Optional[list[:class:`MessageReply`]] replies: Optional[list[:class:`MessageReply`]]
The list of messages to reply to. The list of messages to reply to.
masquerade: Optional[:class:`Masquerade`]
The masquerade for the message, this can overwrite the username and avatar shown
interactions: Optional[:class:`MessageInteractions`]
The interactions for the message
Returns Returns
-------- --------
@@ -58,8 +63,9 @@ class Messageable:
embed_payload = [embed.to_dict() for embed in embeds] if embeds else None embed_payload = [embed.to_dict() for embed in embeds] if embeds else None
reply_payload = [reply.to_dict() for reply in replies] if replies else None reply_payload = [reply.to_dict() for reply in replies] if replies else None
masquerade_payload = masquerade.to_dict() if masquerade else None masquerade_payload = masquerade.to_dict() if masquerade else None
interactions_payload = interactions.to_dict() if interactions else None
message = await self.state.http.send_message(await self._get_channel_id(), content, embed_payload, attachments, reply_payload, masquerade_payload) message = await self.state.http.send_message(await self._get_channel_id(), content, embed_payload, attachments, reply_payload, masquerade_payload, interactions_payload)
return self.state.add_message(message) return self.state.add_message(message)
@@ -76,9 +82,23 @@ class Messageable:
:class:`Message` :class:`Message`
The message with the matching id The message with the matching id
""" """
from .message import Message
payload = await self.state.http.fetch_message(await self._get_channel_id(), message_id) payload = await self.state.http.fetch_message(await self._get_channel_id(), message_id)
return Message(payload, self.state) return Message(payload, self.state)
def _add_missing_users(self, payload: MessageWithUserData):
for user in payload["users"]:
if user["_id"] not in self.state.users:
self.state.add_user(user)
if members := payload.get("members", []):
server = self.state.get_server(members[0]["_id"]["server"])
for member in members:
if member["_id"]["user"] not in server._members:
server._add_member(member)
async def history(self, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None, nearby: Optional[str] = None) -> list[Message]: async def history(self, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None, nearby: Optional[str] = None) -> list[Message]:
"""Fetches multiple messages from the channel's history """Fetches multiple messages from the channel's history
@@ -100,8 +120,12 @@ class Messageable:
list[:class:`Message`] list[:class:`Message`]
The messages found in order of the sort parameter The messages found in order of the sort parameter
""" """
payloads = await self.state.http.fetch_messages(await self._get_channel_id(), sort=sort, limit=limit, before=before, after=after, nearby=nearby) from .message import Message
return [Message(payload, self.state) for payload in payloads]
payload = await self.state.http.fetch_messages(await self._get_channel_id(), sort=sort, limit=limit, before=before, after=after, nearby=nearby, include_users=True)
self._add_missing_users(payload)
return [Message(msg, self.state) for msg in payload["messages"]]
async def search(self, query: str, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None) -> list[Message]: async def search(self, query: str, *, sort: SortType = SortType.latest, limit: int = 100, before: Optional[str] = None, after: Optional[str] = None) -> list[Message]:
"""searches the channel for a query """searches the channel for a query
@@ -124,5 +148,22 @@ class Messageable:
list[:class:`Message`] list[:class:`Message`]
The messages found in order of the sort parameter The messages found in order of the sort parameter
""" """
payloads = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after) from .message import Message
return [Message(payload, self.state) for payload in payloads]
payload = await self.state.http.search_messages(await self._get_channel_id(), query, sort=sort, limit=limit, before=before, after=after, include_users=True)
self._add_missing_users(payload)
return [Message(msg, self.state) for msg in payload["messages"]]
async def delete_messages(self, messages: list[Message]) -> None:
"""Bulk deletes messages from the channel
.. note:: The messages must have been sent in the last 7 days.
Parameters
-----------
messages: list[:class:`Message`]
The messages for deletion, this can be up to 100 messages
"""
await self.state.http.delete_messages(await self._get_channel_id(), [message.id for message in messages])
+65 -32
View File
@@ -1,118 +1,145 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional from typing import TYPE_CHECKING, Any, Optional
from typing_extensions import Self from typing_extensions import Self
from .flags import Flag, Flags
from .types.permissions import Overwrite from .types.permissions import Overwrite
from .flags import Flags, flag_value
__all__ = ("Permissions", "PermissionsOverwrite") __all__ = ("Permissions", "PermissionsOverwrite", "UserPermissions")
class UserPermissions(Flags):
"""Permissions for users"""
@Flag
def access() -> int:
return 1 << 0
@Flag
def view_profile() -> int:
return 1 << 1
@Flag
def send_message() -> int:
return 1 << 2
@Flag
def invite() -> int:
return 1 << 3
@classmethod
def all(cls) -> Self:
return cls(access=True, view_profile=True, send_message=True, invite=True)
class Permissions(Flags): class Permissions(Flags):
@flag_value """Server permissions for members and roles"""
@Flag
def manage_channel() -> int: def manage_channel() -> int:
return 1 << 0 return 1 << 0
@flag_value @Flag
def manage_server() -> int: def manage_server() -> int:
return 1 << 1 return 1 << 1
@flag_value @Flag
def manage_permissions() -> int: def manage_permissions() -> int:
return 1 << 2 return 1 << 2
@flag_value @Flag
def manage_role() -> int: def manage_role() -> int:
return 1 << 3 return 1 << 3
@flag_value @Flag
def kick_members() -> int: def kick_members() -> int:
return 1 << 6 return 1 << 6
@flag_value @Flag
def ban_members() -> int: def ban_members() -> int:
return 1 << 7 return 1 << 7
@flag_value @Flag
def timeout_members() -> int: def timeout_members() -> int:
return 1 << 8 return 1 << 8
@flag_value @Flag
def asign_roles() -> int: def asign_roles() -> int:
return 1 << 9 return 1 << 9
@flag_value @Flag
def change_nickname() -> int: def change_nickname() -> int:
return 1 << 10 return 1 << 10
@flag_value @Flag
def manage_nicknames() -> int: def manage_nicknames() -> int:
return 1 << 11 return 1 << 11
@flag_value @Flag
def change_avatars() -> int: def change_avatars() -> int:
return 1 << 12 return 1 << 12
@flag_value @Flag
def remove_avatars() -> int: def remove_avatars() -> int:
return 1 << 13 return 1 << 13
@flag_value @Flag
def view_channel() -> int: def view_channel() -> int:
return 1 << 20 return 1 << 20
@flag_value @Flag
def read_message_history() -> int: def read_message_history() -> int:
return 1 << 21 return 1 << 21
@flag_value @Flag
def send_messages() -> int: def send_messages() -> int:
return 1 << 22 return 1 << 22
@flag_value @Flag
def manage_messages() -> int: def manage_messages() -> int:
return 1 << 23 return 1 << 23
@flag_value @Flag
def manage_webhooks() -> int: def manage_webhooks() -> int:
return 1 << 24 return 1 << 24
@flag_value @Flag
def invite_others() -> int: def invite_others() -> int:
return 1 << 25 return 1 << 25
@flag_value @Flag
def send_embeds() -> int: def send_embeds() -> int:
return 1 << 26 return 1 << 26
@flag_value @Flag
def upload_files() -> int: def upload_files() -> int:
return 1 << 27 return 1 << 27
@flag_value @Flag
def masquerade() -> int: def masquerade() -> int:
return 1 << 28 return 1 << 28
@flag_value @Flag
def connect() -> int: def connect() -> int:
return 1 << 30 return 1 << 30
@flag_value @Flag
def speak() -> int: def speak() -> int:
return 1 << 31 return 1 << 31
@flag_value @Flag
def video() -> int: def video() -> int:
return 1 << 32 return 1 << 32
@flag_value @Flag
def mute_members() -> int: def mute_members() -> int:
return 1 << 33 return 1 << 33
@flag_value @Flag
def deafen_members() -> int: def deafen_members() -> int:
return 1 << 34 return 1 << 34
@flag_value @Flag
def move_members() -> int: def move_members() -> int:
return 1 << 35 return 1 << 35
@@ -128,7 +155,13 @@ class Permissions(Flags):
def default(cls) -> Self: def default(cls) -> Self:
return cls.default_view_only() | cls(send_messages=True, invite_others=True, send_embeds=True, upload_files=True, connect=True, speak=True) return cls.default_view_only() | cls(send_messages=True, invite_others=True, send_embeds=True, upload_files=True, connect=True, speak=True)
@classmethod
def default_direct_message(cls) -> Self:
return cls.default_view_only() | cls(react=True, manage_channel=True)
class PermissionsOverwrite: class PermissionsOverwrite:
"""A permissions overwrite in a channel"""
def __init__(self, allow: Permissions, deny: Permissions): def __init__(self, allow: Permissions, deny: Permissions):
self._allow = allow self._allow = allow
self._deny = deny self._deny = deny
@@ -143,13 +176,13 @@ class PermissionsOverwrite:
super().__setattr__(perm, value) 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 in Permissions.FLAG_NAMES:
if key is True: if value is True:
setattr(self._allow, key, True) setattr(self._allow, key, True)
super().__setattr__(key, True) super().__setattr__(key, True)
elif key is False: elif value is False:
setattr(self._deny, key, True) setattr(self._deny, key, True)
super().__setattr__(key, False) super().__setattr__(key, False)
+82
View File
@@ -0,0 +1,82 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import TYPE_CHECKING, cast
from revolt.enums import ChannelType
from .permissions import Permissions
if TYPE_CHECKING:
from .channel import Channel, DMChannel, GroupDMChannel, ServerChannel
from .member import Member
from .server import Server
def calculate_permissions(member: Member, target: Server | Channel) -> Permissions:
if member.privileged:
return Permissions.all()
from .server import Server
if isinstance(target, Server):
if target.owner_id == member.id:
return Permissions.all()
permissions = target.default_permissions
for role in member.roles:
permissions = (permissions | role.permissions._allow) & (~role.permissions._deny)
if member.current_timeout and member.current_timeout > datetime.now(timezone.utc):
permissions = permissions & Permissions.default_view_only()
return permissions
else:
channel_type = target.channel_type
if channel_type is ChannelType.saved_messages:
return Permissions.all()
elif channel_type is ChannelType.direct_message:
target = cast("DMChannel", target)
user_permissions = target.recipient.get_permissions()
if user_permissions.send_message:
return Permissions.default_direct_message()
else:
return Permissions.default_view_only()
elif channel_type is ChannelType.group:
target = cast("GroupDMChannel", target)
if target.owner.id != member.id:
return Permissions.default_direct_message()
else:
if target.permissions.value == 0:
return Permissions.default_direct_message()
else:
return target.permissions
else:
target = cast("ServerChannel", target)
server = target.server
if server.owner_id == member.id:
return Permissions.all()
else:
perms = calculate_permissions(member, server)
perms = (perms | target.default_permissions._allow) & (~target.default_permissions._deny)
for role in server.roles[::-1]:
if overwrite :=target.permissions.get(role.id):
perms = (perms | overwrite._allow) & (~overwrite._deny)
if member.current_timeout and member.current_timeout > datetime.now():
perms = perms & Permissions(view_channel=True, read_message_history=True)
return perms
+33 -22
View File
@@ -1,9 +1,9 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Any, Optional
from .permissions import Permissions, PermissionsOverwrite from .permissions import Overwrite, PermissionsOverwrite
from .utils import Missing from .utils import Missing, Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from .server import Server from .server import Server
@@ -13,7 +13,7 @@ if TYPE_CHECKING:
__all__ = ("Role",) __all__ = ("Role",)
class Role: class Role(Ulid):
"""Represents a role """Represents a role
Attributes Attributes
@@ -35,20 +35,20 @@ class Role:
channel_permissions: :class:`ChannelPermissions` channel_permissions: :class:`ChannelPermissions`
The channel permissions for the role 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): def __init__(self, data: RolePayload, role_id: str, server: Server, state: State):
self.state = state self.state: State = state
self.id = role_id self.id: str = role_id
self.name = data["name"] self.name: str = data["name"]
self.colour = None self.colour: str | None = data.get("colour", None)
self.hoist = False self.hoist: bool = data.get("hoist", False)
self.rank = 0 self.rank: int = data["rank"]
self.server = server self.server: Server = server
self.permissions = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0})) self.permissions: PermissionsOverwrite = PermissionsOverwrite._from_overwrite(data.get("permissions", {"a": 0, "d": 0}))
@property @property
def color(self): def color(self) -> str | None:
return self.colour return self.colour
async def set_permissions_overwrite(self, *, permissions: PermissionsOverwrite) -> None: async def set_permissions_overwrite(self, *, permissions: PermissionsOverwrite) -> None:
@@ -63,31 +63,42 @@ class Role:
allow, deny = permissions.to_pair() allow, deny = permissions.to_pair()
await self.state.http.set_server_role_permissions(self.server.id, self.id, allow.value, deny.value) 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): 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: if name is not None:
self.name = name self.name = name
if colour: if colour is not None:
self.colour = colour self.colour = colour
if hoist: if hoist is not None:
self.hoist = hoist self.hoist = hoist
if rank: if rank is not None:
self.rank = rank self.rank = rank
async def delete(self): if permissions is not None:
self.permissions = PermissionsOverwrite._from_overwrite(permissions)
async def delete(self) -> None:
"""Deletes the role""" """Deletes the role"""
await self.state.http.delete_role(self.server.id, self.id) await self.state.http.delete_role(self.server.id, self.id)
async def edit(self, **kwargs): async def edit(self, **kwargs: Any) -> None:
"""Edits the role """Edits the role
Parameters Parameters
----------- -----------
name: str
The name of the role
colour: str
The colour of the role
hoist: bool
Whether the role should make the member display seperately in the member list
rank: int
The position of the role
""" """
if kwargs.get("colour", Missing) is None: if kwargs.get("colour", Missing) is None:
remove = "Colour" remove = ["Colour"]
else: else:
remove = None remove = None
+136 -46
View File
@@ -4,34 +4,45 @@ from typing import TYPE_CHECKING, Optional, cast
from .asset import Asset from .asset import Asset
from .category import Category from .category import Category
from .channel import Channel, VoiceChannel
from .invite import Invite from .invite import Invite
from .permissions import Permissions from .permissions import Permissions
from .role import Role from .role import Role
from .utils import Ulid
from .channel import Channel, TextChannel, VoiceChannel
from .member import Member
if TYPE_CHECKING: if TYPE_CHECKING:
from .channel import TextChannel from .emoji import Emoji
from .member import Member from .file import File
from .state import State from .state import State
from .types import Ban from .types import Ban
from .types import Category as CategoryPayload from .types import Category as CategoryPayload
from .types import File as FilePayload from .types import File as FilePayload
from .types import Server as ServerPayload from .types import Server as ServerPayload
from .types import SystemMessagesConfig from .types import SystemMessagesConfig
from .types import Member as MemberPayload
__all__ = ("Server", "SystemMessages", "ServerBan") __all__ = ("Server", "SystemMessages", "ServerBan")
class SystemMessages: class SystemMessages:
"""Holds all the configuration for the server's system message channels"""
def __init__(self, data: SystemMessagesConfig, state: State): def __init__(self, data: SystemMessagesConfig, state: State):
self.state = state self.state: State = state
self.user_joined_id = data.get("user_joined") self.user_joined_id: str | None = data.get("user_joined")
self.user_left_id = data.get("user_left") self.user_left_id: str | None = data.get("user_left")
self.user_kicked_id = data.get("user_kicked") self.user_kicked_id: str | None = data.get("user_kicked")
self.user_banned_id = data.get("user_banned") self.user_banned_id: str | None = data.get("user_banned")
@property @property
def user_joined(self) -> Optional[TextChannel]: def user_joined(self) -> Optional[TextChannel]:
"""The channel which user join messages get sent in
Returns
--------
Optional[:class:`TextChannel`]
The channel
"""
if not self.user_joined_id: if not self.user_joined_id:
return return
@@ -41,6 +52,13 @@ class SystemMessages:
@property @property
def user_left(self) -> Optional[TextChannel]: def user_left(self) -> Optional[TextChannel]:
"""The channel which user leave messages get sent in
Returns
--------
Optional[:class:`TextChannel`]
The channel
"""
if not self.user_left_id: if not self.user_left_id:
return return
@@ -50,6 +68,13 @@ class SystemMessages:
@property @property
def user_kicked(self) -> Optional[TextChannel]: def user_kicked(self) -> Optional[TextChannel]:
"""The channel which user kick messages get sent in
Returns
--------
Optional[:class:`TextChannel`]
The channel
"""
if not self.user_kicked_id: if not self.user_kicked_id:
return return
@@ -59,6 +84,13 @@ class SystemMessages:
@property @property
def user_banned(self) -> Optional[TextChannel]: def user_banned(self) -> Optional[TextChannel]:
"""The channel which user ban messages get sent in
Returns
--------
Optional[:class:`TextChannel`]
The channel
"""
if not self.user_banned_id: if not self.user_banned_id:
return return
@@ -66,7 +98,7 @@ class SystemMessages:
assert isinstance(channel, TextChannel) assert isinstance(channel, TextChannel)
return channel return channel
class Server: class Server(Ulid):
"""Represents a server """Represents a server
Attributes Attributes
@@ -90,24 +122,28 @@ class Server:
default_permissions: :class:`Permissions` default_permissions: :class:`Permissions`
The permissions for the default role The permissions for the default role
""" """
__slots__ = ("state", "id", "name", "owner_id", "default_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "_categories") __slots__ = ("state", "id", "name", "owner_id", "default_permissions", "_members", "_roles", "_channels", "description", "icon", "banner", "nsfw", "system_messages", "_categories", "_emojis")
def __init__(self, data: ServerPayload, state: State): def __init__(self, data: ServerPayload, state: State):
self.state = state self.state: State = state
self.id = data["_id"] self.id: str = data["_id"]
self.name = data["name"] self.name: str = data["name"]
self.owner_id = data["owner"] self.owner_id: str = data["owner"]
self.description = data.get("description") or None self.description: str | None = data.get("description") or None
self.nsfw = data.get("nsfw", False) self.nsfw: bool = data.get("nsfw", False)
self.system_messages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state) self.system_messages: SystemMessages = SystemMessages(data.get("system_messages", cast("SystemMessagesConfig", {})), state)
self._categories = {data["id"]: Category(data, state) for data in data.get("categories", [])} self._categories: dict[str, Category] = {data["id"]: Category(data, state) for data in data.get("categories", [])}
self.default_permissions = Permissions(data["default_permissions"]) self.default_permissions: Permissions = Permissions(data["default_permissions"])
self.icon: Asset | None
if icon := data.get("icon"): if icon := data.get("icon"):
self.icon = Asset(icon, state) self.icon = Asset(icon, state)
else: else:
self.icon = None self.icon = None
self.banner: Asset | None
if banner := data.get("banner"): if banner := data.get("banner"):
self.banner = Asset(banner, state) self.banner = Asset(banner, state)
else: else:
@@ -116,18 +152,27 @@ class Server:
self._members: dict[str, Member] = {} self._members: dict[str, Member] = {}
self._roles: dict[str, Role] = {role_id: Role(role, role_id, self, state) for role_id, role in data.get("roles", {}).items()} self._roles: dict[str, Role] = {role_id: Role(role, role_id, self, state) for role_id, role in data.get("roles", {}).items()}
self._channels: dict[str, Channel] = {channel_id: state.get_channel(channel_id) for channel_id in data.get("channels", [])} self._channels: dict[str, Channel] = {}
# The api doesnt send us all the channels but sends us all the ids, this is because channels we dont have permissions to see are not sent
# this causes get_channel to error so we have to first check ourself if its in the cache.
for channel_id in data["channels"]:
if channel := state.channels.get(channel_id):
self._channels[channel_id] = channel
self._emojis: dict[str, Emoji] = {}
def _update(self, *, owner: Optional[str] = None, name: Optional[str] = None, description: Optional[str] = None, icon: Optional[FilePayload] = None, banner: Optional[FilePayload] = None, default_permissions: Optional[int] = None, nsfw: Optional[bool] = None, system_messages: Optional[SystemMessagesConfig] = None, categories: Optional[list[CategoryPayload]] = None, channels: Optional[list[str]] = None): def _update(self, *, owner: Optional[str] = None, name: Optional[str] = None, description: Optional[str] = None, icon: Optional[FilePayload] = None, banner: Optional[FilePayload] = None, default_permissions: Optional[int] = None, nsfw: Optional[bool] = None, system_messages: Optional[SystemMessagesConfig] = None, categories: Optional[list[CategoryPayload]] = None, channels: Optional[list[str]] = None):
if owner: if owner is not None:
self.owner_id = owner self.owner_id = owner
if name: if name is not None:
self.name = name self.name = name
if description is not None: if description is not None:
self.description = description or None self.description = description or None
if icon: if icon is not None:
self.icon = Asset(icon, self.state) self.icon = Asset(icon, self.state)
if banner: if banner is not None:
self.banner = Asset(banner, self.state) self.banner = Asset(banner, self.state)
if default_permissions is not None: if default_permissions is not None:
self.default_permissions = Permissions(default_permissions) self.default_permissions = Permissions(default_permissions)
@@ -140,6 +185,12 @@ class Server:
if channels is not None: if channels is not None:
self._channels = {channel_id: self.state.get_channel(channel_id) for channel_id in channels} self._channels = {channel_id: self.state.get_channel(channel_id) for channel_id in channels}
def _add_member(self, payload: MemberPayload) -> Member:
member = Member(payload, self, self.state)
self._members[member.id] = member
return member
@property @property
def roles(self) -> list[Role]: def roles(self) -> list[Role]:
"""list[:class:`Role`] Gets all roles in the server in decending order""" """list[:class:`Role`] Gets all roles in the server in decending order"""
@@ -152,7 +203,7 @@ class Server:
@property @property
def channels(self) -> list[Channel]: def channels(self) -> list[Channel]:
"""list[:class:`Member`] Gets all channels in the server""" """list[:class:`Channel`] Gets all channels in the server"""
return list(self._channels.values()) return list(self._channels.values())
@property @property
@@ -160,6 +211,11 @@ class Server:
"""list[:class:`Category`] Gets all categories in the server""" """list[:class:`Category`] Gets all categories in the server"""
return list(self._categories.values()) return list(self._categories.values())
@property
def emojis(self) -> list[Emoji]:
"""list[:class:`Emoji`] Gets all emojis in the server"""
return list(self._emojis.values())
def get_role(self, role_id: str) -> Role: def get_role(self, role_id: str) -> Role:
"""Gets a role from the cache """Gets a role from the cache
@@ -190,8 +246,8 @@ class Server:
""" """
try: try:
return self._members[member_id] return self._members[member_id]
except KeyError as e: except KeyError:
raise LookupError from e raise LookupError from None
def get_channel(self, channel_id: str) -> Channel: def get_channel(self, channel_id: str) -> Channel:
"""Gets a channel from the cache """Gets a channel from the cache
@@ -208,8 +264,8 @@ class Server:
""" """
try: try:
return self._channels[channel_id] return self._channels[channel_id]
except KeyError as e: except KeyError:
raise LookupError from e raise LookupError from None
def get_category(self, category_id: str) -> Category: def get_category(self, category_id: str) -> Category:
"""Gets a category from the cache """Gets a category from the cache
@@ -226,6 +282,24 @@ class Server:
""" """
try: try:
return self._categories[category_id] return self._categories[category_id]
except KeyError:
raise LookupError from None
def get_emoji(self, emoji_id: str) -> Emoji:
"""Gets a emoji from the cache
Parameters
-----------
id: :class:`str`
The id of the emoji
Returns
--------
:class:`Emoji`
The emoji
"""
try:
return self._emojis[emoji_id]
except KeyError as e: except KeyError as e:
raise LookupError from e raise LookupError from e
@@ -236,21 +310,20 @@ class Server:
async def set_default_permissions(self, permissions: Permissions) -> None: async def set_default_permissions(self, permissions: Permissions) -> None:
"""Sets the default server permissions. """Sets the default server permissions.
Parameters Parameters
----------- -----------
server_permissions: Optional[:class:`ServerPermissions`] permissions: :class:`Permissions`
The new default server permissions The new default server permissions
channel_permissions: Optional[:class:`ChannelPermissions`]
the new default channel permissions
""" """
await self.state.http.set_server_default_permissions(self.id, permissions.value) 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""" """Leaves or deletes the server"""
await self.state.http.delete_leave_server(self.id) 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`""" """Leaves or deletes a server, alias to :meth`Server.leave_server`"""
await self.leave_server() await self.leave_server()
@@ -293,10 +366,10 @@ class Server:
""" """
payload = await self.state.http.create_channel(self.id, "Voice", name, description) payload = await self.state.http.create_channel(self.id, "Voice", name, description)
channel = VoiceChannel(payload, self.state) channel = self.state.add_channel(payload)
self._channels[channel.id] = channel self._channels[channel.id] = channel
return channel return cast(VoiceChannel, channel)
async def fetch_invites(self) -> list[Invite]: async def fetch_invites(self) -> list[Invite]:
"""Fetches all invites in the server """Fetches all invites in the server
@@ -327,11 +400,11 @@ class Server:
return Member(payload, self, self.state) return Member(payload, self, self.state)
async def fetch_bans(self) -> list[ServerBan]: async def fetch_bans(self) -> list[ServerBan]:
"""Fetches all invites in the server """Fetches all bans in the server
Returns Returns
-------- --------
list[:class:`Invite`] list[:class:`ServerBan`]
""" """
payload = await self.state.http.fetch_bans(self.id) payload = await self.state.http.fetch_bans(self.id)
@@ -353,14 +426,31 @@ class Server:
""" """
payload = await self.state.http.create_role(self.id, name) payload = await self.state.http.create_role(self.id, name)
return Role(payload, name, self, self.state) return Role(payload["role"], payload["id"], self, self.state)
async def create_emoji(self, name: str, file: File, *, nsfw: bool = False) -> Emoji:
"""Creates an emoji
Parameters
-----------
name: :class:`str`
The name for the emoji
file: :class:`File`
The image for the emoji
nsfw: :class:`bool`
Whether or not the emoji is nsfw
"""
payload = await self.state.http.create_emoji(name, file, nsfw, {"type": "Server", "id": self.id})
return self.state.add_emoji(payload)
class ServerBan: class ServerBan:
"""Represents a server ban """Represents a server ban
Attributes Attributes
----------- -----------
reason: Optional[:class:str`] reason: Optional[:class:`str`]
The reason the user was banned The reason the user was banned
server: :class:`Server` server: :class:`Server`
The server the user was banned in The server the user was banned in
@@ -371,11 +461,11 @@ class ServerBan:
__slots__ = ("reason", "server", "user_id", "state") __slots__ = ("reason", "server", "user_id", "state")
def __init__(self, ban: Ban, state: State): def __init__(self, ban: Ban, state: State):
self.reason = ban.get("reason") self.reason: str | None = ban.get("reason")
self.server = state.get_server(ban["_id"]["server"]) self.server: Server = state.get_server(ban["_id"]["server"])
self.user_id = ban["_id"]["user"] self.user_id: str = ban["_id"]["user"]
self.state = state self.state: State = state
async def unban(self): async def unban(self) -> None:
"""Unbans the user""" """Unbans the user"""
await self.state.http.unban_member(self.server.id, self.user_id) await self.state.http.unban_member(self.server.id, self.user_id)
+34 -17
View File
@@ -1,9 +1,10 @@
from __future__ import annotations from __future__ import annotations
from collections import deque from collections import deque
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING
from .channel import Channel, channel_factory from .channel import Channel, channel_factory
from .emoji import Emoji
from .member import Member from .member import Member
from .message import Message from .message import Message
from .server import Server from .server import Server
@@ -13,32 +14,35 @@ if TYPE_CHECKING:
from .http import HttpClient from .http import HttpClient
from .types import ApiInfo from .types import ApiInfo
from .types import Channel as ChannelPayload from .types import Channel as ChannelPayload
from .types import Emoji as EmojiPayload
from .types import Member as MemberPayload from .types import Member as MemberPayload
from .types import Message as MessagePayload from .types import Message as MessagePayload
from .types import Server as ServerPayload from .types import Server as ServerPayload
from .types import User as UserPayload from .types import User as UserPayload
__all__ = ("State",) __all__ = ("State",)
class State: class State:
__slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages") __slots__ = ("http", "api_info", "max_messages", "users", "channels", "servers", "messages", "global_emojis", "user_id", "me")
def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int): def __init__(self, http: HttpClient, api_info: ApiInfo, max_messages: int):
self.http = http self.http: HttpClient = http
self.api_info = api_info self.api_info: ApiInfo = api_info
self.max_messages = max_messages self.max_messages: int = max_messages
self.me: User
self.users: dict[str, User] = {} self.users: dict[str, User] = {}
self.channels: dict[str, Channel] = {} self.channels: dict[str, Channel] = {}
self.servers: dict[str, Server] = {} self.servers: dict[str, Server] = {}
self.messages: deque[Message] = deque() self.messages: deque[Message] = deque()
self.global_emojis: list[Emoji] = []
def get_user(self, id: str) -> User: def get_user(self, id: str) -> User:
try: try:
return self.users[id] return self.users[id]
except KeyError as e: except KeyError:
raise LookupError from e raise LookupError from None
def get_member(self, server_id: str, member_id: str) -> Member: def get_member(self, server_id: str, member_id: str) -> Member:
server = self.servers[server_id] server = self.servers[server_id]
@@ -47,26 +51,28 @@ class State:
def get_channel(self, id: str) -> Channel: def get_channel(self, id: str) -> Channel:
try: try:
return self.channels[id] return self.channels[id]
except KeyError as e: except KeyError:
raise LookupError from e raise LookupError from None
def get_server(self, id: str) -> Server: def get_server(self, id: str) -> Server:
try: try:
return self.servers[id] return self.servers[id]
except KeyError as e: except KeyError:
raise LookupError from e raise LookupError from None
def add_user(self, payload: UserPayload) -> User: def add_user(self, payload: UserPayload) -> User:
user = User(payload, self) user = User(payload, self)
if payload.get("relationship") == "User":
self.me = user
self.users[user.id] = user self.users[user.id] = user
return user return user
def add_member(self, server_id: str, payload: MemberPayload) -> Member: def add_member(self, server_id: str, payload: MemberPayload) -> Member:
server = self.get_server(server_id) server = self.get_server(server_id)
member = Member(payload, server, self)
server._members[member.id] = member
return member return server._add_member(payload)
def add_channel(self, payload: ChannelPayload) -> Channel: def add_channel(self, payload: ChannelPayload) -> Channel:
channel = channel_factory(payload, self) channel = channel_factory(payload, self)
@@ -86,6 +92,17 @@ class State:
self.messages.appendleft(message) self.messages.appendleft(message)
return message return message
def add_emoji(self, payload: EmojiPayload) -> Emoji:
emoji = Emoji(payload, self)
if server_id := emoji.server_id:
server = self.get_server(server_id)
server._emojis[emoji.id] = emoji
else:
self.global_emojis.append(emoji)
return emoji
def get_message(self, message_id: str) -> Message: def get_message(self, message_id: str) -> Message:
for msg in self.messages: for msg in self.messages:
if msg.id == message_id: if msg.id == message_id:
@@ -93,7 +110,7 @@ class State:
raise LookupError 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) data = await self.http.fetch_members(server_id)
for user in data["users"]: for user in data["users"]:
@@ -102,6 +119,6 @@ class State:
for member in data["members"]: for member in data["members"]:
self.add_member(server_id, member) 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: for server_id in self.servers:
await self.fetch_server_members(server_id) await self.fetch_server_members(server_id)
+2 -1
View File
@@ -1,13 +1,14 @@
from .category import * from .category import *
from .channel import * from .channel import *
from .embed import * from .embed import *
from .emoji import *
from .file import * from .file import *
from .gateway import * from .gateway import *
from .http import * from .http import *
from .invite import * from .invite import *
from .member import * from .member import *
from .message import * from .message import *
from .permissions import Overwrite from .permissions import *
from .role import * from .role import *
from .server import * from .server import *
from .user import * from .user import *
+1 -5
View File
@@ -1,8 +1,4 @@
from typing import TYPE_CHECKING, TypedDict from typing import TypedDict
if TYPE_CHECKING:
from .channel import Channel
__all__ = ("Category",) __all__ = ("Category",)
+3 -4
View File
@@ -1,12 +1,11 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Literal, Text, TypedDict, Union from typing import TYPE_CHECKING, Literal, TypedDict, Union
from typing_extensions import NotRequired from typing_extensions import NotRequired
if TYPE_CHECKING: if TYPE_CHECKING:
from .file import File from .file import File
from .message import Message
from .permissions import Overwrite from .permissions import Overwrite
__all__ = ( __all__ = (
@@ -15,7 +14,7 @@ __all__ = (
"GroupDMChannel", "GroupDMChannel",
"TextChannel", "TextChannel",
"VoiceChannel", "VoiceChannel",
"GuildChannel", "ServerChannel",
"Channel", "Channel",
) )
@@ -65,5 +64,5 @@ class VoiceChannel(BaseChannel):
role_permissions: NotRequired[dict[str, Overwrite]] role_permissions: NotRequired[dict[str, Overwrite]]
nsfw: NotRequired[bool] nsfw: NotRequired[bool]
GuildChannel = Union[TextChannel, VoiceChannel] ServerChannel = Union[TextChannel, VoiceChannel]
Channel = Union[SavedMessages, DMChannel, GroupDMChannel, TextChannel, VoiceChannel] Channel = Union[SavedMessages, DMChannel, GroupDMChannel, TextChannel, VoiceChannel]
+21
View File
@@ -0,0 +1,21 @@
from typing import Literal, TypedDict, Union
from typing_extensions import NotRequired
class EmojiParentServer(TypedDict):
type: Literal["Server"]
id: str
class EmojiParentDetached(TypedDict):
type: Literal["Detached"]
EmojiParent = Union[EmojiParentServer, EmojiParentDetached]
class Emoji(TypedDict):
_id: str
parent: EmojiParent
creator_id: str
name: str
animated: NotRequired[bool]
nsfw: NotRequired[bool]
+1 -1
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict, Union from typing import Literal, TypedDict, Union
__all__ = ("File",) __all__ = ("File",)
+47 -20
View File
@@ -2,26 +2,26 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Literal, TypedDict, Union from typing import TYPE_CHECKING, Literal, TypedDict, Union
from revolt.types.permissions import Overwrite from typing_extensions import NotRequired
from .channel import (Channel, DMChannel, GroupDMChannel, SavedMessages, from .channel import Channel, DMChannel, GroupDMChannel, SavedMessages, TextChannel, VoiceChannel
TextChannel, VoiceChannel)
from .file import File
from .message import Message from .message import Message
from .user import Status from .permissions import Overwrite
if TYPE_CHECKING: if TYPE_CHECKING:
from .category import Category from .category import Category
from .embed import Embed
from .emoji import Emoji
from .file import File
from .member import Member, MemberID from .member import Member, MemberID
from .server import Server, SystemMessagesConfig from .server import Server, SystemMessagesConfig
from .user import User from .user import Status, User, UserProfile, UserRelation
__all__ = ( __all__ = (
"BasePayload", "BasePayload",
"AuthenticatePayload", "AuthenticatePayload",
"ReadyEventPayload", "ReadyEventPayload",
"MessageEventPayload", "MessageEventPayload",
"MessageUpdateEditedData",
"MessageUpdateData", "MessageUpdateData",
"MessageUpdateEventPayload", "MessageUpdateEventPayload",
"MessageDeleteEventPayload", "MessageDeleteEventPayload",
@@ -39,7 +39,11 @@ __all__ = (
"ServerRoleDeleteEventPayload", "ServerRoleDeleteEventPayload",
"UserUpdateEventPayload", "UserUpdateEventPayload",
"UserRelationshipEventPayload", "UserRelationshipEventPayload",
"ServerCreateEventPayload" "ServerCreateEventPayload",
"MessageReactEventPayload",
"MessageUnreactEventPayload",
"MessageRemoveReactionEventPayload",
"BulkMessageDeleteEventPayload"
) )
class BasePayload(TypedDict): class BasePayload(TypedDict):
@@ -53,15 +57,15 @@ class ReadyEventPayload(BasePayload):
servers: list[Server] servers: list[Server]
channels: list[Channel] channels: list[Channel]
members: list[Member] members: list[Member]
emojis: list[Emoji]
class MessageEventPayload(BasePayload, Message): class MessageEventPayload(BasePayload, Message):
pass pass
MessageUpdateEditedData = TypedDict("MessageUpdateEditedData", {"$date": str})
class MessageUpdateData(TypedDict): class MessageUpdateData(TypedDict):
content: str content: str
edited: MessageUpdateEditedData embeds: list[Embed]
edited: Union[str, int]
class MessageUpdateEventPayload(BasePayload): class MessageUpdateEventPayload(BasePayload):
channel: str channel: str
@@ -140,6 +144,7 @@ class ServerMemberUpdateEventPayloadData(TypedDict, total=False):
nickname: str nickname: str
avatar: File avatar: File
roles: list[str] roles: list[str]
timeout: str | int
class ServerMemberUpdateEventPayload(BasePayload): class ServerMemberUpdateEventPayload(BasePayload):
id: MemberID id: MemberID
@@ -162,20 +167,25 @@ class ServerRoleUpdateEventPayload(BasePayload):
id: str id: str
role_id: str role_id: str
data: ServerRoleUpdateEventPayloadData data: ServerRoleUpdateEventPayloadData
clear: Literal["Color"] clear: Literal["Colour"]
class ServerRoleDeleteEventPayload(BasePayload): class ServerRoleDeleteEventPayload(BasePayload):
id: str id: str
role_id: str role_id: str
UserUpdateEventPayloadData = TypedDict("UserUpdateEventPayloadData", { class UserUpdateEventPayloadData(TypedDict):
"status": Status, status: NotRequired[Status]
"profile.background": File, avatar: NotRequired[File]
"profile.content": str, online: NotRequired[bool]
"avatar": File, profile: NotRequired[UserProfile]
"online": bool username: NotRequired[str]
display_name: NotRequired[str]
}, total=False) relations: NotRequired[list[UserRelation]]
badges: NotRequired[int]
online: NotRequired[bool]
flags: NotRequired[int]
discriminator: NotRequired[str]
privileged: NotRequired[bool]
class UserUpdateEventPayload(BasePayload): class UserUpdateEventPayload(BasePayload):
id: str id: str
@@ -186,3 +196,20 @@ class UserRelationshipEventPayload(BasePayload):
id: str id: str
user: str user: str
status: Status status: Status
class MessageReactEventPayload(BasePayload):
id: str
channel_id: str
user_id: str
emoji_id: str
MessageUnreactEventPayload = MessageReactEventPayload
class MessageRemoveReactionEventPayload(BasePayload):
id: str
channel_id: str
emoji_id: str
class BulkMessageDeleteEventPayload(BasePayload):
channel: str
ids: list[str]
+8 -1
View File
@@ -1,11 +1,13 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired
if TYPE_CHECKING: if TYPE_CHECKING:
from .member import Member from .member import Member
from .message import Message from .message import Message
from .user import User from .user import User
from .role import Role
__all__ = ( __all__ = (
@@ -14,6 +16,7 @@ __all__ = (
"Autumn", "Autumn",
"GetServerMembers", "GetServerMembers",
"MessageWithUserData", "MessageWithUserData",
"CreateRole",
) )
@@ -48,5 +51,9 @@ class GetServerMembers(TypedDict):
class MessageWithUserData(TypedDict): class MessageWithUserData(TypedDict):
messages: list[Message] messages: list[Message]
members: list[Member] members: NotRequired[list[Member]]
users: list[User] users: list[User]
class CreateRole(TypedDict):
id: str
role: Role
+3 -1
View File
@@ -8,7 +8,7 @@ if TYPE_CHECKING:
from .file import File from .file import File
__all__ = ("Member",) __all__ = ("Member", "MemberID")
class MemberID(TypedDict): class MemberID(TypedDict):
server: str server: str
@@ -19,3 +19,5 @@ class Member(TypedDict):
nickname: NotRequired[str] nickname: NotRequired[str]
avatar: NotRequired[File] avatar: NotRequired[File]
roles: NotRequired[list[str]] roles: NotRequired[list[str]]
joined_at: int | str
timeout: NotRequired[str | int]
+23 -5
View File
@@ -10,9 +10,20 @@ if TYPE_CHECKING:
__all__ = ( __all__ = (
"UserAddContent",
"UserRemoveContent",
"UserJoinedContent",
"UserLeftContent",
"UserKickedContent",
"UserBannedContent",
"ChannelRenameContent",
"ChannelDescriptionChangeContent",
"ChannelIconChangeContent",
"Masquerade",
"Interactions",
"Message", "Message",
"MessageReplyPayload", "MessageReplyPayload",
"Masquerade" "SystemMessageContent",
) )
class UserAddContent(TypedDict): class UserAddContent(TypedDict):
@@ -46,24 +57,31 @@ class ChannelDescriptionChangeContent(TypedDict):
class ChannelIconChangeContent(TypedDict): class ChannelIconChangeContent(TypedDict):
by: str by: str
MessageEdited = TypedDict("MessageEdited", {"$date": str})
class Masquerade(TypedDict, total=False): class Masquerade(TypedDict, total=False):
name: str name: str
avatar: str avatar: str
colour: str colour: str
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): class Message(TypedDict):
_id: str _id: str
channel: str channel: str
author: str author: str
content: Union[str, UserAddContent, UserRemoveContent, UserJoinedContent, UserLeftContent, UserKickedContent, UserBannedContent, ChannelRenameContent, ChannelDescriptionChangeContent, ChannelIconChangeContent] content: str
system: NotRequired[SystemMessageContent]
attachments: NotRequired[list[File]] attachments: NotRequired[list[File]]
embeds: NotRequired[list[Embed]] embeds: NotRequired[list[Embed]]
mentions: NotRequired[list[str]] mentions: NotRequired[list[str]]
replies: NotRequired[list[str]] replies: NotRequired[list[str]]
edited: NotRequired[MessageEdited] edited: NotRequired[str | int]
masquerade: NotRequired[Masquerade] masquerade: NotRequired[Masquerade]
interactions: NotRequired[Interactions]
reactions: dict[str, list[str]]
class MessageReplyPayload(TypedDict): class MessageReplyPayload(TypedDict):
id: str id: str
+1
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from typing import TypedDict from typing import TypedDict
class Overwrite(TypedDict): class Overwrite(TypedDict):
a: int a: int
d: int d: int
+1
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, TypedDict from typing import TYPE_CHECKING, TypedDict
from typing_extensions import NotRequired from typing_extensions import NotRequired
if TYPE_CHECKING: if TYPE_CHECKING:
-1
View File
@@ -6,7 +6,6 @@ from typing_extensions import NotRequired
if TYPE_CHECKING: if TYPE_CHECKING:
from .category import Category from .category import Category
from .channel import Channel
from .file import File from .file import File
from .role import Role from .role import Role
+3
View File
@@ -32,6 +32,8 @@ class UserRelation(TypedDict):
class User(TypedDict): class User(TypedDict):
_id: str _id: str
username: str username: str
discriminator: str
display_name: NotRequired[str]
avatar: NotRequired[File] avatar: NotRequired[File]
relations: NotRequired[list[UserRelation]] relations: NotRequired[list[UserRelation]]
badges: NotRequired[int] badges: NotRequired[int]
@@ -40,6 +42,7 @@ class User(TypedDict):
online: NotRequired[bool] online: NotRequired[bool]
flags: NotRequired[int] flags: NotRequired[int]
bot: NotRequired[UserBot] bot: NotRequired[UserBot]
privileged: NotRequired[bool]
class UserProfile(TypedDict, total=False): class UserProfile(TypedDict, total=False):
content: str content: str
+214 -38
View File
@@ -1,19 +1,26 @@
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union from typing import TYPE_CHECKING, NamedTuple, Optional, Union
from weakref import WeakValueDictionary
from revolt.types.user import UserRelation
from .asset import Asset, PartialAsset from .asset import Asset, PartialAsset
from .channel import DMChannel from .channel import DMChannel, GroupDMChannel, SavedMessageChannel
from .enums import PresenceType, RelationshipType from .enums import PresenceType, RelationshipType
from .flags import UserBadges from .flags import UserBadges
from .messageable import Messageable from .messageable import Messageable
from .permissions import UserPermissions
from .utils import Ulid
if TYPE_CHECKING: if TYPE_CHECKING:
from .member import Member
from .state import State from .state import State
from .types import File from .types import File
from .types import Status as StatusPayload from .types import Status as StatusPayload
from .types import User as UserPayload from .types import User as UserPayload
from .member import Member from .types import UserProfile as UserProfileData
from .server import Server
__all__ = ("User", "Status", "Relation", "UserProfile") __all__ = ("User", "Status", "Relation", "UserProfile")
@@ -32,17 +39,21 @@ class UserProfile(NamedTuple):
content: Optional[str] content: Optional[str]
background: Optional[Asset] background: Optional[Asset]
class User(Messageable): class User(Messageable, Ulid):
"""Represents a user """Represents a user
Attributes Attributes
----------- -----------
id: :class:`str` 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` bot: :class:`bool`
Whether or not the user is a bot Whether or not the user is a bot
owner: Optional[:class:`User`] owner_id: Optional[:class:`str`]
The bot's owner if the user is a bot The bot's owner id if the user is a bot
badges: :class:`UserBadges` badges: :class:`UserBadges`
The users badges The users badges
online: :class:`bool` online: :class:`bool`
@@ -57,18 +68,26 @@ class User(Messageable):
The users status The users status
dm_channel: Optional[:class:`DMChannel`] dm_channel: Optional[:class:`DMChannel`]
The dm channel between the client and the user, this will only be set if the client has dm'ed the user or :meth:`User.open_dm` was run The dm channel between the client and the user, this will only be set if the client has dm'ed the user or :meth:`User.open_dm` was run
privileged: :class:`bool`
Whether the user is privileged
""" """
__flattern_attributes__ = ("id", "bot", "owner_id", "badges", "online", "flags", "relations", "relationship", "status", "masquerade_avatar", "masquerade_name", "original_name", "original_avatar", "profile", "dm_channel") __flattern_attributes__: 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__ = (*__flattern_attributes__, "state", "_members") __slots__: tuple[str, ...] = (*__flattern_attributes__, "state", "_members")
def __init__(self, data: UserPayload, state: State): def __init__(self, data: UserPayload, state: State):
self.state = state self.state = state
self._members: list[Member] = [] # we store all member versions of this user to avoid having to check every guild when needing to update. self._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 = data["_id"] self.id: str = data["_id"]
self.original_name = data["username"] self.discriminator: str = data["discriminator"]
self.dm_channel = None self.display_name: str | None = data.get("display_name")
self.original_name: str = data["username"]
self.dm_channel: DMChannel | SavedMessageChannel | None = None
bot = data.get("bot") bot = data.get("bot")
self.bot: bool
self.owner_id: str | None
if bot: if bot:
self.bot = True self.bot = True
self.owner_id = bot["owner"] self.owner_id = bot["owner"]
@@ -76,12 +95,13 @@ class User(Messageable):
self.bot = False self.bot = False
self.owner_id = None self.owner_id = None
self.badges = UserBadges._from_value(data.get("badges", 0)) self.badges: UserBadges = UserBadges._from_value(data.get("badges", 0))
self.online = data.get("online", False) self.online: bool = data.get("online", False)
self.flags = data.get("flags", 0) self.flags: int = data.get("flags", 0)
self.privileged: bool = data.get("privileged", False)
avatar = data.get("avatar") 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] = [] relations: list[Relation] = []
@@ -89,12 +109,15 @@ class User(Messageable):
user = state.get_user(relation["_id"]) user = state.get_user(relation["_id"])
if user: if user:
relations.append(Relation(RelationshipType(relation["status"]), user)) relations.append(Relation(RelationshipType(relation["status"]), user))
self.relations = relations
self.relations: list[Relation] = relations
relationship = data.get("relationship") 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") status = data.get("status")
self.status: Status | None
if status: if status:
presence = status.get("presence") presence = status.get("presence")
self.status = Status(status.get("text"), PresenceType(presence) if presence else None) if status else None self.status = Status(status.get("text"), PresenceType(presence) if presence else None) if status else None
@@ -106,26 +129,76 @@ class User(Messageable):
self.masquerade_avatar: Optional[PartialAsset] = None self.masquerade_avatar: Optional[PartialAsset] = None
self.masquerade_name: Optional[str] = None self.masquerade_name: Optional[str] = None
def get_permissions(self) -> UserPermissions:
"""Gets the permissions for the user
Returns
--------
:class:`UserPermissions`
The users permissions
"""
permissions = UserPermissions()
if self.relationship in [RelationshipType.friend, RelationshipType.user]:
return UserPermissions.all()
elif self.relationship in [RelationshipType.blocked, RelationshipType.blocked_other]:
return UserPermissions(access=True)
elif self.relationship in [RelationshipType.incoming_friend_request, RelationshipType.outgoing_friend_request]:
permissions.access = True
for channel in self.state.channels.values():
if (isinstance(channel, (GroupDMChannel, DMChannel)) and self.id in channel.recipient_ids) or any(self.id in (m.id for m in server.members) for server in self.state.servers.values()):
if self.state.me.bot or self.bot:
permissions.send_message = True
permissions.access = True
permissions.view_profile = True
return permissions
def has_permissions(self, **permissions: bool) -> bool:
"""Computes if the user has the specified permissions
Parameters
-----------
permissions: :class:`bool`
The permissions to check, this also accepted `False` if you need to check if the user does not have the permission
Returns
--------
:class:`bool`
Whether or not they have the permissions
"""
perms = self.get_permissions()
return all([getattr(perms, key, False) == value for key, value in permissions.items()])
async def _get_channel_id(self): async def _get_channel_id(self):
if not self.dm_channel: if not self.dm_channel:
payload = await self.state.http.open_dm(self.id) payload = await self.state.http.open_dm(self.id)
self.dm_channel = DMChannel(payload, self.state)
return self.id if payload["channel_type"] == "SavedMessages":
self.dm_channel = SavedMessageChannel(payload, self.state)
else:
self.dm_channel = DMChannel(payload, self.state)
return self.dm_channel.id
@property @property
def owner(self) -> Optional[User]: def owner(self) -> User:
owner_id = self.owner_id """:class:`User` the owner of the bot account"""
if not owner_id: if not self.owner_id:
return raise LookupError
return self.state.get_user(owner_id) return self.state.get_user(self.owner_id)
@property @property
def name(self) -> str: def name(self) -> str:
""":class:`str` The name the user is displaying, this includes there orginal name and masqueraded name""" """:class:`str` The name the user is displaying, this includes (in order) their masqueraded name, display name and orginal name"""
return self.masquerade_name or self.original_name return self.display_name or self.masquerade_name or self.original_name
@property @property
def avatar(self) -> Union[Asset, PartialAsset, None]: def avatar(self) -> Union[Asset, PartialAsset, None]:
@@ -137,27 +210,85 @@ class User(Messageable):
""":class:`str`: Returns a string that allows you to mention the given user.""" """:class:`str`: Returns a string that allows you to mention the given user."""
return f"<@{self.id}>" return f"<@{self.id}>"
def _update(self, *, status: Optional[StatusPayload] = None, profile_content: Optional[str] = None, profile_background: Optional[File] = None, avatar: Optional[File] = None, online: Optional[bool] = None): def _update(
if status: self,
*,
status: Optional[StatusPayload] = None,
profile: Optional[UserProfileData] = None,
avatar: Optional[File] = None,
online: Optional[bool] = None,
display_name: Optional[str] = None,
relations: Optional[list[UserRelation]] = None,
badges: Optional[int] = None,
flags: Optional[int] = None,
discriminator: Optional[str] = None,
privileged: Optional[bool] = None,
username: Optional[str] = None
) -> None:
if status is not None:
presence = status.get("presence") presence = status.get("presence")
self.status = Status(status.get("text"), PresenceType(presence) if presence else None) self.status = Status(status.get("text"), PresenceType(presence) if presence else None)
if profile_background: if profile is not None:
self.profile = UserProfile(self.profile.content if self.profile else None, Asset(profile_background, self.state)) if background_file := profile.get("background"):
background = Asset(background_file, self.state)
else:
background = None
if profile_content: self.profile = UserProfile(profile.get("content"), background)
self.profile = UserProfile(profile_content, self.profile.background if self.profile else None)
if avatar: if avatar is not None:
self.original_avatar = Asset(avatar, self.state) self.original_avatar = Asset(avatar, self.state)
if online: if online is not None:
self.online = online self.online = online
if display_name is not None:
self.display_name = display_name
if relations is not None:
new_relations: list[Relation] = []
for relation in relations:
user = self.state.get_user(relation["_id"])
if user:
new_relations.append(Relation(RelationshipType(relation["status"]), user))
self.relations = new_relations
if badges is not None:
self.badges = UserBadges(badges)
if flags is not None:
self.flags = flags
if discriminator is not None:
self.discriminator = discriminator
if privileged is not None:
self.privileged = privileged
if username is not None:
self.original_name = username
# update user infomation for all members # update user infomation for all members
for member in self._members: if self.__class__ is User:
User._update(member, status=status, profile_content=profile_content, profile_background=profile_background, avatar=avatar, online=online) for member in self._members.values():
User._update(
member,
status=status,
profile=profile,
avatar=avatar,
online=online,
display_name=display_name,
relations=relations,
badges=badges,
flags=flags,
discriminator=discriminator,
privileged=privileged,
username=username
)
async def default_avatar(self) -> bytes: async def default_avatar(self) -> bytes:
"""Returns the default avatar for this user """Returns the default avatar for this user
@@ -189,3 +320,48 @@ class User(Messageable):
self.profile = UserProfile(payload.get("content"), background) self.profile = UserProfile(payload.get("content"), background)
return self.profile 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
async def open_dm(self) -> DMChannel | SavedMessageChannel:
"""Opens a dm with the user, if this user is the current user this will return :class:`SavedMessageChannel`
.. note:: using this function is discouraged as :meth:`User.send` does this implicitally.
Returns
--------
Union[:class:`DMChannel`, :class:`SavedMessageChannel`]
"""
await self._get_channel_id()
assert self.dm_channel
return self.dm_channel
+55 -7
View File
@@ -1,18 +1,25 @@
from __future__ import annotations
import datetime
import inspect import inspect
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from operator import attrgetter from operator import attrgetter
from typing import Any, Callable, Coroutine, Iterable, TypeVar, Union from typing import Any, Callable, Coroutine, Iterable, Literal, TypeVar, Union
import ulid
from aiohttp import ClientSession from aiohttp import ClientSession
from typing_extensions import ParamSpec from typing_extensions import ParamSpec
__all__ = ("Missing", "copy_doc", "maybe_coroutine", "get", "client_session") __all__ = ("_Missing", "Missing", "copy_doc", "maybe_coroutine", "get", "client_session", "parse_timestamp")
class _Missing: class _Missing:
def __repr__(self): def __repr__(self) -> str:
return "<Missing>" return "<Missing>"
Missing = _Missing() def __bool__(self) -> Literal[False]:
return False
Missing: _Missing = _Missing()
T = TypeVar("T") T = TypeVar("T")
@@ -26,9 +33,9 @@ def copy_doc(from_t: T) -> Callable[[T], T]:
R_T = TypeVar("R_T") R_T = TypeVar("R_T")
P = ParamSpec("P") P = ParamSpec("P")
# it is impossible to type this function correctly for a couple reasons: # it is impossible to type this function correctly as typeguard does not narrow for the negative case,
# 1. isawaitable does not narrow while keeping typevars - there is an open PR for this (typeshed#5658) but it cannot be merged because mypy does not support the feature fully # so `value` would stay being a union even after the if statement (PEP 647 - "The type is not narrowed in the negative case")
# 2. typeguard does not narrow for the negative case which is dumb in my opinion, so `value` would stay being a union even after the if statement (PEP 647 - "The type is not narrowed in the negative case") # see typing#926, typing#930, typing#996
async def maybe_coroutine(func: Callable[P, Union[R_T, Coroutine[Any, Any, R_T]]], *args: P.args, **kwargs: P.kwargs) -> R_T: async def maybe_coroutine(func: Callable[P, Union[R_T, Coroutine[Any, Any, R_T]]], *args: P.args, **kwargs: P.kwargs) -> R_T:
value = func(*args, **kwargs) value = func(*args, **kwargs)
@@ -39,6 +46,41 @@ async def maybe_coroutine(func: Callable[P, Union[R_T, Coroutine[Any, Any, R_T]]
return value # type: ignore return value # type: ignore
class Ulid:
"""Base class for any revolt object with an id
Attributes
-----------
id: :class:`str`
The id of the object
"""
id: str
@property
def created_at(self) -> datetime.datetime:
"""Returns a datetime for when the object was created according to the id
Returns
--------
:class:`datetime.datetime`
The datetime of the creation date and time
"""
return ulid.from_str(self.id).timestamp().datetime
class Object(Ulid):
"""Class to mock objects with an id
.. note::
This does not validate or guarantee the id is correct, you must handle this yourself
Parameters
-----------
id: :class:`str`
The ULID id to mock
"""
def __init__(self, id: str):
self.id = id
def get(iterable: Iterable[T], **attrs: Any) -> T: def get(iterable: Iterable[T], **attrs: Any) -> T:
"""A convenience function to help get a value from an iterable with a specific attribute """A convenience function to help get a value from an iterable with a specific attribute
@@ -103,3 +145,9 @@ async def client_session():
yield session yield session
finally: finally:
await session.close() await session.close()
def parse_timestamp(timestamp: int | str) -> datetime.datetime:
if isinstance(timestamp, int):
return datetime.datetime.fromtimestamp(timestamp / 1000, tz=datetime.timezone.utc)
else:
return datetime.datetime.strptime(timestamp, "%Y-%m-%dT%H:%M:%S.%f%z")
+229 -119
View File
@@ -2,31 +2,43 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
import time
from copy import copy from copy import copy
from traceback import print_exception from typing import TYPE_CHECKING, Callable, NamedTuple, cast
from typing import TYPE_CHECKING, Callable, cast
from .errors import RevoltError
from . import utils
from .channel import GroupDMChannel, TextChannel, VoiceChannel from .channel import GroupDMChannel, TextChannel, VoiceChannel
from .enums import RelationshipType from .enums import RelationshipType
from .types import (ChannelCreateEventPayload, ChannelDeleteEventPayload, from .role import Role
ChannelDeleteTypingEventPayload, from .types import (BulkMessageDeleteEventPayload, ChannelCreateEventPayload,
ChannelDeleteEventPayload, ChannelDeleteTypingEventPayload,
ChannelStartTypingEventPayload, ChannelUpdateEventPayload) ChannelStartTypingEventPayload, ChannelUpdateEventPayload)
from .types import Member as MemberPayload
from .types import MemberID as MemberIDPayload
from .types import Message as MessagePayload from .types import Message as MessagePayload
from .types import (MessageDeleteEventPayload, MessageUpdateEventPayload, from .types import (MessageDeleteEventPayload, MessageReactEventPayload,
ServerDeleteEventPayload, ServerMemberJoinEventPayload, MessageRemoveReactionEventPayload,
MessageUnreactEventPayload, MessageUpdateEventPayload)
from .types import Role as RolePayload
from .types import (ServerCreateEventPayload, ServerDeleteEventPayload,
ServerMemberJoinEventPayload,
ServerMemberLeaveEventPayload, ServerMemberLeaveEventPayload,
ServerCreateEventPayload,
ServerMemberUpdateEventPayload, ServerMemberUpdateEventPayload,
ServerRoleDeleteEventPayload, ServerRoleUpdateEventPayload, ServerRoleDeleteEventPayload, ServerRoleUpdateEventPayload,
ServerUpdateEventPayload, UserRelationshipEventPayload, ServerUpdateEventPayload, UserRelationshipEventPayload,
UserUpdateEventPayload) UserUpdateEventPayload)
from .user import Status, UserProfile from .user import Status, User, UserProfile
import aiohttp
try: try:
import ujson as json import ujson as json
except ImportError: except ImportError:
import json import json
use_msgpack: bool
try: try:
import msgpack import msgpack
use_msgpack = True use_msgpack = True
@@ -37,47 +49,50 @@ if TYPE_CHECKING:
import aiohttp import aiohttp
from .state import State from .state import State
from .types import AuthenticatePayload, BasePayload from .types import (AuthenticatePayload, BasePayload, MessageEventPayload,
from .types import Member as MemberPayload ReadyEventPayload)
from .types import 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: class WebsocketHandler:
__slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready", "server_events") __slots__ = ("session", "token", "ws_url", "dispatch", "state", "websocket", "loop", "user", "ready", "server_events")
def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State): def __init__(self, session: aiohttp.ClientSession, token: str, ws_url: str, dispatch: Callable[..., None], state: State):
self.session = session self.session: aiohttp.ClientSession = session
self.token = token self.token: str = token
self.ws_url = ws_url self.ws_url: str = ws_url
self.dispatch = dispatch self.dispatch: Callable[..., None] = dispatch
self.state = state self.state: State = state
self.websocket: aiohttp.ClientWebSocketResponse self.websocket: aiohttp.ClientWebSocketResponse
self.loop = asyncio.get_running_loop() self.loop: asyncio.AbstractEventLoop = asyncio.get_running_loop()
self.user = None self.user: User | None = None
self.ready = asyncio.Event() self.ready: asyncio.Event = asyncio.Event()
self.server_events: dict[str, 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): if event := self.server_events.get(server_id):
await event.wait() await event.wait()
async def send_payload(self, payload: BasePayload): async def send_payload(self, payload: BasePayload) -> None:
if use_msgpack: if use_msgpack:
await self.websocket.send_bytes(msgpack.packb(payload)) # type: ignore await self.websocket.send_bytes(msgpack.packb(payload)) # type: ignore
else: else:
await self.websocket.send_str(json.dumps(payload)) await self.websocket.send_str(json.dumps(payload))
async def heartbeat(self): async def heartbeat(self) -> None:
while not self.websocket.closed: while not self.websocket.closed:
logger.info("Sending hearbeat") logger.info("Sending hearbeat")
await self.websocket.ping() await self.websocket.ping()
await asyncio.sleep(15) await asyncio.sleep(15)
async def send_authenticate(self): async def send_authenticate(self) -> None:
payload: AuthenticatePayload = { payload: AuthenticatePayload = {
"type": "Authenticate", "type": "Authenticate",
"token": self.token "token": self.token
@@ -85,24 +100,27 @@ class WebsocketHandler:
await self.send_payload(payload) 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() event_type = payload["type"].lower()
logger.debug("Recieved event %s %s", event_type, payload) logger.debug("Recieved event %s %s", event_type, payload)
try: try:
if event_type != "ready": if event_type not in ["ready", "notfound"]:
await self.ready.wait() await self.ready.wait()
func = getattr(self, f"handle_{event_type}") func = getattr(self, f"handle_{event_type}")
except AttributeError: except AttributeError:
logger.debug("Unknown event '%s'", event_type) return logger.debug("Unknown event '%s'", event_type)
return
await func(payload) await func(payload)
async def handle_authenticated(self, _): async def handle_authenticated(self, _: BasePayload) -> None:
logger.info("Successfully authenticated") logger.info("Successfully authenticated")
async def handle_ready(self, payload: ReadyEventPayload): async def handle_notfound(self, _: BasePayload) -> None:
raise RevoltError("Invalid token")
async def handle_ready(self, payload: ReadyEventPayload) -> None:
for user_payload in payload["users"]: for user_payload in payload["users"]:
user = self.state.add_user(user_payload) user = self.state.add_user(user_payload)
@@ -119,67 +137,71 @@ class WebsocketHandler:
for member in payload["members"]: for member in payload["members"]:
self.state.add_member(member["_id"]["server"], member) self.state.add_member(member["_id"]["server"], member)
for emoji in payload["emojis"]:
emoji = self.state.add_emoji(emoji)
await self.state.fetch_all_server_members() await self.state.fetch_all_server_members()
self.ready.set() self.ready.set()
self.dispatch("ready") 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)
message = self.state.add_message(cast(MessagePayload, payload)) message = self.state.add_message(cast(MessagePayload, payload))
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message", message) 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) self.dispatch("raw_message_update", payload)
message = self.state.get_message(payload["id"]) try:
message = self.state.get_message(payload["id"])
except LookupError:
return
data = payload["data"] if server_id := message.channel.server_id:
kwargs = {} await self._wait_for_server_ready(server_id)
if content := data.get("content"): before = copy(message)
kwargs["content"] = content message._update(**payload["data"])
kwargs["edited_at"] = data["edited"]["$date"] self.dispatch("message_update", before, message)
if embeds := data.get("embeds"): async def handle_messagedelete(self, payload: MessageDeleteEventPayload) -> None:
kwargs["embeds"] = embeds
message._update(**kwargs)
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message_update", message)
async def handle_messagedelete(self, payload: MessageDeleteEventPayload):
self.dispatch("raw_message_delete", payload) self.dispatch("raw_message_delete", payload)
try: try:
message = self.state.get_message(payload["id"]) message = self.state.get_message(payload["id"])
except KeyError: except LookupError:
return return
if server_id := message.channel.server_id:
await self._wait_for_server_ready(server_id)
self.state.messages.remove(message) self.state.messages.remove(message)
if server := message.server:
await self._wait_for_server_ready(server.id)
self.dispatch("message_delete", message) 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) channel = self.state.add_channel(payload)
if server := channel.server: if server_id := channel.server_id:
await self._wait_for_server_ready(server.id) await self._wait_for_server_ready(server_id)
self.dispatch("channel_create", channel) self.dispatch("channel_create", channel)
async def handle_channelupdate(self, payload: ChannelUpdateEventPayload): async def handle_channelupdate(self, payload: ChannelUpdateEventPayload) -> None:
channel = self.state.get_channel(payload["id"]) # 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)):
return
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
old_channel = copy(channel) old_channel = copy(channel)
@@ -194,38 +216,40 @@ class WebsocketHandler:
if isinstance(channel, (TextChannel, VoiceChannel, GroupDMChannel)): if isinstance(channel, (TextChannel, VoiceChannel, GroupDMChannel)):
channel.description = None channel.description = None
if server := channel.server:
await self._wait_for_server_ready(server.id)
self.dispatch("channel_update", old_channel, channel) 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"]) channel = self.state.channels.pop(payload["id"])
if server := channel.server: if server_id := channel.server_id:
await self._wait_for_server_ready(server.id) await self._wait_for_server_ready(server_id)
self.dispatch("channel_delete", channel) 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"]) channel = self.state.get_channel(payload["id"])
user = self.state.get_user(payload["user"])
if server := channel.server: if server_id := channel.server_id:
await self._wait_for_server_ready(server.id) await self._wait_for_server_ready(server_id)
user = self.state.get_user(payload["user"])
self.dispatch("typing_start", channel, user) 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"]) channel = self.state.get_channel(payload["id"])
user = self.state.get_user(payload["user"])
if server := channel.server: if server_id := channel.server_id:
await self._wait_for_server_ready(server.id) await self._wait_for_server_ready(server_id)
user = self.state.get_user(payload["user"])
self.dispatch("typing_stop", channel, user) 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"]) server = self.state.get_server(payload["id"])
old_server = copy(server) old_server = copy(server)
@@ -242,11 +266,10 @@ class WebsocketHandler:
elif clear == "Description": elif clear == "Description":
server.description = None server.description = None
await self._wait_for_server_ready(server.id)
self.dispatch("server_update", old_server, server) 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"]) server = self.state.servers.pop(payload["id"])
for channel in server.channels: for channel in server.channels:
@@ -256,7 +279,7 @@ class WebsocketHandler:
self.dispatch("server_delete", server) 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"]: for channel in payload["channels"]:
self.state.add_channel(channel) self.state.add_channel(channel)
@@ -267,9 +290,9 @@ class WebsocketHandler:
await self.state.fetch_server_members(server.id) await self.state.fetch_server_members(server.id)
self.server_events.pop(server.id).set() self.server_events.pop(server.id).set()
self.dispatch("server_create", server) 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"]) await self._wait_for_server_ready(payload["id"]["server"])
member = self.state.get_member(payload["id"]["server"], payload["id"]["user"]) member = self.state.get_member(payload["id"]["server"], payload["id"]["user"])
@@ -285,11 +308,17 @@ class WebsocketHandler:
self.dispatch("member_update", old_member, member) self.dispatch("member_update", old_member, member)
async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload): async def handle_servermemberjoin(self, payload: ServerMemberJoinEventPayload) -> None:
member = self.state.add_member(payload["id"], {"_id": {"server": payload["id"], "user": payload["user"]}}) # avoid an api request if possible
if payload["user"] not in self.state.users:
user = await self.state.http.fetch_user(payload["user"])
self.state.add_user(user)
member = self.state.add_member(payload["id"], MemberPayload(_id=MemberIDPayload(server=payload["id"], user=payload["user"]), joined_at=int(time.time()))) # revolt doesnt give us the joined at time
self.dispatch("member_join", member) self.dispatch("member_join", member)
async def handle_memberleave(self, payload: ServerMemberLeaveEventPayload): async def handle_servermemberleave(self, payload: ServerMemberLeaveEventPayload) -> None:
await self._wait_for_server_ready(payload["id"]) await self._wait_for_server_ready(payload["id"])
server = self.state.get_server(payload["id"]) server = self.state.get_server(payload["id"])
@@ -298,26 +327,34 @@ class WebsocketHandler:
# remove the member from the user # remove the member from the user
user = self.state.get_user(payload["user"]) user = self.state.get_user(payload["user"])
user._members.remove(member) user._members.pop(server.id)
self.dispatch("member_leave", member) self.dispatch("member_leave", member)
async def handle_serveroleupdate(self, payload: ServerRoleUpdateEventPayload): async def handle_serverroleupdate(self, payload: ServerRoleUpdateEventPayload) -> None:
server = self.state.get_server(payload["id"]) server = self.state.get_server(payload["id"])
role = server.get_role(payload["role_id"])
old_role = copy(role)
if clear := payload.get("clear"):
if clear == "Colour":
role.colour = None
role._update(**payload["data"])
await self._wait_for_server_ready(server.id) await self._wait_for_server_ready(server.id)
self.dispatch("role_update", old_role, role) try:
role = server.get_role(payload["role_id"])
except LookupError:
# the role wasnt found meaning it was just created
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload): role = Role(cast(RolePayload, payload["data"]), payload["role_id"], server, self.state)
server._roles[role.id] = role
self.dispatch("role_create", role)
else:
old_role = copy(role)
if clear := payload.get("clear"):
if clear == "Colour":
role.colour = None
role._update(**payload["data"])
self.dispatch("role_update", old_role, role)
async def handle_serverroledelete(self, payload: ServerRoleDeleteEventPayload) -> None:
server = self.state.get_server(payload["id"]) server = self.state.get_server(payload["id"])
role = server._roles.pop(payload["role_id"]) role = server._roles.pop(payload["role_id"])
@@ -325,7 +362,7 @@ class WebsocketHandler:
self.dispatch("role_delete", role) 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"]) user = self.state.get_user(payload["id"])
old_user = copy(user) old_user = copy(user)
@@ -344,43 +381,116 @@ class WebsocketHandler:
elif clear == "Avatar": elif clear == "Avatar":
user.original_avatar = None user.original_avatar = None
# the keys have . in them so I need to replace with _ user._update(**payload["data"])
# type: ignore is for it to stop complaining about the keys not existing in the typeddict
data = payload["data"]
data["profile_content"] = data.pop("profile.content", None) # type: ignore
data["profile_background"] = data.pop("profile.background", None) # type: ignore
user._update(**data) # type: ignore
self.dispatch("user_update", old_user, user) 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"]) user = self.state.get_user(payload["user"])
old_relationship = user.relationship old_relationship = user.relationship
user.relationship = RelationshipType(payload["status"]) user.relationship = RelationshipType(payload["status"])
self.dispatch("user_relationship_update", user, old_relationship, user.relationship) self.dispatch("user_relationship_update", user, old_relationship, user.relationship)
async def start(self): 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)
self.dispatch("raw_reaction_add", payload)
try:
message = utils.get(self.state.messages, id=payload["id"])
except LookupError:
return
user = self.state.get_user(payload["user_id"])
message.reactions.setdefault(payload["emoji_id"], []).append(user)
emoji_id = payload["emoji_id"]
self.dispatch("reaction_add", message, user, emoji_id)
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)
self.dispatch("raw_reaction_remove", payload)
try:
message = utils.get(self.state.messages, id=payload["id"])
except LookupError:
return
user = self.state.get_user(payload["user_id"])
message.reactions[payload["emoji_id"]].remove(user)
self.dispatch("reaction_remove", message, user, payload["emoji_id"])
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)
self.dispatch("raw_reaction_clear", payload)
try:
message = utils.get(self.state.messages, id=payload["id"])
except LookupError:
return
users = message.reactions.pop(payload["emoji_id"])
self.dispatch("reaction_clear", message, users, payload["emoji_id"])
async def handle_bulkmessagedelete(self, payload: BulkMessageDeleteEventPayload) -> None:
channel = self.state.get_channel(payload["channel"])
self.dispatch("raw_bulk_message_delete", payload)
messages: list[Message] = []
for message_id in payload["ids"]:
if server_id := channel.server_id:
await self._wait_for_server_ready(server_id)
self.dispatch("raw_message_delete", MessageDeleteEventPayload(type="messagedelete", channel=payload["channel"], id=message_id))
try:
message = self.state.get_message(message_id)
except LookupError:
pass
else:
self.state.messages.remove(message)
self.dispatch("message_delete", message)
messages.append(message)
self.dispatch("bulk_message_delete", messages)
async def start(self, reconnect: bool) -> None:
if use_msgpack: if use_msgpack:
url = f"{self.ws_url}?format=msgpack" url = f"{self.ws_url}?format=msgpack"
else: else:
url = f"{self.ws_url}?format=json" url = f"{self.ws_url}?format=json"
self.websocket = await self.session.ws_connect(url) while True:
await self.send_authenticate() self.websocket = await self.session.ws_connect(url) # type: ignore
asyncio.create_task(self.heartbeat()) await self.send_authenticate()
hb = asyncio.create_task(self.heartbeat())
async for msg in self.websocket: async for msg in self.websocket:
if use_msgpack: msg = cast(WSMessage, msg) # aiohttp doesnt use NamedTuple so the type info is missing
payload = msgpack.unpackb(msg.data)
else:
payload = json.loads(msg.data)
task = self.loop.create_task(self.handle_event(payload)) if use_msgpack:
# task.add_done_callback(task_done) data = cast(bytes, msg.data)
def task_done(task: asyncio.Task[None]): payload = msgpack.unpackb(data) # type: ignore
if exception := task.exception(): else:
print_exception(type(exception), exception, exception.__traceback__) data = cast(str, msg.data)
payload = json.loads(data)
self.loop.create_task(self.handle_event(payload))
hb.cancel()
if not reconnect:
return
+40
View File
@@ -0,0 +1,40 @@
from __future__ import annotations
from typing import Any, Callable, Dict, List, Optional, Tuple
from typing_extensions import Protocol
class _FileLike(Protocol):
def read(self, n: int) -> bytes: ...
def unpackb(
packed: bytes,
file_like: Optional[_FileLike] = ...,
read_size: int = ...,
use_list: bool = ...,
raw: bool = ...,
timestamp: int = ...,
strict_map_key: bool = ...,
object_hook: Optional[Callable[[Dict[Any, Any]], Any]] = ...,
object_pairs_hook: Optional[Callable[[List[Tuple[Any, Any]]], Any]] = ...,
list_hook: Optional[Callable[[List[Any]], Any]] = ...,
unicode_errors: Optional[str] = ...,
max_buffer_size: int = ...,
ext_hook: Callable[[int, bytes], Any] = ...,
max_str_len: int = ...,
max_bin_len: int = ...,
max_array_len: int = ...,
max_map_len: int = ...,
max_ext_len: int = ...,
) -> Any: ...
def packb(
o: Any,
default: Optional[Callable[[Any], Any]] = ...,
use_single_float: bool = ...,
autoreset: bool = ...,
use_bin_type: bool = ...,
strict_types: bool = ...,
datetime: bool = ...,
unicode_errors: Optional[str] = ...,
) -> bytes: ...
+1
View File
@@ -0,0 +1 @@
def get_html_theme_path() -> str: ...