From d013320bec90117c89e510951abb3951e7c1ecc6 Mon Sep 17 00:00:00 2001 From: Alero Date: Fri, 14 Feb 2025 19:15:19 +0800 Subject: [PATCH] feat: more powerful CustomFilter --- astrbot/api/event/filter/__init__.py | 4 +- astrbot/core/star/filter/command.py | 12 +++- astrbot/core/star/filter/command_group.py | 14 +++- astrbot/core/star/filter/custom_filter.py | 14 ++++ astrbot/core/star/filter/permission.py | 10 +-- astrbot/core/star/register/__init__.py | 4 +- astrbot/core/star/register/star_handler.py | 78 ++++++++++++++++------ 7 files changed, 101 insertions(+), 35 deletions(-) create mode 100644 astrbot/core/star/filter/custom_filter.py diff --git a/astrbot/api/event/filter/__init__.py b/astrbot/api/event/filter/__init__.py index c53536634..89729219d 100644 --- a/astrbot/api/event/filter/__init__.py +++ b/astrbot/api/event/filter/__init__.py @@ -5,7 +5,7 @@ from astrbot.core.star.register import ( register_regex as regex, register_platform_adapter_type as platform_adapter_type, register_permission_type as permission_type, - register_custom_permission_type as custom_permission_type, + register_custom_filter as custom_filter, register_on_llm_request as on_llm_request, register_on_llm_response as on_llm_response, register_llm_tool as llm_tool, @@ -29,7 +29,7 @@ __all__ = [ 'PlatformAdapterTypeFilter', 'PlatformAdapterType', 'PermissionTypeFilter', - 'custom_permission_type', + 'custom_filter', 'PermissionType', 'on_llm_request', 'llm_tool', diff --git a/astrbot/core/star/filter/command.py b/astrbot/core/star/filter/command.py index 3e1394f38..41fdb61dc 100644 --- a/astrbot/core/star/filter/command.py +++ b/astrbot/core/star/filter/command.py @@ -1,10 +1,12 @@ import re import inspect +from typing import List from . import HandlerFilter from astrbot.core.platform.astr_message_event import AstrMessageEvent from astrbot.core.config import AstrBotConfig from astrbot.core.utils.param_validation_mixin import ParameterValidationMixin +from .custom_filter import CustomFilter from ..star_handler import StarHandlerMetadata # 标准指令受到 wake_prefix 的制约。 @@ -14,6 +16,7 @@ class CommandFilter(HandlerFilter, ParameterValidationMixin): self.command_name = command_name if handler_md: self.init_handler_md(handler_md) + self.custom_filter_list: List[CustomFilter] = [] def print_types(self): result = "" @@ -42,10 +45,17 @@ class CommandFilter(HandlerFilter, ParameterValidationMixin): def get_handler_md(self) -> StarHandlerMetadata: return self.handler_md + def add_custom_filter(self, custom_filter: CustomFilter): + self.custom_filter_list.append(custom_filter) + def filter(self, event: AstrMessageEvent, cfg: AstrBotConfig) -> bool: if not event.is_at_or_wake_command: return False - + + for custom_filter in self.custom_filter_list: + if not custom_filter.filter(event, cfg): + raise ValueError(f"没有执行该指令(组)的权限\n") + if event.get_extra("parsing_command"): message_str = event.get_extra("parsing_command").strip() else: diff --git a/astrbot/core/star/filter/command_group.py b/astrbot/core/star/filter/command_group.py index 0dca9916f..0bb298cca 100644 --- a/astrbot/core/star/filter/command_group.py +++ b/astrbot/core/star/filter/command_group.py @@ -6,6 +6,7 @@ from . import HandlerFilter from .command import CommandFilter from astrbot.core.platform.astr_message_event import AstrMessageEvent from astrbot.core.config import AstrBotConfig +from .custom_filter import CustomFilter from ..star_handler import StarHandlerMetadata # 指令组受到 wake_prefix 的制约。 @@ -13,10 +14,14 @@ class CommandGroupFilter(HandlerFilter): def __init__(self, group_name: str): self.group_name = group_name self.sub_command_filters: List[Union[CommandFilter, CommandGroupFilter]] = [] + self.custom_filter_list: List[CustomFilter] = [] def add_sub_command_filter(self, sub_command_filter: Union[CommandFilter, CommandGroupFilter]): self.sub_command_filters.append(sub_command_filter) - + + def add_custom_filter(self, custom_filter: CustomFilter): + self.custom_filter_list.append(custom_filter) + # 以树的形式打印出来 def print_cmd_tree(self, sub_command_filters: List[Union[CommandFilter, CommandGroupFilter]], prefix: str = "") -> str: result = "" @@ -61,7 +66,12 @@ class CommandGroupFilter(HandlerFilter): # 当前还是指令组 tree = self.group_name + "\n" + self.print_cmd_tree(self.sub_command_filters) raise ValueError(f"指令组 {self.group_name} 未填写完全。这个指令组下有如下指令:\n"+tree) - + + # 判断当前指令组的自定义过滤器 + for custom_filter in self.custom_filter_list: + if not custom_filter.filter(event, cfg): + raise ValueError(f"没有执行该指令(组)的权限\n") + child_command_handler_md = None for sub_filter in self.sub_command_filters: if isinstance(sub_filter, CommandFilter): diff --git a/astrbot/core/star/filter/custom_filter.py b/astrbot/core/star/filter/custom_filter.py new file mode 100644 index 000000000..cfe80099e --- /dev/null +++ b/astrbot/core/star/filter/custom_filter.py @@ -0,0 +1,14 @@ +from abc import abstractmethod + +from . import HandlerFilter +from astrbot.core.platform.astr_message_event import AstrMessageEvent +from astrbot.core.config import AstrBotConfig + +class CustomFilter(HandlerFilter): + def __init__(self, raise_error: bool = True): + self.raise_error = raise_error + + @abstractmethod + def filter(self, event: AstrMessageEvent, cfg: AstrBotConfig) -> bool: + ''' 一个用于重写的自定义Filter ''' + raise NotImplementedError diff --git a/astrbot/core/star/filter/permission.py b/astrbot/core/star/filter/permission.py index 23d357338..5e961aa5b 100644 --- a/astrbot/core/star/filter/permission.py +++ b/astrbot/core/star/filter/permission.py @@ -9,15 +9,7 @@ class PermissionType(enum.Flag): ADMIN = enum.auto() MEMBER = enum.auto() -class BasePermissionTypeFilter(HandlerFilter): - def __init__(self, raise_error: bool = True): - self.raise_error = raise_error - - def filter(self, event: AstrMessageEvent, cfg: AstrBotConfig) -> bool: - ''' 一个用于重写的自定义Filter ''' - raise NotImplementedError - -class PermissionTypeFilter(BasePermissionTypeFilter): +class PermissionTypeFilter(HandlerFilter): def __init__(self, permission_type: PermissionType, raise_error: bool = True): self.permission_type = permission_type self.raise_error = raise_error diff --git a/astrbot/core/star/register/__init__.py b/astrbot/core/star/register/__init__.py index 56f40e0e3..ba51f7ab6 100644 --- a/astrbot/core/star/register/__init__.py +++ b/astrbot/core/star/register/__init__.py @@ -6,7 +6,7 @@ from .star_handler import ( register_platform_adapter_type, register_regex, register_permission_type, - register_custom_permission_type, + register_custom_filter, register_on_llm_request, register_on_llm_response, register_llm_tool, @@ -22,7 +22,7 @@ __all__ = [ 'register_platform_adapter_type', 'register_regex', 'register_permission_type', - 'register_custom_permission_type', + 'register_custom_filter', 'register_on_llm_request', 'register_on_llm_response', 'register_llm_tool', diff --git a/astrbot/core/star/register/star_handler.py b/astrbot/core/star/register/star_handler.py index 7ff0412b8..320915a5c 100644 --- a/astrbot/core/star/register/star_handler.py +++ b/astrbot/core/star/register/star_handler.py @@ -1,14 +1,16 @@ from __future__ import annotations import docstring_parser +import traceback from ..star_handler import star_handlers_registry, StarHandlerMetadata, EventType from ..filter.command import CommandFilter from ..filter.command_group import CommandGroupFilter from ..filter.event_message_type import EventMessageTypeFilter, EventMessageType from ..filter.platform_adapter_type import PlatformAdapterTypeFilter, PlatformAdapterType -from ..filter.permission import PermissionTypeFilter, BasePermissionTypeFilter, PermissionType +from ..filter.permission import PermissionTypeFilter, PermissionType +from ..filter.custom_filter import CustomFilter from ..filter.regex import RegexFilter -from typing import Awaitable +from typing import Awaitable, Union from astrbot.core.provider.func_tool_manager import SUPPORTED_TYPES from astrbot.core.provider.register import llm_tools from astrbot.core import logger @@ -80,6 +82,58 @@ def register_command(command_name: str = None, *args, **kwargs): return decorator +def register_custom_filter(custom_type_filter: CustomFilter, *args, **kwargs): + '''注册一个自定义的 CustomFilter + + Args: + cunstom_permission_type_filter: 在裸指令时为CustomFilter对象 + 在指令组时为父指令的RegisteringCommandable对象,即self或者command_group的返回 + raise_error: 如果没有权限,是否抛出错误到消息平台,并且停止事件传播。默认为 True + ''' + add_to_event_filters = False + raise_error = True + + # 判断是否是指令组,指令组则添加到指令组的CommandGroupFilter对象中在waking_check的时候一起判断 + if isinstance(custom_type_filter, RegisteringCommandable): + # 子指令, 此时函数为RegisteringCommandable对象的方法,首位参数为RegisteringCommandable对象的self。 + parent_rigister_commandable = custom_type_filter + custom_filter = args[0] + if len(args) > 1: + raise_error = args[1] + else: + # 裸指令 + add_to_event_filters = True + custom_filter = custom_type_filter + if args: + raise_error = args[0] + + def decorator(awaitable): + # 裸指令,子指令与指令组的区分,指令组会因为标记跳过wake。 + if not add_to_event_filters and isinstance(awaitable, RegisteringCommandable): + # 指令组,添加到本层的grouphandle中一起判断 + awaitable.parent_group.add_custom_filter(custom_filter(raise_error)) + else: + handler_md = get_handler_or_create(awaitable, EventType.AdapterMessageEvent, **kwargs) + + if not add_to_event_filters and not isinstance(awaitable, RegisteringCommandable): + # 底层子指令 + handle_full_name = get_handler_full_name(awaitable) + command_handle = None + for sub_handle in parent_rigister_commandable.parent_group.sub_command_filters: + # 所有符合fullname一致的子指令handle添加自定义过滤器。 + # 不确定是否会有多个子指令有一样的fullname,比如一个方法添加多个command装饰器? + sub_handle_md = sub_handle.get_handler_md() + if sub_handle_md and sub_handle_md.handler_full_name == handle_full_name: + sub_handle.add_custom_filter(custom_filter(raise_error)) + + else: + # 裸指令 + handler_md = get_handler_or_create(awaitable, EventType.AdapterMessageEvent, **kwargs) + handler_md.event_filters.append(custom_filter(raise_error)) + + return awaitable + return decorator + def register_command_group(command_group_name: str = None, *args, **kwargs): '''注册一个 CommandGroup ''' @@ -102,7 +156,7 @@ def register_command_group(command_group_name: str = None, *args, **kwargs): # 根指令组 handler_md = get_handler_or_create(obj, EventType.AdapterMessageEvent, **kwargs) handler_md.event_filters.append(new_group) - + return RegisteringCommandable(new_group) return decorator @@ -111,7 +165,8 @@ class RegisteringCommandable(): '''用于指令组级联注册''' group = register_command_group command = register_command - + custom_filter = register_custom_filter + def __init__(self, parent_group: CommandGroupFilter): self.parent_group = parent_group @@ -156,21 +211,6 @@ def register_permission_type(permission_type: PermissionType, raise_error: bool return decorator -def register_custom_permission_type(cunstom_permission_type_filter: BasePermissionTypeFilter, raise_error: bool = True): - '''注册一个自定义的 PermissionFilter - - Args: - cunstom_permission_type_filter: 一个继承自HandlerFilter - raise_error: 如果没有权限,是否抛出错误到消息平台,并且停止事件传播。默认为 True - ''' - - def decorator(awaitable): - handler_md = get_handler_or_create(awaitable, EventType.AdapterMessageEvent) - handler_md.event_filters.append(cunstom_permission_type_filter()) - return awaitable - - return decorator - def register_on_llm_request(**kwargs): '''当有 LLM 请求时的事件