From 042a7da85ec4cd3934288ec90b197f35271e9e37 Mon Sep 17 00:00:00 2001 From: Zomatree Date: Sat, 24 Jun 2023 01:19:34 +0100 Subject: [PATCH] fix help command argument parsing --- revolt/ext/commands/command.py | 3 +- revolt/ext/commands/group.py | 68 +++++++++++++++++++++++++++++++- revolt/ext/commands/help.py | 72 +++++++++++++++++++--------------- 3 files changed, 108 insertions(+), 35 deletions(-) diff --git a/revolt/ext/commands/command.py b/revolt/ext/commands/command.py index 570aae5..47c7049 100755 --- a/revolt/ext/commands/command.py +++ b/revolt/ext/commands/command.py @@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Annotated, Any, Callable, Coroutine, Generic, Literal, Optional, Union, get_args, get_origin) from typing_extensions import ParamSpec -from revolt.utils import copy_doc, maybe_coroutine +from revolt.utils import maybe_coroutine from .errors import InvalidLiteralArgument, UnionConverterError from .utils import ClientCoT, evaluate_parameters @@ -74,7 +74,6 @@ class Command(Generic[ClientCoT]): except Exception as err: return await self._error_handler(self.cog or context.client, context, err) - @copy_doc(invoke) def __call__(self, context: Context[ClientCoT], *args: Any, **kwargs: Any) -> Any: return self.invoke(context, *args, **kwargs) diff --git a/revolt/ext/commands/group.py b/revolt/ext/commands/group.py index c7494ed..ba7a22b 100755 --- a/revolt/ext/commands/group.py +++ b/revolt/ext/commands/group.py @@ -53,6 +53,10 @@ class Group(Command[ClientCoT]): command = cls(func, name or func.__name__, aliases or []) command.parent = self self.subcommands[command.name] = command + + for alias in command.aliases: + self.subcommands[alias] = command + return command return inner @@ -80,6 +84,10 @@ class Group(Command[ClientCoT]): command = cls(func, name or func.__name__, aliases or []) command.parent = self self.subcommands[command.name] = command + + for alias in command.aliases: + self.subcommands[alias] = command + return command return inner @@ -89,7 +97,65 @@ class Group(Command[ClientCoT]): @property def commands(self) -> list[Command[ClientCoT]]: - return list(self.subcommands.values()) + """Gets all commands registered + + Returns + -------- + list[:class:`Command`] + The registered commands + """ + return list(set(self.subcommands.values())) + + def get_command(self, name: str) -> Command[ClientCoT]: + """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[ClientCoT]) -> 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[ClientCoT]]: + """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]] = Group) -> Callable[[Callable[..., Coroutine[Any, Any, Any]]], Group[ClientT]]: """A decorator that turns a function into a :class:`Group` diff --git a/revolt/ext/commands/help.py b/revolt/ext/commands/help.py index 91fbcff..5d64111 100755 --- a/revolt/ext/commands/help.py +++ b/revolt/ext/commands/help.py @@ -1,7 +1,7 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Generic, Optional, TypedDict, Union +from typing import TYPE_CHECKING, Generic, Optional, TypedDict, Union, cast from typing_extensions import NotRequired @@ -11,9 +11,9 @@ from .context import Context from .group import Group from .utils import ClientCoT, ClientT -if TYPE_CHECKING: - from revolt import File, Message, Messageable, MessageReply, SendableEmbed +from revolt import File, Message, Messageable, MessageReply, SendableEmbed +if TYPE_CHECKING: from .cog import Cog __all__ = ("MessagePayload", "HelpCommand", "DefaultHelpCommand", "help_command_impl") @@ -28,7 +28,7 @@ class MessagePayload(TypedDict): class HelpCommand(ABC, Generic[ClientCoT]): @abstractmethod - async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]: + async def create_global_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]: raise NotImplementedError @abstractmethod @@ -73,19 +73,14 @@ class HelpCommand(ABC, Generic[ClientCoT]): return context @abstractmethod - async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> Any: + async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> Union[str, SendableEmbed, MessagePayload]: raise NotImplementedError - @abstractmethod - async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> Any: - raise NotImplementedError - - class DefaultHelpCommand(HelpCommand[ClientCoT]): def __init__(self, default_cog_name: str = "No Cog"): self.default_cog_name = default_cog_name - async def create_bot_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]: + async def create_global_help(self, context: Context[ClientCoT], commands: dict[Optional[Cog[ClientCoT]], list[Command[ClientCoT]]]) -> Union[str, SendableEmbed, MessagePayload]: lines = ["```"] for cog, cog_commands in commands.items(): @@ -145,14 +140,8 @@ class DefaultHelpCommand(HelpCommand[ClientCoT]): lines.append("```") return "\n".join(lines) - async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> None: - channel = await self.get_channel(context) - await channel.send(f"Command `{name}` not found.") - - async def handle_no_cog_found(self, context: Context[ClientCoT], name: str) -> None: - channel = await self.get_channel(context) - await channel.send(f"Cog `{name}` not found.") - + async def handle_no_command_found(self, context: Context[ClientCoT], name: str) -> str: + return f"Command `{name}` not found." class HelpCommandImpl(Command[ClientCoT]): def __init__(self, client: ClientCoT): @@ -165,34 +154,53 @@ class HelpCommandImpl(Command[ClientCoT]): self.description: str | None = "Shows help for a command, cog or the entire bot" -async def help_command_impl(self: ClientT, context: Context[ClientT], *arguments: str) -> None: - help_command = self.help_command +async def help_command_impl(client: ClientT, context: Context[ClientT], *arguments: str) -> None: + help_command = client.help_command if not help_command: return - filtered_commands = await help_command.filter_commands(context, self.commands) + filtered_commands = await help_command.filter_commands(context, client.commands) commands = await help_command.group_commands(context, filtered_commands) if not arguments: - payload = await help_command.create_bot_help(context, commands) - else: - command_name = arguments[0] + payload = await help_command.create_global_help(context, commands) - try: - command = self.get_command(command_name) - except KeyError: - cog = self.cogs.get(command_name) - if cog: - payload = await help_command.create_cog_help(context, cog) + else: + parent: ClientT | Group[ClientT] = 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 + + if isinstance(command, Group): + command = cast(Group[ClientT], command) + parent = command else: - return await help_command.handle_no_command_found(context, command_name) + payload = await help_command.create_command_help(context, command) + break else: + + if TYPE_CHECKING: + command = cast(Command[ClientT], ...) + 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 if isinstance(payload, str):