From a6d37fb0aa82adca76b3e54b912583ea616624e9 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 24 Jul 2026 11:57:26 +0800 Subject: [PATCH 01/33] :sparkles: APIRequest/Response event for call_api_hook --- arclet/entari/__init__.py | 4 +- arclet/entari/core.py | 2 +- arclet/entari/event/api.py | 107 ++++++++++++++++++++++++++++++++++++ arclet/entari/event/send.py | 81 +++------------------------ arclet/entari/session.py | 46 +++++++++++++++- 5 files changed, 161 insertions(+), 79 deletions(-) create mode 100644 arclet/entari/event/api.py 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/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/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/session.py b/arclet/entari/session.py index 854c2b7..21a686c 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 From ecb380d3356aa5e96fd9065bbed1fb607b1e1924 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 31 Jul 2026 13:56:09 +0800 Subject: [PATCH 02/33] :sparkles: add `startswith` `endswith` filter --- arclet/entari/command/plugin.py | 5 +- arclet/entari/filter/__init__.py | 124 +------------------------ arclet/entari/filter/limit.py | 76 ++++++++++++++++ arclet/entari/filter/message.py | 140 +++++++++++++++++++++++++++++ arclet/entari/filter/permission.py | 50 +++++++++++ example_plugins/example_plugin8.py | 24 +++++ 6 files changed, 297 insertions(+), 122 deletions(-) create mode 100644 arclet/entari/filter/limit.py create mode 100644 arclet/entari/filter/message.py create mode 100644 arclet/entari/filter/permission.py create mode 100644 example_plugins/example_plugin8.py diff --git a/arclet/entari/command/plugin.py b/arclet/entari/command/plugin.py index 28fd654..eba4c65 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) self.supplier = AlconnaSuppiler(command, self.cache, block, skip_for_unmatch) super().__init__(plugin, MessageCreatedEvent, command.path) plugin.collect( diff --git a/arclet/entari/filter/__init__.py b/arclet/entari/filter/__init__.py index dc683a7..9a6a24c 100644 --- a/arclet/entari/filter/__init__.py +++ b/arclet/entari/filter/__init__.py @@ -1,18 +1,16 @@ -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, semaphore as semaphore +from .permission import admins as admins, superusers as superusers +from .message import startswith as startswith, endswith as endswith _SessionFilter: TypeAlias = Callable[[Session], bool] | Callable[[Session], Awaitable[bool]] @@ -62,117 +60,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..57354c0 --- /dev/null +++ b/arclet/entari/filter/limit.py @@ -0,0 +1,76 @@ +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..1c787c4 --- /dev/null +++ b/arclet/entari/filter/message.py @@ -0,0 +1,140 @@ +import re +from typing import Any + +from arclet.letoderea import STOP, Propagator, propagate, Contexts, 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) + + +__all__ = ["startswith", "endswith"] 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/example_plugins/example_plugin8.py b/example_plugins/example_plugin8.py new file mode 100644 index 0000000..aefb9af --- /dev/null +++ b/example_plugins/example_plugin8.py @@ -0,0 +1,24 @@ +from arclet.entari.filter.message import startswith +from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image + + +@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) From 763eb0a6c9ae2ef2c8f6ca73cbcb3268807c59b5 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Mon, 3 Aug 2026 09:53:06 +0800 Subject: [PATCH 03/33] :sparkles: improve MessageChain impl --- arclet/entari/event/base.py | 4 +- arclet/entari/message.py | 635 +++++++++++++++++++++++++----------- arclet/entari/session.py | 2 +- 3 files changed, 454 insertions(+), 187 deletions(-) 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/message.py b/arclet/entari/message.py index 3b3c02e..0be22a4 100644 --- a/arclet/entari/message.py +++ b/arclet/entari/message.py @@ -1,9 +1,9 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Iterable, Sequence +from collections.abc import Awaitable, Callable, Iterable, Sequence, MutableSequence from copy import deepcopy from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar, Union, overload +from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar, Union, overload, Iterator from typing_extensions import Self, SupportsIndex from satori import select as satori_select @@ -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/session.py b/arclet/entari/session.py index 21a686c..45d52a0 100644 --- a/arclet/entari/session.py +++ b/arclet/entari/session.py @@ -163,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) From ed43b459559f26a4f6ab62da5e3b019befdfe8c1 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Mon, 3 Aug 2026 10:31:29 +0800 Subject: [PATCH 04/33] :sparkles: add `fullmatch` `regexmatch` filter --- arclet/entari/command/plugin.py | 2 +- arclet/entari/filter/__init__.py | 9 ++-- arclet/entari/filter/limit.py | 1 + arclet/entari/filter/message.py | 80 +++++++++++++++++++++++++++++- arclet/entari/message.py | 4 +- example_plugins/example_plugin8.py | 19 ++++++- 6 files changed, 105 insertions(+), 10 deletions(-) diff --git a/arclet/entari/command/plugin.py b/arclet/entari/command/plugin.py index eba4c65..70c427b 100644 --- a/arclet/entari/command/plugin.py +++ b/arclet/entari/command/plugin.py @@ -56,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[str, asyncio.Future]" = 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/filter/__init__.py b/arclet/entari/filter/__init__.py index 9a6a24c..039f493 100644 --- a/arclet/entari/filter/__init__.py +++ b/arclet/entari/filter/__init__.py @@ -8,9 +8,12 @@ from ..session import Session from . import common -from .limit import interval as interval, semaphore as semaphore -from .permission import admins as admins, superusers as superusers -from .message import startswith as startswith, endswith as endswith +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]] diff --git a/arclet/entari/filter/limit.py b/arclet/entari/filter/limit.py index 57354c0..5df1890 100644 --- a/arclet/entari/filter/limit.py +++ b/arclet/entari/filter/limit.py @@ -3,6 +3,7 @@ from arclet.letoderea import STOP, Propagator, propagate from arclet.letoderea.utils import TCallable + from ..message import MessageChain from ..session import Session diff --git a/arclet/entari/filter/message.py b/arclet/entari/filter/message.py index 1c787c4..c3f3cda 100644 --- a/arclet/entari/filter/message.py +++ b/arclet/entari/filter/message.py @@ -1,7 +1,7 @@ import re from typing import Any -from arclet.letoderea import STOP, Propagator, propagate, Contexts, provide +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 @@ -137,4 +137,80 @@ def __call__(self, func: TCallable) -> TCallable: return propagate(self)(func) -__all__ = ["startswith", "endswith"] +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/message.py b/arclet/entari/message.py index 0be22a4..7410a91 100644 --- a/arclet/entari/message.py +++ b/arclet/entari/message.py @@ -1,9 +1,9 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Iterable, Sequence, MutableSequence +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, Iterator +from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar, Union, overload from typing_extensions import Self, SupportsIndex from satori import select as satori_select diff --git a/example_plugins/example_plugin8.py b/example_plugins/example_plugin8.py index aefb9af..235022d 100644 --- a/example_plugins/example_plugin8.py +++ b/example_plugins/example_plugin8.py @@ -1,5 +1,5 @@ -from arclet.entari.filter.message import startswith -from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image +from arclet.entari.filter.message import startswith, regexmatch, regex_origin +from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image, Text @listen(MessageCreatedEvent) @@ -22,3 +22,18 @@ async def hello_listener2(sess: Session, message: MessageChain, world: MessageCh 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}") From ee27b19df6618eff6ba210b2c936da0193a6fa52 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 4 Aug 2026 10:26:01 +0800 Subject: [PATCH 05/33] :sparkles: improve plugin import code --- arclet/entari/plugin/module.py | 166 ++++++++++++++------------------- 1 file changed, 71 insertions(+), 95 deletions(-) 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()) From 08f567b75eb39557502fc3080121f1ea74272a96 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 7 Aug 2026 14:37:20 +0800 Subject: [PATCH 06/33] :sparkles: add plugin inspect for skip reload on semantic-identical file changes --- arclet/entari/builtins/auto_reload.py | 19 ++++++++++++++++++ arclet/entari/plugin/model.py | 12 ++++++++++-- arclet/entari/plugin/model.pyi | 8 ++++++++ arclet/entari/plugin/module.py | 28 ++++++++++++++++++--------- 4 files changed, 56 insertions(+), 11 deletions(-) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index 7f99b03..b40a31b 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -1,6 +1,8 @@ +import ast import asyncio from dataclasses import asdict from pathlib import Path +from traceback import format_exception_only from arclet.letoderea import post, publish from launart import Launart, Service, any_completed @@ -17,6 +19,7 @@ from arclet.entari.event.config import ConfigReload from arclet.entari.logger import log from arclet.entari.plugin import PluginRole, find_plugin, find_plugin_by_file, unload_plugin_async +from arclet.entari.utils import escape_tag # declare_static() loguru_logger.disable("watchfiles.main") @@ -92,6 +95,22 @@ async def watch(self): if plugin.is_static: logger.info(f"Plugin {plugin.id!r} is static, ignored.") continue + + if ( + plugin._inspect + and plugin.module.__file__ + and (path := Path(change[1]).resolve()) == Path(plugin.module.__file__).resolve() + ): + try: + nodes = ast.parse(path.read_bytes(), filename=path, type_comments=True) + except (OSError, SyntaxError) as e: + trace = escape_tag("".join(format_exception_only(e))) + logger.error(f"Change in {plugin.id!r} occurred exception, skipped:\n{trace}") + continue + else: + if ast.dump(nodes, include_attributes=False) == plugin._inspect.dump: + logger.debug(f"Change in {plugin.id!r} has no semantic difference, skipped.") + continue logger.info(f"Detected change in {plugin.id!r}, reloading...") pid = plugin.id _conf = plugin.config.copy() diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index c66e495..b3ef33a 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -1,5 +1,6 @@ from __future__ import annotations +import ast import asyncio import inspect import re @@ -216,6 +217,12 @@ def wrapper(func: TCallable) -> TCallable: return wrapper +@dataclass(slots=True) +class PluginInspect: + nodes: ast.AST + dump: str + + @dataclass class Plugin: id: str @@ -226,13 +233,13 @@ class Plugin: is_static: bool = False path: str = field(init=False) uid: str | None = None + _inspect: PluginInspect | None = field(default=None, repr=False) _metadata: PluginMetadata | None = None _is_disposed: bool = False _services: dict[str, Service] = field(init=False, default_factory=dict) _config_key: str = field(init=False) - # _scope: Scope = field(init=False) _extra: dict[str, Any] = field(default_factory=dict, init=False) # extra metadata for inspection - _apply: Callable[[Plugin], Any] | None = field(default=None, init=False) + _apply: Callable[[Plugin], Any] | None = field(default=None, init=False, repr=False) @property def reusable(self) -> bool: @@ -519,6 +526,7 @@ def dispose(self, *, is_cleanup: bool = False): self._scope.propagators.clear() del plugin_service.plugins[self.id] del self.module + del self._inspect return tasks def dispatch(self, event, name: str | None = None): diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index 8935786..f8cea1c 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -1,4 +1,6 @@ +import ast import asyncio + from enum import Enum from collections.abc import Callable from dataclasses import dataclass, field @@ -196,6 +198,11 @@ class PluginMetadata: def inject(*services: type[Service] | str | DependService) -> Callable[[TCallable], TCallable]: ... @overload def inject(*services: type[Service] | str | DependService, _is_global: Literal[True]) -> Check: ... +@dataclass(slots=True) +class PluginInspect: + nodes: ast.AST + dump: str + @dataclass class Plugin: id: str @@ -206,6 +213,7 @@ class Plugin: is_static: bool = ... path: str = ... uid: str | None = ... + _inspect: PluginInspect | None = ... _metadata: PluginMetadata | None = ... _is_disposed: bool = ... _services: dict[str, Service] = field(init=False, default_factory=dict) diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index e7ce798..15b43d9 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -9,6 +9,7 @@ from importlib.metadata import Distribution, PackageNotFoundError, distribution, distributions from importlib.util import module_from_spec, resolve_name from io import BytesIO +from os import PathLike from pathlib import Path from types import ModuleType from typing import Any @@ -21,7 +22,7 @@ from ..event.plugin import PluginLoadedFailed, PluginLoadedSuccess from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError from ..logger import log -from .model import Plugin, PluginMetadata, current_plugin +from .model import Plugin, PluginInspect, PluginMetadata, current_plugin from .service import plugin_service _SUBMODULE_WAITLIST: dict[str, set[str]] = {} @@ -76,10 +77,10 @@ def _ensure_plugin(names: list[str], sub: bool, pid: str, pname: str, prefix="") # fmt: off class _Visitor(ast.NodeVisitor): - def __init__(self, pid: str, pname: str, path: str, plg_lineno: list[int], sub_lineno: list[int], ns_lineno: list[int]): # noqa: E501 + def __init__(self, pid: str, pname: str, path: bytes | str | PathLike[str], plg_lineno: list[int], sub_lineno: list[int], ns_lineno: list[int]): # noqa: E501 self.pid = pid self.pname = pname - self.path = path + self.path = path.decode() if isinstance(path, bytes) else f"{Path(path)}" self.signed_plugin_lineno = plg_lineno self.signed_subplugin_lineno = sub_lineno self.signed_namespace_lineno = ns_lineno @@ -193,6 +194,7 @@ def __init__(self, fullname: str, path: str, plugin_id: str, parent_plugin_id: s self.loaded = False self.plugin_id = plugin_id self.parent_plugin_id = parent_plugin_id + self._inspect: PluginInspect = None # type: ignore super().__init__(fullname, path) def get_code(self, fullname): @@ -208,7 +210,7 @@ def get_code(self, fullname): code_object = self.source_to_code(source_bytes, source_path) return code_object - def source_to_code(self, data, path=""): + def source_to_code(self, data, path="", *, _optimize: int = -1): """Return the code object compiled from source. The 'data' argument can be any object type that compile() supports. @@ -234,13 +236,13 @@ def source_to_code(self, data, path=""): nodes = ast.parse(data, type_comments=True) except SyntaxError: return _bootstrap._call_with_frames_removed( # type: ignore - compile, data, path, "exec", dont_inherit=True, optimize=-1 + compile, data, path, "exec", dont_inherit=True, optimize=_optimize ) visitor = _Visitor(self.plugin_id, name, path, plg_lineno, sub_lineno, ns_lineno) visitor.visit(nodes) - + self._inspect = PluginInspect(nodes, ast.dump(nodes, include_attributes=False)) return _bootstrap._call_with_frames_removed( # type: ignore - compile, nodes, path, "exec", dont_inherit=True, optimize=-1 + compile, nodes, path, "exec", dont_inherit=True, optimize=_optimize ) def create_module(self, spec) -> ModuleType | None: @@ -299,7 +301,10 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) if not plugin.is_static: token1 = scope_ctx.set(plugin._scope) try: - super().exec_module(module) + code = self.get_code(module.__name__) + if code is None: + raise ImportError(f"cannot load module {module.__name__r} when get_code() returns None") + _bootstrap._call_with_frames_removed(exec, code, module.__dict__) # type: ignore except RegisterNotInPluginError as e: deleted = [] for frame in reversed(inspect.trace()): @@ -316,7 +321,10 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) _ensure_plugin(deleted[-1:], False, self.plugin_id, self.name) _ENSURE_IS_PLUGIN.update(deleted[:-1]) try: - super().exec_module(module) + code = self.get_code(module.__name__) + if code is None: + raise ImportError(f"cannot load module {module.__name__r} when get_code() returns None") + _bootstrap._call_with_frames_removed(exec, code, module.__dict__) # type: ignore except Exception as e1: if isinstance(e1, RegisterNotInPluginError): log.plugin.error(f"failed to load plugin {self.plugin_id!r}:\n{e1.msg}") @@ -350,6 +358,8 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) if metadata and not plugin.metadata: plugin.metadata = metadata plugin._apply = getattr(module, "__plugin_apply__", None) + plugin._inspect = self._inspect + del self._inspect if not is_sub: if plugin._apply: log.plugin.success(f"loaded plugin {self.plugin_id!r} partially applied") From 043fb958b4cd1e33eb10a15a7616b4b5f47278ef Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 7 Aug 2026 17:03:33 +0800 Subject: [PATCH 07/33] :sparkles: add Lock in plugin reload (serialize) --- arclet/entari/builtins/auto_reload.py | 87 +++++++++++++++++---------- arclet/entari/plugin/model.py | 2 - arclet/entari/plugin/model.pyi | 1 - 3 files changed, 54 insertions(+), 36 deletions(-) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index b40a31b..59d2964 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -18,7 +18,7 @@ from arclet.entari.config import BasicConfModel, EntariConfig, model_field from arclet.entari.event.config import ConfigReload from arclet.entari.logger import log -from arclet.entari.plugin import PluginRole, find_plugin, find_plugin_by_file, unload_plugin_async +from arclet.entari.plugin import Plugin, PluginRole, find_plugin, find_plugin_by_file, unload_plugin_async from arclet.entari.utils import escape_tag # declare_static() @@ -84,52 +84,72 @@ def stages(self) -> set[Phase]: def __init__(self, config: Config): self.config = config self.fail: dict[str, tuple[str, dict]] = {} + self._locks: dict[str, asyncio.Lock] = {} super().__init__() + def _lock_for(self, plugin_id: str) -> asyncio.Lock: + if plugin_id not in self._locks: + self._locks[plugin_id] = asyncio.Lock() + return self._locks[plugin_id] + async def watch(self): async for event in awatch( *self.config.watch_dirs, debounce=self.config.debounce, step=self.config.step, watch_filter=PythonFilter() ): + pending: dict[str, tuple[str, Plugin]] = {} + failed: list[str] = [] for change in event: if plugin := find_plugin_by_file(change[1]): if plugin.is_static: logger.info(f"Plugin {plugin.id!r} is static, ignored.") continue - - if ( - plugin._inspect - and plugin.module.__file__ - and (path := Path(change[1]).resolve()) == Path(plugin.module.__file__).resolve() - ): - try: - nodes = ast.parse(path.read_bytes(), filename=path, type_comments=True) - except (OSError, SyntaxError) as e: - trace = escape_tag("".join(format_exception_only(e))) - logger.error(f"Change in {plugin.id!r} occurred exception, skipped:\n{trace}") + pending.setdefault(plugin.id, (change[1], plugin)) + elif change[1] in self.fail: + failed.append(change[1]) + for pid, (file_path, plugin) in pending.items(): + if ( + plugin._inspect + and plugin.module.__file__ + and (path := Path(file_path).resolve()) == Path(plugin.module.__file__).resolve() + ): + try: + nodes = ast.parse(path.read_bytes(), filename=path, type_comments=True) + except (OSError, SyntaxError) as e: + trace = escape_tag("".join(format_exception_only(e))) + logger.error(f"Change in {pid!r} occurred exception, skipped:\n{trace}") + continue + else: + if ast.dump(nodes, include_attributes=False) == plugin._inspect.dump: + logger.debug(f"Change in {pid!r} has no semantic difference, skipped.") + self.fail.pop(file_path, None) continue - else: - if ast.dump(nodes, include_attributes=False) == plugin._inspect.dump: - logger.debug(f"Change in {plugin.id!r} has no semantic difference, skipped.") - continue - logger.info(f"Detected change in {plugin.id!r}, reloading...") - pid = plugin.id - _conf = plugin.config.copy() - del plugin + logger.info(f"Detected change in {pid!r}, reloading...") + _conf = plugin.config.copy() + del plugin + async with self._lock_for(pid): await unload_plugin_async(pid) if plugin := load_plugin(pid, _conf): logger.info(f"Reloaded {plugin.id!r}") del plugin + self.fail.pop(file_path, None) else: logger.error(f"Failed to reload {pid!r}") - self.fail[change[1]] = (pid, _conf) - elif change[1] in self.fail: - logger.info(f"Detected change in {change[1]!r} which failed to reload, retrying...") - if plugin := load_plugin(*self.fail[change[1]]): + self.fail[file_path] = (pid, _conf) + pending.clear() + for file_path in failed: + if file_path not in self.fail: + continue + pid, _conf = self.fail[file_path] + async with self._lock_for(pid): + if file_path not in self.fail: + continue + logger.info(f"Detected change in {file_path!r} which failed to reload, retrying...") + if plugin := load_plugin(pid, _conf): logger.info(f"Reloaded {plugin.id!r}") del plugin - del self.fail[change[1]] + del self.fail[file_path] else: - logger.error(f"Failed to reload {self.fail[change[1]][0]!r}") + logger.error(f"Failed to reload {pid!r}") async def watch_config(self): file = EntariConfig.instance.path.resolve() @@ -215,13 +235,14 @@ async def watch_config(self): _conf = plg.config.copy() async def _(): - await unload_plugin_async(pid) - if plg := load_plugin(plugin_name, new_conf): - logger.info(f"Reloaded {plg.id!r}") - del plg - else: - logger.error(f"Failed to reload {plugin_name!r}") - self.fail[plugin_file] = (pid, _conf) + async with self._lock_for(pid): + await unload_plugin_async(pid) + if plg := load_plugin(plugin_name, new_conf): + logger.info(f"Reloaded {plg.id!r}") + del plg + else: + logger.error(f"Failed to reload {plugin_name!r}") + self.fail[plugin_file] = (pid, _conf) await asyncio.shield(_()) else: diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index b3ef33a..692fbd3 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -470,8 +470,6 @@ def dispose(self, *, is_cleanup: bool = False): t.add_done_callback(tasks.discard) tasks.add(t) self._services.clear() - if self.module.__spec__ and self.module.__spec__.cached: - Path(self.module.__spec__.cached).unlink(missing_ok=True) sys.modules.pop(self.module.__name__, None) tasks.update(self.restore()) delattr(self.module, "__plugin__") diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index f8cea1c..f2f5517 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -1,6 +1,5 @@ import ast import asyncio - from enum import Enum from collections.abc import Callable from dataclasses import dataclass, field From 2ab144103d9616dd7cdee70f2aa7749a5140ff99 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 7 Aug 2026 18:38:56 +0800 Subject: [PATCH 08/33] :sparkles: support Hot-Swap function if only fn changes --- arclet/entari/builtins/auto_reload.py | 18 +- arclet/entari/plugin/model.py | 2 +- arclet/entari/plugin/model.pyi | 2 +- arclet/entari/plugin/swap.py | 389 ++++++++++++++++++++++++++ 4 files changed, 408 insertions(+), 3 deletions(-) create mode 100644 arclet/entari/plugin/swap.py diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index 59d2964..fa4cf80 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -19,6 +19,7 @@ from arclet.entari.event.config import ConfigReload from arclet.entari.logger import log from arclet.entari.plugin import Plugin, PluginRole, find_plugin, find_plugin_by_file, unload_plugin_async +from arclet.entari.plugin.swap import classify, swap_functions from arclet.entari.utils import escape_tag # declare_static() @@ -107,6 +108,7 @@ async def watch(self): elif change[1] in self.fail: failed.append(change[1]) for pid, (file_path, plugin) in pending.items(): + nodes: ast.Module | None = None if ( plugin._inspect and plugin.module.__file__ @@ -116,7 +118,7 @@ async def watch(self): nodes = ast.parse(path.read_bytes(), filename=path, type_comments=True) except (OSError, SyntaxError) as e: trace = escape_tag("".join(format_exception_only(e))) - logger.error(f"Change in {pid!r} occurred exception, skipped:\n{trace}") + logger.error(f"Change in {pid!r} occurred exception, skipped:\n{trace}") continue else: if ast.dump(nodes, include_attributes=False) == plugin._inspect.dump: @@ -124,6 +126,20 @@ async def watch(self): self.fail.pop(file_path, None) continue logger.info(f"Detected change in {pid!r}, reloading...") + if plugin._inspect and nodes: + changes = classify(plugin._inspect.nodes, nodes) + if changes is not None and swap_functions(plugin, nodes, changes): + if changes: + logger.info( + f"Hot swapped functions in {pid!r}: " + f"{', '.join(f'{change.qualname}' for change in changes)} " + f"successfully." + ) + else: + logger.debug(f"Change in {pid!r} has no function-level diff, skipped.") + self.fail.pop(file_path, None) + continue + logger.debug(f"Hot swap functions in {pid!r} failed, falling back to full reload.") _conf = plugin.config.copy() del plugin async with self._lock_for(pid): diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 692fbd3..046e8b6 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -219,7 +219,7 @@ def wrapper(func: TCallable) -> TCallable: @dataclass(slots=True) class PluginInspect: - nodes: ast.AST + nodes: ast.Module dump: str diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index f2f5517..5f79361 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -199,7 +199,7 @@ def inject(*services: type[Service] | str | DependService) -> Callable[[TCallabl def inject(*services: type[Service] | str | DependService, _is_global: Literal[True]) -> Check: ... @dataclass(slots=True) class PluginInspect: - nodes: ast.AST + nodes: ast.Module dump: str @dataclass diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py new file mode 100644 index 0000000..fc34ac3 --- /dev/null +++ b/arclet/entari/plugin/swap.py @@ -0,0 +1,389 @@ +import ast +import builtins +import copy +import inspect +import sys +import types +from collections.abc import Iterator, Sequence +from dataclasses import dataclass +from types import ModuleType +from typing import Any + +from ..logger import log +from .model import Plugin, PluginInspect + + +@dataclass +class FunctionChange: + qualname: str + ordinal: int + count: int + node: ast.FunctionDef | ast.AsyncFunctionDef + signature_changed: bool + + +def list_dump(nodes: Sequence[ast.AST]) -> str: + """Copy from ast.py line 167-170""" + return f"[{', '.join(ast.dump(x, include_attributes=False) for x in nodes)}]" + + +def iter_functions(nodes: ast.Module) -> Iterator[tuple[str, int, ast.FunctionDef | ast.AsyncFunctionDef]]: + """按 qualname 收集模块顶层函数与类方法(含嵌套类)""" + + def _walk(body: list[ast.stmt], prefix: str) -> Iterator[tuple[str, int, ast.FunctionDef | ast.AsyncFunctionDef]]: + seen: dict[str, int] = {} + for stmt in body: + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): + qualname = f"{prefix}{stmt.name}" + ordinal = seen.get(stmt.name, 0) + seen[stmt.name] = ordinal + 1 + yield qualname, ordinal, stmt + elif isinstance(stmt, ast.ClassDef): + yield from _walk(stmt.body, f"{prefix}{stmt.name}.") + + yield from _walk(nodes.body, "") + + +def _structurally_same(o: ast.stmt, n: ast.stmt) -> bool: + """结构比较:函数/方法体跳过(由 qualname 配对分析),其余须 dump 相等""" + if isinstance(o, (ast.FunctionDef, ast.AsyncFunctionDef)) and isinstance( + n, (ast.FunctionDef, ast.AsyncFunctionDef) + ): + return True + + if isinstance(o, ast.ClassDef) and isinstance(n, ast.ClassDef): + if list_dump(o.decorator_list) != list_dump(n.decorator_list): + return False + if list_dump(o.bases) != list_dump(n.bases): + return False + if list_dump(o.keywords) != list_dump(n.keywords): + return False + if len(o.body) != len(n.body): + return False + return all(_structurally_same(x, y) for x, y in zip(o.body, n.body)) + return ast.dump(o, include_attributes=False) == ast.dump(n, include_attributes=False) + + +def classify(old_nodes: ast.Module, new_nodes: ast.Module) -> list[FunctionChange] | None: + """比较新旧模块 AST,返回需就地 swap 的函数列表;None 表示需模块层全量重载 + + 函数配对:先按(名称组, 完整节点 dump 相同)锚定未变化函数(容忍位置交换/重排), + 剩余函数按组内顺序配对比较(容忍原地编辑);重命名(组键变化)与增删 + (组内数量不匹配)→ 全量。 + """ + if len(old_nodes.body) != len(new_nodes.body) or not all( + _structurally_same(o, n) for o, n in zip(old_nodes.body, new_nodes.body) + ): + return None + old_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} + new_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} + for qualname, ordinal, fn in iter_functions(old_nodes): + old_groups.setdefault(qualname, []).append((ordinal, fn)) + for qualname, ordinal, fn in iter_functions(new_nodes): + new_groups.setdefault(qualname, []).append((ordinal, fn)) + if set(old_groups.keys()) != set(new_groups.keys()): + return None + changes: list[FunctionChange] = [] + for qualname, old_list in old_groups.items(): + new_list = new_groups[qualname] + if len(old_list) != len(new_list): + return None + used_old: set[int] = set() + used_new: set[int] = set() + for i, (_, o_node) in enumerate(old_list): + for j, (_, n_node) in enumerate(new_list): + if ( + i not in used_old + and j not in used_new + and ast.dump(o_node, include_attributes=False) == ast.dump(n_node, include_attributes=False) + ): + used_old.add(i) + used_new.add(j) + break + rest_old = [old_list[i] for i in range(len(old_list)) if i not in used_old] + rest_new = [new_list[j] for j in range(len(new_list)) if j not in used_new] + for (o_ord, o_node), (n_ord, n_node) in zip(rest_old, rest_new): + if list_dump(n_node.decorator_list) != list_dump(o_node.decorator_list): + return None + signature_changed = n_node.__class__ is not o_node.__class__ or ast.dump( + n_node.args, include_attributes=False + ) != ast.dump(o_node.args, include_attributes=False) + # if signature_changed or list_dump(n_node.body) != list_dump(o_node.body): + changes.append(FunctionChange(qualname, o_ord, len(old_list), n_node, signature_changed)) + return changes + + +def _signature_names(node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> set[str]: + names = {arg.arg for arg in node.args.posonlyargs + node.args.args + node.args.kwonlyargs} + if node.args.vararg: + names.add(node.args.vararg.arg) + if node.args.kwarg: + names.add(node.args.kwarg.arg) + return names + + +def _comp_targets(gen: ast.comprehension) -> set[str]: + targets: set[str] = set() + stack: list[ast.AST] = [gen.target] + while stack: + node = stack.pop() + if isinstance(node, ast.Name): + targets.add(node.id) + elif isinstance(node, (ast.Tuple, ast.List)): + stack.extend(node.elts) + elif isinstance(node, ast.Starred): + stack.append(node.value) + return targets + + +def collect_global_names(fn: ast.FunctionDef | ast.AsyncFunctionDef) -> set[str]: + """收集新函数体中所有会在模块全局或 builtins 解析的名字(含嵌套函数) + + 规则:Name(Load) 在自身与所有外层函数作用域均未绑定 → 全局名。 + """ + used: set[str] = set() + stack: list[set[str]] = [_signature_names(fn)] + + def bound(name: str) -> bool: + return any(name in scope for scope in stack) + + def bind(name: str): + stack[-1].add(name) + + def visit_expr(expr: ast.expr): + if isinstance(expr, ast.Name): + if isinstance(expr.ctx, ast.Load) and not bound(expr.id): + used.add(expr.id) + elif isinstance(expr.ctx, ast.Store): + bind(expr.id) + elif isinstance(expr, ast.Lambda): + stack.append(_signature_names(expr)) + visit_expr(expr.body) + stack.pop() + elif isinstance(expr, (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp)): + comps = expr.generators + if comps: + visit_expr(comps[0].iter) + stack.append({name for gen in comps for name in _comp_targets(gen)}) + for gen in comps: + if gen is not comps[0]: + visit_expr(gen.iter) + for cond in gen.ifs: + visit_expr(cond) + if isinstance(expr, ast.DictComp): + visit_expr(expr.key) + visit_expr(expr.value) + else: + visit_expr(expr.elt) + stack.pop() + else: + for child in ast.iter_child_nodes(expr): + if isinstance(child, ast.expr): + visit_expr(child) + + def visit_stmt(stmt: ast.stmt): + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): + bind(stmt.name) + for dec in stmt.decorator_list: + visit_expr(dec) + for default in [*stmt.args.defaults, *[d for d in stmt.args.kw_defaults if d]]: + visit_expr(default) + for arg in stmt.args.posonlyargs + stmt.args.args + stmt.args.kwonlyargs: + if arg.annotation: + visit_expr(arg.annotation) + if stmt.args.vararg and stmt.args.vararg.annotation: + visit_expr(stmt.args.vararg.annotation) + if stmt.args.kwarg and stmt.args.kwarg.annotation: + visit_expr(stmt.args.kwarg.annotation) + stack.append(_signature_names(stmt)) + for body_stmt in stmt.body: + visit_stmt(body_stmt) + stack.pop() + elif isinstance(stmt, ast.ClassDef): + bind(stmt.name) + for dec in stmt.decorator_list: + visit_expr(dec) + for base in stmt.bases: + visit_expr(base) + for kw in stmt.keywords: + visit_expr(kw.value) + stack.append(set()) + for body_stmt in stmt.body: + visit_stmt(body_stmt) + stack.pop() + elif isinstance(stmt, ast.Assign): + for target in stmt.targets: + visit_expr(target) + visit_expr(stmt.value) + elif isinstance(stmt, ast.AnnAssign): + visit_expr(stmt.target) + visit_expr(stmt.annotation) + if stmt.value: + visit_expr(stmt.value) + elif isinstance(stmt, ast.AugAssign): + visit_expr(stmt.target) + visit_expr(stmt.value) + elif isinstance(stmt, (ast.For, ast.AsyncFor)): + visit_expr(stmt.target) + visit_expr(stmt.iter) + for body_stmt in [*stmt.body, *stmt.orelse]: + visit_stmt(body_stmt) + elif isinstance(stmt, (ast.With, ast.AsyncWith)): + for item in stmt.items: + visit_expr(item.context_expr) + if item.optional_vars: + visit_expr(item.optional_vars) + for body_stmt in stmt.body: + visit_stmt(body_stmt) + elif isinstance(stmt, ast.Try) or (sys.version_info >= (3, 11) and isinstance(stmt, ast.TryStar)): + for body_stmt in [*stmt.body, *stmt.orelse, *stmt.finalbody]: + visit_stmt(body_stmt) + for handler in stmt.handlers: + if handler.type: + visit_expr(handler.type) + if handler.name: + bind(handler.name) + for body_stmt in handler.body: + visit_stmt(body_stmt) + elif isinstance(stmt, ast.Import): + for alias in stmt.names: + bind(alias.asname or alias.name.split(".")[0]) + elif isinstance(stmt, ast.ImportFrom): + for alias in stmt.names: + bind(alias.asname or alias.name) + elif isinstance(stmt, ast.If): + visit_expr(stmt.test) + for body_stmt in [*stmt.body, *stmt.orelse]: + visit_stmt(body_stmt) + elif isinstance(stmt, ast.While): + visit_expr(stmt.test) + for body_stmt in [*stmt.body, *stmt.orelse]: + visit_stmt(body_stmt) + else: + for child in ast.iter_child_nodes(stmt): + if isinstance(child, ast.expr): + visit_expr(child) + elif isinstance(child, ast.stmt): + visit_stmt(child) + + for stmt in fn.body: + visit_stmt(stmt) + return used + + +def _unwrap_fn(target: Any) -> types.FunctionType | None: + """从模块/类绑定或 Subscriber 中取出裸函数对象""" + if isinstance(target, types.FunctionType): + return target + if isinstance(target, (classmethod, staticmethod)): + inner = target.__func__ + return inner if isinstance(inner, types.FunctionType) else None + if hasattr(target, "callable_target"): + inner = getattr(target, "callable_target") + return inner if isinstance(inner, types.FunctionType) else None + return None + + +def _resolve_old_fn(plugin: Plugin, qualname: str, ordinal: int, count: int) -> types.FunctionType | None: + """按 (qualname, 序号) 解析活模块中的旧函数对象;返回 None 表示不可就地替换 + + 最后一个同名定义经模块/类绑定解析(模块执行时同名后者胜出); + 非最后一个同名定义只能经插件 scope 的 Subscriber 按名字+注册序解析 + (即 `_` 惯用法:多个 `@plugin.listen` 装饰的同名函数)。 + """ + module = plugin.module + parts = qualname.split(".") + if len(parts) > 1: + if ordinal != count - 1: + return None + parent = module.__dict__.get(parts[0]) + if parent is None: + return None + for part in parts[1:-1]: + parent = inspect.getattr_static(parent, part, None) + if parent is None: + return None + return _unwrap_fn(inspect.getattr_static(parent, parts[-1], None)) + if ordinal == count - 1: + return _unwrap_fn(module.__dict__.get(parts[0])) + candidates = [ + slot.subscriber.callable_target + for slot in plugin._scope.subscribers + if isinstance(slot.subscriber.callable_target, types.FunctionType) + and slot.subscriber.callable_target.__name__ == parts[0] + ] + if len(candidates) != count: + return None + return candidates[ordinal] + + +def _build_new_fn(new_node: ast.FunctionDef | ast.AsyncFunctionDef, module: ModuleType) -> types.FunctionType | None: + """剥掉装饰器后在临时命名空间执行新函数定义,返回新函数对象""" + node = copy.deepcopy(new_node) + node.decorator_list = [] + ast.fix_missing_locations(node) + try: + code = compile( + ast.Module(body=[node], type_ignores=[]), + module.__name__, + "exec", + dont_inherit=True, + optimize=-1, + ) + ns = dict(module.__dict__) + exec(code, ns) + except Exception as e: + log.plugin.error(f"failed to build new function {node.name!r}: {e!r}") + return None + fn = ns.get(node.name) + return fn if isinstance(fn, types.FunctionType) else None + + +def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[FunctionChange]) -> bool: + """就地替换函数实现,如果会影响插件自身或下游依赖方则返回 False,调用方走全量重载""" + module = plugin.module + resolved: list[tuple[types.FunctionType, types.FunctionType, FunctionChange]] = [] + for change in changes: + new_fn = change.node + free_names = collect_global_names(new_fn) + missing = {name for name in free_names if name not in module.__dict__ and name not in vars(builtins)} + if missing: + log.plugin.warning( + f"cannot hot-swap {change.qualname!r}: free names {missing!r} missing, " + "fallback to full reload" + ) + return False + old_fn = _resolve_old_fn(plugin, change.qualname, change.ordinal, change.count) + if old_fn is None: + log.plugin.warning(f"cannot resolve {change.qualname!r} in module, fallback to full reload") + return False + if change.signature_changed: + from ..command import _commands + + if any(sub.callable_target is old_fn for sub in _commands.subscribers.values()): + log.plugin.warning( + f"signature of command target {change.qualname!r} changed, " "fallback to full reload" + ) + return False + new_function = _build_new_fn(new_fn, module) + if new_function is None: + return False + resolved.append((old_fn, new_function, change)) + for old_fn, new_function, change in resolved: + old_fn.__code__ = new_function.__code__ + old_fn.__defaults__ = new_function.__defaults__ + old_fn.__kwdefaults__ = new_function.__kwdefaults__ + old_fn.__annotations__ = new_function.__annotations__ + if old_fn.__qualname__ != "_": + plugin.module.__dict__[old_fn.__qualname__] = new_function + if change.signature_changed: + for slot in plugin._scope.subscribers: + sub = slot.subscriber + if sub.callable_target is old_fn: + sub.callable_target = new_function + try: + sub._recompile() + except Exception as e: + log.plugin.error(f"failed to recompile subscriber of {change.qualname!r}: {e!r}") + plugin._inspect = PluginInspect(new_nodes, ast.dump(new_nodes, include_attributes=False)) + return True From 7f2e31315f19d928a85aedf67276e979c243d0e7 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Sun, 9 Aug 2026 21:50:06 +0800 Subject: [PATCH 09/33] :sparkles: hot-swap support append fn in file tail --- arclet/entari/plugin/swap.py | 84 +++++++++++++++++++++++++++++++----- 1 file changed, 74 insertions(+), 10 deletions(-) diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py index fc34ac3..17f9c33 100644 --- a/arclet/entari/plugin/swap.py +++ b/arclet/entari/plugin/swap.py @@ -9,17 +9,20 @@ from types import ModuleType from typing import Any +from arclet.letoderea.scope import scope_ctx + from ..logger import log -from .model import Plugin, PluginInspect +from .model import Plugin, PluginInspect, current_plugin -@dataclass +@dataclass(slots=True) class FunctionChange: qualname: str ordinal: int count: int node: ast.FunctionDef | ast.AsyncFunctionDef signature_changed: bool + append: bool = False def list_dump(nodes: Sequence[ast.AST]) -> str: @@ -27,7 +30,7 @@ def list_dump(nodes: Sequence[ast.AST]) -> str: return f"[{', '.join(ast.dump(x, include_attributes=False) for x in nodes)}]" -def iter_functions(nodes: ast.Module) -> Iterator[tuple[str, int, ast.FunctionDef | ast.AsyncFunctionDef]]: +def iter_functions(body: list[ast.stmt]) -> Iterator[tuple[str, int, ast.FunctionDef | ast.AsyncFunctionDef]]: """按 qualname 收集模块顶层函数与类方法(含嵌套类)""" def _walk(body: list[ast.stmt], prefix: str) -> Iterator[tuple[str, int, ast.FunctionDef | ast.AsyncFunctionDef]]: @@ -41,7 +44,7 @@ def _walk(body: list[ast.stmt], prefix: str) -> Iterator[tuple[str, int, ast.Fun elif isinstance(stmt, ast.ClassDef): yield from _walk(stmt.body, f"{prefix}{stmt.name}.") - yield from _walk(nodes.body, "") + yield from _walk(body, "") def _structurally_same(o: ast.stmt, n: ast.stmt) -> bool: @@ -69,17 +72,22 @@ def classify(old_nodes: ast.Module, new_nodes: ast.Module) -> list[FunctionChang 函数配对:先按(名称组, 完整节点 dump 相同)锚定未变化函数(容忍位置交换/重排), 剩余函数按组内顺序配对比较(容忍原地编辑);重命名(组键变化)与增删 - (组内数量不匹配)→ 全量。 + (组内数量不匹配)→ 全量。文件底部追加的新 def 走 append 路径(整句执行, + 含装饰器注册)——要求旧 body 是新的结构前缀且尾部新增全为 def。 """ - if len(old_nodes.body) != len(new_nodes.body) or not all( - _structurally_same(o, n) for o, n in zip(old_nodes.body, new_nodes.body) - ): + old_body = old_nodes.body + new_body = new_nodes.body + if len(new_body) < len(old_body) or not all(_structurally_same(o, n) for o, n in zip(old_body, new_body)): + return None + tail: list[ast.FunctionDef | ast.AsyncFunctionDef] = new_body[len(old_body) :] # type: ignore + if tail and not all(isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)) for stmt in tail): return None + prefix_new = new_body[: len(old_body)] old_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} new_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} - for qualname, ordinal, fn in iter_functions(old_nodes): + for qualname, ordinal, fn in iter_functions(old_body): old_groups.setdefault(qualname, []).append((ordinal, fn)) - for qualname, ordinal, fn in iter_functions(new_nodes): + for qualname, ordinal, fn in iter_functions(prefix_new): new_groups.setdefault(qualname, []).append((ordinal, fn)) if set(old_groups.keys()) != set(new_groups.keys()): return None @@ -110,6 +118,8 @@ def classify(old_nodes: ast.Module, new_nodes: ast.Module) -> list[FunctionChang ) != ast.dump(o_node.args, include_attributes=False) # if signature_changed or list_dump(n_node.body) != list_dump(o_node.body): changes.append(FunctionChange(qualname, o_ord, len(old_list), n_node, signature_changed)) + for stmt in tail: + changes.append(FunctionChange(stmt.name, 0, 1, stmt, signature_changed=False, append=True)) return changes @@ -339,6 +349,43 @@ def _build_new_fn(new_node: ast.FunctionDef | ast.AsyncFunctionDef, module: Modu return fn if isinstance(fn, types.FunctionType) else None +def _decorator_global_names(node: ast.FunctionDef | ast.AsyncFunctionDef) -> set[str]: + """装饰器表达式中的全局引用(append 路径整句执行时需解析)""" + used: set[str] = set() + for dec in node.decorator_list: + for name in ast.walk(dec): + if isinstance(name, ast.Name) and isinstance(name.ctx, ast.Load): + used.add(name.id) + return used + + +def _exec_append(plugin: Plugin, node: ast.FunctionDef | ast.AsyncFunctionDef) -> bool: + """在插件上下文(current_plugin + scope)中整句执行新增语句,完成装饰器注册""" + module = plugin.module + stmt = copy.deepcopy(node) + ast.fix_missing_locations(stmt) + code = compile( + ast.Module(body=[stmt], type_ignores=[]), + module.__name__, + "exec", + dont_inherit=True, + optimize=-1, + ) + token = current_plugin.set(plugin) + try: + if not plugin.is_static: + token1 = scope_ctx.set(plugin._scope) + try: + exec(code, module.__dict__) + finally: + scope_ctx.reset(token1) + else: + exec(code, module.__dict__) + finally: + current_plugin.reset(token) + return True + + def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[FunctionChange]) -> bool: """就地替换函数实现,如果会影响插件自身或下游依赖方则返回 False,调用方走全量重载""" module = plugin.module @@ -353,6 +400,16 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function "fallback to full reload" ) return False + if change.append: + free_names |= _decorator_global_names(new_fn) + missing = {name for name in free_names if name not in module.__dict__ and name not in vars(builtins)} + if missing: + log.plugin.warning( + f"cannot append {change.qualname!r}: free names {missing!r} missing, " + "fallback to full reload" + ) + return False + continue old_fn = _resolve_old_fn(plugin, change.qualname, change.ordinal, change.count) if old_fn is None: log.plugin.warning(f"cannot resolve {change.qualname!r} in module, fallback to full reload") @@ -385,5 +442,12 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function sub._recompile() except Exception as e: log.plugin.error(f"failed to recompile subscriber of {change.qualname!r}: {e!r}") + for change in changes: + if change.append: + try: + _exec_append(plugin, change.node) + except Exception as e: + log.plugin.error(f"failed to append {change.qualname!r}: {e!r}") + return False plugin._inspect = PluginInspect(new_nodes, ast.dump(new_nodes, include_attributes=False)) return True From 9cf1eb7072738c1670d30fb3ac1b465277e53fb9 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Mon, 10 Aug 2026 10:43:25 +0800 Subject: [PATCH 10/33] :recycle: allow import changes when hot-swapping functions --- arclet/entari/plugin/swap.py | 425 +++++++++++++++++++++++------------ 1 file changed, 283 insertions(+), 142 deletions(-) diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py index 17f9c33..9357755 100644 --- a/arclet/entari/plugin/swap.py +++ b/arclet/entari/plugin/swap.py @@ -47,6 +47,29 @@ def _walk(body: list[ast.stmt], prefix: str) -> Iterator[tuple[str, int, ast.Fun yield from _walk(body, "") +def _signature_names(node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> set[str]: + names = {arg.arg for arg in node.args.posonlyargs + node.args.args + node.args.kwonlyargs} + if node.args.vararg: + names.add(node.args.vararg.arg) + if node.args.kwarg: + names.add(node.args.kwarg.arg) + return names + + +def _comp_targets(gen: ast.comprehension) -> set[str]: + targets: set[str] = set() + stack: list[ast.AST] = [gen.target] + while stack: + node = stack.pop() + if isinstance(node, ast.Name): + targets.add(node.id) + elif isinstance(node, (ast.Tuple, ast.List)): + stack.extend(node.elts) + elif isinstance(node, ast.Starred): + stack.append(node.value) + return targets + + def _structurally_same(o: ast.stmt, n: ast.stmt) -> bool: """结构比较:函数/方法体跳过(由 qualname 配对分析),其余须 dump 相等""" if isinstance(o, (ast.FunctionDef, ast.AsyncFunctionDef)) and isinstance( @@ -67,220 +90,299 @@ def _structurally_same(o: ast.stmt, n: ast.stmt) -> bool: return ast.dump(o, include_attributes=False) == ast.dump(n, include_attributes=False) -def classify(old_nodes: ast.Module, new_nodes: ast.Module) -> list[FunctionChange] | None: - """比较新旧模块 AST,返回需就地 swap 的函数列表;None 表示需模块层全量重载 +def _import_names(stmt: ast.stmt) -> dict[str, str]: + """import 语句绑定的 {名字: 目标模块}""" + if isinstance(stmt, ast.Import): + return {alias.asname or alias.name.split(".")[0]: alias.name for alias in stmt.names} + if isinstance(stmt, ast.ImportFrom): + module = "." * (stmt.level or 0) + (stmt.module or "") + return {alias.asname or alias.name: f"{module}.{alias.name}" for alias in stmt.names if alias.name != "*"} + return {} - 函数配对:先按(名称组, 完整节点 dump 相同)锚定未变化函数(容忍位置交换/重排), - 剩余函数按组内顺序配对比较(容忍原地编辑);重命名(组键变化)与增删 - (组内数量不匹配)→ 全量。文件底部追加的新 def 走 append 路径(整句执行, - 含装饰器注册)——要求旧 body 是新的结构前缀且尾部新增全为 def。 - """ - old_body = old_nodes.body - new_body = new_nodes.body - if len(new_body) < len(old_body) or not all(_structurally_same(o, n) for o, n in zip(old_body, new_body)): - return None - tail: list[ast.FunctionDef | ast.AsyncFunctionDef] = new_body[len(old_body) :] # type: ignore - if tail and not all(isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)) for stmt in tail): - return None - prefix_new = new_body[: len(old_body)] - old_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} - new_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} - for qualname, ordinal, fn in iter_functions(old_body): - old_groups.setdefault(qualname, []).append((ordinal, fn)) - for qualname, ordinal, fn in iter_functions(prefix_new): - new_groups.setdefault(qualname, []).append((ordinal, fn)) - if set(old_groups.keys()) != set(new_groups.keys()): - return None - changes: list[FunctionChange] = [] - for qualname, old_list in old_groups.items(): - new_list = new_groups[qualname] - if len(old_list) != len(new_list): - return None - used_old: set[int] = set() - used_new: set[int] = set() - for i, (_, o_node) in enumerate(old_list): - for j, (_, n_node) in enumerate(new_list): - if ( - i not in used_old - and j not in used_new - and ast.dump(o_node, include_attributes=False) == ast.dump(n_node, include_attributes=False) - ): - used_old.add(i) - used_new.add(j) - break - rest_old = [old_list[i] for i in range(len(old_list)) if i not in used_old] - rest_new = [new_list[j] for j in range(len(new_list)) if j not in used_new] - for (o_ord, o_node), (n_ord, n_node) in zip(rest_old, rest_new): - if list_dump(n_node.decorator_list) != list_dump(o_node.decorator_list): - return None - signature_changed = n_node.__class__ is not o_node.__class__ or ast.dump( - n_node.args, include_attributes=False - ) != ast.dump(o_node.args, include_attributes=False) - # if signature_changed or list_dump(n_node.body) != list_dump(o_node.body): - changes.append(FunctionChange(qualname, o_ord, len(old_list), n_node, signature_changed)) - for stmt in tail: - changes.append(FunctionChange(stmt.name, 0, 1, stmt, signature_changed=False, append=True)) - return changes +def _align_prefix( + old_body: list[ast.stmt], new_body: list[ast.stmt] +) -> tuple[dict[str, str], dict[str, str], list[tuple[ast.stmt, ast.stmt]], list[ast.stmt]] | None: + """对齐新旧模块前缀:允许 import 语句就地变化/增删,其余语句须结构相同 -def _signature_names(node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> set[str]: - names = {arg.arg for arg in node.args.posonlyargs + node.args.args + node.args.kwonlyargs} - if node.args.vararg: - names.add(node.args.vararg.arg) - if node.args.kwarg: - names.add(node.args.kwarg.arg) - return names + 返回 (旧 import 绑定, 新 import 绑定, 对齐的非 import 语句对, 新尾部)。 + None 表示结构无法匹配(需全量重载)。 + """ + old_imports: dict[str, str] = {} + new_imports: dict[str, str] = {} + pairs: list[tuple[ast.stmt, ast.stmt]] = [] + i = j = 0 + while i < len(old_body) and j < len(new_body): + o, n = old_body[i], new_body[j] + if isinstance(o, (ast.Import, ast.ImportFrom)) and isinstance(n, (ast.Import, ast.ImportFrom)): + old_imports.update(_import_names(o)) + new_imports.update(_import_names(n)) + i += 1 + j += 1 + elif isinstance(o, (ast.Import, ast.ImportFrom)): + old_imports.update(_import_names(o)) + i += 1 + elif isinstance(n, (ast.Import, ast.ImportFrom)): + new_imports.update(_import_names(n)) + j += 1 + elif _structurally_same(o, n): + pairs.append((o, n)) + i += 1 + j += 1 + else: + return None + if i != len(old_body): + return None + return old_imports, new_imports, pairs, new_body[j:] -def _comp_targets(gen: ast.comprehension) -> set[str]: - targets: set[str] = set() - stack: list[ast.AST] = [gen.target] - while stack: - node = stack.pop() - if isinstance(node, ast.Name): - targets.add(node.id) - elif isinstance(node, (ast.Tuple, ast.List)): - stack.extend(node.elts) - elif isinstance(node, ast.Starred): - stack.append(node.value) - return targets +def _changed_import_names(old_imports: dict[str, str], new_imports: dict[str, str]) -> set[str]: + """新旧 import 绑定的变更名集合(新增/移除/改绑)""" + common = set(old_imports) & set(new_imports) + return (set(old_imports) ^ set(new_imports)) | {n for n in common if old_imports[n] != new_imports[n]} -def collect_global_names(fn: ast.FunctionDef | ast.AsyncFunctionDef) -> set[str]: - """收集新函数体中所有会在模块全局或 builtins 解析的名字(含嵌套函数) +class _GlobalNameCollector: + """作用域感知的名字收集器:收集会在模块全局或 builtins 解析的名字 规则:Name(Load) 在自身与所有外层函数作用域均未绑定 → 全局名。 """ - used: set[str] = set() - stack: list[set[str]] = [_signature_names(fn)] - def bound(name: str) -> bool: - return any(name in scope for scope in stack) + __slots__ = ("used", "stack") + + def __init__(self): + self.used: set[str] = set() + self.stack: list[set[str]] = [set()] + + def _bound(self, name: str) -> bool: + return any(name in scope for scope in self.stack) + + def _bind(self, name: str): + self.stack[-1].add(name) - def bind(name: str): - stack[-1].add(name) + def collect(self, stmts: list[ast.stmt], initial: set[str] | None = None) -> set[str]: + self.used = set() + self.stack = [initial or set()] + for stmt in stmts: + self.visit_stmt(stmt) + return self.used - def visit_expr(expr: ast.expr): + def visit_expr(self, expr: ast.expr): if isinstance(expr, ast.Name): - if isinstance(expr.ctx, ast.Load) and not bound(expr.id): - used.add(expr.id) + if isinstance(expr.ctx, ast.Load) and not self._bound(expr.id): + self.used.add(expr.id) elif isinstance(expr.ctx, ast.Store): - bind(expr.id) + self._bind(expr.id) elif isinstance(expr, ast.Lambda): - stack.append(_signature_names(expr)) - visit_expr(expr.body) - stack.pop() + self.stack.append(_signature_names(expr)) + self.visit_expr(expr.body) + self.stack.pop() elif isinstance(expr, (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp)): comps = expr.generators if comps: - visit_expr(comps[0].iter) - stack.append({name for gen in comps for name in _comp_targets(gen)}) + self.visit_expr(comps[0].iter) + self.stack.append({name for gen in comps for name in _comp_targets(gen)}) for gen in comps: if gen is not comps[0]: - visit_expr(gen.iter) + self.visit_expr(gen.iter) for cond in gen.ifs: - visit_expr(cond) + self.visit_expr(cond) if isinstance(expr, ast.DictComp): - visit_expr(expr.key) - visit_expr(expr.value) + self.visit_expr(expr.key) + self.visit_expr(expr.value) else: - visit_expr(expr.elt) - stack.pop() + self.visit_expr(expr.elt) + self.stack.pop() else: for child in ast.iter_child_nodes(expr): if isinstance(child, ast.expr): - visit_expr(child) + self.visit_expr(child) - def visit_stmt(stmt: ast.stmt): + def visit_stmt(self, stmt: ast.stmt): if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): - bind(stmt.name) + self._bind(stmt.name) for dec in stmt.decorator_list: - visit_expr(dec) + self.visit_expr(dec) for default in [*stmt.args.defaults, *[d for d in stmt.args.kw_defaults if d]]: - visit_expr(default) + self.visit_expr(default) for arg in stmt.args.posonlyargs + stmt.args.args + stmt.args.kwonlyargs: if arg.annotation: - visit_expr(arg.annotation) + self.visit_expr(arg.annotation) if stmt.args.vararg and stmt.args.vararg.annotation: - visit_expr(stmt.args.vararg.annotation) + self.visit_expr(stmt.args.vararg.annotation) if stmt.args.kwarg and stmt.args.kwarg.annotation: - visit_expr(stmt.args.kwarg.annotation) - stack.append(_signature_names(stmt)) + self.visit_expr(stmt.args.kwarg.annotation) + self.stack.append(_signature_names(stmt)) for body_stmt in stmt.body: - visit_stmt(body_stmt) - stack.pop() + self.visit_stmt(body_stmt) + self.stack.pop() elif isinstance(stmt, ast.ClassDef): - bind(stmt.name) + self._bind(stmt.name) for dec in stmt.decorator_list: - visit_expr(dec) + self.visit_expr(dec) for base in stmt.bases: - visit_expr(base) + self.visit_expr(base) for kw in stmt.keywords: - visit_expr(kw.value) - stack.append(set()) + self.visit_expr(kw.value) + self.stack.append(set()) for body_stmt in stmt.body: - visit_stmt(body_stmt) - stack.pop() + self.visit_stmt(body_stmt) + self.stack.pop() elif isinstance(stmt, ast.Assign): for target in stmt.targets: - visit_expr(target) - visit_expr(stmt.value) + self.visit_expr(target) + self.visit_expr(stmt.value) elif isinstance(stmt, ast.AnnAssign): - visit_expr(stmt.target) - visit_expr(stmt.annotation) + self.visit_expr(stmt.target) + self.visit_expr(stmt.annotation) if stmt.value: - visit_expr(stmt.value) + self.visit_expr(stmt.value) elif isinstance(stmt, ast.AugAssign): - visit_expr(stmt.target) - visit_expr(stmt.value) + self.visit_expr(stmt.target) + self.visit_expr(stmt.value) elif isinstance(stmt, (ast.For, ast.AsyncFor)): - visit_expr(stmt.target) - visit_expr(stmt.iter) + self.visit_expr(stmt.target) + self.visit_expr(stmt.iter) for body_stmt in [*stmt.body, *stmt.orelse]: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) elif isinstance(stmt, (ast.With, ast.AsyncWith)): for item in stmt.items: - visit_expr(item.context_expr) + self.visit_expr(item.context_expr) if item.optional_vars: - visit_expr(item.optional_vars) + self.visit_expr(item.optional_vars) for body_stmt in stmt.body: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) elif isinstance(stmt, ast.Try) or (sys.version_info >= (3, 11) and isinstance(stmt, ast.TryStar)): for body_stmt in [*stmt.body, *stmt.orelse, *stmt.finalbody]: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) for handler in stmt.handlers: if handler.type: - visit_expr(handler.type) + self.visit_expr(handler.type) if handler.name: - bind(handler.name) + self._bind(handler.name) for body_stmt in handler.body: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) elif isinstance(stmt, ast.Import): for alias in stmt.names: - bind(alias.asname or alias.name.split(".")[0]) + self._bind(alias.asname or alias.name.split(".")[0]) elif isinstance(stmt, ast.ImportFrom): for alias in stmt.names: - bind(alias.asname or alias.name) + self._bind(alias.asname or alias.name) elif isinstance(stmt, ast.If): - visit_expr(stmt.test) + self.visit_expr(stmt.test) for body_stmt in [*stmt.body, *stmt.orelse]: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) elif isinstance(stmt, ast.While): - visit_expr(stmt.test) + self.visit_expr(stmt.test) for body_stmt in [*stmt.body, *stmt.orelse]: - visit_stmt(body_stmt) + self.visit_stmt(body_stmt) else: for child in ast.iter_child_nodes(stmt): if isinstance(child, ast.expr): - visit_expr(child) + self.visit_expr(child) elif isinstance(child, ast.stmt): - visit_stmt(child) + self.visit_stmt(child) - for stmt in fn.body: - visit_stmt(stmt) + +def _collect_unchanged_references( + pairs: list[tuple[ast.stmt, ast.stmt]], changes: list[FunctionChange], new_nodes: ast.Module +) -> set[str]: + """收集未变更代码(顶层语句、类级语句、未变更函数/方法体)引用的全局名""" + collector = _GlobalNameCollector() + used: set[str] = set() + for _, new_stmt in pairs: + if isinstance(new_stmt, (ast.Import, ast.ImportFrom, ast.FunctionDef, ast.AsyncFunctionDef)): + continue + if isinstance(new_stmt, ast.ClassDef): + class_stmts = [s for s in new_stmt.body if not isinstance(s, (ast.FunctionDef, ast.AsyncFunctionDef))] + if class_stmts: + used |= collector.collect(class_stmts) + continue + used |= collector.collect([new_stmt]) + changed_nodes = {c.node for c in changes} + for _, _, fn in iter_functions(new_nodes.body): + if fn in changed_nodes: + continue + used |= collector.collect(fn.body, _signature_names(fn)) return used +def classify(old_nodes: ast.Module, new_nodes: ast.Module) -> list[FunctionChange] | None: + """比较新旧模块的 AST,返回可以热替换的函数列表 + + 函数配对:先按(名称组, 完整节点 dump 相同)锚定未变化函数(容忍位置交换/重排), + 剩余函数按组内顺序配对比较(容忍原地编辑);重命名(组键变化)与增删(组内数量不匹配)走全量重载。 + 文件底部追加的新函数 走 append 路径(整句执行,含装饰器注册,要求旧 body 是新的结构前缀且尾部新增全为函数定义)。 + 前缀中的 import 语句允许就地变化/增删,条件是变更的绑定名不被任何未变更代码引用(否则走全量重载)。 + + Returns: + 可就地替换的函数列表,None 表示模块整体需要全量重载 + """ + old_body = old_nodes.body + new_body = new_nodes.body + aligned = _align_prefix(old_body, new_body) + if aligned is None: + return None + old_imports, new_imports, pairs, tail = aligned + if tail and not all( + isinstance(s, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Import, ast.ImportFrom)) for s in tail + ): + return None + old_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} + new_groups: dict[str, list[tuple[int, ast.FunctionDef | ast.AsyncFunctionDef]]] = {} + for qualname, ordinal, fn in iter_functions([o for o, _ in pairs]): + old_groups.setdefault(qualname, []).append((ordinal, fn)) + for qualname, ordinal, fn in iter_functions([n for _, n in pairs]): + new_groups.setdefault(qualname, []).append((ordinal, fn)) + if set(old_groups.keys()) != set(new_groups.keys()): + return None + changes: list[FunctionChange] = [] + for qualname, old_list in old_groups.items(): + new_list = new_groups[qualname] + if len(old_list) != len(new_list): + return None + used_old: set[int] = set() + used_new: set[int] = set() + for i, (_, o_node) in enumerate(old_list): + for j, (_, n_node) in enumerate(new_list): + if ( + i not in used_old + and j not in used_new + and ast.dump(o_node, include_attributes=False) == ast.dump(n_node, include_attributes=False) + ): + used_old.add(i) + used_new.add(j) + break + rest_old = [old_list[i] for i in range(len(old_list)) if i not in used_old] + rest_new = [new_list[j] for j in range(len(new_list)) if j not in used_new] + for (o_ord, o_node), (n_ord, n_node) in zip(rest_old, rest_new): + if list_dump(n_node.decorator_list) != list_dump(o_node.decorator_list): + return None + signature_changed = n_node.__class__ is not o_node.__class__ or ast.dump( + n_node.args, include_attributes=False + ) != ast.dump(o_node.args, include_attributes=False) + # if signature_changed or list_dump(n_node.body) != list_dump(o_node.body): + changes.append(FunctionChange(qualname, o_ord, len(old_list), n_node, signature_changed)) + for stmt in tail: + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): + changes.append(FunctionChange(stmt.name, 0, 1, stmt, signature_changed=False, append=True)) + if changed_import_names := _changed_import_names(old_imports, new_imports): + used = _collect_unchanged_references(pairs, changes, new_nodes) + if changed_import_names & used: + log.plugin.debug( + f"changed import names {sorted(changed_import_names & used)!r} are referenced " + f"by unchanged code, fallback to full reload" + ) + return None + return changes + + +def collect_global_names(fn: ast.FunctionDef | ast.AsyncFunctionDef) -> set[str]: + """收集新函数体中所有会在模块全局或 builtins 解析的名字(含嵌套函数)""" + collector = _GlobalNameCollector() + return collector.collect(fn.body, initial=_signature_names(fn)) + + def _unwrap_fn(target: Any) -> types.FunctionType | None: """从模块/类绑定或 Subscriber 中取出裸函数对象""" if isinstance(target, types.FunctionType): @@ -359,6 +461,29 @@ def _decorator_global_names(node: ast.FunctionDef | ast.AsyncFunctionDef) -> set return used +def _exec_imports(plugin: Plugin, nodes: ast.Module) -> bool: + """将新源码中的 import 语句按源码顺序执行并写入模块中""" + module = plugin.module + for stmt in nodes.body: + if not isinstance(stmt, (ast.Import, ast.ImportFrom)): + continue + stmt_cp = copy.deepcopy(stmt) + ast.fix_missing_locations(stmt_cp) + code = compile( + ast.Module(body=[stmt_cp], type_ignores=[]), + module.__name__, + "exec", + dont_inherit=True, + optimize=-1, + ) + try: + exec(code, module.__dict__) + except Exception as e: + log.plugin.error(f"failed to exec import statement at line {stmt.lineno}: {e!r}") + return False + return True + + def _exec_append(plugin: Plugin, node: ast.FunctionDef | ast.AsyncFunctionDef) -> bool: """在插件上下文(current_plugin + scope)中整句执行新增语句,完成装饰器注册""" module = plugin.module @@ -389,11 +514,21 @@ def _exec_append(plugin: Plugin, node: ast.FunctionDef | ast.AsyncFunctionDef) - def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[FunctionChange]) -> bool: """就地替换函数实现,如果会影响插件自身或下游依赖方则返回 False,调用方走全量重载""" module = plugin.module + import_names = { + name + for stmt in new_nodes.body + if isinstance(stmt, (ast.Import, ast.ImportFrom)) + for name in _import_names(stmt) + } resolved: list[tuple[types.FunctionType, types.FunctionType, FunctionChange]] = [] for change in changes: new_fn = change.node free_names = collect_global_names(new_fn) - missing = {name for name in free_names if name not in module.__dict__ and name not in vars(builtins)} + missing = { + name + for name in free_names + if name not in module.__dict__ and name not in vars(builtins) and name not in import_names + } if missing: log.plugin.warning( f"cannot hot-swap {change.qualname!r}: free names {missing!r} missing, " @@ -402,7 +537,11 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function return False if change.append: free_names |= _decorator_global_names(new_fn) - missing = {name for name in free_names if name not in module.__dict__ and name not in vars(builtins)} + missing = { + name + for name in free_names + if name not in module.__dict__ and name not in vars(builtins) and name not in import_names + } if missing: log.plugin.warning( f"cannot append {change.qualname!r}: free names {missing!r} missing, " @@ -426,6 +565,8 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function if new_function is None: return False resolved.append((old_fn, new_function, change)) + if not _exec_imports(plugin, new_nodes): + return False for old_fn, new_function, change in resolved: old_fn.__code__ = new_function.__code__ old_fn.__defaults__ = new_function.__defaults__ From 45f544dcec10aca245c6046b3e56c5c1233689ff Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 09:31:31 +0800 Subject: [PATCH 11/33] :sparkles: replace referent cascade with rebind driver and surface fingerprint --- arclet/entari/plugin/__init__.py | 188 +++++-------------- arclet/entari/plugin/dependent.py | 297 ++++++++++++++++++++++++++++++ arclet/entari/plugin/model.py | 92 +++++---- arclet/entari/plugin/model.pyi | 2 + arclet/entari/plugin/module.py | 32 ++++ arclet/entari/plugin/service.py | 18 +- 6 files changed, 445 insertions(+), 184 deletions(-) create mode 100644 arclet/entari/plugin/dependent.py diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index d0cb445..f9b9d03 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -7,7 +7,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload -from arclet.letoderea import Subscriber, on, publish +from arclet.letoderea import Subscriber, on from arclet.letoderea.effect import AsyncDisposable, Disposable from arclet.letoderea.utils import Resultable from tarina import init_spec @@ -15,19 +15,17 @@ from ..config import EntariConfig, config_model_keys, config_model_validate from ..event.config import ConfigReload -from ..event.lifespan import Ready -from ..event.plugin import PluginLoadedFailed -from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError -from ..logger import log +from ..exceptions import StaticPluginDispatchError from ..message import Fragment, MessageChain, Render from ..session import COMPONENTS, Session, component_transform -from ..utils import escape_tag +from .dependent import find_plugin as find_plugin +from .dependent import load_plugin as load_plugin +from .dependent import unload_plugin as unload_plugin from .model import TS, Plugin, PluginDispatcher, current_plugin from .model import PluginMetadata as PluginMetadata from .model import PluginRole as PluginRole from .model import RootlessPlugin as RootlessPlugin from .model import keeping as keeping -from .module import import_plugin from .module import package as package from .module import requires as requires from .service import plugin_service @@ -200,108 +198,6 @@ def dispatch(event: type, name: str | None = None) -> PluginDispatcher: return get_plugin(1).dispatch(event, name=name) -def load_plugin( - path: str, config: dict | None = None, recursive_guard: set[str] | None = None, prelude: bool = False -) -> Plugin | None: - """ - 以导入路径方式加载模块 - - Args: - path (str): 模块路径 - config (dict): 模块配置 - recursive_guard (set[str]): 递归保护 - prelude (bool): 是否为前置插件 - """ - if config is not None: - config["$path"] = path - else: - for k, names in EntariConfig.instance._plugin_names.items(): - if path in names: - config = EntariConfig.instance.plugin.get(k, {}) - config["$path"] = k - break - else: - config = {"$path": path} - if prelude: - config["$static"] = True - if recursive_guard is None: - recursive_guard = set() - path = path.replace("::", "arclet.entari.builtins.") - while path in plugin_service._subplugined: - path = plugin_service._subplugined[path] - if path in plugin_service._apply: - if path in plugin_service.plugins: - return plugin_service.plugins[path] - log.plugin.trace(f"loaded rootless plugin {path!r}") - return plugin_service._apply[path][0](config) - if plug := find_plugin(path): - plugin_service._direct_plugins.add(plug.path) - return plug - try: - mod = import_plugin(path, config=config) - if not mod: - mod = next( - (import_plugin(_path, config=config) for _path in EntariConfig.instance._plugin_names.get(path, [])), - None, - ) - if not mod: - log.plugin.error(f"cannot found plugin {path!r}") - publish(PluginLoadedFailed(path)) - return - plugin_service._direct_plugins.add(mod.__name__) - if mod.__name__ in plugin_service.referents and plugin_service.referents[mod.__name__]: - referents = plugin_service.referents[mod.__name__].copy() - # plugin_service.referents[mod.__name__].clear() - for referent in referents: - if referent in recursive_guard: - continue - if referent.startswith(mod.__name__): - continue - if referent in plugin_service._subplugined and mod.__name__.startswith( - plugin_service._subplugined[referent] - ): - continue - if referent in plugin_service.plugins: - plugin_service.referents[mod.__name__].discard(referent) - log.plugin.debug(f"reloading {escape_tag(mod.__name__)}'s referent {referent!r}") - unload_plugin(referent) - if not (plug := load_plugin(referent)): - plugin_service.referents[mod.__name__].add(referent) - else: - publish(Ready(), plug._scope) - recursive_guard.add(referent) - - return mod.__plugin__ - except (ImportError, RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError): - return - except Exception as e: - log.plugin.exception(f"failed to load plugin {path!r}: {e}", exc_info=e) - publish(PluginLoadedFailed(path)) - return - - -def load_plugins(dir_: str | os.PathLike | Path): - """加载指定目录下的所有插件""" - path = dir_ if isinstance(dir_, Path) else Path(dir_) - if not path.is_dir(): - raise NotADirectoryError(f"{path} is not a directory") - path: Path = path.resolve() # .relative_to(Path.cwd()) - syspaths = [Path(p).resolve() for p in sys.path if p] - prefixes = [p for p in syspaths if path.is_relative_to(p)] - if prefixes: - prefix = max(prefixes, key=lambda p: len(p.parts)) - else: - prefix = Path.cwd() - for p in path.iterdir(): - if p.suffix in (".py", "") and p.stem not in {"__init__", "__pycache__"}: - p = p.resolve().relative_to(prefix) - if len(p.parts) > 1: - plg = ".".join(p.parts[:-1:1]) + "." + p.stem - else: - plg = p.stem - load_plugin(plg) - - if TYPE_CHECKING: @init_spec(PluginMetadata) @@ -370,6 +266,43 @@ def _reload(event: ConfigReload): get_config = plugin_config +def load_plugins(dir_: str | os.PathLike | Path): + """加载指定目录下的所有插件""" + path = dir_ if isinstance(dir_, Path) else Path(dir_) + if not path.is_dir(): + raise NotADirectoryError(f"{path} is not a directory") + path: Path = path.resolve() # .relative_to(Path.cwd()) + syspaths = [Path(p).resolve() for p in sys.path if p] + prefixes = [p for p in syspaths if path.is_relative_to(p)] + if prefixes: + prefix = max(prefixes, key=lambda p: len(p.parts)) + else: + prefix = Path.cwd() + for p in path.iterdir(): + if p.suffix in (".py", "") and p.stem not in {"__init__", "__pycache__"}: + p = p.resolve().relative_to(prefix) + if len(p.parts) > 1: + plg = ".".join(p.parts[:-1:1]) + "." + p.stem + else: + plg = p.stem + load_plugin(plg) + + +def find_plugin_by_file(file: str) -> Plugin | None: + path = Path(file).resolve() + for plugin in plugin_service.plugins.values(): + if plugin.module.__file__ == str(path): + return plugin + if plugin.module.__file__ and Path(plugin.module.__file__).parent == path: + return plugin + path1 = Path(path) + while path1.parent != path1: + if str(path1) == plugin.module.__file__: + return plugin + path1 = path1.parent + return None + + def declare_static(): """声明当前插件为静态插件""" _plugin = get_plugin(1) @@ -400,43 +333,6 @@ def restore(): return get_plugin(1).restore() -def find_plugin(name: str) -> Plugin | None: - """根据插件名称查找插件""" - if name in plugin_service.plugins: - return plugin_service.plugins[name] - if name in EntariConfig.instance.plugin_prefixes: - for prefix in EntariConfig.instance.plugin_prefixes[name]: - if f"{prefix}{name}" in plugin_service.plugins: - return plugin_service.plugins[f"{prefix}{name}"] - if not name.count(".") and f"entari_plugin_{name}" in plugin_service.plugins: - return plugin_service.plugins[f"entari_plugin_{name}"] - - -def find_plugin_by_file(file: str) -> Plugin | None: - path = Path(file).resolve() - for plugin in plugin_service.plugins.values(): - if plugin.module.__file__ == str(path): - return plugin - if plugin.module.__file__ and Path(plugin.module.__file__).parent == path: - return plugin - path1 = Path(path) - while path1.parent != path1: - if str(path1) == plugin.module.__file__: - return plugin - path1 = path1.parent - return None - - -def unload_plugin(plugin: str): - plugin = plugin.replace("::", "arclet.entari.builtins.") - while plugin in plugin_service._subplugined: - plugin = plugin_service._subplugined[plugin] - if not (_plugin := find_plugin(plugin)): - return False - _plugin.dispose() - return True - - async def unload_plugin_async(plugin: str): plugin = plugin.replace("::", "arclet.entari.builtins.") while plugin in plugin_service._subplugined: diff --git a/arclet/entari/plugin/dependent.py b/arclet/entari/plugin/dependent.py new file mode 100644 index 0000000..0286078 --- /dev/null +++ b/arclet/entari/plugin/dependent.py @@ -0,0 +1,297 @@ +import ast +from typing import Any + +from arclet.letoderea import publish + +from ..config import EntariConfig +from ..event.lifespan import Ready +from ..event.plugin import PluginLoadedFailed +from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError +from ..logger import log +from .model import Plugin +from .module import import_plugin +from .service import plugin_service + + +def collect_module_level_names(nodes: ast.Module) -> set[str]: + """收集模块层表达式引用的名字(函数体除外;含装饰器/默认值/注解/基类/类级语句)""" + used: set[str] = set() + + def visit_expr(expr: ast.expr): + if isinstance(expr, ast.Name) and isinstance(expr.ctx, ast.Load): + used.add(expr.id) + return + for child in ast.iter_child_nodes(expr): + if isinstance(child, ast.expr): + visit_expr(child) + + def visit_stmt(stmt: ast.stmt): + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): + for dec in stmt.decorator_list: + visit_expr(dec) + for default in [*stmt.args.defaults, *[d for d in stmt.args.kw_defaults if d]]: + visit_expr(default) + for arg in stmt.args.posonlyargs + stmt.args.args + stmt.args.kwonlyargs: + if arg.annotation: + visit_expr(arg.annotation) + if stmt.args.vararg and stmt.args.vararg.annotation: + visit_expr(stmt.args.vararg.annotation) + if stmt.args.kwarg and stmt.args.kwarg.annotation: + visit_expr(stmt.args.kwarg.annotation) + return + if isinstance(stmt, ast.ClassDef): + for dec in stmt.decorator_list: + visit_expr(dec) + for base in stmt.bases: + visit_expr(base) + for kw in stmt.keywords: + visit_expr(kw.value) + for s in stmt.body: + visit_stmt(s) + return + for child in ast.iter_child_nodes(stmt): + if isinstance(child, ast.expr): + visit_expr(child) + elif isinstance(child, ast.stmt): + visit_stmt(child) + + for stmt in nodes.body: + visit_stmt(stmt) + return used + + +def _uses_module_level(plugin: Plugin, path: str) -> bool: + """检测插件是否在模块层使用 path 的绑定名(基类/模块级装饰器/顶层实例化等)""" + if not plugin._inspect: + return True + bound = {name for name, (target, _) in plugin.bindings.items() if target == path or target.startswith(path + ".")} + if not bound: + return False + return bool(collect_module_level_names(plugin._inspect.nodes) & bound) + + +_MISSING = object() + + +def _rebind_imports(plugin: Plugin, path: str) -> bool: + """将插件对 path(或其子树)的 import 绑定改写为新对象;False 表示有绑定无法满足(升级全量) + + 模块绑定(attr 为 None)仅在名字等同于目标模块或其末组件时写回:`import a.b` 风格的名字是顶层包,其子模块链由 promote 的父属性写回维护,不在此处理。 + """ + module = plugin.module + for name, (target, attr) in plugin.bindings.items(): + if not (target == path or target.startswith(path + ".")): + continue + if target not in plugin_service.plugins: + continue + new_module = plugin_service.plugins[target].module + if attr is None: + if name != target and name != target.rpartition(".")[-1]: + continue + value: Any = new_module + else: + value = getattr(new_module, attr, _MISSING) + if value is _MISSING: + return False + module.__dict__[name] = value + return True + + +def public_fingerprint(plugin: Plugin) -> str | None: + """插件公开面指纹:排序后的导出名 + 各名字的定义节点的dump(导入名仅记录类型) + + 无 `__all__` 时导出名 = 非下划线开头的模块 dict 名字;定义节点 dump 覆盖函数体与类体,任何行为差异都反映为指纹差异。 + """ + if not plugin._inspect: + return None + module = plugin.module + names = getattr(module, "__all__", None) + if names is None: + names = sorted(n for n in module.__dict__ if not n.startswith("_")) + else: + names = sorted(names) + node_map: dict[str, ast.stmt] = {} + for stmt in plugin._inspect.nodes.body: + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + node_map[stmt.name] = stmt + elif isinstance(stmt, ast.Assign): + for target in stmt.targets: + if isinstance(target, ast.Name): + node_map[target.id] = stmt + elif isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Name): + node_map[stmt.target.id] = stmt + parts: list[str] = [] + for name in names: + if node := node_map.get(name): + parts.append(f"{name}:{ast.dump(node, include_attributes=False)}") + else: + value = module.__dict__.get(name, None) + parts.append(f"{name}:imported:{type(value).__name__}") + return "\n".join(parts) + + +def dependents_of(path: str) -> list[str]: + """path 的依赖方插件:直接导入者 + 经父包再导出链一层(F5)""" + parent_pkg = path.rpartition(".")[0] + result: list[str] = [] + for plug_id, bindings in plugin_service.bindings.items(): + for name, (target, attr) in bindings.items(): + if target == path: + result.append(plug_id) + break + if parent_pkg and target == parent_pkg and parent_pkg in plugin_service.bindings: + parent_chain = plugin_service.bindings[parent_pkg] + if parent_chain.get(attr or name, (None, None))[0] == path: + result.append(plug_id) + break + return result + + +def _fingerprint_changed(plugin: Plugin) -> bool: + """比较插件公开面指纹(存储的旧指纹 vs 现算新指纹),并更新存储""" + path = plugin.path + old = plugin_service.fingerprints.get(path) + new = public_fingerprint(plugin) + plugin_service.fingerprints[path] = new or "" + return bool(old) and old != new + + +def _rebind_dep(dep: Plugin, path: str): + """重绑依赖方对 path 的 import 绑定并恢复其可用状态""" + if _rebind_imports(dep, path): + log.plugin.debug(f"rebound {dep.id!r}'s imports on {path}") + dep.enable() + else: + log.plugin.warning(f"cannot satisfy {dep.id!r}'s imports on {path}, falling back to full reload") + _cascade_dep(dep.id, plugin_service.referents.get(path, set()), set()) + + +def _cascade_dep(dep_id: str, referent_set: set[str], recursive_guard: set[str]): + """对依赖方执行全量级联(unload + load + Ready),失败时恢复 referent 边""" + referent_set.discard(dep_id) + log.plugin.debug(f"reloading {dep_id!r}") + unload_plugin(dep_id) + if not (plug := load_plugin(dep_id)): + referent_set.add(dep_id) + else: + publish(Ready(), plug._scope) + recursive_guard.add(dep_id) + + +def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): + """插件完成(重)加载后处理依赖方:公开面未变 → 全部重绑;否则按模块层使用细粒度重绑或级联 + + 依赖方 = referents 图与绑定索引的并集;公开面指纹未变时新对象与旧对象行为等价,模块层绑定无需重执行; + 指纹变化时仅模块层使用方升级全量,惰性依赖方重绑即可。 + """ + if recursive_guard is None: + recursive_guard = set() + path = plugin.path + dependents = set(plugin_service.referents.get(path, ())) | set(dependents_of(path)) + if not dependents: + return + surface_changed = _fingerprint_changed(plugin) + referent_set = plugin_service.referents.setdefault(path, set()) + for dep_id in sorted(dependents): + if dep_id in recursive_guard: + continue + if dep_id.startswith(path): + continue + if dep_id in plugin_service._subplugined and path.startswith(plugin_service._subplugined[dep_id]): + continue + dep = plugin_service.plugins.get(dep_id) + if dep is None: + if not (plug := load_plugin(dep_id)): + referent_set.add(dep_id) + continue + publish(Ready(), plug._scope) + recursive_guard.add(dep_id) + continue + if not surface_changed: + _rebind_dep(dep, path) + continue + if _uses_module_level(dep, path): + _cascade_dep(dep_id, referent_set, recursive_guard) + else: + _rebind_dep(dep, path) + + +def load_plugin( + path: str, config: dict | None = None, recursive_guard: set[str] | None = None, prelude: bool = False +) -> Plugin | None: + """ + 以导入路径方式加载模块 + + Args: + path (str): 模块路径 + config (dict): 模块配置 + recursive_guard (set[str]): 递归保护 + prelude (bool): 是否为前置插件 + """ + if config is not None: + config["$path"] = path + else: + for k, names in EntariConfig.instance._plugin_names.items(): + if path in names: + config = EntariConfig.instance.plugin.get(k, {}) + config["$path"] = k + break + else: + config = {"$path": path} + if prelude: + config["$static"] = True + if recursive_guard is None: + recursive_guard = set() + path = path.replace("::", "arclet.entari.builtins.") + while path in plugin_service._subplugined: + path = plugin_service._subplugined[path] + if path in plugin_service._apply: + if path in plugin_service.plugins: + return plugin_service.plugins[path] + log.plugin.trace(f"loaded rootless plugin {path!r}") + return plugin_service._apply[path][0](config) + if plug := find_plugin(path): + plugin_service._direct_plugins.add(plug.path) + return plug + try: + mod = import_plugin(path, config=config) + if not mod: + mod = next( + (import_plugin(_path, config=config) for _path in EntariConfig.instance._plugin_names.get(path, [])), + None, + ) + if not mod: + log.plugin.error(f"cannot found plugin {path!r}") + publish(PluginLoadedFailed(path)) + return + plugin_service._direct_plugins.add(mod.__name__) + _handle_dependents(mod.__plugin__, recursive_guard) + return mod.__plugin__ + except (ImportError, RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError): + return + except Exception as e: + log.plugin.exception(f"failed to load plugin {path!r}: {e}", exc_info=e) + publish(PluginLoadedFailed(path)) + return + + +def find_plugin(name: str) -> Plugin | None: + """根据插件名称查找插件""" + if name in plugin_service.plugins: + return plugin_service.plugins[name] + if name in EntariConfig.instance.plugin_prefixes: + for prefix in EntariConfig.instance.plugin_prefixes[name]: + if f"{prefix}{name}" in plugin_service.plugins: + return plugin_service.plugins[f"{prefix}{name}"] + if not name.count(".") and f"entari_plugin_{name}" in plugin_service.plugins: + return plugin_service.plugins[f"entari_plugin_{name}"] + + +def unload_plugin(plugin: str): + plugin = plugin.replace("::", "arclet.entari.builtins.") + while plugin in plugin_service._subplugined: + plugin = plugin_service._subplugined[plugin] + if not (_plugin := find_plugin(plugin)): + return False + _plugin.dispose() + return True diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 046e8b6..5e85330 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -241,6 +241,11 @@ class Plugin: _extra: dict[str, Any] = field(default_factory=dict, init=False) # extra metadata for inspection _apply: Callable[[Plugin], Any] | None = field(default=None, init=False, repr=False) + @property + def bindings(self) -> dict[str, tuple[str, str | None]]: + """本插件模块导入层绑定的 {名字: (目标模块, 属性)};attr 为 None 表示模块绑定""" + return plugin_service.bindings.get(self.id, {}) + @property def reusable(self) -> bool: return self.uid is not None and self.uid != "" @@ -278,6 +283,46 @@ def metadata(self, value: PluginMetadata): ): value.config.__doc__ = value.description or value.name + def __post_init__(self): + uid_index = self.id.rfind("@") + self.path = self.id[:uid_index] if uid_index != -1 else self.id + self.uid = self.id[uid_index + 1 :] if uid_index != -1 else None + if self.id in plugin_service.plugins and not self.id.startswith("."): + # build-then-swap 暂存:id 冲突时注册进 _staged,scope 用唯一 id 并置 disabled + self._scope = _make_scope(self).of(f"{self.id}@staging") + self._scope.disable() + plugin_service._staged[self.id] = self # type: ignore + else: + self._scope = _make_scope(self).of(self.id) + plugin_service.plugins[self.id] = self # type: ignore + self.effect = self._scope.effect + self._config_key = self.config.pop("$path", self.id) + if filter_expr := self.config.get("$filter", ""): + self._scope.propagators.append(FilterPropagator(filter_expr)) + # if self._metadata and self._metadata.depend_services: + # self._scope.propagators.append(inject(*self._metadata.depend_services, _is_global=True)) # type: ignore + # self._extra["injected_services"] = [ + # s.id if isinstance(s, type) else s for s in self._metadata.depend_services + # ] + if "$disable" in self.config and isinstance(self.config["$disable"], str): + + async def _check_reload(event: ConfigReload): + if event.scope == "basic": + self.check_disable() + + sub = on(ConfigReload, _check_reload) + self.collect(sub.dispose) + + self.is_static = self.config.pop("$static", False) + if self.id not in plugin_service._keep_values: + plugin_service._keep_values[self.id] = {} + if self.path not in plugin_service.referents: + plugin_service.referents[self.path] = set() + if self.path not in plugin_service.references: + plugin_service.references[self.path] = set() + plugin_service._unloaded.discard(self.id) + finalize(self, self.dispose, is_cleanup=True) + def exec_apply(self): if not self._apply: return @@ -395,40 +440,6 @@ def restore(self): """回收所有副作用""" return self._scope._effect_manager.dispose() - def __post_init__(self): - uid_index = self.id.rfind("@") - self.path = self.id[:uid_index] if uid_index != -1 else self.id - self.uid = self.id[uid_index + 1 :] if uid_index != -1 else None - self._scope = _make_scope(self).of(self.id) - self.effect = self._scope.effect - plugin_service.plugins[self.id] = self # type: ignore - self._config_key = self.config.pop("$path", self.id) - if filter_expr := self.config.get("$filter", ""): - self._scope.propagators.append(FilterPropagator(filter_expr)) - # if self._metadata and self._metadata.depend_services: - # self._scope.propagators.append(inject(*self._metadata.depend_services, _is_global=True)) # type: ignore - # self._extra["injected_services"] = [ - # s.id if isinstance(s, type) else s for s in self._metadata.depend_services - # ] - if "$disable" in self.config and isinstance(self.config["$disable"], str): - - async def _check_reload(event: ConfigReload): - if event.scope == "basic": - self.check_disable() - - sub = on(ConfigReload, _check_reload) - self.collect(sub.dispose) - - self.is_static = self.config.pop("$static", False) - if self.id not in plugin_service._keep_values: - plugin_service._keep_values[self.id] = {} - if self.path not in plugin_service.referents: - plugin_service.referents[self.path] = set() - if self.path not in plugin_service.references: - plugin_service.references[self.path] = set() - plugin_service._unloaded.discard(self.id) - finalize(self, self.dispose, is_cleanup=True) - def _clean_service(self): manager = it(Launart) @@ -464,6 +475,7 @@ def dispose(self, *, is_cleanup: bool = False): return if not self.id.startswith(".") and self.id not in plugin_service._subplugined: log.plugin.debug(f"disposing plugin {self.id}") + _was_staged = self.id in plugin_service._staged self._is_disposed = True tasks = set() t = self._clean_service() @@ -479,7 +491,10 @@ def dispose(self, *, is_cleanup: bool = False): log.plugin.trace(f"disposing sub-plugin {', '.join(subplugs)} of {self.id}") for subplug in self.subplugins: if subplug not in plugin_service.plugins: - plugin_service._subplugined.pop(subplug, None) + if subplug in plugin_service._staged: + tasks.update(plugin_service._staged[subplug].dispose(is_cleanup=is_cleanup)) + else: + plugin_service._subplugined.pop(subplug, None) continue try: tasks.update(plugin_service.plugins[subplug].dispose(is_cleanup=is_cleanup)) @@ -488,7 +503,7 @@ def dispose(self, *, is_cleanup: bool = False): log.plugin.error(f"failed to dispose sub-plugin {subplug} caused by {e!r}") plugin_service.plugins.pop(subplug, None) self.subplugins.clear() - if not is_cleanup: + if not is_cleanup and not _was_staged: publish(PluginUnloaded(self.id)) for ref in plugin_service.references.pop(self.path): if ref not in plugin_service.plugins: @@ -522,7 +537,10 @@ def dispose(self, *, is_cleanup: bool = False): tasks.update(plugin_service.plugins[ret].disable()) self._scope.dispose() self._scope.propagators.clear() - del plugin_service.plugins[self.id] + if self.id in plugin_service.plugins: + del plugin_service.plugins[self.id] + else: + plugin_service._staged.pop(self.id, None) del self.module del self._inspect return tasks diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index 5f79361..eed01b5 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -221,6 +221,8 @@ class Plugin: _extra: dict[str, Any] = field(default_factory=dict, init=False) # extra metadata for inspection _apply: Callable[[Plugin], Any] | None = field(default=None, init=False) + @property + def bindings(self) -> dict[str, tuple[str, str | None]]: ... @property def reusable(self) -> bool: ... @property diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 15b43d9..6d07204 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -75,6 +75,27 @@ def _ensure_plugin(names: list[str], sub: bool, pid: str, pname: str, prefix="") _IMPORTING.add(f"{prefix}{name}") +def _resolve_from_target(node: ast.ImportFrom, pname: str, is_init: bool) -> str | None: + """from-import 的目标模块全限定名(相对导入按当前模块解析)""" + if node.level == 0: + return node.module + parts = pname.split(".") + pkg = parts if is_init else parts[:-1] + if node.level > len(pkg) + 1: + return None + base = pkg if node.level == 1 else pkg[: 1 - node.level] + if node.module: + return ".".join([*base, node.module]) + return ".".join(base) if base else None + + +def _record_binding(pid: str, name: str, target: str, attr: str | None): + """记录名字级 import 绑定(name → (target, attr)),供重载侧查询与改写""" + if not target: + return + plugin_service.bindings.setdefault(pid, {})[name] = (target, attr) + + # fmt: off class _Visitor(ast.NodeVisitor): def __init__(self, pid: str, pname: str, path: bytes | str | PathLike[str], plg_lineno: list[int], sub_lineno: list[int], ns_lineno: list[int]): # noqa: E501 @@ -96,6 +117,8 @@ def visit_Import(self, node: ast.Import): if self._in_type_checking(): return + for alias in node.names: + _record_binding(self.pid, alias.asname or alias.name.split(".")[0], alias.name, None) if node.lineno in self.signed_plugin_lineno or all(x.name in _ENSURE_IS_PLUGIN for x in node.names): _ensure_plugin([alias.name for alias in node.names], False, self.pid, self.pname) elif node.lineno in self.signed_subplugin_lineno or all(x.name in _SUBMODULE_WAITLIST.get(self.pname, ()) for x in node.names): # noqa: E501 @@ -105,6 +128,15 @@ def visit_ImportFrom(self, node: ast.ImportFrom): name = self.pname if self._in_type_checking(): return + target = _resolve_from_target(node, name, self.path.endswith("__init__.py")) + if target: + for alias in node.names: + if alias.name == "*": + continue + if node.level == 1 and node.module is None: + _record_binding(self.pid, alias.asname or alias.name, f"{target}.{alias.name}", None) + else: + _record_binding(self.pid, alias.asname or alias.name, target, alias.name) if node.module is None: # from . import xxx _ensure_plugin([alias.name for alias in node.names], node.lineno not in self.signed_plugin_lineno, self.pid, name, f"{name}.") # noqa: E501 elif node.level == 0: # from xxx import xxx diff --git a/arclet/entari/plugin/service.py b/arclet/entari/plugin/service.py index 680f8ab..c591499 100644 --- a/arclet/entari/plugin/service.py +++ b/arclet/entari/plugin/service.py @@ -42,13 +42,26 @@ class PluginManagerService(Service): id = "entari.plugin.manager" plugins: dict[str, Plugin] - _keep_values: dict[str, dict[str, KeepingVariable]] + """插件字典,键为插件ID,值为插件对象""" referents: dict[str, set[str]] + """插件引用字典,键为插件ID,值为引用该插件的其他插件ID集合""" references: dict[str, set[str]] + """插件被引用字典,键为插件ID,值为该插件引用的其他插件ID集合""" + bindings: dict[str, dict[str, tuple[str, str | None]]] + """插件导入绑定字典,键为插件ID,值为该插件中导入的其他插件的绑定信息 {名字: (目标模块, 属性)}""" + fingerprints: dict[str, str] + """插件指纹字典,键为插件ID,值为该插件的指纹字符串""" + service_waiter: ServiceWaiters + _keep_values: dict[str, dict[str, KeepingVariable]] _direct_plugins: set[str] + """直接插件集合,存储所有直接加载(反过来即只由插件导入的插件)的插件ID""" _unloaded: set[str] + """卸载插件集合,存储所有已卸载的插件ID""" _subplugined: dict[str, str] + """子插件字典,键为子插件ID,值为父插件ID""" _apply: dict[str, tuple[Callable[[dict[str, Any]], RootlessPlugin], bool]] + _staged: dict[str, Plugin] + """插件暂存""" def __init__(self): super().__init__() @@ -60,6 +73,9 @@ def __init__(self): self._unloaded = set() self._subplugined = {} self._apply = {} + self.bindings = {} + self.fingerprints = {} + self._staged = {} self.service_waiter = ServiceWaiters() @property From bb7d5e297031e1a7151cb39fcafb4b5134c8c82b Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 10:56:08 +0800 Subject: [PATCH 12/33] :sparkles: support staged plugin registration and atomic reload --- arclet/entari/builtins/auto_reload.py | 60 ++++++++++++----------- arclet/entari/plugin/__init__.py | 1 + arclet/entari/plugin/dependent.py | 65 ++++++++++++++++++++++--- arclet/entari/plugin/model.py | 2 +- arclet/entari/plugin/module.py | 69 ++++++++++++++++++++------- 5 files changed, 146 insertions(+), 51 deletions(-) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index fa4cf80..9f7929a 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -18,7 +18,14 @@ from arclet.entari.config import BasicConfModel, EntariConfig, model_field from arclet.entari.event.config import ConfigReload from arclet.entari.logger import log -from arclet.entari.plugin import Plugin, PluginRole, find_plugin, find_plugin_by_file, unload_plugin_async +from arclet.entari.plugin import ( + Plugin, + PluginRole, + find_plugin, + find_plugin_by_file, + reload_plugin, + unload_plugin_async, +) from arclet.entari.plugin.swap import classify, swap_functions from arclet.entari.utils import escape_tag @@ -93,6 +100,10 @@ def _lock_for(self, plugin_id: str) -> asyncio.Lock: self._locks[plugin_id] = asyncio.Lock() return self._locks[plugin_id] + async def _reload(self, pid: str, cfg: dict) -> bool: + async with self._lock_for(pid): + return await reload_plugin(pid, cfg) + async def watch(self): async for event in awatch( *self.config.watch_dirs, debounce=self.config.debounce, step=self.config.step, watch_filter=PythonFilter() @@ -142,30 +153,25 @@ async def watch(self): logger.debug(f"Hot swap functions in {pid!r} failed, falling back to full reload.") _conf = plugin.config.copy() del plugin - async with self._lock_for(pid): - await unload_plugin_async(pid) - if plugin := load_plugin(pid, _conf): - logger.info(f"Reloaded {plugin.id!r}") - del plugin - self.fail.pop(file_path, None) - else: - logger.error(f"Failed to reload {pid!r}") - self.fail[file_path] = (pid, _conf) + if await self._reload(pid, _conf): + logger.info(f"Reloaded {pid!r}") + self.fail.pop(file_path, None) + else: + logger.error(f"Failed to reload {pid!r}") + self.fail[file_path] = (pid, _conf) pending.clear() for file_path in failed: if file_path not in self.fail: continue pid, _conf = self.fail[file_path] - async with self._lock_for(pid): - if file_path not in self.fail: - continue - logger.info(f"Detected change in {file_path!r} which failed to reload, retrying...") - if plugin := load_plugin(pid, _conf): - logger.info(f"Reloaded {plugin.id!r}") - del plugin - del self.fail[file_path] - else: - logger.error(f"Failed to reload {pid!r}") + if file_path not in self.fail: + continue + logger.info(f"Detected change in {file_path!r} which failed to reload, retrying...") + if await self._reload(pid, _conf): + logger.info(f"Reloaded {pid!r}") + del self.fail[file_path] + else: + logger.error(f"Failed to reload {pid!r}") async def watch_config(self): file = EntariConfig.instance.path.resolve() @@ -251,14 +257,12 @@ async def watch_config(self): _conf = plg.config.copy() async def _(): - async with self._lock_for(pid): - await unload_plugin_async(pid) - if plg := load_plugin(plugin_name, new_conf): - logger.info(f"Reloaded {plg.id!r}") - del plg - else: - logger.error(f"Failed to reload {plugin_name!r}") - self.fail[plugin_file] = (pid, _conf) + if await self._reload(pid, new_conf): + logger.info(f"Reloaded {pid!r}") + self.fail.pop(plugin_file, None) + else: + logger.error(f"Failed to reload {plugin_name!r}") + self.fail[plugin_file] = (pid, _conf) await asyncio.shield(_()) else: diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index f9b9d03..c40de30 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -20,6 +20,7 @@ from ..session import COMPONENTS, Session, component_transform from .dependent import find_plugin as find_plugin from .dependent import load_plugin as load_plugin +from .dependent import reload_plugin as reload_plugin from .dependent import unload_plugin as unload_plugin from .model import TS, Plugin, PluginDispatcher, current_plugin from .model import PluginMetadata as PluginMetadata diff --git a/arclet/entari/plugin/dependent.py b/arclet/entari/plugin/dependent.py index 0286078..5f799b4 100644 --- a/arclet/entari/plugin/dependent.py +++ b/arclet/entari/plugin/dependent.py @@ -1,4 +1,5 @@ import ast +import asyncio from typing import Any from arclet.letoderea import publish @@ -76,7 +77,8 @@ def _uses_module_level(plugin: Plugin, path: str) -> bool: def _rebind_imports(plugin: Plugin, path: str) -> bool: """将插件对 path(或其子树)的 import 绑定改写为新对象;False 表示有绑定无法满足(升级全量) - 模块绑定(attr 为 None)仅在名字等同于目标模块或其末组件时写回:`import a.b` 风格的名字是顶层包,其子模块链由 promote 的父属性写回维护,不在此处理。 + 模块绑定(attr 为 None)仅在名字等同于目标模块或其末组件时写回:`import a.b` 风格的名字是顶层包, + 其子模块链由 promote 的父属性写回维护,不在此处理。 """ module = plugin.module for name, (target, attr) in plugin.bindings.items(): @@ -217,7 +219,11 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): def load_plugin( - path: str, config: dict | None = None, recursive_guard: set[str] | None = None, prelude: bool = False + path: str, + config: dict | None = None, + recursive_guard: set[str] | None = None, + prelude: bool = False, + staged: bool = False, ) -> Plugin | None: """ 以导入路径方式加载模块 @@ -227,6 +233,7 @@ def load_plugin( config (dict): 模块配置 recursive_guard (set[str]): 递归保护 prelude (bool): 是否为前置插件 + staged (bool): 是否为暂存加载 """ if config is not None: config["$path"] = path @@ -250,14 +257,17 @@ def load_plugin( return plugin_service.plugins[path] log.plugin.trace(f"loaded rootless plugin {path!r}") return plugin_service._apply[path][0](config) - if plug := find_plugin(path): + if not staged and (plug := find_plugin(path)): plugin_service._direct_plugins.add(plug.path) return plug try: - mod = import_plugin(path, config=config) + mod = import_plugin(path, config=config, staged=staged) if not mod: mod = next( - (import_plugin(_path, config=config) for _path in EntariConfig.instance._plugin_names.get(path, [])), + ( + import_plugin(_path, config=config, staged=staged) + for _path in EntariConfig.instance._plugin_names.get(path, []) + ), None, ) if not mod: @@ -265,7 +275,8 @@ def load_plugin( publish(PluginLoadedFailed(path)) return plugin_service._direct_plugins.add(mod.__name__) - _handle_dependents(mod.__plugin__, recursive_guard) + if not staged: + _handle_dependents(mod.__plugin__, recursive_guard) return mod.__plugin__ except (ImportError, RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError): return @@ -288,6 +299,7 @@ def find_plugin(name: str) -> Plugin | None: def unload_plugin(plugin: str): + """卸载插件及其子插件""" plugin = plugin.replace("::", "arclet.entari.builtins.") while plugin in plugin_service._subplugined: plugin = plugin_service._subplugined[plugin] @@ -295,3 +307,44 @@ def unload_plugin(plugin: str): return False _plugin.dispose() return True + + +def promote_staged(plugin: Plugin): + """将暂存插件及其子插件正式注册进 plugin_service,并恢复启用状态与父绑定 + + 子插件按加载顺序已记入父插件 subplugins 列表; + 父属性写回(插件导入链)使模块中的属性指向新模块对象。 + """ + for sid in [plugin.id, *plugin.subplugins]: + if staged := plugin_service._staged.pop(sid, None): + plugin_service.plugins[sid] = staged + if sid != plugin.id: + plugin_service._subplugined[sid] = plugin.id + staged.check_disable() + plugin_service._unloaded.discard(plugin.id) + for sid in plugin.subplugins: + if sid not in plugin_service.plugins or sid not in plugin_service._subplugined: + continue + parent = plugin_service.plugins.get(plugin_service._subplugined[sid]) + if parent is not None: + parent.module.__dict__[sid.rpartition(".")[-1]] = plugin_service.plugins[sid].module + + +async def reload_plugin(path: str, conf: dict | None = None) -> bool: + """原子重载:导入失败时旧插件继续运行;成功后处理依赖方""" + path = path.replace("::", "arclet.entari.builtins.") + while path in plugin_service._subplugined: + path = plugin_service._subplugined[path] + if not (plugin := find_plugin(path)): + return False + if plugin.is_static: + return False + _conf = conf if conf is not None else plugin.config.copy() + if not (new_plugin := load_plugin(path, _conf, staged=True)): + log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") + return False + if tasks := plugin.dispose(): + await asyncio.wait(tasks) + promote_staged(new_plugin) + _handle_dependents(new_plugin) + return True diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 5e85330..90b4d98 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -288,7 +288,7 @@ def __post_init__(self): self.path = self.id[:uid_index] if uid_index != -1 else self.id self.uid = self.id[uid_index + 1 :] if uid_index != -1 else None if self.id in plugin_service.plugins and not self.id.startswith("."): - # build-then-swap 暂存:id 冲突时注册进 _staged,scope 用唯一 id 并置 disabled + # 原子重载暂存:id 冲突时注册进 _staged,scope 用唯一 id 并置 disabled self._scope = _make_scope(self).of(f"{self.id}@staging") self._scope.disable() plugin_service._staged[self.id] = self # type: ignore diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 6d07204..ceeec56 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -227,6 +227,7 @@ def __init__(self, fullname: str, path: str, plugin_id: str, parent_plugin_id: s self.plugin_id = plugin_id self.parent_plugin_id = parent_plugin_id self._inspect: PluginInspect = None # type: ignore + self.staged = False super().__init__(fullname, path) def get_code(self, fullname): @@ -278,6 +279,8 @@ def source_to_code(self, data, path="", *, _optimize: int = -1): ) def create_module(self, spec) -> ModuleType | None: + if self.staged: + return super().create_module(spec) if self.name in plugin_service.plugins: self.loaded = True return plugin_service.plugins[self.name].proxy() @@ -289,8 +292,14 @@ def create_module(self, spec) -> ModuleType | None: def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) -> None: is_sub = False - if plugin := plugin_service.plugins.get(self.parent_plugin_id) if self.parent_plugin_id else None: - plugin.subplugins.append(self.plugin_id) + plugin = ( + (plugin_service.plugins.get(self.parent_plugin_id) or plugin_service._staged.get(self.parent_plugin_id)) + if self.parent_plugin_id + else None + ) + if plugin: + if self.plugin_id not in plugin.subplugins: + plugin.subplugins.append(self.plugin_id) plugin_service._subplugined[self.plugin_id] = plugin.id is_sub = True if config is None or not {k: v for k, v in config.items() if k not in ("$path", "$static")}: @@ -441,15 +450,21 @@ def _path_find_spec(fullname, path=None, target=None) -> ModuleSpec | None: return spec -def _as_plugin(module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str) -> ModuleSpec: - module_spec.loader = PluginLoader(fullname, module_origin, plugin_id) +def _as_plugin( + module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str, staged: bool = False +) -> ModuleSpec: + loader = PluginLoader(fullname, module_origin, plugin_id) + loader.staged = staged + module_spec.loader = loader return module_spec def _as_submodule( - module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str, parent: str + module_spec: ModuleSpec, fullname: str, module_origin: str, plugin_id: str, parent: str, staged: bool = False ) -> ModuleSpec: - module_spec.loader = PluginLoader(fullname, module_origin, plugin_id, parent) + loader = PluginLoader(fullname, module_origin, plugin_id, parent) + loader.staged = staged + module_spec.loader = loader return module_spec @@ -467,6 +482,7 @@ def find_spec( target: ModuleType | None = None, origin_id_: str | None = None, force: bool = False, + staged: bool = False, ) -> ModuleSpec | None: # get the module spec using the default path-finder module_spec = _path_find_spec(fullname, path, target) @@ -504,11 +520,27 @@ def find_spec( ) or module_spec.name in _SUBMODULE_WAITLIST.get( # noqa: E501 parent_id, () ): - return _as_submodule(module_spec, fullname, module_origin, plugin_id, parent_id) + return _as_submodule( + module_spec, + fullname, + module_origin, + plugin_id, + parent_id, + staged=staged or plug.id in plugin_service._staged, + ) # noqa: E501 # 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: - return _as_plugin(module_spec, fullname, module_origin, plugin_id) + if module_spec.name in plugin_service._subplugined: + return _as_submodule( + module_spec, + fullname, + module_origin, + plugin_id, + plugin_service._subplugined[module_spec.name], + staged=staged, + ) # noqa: E501 + return _as_plugin(module_spec, fullname, module_origin, plugin_id, staged=staged) # 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 @@ -539,7 +571,7 @@ def find_spec( except (KeyError, ValueError): pass if marked: - _as_plugin(module_spec, fullname, module_origin, plugin_id) + _as_plugin(module_spec, fullname, module_origin, plugin_id, staged=staged) # if there already exists a plugin that is importing this module, # we should add the plugin as a referent of this module if plug: @@ -547,23 +579,28 @@ 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: - return _as_submodule( # noqa: E501 - module_spec, fullname, module_origin, plugin_id, plugin_service._subplugined[module_spec.name] - ) + return _as_submodule( + module_spec, + fullname, + module_origin, + plugin_id, + plugin_service._subplugined[module_spec.name], + staged=staged, + ) # noqa: E501 # 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. _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: - return _as_submodule(module_spec, fullname, module_origin, plugin_id, module_spec.parent) + return _as_submodule(module_spec, fullname, module_origin, plugin_id, module_spec.parent, staged=staged) # 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 _as_plugin(module_spec, fullname, module_origin, plugin_id, staged=staged) return -def import_plugin(id_, package=None, config: dict | None = None): +def import_plugin(id_, package=None, config: dict | None = None, staged: bool = False) -> ModuleType | 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 @@ -601,7 +638,7 @@ def import_plugin(id_, package=None, config: dict | None = None): parent_path = parent.__path__ else: parent_path = None - spec = _PluginFinder.find_spec(fullname, parent_path, origin_id_=id_, force=True) + spec = _PluginFinder.find_spec(fullname, parent_path, origin_id_=id_, force=True, staged=staged) if not spec: return mod = module_from_spec(spec) From 090c24d70da79f3801bc7fe2ece2eceefe471990 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 11:56:17 +0800 Subject: [PATCH 13/33] :sparkles: reload sub-plugins at sub-plugin granularity --- arclet/entari/builtins/auto_reload.py | 4 + arclet/entari/plugin/__init__.py | 9 +- .../entari/plugin/{dependent.py => loader.py} | 82 +++++++++++++++---- arclet/entari/plugin/module.py | 10 ++- example_plugin.py | 6 +- 5 files changed, 84 insertions(+), 27 deletions(-) rename arclet/entari/plugin/{dependent.py => loader.py} (80%) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index 9f7929a..24a40c9 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -23,7 +23,9 @@ PluginRole, find_plugin, find_plugin_by_file, + plugin_service, reload_plugin, + reload_subplugin, unload_plugin_async, ) from arclet.entari.plugin.swap import classify, swap_functions @@ -102,6 +104,8 @@ def _lock_for(self, plugin_id: str) -> asyncio.Lock: async def _reload(self, pid: str, cfg: dict) -> bool: async with self._lock_for(pid): + if pid in plugin_service._subplugined: + return await reload_subplugin(pid, cfg) return await reload_plugin(pid, cfg) async def watch(self): diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index c40de30..db6c173 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -18,10 +18,11 @@ from ..exceptions import StaticPluginDispatchError from ..message import Fragment, MessageChain, Render from ..session import COMPONENTS, Session, component_transform -from .dependent import find_plugin as find_plugin -from .dependent import load_plugin as load_plugin -from .dependent import reload_plugin as reload_plugin -from .dependent import unload_plugin as unload_plugin +from .loader import find_plugin as find_plugin +from .loader import load_plugin as load_plugin +from .loader import reload_plugin as reload_plugin +from .loader import reload_subplugin as reload_subplugin +from .loader import unload_plugin as unload_plugin from .model import TS, Plugin, PluginDispatcher, current_plugin from .model import PluginMetadata as PluginMetadata from .model import PluginRole as PluginRole diff --git a/arclet/entari/plugin/dependent.py b/arclet/entari/plugin/loader.py similarity index 80% rename from arclet/entari/plugin/dependent.py rename to arclet/entari/plugin/loader.py index 5f799b4..80b011f 100644 --- a/arclet/entari/plugin/dependent.py +++ b/arclet/entari/plugin/loader.py @@ -79,23 +79,37 @@ def _rebind_imports(plugin: Plugin, path: str) -> bool: 模块绑定(attr 为 None)仅在名字等同于目标模块或其末组件时写回:`import a.b` 风格的名字是顶层包, 其子模块链由 promote 的父属性写回维护,不在此处理。 + 绑定目标为 path 的父包、且父包绑定记录将该名字解析到 path 时, + 经父模块 getattr 取新值(父包在 sorted 依赖序中先于其子模块处理,故已重绑完成)。 """ module = plugin.module + parent_pkg = path.rpartition(".")[0] for name, (target, attr) in plugin.bindings.items(): - if not (target == path or target.startswith(path + ".")): - continue - if target not in plugin_service.plugins: - continue - new_module = plugin_service.plugins[target].module - if attr is None: - if name != target and name != target.rpartition(".")[-1]: + if target == path or target.startswith(path + "."): + if target not in plugin_service.plugins: continue - value: Any = new_module - else: - value = getattr(new_module, attr, _MISSING) - if value is _MISSING: - return False - module.__dict__[name] = value + new_module = plugin_service.plugins[target].module + if attr is None: + if name != target and name != target.rpartition(".")[-1]: + continue + value: Any = new_module + else: + value = getattr(new_module, attr, _MISSING) + if value is _MISSING: + return False + module.__dict__[name] = value + elif ( + parent_pkg + and target == parent_pkg + and parent_pkg in plugin_service.bindings + and parent_pkg in plugin_service.plugins + ): + parent_chain = plugin_service.bindings[parent_pkg] + if parent_chain.get(attr or name, (None, None))[0] == path: + value = getattr(plugin_service.plugins[parent_pkg].module, attr or name, _MISSING) + if value is _MISSING: + return False + module.__dict__[name] = value return True @@ -150,12 +164,15 @@ def dependents_of(path: str) -> list[str]: def _fingerprint_changed(plugin: Plugin) -> bool: - """比较插件公开面指纹(存储的旧指纹 vs 现算新指纹),并更新存储""" + """比较插件公开面指纹(存储的旧指纹 vs 现算新指纹),并更新存储 + + 旧指纹缺失(未存储过)时视为已变化:无法证明公开面未变 → 保守级联。 + """ path = plugin.path old = plugin_service.fingerprints.get(path) new = public_fingerprint(plugin) plugin_service.fingerprints[path] = new or "" - return bool(old) and old != new + return old is None or old != new def _rebind_dep(dep: Plugin, path: str): @@ -197,7 +214,7 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): for dep_id in sorted(dependents): if dep_id in recursive_guard: continue - if dep_id.startswith(path): + if dep_id.startswith(path + "."): continue if dep_id in plugin_service._subplugined and path.startswith(plugin_service._subplugined[dep_id]): continue @@ -321,6 +338,8 @@ def promote_staged(plugin: Plugin): if sid != plugin.id: plugin_service._subplugined[sid] = plugin.id staged.check_disable() + if plugin_service.status.blocking: + publish(Ready(), staged._scope) plugin_service._unloaded.discard(plugin.id) for sid in plugin.subplugins: if sid not in plugin_service.plugins or sid not in plugin_service._subplugined: @@ -340,6 +359,7 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: if plugin.is_static: return False _conf = conf if conf is not None else plugin.config.copy() + log.plugin.debug(f"staged loading {path!r}, old plugin keeps running until swap") if not (new_plugin := load_plugin(path, _conf, staged=True)): log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") return False @@ -348,3 +368,33 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: promote_staged(new_plugin) _handle_dependents(new_plugin) return True + + +async def reload_subplugin(path: str, conf: dict | None = None) -> bool: + """子插件粒度重载:仅替换子插件自身,父插件与兄弟经绑定索引重绑 + + 父插件在模块层使用该子插件名字(基类/模块级装饰器/顶层实例化)→ 回退整树重载; + 子插件自身经暂存机制原子替换(失败时旧子插件继续运行)。 + """ + path = path.replace("::", "arclet.entari.builtins.") + if path not in plugin_service.plugins or path not in plugin_service._subplugined: + return False + plugin = plugin_service.plugins[path] + parent_id = plugin_service._subplugined[path] + parent = plugin_service.plugins.get(parent_id) + if parent and _uses_module_level(parent, path): + log.plugin.debug(f"parent {parent_id!r} uses {path!r} at module level, full tree reload") + return await reload_plugin(parent_id, parent.config.copy()) + _conf = conf if conf is not None else plugin.config.copy() + log.plugin.debug( + f"build-then-swap: staged loading sub-plugin {path!r}, old sub-plugin keeps running until swap" + ) + if not (mod := import_plugin(path, config=_conf, staged=True)): + log.plugin.error(f"failed to load staged sub-plugin {path!r}, old sub-plugin keeps running") + return False + if tasks := plugin.dispose(): + await asyncio.wait(tasks) + new_plugin = mod.__plugin__ # type: ignore + promote_staged(new_plugin) + _handle_dependents(new_plugin) + return True diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index ceeec56..31e6a4c 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -401,16 +401,17 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) plugin._apply = getattr(module, "__plugin_apply__", None) plugin._inspect = self._inspect del self._inspect + staged = "staged " if self.staged else "" if not is_sub: if plugin._apply: - log.plugin.success(f"loaded plugin {self.plugin_id!r} partially applied") + log.plugin.success(f"{staged}loaded plugin {self.plugin_id!r} partially applied") else: - log.plugin.success(f"loaded plugin {self.plugin_id!r}") + log.plugin.success(f"{staged}loaded plugin {self.plugin_id!r}") else: - log.plugin.trace(f"loaded sub-plugin {plugin.id!r} of {self.parent_plugin_id!r}") + log.plugin.trace(f"{staged}loaded sub-plugin {plugin.id!r} of {self.parent_plugin_id!r}") if not plugin._apply: publish(PluginLoadedSuccess(self.plugin_id)) - if plugin_service.status.blocking: + if plugin_service.status.blocking and not self.staged: if plugin._apply: plugin.exec_apply() plugin.check_disable() @@ -504,6 +505,7 @@ def find_spec( return # current import statement is within a plugin. if plug := current_plugin.get(None): + _record_binding(plug.id, module_spec.name.split(".")[0], module_spec.name, 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_origin: diff --git a/example_plugin.py b/example_plugin.py index f51ddf8..7232de0 100644 --- a/example_plugin.py +++ b/example_plugin.py @@ -32,7 +32,7 @@ async def cleanup(): @plug.dispatch(MessageCreatedEvent) @filter_.public -async def _(session: Session): +async def invoke_test(session: Session): if session.content == "test": resp = await session.send("This message will recall in 5s...", at_sender=True) @@ -49,13 +49,13 @@ async def filter_content(session: Session): @disp_message.on().if_(filter_content) -async def _(session: Session): +async def filter_content1(session: Session): return await session.send("Filter: public message, to me, and content is 'aaa'") @disp_message @filter_.public & filter_.to_me & filter_(lambda sess: str(sess.content) != "aaa") -async def _(session: Session): +async def filter_content2(session: Session): return await session.send("Filter: public message, to me, but content is not 'aaa'") From 01ce828424d343bc97aae7060b4f29cb7a827a47 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 12:08:43 +0800 Subject: [PATCH 14/33] :bug: fix missing subtree match --- arclet/entari/plugin/loader.py | 22 ++++++++++++++++++++-- 1 file changed, 20 insertions(+), 2 deletions(-) diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index 80b011f..b61cc05 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -147,12 +147,15 @@ def public_fingerprint(plugin: Plugin) -> str | None: def dependents_of(path: str) -> list[str]: - """path 的依赖方插件:直接导入者 + 经父包再导出链一层(F5)""" + """path(及其子树)的依赖方插件:直接导入者 + 经父包再导出链一层(F5) + + 子树匹配使整树重载时,依赖子插件的下游插件同样被处理(子插件随树重建,绑定需重绑/级联)。 + """ parent_pkg = path.rpartition(".")[0] result: list[str] = [] for plug_id, bindings in plugin_service.bindings.items(): for name, (target, attr) in bindings.items(): - if target == path: + if target == path or target.startswith(path + "."): result.append(plug_id) break if parent_pkg and target == parent_pkg and parent_pkg in plugin_service.bindings: @@ -363,9 +366,24 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: if not (new_plugin := load_plugin(path, _conf, staged=True)): log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") return False + old_subplugins = list(plugin.subplugins) if tasks := plugin.dispose(): await asyncio.wait(tasks) promote_staged(new_plugin) + # 恢复未随 staged exec 重新导入的子插件(load_plugins/config 方式加载的模块在 + # 暂存加载期间 load_plugin 上溯命中旧插件早退,且不记入新插件 subplugins) + for sub_id in old_subplugins: + if sub_id in new_plugin.subplugins or sub_id in plugin_service.plugins: + continue + try: + mod = import_plugin(sub_id) + except Exception as e: + log.plugin.error(f"failed to restore sub-plugin {sub_id!r}: {e!r}") + continue + if mod is None: + log.plugin.error(f"cannot restore sub-plugin {sub_id!r}: module not found") + continue + log.plugin.debug(f"restored sub-plugin {sub_id!r} not re-imported by staged reload") _handle_dependents(new_plugin) return True From 071a5e2ba8d7f1a3bca327fb4ad32766f9f02b33 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 12:10:20 +0800 Subject: [PATCH 15/33] :zap: cascade dependents in topological order and parallelize service cleanup in plugin dispose --- arclet/entari/plugin/loader.py | 25 ++++++++++++++++++++++++- arclet/entari/plugin/model.py | 24 ++++++++++++++---------- 2 files changed, 38 insertions(+), 11 deletions(-) diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index b61cc05..952caf0 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -200,6 +200,29 @@ def _cascade_dep(dep_id: str, referent_set: set[str], recursive_guard: set[str]) recursive_guard.add(dep_id) +def _topo_dependents(dependents: set[str]) -> list[str]: + """依赖方按 references 图拓扑排序:被依赖者先于依赖者重载 + + 依赖者 C(`from B import x`)若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 + 拓扑序保证 B 先重载。 + """ + ordered: list[str] = [] + visited: set[str] = set() + + def visit(dep_id: str): + if dep_id in visited: + return + visited.add(dep_id) + for ref in plugin_service.references.get(dep_id, ()): + if ref in dependents: + visit(ref) + ordered.append(dep_id) + + for dep_id in sorted(dependents): + visit(dep_id) + return ordered + + def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): """插件完成(重)加载后处理依赖方:公开面未变 → 全部重绑;否则按模块层使用细粒度重绑或级联 @@ -214,7 +237,7 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): return surface_changed = _fingerprint_changed(plugin) referent_set = plugin_service.referents.setdefault(path, set()) - for dep_id in sorted(dependents): + for dep_id in _topo_dependents(dependents): if dep_id in recursive_guard: continue if dep_id.startswith(path + "."): diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 90b4d98..b53faa0 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -451,19 +451,23 @@ def _gen(service: Service): _services = [s for serv in self._services.values() for s in _gen(serv)] + async def _clean_one(service: Service): + if not manager.task_group: + return + plugin_service.service_waiter.clear(service.id) + if service.id not in manager.task_group.sideload_trackers: + return + try: + tracker = manager.task_group.sideload_trackers[service.id] + manager.remove_component(service) + await asyncio.wait([tracker, add_task(service.status.wait_for("finished"))]) + except (ValueError, KeyError): + pass + async def _clean(services: list[Service]): if not manager.task_group: return - for serv in services: - plugin_service.service_waiter.clear(serv.id) - if serv.id not in manager.task_group.sideload_trackers: - continue - try: - tracker = manager.task_group.sideload_trackers[serv.id] - manager.remove_component(serv) - await asyncio.wait([tracker, add_task(serv.status.wait_for("finished"))]) - except (ValueError, KeyError): - pass + await asyncio.gather(*(_clean_one(serv) for serv in services)) return add_task(_clean(_services)) From 19766905fbca4f85d0b7af4b25aebb2135581c83 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 11 Aug 2026 12:46:02 +0800 Subject: [PATCH 16/33] :sparkles: keep module-level mutable state across full reload --- arclet/entari/plugin/model.py | 13 ++++++++++--- arclet/entari/plugin/model.pyi | 18 ++++++++++++++++-- arclet/entari/plugin/module.py | 1 + 3 files changed, 27 insertions(+), 5 deletions(-) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index b53faa0..3b6e6c9 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -603,6 +603,12 @@ def service(self, serv: TS | type[TS]) -> TS: plugin_service.service_waiter.assign(serv.id) return serv + def restore_kept_state(self): + """重载后将保持的模块级可变对象重新绑定到模块 dict""" + for kept in plugin_service._keep_values.get(self.id, {}).values(): + if kept.module_attr: + self.module.__dict__[kept.module_attr] = kept.obj + class RootlessPlugin(Plugin): # fmt: off @@ -643,9 +649,10 @@ def validate(self, func): class KeepingVariable(Generic[T]): - def __init__(self, obj: T, dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None): + def __init__(self, obj: T, dispose=None, module_attr=None): self.obj = obj self._dispose = None + self.module_attr = module_attr if hasattr(self.obj, "dispose"): _dispose = self.obj.dispose.__func__ # type: ignore if _is_awaitable(_dispose): @@ -669,7 +676,7 @@ async def dispose(self): # fmt: off -def keeping(id_: str, obj: T | None = None, obj_factory: Callable[[], T] | None = None, dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None) -> T: # noqa: E501 +def keeping(id_: str, obj: T | None = None, obj_factory: Callable[[], T] | None = None, dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None, module_attr: str | None = None) -> T: # noqa: E501 # fmt: on if not (plug := current_plugin.get(None)): raise LookupError("no plugin context found") @@ -678,5 +685,5 @@ def keeping(id_: str, obj: T | None = None, obj_factory: Callable[[], T] | None raise ValueError("Either `obj` or `obj_factory` must be provided") _obj = obj_factory() if obj_factory else obj plug._extra.setdefault("kept_variables", []).append(id_) - plugin_service._keep_values[plug.id][id_] = KeepingVariable(cast(T, _obj), dispose) # type: ignore + plugin_service._keep_values[plug.id][id_] = KeepingVariable(cast(T, _obj), dispose, module_attr) # type: ignore return plugin_service._keep_values[plug.id][id_].obj # type: ignore diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index eed01b5..a18bf95 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -355,6 +355,7 @@ class Plugin: def proxy(self) -> ModuleType: ... def subproxy(self, sub_id: str) -> ModuleType: ... def service(self, serv: TS | type[TS]) -> TS: ... + def restore_kept_state(self) -> None: ... class RootlessPlugin(Plugin): @classmethod @@ -370,15 +371,28 @@ class RootlessPlugin(Plugin): class KeepingVariable(Generic[T]): obj: T _dispose: Callable[[T], Awaitable[None]] | None - def __init__(self, obj: T, dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None): ... + module_attr: str | None + def __init__( + self, + obj: T, + dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None, + module_attr: str | None = None, + ): ... async def dispose(self): ... @overload -def keeping(id_: str, obj: T, *, dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None) -> T: ... +def keeping( + id_: str, + obj: T, + *, + dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None, + module_attr: str | None = None, +) -> T: ... @overload def keeping( id_: str, *, obj_factory: Callable[[], T], dispose: Callable[[T], None] | Callable[[T], Awaitable[None]] | None = None, + module_attr: str | None = None, ) -> T: ... diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 31e6a4c..aeec55e 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -400,6 +400,7 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) plugin.metadata = metadata plugin._apply = getattr(module, "__plugin_apply__", None) plugin._inspect = self._inspect + plugin.restore_kept_state() del self._inspect staged = "staged " if self.staged else "" if not is_sub: From 73b2a969165948e9c6a611d60a7a491ac56e530d Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Wed, 12 Aug 2026 14:30:38 +0800 Subject: [PATCH 17/33] :sparkles: reload upstream common modules and their dependent plugins --- arclet/entari/builtins/auto_reload.py | 68 +++++++++++++++++++++++++++ 1 file changed, 68 insertions(+) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index 24a40c9..a6c3c2b 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -1,8 +1,11 @@ import ast import asyncio +import importlib +import sys from dataclasses import asdict from pathlib import Path from traceback import format_exception_only +from types import ModuleType from arclet.letoderea import post, publish from launart import Launart, Service, any_completed @@ -28,6 +31,7 @@ reload_subplugin, unload_plugin_async, ) +from arclet.entari.plugin.loader import dependents_of, topo_dependents from arclet.entari.plugin.swap import classify, swap_functions from arclet.entari.utils import escape_tag @@ -80,6 +84,29 @@ class Config(BasicConfModel): logger = log.wrapper("[AutoReload]").opt(colors=True) +def module_name_from_path(path: Path) -> str | None: + path = Path(path).resolve() + + for entry in map(Path, sys.path): + try: + relative = path.relative_to(entry.resolve()) + except ValueError: + continue + + if relative.suffix == ".py": + relative = relative.with_suffix("") + + parts = list(relative.parts) + + # __init__.py 对应包本身 + if parts[-1] == "__init__": + parts.pop() + + return ".".join(parts) + + return None + + class Watcher(Service): id = "entari.plugin.auto_reload/watcher" @@ -108,12 +135,45 @@ async def _reload(self, pid: str, cfg: dict) -> bool: return await reload_subplugin(pid, cfg) return await reload_plugin(pid, cfg) + async def _reload_upstream(self, module_name: str) -> list[str]: + """刷新非插件上游模块,并重载依赖它的插件 + + 必须先刷新 sys.modules 中的内容, 否则依赖插件重载时 import 链命中 sys.modules 缓存,拿到的仍是旧模块。 + + Returns: + 实际重载的插件 id 列表(供同批次去重)。 + """ + if module_name in plugin_service.plugins or module_name in plugin_service._subplugined: + return [] + mod = sys.modules.get(module_name) + if mod is None or not isinstance(mod, ModuleType): + return [] + dependents = [dep for dep in dependents_of(module_name) if dep in plugin_service.plugins] + if not dependents: + return [] + plugins = ", ".join(sorted(dependents)) + logger.debug(f"Reloading upstream module {module_name!r}, affected plugins: {plugins}") + try: + importlib.reload(mod) + except Exception as e: + logger.error(f"Failed to reload upstream module {module_name!r}: {e!r}") + return [] + reloaded: list[str] = [] + for dep_id in topo_dependents(set(dependents)): + async with self._lock_for(dep_id): + if await reload_plugin(dep_id): + reloaded.append(dep_id) + else: + logger.error(f"Failed to reload plugin {dep_id!r} after upstream module reload") + return reloaded + async def watch(self): async for event in awatch( *self.config.watch_dirs, debounce=self.config.debounce, step=self.config.step, watch_filter=PythonFilter() ): pending: dict[str, tuple[str, Plugin]] = {} failed: list[str] = [] + upstream: set[str] = set() for change in event: if plugin := find_plugin_by_file(change[1]): if plugin.is_static: @@ -122,7 +182,15 @@ async def watch(self): pending.setdefault(plugin.id, (change[1], plugin)) elif change[1] in self.fail: failed.append(change[1]) + elif module_name := module_name_from_path(Path(change[1])): + upstream.add(module_name) + reloaded: set[str] = set() + for module_name in upstream: + reloaded.update(await self._reload_upstream(module_name)) for pid, (file_path, plugin) in pending.items(): + if pid in reloaded: + self.fail.pop(file_path, None) + continue nodes: ast.Module | None = None if ( plugin._inspect From a2c8727471599b5cd70b6aed70f3870d39efb465 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Wed, 12 Aug 2026 14:34:24 +0800 Subject: [PATCH 18/33] :bug: fix problem in subplugin-reload --- arclet/entari/plugin/loader.py | 23 ++++++++++++++++------- arclet/entari/plugin/module.py | 4 +++- arclet/entari/plugin/swap.py | 2 ++ 3 files changed, 21 insertions(+), 8 deletions(-) diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index 952caf0..d3c98d2 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -1,5 +1,6 @@ import ast import asyncio +import sys from typing import Any from arclet.letoderea import publish @@ -25,6 +26,8 @@ def visit_expr(expr: ast.expr): for child in ast.iter_child_nodes(expr): if isinstance(child, ast.expr): visit_expr(child) + elif isinstance(child, ast.keyword) and child.value is not None: + visit_expr(child.value) def visit_stmt(stmt: ast.stmt): if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): @@ -55,6 +58,8 @@ def visit_stmt(stmt: ast.stmt): visit_expr(child) elif isinstance(child, ast.stmt): visit_stmt(child) + elif isinstance(child, ast.keyword) and child.value is not None: + visit_expr(child.value) for stmt in nodes.body: visit_stmt(stmt) @@ -100,6 +105,7 @@ def _rebind_imports(plugin: Plugin, path: str) -> bool: module.__dict__[name] = value elif ( parent_pkg + and attr is not None and target == parent_pkg and parent_pkg in plugin_service.bindings and parent_pkg in plugin_service.plugins @@ -150,15 +156,18 @@ def dependents_of(path: str) -> list[str]: """path(及其子树)的依赖方插件:直接导入者 + 经父包再导出链一层(F5) 子树匹配使整树重载时,依赖子插件的下游插件同样被处理(子插件随树重建,绑定需重绑/级联)。 + 插件(或其子模块)对自身子树的绑定属内部边,不构成依赖方,直接跳过。 """ parent_pkg = path.rpartition(".")[0] result: list[str] = [] for plug_id, bindings in plugin_service.bindings.items(): + if plug_id == path or plug_id.startswith(path + "."): + continue for name, (target, attr) in bindings.items(): if target == path or target.startswith(path + "."): result.append(plug_id) break - if parent_pkg and target == parent_pkg and parent_pkg in plugin_service.bindings: + if parent_pkg and attr is not None and target == parent_pkg and parent_pkg in plugin_service.bindings: parent_chain = plugin_service.bindings[parent_pkg] if parent_chain.get(attr or name, (None, None))[0] == path: result.append(plug_id) @@ -200,7 +209,7 @@ def _cascade_dep(dep_id: str, referent_set: set[str], recursive_guard: set[str]) recursive_guard.add(dep_id) -def _topo_dependents(dependents: set[str]) -> list[str]: +def topo_dependents(dependents: set[str]) -> list[str]: """依赖方按 references 图拓扑排序:被依赖者先于依赖者重载 依赖者 C(`from B import x`)若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 @@ -237,12 +246,10 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): return surface_changed = _fingerprint_changed(plugin) referent_set = plugin_service.referents.setdefault(path, set()) - for dep_id in _topo_dependents(dependents): + for dep_id in topo_dependents(dependents): if dep_id in recursive_guard: continue - if dep_id.startswith(path + "."): - continue - if dep_id in plugin_service._subplugined and path.startswith(plugin_service._subplugined[dep_id]): + if dep_id == path or dep_id.startswith(path + "."): continue dep = plugin_service.plugins.get(dep_id) if dep is None: @@ -385,6 +392,8 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: if plugin.is_static: return False _conf = conf if conf is not None else plugin.config.copy() + for name in plugin.subplugins: + sys.modules.pop(name, None) log.plugin.debug(f"staged loading {path!r}, old plugin keeps running until swap") if not (new_plugin := load_plugin(path, _conf, staged=True)): log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") @@ -428,7 +437,7 @@ async def reload_subplugin(path: str, conf: dict | None = None) -> bool: return await reload_plugin(parent_id, parent.config.copy()) _conf = conf if conf is not None else plugin.config.copy() log.plugin.debug( - f"build-then-swap: staged loading sub-plugin {path!r}, old sub-plugin keeps running until swap" + f"staged loading sub-plugin {path!r}, old sub-plugin keeps running until swap" ) if not (mod := import_plugin(path, config=_conf, staged=True)): log.plugin.error(f"failed to load staged sub-plugin {path!r}, old sub-plugin keeps running") diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index aeec55e..714bb9a 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -293,7 +293,9 @@ def create_module(self, spec) -> ModuleType | None: def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) -> None: is_sub = False plugin = ( - (plugin_service.plugins.get(self.parent_plugin_id) or plugin_service._staged.get(self.parent_plugin_id)) + # 暂存期间 plugins 中仍是旧顶插件,须优先取 _staged 中的新插件, + # 否则新子插件的父链接指向旧插件 + (plugin_service._staged.get(self.parent_plugin_id) or plugin_service.plugins.get(self.parent_plugin_id)) if self.parent_plugin_id else None ) diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py index 9357755..3d22e3b 100644 --- a/arclet/entari/plugin/swap.py +++ b/arclet/entari/plugin/swap.py @@ -197,6 +197,8 @@ def visit_expr(self, expr: ast.expr): for child in ast.iter_child_nodes(expr): if isinstance(child, ast.expr): self.visit_expr(child) + elif isinstance(child, ast.keyword) and child.value is not None: + self.visit_expr(child.value) def visit_stmt(self, stmt: ast.stmt): if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): From fc068bd29e329dbfacf31ed0a55b3756550d6df8 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Wed, 12 Aug 2026 16:47:27 +0800 Subject: [PATCH 19/33] :bug: fix more problem in subplugin-reload and adapt command delete in staged reload --- arclet/entari/command/__init__.py | 20 ++++- arclet/entari/command/plugin.py | 10 ++- arclet/entari/plugin/loader.py | 33 ++++++-- arclet/entari/plugin/model.py | 14 ++-- arclet/entari/plugin/model.pyi | 2 +- arclet/entari/plugin/module.py | 14 +++- entari.schema.json | 87 ++++++++++++++++++---- entari.yml | 4 +- example_plugins/example_plugin5/foo/bar.py | 4 +- 9 files changed, 148 insertions(+), 40 deletions(-) diff --git a/arclet/entari/command/__init__.py b/arclet/entari/command/__init__.py index bb4494b..5a3aa3f 100644 --- a/arclet/entari/command/__init__.py +++ b/arclet/entari/command/__init__.py @@ -166,7 +166,15 @@ def wrapper(func: Callable[..., TM]) -> Subscriber[TM]: self.subscribers[target.id] = target def _remove(_): - command_manager.delete(get_cmd(_)) + _cmd = get_cmd(_) + # if returned, it means the subscriber is already staged reload. + try: + record = command_manager._resolve(_cmd._hash) + except KeyError: + pass + else: + if id(_cmd) == id(record): + command_manager.delete(_cmd) self.trie[key].remove(target.id) # type: ignore if not self.trie[key]: self.trie.pop(key, None) # type: ignore @@ -203,7 +211,15 @@ def _remove(_): self.trie.setdefault(_key, []).append(target.id) def _remove(_): - command_manager.delete(get_cmd(_)) + _cmd = get_cmd(_) + # if returned, it means the subscriber is already staged reload. + try: + record = command_manager._resolve(_cmd._hash) + except KeyError: + pass + else: + if id(_cmd) == id(record): + command_manager.delete(_cmd) self.subscribers.pop(target.id, None) for _key in keys: self.trie[_key].remove(target.id) # type: ignore diff --git a/arclet/entari/command/plugin.py b/arclet/entari/command/plugin.py index 70c427b..155b75c 100644 --- a/arclet/entari/command/plugin.py +++ b/arclet/entari/command/plugin.py @@ -69,7 +69,15 @@ def __init__(self, plugin: Plugin, command: Alconna, need_reply_me: bool = False @plugin.collect def dispose(): - command_manager.delete(self.supplier.cmd) + _cmd = self.supplier.cmd + # if returned, it means the subscriber is already staged reload. + try: + record = command_manager._resolve(_cmd._hash) + except KeyError: + pass + else: + if id(_cmd) == id(record): + command_manager.delete(_cmd) del self.supplier.cmd del self.supplier diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index d3c98d2..16ab21a 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -10,7 +10,7 @@ from ..event.plugin import PluginLoadedFailed from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError from ..logger import log -from .model import Plugin +from .model import Plugin, current_plugin from .module import import_plugin from .service import plugin_service @@ -246,9 +246,14 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): return surface_changed = _fingerprint_changed(plugin) referent_set = plugin_service.referents.setdefault(path, set()) + current = current_plugin.get(None) for dep_id in topo_dependents(dependents): if dep_id in recursive_guard: continue + if current is not None and (dep_id == current.id or dep_id.startswith(current.id + ".")): + # 如果依赖方是当前正在加载的插件或其子插件,说明它们在同一 load_plugins 调用链中被导入, + # 且已在当前上下文中处理过依赖关系。 + continue if dep_id == path or dep_id.startswith(path + "."): continue dep = plugin_service.plugins.get(dep_id) @@ -382,6 +387,24 @@ def promote_staged(plugin: Plugin): parent.module.__dict__[sid.rpartition(".")[-1]] = plugin_service.plugins[sid].module +def _collect_subtree(plugin: Plugin) -> list[str]: + """DFS 收集插件树的全部子插件 id(父在前,去重)""" + ids: list[str] = [] + seen: set[str] = set() + + def walk(plug: Plugin): + for sid in plug.subplugins: + if sid in seen: + continue + seen.add(sid) + ids.append(sid) + if sub := plugin_service.plugins.get(sid): + walk(sub) + + walk(plugin) + return ids + + async def reload_plugin(path: str, conf: dict | None = None) -> bool: """原子重载:导入失败时旧插件继续运行;成功后处理依赖方""" path = path.replace("::", "arclet.entari.builtins.") @@ -392,13 +415,13 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: if plugin.is_static: return False _conf = conf if conf is not None else plugin.config.copy() - for name in plugin.subplugins: + old_subplugins = _collect_subtree(plugin) + for name in old_subplugins: sys.modules.pop(name, None) log.plugin.debug(f"staged loading {path!r}, old plugin keeps running until swap") if not (new_plugin := load_plugin(path, _conf, staged=True)): log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") return False - old_subplugins = list(plugin.subplugins) if tasks := plugin.dispose(): await asyncio.wait(tasks) promote_staged(new_plugin) @@ -436,9 +459,7 @@ async def reload_subplugin(path: str, conf: dict | None = None) -> bool: log.plugin.debug(f"parent {parent_id!r} uses {path!r} at module level, full tree reload") return await reload_plugin(parent_id, parent.config.copy()) _conf = conf if conf is not None else plugin.config.copy() - log.plugin.debug( - f"staged loading sub-plugin {path!r}, old sub-plugin keeps running until swap" - ) + log.plugin.debug(f"staged loading sub-plugin {path!r}, old sub-plugin keeps running until swap") if not (mod := import_plugin(path, config=_conf, staged=True)): log.plugin.error(f"failed to load staged sub-plugin {path!r}, old sub-plugin keeps running") return False diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 3b6e6c9..47b8b10 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -397,9 +397,9 @@ def disable(self): continue plugin_service.plugins[ret].disable() tasks = set() - t = self._clean_service() - t.add_done_callback(tasks.discard) - tasks.add(t) + if (t := self._clean_service()) is not None: + t.add_done_callback(tasks.discard) + tasks.add(t) self._scope.disable() if "$disable" not in self.config or isinstance(self.config["$disable"], bool): self.config["$disable"] = True @@ -450,6 +450,8 @@ def _gen(service: Service): yield service _services = [s for serv in self._services.values() for s in _gen(serv)] + if not _services: + return async def _clean_one(service: Service): if not manager.task_group: @@ -482,9 +484,9 @@ def dispose(self, *, is_cleanup: bool = False): _was_staged = self.id in plugin_service._staged self._is_disposed = True tasks = set() - t = self._clean_service() - t.add_done_callback(tasks.discard) - tasks.add(t) + if (t := self._clean_service()) is not None: + t.add_done_callback(tasks.discard) + tasks.add(t) self._services.clear() sys.modules.pop(self.module.__name__, None) tasks.update(self.restore()) diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index a18bf95..b3aafa8 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -249,7 +249,7 @@ class Plugin: def effect(self, execute: Callable[[], AsyncEffect], label: str = "") -> AsyncDisposable[Awaitable[None]]: ... def collect(self, *disposes: Disposable | AsyncDisposable) -> Self: ... def restore(self) -> set[asyncio.Task]: ... - def _clean_service(self) -> asyncio.Task: ... + def _clean_service(self) -> asyncio.Task | None: ... def dispose(self, *, is_cleanup: bool = False) -> set[asyncio.Task]: ... @overload def dispatch(self, event: type[Resultable[T]], name: str | None = None) -> PluginDispatcher[T]: ... diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 714bb9a..0ae0984 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -380,7 +380,7 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) if isinstance(e, (ImportError, StaticPluginDispatchError, ReusablePluginError)): raise e1 from None else: - raise ImportError(f"{e1!r} in {self.name!r}", name=self.name, path=self.path) from None + raise ImportError(f"{e1!r} in {self.name!r}", name=self.name, path=self.path) except Exception as e: log.plugin.exception(f"failed to load plugin {self.plugin_id!r} caused by {e!r}", exc_info=e) plugin.dispose() @@ -388,7 +388,7 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) if isinstance(e, (ImportError, StaticPluginDispatchError, ReusablePluginError)): raise else: - raise ImportError(f"{e!r} in {self.name!r}", name=self.name, path=self.path) from None + raise ImportError(f"{e!r} in {self.name!r}", name=self.name, path=self.path) finally: # leave plugin context delattr(module, "__cached__") @@ -508,7 +508,9 @@ def find_spec( return # current import statement is within a plugin. if plug := current_plugin.get(None): - _record_binding(plug.id, module_spec.name.split(".")[0], module_spec.name, None) + # only record outside import, as inside import are already recorded by the plugin's loader. + if module_spec.name != plug.module.__name__ and not module_spec.name.startswith(plug.module.__name__ + "."): + _record_binding(plug.id, module_spec.name.split(".")[0], module_spec.name, 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_origin: @@ -596,9 +598,13 @@ def find_spec( # it cannot be imported directly, otherwise it will break the uniqueness of the plugin instance. _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 + # we should still treat it as a submodule of the plugin to avoid breaking existing plugins. + # notice: cannot merge two conditions below, because some spec with submodule_search_locations (Namespace), + # their parent is the spec itself, not the parent module name. if module_spec.parent and module_spec.parent in plugin_service.plugins: return _as_submodule(module_spec, fullname, module_origin, plugin_id, module_spec.parent, staged=staged) + if (parent_name := module_spec.name.rpartition(".")[0]) and parent_name in plugin_service.plugins: + return _as_submodule(module_spec, fullname, module_origin, plugin_id, parent_name, staged=staged) # 6. force-wrap as a plugin when explicitly requested by import_plugin. if force: return _as_plugin(module_spec, fullname, module_origin, plugin_id, staged=staged) diff --git a/entari.schema.json b/entari.schema.json index 46c7d76..8872b08 100644 --- a/entari.schema.json +++ b/entari.schema.json @@ -77,6 +77,23 @@ "default": false, "description": "是否利用元数据进行插件导入检测(可能会增加启动时间)", "title": "Check Metadata" + }, + "str_as_message": { + "type": "boolean", + "default": true, + "description": "发送字符串时是否自动转换为消息链", + "title": "Str As Message" + }, + "superusers": { + "type": "object", + "additionalProperties": { + "type": "array", + "items": { + "type": "string" + } + }, + "description": "超级用户配置,键为平台名称,值为该平台的超级用户 ID 列表", + "title": "Superusers" } }, "additionalProperties": false, @@ -102,7 +119,14 @@ "title": "Host" }, "port": { - "type": "integer", + "oneOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ], "default": 5140, "description": "服务器端口", "title": "Port" @@ -156,7 +180,14 @@ "title": "Host" }, "port": { - "type": "integer", + "oneOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ], "default": 8080, "description": "本机 Webhook 服务器端口", "title": "Port" @@ -185,9 +216,16 @@ "title": "Server Host" }, "server_port": { - "type": "integer", + "oneOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ], "default": 5140, - "description": "发送请求的目标服务器端口", + "description": "发送请求的目标服务器端port", "title": "Server Port" }, "server_path": { @@ -306,14 +344,21 @@ "type": "string" }, "plugins": { - "description": "List of plugins under the prefix", - "items": { - "type": "string", - "description": "Plugin name" - }, - "title": "Plugins", - "type": "array", - "uniqueItems": true + "anyOf": [ + { + "type": "string" + }, + { + "items": { + "type": "string", + "description": "Plugin name" + }, + "type": "array", + "uniqueItems": true + } + ], + "description": "List of plugins under the prefix, or select an item of $files to apply plugins", + "title": "Plugins" } }, "required": [ @@ -425,6 +470,12 @@ "description": "是否记录发送的消息", "title": "Record Send" }, + "short_message": { + "type": "boolean", + "default": false, + "description": "是否在日志中使用简短的消息内容", + "title": "Short Message" + }, "$disable": { "type": "string", "description": "Expression for whether disable this plugin" @@ -443,9 +494,13 @@ }, ".scheduler": { "type": "object", - "description": "Simple Scheduler with interval / crontab task; no configuration required", - "additionalProperties": true, + "title": "_SchedulerConf", "properties": { + "debug": { + "type": "boolean", + "default": true, + "title": "Debug" + }, "$disable": { "type": "string", "description": "Expression for whether disable this plugin" @@ -458,7 +513,9 @@ "type": "string", "description": "Plugin filter expression, which will be evaluated in the context of the plugin" } - } + }, + "additionalProperties": false, + "description": "Simple Scheduler with interval / crontab task" }, "reusable@1": { "type": "object", diff --git a/entari.yml b/entari.yml index 5c21ebd..3836cc7 100644 --- a/entari.yml +++ b/entari.yml @@ -42,9 +42,7 @@ plugins: output: bar! example_plugin2: $filter: > - (platform in ("onebot", "milky") and member.roles[0].id ne 'member') - or - platform eq "console" + (platform in ("onebot", "milky") and member.roles[0].id ne 'member') or platform eq "console" example_plugin3: $filter: message == "test_plugin3" ?example_plugin4: {} diff --git a/example_plugins/example_plugin5/foo/bar.py b/example_plugins/example_plugin5/foo/bar.py index 23f5ad6..58e104f 100644 --- a/example_plugins/example_plugin5/foo/bar.py +++ b/example_plugins/example_plugin5/foo/bar.py @@ -6,6 +6,6 @@ metadata(__file__, description="A test plugin 5") -@command.on("exam5 ") -def exam5(x: int, y: int): +@command.on("exam5 {x} {y}") +async def exam5(x: int, y: int): return f"example_plugin5: {x} * {y} = {calc(x, y)}" From 7b876bce57be10fa571202fa7c399269592deb04 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 13 Aug 2026 09:35:01 +0800 Subject: [PATCH 20/33] :bug: fix dispose referent handle, avoid dangling-binding --- arclet/entari/builtins/auto_reload.py | 5 +-- arclet/entari/plugin/loader.py | 54 ++------------------------- arclet/entari/plugin/model.py | 18 +++++++-- arclet/entari/plugin/model.pyi | 2 +- arclet/entari/plugin/service.py | 53 ++++++++++++++++++++++++++ 5 files changed, 74 insertions(+), 58 deletions(-) diff --git a/arclet/entari/builtins/auto_reload.py b/arclet/entari/builtins/auto_reload.py index a6c3c2b..406043e 100644 --- a/arclet/entari/builtins/auto_reload.py +++ b/arclet/entari/builtins/auto_reload.py @@ -31,7 +31,6 @@ reload_subplugin, unload_plugin_async, ) -from arclet.entari.plugin.loader import dependents_of, topo_dependents from arclet.entari.plugin.swap import classify, swap_functions from arclet.entari.utils import escape_tag @@ -148,7 +147,7 @@ async def _reload_upstream(self, module_name: str) -> list[str]: mod = sys.modules.get(module_name) if mod is None or not isinstance(mod, ModuleType): return [] - dependents = [dep for dep in dependents_of(module_name) if dep in plugin_service.plugins] + dependents = plugin_service.dependents_of(module_name, ensure=True) if not dependents: return [] plugins = ", ".join(sorted(dependents)) @@ -159,7 +158,7 @@ async def _reload_upstream(self, module_name: str) -> list[str]: logger.error(f"Failed to reload upstream module {module_name!r}: {e!r}") return [] reloaded: list[str] = [] - for dep_id in topo_dependents(set(dependents)): + for dep_id in plugin_service.topo_dependents(set(dependents)): async with self._lock_for(dep_id): if await reload_plugin(dep_id): reloaded.append(dep_id) diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index 16ab21a..28736c8 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -152,29 +152,6 @@ def public_fingerprint(plugin: Plugin) -> str | None: return "\n".join(parts) -def dependents_of(path: str) -> list[str]: - """path(及其子树)的依赖方插件:直接导入者 + 经父包再导出链一层(F5) - - 子树匹配使整树重载时,依赖子插件的下游插件同样被处理(子插件随树重建,绑定需重绑/级联)。 - 插件(或其子模块)对自身子树的绑定属内部边,不构成依赖方,直接跳过。 - """ - parent_pkg = path.rpartition(".")[0] - result: list[str] = [] - for plug_id, bindings in plugin_service.bindings.items(): - if plug_id == path or plug_id.startswith(path + "."): - continue - for name, (target, attr) in bindings.items(): - if target == path or target.startswith(path + "."): - result.append(plug_id) - break - if parent_pkg and attr is not None and target == parent_pkg and parent_pkg in plugin_service.bindings: - parent_chain = plugin_service.bindings[parent_pkg] - if parent_chain.get(attr or name, (None, None))[0] == path: - result.append(plug_id) - break - return result - - def _fingerprint_changed(plugin: Plugin) -> bool: """比较插件公开面指纹(存储的旧指纹 vs 现算新指纹),并更新存储 @@ -209,29 +186,6 @@ def _cascade_dep(dep_id: str, referent_set: set[str], recursive_guard: set[str]) recursive_guard.add(dep_id) -def topo_dependents(dependents: set[str]) -> list[str]: - """依赖方按 references 图拓扑排序:被依赖者先于依赖者重载 - - 依赖者 C(`from B import x`)若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 - 拓扑序保证 B 先重载。 - """ - ordered: list[str] = [] - visited: set[str] = set() - - def visit(dep_id: str): - if dep_id in visited: - return - visited.add(dep_id) - for ref in plugin_service.references.get(dep_id, ()): - if ref in dependents: - visit(ref) - ordered.append(dep_id) - - for dep_id in sorted(dependents): - visit(dep_id) - return ordered - - def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): """插件完成(重)加载后处理依赖方:公开面未变 → 全部重绑;否则按模块层使用细粒度重绑或级联 @@ -241,13 +195,13 @@ def _handle_dependents(plugin: Plugin, recursive_guard: set[str] | None = None): if recursive_guard is None: recursive_guard = set() path = plugin.path - dependents = set(plugin_service.referents.get(path, ())) | set(dependents_of(path)) + dependents = set(plugin_service.referents.get(path, ())) | set(plugin_service.dependents_of(path, ensure=False)) if not dependents: return surface_changed = _fingerprint_changed(plugin) referent_set = plugin_service.referents.setdefault(path, set()) current = current_plugin.get(None) - for dep_id in topo_dependents(dependents): + for dep_id in plugin_service.topo_dependents(dependents): if dep_id in recursive_guard: continue if current is not None and (dep_id == current.id or dep_id.startswith(current.id + ".")): @@ -422,7 +376,7 @@ async def reload_plugin(path: str, conf: dict | None = None) -> bool: if not (new_plugin := load_plugin(path, _conf, staged=True)): log.plugin.error(f"failed to load staged plugin {path!r}, old plugin keeps running") return False - if tasks := plugin.dispose(): + if tasks := plugin.dispose(replacing=True): await asyncio.wait(tasks) promote_staged(new_plugin) # 恢复未随 staged exec 重新导入的子插件(load_plugins/config 方式加载的模块在 @@ -463,7 +417,7 @@ async def reload_subplugin(path: str, conf: dict | None = None) -> bool: if not (mod := import_plugin(path, config=_conf, staged=True)): log.plugin.error(f"failed to load staged sub-plugin {path!r}, old sub-plugin keeps running") return False - if tasks := plugin.dispose(): + if tasks := plugin.dispose(replacing=True): await asyncio.wait(tasks) new_plugin = mod.__plugin__ # type: ignore promote_staged(new_plugin) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 47b8b10..dc49206 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -473,7 +473,11 @@ async def _clean(services: list[Service]): return add_task(_clean(_services)) - def dispose(self, *, is_cleanup: bool = False): + def dispose(self, *, is_cleanup: bool = False, replacing: bool = False): + """拆卸插件 + + replacing=True 表示是重载插件下的卸载(reload_plugin/reload_subplugin) + """ if not is_cleanup and self.is_static: return # static plugin can only be disposed in cleanup phase plugin_service._unloaded.add(self.id) @@ -498,12 +502,14 @@ def dispose(self, *, is_cleanup: bool = False): for subplug in self.subplugins: if subplug not in plugin_service.plugins: if subplug in plugin_service._staged: - tasks.update(plugin_service._staged[subplug].dispose(is_cleanup=is_cleanup)) + tasks.update( + plugin_service._staged[subplug].dispose(is_cleanup=is_cleanup, replacing=replacing) + ) else: plugin_service._subplugined.pop(subplug, None) continue try: - tasks.update(plugin_service.plugins[subplug].dispose(is_cleanup=is_cleanup)) + tasks.update(plugin_service.plugins[subplug].dispose(is_cleanup=is_cleanup, replacing=replacing)) plugin_service._subplugined.pop(subplug, None) except Exception as e: log.plugin.error(f"failed to dispose sub-plugin {subplug} caused by {e!r}") @@ -511,6 +517,7 @@ def dispose(self, *, is_cleanup: bool = False): self.subplugins.clear() if not is_cleanup and not _was_staged: publish(PluginUnloaded(self.id)) + if not is_cleanup and not _was_staged and not replacing: for ref in plugin_service.references.pop(self.path): if ref not in plugin_service.plugins: continue @@ -531,7 +538,10 @@ def dispose(self, *, is_cleanup: bool = False): except Exception as e: log.plugin.error(f"failed to dispose referent plugin {ref} caused by {e!r}") plugin_service.plugins.pop(ref, None) - for ret in plugin_service.referents[self.path].copy(): + # bindings-only 依赖方在卸载时同样需要停用,否则 A 永久卸载后其持有僵尸绑定继续运行 + _dependents = set(plugin_service.referents[self.path]) + _dependents.update(plugin_service.dependents_of(self.path, ensure=True)) + for ret in _dependents: if ret not in plugin_service.plugins: continue if ( diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index b3aafa8..536feb9 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -250,7 +250,7 @@ class Plugin: def collect(self, *disposes: Disposable | AsyncDisposable) -> Self: ... def restore(self) -> set[asyncio.Task]: ... def _clean_service(self) -> asyncio.Task | None: ... - def dispose(self, *, is_cleanup: bool = False) -> set[asyncio.Task]: ... + def dispose(self, *, is_cleanup: bool = False, replacing: bool = False) -> set[asyncio.Task]: ... @overload def dispatch(self, event: type[Resultable[T]], name: str | None = None) -> PluginDispatcher[T]: ... @overload diff --git a/arclet/entari/plugin/service.py b/arclet/entari/plugin/service.py index c591499..c8760be 100644 --- a/arclet/entari/plugin/service.py +++ b/arclet/entari/plugin/service.py @@ -86,6 +86,59 @@ def required(self) -> set[str]: def stages(self) -> set[Phase]: return {"preparing", "cleanup", "blocking"} + def dependents_of(self, path: str, ensure: bool = True) -> list[str]: + """path(及其子树)的依赖方插件:直接导入者 + 经父包再导出链一层 + + 子树匹配使整树重载时,依赖子插件的下游插件同样被处理(子插件随树重建,绑定需重绑/级联)。 + 插件(或其子模块)对自身子树的绑定属内部边,不构成依赖方,直接跳过。 + + Args: + path (str): 插件ID或其子模块路径 + ensure (bool, optional): 是否确保返回的插件ID存在于已加载插件中. Defaults to True. + + Returns: + list[str]: 依赖方插件ID列表 + """ + parent_pkg = path.rpartition(".")[0] + result: list[str] = [] + for plug_id, bindings in self.bindings.items(): + if plug_id == path or plug_id.startswith(path + "."): + continue + for name, (target, attr) in bindings.items(): + if target == path or target.startswith(path + "."): + result.append(plug_id) + break + if parent_pkg and attr is not None and target == parent_pkg and parent_pkg in self.bindings: + parent_chain = self.bindings[parent_pkg] + if parent_chain.get(attr or name, (None, None))[0] == path: + result.append(plug_id) + break + if ensure: + result = [r for r in result if r in self.plugins] + return result + + def topo_dependents(self, dependents: set[str]) -> list[str]: + """依赖方按 references 图拓扑排序:上游(被依赖者)优先于下游(依赖者) + + 依赖者 C(`from B import x`)若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 + 拓扑序保证 B 先重载。 + """ + ordered: list[str] = [] + visited: set[str] = set() + + def visit(dep_id: str): + if dep_id in visited: + return + visited.add(dep_id) + for ref in self.references.get(dep_id, ()): + if ref in dependents: + visit(ref) + ordered.append(dep_id) + + for dep_id in sorted(dependents): + visit(dep_id) + return ordered + async def launch(self, manager: Launart): servs = [] From f0985c2e46023e61541ca4e85d78872a00d80530 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 13 Aug 2026 14:02:38 +0800 Subject: [PATCH 21/33] :bug: hot-swap do not replace old-fn object to new-fn --- arclet/entari/plugin/swap.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py index 3d22e3b..b57c5ac 100644 --- a/arclet/entari/plugin/swap.py +++ b/arclet/entari/plugin/swap.py @@ -574,8 +574,6 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function old_fn.__defaults__ = new_function.__defaults__ old_fn.__kwdefaults__ = new_function.__kwdefaults__ old_fn.__annotations__ = new_function.__annotations__ - if old_fn.__qualname__ != "_": - plugin.module.__dict__[old_fn.__qualname__] = new_function if change.signature_changed: for slot in plugin._scope.subscribers: sub = slot.subscriber From 93b9b94cefa1fb19405aac7868c653877e7ebdee Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 14 Aug 2026 17:07:20 +0800 Subject: [PATCH 22/33] :arrow_up: use letoderea 0.22.0 --- arclet/entari/plugin/model.py | 3 +-- arclet/entari/plugin/service.py | 2 +- arclet/entari/plugin/swap.py | 20 +++++++++++--------- example_plugins/example_plugin8.py | 18 +++++++----------- pdm.lock | 8 ++++---- pyproject.toml | 2 +- 6 files changed, 25 insertions(+), 28 deletions(-) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index dc49206..4f0888b 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -493,7 +493,6 @@ def dispose(self, *, is_cleanup: bool = False, replacing: bool = False): tasks.add(t) self._services.clear() sys.modules.pop(self.module.__name__, None) - tasks.update(self.restore()) delattr(self.module, "__plugin__") if self.subplugins: subplugs = [i.removeprefix(self.id)[1:] for i in self.subplugins] @@ -551,7 +550,7 @@ def dispose(self, *, is_cleanup: bool = False, replacing: bool = False): ): continue tasks.update(plugin_service.plugins[ret].disable()) - self._scope.dispose() + tasks.update(self._scope.dispose()) self._scope.propagators.clear() if self.id in plugin_service.plugins: del plugin_service.plugins[self.id] diff --git a/arclet/entari/plugin/service.py b/arclet/entari/plugin/service.py index c8760be..02bc593 100644 --- a/arclet/entari/plugin/service.py +++ b/arclet/entari/plugin/service.py @@ -120,7 +120,7 @@ def dependents_of(self, path: str, ensure: bool = True) -> list[str]: def topo_dependents(self, dependents: set[str]) -> list[str]: """依赖方按 references 图拓扑排序:上游(被依赖者)优先于下游(依赖者) - 依赖者 C(`from B import x`)若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 + 依赖者 C 若在 B 之前级联,会绑定旧 B,随后 B 重载时, C 已在 recursive_guard 中被跳过 → 静默。 拓扑序保证 B 先重载。 """ ordered: list[str] = [] diff --git a/arclet/entari/plugin/swap.py b/arclet/entari/plugin/swap.py index b57c5ac..4b22d99 100644 --- a/arclet/entari/plugin/swap.py +++ b/arclet/entari/plugin/swap.py @@ -574,15 +574,17 @@ def swap_functions(plugin: Plugin, new_nodes: ast.Module, changes: list[Function old_fn.__defaults__ = new_function.__defaults__ old_fn.__kwdefaults__ = new_function.__kwdefaults__ old_fn.__annotations__ = new_function.__annotations__ - if change.signature_changed: - for slot in plugin._scope.subscribers: - sub = slot.subscriber - if sub.callable_target is old_fn: - sub.callable_target = new_function - try: - sub._recompile() - except Exception as e: - log.plugin.error(f"failed to recompile subscriber of {change.qualname!r}: {e!r}") + for slot in plugin._scope.subscribers: + if slot.subscriber.callable_target is not old_fn: + continue + sub = slot.subscriber + sub.cancel_running() + if change.signature_changed: + sub.callable_target = new_function + try: + sub._recompile() + except Exception as e: + log.plugin.error(f"failed to recompile subscriber of {change.qualname!r}: {e!r}") for change in changes: if change.append: try: diff --git a/example_plugins/example_plugin8.py b/example_plugins/example_plugin8.py index 235022d..7ae902c 100644 --- a/example_plugins/example_plugin8.py +++ b/example_plugins/example_plugin8.py @@ -1,31 +1,27 @@ from arclet.entari.filter.message import startswith, regexmatch, regex_origin -from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image, Text +from arclet.entari import MessageCreatedEvent, MessageChain, Session, listen, Image, Text, User -@listen(MessageCreatedEvent) -@startswith("!hello") -async def hello_listener1(sess: Session, message: MessageChain): - await sess.send("Hello! This is a response from the hello_listener.") +@listen(MessageCreatedEvent).if_(startswith("!hello")) +async def hello_listener1(sess: Session, message: MessageChain, user: User): + await sess.send(f"Hello! This is a response from the hello_listener. {user}") await sess.send(message) -@listen(MessageCreatedEvent) -@startswith(Image, include=True) +@listen(MessageCreatedEvent).if_(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") +@listen(MessageCreatedEvent).if_(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) +@listen(MessageCreatedEvent).if_(regexmatch(r"test (\d+)", flags=2)) async def regex_listener( sess: Session, message: MessageChain, diff --git a/pdm.lock b/pdm.lock index 57651c3..eaae681 100644 --- a/pdm.lock +++ b/pdm.lock @@ -5,7 +5,7 @@ groups = ["default", "cli", "cron", "dev", "dotenv", "full", "msgspec", "pydantic", "reload", "rich", "toml", "yaml"] strategy = ["inherit_metadata"] lock_version = "4.5.0" -content_hash = "sha256:ff5971fff00cca2f21031effa0c543be4be9b16d453f356335f0dbdc7d66a559" +content_hash = "sha256:f182f588fcce2a6a56161db2f42695f7ec5769d67554dee2ebc73a2390b5bce5" [[metadata.targets]] requires_python = ">=3.10" @@ -221,7 +221,7 @@ files = [ [[package]] name = "arclet-letoderea" -version = "0.21.3" +version = "0.22.0" requires_python = "<4.0,>=3.10" summary = "A high-performance, simple-structured event system, relies on asyncio" groups = ["default"] @@ -230,8 +230,8 @@ dependencies = [ "typing-extensions>=4.12.0", ] files = [ - {file = "arclet_letoderea-0.21.3-py3-none-any.whl", hash = "sha256:ea1e235a037a193c734fbbe754969db9242d551290b11a0c0babae4e17d56d9d"}, - {file = "arclet_letoderea-0.21.3.tar.gz", hash = "sha256:4a2588d521afcd18d356804dc4289209be07b885e1d7851c8e37c8c03b954af4"}, + {file = "arclet_letoderea-0.22.0-py3-none-any.whl", hash = "sha256:7b0abe3e5d532cc4db8527f97bafeefb698fcb83d5094a41959eb3d0f5d7c847"}, + {file = "arclet_letoderea-0.22.0.tar.gz", hash = "sha256:eab31c29a42aa61265944ec6aab31f94fb7465df77b53408875869cd70c4009c"}, ] [[package]] diff --git a/pyproject.toml b/pyproject.toml index 16fe674..46e231a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,7 +6,7 @@ authors = [ {name = "RF-Tar-Railt",email = "rf_tar_railt@qq.com"}, ] dependencies = [ - "arclet-letoderea<0.22.0,>=0.21.3", + "arclet-letoderea<0.23.0,>=0.22.0", "arclet-alconna<2.0.0,>=1.8.44", "satori-python-core<1.4.0,>=1.3.5", "satori-python-client<1.4.0,>=1.3.2", From bfa96739dda21d125c315e90cfb37297bb26081b Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 14 Aug 2026 17:47:54 +0800 Subject: [PATCH 23/33] :bookmark: version 0.19.0rc1 --- arclet/entari/__init__.py | 2 +- pyproject.toml | 19 +++++++++++++++++++ 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index cbad7d8..aecc3e5 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -91,4 +91,4 @@ WS = WebsocketsInfo WH = WebhookInfo -__version__ = "0.18.5" +__version__ = "0.19.0rc1" diff --git a/pyproject.toml b/pyproject.toml index 46e231a..b7e053b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -32,6 +32,25 @@ classifiers = [ "Framework :: AsyncIO", "Operating System :: OS Independent", ] +keywords = [ + "asyncio", + "framework", + "instant-messaging", + "satori", + "entari", + "plugin-system", + "hmr", + "hot", + "module", + "replacement", + "reload", + "watch", + "development", + "hot-reload", + "hot-module-replacement", + "high-performance", + "event-driven", +] [project.urls] homepage = "https://arclet.top" From d03d80046a04c9e6f6e0da828c43dfdee9b63523 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 20 Aug 2026 16:31:58 +0800 Subject: [PATCH 24/33] :bug: fix namespace module import --- arclet/entari/plugin/module.py | 43 ++++++++++++++++++++-------------- 1 file changed, 26 insertions(+), 17 deletions(-) diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 0ae0984..031aed9 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -653,24 +653,33 @@ def import_plugin(id_, package=None, config: dict | None = None, staged: bool = if not spec: 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() + spec._initializing = True # type: ignore + try: + if spec.loader is None: + if spec.submodule_search_locations is None: + raise ImportError("missing loader", name=spec.name) + # A namespace package so do nothing. else: - spec.loader.exec_module(mod) + 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) + sys.modules[mod.__name__] = mod + finally: + spec._initializing = False # type: ignore return mod From c7b57c44e1b72d79ed1fb4846fb4726f18e08335 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 20 Aug 2026 17:10:52 +0800 Subject: [PATCH 25/33] :label: improve MessageChain typing --- arclet/entari/message.py | 56 ++++++++++++---------------------------- 1 file changed, 17 insertions(+), 39 deletions(-) diff --git a/arclet/entari/message.py b/arclet/entari/message.py index 7410a91..0793e60 100644 --- a/arclet/entari/message.py +++ b/arclet/entari/message.py @@ -34,22 +34,13 @@ class MessageChain(MutableSequence[TE]): def __init__(self): ... @overload - def __init__(self: MessageChain[Text], message: str): ... + def __init__(self, message: TE | Iterable[TE] | Sequence[TE]): ... @overload - def __init__(self, message: TE): ... + def __init__(self: MessageChain[TE1], message: TE1 | Iterable[TE1] | Sequence[TE1]): ... @overload - def __init__(self: MessageChain[TE1], message: TE1): ... - - @overload - def __init__(self, message: Iterable[TE]): ... - - @overload - def __init__(self: MessageChain[TE1], message: Iterable[TE1]): ... - - @overload - def __init__(self: MessageChain[Text], message: Iterable[str]): ... + def __init__(self: MessageChain[Text], message: str | Sequence[str]): ... @overload def __init__(self: MessageChain[Text | TE1], message: Iterable[str | TE1]): ... @@ -90,10 +81,10 @@ def __repr__(self) -> str: def __add__(self, other: str) -> MessageChain[TE | Text]: ... @overload - def __add__(self, other: TE | Iterable[TE]) -> MessageChain[TE]: ... + def __add__(self, other: TE | Iterable[TE] | Sequence[TE]) -> MessageChain[TE]: ... @overload - def __add__(self, other: TE1 | Iterable[TE1]) -> MessageChain[TE | TE1]: ... + def __add__(self, other: TE1 | Iterable[TE1] | Sequence[TE1]) -> MessageChain[TE | TE1]: ... def __add__(self, other: str | TE | TE1 | Iterable[TE | TE1]) -> MessageChain: """将另一个消息段或消息链添加到当前消息链. @@ -126,10 +117,10 @@ def __add__(self, other: str | TE | TE1 | Iterable[TE | TE1]) -> MessageChain: def __radd__(self, other: str) -> MessageChain[Text | TE]: ... @overload - def __radd__(self, other: TE | Iterable[TE]) -> MessageChain[TE]: ... + def __radd__(self, other: TE | Iterable[TE] | Sequence[TE]) -> MessageChain[TE]: ... @overload - def __radd__(self, other: TE1 | Iterable[TE1]) -> MessageChain[TE1 | TE]: ... + def __radd__(self, other: TE1 | Iterable[TE1] | Sequence[TE1]) -> MessageChain[TE1 | TE]: ... def __radd__(self, other: str | TE1 | Iterable[TE1]) -> MessageChain: result = MessageChain(other) @@ -483,7 +474,7 @@ def filter(self, predicate: Callable[[TE], bool]) -> MessageChain[TE]: """ return MessageChain(elem for elem in self.content if predicate(elem)) - def __iter__(self) -> Iterator[Element]: + def __iter__(self) -> Iterator[TE]: yield from self.content def __len__(self) -> int: @@ -607,11 +598,7 @@ def split(self, pattern: str = " ") -> list[Self]: tmp = [] return result - def replace( - self, - old: str, - new: str, - ) -> Self: + def replace(self, old: str, new: str) -> Self: """替换消息中有关的文本 Args: @@ -639,9 +626,9 @@ def startswith(self, string: str) -> bool: bool: 是否以给出的字符串开头 """ - if not self.content or not isinstance(self.content[0], Text): + if not self.content or not isinstance(text := self.content[0], Text): return False - return self.content[0].text.startswith(string) + return text.text.startswith(string) def endswith(self, string: str) -> bool: """判断消息链是否以给出的字符串结尾 @@ -653,11 +640,11 @@ def endswith(self, string: str) -> bool: bool: 是否以给出的字符串结尾 """ - if not self.content or not isinstance(self.content[-1], Text): + if not self.content or not isinstance(text := self.content[-1], Text): return False - return self.content[-1].text.endswith(string) + return text.text.endswith(string) - def append(self, element: Element | str) -> None: + def append(self, element: TE | str) -> None: """ 向消息链最后追加单个元素 @@ -667,19 +654,14 @@ def append(self, element: Element | str) -> None: Returns: None """ - if isinstance(element, str): - element = Text(element) - self.content.append(element) # type: ignore + self.content.append(Text(element) if isinstance(element, str) else 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: + def extend(self, values: Iterable[Self | TE | Sequence[TE | str]]) -> None: """ 向消息链最后添加元素/元素列表/消息链 @@ -831,11 +813,7 @@ def rstrip(self, *elements: str | type[Element] | Element, copy: bool = True) -> self.content.extend(content) return self - def replace_chain( - self, - old: MessageChain | list[Element], - new: MessageChain | list[Element], - ) -> Self: + def replace_chain(self, old: Sequence[Element], new: Sequence[Element]) -> Self: """替换消息链中的一部分. (在副本上操作) Args: From 2e05f3d9b9787b03272c22e5dbb73a154d197d07 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 20 Aug 2026 17:12:36 +0800 Subject: [PATCH 26/33] :bookmark: version 0.19.0rc2 --- arclet/entari/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index aecc3e5..88f2f97 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -91,4 +91,4 @@ WS = WebsocketsInfo WH = WebhookInfo -__version__ = "0.19.0rc1" +__version__ = "0.19.0rc2" From 7f96db3d016a25c8698d8d732541dd7daf7ddaca Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 4 Sep 2026 12:03:01 +0800 Subject: [PATCH 27/33] :sparkles: add config schema fragment registry --- arclet/entari/config/models/default.py | 2 +- arclet/entari/config/schema/__init__.py | 4 + .../{dc_schema.py => schema/dataclass.py} | 2 +- arclet/entari/config/schema/registry.py | 125 ++++++++++++++++++ 4 files changed, 131 insertions(+), 2 deletions(-) create mode 100644 arclet/entari/config/schema/__init__.py rename arclet/entari/config/{dc_schema.py => schema/dataclass.py} (99%) create mode 100644 arclet/entari/config/schema/registry.py diff --git a/arclet/entari/config/models/default.py b/arclet/entari/config/models/default.py index 6e196c0..04f96d6 100644 --- a/arclet/entari/config/models/default.py +++ b/arclet/entari/config/models/default.py @@ -14,7 +14,7 @@ from tarina.signature import merge_cls_and_parent_ns from ..action import ConfigModelAction -from ..dc_schema import SchemaGenerator +from ..schema import SchemaGenerator from ..util import store_field_description _available_dc_attrs = set(Signature.from_callable(dataclass).parameters.keys()) diff --git a/arclet/entari/config/schema/__init__.py b/arclet/entari/config/schema/__init__.py new file mode 100644 index 0000000..340a8f6 --- /dev/null +++ b/arclet/entari/config/schema/__init__.py @@ -0,0 +1,4 @@ +from .dataclass import SchemaGenerator as SchemaGenerator +from .registry import schema_fragment as schema_fragment +from .registry import purge_schema_fragments as purge_schema_fragments +from .registry import has_schema_fragments as has_schema_fragments diff --git a/arclet/entari/config/dc_schema.py b/arclet/entari/config/schema/dataclass.py similarity index 99% rename from arclet/entari/config/dc_schema.py rename to arclet/entari/config/schema/dataclass.py index 941db37..7d8c0c7 100644 --- a/arclet/entari/config/dc_schema.py +++ b/arclet/entari/config/schema/dataclass.py @@ -39,7 +39,7 @@ from tarina.generic import Unions, is_typed_dict -from .util import store_field_description +from ..util import store_field_description _MISSING = dataclasses.MISSING diff --git a/arclet/entari/config/schema/registry.py b/arclet/entari/config/schema/registry.py new file mode 100644 index 0000000..70f6617 --- /dev/null +++ b/arclet/entari/config/schema/registry.py @@ -0,0 +1,125 @@ +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + +_SENTINEL = object() + +Fragment = dict[str, Any] | type | Callable[[Any], Any] + + +@dataclass(frozen=True) +class _FragmentRecord: + """已注册的 schema 片段。 + + Attributes: + path: 规范化后的路径段(空元组 = 整个配置 schema) + fragment: dict / config 模型类型 / callable + replace: True 时目标节点整体替换(跳过深合并) + origin: 注册来源插件 id;None = 无插件上下文(仅 dispose 清除会用它) + """ + + path: tuple[str, ...] + """规范化后的路径段(空元组 = 整个配置 schema)""" + fragment: Fragment + """dict / config 模型类型 / callable""" + replace: bool + """True 时目标节点整体替换(跳过深合并)""" + origin: str | None + """注册来源插件 id;None = 无插件上下文(仅 dispose 清除会用它)""" + + +_schema_fragments: dict[str, list[_FragmentRecord]] = {} + + +def _parse_path(path: str | tuple[str, ...]) -> tuple[str, ...]: + """把注册 的 path 参数规范化为非空字符串段元组。 + + 字符串按 `.` 拆段;前导点/连续点/尾随点(空段)一律抛 ValueError + + dotted 子插件键以单元素元组 `(".a.b",)` 传入 + """ + if isinstance(path, str): + if not path: + return () + if path.startswith("."): + raise ValueError( + "dotted sub-plugin config keys must be passed as a single-element tuple, " + f"e.g. schema_fragment(({path!r},), ...); a string path must not start with '.'" + ) + parts = tuple(path.split(".")) + if any(not part for part in parts): + raise ValueError(f"invalid schema fragment path {path!r}: empty segment") + return parts + parts = tuple(path) + if any(not isinstance(part, str) or not part for part in parts): + raise ValueError(f"invalid schema fragment path {path!r}: segments must be non-empty strings") + return parts + + +def schema_fragment( + path: str | tuple[str, ...] = "", + fragment: dict | type | Callable[[Any], Any] | None = None, + *, + config_key: str | None = None, + replace: bool = False, +) -> None: + """注册一个配置 schema 片段(spec §4.1)。 + + Args: + path (str | tuple[str, ...]): 点分字符串或段元组;空字符串意味着作用于整个配置 schema。 + 以 `.` 开头的 dotted 子插件键必须用单元素元组,如 `(".a.b",)`。 + fragment: dict(深合并);config 模型类型(应用时才按最终 ref_root 生成); + 或 callable(接收目标节点当前值,返回 dict 深合并 / 非 dict 整体替换)。 + config_key (str, optional): 目标插件的配置键。为空时取 ``Plugin.current()._config_key``, + 无插件上下文且未传则抛 LookupError。 + replace (bool, optional): 目标节点是否整体替换为片段,不做深合并。默认 False。 + Raises: + TypeError: fragment 类型不合法 + LookupError: 无插件上下文且未传 config_key + """ + from arclet.entari.plugin import get_plugin # 延迟导入避免 config<->plugin 环 + + if fragment is None: + raise TypeError("schema_fragment requires a `fragment` (dict / config model type / callable)") + if not isinstance(fragment, dict) and not isinstance(fragment, type) and not callable(fragment): + raise TypeError(f"unsupported fragment type: {type(fragment).__name__}") + parts = _parse_path(path) + origin: str | None = None + try: + plugin = get_plugin(1) + except ValueError: + plugin = None + if config_key is None: + if plugin is None: + raise LookupError("no plugin context found; pass `config_key` explicitly") + config_key = plugin._config_key + origin = plugin.id + else: + if plugin is not None: + origin = plugin.id + records = _schema_fragments.setdefault(config_key, []) + for record in records: + if ( + record.origin == origin + and record.path == parts + and record.fragment == fragment + and record.replace == replace + ): + return + records.append(_FragmentRecord(parts, fragment, replace, origin)) + + +def purge_schema_fragments(plugin_id: str) -> None: + """清除某插件注册的全部 schema 片段""" + for key in [k for k, v in _schema_fragments.items() if any(r.origin == plugin_id for r in v)]: + remained = [r for r in _schema_fragments[key] if r.origin != plugin_id] + if remained: + _schema_fragments[key] = remained + else: + del _schema_fragments[key] + + +def has_schema_fragments(config_key: str) -> bool: + return bool(_schema_fragments.get(config_key)) From 3dd0b03bb552849faf52aefe01dcdc9b706bcd8a Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 4 Sep 2026 14:53:42 +0800 Subject: [PATCH 28/33] :sparkles: add schema fragment merge engine --- arclet/entari/config/action.py | 1 + arclet/entari/config/schema/rebase.py | 157 ++++++++++++++++++++++++ arclet/entari/config/schema/registry.py | 36 +++--- 3 files changed, 179 insertions(+), 15 deletions(-) create mode 100644 arclet/entari/config/schema/rebase.py diff --git a/arclet/entari/config/action.py b/arclet/entari/config/action.py index c5480c9..ee3de0d 100644 --- a/arclet/entari/config/action.py +++ b/arclet/entari/config/action.py @@ -135,6 +135,7 @@ def config_model_keys(obj: Any) -> list[str]: def config_model_schema(base: type[C], ref_root: str = "/") -> dict[str, Any]: for b in base.__mro__[-2::-1]: + _ensure_actions(b) if b in _config_model_actions: return _config_model_actions[b].schema(base, ref_root) return {field_.name: field_.type for field_ in fields(base)} # type: ignore diff --git a/arclet/entari/config/schema/rebase.py b/arclet/entari/config/schema/rebase.py new file mode 100644 index 0000000..2a8ff5b --- /dev/null +++ b/arclet/entari/config/schema/rebase.py @@ -0,0 +1,157 @@ +import copy + +from collections.abc import Callable +from typing import Any, TypeAlias + +from ..action import config_model_schema + +_SENTINEL = object() + +Fragment: TypeAlias = dict[str, Any] | type | Callable[[Any], Any] + + +def _deep_merge(target: dict, source: dict) -> None: + for key, value in source.items(): + if key in target and isinstance(target[key], dict) and isinstance(value, dict): + _deep_merge(target[key], value) + else: + target[key] = copy.deepcopy(value) + + +def _pointer_path(parts: tuple[str, ...]) -> str: + """拼接 JSON Pointer 风格的路径字符串(用于 schema 的 $ref)""" + segments: list[str] = [] + for part in parts: + if part.startswith("."): + segments.append("properties") + segments.append(part) + return "/" + "/".join(segments) if segments else "" + + +def _resolve_target( + schema: dict, parts: tuple[str, ...] +) -> tuple[dict | None, str | None, dict | None]: + """沿 parts 在 schema 内导航到目标节点,返回 (holder, key, target) + - target: 已存在的 dict 目标节点(可原地深合并); + - target 为 None 且 holder/key 非空:目标缺失或为非 dict,调用方应整体写入 holder[key]; + - holder 为 None:导航中途遇到已存在但非 dict 的节点,已告警,调用方跳过该片段。 + + Raises: + ValueError: 导航中途遇到已存在但非 dict 的节点 + """ + node: Any = schema + holder: dict | None = None + key: str | None = None + last = len(parts) - 1 + if not parts: + return schema, None, schema # 空 path:目标即 schema 根 + for index, part in enumerate(parts): + if part.startswith("."): + if not isinstance(node, dict): + raise ValueError(f"schema fragment path crosses non-dict node {part!r}, skipped") + props = node.setdefault("properties", {}) + if not isinstance(props, dict): + raise ValueError(f"schema fragment path crosses non-dict 'properties' node, skipped") + holder, key = props, part + value = props.get(part, _SENTINEL) + else: + if not isinstance(node, dict): + raise ValueError(f"schema fragment path crosses non-dict node {part!r}, skipped") + holder, key = node, part + value = node.get(part, _SENTINEL) + if index == last: + if isinstance(value, dict): + return holder, key, value + return holder, key, None + if isinstance(value, dict): + node = value + elif value is _SENTINEL: + node = holder[key] = {"type": "object"} + else: + raise ValueError(f"schema fragment path crosses existing non-dict node {part!r}, skipped") + return holder, key, None + + +def _rewrite_refs(node: Any, renamed: dict[str, str], model_prefix: str, ref_root: str) -> None: + """把片段内 $ref 就地重写为提升后的位置。 + + 匹配两种源形态(其余前缀视为指向片段外,保持原样并告警): + - `#{最终 ref_root}$defs/X`:模型片段生成时带完整前缀; + - `#/$defs/X`:手写 dict 片段相对片段根的形式。 + """ + if isinstance(node, dict): + ref = node.get("$ref") + if isinstance(ref, str): + name: str | None = None + if ref.startswith(model_prefix): + name = ref[len(model_prefix) :] + elif ref.startswith("#/$defs/"): + name = ref[len("#/$defs/") :] + if name is not None: + if name in renamed: + node["$ref"] = f"#{ref_root}$defs/{renamed[name]}" + else: + # log.warning(f"schema fragment $ref {ref!r} points to unknown local $defs {name!r}, kept as-is") + pass + for value in node.values(): + _rewrite_refs(value, renamed, model_prefix, ref_root) + elif isinstance(node, list): + for item in node: + _rewrite_refs(item, renamed, model_prefix, ref_root) + + +def _collect_defs_and_rebase( + content: dict, *, config_key: str, parts: tuple[str, ...], ref_root: str, final_ref_root: str +) -> dict[str, dict]: + """提取片段根级 $defs,唯一化重命名并就地重写 content 内 $ref。 + + 返回 {新名: defs schema 深拷贝}。 + """ + defs = content.pop("$defs", None) + if not isinstance(defs, dict) or not defs: + return {} + path_token = ".".join(parts).strip(".") or "root" + renamed: dict[str, str] = {} + hoisted: dict[str, dict] = {} + for name, def_schema in defs.items(): + new_name = f"{config_key}.{path_token}.{name}" + renamed[name] = new_name + hoisted[new_name] = copy.deepcopy(def_schema) + _rewrite_refs(content, renamed, f"#{final_ref_root}$defs/", ref_root) + return hoisted + + +def update_schema(schema: dict, config_key: str, ref_root: str, parts: tuple[str, ...], payload: Fragment, replace: bool = False) -> None: + final_ref_root = f"{ref_root}{_pointer_path(parts)}/" if parts else ref_root + if isinstance(payload, type): + # 模型片段:以最终 ref_root 生成,再走统一的 defs 处理 + try: + payload = config_model_schema(payload, ref_root=final_ref_root) + except Exception as e: + raise RuntimeError(f"failed to resolve config model fragment {payload!r}: {e!r}, skipped") from e + holder, key, target = _resolve_target(schema, parts) + if holder is None: + return + if callable(payload) and not isinstance(payload, type): + try: + current = copy.deepcopy(target) if target is not None else None + payload = payload(current) + except Exception as e: + raise RuntimeError(f"schema fragment callable raised {e!r}, skipped") from e + if not isinstance(payload, dict): + holder[key] = copy.deepcopy(payload) + return + content = copy.deepcopy(payload) # 绝不原地改写注册的片段:热重载后再次应用时内容必须仍是原样 + hoisted = _collect_defs_and_rebase( + content, config_key=config_key, parts=parts, ref_root=ref_root, final_ref_root=final_ref_root + ) # 先剥离 content 的 $defs 并重写 ref,再合并,避免原样 defs 先落入目标 + if target is not None and not replace: + _deep_merge(target, content) + elif target is not None: + target.clear() # replace 空 path:整体替换基础 schema + _deep_merge(target, content) + else: + holder[key] = content + # replace 空 path 会清空 schema 根,不能被 clear 冲掉 + if hoisted: + schema.setdefault("$defs", {}).update(hoisted) diff --git a/arclet/entari/config/schema/registry.py b/arclet/entari/config/schema/registry.py index 70f6617..802ed7e 100644 --- a/arclet/entari/config/schema/registry.py +++ b/arclet/entari/config/schema/registry.py @@ -4,21 +4,12 @@ from dataclasses import dataclass from typing import Any -_SENTINEL = object() - -Fragment = dict[str, Any] | type | Callable[[Any], Any] +from .rebase import Fragment, update_schema @dataclass(frozen=True) class _FragmentRecord: - """已注册的 schema 片段。 - - Attributes: - path: 规范化后的路径段(空元组 = 整个配置 schema) - fragment: dict / config 模型类型 / callable - replace: True 时目标节点整体替换(跳过深合并) - origin: 注册来源插件 id;None = 无插件上下文(仅 dispose 清除会用它) - """ + """已注册的 schema 片段""" path: tuple[str, ...] """规范化后的路径段(空元组 = 整个配置 schema)""" @@ -44,10 +35,7 @@ def _parse_path(path: str | tuple[str, ...]) -> tuple[str, ...]: if not path: return () if path.startswith("."): - raise ValueError( - "dotted sub-plugin config keys must be passed as a single-element tuple, " - f"e.g. schema_fragment(({path!r},), ...); a string path must not start with '.'" - ) + raise ValueError("sub-plugin config keys must be passed as a single-element tuple") parts = tuple(path.split(".")) if any(not part for part in parts): raise ValueError(f"invalid schema fragment path {path!r}: empty segment") @@ -123,3 +111,21 @@ def purge_schema_fragments(plugin_id: str) -> None: def has_schema_fragments(config_key: str) -> bool: return bool(_schema_fragments.get(config_key)) + + +def apply_schema_fragments(schema: dict, config_key: str, ref_root: str = "/") -> dict: + """按注册顺序把 config_key 的片段应用到 schema(原地合并)并返回 schema。 + + 单轮应用内对等价片段去重:跨 origin 的重复只生效一次。 + """ + records = _schema_fragments.get(config_key) + if not records: + return schema + applied: list[tuple[tuple[str, ...], bool, Fragment]] = [] + for record in records: + mark = (record.path, record.replace, record.fragment) + if any(mark == old for old in applied): + continue + applied.append(mark) + update_schema(schema, config_key, ref_root, record.path, record.fragment, record.replace) + return schema From ed5d83e6533f59f0b9d23470fc3699640a67ae8f Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 4 Sep 2026 15:19:32 +0800 Subject: [PATCH 29/33] :sparkles: Plugin.register_schema and apply schema fragments in Plugin.config_schema --- arclet/entari/config/schema/__init__.py | 3 ++- arclet/entari/config/schema/registry.py | 33 +++++++---------------- arclet/entari/plugin/__init__.py | 12 ++++++++- arclet/entari/plugin/model.py | 35 ++++++++++++++++++++----- arclet/entari/plugin/model.pyi | 4 ++- 5 files changed, 54 insertions(+), 33 deletions(-) diff --git a/arclet/entari/config/schema/__init__.py b/arclet/entari/config/schema/__init__.py index 340a8f6..47e9c7a 100644 --- a/arclet/entari/config/schema/__init__.py +++ b/arclet/entari/config/schema/__init__.py @@ -1,4 +1,5 @@ from .dataclass import SchemaGenerator as SchemaGenerator -from .registry import schema_fragment as schema_fragment +from .registry import add_schema_fragment as add_schema_fragment from .registry import purge_schema_fragments as purge_schema_fragments from .registry import has_schema_fragments as has_schema_fragments +from .registry import apply_schema_fragments as apply_schema_fragments diff --git a/arclet/entari/config/schema/registry.py b/arclet/entari/config/schema/registry.py index 802ed7e..d568a0a 100644 --- a/arclet/entari/config/schema/registry.py +++ b/arclet/entari/config/schema/registry.py @@ -29,7 +29,7 @@ def _parse_path(path: str | tuple[str, ...]) -> tuple[str, ...]: 字符串按 `.` 拆段;前导点/连续点/尾随点(空段)一律抛 ValueError - dotted 子插件键以单元素元组 `(".a.b",)` 传入 + 子插件键以单元素元组 `(".a.b",)` 传入 """ if isinstance(path, str): if not path: @@ -46,53 +46,38 @@ def _parse_path(path: str | tuple[str, ...]) -> tuple[str, ...]: return parts -def schema_fragment( +def add_schema_fragment( + config_key: str, path: str | tuple[str, ...] = "", fragment: dict | type | Callable[[Any], Any] | None = None, - *, - config_key: str | None = None, replace: bool = False, + origin: str | None = None, ) -> None: - """注册一个配置 schema 片段(spec §4.1)。 + """注册一个配置 schema 片段。 Args: + config_key (str): 目标插件的配置键。 path (str | tuple[str, ...]): 点分字符串或段元组;空字符串意味着作用于整个配置 schema。 - 以 `.` 开头的 dotted 子插件键必须用单元素元组,如 `(".a.b",)`。 + 以 `.` 开头的子插件键必须用单元素元组,如 `(".a.b",)`。 fragment: dict(深合并);config 模型类型(应用时才按最终 ref_root 生成); 或 callable(接收目标节点当前值,返回 dict 深合并 / 非 dict 整体替换)。 - config_key (str, optional): 目标插件的配置键。为空时取 ``Plugin.current()._config_key``, - 无插件上下文且未传则抛 LookupError。 replace (bool, optional): 目标节点是否整体替换为片段,不做深合并。默认 False。 + origin (str, optional): 注册来源插件 id;空表示无插件上下文(仅 purge 清除会用它)。默认 None。 Raises: TypeError: fragment 类型不合法 - LookupError: 无插件上下文且未传 config_key """ - from arclet.entari.plugin import get_plugin # 延迟导入避免 config<->plugin 环 if fragment is None: raise TypeError("schema_fragment requires a `fragment` (dict / config model type / callable)") if not isinstance(fragment, dict) and not isinstance(fragment, type) and not callable(fragment): raise TypeError(f"unsupported fragment type: {type(fragment).__name__}") parts = _parse_path(path) - origin: str | None = None - try: - plugin = get_plugin(1) - except ValueError: - plugin = None - if config_key is None: - if plugin is None: - raise LookupError("no plugin context found; pass `config_key` explicitly") - config_key = plugin._config_key - origin = plugin.id - else: - if plugin is not None: - origin = plugin.id records = _schema_fragments.setdefault(config_key, []) for record in records: if ( record.origin == origin and record.path == parts - and record.fragment == fragment + and record.fragment is fragment and record.replace == replace ): return diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index db6c173..b8f7258 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -5,7 +5,7 @@ import sys from collections.abc import Awaitable from pathlib import Path -from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload +from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload, Callable from arclet.letoderea import Subscriber, on from arclet.letoderea.effect import AsyncDisposable, Disposable @@ -268,6 +268,16 @@ def _reload(event: ConfigReload): get_config = plugin_config +def register_schema(path: str | tuple[str, ...] = "", fragment: dict | type | Callable[[dict], dict] | None = None, *, replace: bool = False) -> None: + plugin = get_plugin(1) + plugin.register_schema(path, fragment, replace=replace) + + +def get_plugin_schema() -> dict: + plugin = get_plugin(1) + return plugin.config_schema() + + def load_plugins(dir_: str | os.PathLike | Path): """加载指定目录下的所有插件""" path = dir_ if isinstance(dir_, Path) else Path(dir_) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 4f0888b..29d9533 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -8,6 +8,7 @@ from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from enum import Enum +from itertools import chain from pathlib import Path from types import ModuleType from typing import Any, Generic, TypeVar, cast @@ -39,6 +40,7 @@ from tarina.tools import TCallable, run_sync from ..config import config_model_schema +from ..config.schema import purge_schema_fragments, add_schema_fragment, apply_schema_fragments, has_schema_fragments from ..event.config import ConfigReload from ..event.plugin import PluginLoadedFailed, PluginLoadedSuccess, PluginUnloaded from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError @@ -174,12 +176,6 @@ class PluginMetadata: config: Any | None = None """插件配置模型""" - def get_config_schema(self) -> dict[str, Any]: - """获取插件配置模型的 JSON Schema""" - if self.config is None: - return {} - return config_model_schema(self.config) - def inject(*services: type[Service] | str | dict, _is_global: bool = False): @@ -323,6 +319,32 @@ async def _check_reload(event: ConfigReload): plugin_service._unloaded.discard(self.id) finalize(self, self.dispose, is_cleanup=True) + def register_schema( + self, + path: str | tuple[str, ...] = "", + fragment: dict | type | Callable[[dict], dict] | None = None, + *, + replace: bool = False, + ) -> None: + """注册作用于本插件配置的 schema 片段""" + add_schema_fragment(self._config_key, path, fragment, replace, self.id) + + def config_schema(self, config_key: str | None = None) -> dict[str, Any]: + """获取插件配置模型的 JSON Schema + + Args: + config_key (str, optional): 本插件的配置键。 + """ + metadata = self._metadata + if metadata is None or metadata.config is None: + if not config_key or not has_schema_fragments(config_key): + return {} + return apply_schema_fragments({"type": "object", "additionalProperties": True}, config_key, ref_root="/") + schema = config_model_schema(metadata.config) + if config_key: + return apply_schema_fragments(schema, config_key, ref_root="/") + return schema + def exec_apply(self): if not self._apply: return @@ -483,6 +505,7 @@ def dispose(self, *, is_cleanup: bool = False, replacing: bool = False): plugin_service._unloaded.add(self.id) if self._is_disposed: return + purge_schema_fragments(self.id) if not self.id.startswith(".") and self.id not in plugin_service._subplugined: log.plugin.debug(f"disposing plugin {self.id}") _was_staged = self.id in plugin_service._staged diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index 536feb9..75a28f3 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -191,7 +191,7 @@ class PluginMetadata: depend_services: list[type[Service] | str | DependService] = ... config: Any | None = ... - def get_config_schema(self) -> dict[str, Any]: ... + def get_config_schema(self, config_key: str | None = None) -> dict[str, Any]: ... @overload def inject(*services: type[Service] | str | DependService) -> Callable[[TCallable], TCallable]: ... @@ -233,6 +233,8 @@ class Plugin: def metadata(self) -> PluginMetadata | None: ... @metadata.setter def metadata(self, value: PluginMetadata): ... + def register_schema(self, path: str | tuple[str, ...] = "", fragment: dict | type | Callable[[dict], dict] | None = None, *, replace: bool = False) -> None: ... + def config_schema(self, config_key: str | None = None) -> dict[str, Any]: ... def exec_apply(self) -> None: ... @property def is_available(self) -> bool: ... From 502be180cc67177159841baf7f14c66214c547cf Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Fri, 4 Sep 2026 17:02:39 +0800 Subject: [PATCH 30/33] :sparkles: apply fragment pipeline in EntariConfig.generate_schema --- arclet/entari/__init__.py | 1 + arclet/entari/config/file.py | 37 +++------ arclet/entari/config/schema/__init__.py | 5 +- arclet/entari/config/schema/plugin.py | 91 +++++++++++++++++++++ arclet/entari/config/schema/rebase.py | 11 ++- arclet/entari/config/schema/registry.py | 7 +- arclet/entari/plugin/__init__.py | 11 ++- arclet/entari/plugin/model.py | 17 +--- arclet/entari/plugin/model.pyi | 8 +- example_plugins/example_plugin7/__init__.py | 34 ++++++++ example_plugins/example_plugin7/config.py | 7 ++ example_plugins/example_plugin7/listener.py | 15 ++++ 12 files changed, 192 insertions(+), 52 deletions(-) create mode 100644 arclet/entari/config/schema/plugin.py create mode 100644 example_plugins/example_plugin7/__init__.py create mode 100644 example_plugins/example_plugin7/config.py create mode 100644 example_plugins/example_plugin7/listener.py diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index 88f2f97..db1b289 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -82,6 +82,7 @@ from .plugin import metadata as metadata from .plugin import package as package from .plugin import plugin_config as plugin_config +from .plugin import register_schema as register_schema from .plugin import requires as requires from .plugin import unload_plugin as unload_plugin from .plugin.model import inject as inject diff --git a/arclet/entari/config/file.py b/arclet/entari/config/file.py index c0ca546..e5d4e8d 100644 --- a/arclet/entari/config/file.py +++ b/arclet/entari/config/file.py @@ -15,6 +15,8 @@ from .action import Proxy, config_model_dump, config_model_schema, config_model_validate from .model import BasicConfig +from .schema import apply_schema_fragments +from .schema.plugin import PLUGIN_META_PROPERTIES_EXTRA, plugin_config_schema from .util import GetattrDict try: @@ -336,58 +338,41 @@ def updater(target): def generate_schema(self, plugins: list["Plugin"]): plugins_properties = {} # fmt: off - plugin_meta_properties = {"$disable": {"type": "string", "description": "Expression for whether disable this plugin"}, "$priority": {"type": "integer", "description": "Plugin loading priority, lower value means higher priority (default: 16)"}, "$filter": {"type": "string", "description": "Plugin filter expression, which will be evaluated in the context of the plugin"}} # noqa: E501 # Build a mapping from plugin config key to plugin object for $files schema generation plugin_map: dict[str, "Plugin"] = {} # noqa: UP037 for plug in plugins: plugin_map[plug._config_key] = plug - if plug.metadata is not None: - if plug.metadata.config: - schema = config_model_schema(plug.metadata.config, ref_root=f"/properties/plugins/properties/{plug._config_key}/") # noqa: E501 - schema["properties"].update(plugin_meta_properties) - plugins_properties[plug._config_key] = schema - else: - plugins_properties[plug._config_key] = {"type": "object", "description": f"{plug.metadata.description or plug.metadata.name}; no configuration required", "additionalProperties": True, "properties": plugin_meta_properties} # noqa: E501 - else: - plugins_properties[plug._config_key] = {"type": "object", "description": "No configuration required", "additionalProperties": True, "properties": plugin_meta_properties} # noqa: E501 + plugins_properties[plug._config_key] = plugin_config_schema(plug, ref_root=f"/properties/plugins/properties/{plug._config_key}") # noqa: E501 schemas = { "basic": config_model_schema(BasicConfig, ref_root="/properties/basic/"), "plugins": {"type": "object", "description": "Plugin configurations", "properties": {"$prefix": {"description": "List of prefix config", "items": {"properties": {"key": {"description": "Prefix key", "title": "Key", "type": "string"}, "plugins": {"anyOf": [{"type": "string"}, {"items": {"type": "string", "description": "Plugin name"}, "type": "array", "uniqueItems": True}], "description": "List of plugins under the prefix, or select an item of $files to apply plugins", "title": "Plugins"}}, "required": ["key"], "title": "Prefix Config", "type": "object"}, "type": "array"}, "$prelude": {"type": "array", "items": {"type": "string", "description": "Plugin name"}, "description": "List of prelude plugins to load", "default": [], "uniqueItems": True}, "$files": {"type": "array", "items": {"type": "string", "description": "File path"}, "description": "List of configuration files to load", "default": [], "uniqueItems": True}, **plugins_properties}}, "adapters": {"type": "array", "description": "Adapter configurations", "items": {"type": "object", "description": "Adapter configuration", "properties": {"$path": {"type": "string", "description": "Adapter Module Path"}}, "required": ["$path"], "additionalProperties": True}} # noqa: E501 } with open(f"{self.path.stem}.schema.json", "w", encoding="utf-8") as f: json.dump({"$schema": "https://json-schema.org/draft/2020-12/schema", "type": "object", "properties": schemas, "additionalProperties": False, "required": ["basic"]}, f, indent=2, ensure_ascii=False) # noqa: E501 + # fmt: on # Generate schema for each file in $files for file in self.plugin_extra_files: path = Path(file) if path.is_file() and not path.name.endswith(".schema.json"): - self._generate_extra_file_schema(path, plugin_map, plugin_meta_properties) + self._generate_extra_file_schema(path, plugin_map) elif path.is_dir(): for _path in path.iterdir(): if _path.is_file() and not _path.name.endswith(".schema.json"): - self._generate_extra_file_schema(_path, plugin_map, plugin_meta_properties) - # fmt: on + self._generate_extra_file_schema(_path, plugin_map) plugin_map.clear() - def _generate_extra_file_schema(self, path: Path, plugin_map: dict[str, "Plugin"], plugin_meta_properties: dict): + def _generate_extra_file_schema(self, path: Path, plugin_map: dict[str, "Plugin"]): """Generate schema for an extra config file from $files.""" - # fmt: off plugin_key = path.stem schema_file = path.with_suffix(".schema.json") - plugin_meta_properties = {**plugin_meta_properties, "$optional": {"type": "boolean", "description": "Whether this plugin is optional"}} # noqa: E501 - + # fmt: off # Check if we have a matching plugin with config if plugin_key in plugin_map: - plug = plugin_map[plugin_key] - if plug.metadata is not None and plug.metadata.config: - plugin_schema = config_model_schema(plug.metadata.config, ref_root="/") - plugin_schema["properties"].update(plugin_meta_properties) - elif plug.metadata is not None: - plugin_schema = {"type": "object", "description": f"{plug.metadata.description or plug.metadata.name}; no configuration required", "additionalProperties": True, "properties": plugin_meta_properties} # noqa: E501 - else: - plugin_schema = {"type": "object", "description": "No configuration required", "additionalProperties": True, "properties": plugin_meta_properties} # noqa: E501 + plugin_schema = plugin_config_schema(plugin_map[plugin_key], use_extra_meta=True) else: # Plugin not found, generate a generic schema - plugin_schema = {"type": "object", "description": f"Configuration for {plugin_key}", "additionalProperties": True, "properties": plugin_meta_properties} # noqa: E501 + plugin_schema = {"type": "object", "description": f"Configuration for {plugin_key}", "additionalProperties": True, "properties": {k: dict(v) for k, v in PLUGIN_META_PROPERTIES_EXTRA.items()}} # noqa: E501 + apply_schema_fragments(plugin_schema, plugin_key, ref_root="/") with open(schema_file, "w", encoding="utf-8") as f: json.dump({"$schema": "https://json-schema.org/draft/2020-12/schema", **plugin_schema}, f, indent=2, ensure_ascii=False) # noqa: E501 diff --git a/arclet/entari/config/schema/__init__.py b/arclet/entari/config/schema/__init__.py index 47e9c7a..77d2893 100644 --- a/arclet/entari/config/schema/__init__.py +++ b/arclet/entari/config/schema/__init__.py @@ -1,5 +1,6 @@ from .dataclass import SchemaGenerator as SchemaGenerator +from .plugin import plugin_config_schema as plugin_config_schema from .registry import add_schema_fragment as add_schema_fragment -from .registry import purge_schema_fragments as purge_schema_fragments -from .registry import has_schema_fragments as has_schema_fragments from .registry import apply_schema_fragments as apply_schema_fragments +from .registry import has_schema_fragments as has_schema_fragments +from .registry import purge_schema_fragments as purge_schema_fragments diff --git a/arclet/entari/config/schema/plugin.py b/arclet/entari/config/schema/plugin.py new file mode 100644 index 0000000..c41a376 --- /dev/null +++ b/arclet/entari/config/schema/plugin.py @@ -0,0 +1,91 @@ +from typing import TYPE_CHECKING, Any + +from ..action import config_model_schema +from .registry import apply_schema_fragments + +if TYPE_CHECKING: + from arclet.entari.plugin import Plugin + + +# fmt: off +PLUGIN_META_PROPERTIES = {"$disable": {"type": "string", "description": "Expression for whether disable this plugin"}, "$priority": {"type": "integer", "description": "Plugin loading priority, lower value means higher priority (default: 16)"}, "$filter": {"type": "string", "description": "Plugin filter expression, which will be evaluated in the context of the plugin"}} # noqa: E501 +PLUGIN_META_PROPERTIES_EXTRA = PLUGIN_META_PROPERTIES | {"$optional": {"type": "boolean", "description": "Whether this plugin is optional"}} # noqa: E501 +# fmt: on + + +def _synthesize_subplugins(plug: "Plugin", schema: dict, ref_root: str): + from arclet.entari.plugin.service import plugin_service + + properties = schema.setdefault("properties", {}) + if not isinstance(properties, dict): + # log.warning("plugin schema 'properties' is not a dict, skip sub-plugin synthesis") + return + for sid in plug.subplugins: + if not sid.startswith(plug.id): + continue + dotted = sid[len(plug.id) :] + if (at := dotted.rfind("@")) != -1: # A.a@uid -> .a(与 module.py:310 加载键一致) + dotted = dotted[:at] + meta = {k: dict(v) for k, v in PLUGIN_META_PROPERTIES.items()} + sub = plugin_service.plugins.get(sid) + if sub is not None and sub.metadata is not None and sub.metadata.config: + sub_ref_root = f"{ref_root}properties/{dotted}/" + sub_schema = config_model_schema(sub.metadata.config, ref_root=sub_ref_root) + sub_props = sub_schema.setdefault("properties", {}) + if not isinstance(sub_props, dict): + sub_props = sub_schema["properties"] = {} + sub_props.update(meta) + if sub.metadata.description or sub.metadata.name: + sub_schema.setdefault("description", sub.metadata.description or sub.metadata.name) + # else: + # if sub is not None and sub.metadata is not None: + # desc = f"{sub.metadata.description or sub.metadata.name}; no configuration required" + # else: + # desc = "No configuration required" + # sub_schema = {"type": "object", "description": desc, "additionalProperties": True, "properties": meta} + properties[dotted] = sub_schema + + +def plugin_config_schema( + plugin: "Plugin", + config_key: str | None = None, + *, + ref_root: str = "/", + expand_subplugins: bool = True, + use_extra_meta: bool = False, +) -> dict[str, Any]: + """获取插件配置模型的 JSON Schema + + Args: + plugin (Plugin): 插件实例。 + config_key (str, optional): 插件实例的配置键。 + ref_root (str, optional): JSON Schema $ref 的根路径,默认为 "/"。 + expand_subplugins (bool, optional): 是否展开子插件的配置模型,默认为 True。 + use_extra_meta (bool, optional): 是否使用额外的元属性($optional),默认为 False。 + """ + meta_properties = { + k: dict(v) for k, v in (PLUGIN_META_PROPERTIES_EXTRA if use_extra_meta else PLUGIN_META_PROPERTIES).items() + } + metadata = plugin.metadata + if metadata and metadata.config is not None: + schema = config_model_schema(metadata.config, ref_root=ref_root) + properties = schema.setdefault("properties", {}) + if not isinstance(properties, dict): + properties = schema["properties"] = {} + properties.update(meta_properties) + elif metadata is not None: + desc = f"{metadata.description or metadata.name}; no configuration required" + schema = {"type": "object", "description": desc, "additionalProperties": True, "properties": meta_properties} + else: + schema = { + "type": "object", + "description": "No configuration required", + "additionalProperties": True, + "properties": meta_properties, + } + if expand_subplugins: + _synthesize_subplugins(plugin, schema, ref_root) + config_key = config_key or getattr(plugin, "_config_key", None) + if config_key: + apply_schema_fragments(schema, config_key, ref_root) + return schema diff --git a/arclet/entari/config/schema/rebase.py b/arclet/entari/config/schema/rebase.py index 2a8ff5b..6fa412b 100644 --- a/arclet/entari/config/schema/rebase.py +++ b/arclet/entari/config/schema/rebase.py @@ -1,5 +1,4 @@ import copy - from collections.abc import Callable from typing import Any, TypeAlias @@ -28,9 +27,7 @@ def _pointer_path(parts: tuple[str, ...]) -> str: return "/" + "/".join(segments) if segments else "" -def _resolve_target( - schema: dict, parts: tuple[str, ...] -) -> tuple[dict | None, str | None, dict | None]: +def _resolve_target(schema: dict, parts: tuple[str, ...]) -> tuple[dict | None, str | None, dict | None]: """沿 parts 在 schema 内导航到目标节点,返回 (holder, key, target) - target: 已存在的 dict 目标节点(可原地深合并); - target 为 None 且 holder/key 非空:目标缺失或为非 dict,调用方应整体写入 holder[key]; @@ -51,7 +48,7 @@ def _resolve_target( raise ValueError(f"schema fragment path crosses non-dict node {part!r}, skipped") props = node.setdefault("properties", {}) if not isinstance(props, dict): - raise ValueError(f"schema fragment path crosses non-dict 'properties' node, skipped") + raise ValueError("schema fragment path crosses non-dict 'properties' node, skipped") holder, key = props, part value = props.get(part, _SENTINEL) else: @@ -121,7 +118,9 @@ def _collect_defs_and_rebase( return hoisted -def update_schema(schema: dict, config_key: str, ref_root: str, parts: tuple[str, ...], payload: Fragment, replace: bool = False) -> None: +def update_schema( + schema: dict, config_key: str, ref_root: str, parts: tuple[str, ...], payload: Fragment, replace: bool = False +) -> None: final_ref_root = f"{ref_root}{_pointer_path(parts)}/" if parts else ref_root if isinstance(payload, type): # 模型片段:以最终 ref_root 生成,再走统一的 defs 处理 diff --git a/arclet/entari/config/schema/registry.py b/arclet/entari/config/schema/registry.py index d568a0a..54001d3 100644 --- a/arclet/entari/config/schema/registry.py +++ b/arclet/entari/config/schema/registry.py @@ -112,5 +112,10 @@ def apply_schema_fragments(schema: dict, config_key: str, ref_root: str = "/") - if any(mark == old for old in applied): continue applied.append(mark) - update_schema(schema, config_key, ref_root, record.path, record.fragment, record.replace) + try: + update_schema(schema, config_key, ref_root, record.path, record.fragment, record.replace) + except (RuntimeError, ValueError) as e: + from ...logger import log + + log.plugin.warning(str(e)) return schema diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index b8f7258..bdb9830 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -3,9 +3,9 @@ import itertools import os import sys -from collections.abc import Awaitable +from collections.abc import Awaitable, Callable from pathlib import Path -from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload, Callable +from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload from arclet.letoderea import Subscriber, on from arclet.letoderea.effect import AsyncDisposable, Disposable @@ -268,7 +268,12 @@ def _reload(event: ConfigReload): get_config = plugin_config -def register_schema(path: str | tuple[str, ...] = "", fragment: dict | type | Callable[[dict], dict] | None = None, *, replace: bool = False) -> None: +def register_schema( + path: str | tuple[str, ...] = "", + fragment: dict | type | Callable[[dict], dict] | None = None, + *, + replace: bool = False, +) -> None: plugin = get_plugin(1) plugin.register_schema(path, fragment, replace=replace) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 29d9533..ee023af 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -8,7 +8,6 @@ from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from enum import Enum -from itertools import chain from pathlib import Path from types import ModuleType from typing import Any, Generic, TypeVar, cast @@ -39,8 +38,7 @@ from tarina import ContextModel from tarina.tools import TCallable, run_sync -from ..config import config_model_schema -from ..config.schema import purge_schema_fragments, add_schema_fragment, apply_schema_fragments, has_schema_fragments +from ..config.schema import add_schema_fragment, plugin_config_schema, purge_schema_fragments from ..event.config import ConfigReload from ..event.plugin import PluginLoadedFailed, PluginLoadedSuccess, PluginUnloaded from ..exceptions import RegisterNotInPluginError, ReusablePluginError, StaticPluginDispatchError @@ -329,21 +327,14 @@ def register_schema( """注册作用于本插件配置的 schema 片段""" add_schema_fragment(self._config_key, path, fragment, replace, self.id) - def config_schema(self, config_key: str | None = None) -> dict[str, Any]: + def config_schema(self, config_key: str | None = None, ref_root: str = "/") -> dict[str, Any]: """获取插件配置模型的 JSON Schema Args: config_key (str, optional): 本插件的配置键。 + ref_root (str, optional): JSON Schema $ref 的根路径,默认为 "/"。 """ - metadata = self._metadata - if metadata is None or metadata.config is None: - if not config_key or not has_schema_fragments(config_key): - return {} - return apply_schema_fragments({"type": "object", "additionalProperties": True}, config_key, ref_root="/") - schema = config_model_schema(metadata.config) - if config_key: - return apply_schema_fragments(schema, config_key, ref_root="/") - return schema + return plugin_config_schema(self, config_key, ref_root=ref_root) # type: ignore def exec_apply(self): if not self._apply: diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index 75a28f3..c37dc23 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -233,7 +233,13 @@ class Plugin: def metadata(self) -> PluginMetadata | None: ... @metadata.setter def metadata(self, value: PluginMetadata): ... - def register_schema(self, path: str | tuple[str, ...] = "", fragment: dict | type | Callable[[dict], dict] | None = None, *, replace: bool = False) -> None: ... + def register_schema( + self, + path: str | tuple[str, ...] = "", + fragment: dict | type | Callable[[dict], dict] | None = None, + *, + replace: bool = False, + ) -> None: ... def config_schema(self, config_key: str | None = None) -> dict[str, Any]: ... def exec_apply(self) -> None: ... @property diff --git a/example_plugins/example_plugin7/__init__.py b/example_plugins/example_plugin7/__init__.py new file mode 100644 index 0000000..b184ee4 --- /dev/null +++ b/example_plugins/example_plugin7/__init__.py @@ -0,0 +1,34 @@ +from arclet.entari import BasicConfModel, metadata, plugin_config, register_schema + +from .config import Config +from . import listener + +metadata(__file__, description="1", config=Config) +conf = plugin_config(Config) + + +class Foo(BasicConfModel): + abc: int + + +class Bar(BasicConfModel): + xyz: str + + +if conf.tools: + register_schema( + "properties.tools", + {"type": "object", "description": "Tools configuration", "properties": {}}, + replace=True + ) + for tool in conf.tools: + if tool == "foo": + register_schema( + f"properties.tools.properties.{tool}", + Foo + ) + elif tool == "bar": + register_schema( + f"properties.tools.properties.{tool}", + Bar + ) diff --git a/example_plugins/example_plugin7/config.py b/example_plugins/example_plugin7/config.py new file mode 100644 index 0000000..edb8f3c --- /dev/null +++ b/example_plugins/example_plugin7/config.py @@ -0,0 +1,7 @@ +from arclet.entari.config import BasicConfModel, model_field + + +class Config(BasicConfModel): + input: str + output: str + tools: dict[str, dict] = model_field(default_factory=dict) diff --git a/example_plugins/example_plugin7/listener.py b/example_plugins/example_plugin7/listener.py new file mode 100644 index 0000000..75a9d18 --- /dev/null +++ b/example_plugins/example_plugin7/listener.py @@ -0,0 +1,15 @@ +from arclet.entari import BasicConfModel, MessageCreatedEvent, plugin + + +class Config(BasicConfModel): + foo: str + bar: str + + +plugin.metadata(__file__, config=Config) +conf = plugin.get_config(Config) + + +@plugin.listen(MessageCreatedEvent) +async def _(): + print(conf, 2) From f3970b6658ddb57d30780d58ee3d7669e26a0aad Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Mon, 21 Sep 2026 17:48:15 +0800 Subject: [PATCH 31/33] :boom: `plugin.isolate` --- arclet/entari/__init__.py | 2 + arclet/entari/command/__init__.py | 4 +- arclet/entari/config/file.py | 11 ++++- arclet/entari/core.py | 21 +++++----- arclet/entari/localdata.py | 6 +-- arclet/entari/plugin/__init__.py | 12 +++++- arclet/entari/plugin/loader.py | 9 ++--- arclet/entari/plugin/model.py | 67 +++++++++++++------------------ arclet/entari/plugin/model.pyi | 12 +----- arclet/entari/plugin/module.py | 31 +++++++------- arclet/entari/plugin/service.py | 9 ++--- arclet/entari/scheduler.py | 6 +-- 12 files changed, 93 insertions(+), 97 deletions(-) diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index db1b289..ecb1ec8 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -68,6 +68,7 @@ from .filter import filter_ as filter_ from .localdata import local_data as local_data from .message import MessageChain as MessageChain +from .plugin import ROOT as ROOT from .plugin import Plugin as Plugin from .plugin import PluginMetadata as PluginMetadata from .plugin import add_service as add_service @@ -75,6 +76,7 @@ from .plugin import declare_static as declare_static from .plugin import disable_plugin as disable_plugin from .plugin import enable_plugin as enable_plugin +from .plugin import isolate as isolate from .plugin import keeping as keeping from .plugin import listen as listen from .plugin import load_plugin as load_plugin diff --git a/arclet/entari/command/__init__.py b/arclet/entari/command/__init__.py index 5a3aa3f..bb029ac 100644 --- a/arclet/entari/command/__init__.py +++ b/arclet/entari/command/__init__.py @@ -23,7 +23,7 @@ from ..event.config import ConfigReload from ..logger import DEBUG_NO, log from ..message import MessageChain -from ..plugin import PluginRole, RootlessPlugin, get_plugin, metadata, plugin_config +from ..plugin import ROOT, PluginRole, get_plugin, metadata, plugin_config from ..session import Session from .argv import MessageArgv # noqa: F401 from .model import CommandResult, Match, Query @@ -268,7 +268,7 @@ class CommandsConfig(BasicConfModel): ) -@RootlessPlugin.apply("commands", default=True) +@ROOT.isolate("commands") def command_plugin(_): metadata( "Commands Plugin", diff --git a/arclet/entari/config/file.py b/arclet/entari/config/file.py index e5d4e8d..6ae4a7e 100644 --- a/arclet/entari/config/file.py +++ b/arclet/entari/config/file.py @@ -341,8 +341,15 @@ def generate_schema(self, plugins: list["Plugin"]): # Build a mapping from plugin config key to plugin object for $files schema generation plugin_map: dict[str, "Plugin"] = {} # noqa: UP037 for plug in plugins: - plugin_map[plug._config_key] = plug - plugins_properties[plug._config_key] = plugin_config_schema(plug, ref_root=f"/properties/plugins/properties/{plug._config_key}") # noqa: E501 + if plug.id == "$": + schema = plugin_config_schema(plug, ref_root="/properties/plugins") + for key, property_ in schema["properties"].items(): + if key.startswith("$"): + continue + plugins_properties[key] = property_ + else: + plugin_map[plug._config_key] = plug + plugins_properties[plug._config_key] = plugin_config_schema(plug, ref_root=f"/properties/plugins/properties/{plug._config_key}") # noqa: E501 schemas = { "basic": config_model_schema(BasicConfig, ref_root="/properties/basic/"), "plugins": {"type": "object", "description": "Plugin configurations", "properties": {"$prefix": {"description": "List of prefix config", "items": {"properties": {"key": {"description": "Prefix key", "title": "Key", "type": "string"}, "plugins": {"anyOf": [{"type": "string"}, {"items": {"type": "string", "description": "Plugin name"}, "type": "array", "uniqueItems": True}], "description": "List of plugins under the prefix, or select an item of $files to apply plugins", "title": "Plugins"}}, "required": ["key"], "title": "Prefix Config", "type": "object"}, "type": "array"}, "$prelude": {"type": "array", "items": {"type": "string", "description": "Plugin name"}, "description": "List of prelude plugins to load", "default": [], "uniqueItems": True}, "$files": {"type": "array", "items": {"type": "string", "description": "File path"}, "description": "List of configuration files to load", "default": [], "uniqueItems": True}, **plugins_properties}}, "adapters": {"type": "array", "description": "Adapter configurations", "items": {"type": "object", "description": "Adapter configuration", "properties": {"$path": {"type": "string", "description": "Adapter Module Path"}}, "required": ["$path"], "additionalProperties": True}} # noqa: E501 } diff --git a/arclet/entari/core.py b/arclet/entari/core.py index 1dd2c90..645bdfb 100644 --- a/arclet/entari/core.py +++ b/arclet/entari/core.py @@ -49,12 +49,12 @@ from .event.api import SendResponse from .event.base import MessageCreatedEvent, event_parse from .event.config import ConfigReload -from .event.lifespan import AccountUpdate +from .event.lifespan import AccountUpdate, Ready from .localdata import local_data from .logger import apply_log_save, enable_rich_except, log from .message import MessageChain -from .plugin import get_plugins, load_plugin, plugin_config, requires -from .plugin.model import PluginMetadata, PluginRole, RootlessPlugin +from .plugin import ROOT, Plugin, get_plugins, load_plugin, plugin_config, requires +from .plugin.model import PluginMetadata, PluginRole from .plugin.service import plugin_service from .session import EntariProtocol, Session @@ -254,8 +254,8 @@ class RecordConfig(BasicConfModel): """是否在日志中使用简短的消息内容""" -@RootlessPlugin.apply("record_message", default=True) -def record(plg: RootlessPlugin): +@ROOT.isolate("record_message") +def record(plg: Plugin): plg.metadata = PluginMetadata( "记录消息", PluginRole.UTILITY, @@ -303,10 +303,6 @@ async def log_send(event: SendResponse): class Entari(App): id = "entari.service" - @property - def main_plugin(self): - return RootlessPlugin.apply("main", default=True) - @classmethod def load(cls, path: str | os.PathLike[str] | None = None): return cls.from_config(EntariConfig.load(path)) @@ -455,9 +451,7 @@ def ensure_manager(self, manager: Launart): load_plugin(plug, prelude=True) plugins = EntariConfig.instance.plugin_names requires(*plugins) - for apply, slot in plugin_service._apply.items(): - if slot[1] and apply not in EntariConfig.instance.plugin: - plugins.append(apply) + for plug in plugins: load_plugin(plug) @@ -465,6 +459,9 @@ def ensure_manager(self, manager: Launart): if entry.name not in plugins and entry.module not in plugin_service.plugins: load_plugin(entry.module) + le.on(Ready, self._generate_schema, priority=0) + + def _generate_schema(self): if self.gen_schema and EntariConfig.instance.path.exists(): EntariConfig.instance.generate_schema(get_plugins()) diff --git a/arclet/entari/localdata.py b/arclet/entari/localdata.py index 3bacd52..0e396cf 100644 --- a/arclet/entari/localdata.py +++ b/arclet/entari/localdata.py @@ -10,7 +10,7 @@ from .config import BasicConfModel, config_model_validate, model_field from .event.config import ConfigReload -from .plugin import PluginRole, RootlessPlugin, metadata, plugin_config +from .plugin import ROOT, PluginRole, metadata, plugin_config P = ParamSpec("P") @@ -108,8 +108,8 @@ class Config(BasicConfModel): base_dir: str | None = model_field(default=None, description="基础目录,默认为空,表示使用 `app_name` 作为目录名,") -@RootlessPlugin.apply("localdata", default=True) -def localdata_apply(plg: RootlessPlugin): +@ROOT.isolate("localdata") +def localdata_apply(plg): metadata( "Data store interface with local directories", PluginRole.LIBRARY, diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index bdb9830..aff43ad 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -5,6 +5,7 @@ import sys from collections.abc import Awaitable, Callable from pathlib import Path +from types import ModuleType from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload from arclet.letoderea import Subscriber, on @@ -26,7 +27,6 @@ from .model import TS, Plugin, PluginDispatcher, current_plugin from .model import PluginMetadata as PluginMetadata from .model import PluginRole as PluginRole -from .model import RootlessPlugin as RootlessPlugin from .model import keeping as keeping from .module import package as package from .module import requires as requires @@ -404,3 +404,13 @@ def _(): return func return wrapper + + +def isolate(label: str): + """创建一个隔离的插件上下文, 用于在当前插件中创建一个子插件""" + _plugin = get_plugin(1) + return _plugin.isolate(label) + + +ROOT = plugin_service.root = Plugin("$", ModuleType("$"), config={"$path": "$"}) +ROOT.module.__file__ = "__main__" diff --git a/arclet/entari/plugin/loader.py b/arclet/entari/plugin/loader.py index 28736c8..d7ad92a 100644 --- a/arclet/entari/plugin/loader.py +++ b/arclet/entari/plugin/loader.py @@ -259,13 +259,12 @@ def load_plugin( if recursive_guard is None: recursive_guard = set() path = path.replace("::", "arclet.entari.builtins.") + if path.startswith(".") and f"${path}" in plugin_service.root.subplugins: + rootless = plugin_service.plugins[f"${path}"] + rootless.config = config + return rootless while path in plugin_service._subplugined: path = plugin_service._subplugined[path] - if path in plugin_service._apply: - if path in plugin_service.plugins: - return plugin_service.plugins[path] - log.plugin.trace(f"loaded rootless plugin {path!r}") - return plugin_service._apply[path][0](config) if not staged and (plug := find_plugin(path)): plugin_service._direct_plugins.add(plug.path) return plug diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index ee023af..8faa4af 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -31,7 +31,7 @@ from arclet.letoderea.effect import AsyncDisposable, Disposable from arclet.letoderea.provider import Provider, ProviderFactory, TProviders from arclet.letoderea.publisher import Publisher, _publishers, filter_publisher -from arclet.letoderea.scope import RegisterWrapper +from arclet.letoderea.scope import RegisterWrapper, scope_ctx from arclet.letoderea.utils import DisposableList, add_task from creart import it from launart import Launart, Service @@ -278,10 +278,12 @@ def metadata(self, value: PluginMetadata): value.config.__doc__ = value.description or value.name def __post_init__(self): + if not hasattr(self.module, "__plugin__"): + setattr(self.module, "__plugin__", self) uid_index = self.id.rfind("@") self.path = self.id[:uid_index] if uid_index != -1 else self.id self.uid = self.id[uid_index + 1 :] if uid_index != -1 else None - if self.id in plugin_service.plugins and not self.id.startswith("."): + if self.id in plugin_service.plugins and not self.id.startswith("$"): # 原子重载暂存:id 冲突时注册进 _staged,scope 用唯一 id 并置 disabled self._scope = _make_scope(self).of(f"{self.id}@staging") self._scope.disable() @@ -341,9 +343,14 @@ def exec_apply(self): return log.plugin.trace(f"applying plugin {self.id!r}") token = current_plugin.set(self) + if not self.is_static: + token1 = scope_ctx.set(self._scope) try: self._apply(self) - log.plugin.success(f"plugin {self.id!r} fully applied") + if self.id in plugin_service._subplugined: + log.plugin.trace(f"sub-plugin {self.id!r} fully applied") + else: + log.plugin.success(f"plugin {self.id!r} fully applied") publish(PluginLoadedSuccess(self.id)) except (ImportError, RegisterNotInPluginError, StaticPluginDispatchError, ReusablePluginError) as e: log.plugin.error(f"failed to load plugin {self.id!r}: {e.args[0]}") @@ -357,6 +364,8 @@ def exec_apply(self): raise finally: current_plugin.reset(token) + if not self.is_static or "token1" in locals(): + scope_ctx.reset(token1) # type: ignore @property def is_available(self) -> bool: @@ -497,7 +506,7 @@ def dispose(self, *, is_cleanup: bool = False, replacing: bool = False): if self._is_disposed: return purge_schema_fragments(self.id) - if not self.id.startswith(".") and self.id not in plugin_service._subplugined: + if not self.id.startswith("$") and self.id not in plugin_service._subplugined: log.plugin.debug(f"disposing plugin {self.id}") _was_staged = self.id in plugin_service._staged self._is_disposed = True @@ -634,44 +643,24 @@ def restore_kept_state(self): if kept.module_attr: self.module.__dict__[kept.module_attr] = kept.obj + def isolate(self, label: str): + """创建一个隔离的子插件,子插件的生命周期与父插件绑定""" + sub_id = f"{self.id}.{label}" + if sub_id in plugin_service.plugins: + raise ValueError(f"sub-plugin {sub_id} already exists") + subplug = Plugin(sub_id, ModuleType(sub_id), config=self.config.copy()) + setattr(subplug.module, "__plugin__", subplug) + self.subplugins.append(sub_id) + plugin_service._subplugined[sub_id] = self.id + + def wrapper(func: Callable[[Plugin], Any], /): + setattr(func, "__plugin__", subplug) + setattr(subplug.module, "__file__", func.__code__.co_filename) + subplug._apply = func + return subplug -class RootlessPlugin(Plugin): - # fmt: off - @classmethod - def apply(cls: type[RootlessPlugin], id: str, func: Callable[[RootlessPlugin], Any] | None = None, *, default: bool = False) -> Any: # noqa: E501 - # fmt: on - if not id.startswith("."): - id = f".{id}" - - def dispose(): - if id in plugin_service.plugins: - plugin_service.plugins[id].dispose() - else: - plugin_service._apply.pop(id, None) - - def wrapper(func: Callable[[RootlessPlugin], Any]): - plugin_service._apply[id] = (lambda config: cls(id, func, config), default) # type: ignore - return dispose - - if func: - return wrapper(func) return wrapper - def __init__(self, id: str, func: Callable[[RootlessPlugin], Any], config: dict): - super().__init__(id, ModuleType(id), config=config) - setattr(self.module, "__plugin__", self) - setattr(self.module, "__file__", func.__code__.co_filename) - self.func = func - setattr(self.func, "__plugin__", self) - token = current_plugin.set(self) - try: - func(self) - finally: - current_plugin.reset(token) - - def validate(self, func): - pass - class KeepingVariable(Generic[T]): def __init__(self, obj: T, dispose=None, module_attr=None): diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index c37dc23..d71e468 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -364,17 +364,7 @@ class Plugin: def subproxy(self, sub_id: str) -> ModuleType: ... def service(self, serv: TS | type[TS]) -> TS: ... def restore_kept_state(self) -> None: ... - -class RootlessPlugin(Plugin): - @classmethod - @overload - def apply( - cls, id: str, *, default: bool = False - ) -> Callable[[Callable[[RootlessPlugin], Any]], Callable[[], None]]: ... - @classmethod - @overload - def apply(cls, id: str, func: Callable[[RootlessPlugin], Any], *, default: bool = False) -> Callable[[], None]: ... - def __init__(self, id: str, func: Callable[[RootlessPlugin], Any], config: dict): ... + def isolate(self, label: str) -> Callable[[Callable[[Plugin], Any]], Plugin]: ... class KeepingVariable(Generic[T]): obj: T diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index 031aed9..ac0eb5b 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -318,21 +318,24 @@ def exec_module(self, module: ModuleType, config: dict[str, Any] | None = None) if config is None or (not is_sub and not {k: v for k, v in config.items() if k not in ("$path", "$static")}): key = module.__name__ - if key.startswith("arclet.entari.builtins.") and f"::{key[23:]}" in EntariConfig.instance.plugin: - key = f"::{key[23:]}" - config = EntariConfig.instance.plugin.get(key) - if config is not None: - config["$path"] = key - else: - for k, names in EntariConfig.instance._plugin_names.items(): - if key in names: - config = EntariConfig.instance.plugin.get(k, {}) - config["$path"] = k - break + if hasattr(EntariConfig, "instance"): + if key.startswith("arclet.entari.builtins.") and f"::{key[23:]}" in EntariConfig.instance.plugin: + key = f"::{key[23:]}" + config = EntariConfig.instance.plugin.get(key) + if config is not None: + config["$path"] = key else: - config = {"$path": key} - if key in EntariConfig.instance.prelude_plugin: - config["$static"] = True # type: ignore + for k, names in EntariConfig.instance._plugin_names.items(): + if key in names: + config = EntariConfig.instance.plugin.get(k, {}) + config["$path"] = k + break + else: + config = {"$path": key} + if key in EntariConfig.instance.prelude_plugin: + config["$static"] = True # type: ignore + else: + config = {"$path": key} # create plugin before executing plugin = Plugin(self.plugin_id, module, config=config) # for `dataclasses` module diff --git a/arclet/entari/plugin/service.py b/arclet/entari/plugin/service.py index 02bc593..8b5f672 100644 --- a/arclet/entari/plugin/service.py +++ b/arclet/entari/plugin/service.py @@ -1,8 +1,7 @@ from __future__ import annotations import asyncio -from collections.abc import Callable -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from arclet.letoderea import es from launart import Launart, Service @@ -14,7 +13,7 @@ from ..logger import log if TYPE_CHECKING: - from .model import KeepingVariable, Plugin, RootlessPlugin + from .model import KeepingVariable, Plugin class ServiceWaiters: @@ -59,9 +58,10 @@ class PluginManagerService(Service): """卸载插件集合,存储所有已卸载的插件ID""" _subplugined: dict[str, str] """子插件字典,键为子插件ID,值为父插件ID""" - _apply: dict[str, tuple[Callable[[dict[str, Any]], RootlessPlugin], bool]] _staged: dict[str, Plugin] """插件暂存""" + root: Plugin + """根插件对象,作为基础插件的父插件""" def __init__(self): super().__init__() @@ -72,7 +72,6 @@ def __init__(self): self.references = {} self._unloaded = set() self._subplugined = {} - self._apply = {} self.bindings = {} self.fingerprints = {} self._staged = {} diff --git a/arclet/entari/scheduler.py b/arclet/entari/scheduler.py index 060739a..07b6fcd 100644 --- a/arclet/entari/scheduler.py +++ b/arclet/entari/scheduler.py @@ -13,7 +13,7 @@ from .config import BasicConfModel from .event.config import ConfigReload from .logger import log -from .plugin import PluginRole, RootlessPlugin, get_plugin, metadata, plugin_config +from .plugin import ROOT, PluginRole, get_plugin, metadata, plugin_config @make_event(name="entari.event/internal/schedule") @@ -137,8 +137,8 @@ class _SchedulerConf(BasicConfModel): schedule = scheduler.schedule -@RootlessPlugin.apply("scheduler", default=True) -def _(plg: RootlessPlugin): +@ROOT.isolate("scheduler") +def _(plg): metadata( "Simple Scheduler with interval / crontab task", PluginRole.LIBRARY, From 0aabddacd66109c2462e00f57df2ec195f8382d1 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Tue, 22 Sep 2026 16:54:50 +0800 Subject: [PATCH 32/33] :beers: `plugin.isolate` -> `plugin.niche` --- arclet/entari/__init__.py | 2 +- arclet/entari/command/__init__.py | 4 ++-- arclet/entari/core.py | 4 ++-- arclet/entari/localdata.py | 4 ++-- arclet/entari/plugin/__init__.py | 9 ++++++--- arclet/entari/plugin/model.py | 12 ++++++------ arclet/entari/plugin/model.pyi | 4 +++- arclet/entari/scheduler.py | 4 ++-- 8 files changed, 24 insertions(+), 19 deletions(-) diff --git a/arclet/entari/__init__.py b/arclet/entari/__init__.py index ecb1ec8..10b0e4b 100644 --- a/arclet/entari/__init__.py +++ b/arclet/entari/__init__.py @@ -76,12 +76,12 @@ from .plugin import declare_static as declare_static from .plugin import disable_plugin as disable_plugin from .plugin import enable_plugin as enable_plugin -from .plugin import isolate as isolate from .plugin import keeping as keeping from .plugin import listen as listen from .plugin import load_plugin as load_plugin from .plugin import load_plugins as load_plugins from .plugin import metadata as metadata +from .plugin import niche as niche from .plugin import package as package from .plugin import plugin_config as plugin_config from .plugin import register_schema as register_schema diff --git a/arclet/entari/command/__init__.py b/arclet/entari/command/__init__.py index bb029ac..871271e 100644 --- a/arclet/entari/command/__init__.py +++ b/arclet/entari/command/__init__.py @@ -23,7 +23,7 @@ from ..event.config import ConfigReload from ..logger import DEBUG_NO, log from ..message import MessageChain -from ..plugin import ROOT, PluginRole, get_plugin, metadata, plugin_config +from ..plugin import PluginRole, get_plugin, metadata, niche, plugin_config from ..session import Session from .argv import MessageArgv # noqa: F401 from .model import CommandResult, Match, Query @@ -268,7 +268,7 @@ class CommandsConfig(BasicConfModel): ) -@ROOT.isolate("commands") +@niche("commands") def command_plugin(_): metadata( "Commands Plugin", diff --git a/arclet/entari/core.py b/arclet/entari/core.py index 645bdfb..66a1b52 100644 --- a/arclet/entari/core.py +++ b/arclet/entari/core.py @@ -53,7 +53,7 @@ from .localdata import local_data from .logger import apply_log_save, enable_rich_except, log from .message import MessageChain -from .plugin import ROOT, Plugin, get_plugins, load_plugin, plugin_config, requires +from .plugin import Plugin, get_plugins, load_plugin, niche, plugin_config, requires from .plugin.model import PluginMetadata, PluginRole from .plugin.service import plugin_service from .session import EntariProtocol, Session @@ -254,7 +254,7 @@ class RecordConfig(BasicConfModel): """是否在日志中使用简短的消息内容""" -@ROOT.isolate("record_message") +@niche("record_message") def record(plg: Plugin): plg.metadata = PluginMetadata( "记录消息", diff --git a/arclet/entari/localdata.py b/arclet/entari/localdata.py index 0e396cf..c2beef4 100644 --- a/arclet/entari/localdata.py +++ b/arclet/entari/localdata.py @@ -10,7 +10,7 @@ from .config import BasicConfModel, config_model_validate, model_field from .event.config import ConfigReload -from .plugin import ROOT, PluginRole, metadata, plugin_config +from .plugin import PluginRole, metadata, niche, plugin_config P = ParamSpec("P") @@ -108,7 +108,7 @@ class Config(BasicConfModel): base_dir: str | None = model_field(default=None, description="基础目录,默认为空,表示使用 `app_name` 作为目录名,") -@ROOT.isolate("localdata") +@niche("localdata") def localdata_apply(plg): metadata( "Data store interface with local directories", diff --git a/arclet/entari/plugin/__init__.py b/arclet/entari/plugin/__init__.py index aff43ad..5145c00 100644 --- a/arclet/entari/plugin/__init__.py +++ b/arclet/entari/plugin/__init__.py @@ -406,10 +406,13 @@ def _(): return wrapper -def isolate(label: str): +def niche(label: str, config: dict[str, Any] | None = None): """创建一个隔离的插件上下文, 用于在当前插件中创建一个子插件""" - _plugin = get_plugin(1) - return _plugin.isolate(label) + try: + _plugin = get_plugin(1) + except LookupError: + _plugin = ROOT + return _plugin.niche(label, config) ROOT = plugin_service.root = Plugin("$", ModuleType("$"), config={"$path": "$"}) diff --git a/arclet/entari/plugin/model.py b/arclet/entari/plugin/model.py index 8faa4af..44d942c 100644 --- a/arclet/entari/plugin/model.py +++ b/arclet/entari/plugin/model.py @@ -292,7 +292,7 @@ def __post_init__(self): self._scope = _make_scope(self).of(self.id) plugin_service.plugins[self.id] = self # type: ignore self.effect = self._scope.effect - self._config_key = self.config.pop("$path", self.id) + self._config_key = self.config.pop("$path", self.id.lstrip("$")) if filter_expr := self.config.get("$filter", ""): self._scope.propagators.append(FilterPropagator(filter_expr)) # if self._metadata and self._metadata.depend_services: @@ -643,20 +643,20 @@ def restore_kept_state(self): if kept.module_attr: self.module.__dict__[kept.module_attr] = kept.obj - def isolate(self, label: str): + def niche(self, label: str, config: dict[str, Any] | None = None): """创建一个隔离的子插件,子插件的生命周期与父插件绑定""" sub_id = f"{self.id}.{label}" if sub_id in plugin_service.plugins: raise ValueError(f"sub-plugin {sub_id} already exists") - subplug = Plugin(sub_id, ModuleType(sub_id), config=self.config.copy()) - setattr(subplug.module, "__plugin__", subplug) - self.subplugins.append(sub_id) - plugin_service._subplugined[sub_id] = self.id def wrapper(func: Callable[[Plugin], Any], /): + subplug = Plugin(sub_id, ModuleType(sub_id), config=self.config.copy() if config is None else config) + setattr(subplug.module, "__plugin__", subplug) setattr(func, "__plugin__", subplug) setattr(subplug.module, "__file__", func.__code__.co_filename) subplug._apply = func + self.subplugins.append(sub_id) + plugin_service._subplugined[sub_id] = self.id return subplug return wrapper diff --git a/arclet/entari/plugin/model.pyi b/arclet/entari/plugin/model.pyi index d71e468..544d97d 100644 --- a/arclet/entari/plugin/model.pyi +++ b/arclet/entari/plugin/model.pyi @@ -364,7 +364,9 @@ class Plugin: def subproxy(self, sub_id: str) -> ModuleType: ... def service(self, serv: TS | type[TS]) -> TS: ... def restore_kept_state(self) -> None: ... - def isolate(self, label: str) -> Callable[[Callable[[Plugin], Any]], Plugin]: ... + def niche( + self, label: str, config: dict[str, Any] | None = None + ) -> Callable[[Callable[[Plugin], Any]], Plugin]: ... class KeepingVariable(Generic[T]): obj: T diff --git a/arclet/entari/scheduler.py b/arclet/entari/scheduler.py index 07b6fcd..66549b4 100644 --- a/arclet/entari/scheduler.py +++ b/arclet/entari/scheduler.py @@ -13,7 +13,7 @@ from .config import BasicConfModel from .event.config import ConfigReload from .logger import log -from .plugin import ROOT, PluginRole, get_plugin, metadata, plugin_config +from .plugin import PluginRole, get_plugin, metadata, niche, plugin_config @make_event(name="entari.event/internal/schedule") @@ -137,7 +137,7 @@ class _SchedulerConf(BasicConfModel): schedule = scheduler.schedule -@ROOT.isolate("scheduler") +@niche("scheduler") def _(plg): metadata( "Simple Scheduler with interval / crontab task", From ad70f1d351727beaa449d03136087192c53d8b90 Mon Sep 17 00:00:00 2001 From: RF-Tar-Railt Date: Thu, 24 Sep 2026 14:45:31 +0800 Subject: [PATCH 33/33] :sparkles: improve plugin import code (align bootstrap) --- arclet/entari/plugin/module.py | 106 +++++++++++++++++++++++---------- 1 file changed, 73 insertions(+), 33 deletions(-) diff --git a/arclet/entari/plugin/module.py b/arclet/entari/plugin/module.py index ac0eb5b..62be3f7 100644 --- a/arclet/entari/plugin/module.py +++ b/arclet/entari/plugin/module.py @@ -614,18 +614,68 @@ def find_spec( return +def _load_module(spec: ModuleSpec, config: dict | None = None): + mod = module_from_spec(spec) + spec._initializing = True # type: ignore + try: + sys.modules[spec.name] = mod + try: + if spec.loader is None: + if spec.submodule_search_locations is None: + raise ImportError("missing loader", name=spec.name) + # A namespace package so do nothing. + else: + if isinstance(spec.loader, PluginLoader): + spec.loader.exec_module(mod, config=config) + else: + spec.loader.exec_module(mod) + except: + sys.modules.pop(spec.name, None) + raise + if isinstance(spec.loader, PluginLoader): + 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: + # Move the module to the end of sys.modules. + # We don't ensure that the import-related module attributes get + # set in the sys.modules replacement case. Such modules are on + # their own. + mod = sys.modules.pop(spec.name) + sys.modules[mod.__name__] = mod + finally: + spec._initializing = False # type: ignore + return mod + + def import_plugin(id_, package=None, config: dict | None = None, staged: bool = False) -> ModuleType | 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 parent_name = fullname.rpartition(".")[0] + child = fullname.rpartition(".")[2] + parent_path: Sequence[str] | None = None + parent_spec: ModuleSpec | None = None if parent_name: - parent: ModuleType | None + parent: ModuleType parts = parent_name.split(".") _current = parts[0] if _current in plugin_service.plugins: parent = plugin_service.plugins[_current].module enter_plugin = True + elif _current in sys.modules: + parent = sys.modules[_current] + enter_plugin = hasattr(parent, "__plugin__") else: parent = __import__(_current, fromlist=["__path__"]) enter_plugin = False @@ -635,54 +685,44 @@ def import_plugin(id_, package=None, config: dict | None = None, staged: bool = if _current in plugin_service.plugins: parent = plugin_service.plugins[_current].module enter_plugin = True + elif _current in sys.modules: + parent = sys.modules[_current] + enter_plugin = hasattr(parent, "__plugin__") elif _current in _ENSURE_IS_PLUGIN or enter_plugin: - if parent := import_plugin(_current): + tmp = import_plugin(_current) + if tmp is not None: + parent = tmp enter_plugin = True + del tmp else: parent = __import__(_current, fromlist=["__path__"]) enter_plugin = False else: parent = __import__(_current, fromlist=["__path__"]) _current += "." - if parent is None: + try: + parent_path = parent.__path__ + except AttributeError: raise ModuleNotFoundError( - f"parent module {parent_name!r} does not have __path__ attribute " f"while trying to find {fullname!r}", + f"parent module {parent_name!r} does not have __path__ attribute while trying to find {fullname!r}", name=fullname, ) - parent_path = parent.__path__ - else: - parent_path = None + parent_spec = parent.__spec__ + # _find_spec spec = _PluginFinder.find_spec(fullname, parent_path, origin_id_=id_, force=True, staged=staged) if not spec: return - mod = module_from_spec(spec) - spec._initializing = True # type: ignore + if parent_spec: + # Temporarily add child we are currently importing to parent's + # _uninitialized_submodules for circular import tracking. + parent_spec._uninitialized_submodules.append(child) # type: ignore try: - if spec.loader is None: - if spec.submodule_search_locations is None: - raise ImportError("missing loader", name=spec.name) - # A namespace package so do nothing. - else: - 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) - sys.modules[mod.__name__] = mod + mod = _load_module(spec, config=config) finally: - spec._initializing = False # type: ignore + if parent_spec: + parent_spec._uninitialized_submodules.pop() # type: ignore + if parent_name: + setattr(parent, child, mod.__plugin__.proxy()) return mod