diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index 5ddf75c..cbad7d8 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -55,6 +55,8 @@ from .event import BaseEvent as BaseEvent from .event import attr as attr from .event import register_internal_event as register_internal_event +from .event.api import SendRequest as SendRequest +from .event.api import SendResponse as SendResponse from .event.base import MessageCreatedEvent as MessageCreatedEvent from .event.base import MessageEvent as MessageEvent from .event.base import Reply as Reply @@ -63,8 +65,6 @@ from .event.lifespan import Cleanup as Cleanup from .event.lifespan import Ready as Ready from .event.lifespan import Startup as Startup -from .event.send import SendRequest as SendRequest -from .event.send import SendResponse as SendResponse from .filter import filter_ as filter_ from .localdata import local_data as local_data from .message import MessageChain as MessageChain diff --git a/arclet/entari/command/plugin.py b/arclet/entari/command/plugin.py index 28fd654..70c427b 100644 --- a/arclet/entari/command/plugin.py +++ b/arclet/entari/command/plugin.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from typing import Any from typing_extensions import TypeVar, deprecated @@ -28,7 +29,7 @@ async def _after_execute(ctx: Contexts, session: Session | None = None): - result = ctx[RESULT] + result: str | MessageChain | _ExitException | None = ctx[RESULT] event = ctx[EVENT] if result is not None: if isinstance(result, _ExitException): @@ -55,7 +56,7 @@ def assign(self, path: str, value: Any = _seminal, or_not: bool = False, priorit class AlconnaPluginDispatcher(PluginDispatcher[T]): def __init__(self, plugin: Plugin, command: Alconna, need_reply_me: bool = False, need_notice_me: bool = False, use_config_prefix: bool = True, block: bool = True, skip_for_unmatch: bool = True): # noqa: E501 plugin._extra.setdefault("commands", []).append((command.prefixes, command.command)) - self.cache = LRU(10) + self.cache: "LRU[str, asyncio.Future]" = LRU(10) # noqa: UP037 self.supplier = AlconnaSuppiler(command, self.cache, block, skip_for_unmatch) super().__init__(plugin, MessageCreatedEvent, command.path) plugin.collect( diff --git a/arclet/entari/core.py b/arclet/entari/core.py index ce2cc35..1dd2c90 100644 --- a/arclet/entari/core.py +++ b/arclet/entari/core.py @@ -46,10 +46,10 @@ ITEM_SESSION, ITEM_USER, ) +from .event.api import SendResponse from .event.base import MessageCreatedEvent, event_parse from .event.config import ConfigReload from .event.lifespan import AccountUpdate -from .event.send import SendResponse from .localdata import local_data from .logger import apply_log_save, enable_rich_except, log from .message import MessageChain diff --git a/arclet/entari/event/api.py b/arclet/entari/event/api.py new file mode 100644 index 0000000..a5c392c --- /dev/null +++ b/arclet/entari/event/api.py @@ -0,0 +1,107 @@ +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from arclet.letoderea import Contexts, Result, define, provide +from satori import ChannelType +from satori.client import Account +from satori.exception import ActionFailed +from satori.model import Channel, MessageObject + +from ..const import ITEM_ACCOUNT, ITEM_CHANNEL, ITEM_MESSAGE_CONTENT, ITEM_SESSION +from ..message import MessageChain + +if TYPE_CHECKING: + from ..session import Session + + +@dataclass +class SendRequest: + account: Account + channel: str + message: MessageChain + session: "Session | None" = None + + def check_result(self, value) -> Result[bool | MessageChain] | None: + if isinstance(value, bool | MessageChain): + return Result(value) + + +before_send_pub = define(SendRequest, name="entari.event/before_send") + + +@before_send_pub.gather +async def send_req_gather(req: SendRequest, context: Contexts): + context[ITEM_ACCOUNT] = req.account + context[ITEM_MESSAGE_CONTENT] = req.message + if req.session: + context[ITEM_SESSION] = req.session + context[ITEM_CHANNEL] = req.session.channel + else: + try: + context[ITEM_CHANNEL] = await req.account.channel_get(req.channel) + except ActionFailed: + context[ITEM_CHANNEL] = Channel( + req.channel, ChannelType.DIRECT if req.channel.startswith("private:") else ChannelType.TEXT + ) + + +@dataclass +class SendResponse: + account: Account + channel: str + message: MessageChain + result: list[MessageObject] + session: "Session | None" = None + + +send_pub = define(SendResponse, name="entari.event/after_send") +send_pub.providers.append(provide(list[MessageObject], call="$resp_result")) + + +@send_pub.gather +async def send_resp_gather(resp: SendResponse, context: Contexts): + context[ITEM_ACCOUNT] = resp.account + context[ITEM_MESSAGE_CONTENT] = resp.message + context["$resp_result"] = resp.result + if resp.session: + context[ITEM_SESSION] = resp.session + context[ITEM_CHANNEL] = resp.session.channel + else: + try: + context[ITEM_CHANNEL] = await resp.account.channel_get(resp.channel) + except ActionFailed: + context[ITEM_CHANNEL] = Channel( + resp.channel, ChannelType.DIRECT if resp.channel.startswith("private:") else ChannelType.TEXT + ) + + +@dataclass +class APIRequest: + account: Account + name: str + params: dict[str, Any] + + +before_api_pub = define(APIRequest, name="entari.event/before_api_call") + + +@before_api_pub.gather +async def call_req_gather(req: APIRequest, context: Contexts): + context[ITEM_ACCOUNT] = req.account + + +@dataclass +class APIResponse: + account: Account + name: str + params: dict[str, Any] + success: bool + result: Any + + +after_api_pub = define(APIResponse, name="entari.event/after_api_call") + + +@after_api_pub.gather +async def call_resp_gather(resp: APIResponse, context: Contexts): + context[ITEM_ACCOUNT] = resp.account diff --git a/arclet/entari/event/base.py b/arclet/entari/event/base.py index 60261eb..f3234da 100644 --- a/arclet/entari/event/base.py +++ b/arclet/entari/event/base.py @@ -53,7 +53,7 @@ def _is_notice_me(message: MessageChain, account: Account): def _remove_notice_me(message: MessageChain, account: Account): - message = message.copy() + message = message.fork() message.pop(0) if _is_notice_me(message, account): message.pop(0) @@ -316,7 +316,7 @@ def __init__(self, account: Account, origin: OriginEvent): super().__init__(account, origin) self.content = MessageChain(self.message.message) if self.content.has(Quote): - self.quote = self.content.get(Quote, 1)[0] + self.quote = self.content.get_first(Quote) self.content = self.content.exclude(Quote) async def gather(self, context: Contexts): diff --git a/arclet/entari/event/send.py b/arclet/entari/event/send.py index de8991e..ebbb250 100644 --- a/arclet/entari/event/send.py +++ b/arclet/entari/event/send.py @@ -1,75 +1,10 @@ -from dataclasses import dataclass -from typing import TYPE_CHECKING +from warnings import warn -from arclet.letoderea import Contexts, Result, define, provide -from satori import ChannelType -from satori.client import Account -from satori.exception import ActionFailed -from satori.model import Channel, MessageObject +warn( + "arclet.entari.event.send is deprecated, please use arclet.entari.event.api instead", + DeprecationWarning, + stacklevel=2, +) -from ..const import ITEM_ACCOUNT, ITEM_CHANNEL, ITEM_MESSAGE_CONTENT, ITEM_SESSION -from ..message import MessageChain - -if TYPE_CHECKING: - from ..session import Session - - -@dataclass -class SendRequest: - account: Account - channel: str - message: MessageChain - session: "Session | None" = None - - def check_result(self, value) -> Result[bool | MessageChain] | None: - if isinstance(value, bool | MessageChain): - return Result(value) - - -before_send_pub = define(SendRequest, name="entari.event/before_send") - - -@before_send_pub.gather -async def req_gather(req: SendRequest, context: Contexts): - context[ITEM_ACCOUNT] = req.account - context[ITEM_MESSAGE_CONTENT] = req.message - if req.session: - context[ITEM_SESSION] = req.session - context[ITEM_CHANNEL] = req.session.channel - else: - try: - context[ITEM_CHANNEL] = await req.account.channel_get(req.channel) - except ActionFailed: - context[ITEM_CHANNEL] = Channel( - req.channel, ChannelType.DIRECT if req.channel.startswith("private:") else ChannelType.TEXT - ) - - -@dataclass -class SendResponse: - account: Account - channel: str - message: MessageChain - result: list[MessageObject] - session: "Session | None" = None - - -send_pub = define(SendResponse, name="entari.event/after_send") -send_pub.providers.append(provide(list[MessageObject], call="$resp_result")) - - -@send_pub.gather -async def resp_gather(resp: SendResponse, context: Contexts): - context[ITEM_ACCOUNT] = resp.account - context[ITEM_MESSAGE_CONTENT] = resp.message - context["$resp_result"] = resp.result - if resp.session: - context[ITEM_SESSION] = resp.session - context[ITEM_CHANNEL] = resp.session.channel - else: - try: - context[ITEM_CHANNEL] = await resp.account.channel_get(resp.channel) - except ActionFailed: - context[ITEM_CHANNEL] = Channel( - resp.channel, ChannelType.DIRECT if resp.channel.startswith("private:") else ChannelType.TEXT - ) +from .api import SendRequest as SendRequest # noqa: F401 +from .api import SendResponse as SendResponse # noqa: F401 diff --git a/arclet/entari/filter/__init__.py b/arclet/entari/filter/__init__.py index dc683a7..039f493 100644 --- a/arclet/entari/filter/__init__.py +++ b/arclet/entari/filter/__init__.py @@ -1,18 +1,19 @@ -import asyncio import inspect from collections.abc import Awaitable, Callable -from datetime import datetime from typing import Final, TypeAlias from typing_extensions import ParamSpec -from arclet.letoderea import STOP, Propagator, enter_if, propagate -from arclet.letoderea.utils import TCallable +from arclet.letoderea import enter_if from tarina import is_coroutinefunction -from ..config import EntariConfig -from ..message import MessageChain from ..session import Session from . import common +from .limit import interval as interval +from .limit import semaphore as semaphore +from .message import endswith as endswith +from .message import startswith as startswith +from .permission import admins as admins +from .permission import superusers as superusers _SessionFilter: TypeAlias = Callable[[Session], bool] | Callable[[Session], Awaitable[bool]] @@ -62,117 +63,3 @@ async def _(*args, _func=func, **kwargs): filter_: Final[_Filter] = _Filter() F = filter_ - - -class interval(Propagator): - def __init__(self, value: float, limit_prompt: str | MessageChain | None = None, priority: int = 80): - self.success = True - self.value = value - self.priority = priority - self.limit_prompt = limit_prompt - self.last_times: dict[str, datetime] = {} - - async def before(self, session: Session | None = None): - session_id = ( - "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" - ) - last_time = self.last_times.get(session_id, None) - if not last_time: - return - self.success = (datetime.now() - last_time).total_seconds() > self.value - if not self.success: - if session and self.limit_prompt: - await session.send(self.limit_prompt) - return STOP - - async def after(self, session: Session | None = None): - session_id = ( - "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" - ) - self.last_times[session_id] = datetime.now() - - def compose(self): - yield self.before, True, self.priority - yield self.after, False, self.priority - - def __call__(self, func: TCallable) -> TCallable: - return propagate(self)(func) - - -class semaphore(Propagator): - def __init__(self, count: int, limit_prompt: str | MessageChain | None = None, priority: int = 80): - self.count = count - self.limit_prompt = limit_prompt - self.priority = priority - self.semaphores: dict[str, asyncio.Semaphore] = {} - - async def before(self, session: Session | None = None): - session_id = ( - "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" - ) - if session_id not in self.semaphores: - self.semaphores[session_id] = asyncio.Semaphore(self.count) - if not await self.semaphores[session_id].acquire(): - if session and self.limit_prompt: - await session.send(self.limit_prompt) - return STOP - - async def after(self, session: Session | None = None): - session_id = ( - "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" - ) - if session_id not in self.semaphores: - self.semaphores[session_id] = asyncio.Semaphore(self.count) - self.semaphores[session_id].release() - - def compose(self): - yield self.before, True, self.priority - yield self.after, False, self.priority - - def __call__(self, func: TCallable) -> TCallable: - return propagate(self)(func) - - -class superusers(Propagator): - - async def check(self, session: Session | None = None): - if not session: - return STOP - config = EntariConfig.instance.basic.superusers - if session.account.platform not in config: - return STOP - if not session.event.user: - return STOP - if session.event.user.id not in config[session.account.platform]: - return STOP - - def compose(self): - yield self.check, True, 50 - - def __call__(self, func: TCallable) -> TCallable: - return propagate(self)(func) - - -class admins(Propagator): - - async def check(self, session: Session | None = None): - if not session: - return STOP - if session.event.member and session.event.member.roles: - for role in session.event.member.roles: - if any(keyword in role.id.lower() for keyword in ("admin", "administrator", "owner")): - return - config = EntariConfig.instance.basic.superusers - if ( - session.account.platform in config - and session.event.user - and session.event.user.id in config[session.account.platform] - ): - return - return STOP - - def compose(self): - yield self.check, True, 50 - - def __call__(self, func: TCallable) -> TCallable: - return propagate(self)(func) diff --git a/arclet/entari/filter/limit.py b/arclet/entari/filter/limit.py new file mode 100644 index 0000000..5df1890 --- /dev/null +++ b/arclet/entari/filter/limit.py @@ -0,0 +1,77 @@ +import asyncio +from datetime import datetime + +from arclet.letoderea import STOP, Propagator, propagate +from arclet.letoderea.utils import TCallable + +from ..message import MessageChain +from ..session import Session + + +class interval(Propagator): + def __init__(self, value: float, limit_prompt: str | MessageChain | None = None, priority: int = 80): + self.success = True + self.value = value + self.priority = priority + self.limit_prompt = limit_prompt + self.last_times: dict[str, datetime] = {} + + async def before(self, session: Session | None = None): + session_id = ( + "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" + ) + last_time = self.last_times.get(session_id, None) + if not last_time: + return + self.success = (datetime.now() - last_time).total_seconds() > self.value + if not self.success: + if session and self.limit_prompt: + await session.send(self.limit_prompt) + return STOP + + async def after(self, session: Session | None = None): + session_id = ( + "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" + ) + self.last_times[session_id] = datetime.now() + + def compose(self): + yield self.before, True, self.priority + yield self.after, False, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +class semaphore(Propagator): + def __init__(self, count: int, limit_prompt: str | MessageChain | None = None, priority: int = 80): + self.count = count + self.limit_prompt = limit_prompt + self.priority = priority + self.semaphores: dict[str, asyncio.Semaphore] = {} + + async def before(self, session: Session | None = None): + session_id = ( + "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" + ) + if session_id not in self.semaphores: + self.semaphores[session_id] = asyncio.Semaphore(self.count) + if not await self.semaphores[session_id].acquire(): + if session and self.limit_prompt: + await session.send(self.limit_prompt) + return STOP + + async def after(self, session: Session | None = None): + session_id = ( + "$global" if not session else f"{session.account.platform}/{session.account.self_id}/{session.channel.id}" + ) + if session_id not in self.semaphores: + self.semaphores[session_id] = asyncio.Semaphore(self.count) + self.semaphores[session_id].release() + + def compose(self): + yield self.before, True, self.priority + yield self.after, False, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) diff --git a/arclet/entari/filter/message.py b/arclet/entari/filter/message.py new file mode 100644 index 0000000..c3f3cda --- /dev/null +++ b/arclet/entari/filter/message.py @@ -0,0 +1,216 @@ +import re +from typing import Any + +from arclet.letoderea import STOP, Contexts, Propagator, deref, propagate, provide +from arclet.letoderea.utils import TCallable +from nepattern import ANY, BasePattern, MatchMode, parser +from satori import Text +from tarina import Empty + +from ..const import ITEM_MESSAGE_CONTENT +from ..message import MessageChain + + +def _prefixed(pat: BasePattern): + if pat.mode not in (MatchMode.REGEX_MATCH, MatchMode.REGEX_CONVERT): + return pat + new_pat = pat.copy() + new_pat.regex_pattern = re.compile(f"^{new_pat.pattern}") + return new_pat + + +def _suffixed(pat: BasePattern): + if pat.mode not in (MatchMode.REGEX_MATCH, MatchMode.REGEX_CONVERT): + return pat + new_pat = pat.copy() + new_pat.regex_pattern = re.compile(f"{new_pat.pattern}$") + return new_pat + + +class startswith(Propagator): + def __init__(self, prefix: Any, include: bool = False, bind: str | None = None, priority: int = 80): + """ + 前缀匹配 + + Args: + prefix: 需要匹配的前缀, 支持格式有 a|b , ['a', At(...)] 等 + include: 指示消息链是否仅返回前缀被匹配的部分, 默认为 False + bind: 指定注入返回值的参数名称,未指定则注入到所有的 MessageChain 参数中 + priority: 优先级 + """ + self.prefix = prefix + self.priority = priority + self.include = include + self.bind = bind + + pattern = BasePattern(prefix, mode=MatchMode.REGEX_MATCH) if isinstance(prefix, str) else parser(prefix) + if pattern in (ANY, Empty): + raise ValueError(prefix) + self.pattern = _prefixed(pattern) + + def providers(self): + if self.bind: + return [provide(MessageChain, self.bind, call=f"$startswith_{self.bind}", priority=4)] + return [] + + async def before(self, ctx: Contexts, message: MessageChain): + message = message.fork() + if message: + elem = message[0] + if isinstance(elem, Text) and (res := self.pattern.validate(elem.text)).success: + if self.include: + message = MessageChain(Text(str(res.value()))) + else: + message[0] = Text(elem.text[len(str(res.value())) :].lstrip()) + elif self.pattern.validate(elem).success: + if self.include: + message = MessageChain(elem) + else: + message.remove(elem) + else: + return STOP + if self.bind: + return {f"$startswith_{self.bind}": message} + if ITEM_MESSAGE_CONTENT in ctx: + return {ITEM_MESSAGE_CONTENT: message} + return {"$message": message} + + def compose(self): + yield self.before, True, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +class endswith(Propagator): + def __init__(self, suffix: Any, include: bool = False, bind: str | None = None, priority: int = 80): + """ + 后缀匹配 + + Args: + suffix: 需要匹配的后缀, 支持格式有 a|b , ['a', At(...)] 等 + include: 指示消息链是否仅返回后缀被匹配的部分, 默认为 False + bind: 指定注入返回值的参数名称,未指定则注入到所有的 MessageChain 参数中 + priority: 优先级 + """ + self.suffix = suffix + self.priority = priority + self.include = include + self.bind = bind + + pattern = BasePattern(suffix, mode=MatchMode.REGEX_MATCH) if isinstance(suffix, str) else parser(suffix) + if pattern in (ANY, Empty): + raise ValueError(suffix) + self.pattern = _suffixed(pattern) + + def providers(self): + if self.bind: + return [provide(MessageChain, self.bind, call=f"$endswith_{self.bind}", priority=4)] + return [] + + async def before(self, ctx: Contexts, message: MessageChain): + message = message.fork() + if message: + elem = message[-1] + if isinstance(elem, Text) and (res := self.pattern.validate(elem.text)).success: + if self.include: + message = MessageChain(Text(str(res.value()))) + else: + message[-1] = Text(elem.text[: elem.text.rfind(str(res.value()))].rstrip()) + elif self.pattern.validate(elem).success: + if self.include: + message = MessageChain(elem) + else: + message.remove(elem) + else: + return STOP + if self.bind: + return {f"$endswith_{self.bind}": message} + if ITEM_MESSAGE_CONTENT in ctx: + return {ITEM_MESSAGE_CONTENT: message} + return {"$message": message} + + def compose(self): + yield self.before, True, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +class fullmatch(Propagator): + def __init__( + self, pattern: str | tuple[str, ...], ignorecase: bool = False, bind: str = "fullmatch", priority: int = 80 + ): + """ + 完全匹配 + + Args: + pattern: 指定消息全匹配字符串元组 + ignorecase: 是否忽略大小写, 默认为 False + bind: 指定注入返回值的参数名称,默认为 "fullmatch" + priority: 优先级 + """ + if isinstance(pattern, str): + pattern = (pattern,) + self.pattern = tuple(map(str.casefold, pattern)) if ignorecase else pattern + self.ignorecase = ignorecase + self.priority = priority + self.bind = bind + + def providers(self): + if self.bind: + return [provide(str, self.bind, call=f"$fullmatch_{self.bind}", priority=4)] + return [] + + async def before(self, ctx: Contexts, message: MessageChain): + text = message.extract_plain_text() + if not text: + return STOP + text = text.casefold() if self.ignorecase else text + if text in self.pattern: + return {f"$fullmatch_{self.bind}": text} + return STOP + + def compose(self): + yield self.before, True, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +class regexmatch(Propagator): + def __init__(self, pattern: str, flags: int | re.RegexFlag = 0, priority: int = 80): + """ + 正则匹配,注意正则表达式匹配使用 search 而非 match,如需从头匹配请使用 `r"^xxx"` 来确保匹配开头 + + Args: + pattern: 需要匹配的正则表达式 + flags: 正则匹配标志, 默认为 0 + priority: 优先级 + """ + self.pattern = re.compile(pattern, flags) + self.priority = priority + + def providers(self): + return [provide(re.Match, call="$regexmatch", priority=4)] + + async def before(self, ctx: Contexts, message: MessageChain): + text = message.extract_plain_text() + if not text: + return STOP + if matched := self.pattern.search(text): + return {"$regexmatch": matched} + return STOP + + def compose(self): + yield self.before, True, self.priority + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +def regex_origin(): + return deref(re.Match) + + +__all__ = ["startswith", "endswith", "fullmatch", "regexmatch", "regex_origin"] diff --git a/arclet/entari/filter/permission.py b/arclet/entari/filter/permission.py new file mode 100644 index 0000000..d337b74 --- /dev/null +++ b/arclet/entari/filter/permission.py @@ -0,0 +1,50 @@ +from arclet.letoderea import STOP, Propagator, propagate +from arclet.letoderea.utils import TCallable + +from ..config import EntariConfig +from ..session import Session + + +class superusers(Propagator): + + async def check(self, session: Session | None = None): + if not session: + return STOP + config = EntariConfig.instance.basic.superusers + if session.account.platform not in config: + return STOP + if not session.event.user: + return STOP + if session.event.user.id not in config[session.account.platform]: + return STOP + + def compose(self): + yield self.check, True, 50 + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) + + +class admins(Propagator): + + async def check(self, session: Session | None = None): + if not session: + return STOP + if session.event.member and session.event.member.roles: + for role in session.event.member.roles: + if any(keyword in role.id.lower() for keyword in ("admin", "administrator", "owner")): + return + config = EntariConfig.instance.basic.superusers + if ( + session.account.platform in config + and session.event.user + and session.event.user.id in config[session.account.platform] + ): + return + return STOP + + def compose(self): + yield self.check, True, 50 + + def __call__(self, func: TCallable) -> TCallable: + return propagate(self)(func) diff --git a/arclet/entari/message.py b/arclet/entari/message.py index 3b3c02e..7410a91 100644 --- a/arclet/entari/message.py +++ b/arclet/entari/message.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Iterable, Sequence +from collections.abc import Awaitable, Callable, Iterable, Iterator, MutableSequence, Sequence from copy import deepcopy from dataclasses import dataclass from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar, Union, overload @@ -27,12 +27,8 @@ MessageContainer = Union[str, Element, Sequence["MessageContainer"], "MessageChain[Element]"] -class MessageChain(list[TE]): - """消息序列 - - Args: - message: 消息内容 - """ +class MessageChain(MutableSequence[TE]): + """消息链, 被用于承载整个消息内容的数据结构, 包含有一有序列表, 包含有继承了 Element 的各式类实例.""" @overload def __init__(self): ... @@ -62,7 +58,13 @@ def __init__( self: MessageChain[Element], message: Iterable[str | TE] | str | TE | None = None, ): - super().__init__() + """从传入的序列(可以是元组 tuple, 也可以是列表 list) 创建消息链. + Args: + message (Iterable[str | TE] | str | TE): 包含且仅包含消息元素和字符串的序列 + Returns: + MessageChain: 以传入的序列作为所承载消息的消息链 + """ + self.content: list[TE] = [] if message: if isinstance(message, (str, Element)): self.__iadd__(message) @@ -71,10 +73,18 @@ def __init__( self.__iadd__(i) def __str__(self) -> str: - return "".join(str(elem) for elem in self) + """获取以字符串形式表示的消息链, 且趋于通常你见到的样子. + Returns: + str: 以字符串形式表示的消息链 + """ + return "".join(str(elem) for elem in self.content) def __repr__(self) -> str: - return "[" + ", ".join(repr(elem) for elem in self) + "]" + """获取以字符串形式表示的消息链的详细信息. + Returns: + str: 以字符串形式表示的消息链的详细信息 + """ + return "[" + ", ".join(repr(elem) for elem in self.content) + "]" @overload def __add__(self, other: str) -> MessageChain[TE | Text]: ... @@ -86,17 +96,25 @@ def __add__(self, other: TE | Iterable[TE]) -> MessageChain[TE]: ... def __add__(self, other: TE1 | Iterable[TE1]) -> MessageChain[TE | TE1]: ... def __add__(self, other: str | TE | TE1 | Iterable[TE | TE1]) -> MessageChain: - result: MessageChain = self.fork() + """将另一个消息段或消息链添加到当前消息链. + + Args: + other: 要添加的消息段或消息链 + + Returns: + 添加后的消息链 + """ + result: MessageChain[Element] = self.fork() # type: ignore if isinstance(other, str): - if result and isinstance(text := result[-1], Text): - result[-1] = Text(text.text + other) + if result.content and isinstance(text := result[-1], Text): + result.content[-1] = Text(text.text + other) else: - result.append(Text(other)) + result.content.append(Text(other)) elif isinstance(other, Element): - if result and isinstance(result[-1], Text) and isinstance(other, Text): - result[-1] = Text(result[-1].text + other.text) + if result.content and isinstance(text := result[-1], Text) and isinstance(other, Text): + result.content[-1] = Text(text.text + other.text) else: - result.append(other) + result.content.append(other) elif isinstance(other, Iterable): for elem in other: result += elem @@ -119,15 +137,15 @@ def __radd__(self, other: str | TE1 | Iterable[TE1]) -> MessageChain: def __iadd__(self, other: str | TE | Iterable[TE]) -> Self: if isinstance(other, str): - if self and isinstance(text := self[-1], Text): - list.__setitem__(self, -1, Text(text.text + other)) + if self.content and isinstance(text := self[-1], Text): + self.content[-1] = Text(text.text + other) # type: ignore else: - self.append(Text(other)) # type: ignore + self.content.append(Text(other)) # type: ignore elif isinstance(other, Element): - if self and (isinstance(text := self[-1], Text) and isinstance(other, Text)): - list.__setitem__(self, -1, Text(text.text + other.text)) + if self.content and (isinstance(text := self[-1], Text) and isinstance(other, Text)): + self.content[-1] = Text(text.text + other.text) # type: ignore else: - self.append(other) + self.content.append(other) elif other: for elem in other: self.__iadd__(elem) @@ -136,7 +154,7 @@ def __iadd__(self, other: str | TE | Iterable[TE]) -> Self: return self @overload - def __getitem__(self, args: type[TE1]) -> MessageChain[TE1]: + def __getitem__(self, args: type[TE1], /) -> MessageChain[TE1]: """获取仅包含指定消息段类型的消息 Args: @@ -147,7 +165,7 @@ def __getitem__(self, args: type[TE1]) -> MessageChain[TE1]: """ @overload - def __getitem__(self, args: tuple[type[TE1], int]) -> TE1: + def __getitem__(self, args: tuple[type[TE1], int], /) -> TE1: """索引指定类型的消息段 Args: @@ -158,7 +176,7 @@ def __getitem__(self, args: tuple[type[TE1], int]) -> TE1: """ @overload - def __getitem__(self, args: tuple[type[TE1], slice]) -> MessageChain[TE1]: + def __getitem__(self, args: tuple[type[TE1], slice], /) -> MessageChain[TE1]: """切片指定类型的消息段 Args: @@ -169,7 +187,7 @@ def __getitem__(self, args: tuple[type[TE1], slice]) -> MessageChain[TE1]: """ @overload - def __getitem__(self, args: int) -> TE: + def __getitem__(self, args: int, /) -> TE: """索引消息段 Args: @@ -180,7 +198,7 @@ def __getitem__(self, args: int) -> TE: """ @overload - def __getitem__(self, args: slice) -> Self: + def __getitem__(self, args: slice, /) -> Self: """切片消息段 Args: @@ -196,35 +214,72 @@ def __getitem__( ) -> TE | TE1 | MessageChain[TE1] | Self: arg1, arg2 = args if isinstance(args, tuple) else (args, None) if isinstance(arg1, int) and arg2 is None: - return super().__getitem__(arg1) + return self.content[arg1] if isinstance(arg1, slice) and arg2 is None: - return MessageChain(super().__getitem__(arg1)) # type: ignore + return MessageChain(self.content[arg1]) # type: ignore if TYPE_CHECKING: assert not isinstance(arg1, slice | int) if issubclass(arg1, Element) and arg2 is None: - return MessageChain(elem for elem in self if isinstance(elem, arg1)) # type: ignore + return MessageChain(elem for elem in self.content if isinstance(elem, arg1)) # type: ignore if issubclass(arg1, Element) and isinstance(arg2, int): - return [elem for elem in self if isinstance(elem, arg1)][arg2] + return [elem for elem in self.content if isinstance(elem, arg1)][arg2] if issubclass(arg1, Element) and isinstance(arg2, slice): - return MessageChain([elem for elem in self if isinstance(elem, arg1)][arg2]) # type: ignore + return MessageChain([elem for elem in self.content if isinstance(elem, arg1)][arg2]) # type: ignore raise ValueError("Incorrect arguments to slice") # pragma: no cover - def __contains__(self, value: str | Element | type[Element]) -> bool: - """检查消息段是否存在 + def __setitem__(self, index: int, value: TE | str, /) -> None: + if isinstance(value, str): + value = Text(value) # type: ignore + self.content[index] = value # type: ignore + + def __delitem__(self, index: int, /) -> None: + del self.content[index] + + def __contains__(self, item: str | Element | type[Element] | Self | Sequence[str | Element]) -> bool: + """判断消息链中是否含有特定的内容. Args: - value: 消息段或消息段类型 + item (str | Element | type[Element] | Self | Sequence[str | Element]): 需判断内容. Returns: 消息内是否存在给定消息段或给定类型的消息段 """ - if isinstance(value, type): - return not not next((elem for elem in self if isinstance(elem, value)), None) - if isinstance(value, str): - value = Text(value) - return super().__contains__(value) + if isinstance(item, type): + return not not next((elem for elem in self.content if isinstance(elem, item)), None) + if isinstance(item, Element): + return item in self.merge().content + if isinstance(item, (MessageChain, Sequence)): + return not not self.index_sub(item) + + raise ValueError(f"{item} is not an acceptable argument!") + + def merge(self, *, copy: bool = True) -> Self: + """合并相邻的 Text 项, 选择返回一个新的消息链实例 + + Returns: + MessageChain: 得到的新的消息链实例, 里面不应存在有任何的相邻的 Text 元素. + """ - def has(self, value: str | Element | type[Element]) -> bool: - return value in self + result = [] + + texts = [] + for i in self.content: + if not isinstance(i, Text): + if texts: + result.append(Text("".join(texts))) + texts.clear() # 清空缓存 + result.append(i) + else: + texts.append(i.text) + if texts: + result.append(Text("".join(texts))) + texts.clear() # 清空缓存 + if copy: + return self.__class__(result) + self.content.clear() + self.content.extend(result) + return self + + has = __contains__ def index(self, value: str | Element | type[Element], *args: SupportsIndex) -> int: """索引消息段 @@ -243,114 +298,177 @@ def index(self, value: str | Element | type[Element], *args: SupportsIndex) -> i first_elemment = next((elem for elem in self if isinstance(elem, value)), None) if first_elemment is None: raise ValueError(f"Element with type {value!r} is not in message") - return super().index(first_elemment, *args) + return self.content.index(first_elemment, *args) # type: ignore if isinstance(value, str): value = Text(value) - return super().index(value, *args) # type: ignore + return self.content.index(value, *args) # type: ignore + + def index_sub(self, sub: MessageChain | Sequence[str | Element]) -> list[int]: + """判断消息链是否含有子链. 使用 KMP 算法. + + Args: + sub (MessageChain | Sequence[str | Element]): 要判断的子链. - def get(self, type_: type[TE], count: int | None = None) -> MessageChain[TE]: - """获取指定类型的消息段 + Returns: + List[int]: 所有找到的下标. + """ + + def unzip(seq: Sequence[str | Element]) -> list[str | Element]: + res: list[str | Element] = [] + for e in seq: + if isinstance(e, Text): + res.extend(e.text) + elif isinstance(e, str): + res.extend(e) + else: + res.append(e) + return res + + pattern: list[str | Element] = unzip(sub.content) if isinstance(sub, MessageChain) else unzip(sub) + + match_target: list[str | Element] = unzip(self.content) + + if len(match_target) < len(pattern): + return [] + + fallback: list[int] = [0 for _ in pattern] + current_fb: int = 0 # current fallback index + for i in range(1, len(pattern)): + while current_fb and pattern[i] != pattern[current_fb]: + current_fb = fallback[current_fb - 1] + if pattern[i] == pattern[current_fb]: + current_fb += 1 + fallback[i] = current_fb + + match_index: list[int] = [] + ptr = 0 + for i, e in enumerate(match_target): + while ptr and e != pattern[ptr]: + ptr = fallback[ptr - 1] + if e == pattern[ptr]: + ptr += 1 + if ptr == len(pattern): + match_index.append(i - ptr + 1) + ptr = fallback[ptr - 1] + return match_index + + def get(self, element_class: type[TE1], count: int | None = None) -> MessageChain[TE1]: + """ + 获取消息链中所有特定类型的消息元素 Args: - type_: 消息段类型 - count: 获取个数 + element_class (type[E]): 指定的消息元素的类型, 例如 "Text", "At", "Image" 等. + count (int, optional): 至多获取的元素个数 Returns: - 构建的新消息 + MessageChain[E]: 获取到的符合要求的所有消息元素; 另: 可能是空列表([]). """ if count is None: - return self[type_] + return self[element_class] - iterator, filtered = (elem for elem in self if isinstance(elem, type_)), MessageChain() - for _ in range(count): - elem = next(iterator, None) - if elem is None: - break - filtered.append(elem) - return filtered # type: ignore + return MessageChain(elem for elem in self.content if isinstance(elem, element_class))[:count] # type: ignore + + def get_one(self, element_class: type[TE1], index: int) -> TE1: + """获取消息链中第 index + 1 个特定类型的消息元素 + Args: + element_class (type[Element]): 指定的消息元素的类型, 例如 "Text", "At", "Image" 等. + index (int): 索引, 从 0 开始数 + Returns: + T: 消息链第 index + 1 个特定类型的消息元素 + """ + return self.get(element_class)[index] + + def get_first(self, element_class: type[TE1]) -> TE1: + """获取消息链中第 1 个特定类型的消息元素 + Args: + element_class (type[Element]): 指定的消息元素的类型, 例如 "Text", "At", "Image" 等. + Returns: + T: 消息链第 1 个特定类型的消息元素 + """ + return self.get(element_class)[0] + + def join(self, *chains: Self | Iterable[Self]) -> Self: + """将多个消息链连接起来, 并在其中插入自身. + + Args: + *chains (Iterable[MessageChain]): 要连接的消息链. + + Returns: + MessageChain: 连接后的消息链, 已对文本进行合并. + """ + result: list[TE] = [] + list_chains: list[MessageChain] = [] + for chain in chains: + if isinstance(chain, MessageChain): + list_chains.append(chain) + else: + list_chains.extend(chain) + + for chain in list_chains: + if chain is not list_chains[0]: + result.extend(deepcopy(self.content)) + result.extend(deepcopy(chain.content)) + return self.__class__(result).merge() def count(self, value: type[Element] | str | Element) -> int: - """计算指定消息段的个数 + """计算指定消息元素的个数 Args: - value: 消息段或消息段类型 + value (str | Element | type[Element]): 消息元素或消息元素类型 Returns: - 个数 + int: 消息元素的个数 """ if isinstance(value, str): value = Text(value) return ( len(self[value]) # type: ignore if isinstance(value, type) - else super().count(value) # type: ignore + else self.content.count(value) # type: ignore ) def only(self, value: type[Element] | str | Element) -> bool: - """检查消息中是否仅包含指定消息段 + """检查消息中是否仅包含指定消息元素 Args: - value: 指定消息段或消息段类型 + value: 指定消息元素或消息元素类型 Returns: - 是否仅包含指定消息段 + bool: 是否仅包含指定消息元素 """ if isinstance(value, type): - return all(isinstance(elem, value) for elem in self) + return all(isinstance(elem, value) for elem in self.content) if isinstance(value, str): value = Text(value) - return all(elem == value for elem in self) - - def join(self, iterable: Iterable[TE1 | MessageChain[TE1]]) -> MessageChain[TE | TE1]: - """将多个消息连接并将自身作为分割 - - Args: - iterable: 要连接的消息 - - Returns: - 连接后的消息 - """ - ret = MessageChain() - for index, msg in enumerate(iterable): - if index != 0: - ret.extend(self) - if isinstance(msg, Element): - ret.append(msg) - else: - ret.extend(msg.copy()) - return ret # type: ignore + return all(elem == value for elem in self.content) - def copy(self) -> MessageChain[TE]: + def copy(self) -> Self: """深拷贝消息""" return deepcopy(self) - def fork(self) -> MessageChain[TE]: + def fork(self) -> Self: """浅拷贝消息""" new = self.__class__() - list.extend(new, self) + new.content = self.content[:] return new - def include(self, *types: type[Element]) -> MessageChain: - """过滤消息 - + def exclude(self, *types: type[Element]) -> Self: + """将除了在给出的消息元素类型中符合的消息元素重新包装为一个新的消息链 Args: - types: 包含的消息段类型 - + *types (type[Element]): 将排除在外的消息元素类型 Returns: - 新构造的消息 + MessageChain: 返回的消息链中不包含参数中给出的消息元素类型 """ - return MessageChain(elem for elem in self if elem.__class__ in types) - - def exclude(self, *types: type[Element]) -> MessageChain: - """过滤消息 + return self.__class__([i for i in self.content if not isinstance(i, types)]) + def include(self, *types: type[Element]) -> Self: + """将只在给出的消息元素类型中符合的消息元素重新包装为一个新的消息链 Args: - types: 不包含的消息段类型 - + *types (type[Element]): 将只包含在内的消息元素类型 Returns: - 新构造的消息 + MessageChain: 返回的消息链中只包含参数中给出的消息元素类型 """ - return MessageChain(elem for elem in self if elem.__class__ not in types) + return self.__class__([i for i in self.content if isinstance(i, types)]) def extract_plain_text(self) -> str: """提取消息内纯文本消息""" @@ -363,7 +481,13 @@ def filter(self, predicate: Callable[[TE], bool]) -> MessageChain[TE]: Args: predicate: 过滤函数 """ - return MessageChain(elem for elem in self if predicate(elem)) + return MessageChain(elem for elem in self.content if predicate(elem)) + + def __iter__(self) -> Iterator[Element]: + yield from self.content + + def __len__(self) -> int: + return len(self.content) @overload def map(self, func: Callable[[TE], TE1]) -> MessageChain[TE1]: ... @@ -374,7 +498,7 @@ def map(self, func: Callable[[TE], T]) -> list[T]: ... def map(self, func: Callable[[TE], TE1] | Callable[[TE], T]) -> MessageChain[TE1] | list[T]: result1 = [] result2 = [] - for elem in self: + for elem in self.content: result = func(elem) if isinstance(result, Element): result1.append(result) @@ -418,7 +542,7 @@ def transform(self, rules: SyncVisitor[S], session: S = None) -> MessageChain: 转换后的消息 """ output = MessageChain() - for elem in self: + for elem in self.content: result = self._visit_sync(elem, rules, session) if result is True: children = MessageChain(elem.children) @@ -428,7 +552,7 @@ def transform(self, rules: SyncVisitor[S], session: S = None) -> MessageChain: if isinstance(result, str | Element): output += result else: - output.extend(result) + output.content.extend(result) return output async def transform_async(self, rules: AsyncVisitor[S], session: S = None) -> MessageChain: @@ -442,7 +566,7 @@ async def transform_async(self, rules: AsyncVisitor[S], session: S = None) -> Me 转换后的消息 """ output = MessageChain() - for elem in self: + for elem in self.content: result = await self._visit_async(elem, rules, session) if result is True: children = MessageChain(elem.children) @@ -467,7 +591,7 @@ def split(self, pattern: str = " ") -> list[Self]: result: list[Self] = [] tmp = [] - for seg in self: + for seg in self.content: if isinstance(seg, Text): split_result = seg.text.split(pattern) for index, split_text in enumerate(split_result): @@ -498,7 +622,7 @@ def replace( UniMessage: 修改后的消息链, 若未替换则原样返回. """ result_list: list[TE] = [] - for seg in self: + for seg in self.content: if isinstance(seg, Text): result_list.append(seg.__class__(seg.text.replace(old, new))) else: @@ -515,9 +639,9 @@ def startswith(self, string: str) -> bool: bool: 是否以给出的字符串开头 """ - if not self or not isinstance(self[0], Text): + if not self.content or not isinstance(self.content[0], Text): return False - return list.__getitem__(self, 0).text.startswith(string) + return self.content[0].text.startswith(string) def endswith(self, string: str) -> bool: """判断消息链是否以给出的字符串结尾 @@ -529,102 +653,245 @@ def endswith(self, string: str) -> bool: bool: 是否以给出的字符串结尾 """ - if not self or not isinstance(self[-1], Text): + if not self.content or not isinstance(self.content[-1], Text): return False - return list.__getitem__(self, -1).text.endswith(string) + return self.content[-1].text.endswith(string) + + def append(self, element: Element | str) -> None: + """ + 向消息链最后追加单个元素 + + Args: + element (Element): 要添加的元素 + + Returns: + None + """ + if isinstance(element, str): + element = Text(element) + self.content.append(element) # type: ignore + + def insert(self, index: int, value: Element | str, /) -> None: + if isinstance(value, str): + value = Text(value) + self.content.insert(index, value) # type: ignore + + def extend( + self, + values: Iterable[Self | Element | list[Element | str]], + ) -> None: + """ + 向消息链最后添加元素/元素列表/消息链 + + Args: + *values (MessageChain | Element | list[Element | str]): 要添加的元素/元素容器. + + Returns: + MessageChain: copy = True 时返回副本, 否则返回自己的引用. + """ + result = [] + for i in values: + if isinstance(i, Element): + result.append(i) + elif isinstance(i, str): + result.append(Text(i)) + elif isinstance(i, MessageChain): + result.extend(i.content) + else: + for e in i: + if isinstance(e, str): + result.append(Text(e)) + else: + result.append(e) + self.content.extend(result) + + def empty(self) -> bool: + """ + 判断消息链是否为空,包括判断是否仅包含空字符串。 + + Returns: + bool: 判断结果。 + """ - def removeprefix(self, prefix: str) -> Self: + return not bool(self.content and str(self)) + + def pop(self, index: int = -1, /) -> TE: + """移除并返回指定位置的元素,默认移除最后一个元素。 + + Args: + index (int, optional): 要移除的元素的索引,默认为 -1(最后一个元素)。 + + Returns: + TE: 被移除的元素。 + """ + return self.content.pop(index) # type: ignore + + def removeprefix(self, prefix: str, *, copy: bool = True) -> Self: """移除消息链前缀. Args: prefix (str): 要移除的前缀. + copy (bool, optional): 是否在副本上修改, 默认为 True. Returns: - UniMessage: 修改后的消息链. + MessageChain: 修改后的消息链, 若未移除则原样返回. """ - copy = list.copy(self) - if not copy: - return self.__class__(copy) - seg = copy[0] - if not isinstance(seg, Text): - return self.__class__(copy) - if seg.text.startswith(prefix): - seg = seg.__class__(seg.text[len(prefix) :]) - if not seg.text: - copy.pop(0) - else: - copy[0] = seg - return self.__class__(copy) + elements = deepcopy(self.content) if copy else self.content + if not elements: + return self.copy() if copy else self + elem = elements[0] + if not isinstance(elem, Text): + return self.copy() if copy else self + if elem.text.startswith(prefix): + elem.text = elem.text[len(prefix) :] + if not elem.text: + elements.pop(0) + if copy: + return self.__class__(elements) + self.content.clear() + self.content.extend(elements) + return self - def removesuffix(self, suffix: str) -> Self: + def removesuffix(self, suffix: str, *, copy: bool = True) -> Self: """移除消息链后缀. Args: suffix (str): 要移除的后缀. + copy (bool, optional): 是否在副本上修改, 默认为 True. Returns: - UniMessage: 修改后的消息链. + MessageChain: 修改后的消息链, 若未移除则原样返回. """ - copy = list.copy(self) - if not copy: - return self.__class__(copy) - seg = copy[-1] - if not isinstance(seg, Text): - return self.__class__(copy) - if seg.text.endswith(suffix): - seg = seg.__class__(seg.text[: -len(suffix)]) - if not seg.text: - copy.pop(-1) - else: - copy[-1] = seg - return self.__class__(copy) - - def strip(self, *segments: str | Element | type[Element]) -> Self: - return self.lstrip(*segments).rstrip(*segments) - - def lstrip(self, *segments: str | Element | type[Element]) -> Self: - types = [i for i in segments if not isinstance(i, str)] or [] - chars = "".join([i for i in segments if isinstance(i, str)]) or None - copy = list.copy(self) - if not copy: - return self.__class__(copy) - while copy: - seg = copy[0] - if seg in types or seg.__class__ in types: - copy.pop(0) - elif isinstance(seg, Text): - seg = seg.__class__(seg.text.lstrip(chars)) - if not seg.text: - copy.pop(0) + elements = deepcopy(self.content) if copy else self.content + if not elements: + return self.copy() if copy else self + elem = elements[-1] + if not isinstance(elem, Text): + return self.copy() if copy else self + if elem.text.endswith(suffix): + elem.text = elem.text[: -len(suffix)] + if not elem.text: + elements.pop(-1) + if copy: + return self.__class__(elements) + self.content.clear() + self.content.extend(elements) + return self + + def strip(self, *elements: str | type[Element] | Element, copy: bool = True) -> Self: + return self.lstrip(*elements, copy=copy).rstrip(*elements, copy=copy) + + def lstrip(self, *elements: str | type[Element] | Element, copy: bool = True) -> Self: + types = [i for i in elements if not isinstance(i, str)] or [] + chars = "".join([i for i in elements if isinstance(i, str)]) or None + content = deepcopy(self.content) if copy else self.content + if not content: + return self.copy() if copy else self + while content: + elem = content[0] + if elem in types or elem.__class__ in types: + content.pop(0) + elif isinstance(elem, Text): + text = elem.text.lstrip(chars) + if not text: + content.pop(0) continue - else: - copy[0] = seg + elem.text = text break else: break - return self.__class__(copy) - - def rstrip(self, *segments: str | Element | type[Element]) -> Self: - types = [i for i in segments if not isinstance(i, str)] or [] - chars = "".join([i for i in segments if isinstance(i, str)]) or None - copy = list.copy(self) - if not copy: - return self.__class__(copy) - while copy: - seg = copy[-1] - if seg in types or seg.__class__ in types: - copy.pop(-1) - elif isinstance(seg, Text): - seg = seg.__class__(seg.text.rstrip(chars)) - if not seg.text: - copy.pop(-1) + if copy: + return self.__class__(content) + self.content.clear() + self.content.extend(content) + return self + + def rstrip(self, *elements: str | type[Element] | Element, copy: bool = True) -> Self: + types = [i for i in elements if not isinstance(i, str)] or [] + chars = "".join([i for i in elements if isinstance(i, str)]) or None + content = deepcopy(self.content) if copy else self.content + if not content: + return self.copy() if copy else self + while content: + elem = content[-1] + if elem in types or elem.__class__ in types: + content.pop(-1) + elif isinstance(elem, Text): + text = elem.text.rstrip(chars) + if not text: + content.pop(-1) continue - else: - copy[-1] = seg + elem.text = text break else: break - return self.__class__(copy) + if copy: + return self.__class__(content) + self.content.clear() + self.content.extend(content) + return self + + def replace_chain( + self, + old: MessageChain | list[Element], + new: MessageChain | list[Element], + ) -> Self: + """替换消息链中的一部分. (在副本上操作) + + Args: + old (MessageChain): 要替换的消息链. + new (MessageChain): 替换后的消息链. + + Returns: + MessageChain: 修改后的消息链, 若未替换则原样返回. + """ + if not isinstance(old, MessageChain): + old = MessageChain(old) + if not isinstance(new, MessageChain): + new = MessageChain(new) + index_list: list[int] = self.index_sub(old) + + def unzip(chain: MessageChain) -> list[str | Element]: + unzipped: list[str | Element] = [] + for e in chain.content: + if isinstance(e, Text): + unzipped.extend(e.text) + else: + unzipped.append(e) + return unzipped + + unzipped_new: list[str | Element] = unzip(new) + unzipped_old: list[str | Element] = unzip(old) + unzipped_self: list[str | Element] = unzip(self) + unzipped_result: list[str | Element] = [] + last_end: int = 0 + for start in index_list: + unzipped_result.extend(unzipped_self[last_end:start]) + last_end = start + len(unzipped_old) + unzipped_result.extend(unzipped_new) + unzipped_result.extend(unzipped_self[last_end:]) + + # Merge result + result_list: list[TE] = [] + char_stk: list[str] = [] + for v in unzipped_result: + if isinstance(v, str): + char_stk.append(v) + else: + result_list.append(Text("".join(char_stk))) # type: ignore + char_stk = [] + result_list.append(v) # type: ignore + if char_stk: + result_list.append(Text("".join(char_stk))) # type: ignore + return self.__class__(result_list) + + def __bool__(self): + return bool(self.content and str(self)) + + def __eq__(self, value: object, /): + if not isinstance(value, MessageChain): + return False + return value.content == self.content def display(self): texts = [] diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 2e5b8eb..e7ce798 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -5,7 +5,6 @@ import tokenize from collections.abc import Sequence from importlib import _bootstrap, _bootstrap_external # type: ignore -from importlib.abc import MetaPathFinder from importlib.machinery import ExtensionFileLoader, ModuleSpec, PathFinder, SourceFileLoader from importlib.metadata import Distribution, PackageNotFoundError, distribution, distributions from importlib.util import module_from_spec, resolve_name @@ -204,9 +203,8 @@ def get_code(self, fullname): """ source_path = self.get_filename(fullname) - source_bytes = None - if source_bytes is None: - source_bytes = self.get_data(source_path) + # --- SourceFileLoader's cache handler removed --- + source_bytes = self.get_data(source_path) code_object = self.source_to_code(source_bytes, source_path) return code_object @@ -252,11 +250,7 @@ def create_module(self, spec) -> ModuleType | None: if self.name in plugin_service._subplugined: self.loaded = True return plugin_service.plugins[plugin_service._subplugined[self.name]].subproxy(self.name) - if ( - any((k.startswith(self.name) and k.rfind("@") != -1) for k in plugin_service.plugins) - and self.plugin_id.rfind("@") == -1 - ): - raise ReusablePluginError(f"reusable plugin {self.name!r} cannot be imported directly") + _check_reusable(self.name, self.plugin_id) return super().create_module(spec) def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) -> None: @@ -405,20 +399,39 @@ def _path_find_spec(fullname, path=None, target=None) -> ModuleSpec | None: return spec -class _PluginFinder(MetaPathFinder): +def _as_plugin(module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str) -> ModuleSpec: + module_spec.loader = PluginLoader(fullname, module_origin, plugin_id) + return module_spec + + +def _as_submodule( + module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str, parent: str +) -> ModuleSpec: + module_spec.loader = PluginLoader(fullname, module_origin, plugin_id, parent) + return module_spec + + +def _check_reusable(name: str, plugin_id: str) -> None: + if any(k.startswith(name) and k.rfind("@") != -1 for k in plugin_service.plugins) and plugin_id.rfind("@") == -1: + raise ReusablePluginError(f"reusable plugin {name!r} cannot be imported directly") + + +class _PluginFinder(PathFinder): @classmethod def find_spec( cls, fullname: str, - path: Sequence[str] | None, + path: Sequence[str] | None = None, target: ModuleType | None = None, origin_id_: str | None = None, - ): + force: bool = False, + ) -> ModuleSpec | None: # get the module spec using the default path-finder module_spec = _path_find_spec(fullname, path, target) if not module_spec: return module_origin = module_spec.origin + plugin_id = origin_id_ or fullname # if the module has no origin, it might be a namespace package or a built-in module. # We only care about namespace packages here, as built-in modules should not be treated as plugins. # For namespace packages, we can still return the spec without modification, @@ -435,28 +448,25 @@ def find_spec( if plug := current_plugin.get(None): # if the module being imported is the same as the plugin's module, # return the plugin's module spec directly to avoid infinite recursion. - if plug.module.__spec__ and plug.module.__spec__.origin == module_spec.origin: + if plug.module.__spec__ and plug.module.__spec__.origin == module_origin: return plug.module.__spec__ # get the top-level plugin id (the parent) of the current plugin - plugin_id = plug.id - while plugin_id in plugin_service._subplugined: - plugin_id = plugin_service._subplugined[plugin_id] + parent_id = plug.id + while parent_id in plugin_service._subplugined: + parent_id = plugin_service._subplugined[parent_id] # if the module being imported is a submodule of the top-level plugin, - if module_spec.name.startswith(plugin_service.plugins[plugin_id].module.__name__ + "."): - module_spec.loader = PluginLoader(fullname, module_origin, origin_id_ or fullname, plugin_id) - return module_spec - # if the module being imported is in the waitlist of the top-level plugin, + # or if the module being imported is in the waitlist of the top-level plugin, # it means it is marked as a submodule by the plugin author. - if module_spec.name in _SUBMODULE_WAITLIST.get(plugin_id, ()): - module_spec.loader = PluginLoader(fullname, module_origin, origin_id_ or fullname, plugin_id) - # plugin_service.referents.setdefault(module_spec.name, set()).add(plug.id) - # _SUBMODULE_WAITLIST[plug.module.__name__].remove(module_spec.name) - return module_spec + if module_spec.name.startswith( + plugin_service.plugins[parent_id].module.__name__ + "." + ) or module_spec.name in _SUBMODULE_WAITLIST.get( # noqa: E501 + parent_id, () + ): + return _as_submodule(module_spec, fullname, module_origin, plugin_id, parent_id) # in the following cases, the module is imported directly (probably from Entari App) # 1. the module is already a plugin. if module_spec.name in plugin_service.plugins: - module_spec.loader = PluginLoader(fullname, module_origin, origin_id_ or fullname) - return module_spec + return _as_plugin(module_spec, fullname, module_origin, plugin_id) # 2. the module is marked as a plugin by the plugin author, or followed the naming convention for plugins. marked = ( module_spec.name in _ENSURE_IS_PLUGIN @@ -487,7 +497,7 @@ def find_spec( except (KeyError, ValueError): pass if marked: - module_spec.loader = PluginLoader(fullname, module_origin, origin_id_ or fullname) + _as_plugin(module_spec, fullname, module_origin, plugin_id) # if there already exists a plugin that is importing this module, # we should add the plugin as a referent of this module if plug: @@ -495,33 +505,23 @@ def find_spec( return module_spec # 3. the module is marked as a submodule by other plugin, or it is a submodule of a plugin. if module_spec.name in plugin_service._subplugined: - module_spec.loader = PluginLoader( - fullname, module_origin, origin_id_ or fullname, plugin_service._subplugined[module_spec.name] + return _as_submodule( # noqa: E501 + module_spec, fullname, module_origin, plugin_id, plugin_service._subplugined[module_spec.name] ) - return module_spec # 4. if the module is already a plugin, but it is assigned an unique id (usage of reusable plugin), # it cannot be imported directly, otherwise it will break the uniqueness of the plugin instance. - if ( - any(k.startswith(module_spec.name) and k.rfind("@") != -1 for k in plugin_service.plugins) - and (origin_id_ or fullname).rfind("@") == -1 - ): - raise ReusablePluginError(f"reusable plugin {module_spec.name!r} cannot be imported directly") + _check_reusable(module_spec.name, plugin_id) # 5. the module is a submodule of a plugin, but it is not marked as a submodule by the plugin author, # we should still treat it as a submodule of the plugin to avoid breaking existing plugins if module_spec.parent and module_spec.parent in plugin_service.plugins: - module_spec.loader = PluginLoader(fullname, module_origin, origin_id_ or fullname, module_spec.parent) - return module_spec - # 6. the module is a submodule of a plugin, but it is not marked as a submodule by the plugin author, - # we should still treat it as a submodule of the plugin to avoid breaking existing plugins - if module_spec.name.rpartition(".")[0] in plugin_service.plugins: - module_spec.loader = PluginLoader( - fullname, module_origin, origin_id_ or fullname, module_spec.name.rpartition(".")[0] - ) - return module_spec + return _as_submodule(module_spec, fullname, module_origin, plugin_id, module_spec.parent) + # 6. force-wrap as a plugin when explicitly requested by import_plugin. + if force: + return _as_plugin(module_spec, fullname, module_origin, plugin_id) return -def find_spec(id_, package=None) -> ModuleSpec | None: +def import_plugin(id_, package=None, config: dict | None = None): uid_index = id_.rfind("@") name = id_ if uid_index == -1 else id_[:uid_index] fullname = resolve_name(name, package) if name.startswith(".") else name @@ -542,20 +542,13 @@ def find_spec(id_, package=None) -> ModuleSpec | None: if _current in plugin_service.plugins: parent = plugin_service.plugins[_current].module enter_plugin = True - _current += "." - continue - if _current in _ENSURE_IS_PLUGIN: - parent = import_plugin(_current) - if parent: + elif _current in _ENSURE_IS_PLUGIN or enter_plugin: + if parent := import_plugin(_current): enter_plugin = True else: parent = __import__(_current, fromlist=["__path__"]) - _current += "." - continue - if enter_plugin and (parent := import_plugin(_current)): - pass + enter_plugin = False else: - enter_plugin = False parent = __import__(_current, fromlist=["__path__"]) _current += "." if parent is None: @@ -566,46 +559,29 @@ def find_spec(id_, package=None) -> ModuleSpec | None: parent_path = parent.__path__ else: parent_path = None - if isinstance(parent_path, _bootstrap_external._NamespacePath): # type: ignore - parent_path = _NamespacePath(parent_path._name, parent_path._path, PathFinder._get_spec) # type: ignore - if spec := _PluginFinder.find_spec(fullname, parent_path, origin_id_=id_): - return spec - module_spec = _path_find_spec(fullname, parent_path, None) - if not module_spec: + spec = _PluginFinder.find_spec(fullname, parent_path, origin_id_=id_, force=True) + if not spec: return - module_origin = module_spec.origin - if not module_origin: - return - if isinstance(module_spec.loader, ExtensionFileLoader): - return - module_spec.loader = PluginLoader(fullname, module_origin, id_) - return module_spec - - -def import_plugin(id_, package=None, config: dict | None = None): - spec = find_spec(id_, package) - if spec: - mod = module_from_spec(spec) - if spec.loader: - if isinstance(spec.loader, PluginLoader): - spec.loader.exec_module(mod, config=config) - protected_modules = set() - module_name = mod.__name__ - if module_name: - prefix = [] - for part in module_name.split("."): - prefix.append(part) - protected_modules.add(".".join(prefix)) - sys.modules.pop(module_name, None) - for _imported in _IMPORTING: - if _imported in protected_modules or _imported in plugin_service.plugins: - continue - sys.modules.pop(_imported, None) - _IMPORTING.clear() - else: - spec.loader.exec_module(mod) - return mod - return + mod = module_from_spec(spec) + if spec.loader: + if isinstance(spec.loader, PluginLoader): + spec.loader.exec_module(mod, config=config) + protected_modules = set() + module_name = mod.__name__ + if module_name: + prefix = [] + for part in module_name.split("."): + prefix.append(part) + protected_modules.add(".".join(prefix)) + sys.modules.pop(module_name, None) + for _imported in _IMPORTING: + if _imported in protected_modules or _imported in plugin_service.plugins: + continue + sys.modules.pop(_imported, None) + _IMPORTING.clear() + else: + spec.loader.exec_module(mod) + return mod sys.meta_path.insert(0, _PluginFinder()) diff --git a/arclet/entari/session.py b/arclet/entari/session.py index 854c2b7..45d52a0 100644 --- a/arclet/entari/session.py +++ b/arclet/entari/session.py @@ -1,6 +1,9 @@ import asyncio +import inspect import secrets from collections.abc import Awaitable, Callable, Iterable +from functools import wraps +from types import MethodType from typing import Any, Generic, NoReturn, cast, overload from typing_extensions import TypeVar @@ -26,6 +29,7 @@ from . import command from .config import EntariConfig +from .event.api import APIRequest, APIResponse, SendRequest, SendResponse from .event.base import ( FriendRequestEvent, GuildMemberRequestEvent, @@ -35,7 +39,6 @@ Reply, SatoriEvent, ) -from .event.send import SendRequest, SendResponse from .message import MessageChain, Render TEvent = TypeVar("TEvent", bound=SatoriEvent, default=SatoriEvent) @@ -55,14 +58,51 @@ async def rule(elem: Element, sess: "Session"): return await content.transform_async(rule, session) +STATIC_METHODS = frozenset( + { + "__init__", + "call_api", + "request_internal", + "send", + "send_message", + "send_private_message", + "update_message", + "message_create", + } +) + + class EntariProtocol(ApiProtocol): # fmt: off + def __init__(self, account: Account["EntariProtocol"]): + super().__init__(account) + funcs = inspect.getmembers(self, predicate=lambda x: isinstance(x, MethodType)) + for name, func in funcs: + if name in STATIC_METHODS: + continue + + @wraps(func) + async def wrapper(*args, _func=func, _sig=inspect.signature(func), **kwargs): + bounds = _sig.bind(*args, **kwargs) + bounds.apply_defaults() + try: + if result := await es.post(APIRequest(self.account, _func.__name__, bounds.arguments)): + ans = result.value + else: + ans = await _func(**bounds.arguments) + success = True + except Exception as e: + ans = e + success = False + await es.publish(APIResponse(self.account, _func.__name__, bounds.arguments, success, ans)) + return ans + setattr(self, name, wrapper) + async def send_message(self, channel: str | Channel, message: str | Iterable[str | Element], at_sender: At | None = None, reply_to: Quote | None = None, referrer: dict[str, Any] | None = None) -> list[MessageObject]: # noqa: E501 """发送消息。返回一个 `MessageReceipt` 对象构成的数组。 - Args: - channel (str | Channel): 要发送的频道 ID + Args: channel (str | Channel): 要发送的频道 ID message (str | Iterable[str | Element]): 要发送的消息 at_sender (At | None): 是否 @ 发送者,默认为 None reply_to (Quote | None): 是否作为回复发送,默认为 None @@ -123,7 +163,7 @@ async def message_create(self, channel_id: str, content: str | Iterable[str | El msg = await component_transform(sess, msg) referrer = {k: v for k, v in referrer.items() if k != "source"} sess.elements = msg - btns = select(msg, Button) + btns = select(msg.content, Button) for btn in btns: if btn.type != "link" and not btn.id: btn.id = secrets.token_urlsafe(16) diff --git a/example_plugins/example_plugin8.py b/example_plugins/example_plugin8.py new file mode 100644 index 0000000..235022d --- /dev/null +++ b/example_plugins/example_plugin8.py @@ -0,0 +1,39 @@ +from arclet.entari.filter.message import startswith, regexmatch, regex_origin +from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image, Text + + +@listen(MessageCreatedEvent) +@startswith("!hello") +async def hello_listener1(sess: Session, message: MessageChain): + await sess.send("Hello! This is a response from the hello_listener.") + await sess.send(message) + + +@listen(MessageCreatedEvent) +@startswith(Image, include=True) +async def image_listener(sess: Session, message: MessageChain): + await sess.send("Hello! This is a response from the image_listener.") + await sess.send(message) + + +@listen(MessageCreatedEvent) +@startswith("!world", bind="world") +async def hello_listener2(sess: Session, message: MessageChain, world: MessageChain): + await sess.send("Hello! This is a response from the hello_listener2.") + await sess.send(message) + await sess.send(world) + + +@listen(MessageCreatedEvent) +@regexmatch(r"test (\d+)", flags=2) +async def regex_listener( + sess: Session, + message: MessageChain, + match = regex_origin(), + group1: str = regex_origin().group(1), + dicts: dict = regex_origin().groupdict(), +): + await sess.send(f"Hello! This is a response from the regex_listener. You said: {message}") + await sess.send(f"Matched: {Text(str(match))}") + await sess.send(f"Matched group 1: {group1}") + await sess.send(f"Matched dicts: {dicts}")