Add v4 compat layer and legacy shims

- Introduce private v4 compatibility surface using
  _legacy_api.py, _legacy_runtime.py, _legacy_loader.py plus new
  _legacy_context.py and _legacy_star.py to centralize legacy adapters
  while keeping public APIs thin.
- Extend InitializeOutput to carry protocol_version for negotiated
  protocol, enabling runtime to adapt to the chosen v4 version.
- Add lightweight legacy support for Star/Context via new LegacyStar and
  LegacyContext shims and expose legacy API through the aggregate
  _legacy_api entry point.
- Ensure legacy loader preserves class declaration order by iterating
  module.__dict__ instead of relying on alphabetical sorting.
- Add tests: protocol_version handling in InitializeOutput, legacy
  main component order preservation, and embedded-newline framing in
  transport tests.
This commit is contained in:
whatevertogo
2026-03-14 17:16:42 +08:00
parent 142d3ad747
commit 1672bc3227
21 changed files with 4298 additions and 2650 deletions
+1
View File
@@ -70,3 +70,4 @@ old文件夹是兼容旧插件的测试,旧插件全部放进old文件夹
- 2026-03-13: 不要再维护第二套 `_legacy/` 并行目录。private compat 以顶层 `_legacy_api.py`、`_legacy_runtime.py`、`_legacy_loader.py`、`_session_waiter.py`、`_shared_preferences.py` 为唯一实现位置,同时保留公开兼容面 `astrbot_sdk.api`、`astrbot_sdk.compat` 和 `src-new/astrbot` facade。
- 2026-03-14: `test_plugin/old/` 和 `test_plugin/new/` 里可能带着已生成的 `__pycache__` / `*.pyc`。测试夹具复制示例插件时必须显式忽略这些缓存文件,否则临时插件目录、断言结果和 `git status` 都可能被污染。
- 2026-03-14: grouped worker / grouped env 路径不要再复制单 worker 的 compat 生命周期和 legacy runtime 绑定逻辑。优先复用 `_legacy_runtime.py` 里的 `bind_legacy_runtime_contexts()`、`run_legacy_worker_startup_hooks()`、`run_legacy_worker_shutdown_hooks()` 以及 `resolve_plugin_lifecycle_hook()`,否则很容易出现“普通 worker 测试通过,但真正的 grouped subprocess 路径在运行时 NameError/行为漂移”的回归。
- 2026-03-14: `inspect.getmembers(module, inspect.isclass)` 会按属性名排序,所以 legacy `main.py` 组件发现若要保留声明顺序,必须遍历 `module.__dict__`;只删除后面的 `.sort()` 仍然不够。
+1
View File
@@ -40,6 +40,7 @@
- 2026-03-13: Real legacy plugins may still load through deep `astrbot.core.*` imports even when their public entrypoint only looks like `astrbot.api.*`. `astrbot_plugin_self_learning` hits `astrbot.core.utils.astrbot_path`, `astrbot.core.provider.*`, `astrbot.core.agent.message`, and `astrbot.core.db.po` during load; keep those deep-path shims minimal and whitelist-driven, but do not assume the `api` facade alone is enough.
- 2026-03-13: `ARCHITECTURE.md` and `refactor.md` are no longer a full source of truth for the current runtime/compat surface. The shipped code also includes `runtime.environment_groups`, `_session_waiter`, the controlled `src-new/astrbot` alias facade, compat hook execution, and extra DB capabilities such as `db.get_many` / `db.set_many` / `db.watch`. Verify architectural claims against code and tests before declaring drift or completeness.
- 2026-03-13: Duplicating private compat logic into a second `_legacy/` package added import-order risk and architectural noise. Keep one canonical set of top-level private compat modules (`_legacy_api.py`, `_legacy_runtime.py`, `_legacy_loader.py`, `_session_waiter.py`, `_shared_preferences.py`) while preserving public `astrbot_sdk.api`, `astrbot_sdk.compat`, and `src-new/astrbot` facades.
- 2026-03-14: `inspect.getmembers(module, inspect.isclass)` sorts legacy `main.py` classes alphabetically by attribute name. Preserving old-plugin declaration order requires iterating `module.__dict__` directly; deleting a later explicit `.sort()` is insufficient.
# 开发命令
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+14 -9
View File
@@ -84,6 +84,19 @@ def _prepare_legacy_package(package_name: str, plugin_dir: Path) -> None:
importlib.invalidate_caches()
def _iter_main_module_component_classes(module: types.ModuleType) -> list[type[Any]]:
component_classes: list[type[Any]] = []
for candidate in module.__dict__.values():
if not inspect.isclass(candidate):
continue
if candidate.__module__ != module.__name__:
continue
if not issubclass(candidate, Star) or candidate is Star:
continue
component_classes.append(candidate)
return component_classes
def load_legacy_main_component_classes(
*,
plugin_name: str,
@@ -99,15 +112,7 @@ def load_legacy_main_component_classes(
module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = module
spec.loader.exec_module(module)
component_classes: list[type[Any]] = []
for _, candidate in inspect.getmembers(module, inspect.isclass):
if candidate.__module__ != module.__name__:
continue
if not issubclass(candidate, Star) or candidate is Star:
continue
component_classes.append(candidate)
component_classes.sort(key=lambda cls: cls.__name__)
return component_classes
return _iter_main_module_component_classes(module)
def resolve_plugin_component_classes(
+35
View File
@@ -463,3 +463,38 @@ def resolve_plugin_lifecycle_hook(
if callable(hook):
return hook
return None
async def run_plugin_lifecycle(
instances: list[Any],
method_name: str,
context: Any,
) -> None:
"""执行插件实例列表的生命周期钩子。
对每个实例查找对应的生命周期方法,按签名决定是否注入 context,然后调用。
"""
for instance in instances:
hook = resolve_plugin_lifecycle_hook(instance, method_name)
if hook is None:
continue
args: list[Any] = []
try:
signature = inspect.signature(hook)
except (TypeError, ValueError):
signature = None
if signature is not None:
positional_params = [
parameter
for parameter in signature.parameters.values()
if parameter.kind
in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
]
if positional_params:
args.append(context)
result = hook(*args)
if inspect.isawaitable(result):
await result
+177
View File
@@ -0,0 +1,177 @@
"""旧版 API 兼容层 — 插件基类与注册装饰器。
这个模块承接旧 ``Star`` / ``CommandComponent`` / ``register`` 的实现,
供旧版插件在不修改代码的情况下继续运行。
依赖关系:
- ``_legacy_context`` 提供 ``LegacyContext``(单向依赖,本模块不被 ``_legacy_context`` 导入)
- ``_legacy_llm`` 提供 ``CompatLLMToolManager``
外部代码应通过 ``_legacy_api`` 聚合入口导入,而不是直接导入本模块。
"""
from __future__ import annotations
import inspect
from collections.abc import Callable
from pathlib import Path
from typing import Any
from ._legacy_context import LegacyContext
from ._legacy_llm import CompatLLMToolManager
from .star import Star
class StarTools:
"""旧版 ``StarTools`` 的最小兼容实现。"""
@staticmethod
def get_data_dir() -> Path:
frame = inspect.currentframe()
caller = frame.f_back if frame is not None else None
try:
while caller is not None:
caller_file = caller.f_globals.get("__file__")
if isinstance(caller_file, str) and caller_file:
data_dir = Path(caller_file).resolve().parent / "data"
data_dir.mkdir(parents=True, exist_ok=True)
return data_dir
caller = caller.f_back
finally:
del frame
data_dir = Path.cwd() / "data"
data_dir.mkdir(parents=True, exist_ok=True)
return data_dir
class LegacyStar(Star):
"""旧版 ``astrbot.api.star.Star`` 兼容基类。"""
def __init__(self, context: LegacyContext | None = None, config: Any | None = None):
self.context = context
if config is not None:
self.config = config
def _require_legacy_context(self) -> LegacyContext:
if self.context is None:
raise RuntimeError("LegacyStar 尚未绑定 compat Context")
return self.context
async def put_kv_data(self, key: str, value: Any) -> None:
await self._require_legacy_context().put_kv_data(key, value)
async def get_kv_data(self, key: str, default: Any = None) -> Any:
return await self._require_legacy_context().get_kv_data(key, default)
async def delete_kv_data(self, key: str) -> None:
await self._require_legacy_context().delete_kv_data(key)
async def send_message(self, session: str, message_chain: Any) -> None:
await self._require_legacy_context().send_message(session, message_chain)
async def llm_generate(
self,
chat_provider_id: str,
*args: Any,
**kwargs: Any,
) -> Any:
return await self._require_legacy_context().llm_generate(
chat_provider_id,
*args,
**kwargs,
)
async def tool_loop_agent(
self,
chat_provider_id: str,
*args: Any,
**kwargs: Any,
) -> Any:
return await self._require_legacy_context().tool_loop_agent(
chat_provider_id,
*args,
**kwargs,
)
async def add_llm_tools(self, *tools: Any) -> None:
await self._require_legacy_context().add_llm_tools(*tools)
def get_llm_tool_manager(self) -> CompatLLMToolManager:
return self._require_legacy_context().get_llm_tool_manager()
def activate_llm_tool(self, name: str) -> bool:
return self._require_legacy_context().activate_llm_tool(name)
def deactivate_llm_tool(self, name: str) -> bool:
return self._require_legacy_context().deactivate_llm_tool(name)
def register_llm_tool(
self,
name: str,
func_args: list[dict[str, Any]],
desc: str,
func_obj: Callable[..., Any],
) -> None:
self._require_legacy_context().register_llm_tool(
name,
func_args,
desc,
func_obj,
)
def unregister_llm_tool(self, name: str) -> None:
self._require_legacy_context().unregister_llm_tool(name)
def get_config(self) -> dict[str, Any]:
return self._require_legacy_context().get_config()
@classmethod
def __astrbot_is_new_star__(cls) -> bool:
return False
@classmethod
def _astrbot_create_legacy_context(cls, plugin_id: str) -> LegacyContext:
return LegacyContext(plugin_id)
class CommandComponent(LegacyStar):
@classmethod
def __astrbot_is_new_star__(cls) -> bool:
return False
@classmethod
def _astrbot_create_legacy_context(cls, plugin_id: str) -> LegacyContext:
# Loader 通过这个工厂拿到旧 Context,避免核心运行时直接依赖 compat 实现。
return LegacyContext(plugin_id)
def register(
name: str | None = None,
author: str | None = None,
desc: str | None = None,
version: str | None = None,
repo: str | None = None,
):
"""旧版插件元数据装饰器兼容入口。"""
metadata = {
"name": name,
"author": author,
"desc": desc,
"version": version,
"repo": repo,
}
def decorator(cls):
existing = getattr(cls, "__astrbot_plugin_metadata__", {})
setattr(
cls,
"__astrbot_plugin_metadata__",
{
**existing,
**{key: value for key, value in metadata.items() if value is not None},
},
)
return cls
return decorator
+2
View File
@@ -110,11 +110,13 @@ class InitializeOutput(_MessageBase):
Attributes:
peer: 接收方(核心)节点信息
protocol_version: 协商后的协议版本;未协商时可为空
capabilities: 核心提供的能力描述符列表
metadata: 扩展元数据
"""
peer: PeerInfo
protocol_version: str | None = None
capabilities: list[CapabilityDescriptor] = Field(default_factory=list)
metadata: dict[str, Any] = Field(default_factory=dict)
File diff suppressed because it is too large Load Diff
+312 -273
View File
@@ -98,8 +98,8 @@ from typing import Any
from ..errors import AstrBotError
from ..protocol.descriptors import (
BUILTIN_CAPABILITY_SCHEMAS,
CapabilityDescriptor,
RESERVED_CAPABILITY_PREFIXES,
CapabilityDescriptor,
SessionRef,
)
@@ -244,332 +244,371 @@ class CapabilityRouter:
collect_chunks=execution.collect_chunks,
)
# ------------------------------------------------------------------
# Built-in capability registration
# ------------------------------------------------------------------
def _register_builtin_capabilities(self) -> None:
def resolve_target(
payload: dict[str, Any],
) -> tuple[str, dict[str, Any] | None]:
target_payload = payload.get("target")
if isinstance(target_payload, dict):
target = SessionRef.model_validate(target_payload)
return target.session, target.to_payload()
return str(payload.get("session", "")), None
"""注册全部 18 条内建 capability。"""
self._register_llm_capabilities()
self._register_memory_capabilities()
self._register_db_capabilities()
self._register_platform_capabilities()
def builtin_descriptor(
name: str,
description: str,
*,
supports_stream: bool = False,
cancelable: bool = False,
) -> CapabilityDescriptor:
schema = BUILTIN_CAPABILITY_SCHEMAS[name]
return CapabilityDescriptor(
name=name,
description=description,
input_schema=copy.deepcopy(schema["input"]),
output_schema=copy.deepcopy(schema["output"]),
supports_stream=supports_stream,
cancelable=cancelable,
)
def _builtin_descriptor(
self,
name: str,
description: str,
*,
supports_stream: bool = False,
cancelable: bool = False,
) -> CapabilityDescriptor:
"""构建内建 capability 描述符,schema 从注册表读取。"""
schema = BUILTIN_CAPABILITY_SCHEMAS[name]
return CapabilityDescriptor(
name=name,
description=description,
input_schema=copy.deepcopy(schema["input"]),
output_schema=copy.deepcopy(schema["output"]),
supports_stream=supports_stream,
cancelable=cancelable,
)
async def llm_chat(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prompt = str(payload.get("prompt", ""))
return {"text": f"Echo: {prompt}"}
def _resolve_target(
self, payload: dict[str, Any]
) -> tuple[str, dict[str, Any] | None]:
"""从 payload 解析 session + target。"""
target_payload = payload.get("target")
if isinstance(target_payload, dict):
target = SessionRef.model_validate(target_payload)
return target.session, target.to_payload()
return str(payload.get("session", "")), None
async def llm_chat_raw(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prompt = str(payload.get("prompt", ""))
text = f"Echo: {prompt}"
return {
"text": text,
"usage": {
"input_tokens": len(prompt),
"output_tokens": len(text),
},
"finish_reason": "stop",
"tool_calls": [],
}
# ------------------------------------------------------------------
# LLM handlers
# ------------------------------------------------------------------
async def llm_stream(
_request_id: str,
payload: dict[str, Any],
token,
) -> AsyncIterator[dict[str, Any]]:
text = f"Echo: {str(payload.get('prompt', ''))}"
for char in text:
token.raise_if_cancelled()
await asyncio.sleep(0)
yield {"text": char}
async def _llm_chat(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prompt = str(payload.get("prompt", ""))
return {"text": f"Echo: {prompt}"}
async def memory_search(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
query = str(payload.get("query", ""))
items = [
{"key": key, "value": value}
for key, value in self.memory_store.items()
if query in key or query in json.dumps(value, ensure_ascii=False)
]
return {"items": items}
async def _llm_chat_raw(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prompt = str(payload.get("prompt", ""))
text = f"Echo: {prompt}"
return {
"text": text,
"usage": {
"input_tokens": len(prompt),
"output_tokens": len(text),
},
"finish_reason": "stop",
"tool_calls": [],
}
async def memory_save(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
if not isinstance(value, dict):
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
self.memory_store[key] = value
return {}
async def memory_get(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.memory_store.get(str(payload.get("key", "")))}
async def memory_delete(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.memory_store.pop(str(payload.get("key", "")), None)
return {}
async def db_get(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.db_store.get(str(payload.get("key", "")))}
async def db_set(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def db_delete(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
self.db_store.pop(key, None)
self._emit_db_change(op="delete", key=key, value=None)
return {}
async def db_list(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prefix = payload.get("prefix")
keys = sorted(self.db_store.keys())
if isinstance(prefix, str):
keys = [item for item in keys if item.startswith(prefix)]
return {"keys": keys}
async def db_get_many(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
return {"items": items}
async def db_set_many(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
items_payload = payload.get("items")
if not isinstance(items_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
for entry in items_payload:
if not isinstance(entry, dict):
raise AstrBotError.invalid_input(
"db.set_many 的 items 必须是 object 数组"
)
key = str(entry.get("key", ""))
value = entry.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def db_watch(
request_id: str, payload: dict[str, Any], _token
) -> StreamExecution:
prefix = payload.get("prefix")
prefix_value: str | None
if isinstance(prefix, str):
prefix_value = prefix
elif prefix is None:
prefix_value = None
else:
raise AstrBotError.invalid_input(
"db.watch 的 prefix 必须是 string 或 null"
)
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
async def iterator() -> AsyncIterator[dict[str, Any]]:
try:
while True:
yield await queue.get()
finally:
self._db_watch_subscriptions.pop(request_id, None)
return StreamExecution(
iterator=iterator(),
finalize=lambda _chunks: {},
collect_chunks=False,
)
async def platform_send(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = resolve_target(payload)
text = str(payload.get("text", ""))
message_id = f"msg_{len(self.sent_messages) + 1}"
sent = {"message_id": message_id, "session": session, "text": text}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def platform_send_image(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = resolve_target(payload)
image_url = str(payload.get("image_url", ""))
message_id = f"img_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"image_url": image_url,
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def platform_send_chain(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = resolve_target(payload)
chain = payload.get("chain")
if not isinstance(chain, list) or not all(
isinstance(item, dict) for item in chain
):
raise AstrBotError.invalid_input(
"platform.send_chain 的 chain 必须是 object 数组"
)
message_id = f"chain_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"chain": [dict(item) for item in chain],
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def platform_get_members(
_request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, _target = resolve_target(payload)
return {
"members": [
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
]
}
async def _llm_stream(
self,
_request_id: str,
payload: dict[str, Any],
token,
) -> AsyncIterator[dict[str, Any]]: # type: ignore[override]
text = f"Echo: {str(payload.get('prompt', ''))}"
for char in text:
token.raise_if_cancelled()
await asyncio.sleep(0)
yield {"text": char}
def _register_llm_capabilities(self) -> None:
self.register(
builtin_descriptor("llm.chat", "发送对话请求,返回文本"),
call_handler=llm_chat,
self._builtin_descriptor("llm.chat", "发送对话请求,返回文本"),
call_handler=self._llm_chat,
)
self.register(
builtin_descriptor("llm.chat_raw", "发送对话请求,返回完整响应"),
call_handler=llm_chat_raw,
self._builtin_descriptor("llm.chat_raw", "发送对话请求,返回完整响应"),
call_handler=self._llm_chat_raw,
)
self.register(
builtin_descriptor(
self._builtin_descriptor(
"llm.stream_chat",
"流式对话",
supports_stream=True,
cancelable=True,
),
stream_handler=llm_stream,
stream_handler=self._llm_stream,
finalize=lambda chunks: {
"text": "".join(item.get("text", "") for item in chunks)
},
)
# ------------------------------------------------------------------
# Memory handlers
# ------------------------------------------------------------------
async def _memory_search(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
query = str(payload.get("query", ""))
items = [
{"key": key, "value": value}
for key, value in self.memory_store.items()
if query in key or query in json.dumps(value, ensure_ascii=False)
]
return {"items": items}
async def _memory_save(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
if not isinstance(value, dict):
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
self.memory_store[key] = value
return {}
async def _memory_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.memory_store.get(str(payload.get("key", "")))}
async def _memory_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.memory_store.pop(str(payload.get("key", "")), None)
return {}
def _register_memory_capabilities(self) -> None:
self.register(
builtin_descriptor("memory.search", "搜索记忆"),
call_handler=memory_search,
self._builtin_descriptor("memory.search", "搜索记忆"),
call_handler=self._memory_search,
)
self.register(
builtin_descriptor("memory.save", "保存记忆"),
call_handler=memory_save,
self._builtin_descriptor("memory.save", "保存记忆"),
call_handler=self._memory_save,
)
self.register(
builtin_descriptor("memory.get", "读取单条记忆"),
call_handler=memory_get,
self._builtin_descriptor("memory.get", "读取单条记忆"),
call_handler=self._memory_get,
)
self.register(
builtin_descriptor("memory.delete", "删除记忆"),
call_handler=memory_delete,
self._builtin_descriptor("memory.delete", "删除记忆"),
call_handler=self._memory_delete,
)
# ------------------------------------------------------------------
# DB handlers
# ------------------------------------------------------------------
async def _db_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.db_store.get(str(payload.get("key", "")))}
async def _db_set(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
self.db_store.pop(key, None)
self._emit_db_change(op="delete", key=key, value=None)
return {}
async def _db_list(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prefix = payload.get("prefix")
keys = sorted(self.db_store.keys())
if isinstance(prefix, str):
keys = [item for item in keys if item.startswith(prefix)]
return {"keys": keys}
async def _db_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
return {"items": items}
async def _db_set_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
items_payload = payload.get("items")
if not isinstance(items_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
for entry in items_payload:
if not isinstance(entry, dict):
raise AstrBotError.invalid_input(
"db.set_many 的 items 必须是 object 数组"
)
key = str(entry.get("key", ""))
value = entry.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_watch(
self, request_id: str, payload: dict[str, Any], _token
) -> StreamExecution:
prefix = payload.get("prefix")
prefix_value: str | None
if isinstance(prefix, str):
prefix_value = prefix
elif prefix is None:
prefix_value = None
else:
raise AstrBotError.invalid_input("db.watch 的 prefix 必须是 string 或 null")
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
async def iterator() -> AsyncIterator[dict[str, Any]]:
try:
while True:
yield await queue.get()
finally:
self._db_watch_subscriptions.pop(request_id, None)
return StreamExecution(
iterator=iterator(),
finalize=lambda _chunks: {},
collect_chunks=False,
)
def _register_db_capabilities(self) -> None:
self.register(
self._builtin_descriptor("db.get", "读取 KV"),
call_handler=self._db_get,
)
self.register(
builtin_descriptor("db.get", "读取 KV"),
call_handler=db_get,
self._builtin_descriptor("db.set", "写入 KV"),
call_handler=self._db_set,
)
self.register(
builtin_descriptor("db.set", "写入 KV"),
call_handler=db_set,
self._builtin_descriptor("db.delete", "删除 KV"),
call_handler=self._db_delete,
)
self.register(
builtin_descriptor("db.delete", "删除 KV"),
call_handler=db_delete,
self._builtin_descriptor("db.list", "列出 KV"),
call_handler=self._db_list,
)
self.register(
builtin_descriptor("db.list", "列出 KV"),
call_handler=db_list,
self._builtin_descriptor("db.get_many", "批量读取 KV"),
call_handler=self._db_get_many,
)
self.register(
builtin_descriptor("db.get_many", "批量读取 KV"),
call_handler=db_get_many,
self._builtin_descriptor("db.set_many", "批量写入 KV"),
call_handler=self._db_set_many,
)
self.register(
builtin_descriptor("db.set_many", "批量写入 KV"),
call_handler=db_set_many,
)
self.register(
builtin_descriptor(
self._builtin_descriptor(
"db.watch",
"订阅 KV 变更",
supports_stream=True,
cancelable=True,
),
stream_handler=db_watch,
stream_handler=self._db_watch,
)
# ------------------------------------------------------------------
# Platform handlers
# ------------------------------------------------------------------
async def _platform_send(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
text = str(payload.get("text", ""))
message_id = f"msg_{len(self.sent_messages) + 1}"
sent = {"message_id": message_id, "session": session, "text": text}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_image(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
image_url = str(payload.get("image_url", ""))
message_id = f"img_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"image_url": image_url,
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_chain(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
chain = payload.get("chain")
if not isinstance(chain, list) or not all(
isinstance(item, dict) for item in chain
):
raise AstrBotError.invalid_input(
"platform.send_chain 的 chain 必须是 object 数组"
)
message_id = f"chain_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"chain": [dict(item) for item in chain],
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_get_members(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, _target = self._resolve_target(payload)
return {
"members": [
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
]
}
def _register_platform_capabilities(self) -> None:
self.register(
self._builtin_descriptor("platform.send", "发送消息"),
call_handler=self._platform_send,
)
self.register(
builtin_descriptor("platform.send", "发送消息"),
call_handler=platform_send,
self._builtin_descriptor("platform.send_image", "发送图片"),
call_handler=self._platform_send_image,
)
self.register(
builtin_descriptor("platform.send_image", "发送图片"),
call_handler=platform_send_image,
self._builtin_descriptor("platform.send_chain", "发送消息链"),
call_handler=self._platform_send_chain,
)
self.register(
builtin_descriptor("platform.send_chain", "发送消息链"),
call_handler=platform_send_chain,
)
self.register(
builtin_descriptor("platform.get_members", "获取群成员"),
call_handler=platform_get_members,
self._builtin_descriptor("platform.get_members", "获取群成员"),
call_handler=self._platform_get_members,
)
# ------------------------------------------------------------------
# Schema validation
# ------------------------------------------------------------------
def _validate_schema(
self,
schema: dict[str, Any] | None,
+4 -29
View File
@@ -99,7 +99,7 @@ from typing import Any
import yaml
from .._legacy_loader import (
build_legacy_manifest,
PLUGIN_MANIFEST_FILE,
load_legacy_main_component_classes,
load_plugin_manifest_payload,
looks_like_legacy_plugin,
@@ -109,7 +109,6 @@ from .._legacy_runtime import (
LegacyRuntimeAdapter,
build_capability_legacy_runtime,
build_handler_legacy_runtime,
create_legacy_component_context,
finalize_legacy_component_instance,
is_new_star_component,
plan_legacy_component_construction,
@@ -119,15 +118,12 @@ from ..decorators import get_capability_meta, get_handler_meta
from ..protocol.descriptors import CapabilityDescriptor, HandlerDescriptor
from .environment_groups import (
EnvironmentGroup,
EnvironmentPlanResult,
EnvironmentPlanner,
EnvironmentPlanResult,
GroupEnvironmentManager,
)
STATE_FILE_NAME = ".astrbot-worker-state.json"
PLUGIN_MANIFEST_FILE = "plugin.yaml"
LEGACY_METADATA_FILE = "metadata.yaml"
LEGACY_MAIN_FILE = "main.py"
CONFIG_SCHEMA_FILE = "_conf_schema.json"
LEGACY_MAIN_MANIFEST_KEY = "__legacy_main__"
PLUGIN_METADATA_ATTR = "__astrbot_plugin_metadata__"
@@ -207,14 +203,6 @@ class LoadedPlugin:
instances: list[Any] = field(default_factory=list)
def _is_new_star_component(component_cls: Any) -> bool:
return is_new_star_component(component_cls)
def _create_legacy_context(component_cls: Any, plugin_name: str) -> Any:
return create_legacy_component_context(component_cls, plugin_name)
def _iter_handler_names(instance: Any) -> list[str]:
handler_names = getattr(instance.__class__, "__handlers__", ())
if handler_names:
@@ -277,19 +265,6 @@ def _read_requirements_text(path: Path) -> str:
return path.read_text(encoding="utf-8")
def _looks_like_legacy_plugin(plugin_dir: Path) -> bool:
return looks_like_legacy_plugin(plugin_dir)
def _build_legacy_manifest(plugin_dir: Path) -> tuple[Path, dict[str, Any]]:
return build_legacy_manifest(
plugin_dir,
read_yaml=_read_yaml,
default_python_version=_default_python_version(),
manifest_flag_key=LEGACY_MAIN_MANIFEST_KEY,
)
def _plugin_config_dir(plugin_dir: Path) -> Path:
if plugin_dir.parent.name == "plugins" and plugin_dir.parent.parent.exists():
return plugin_dir.parent.parent / "config"
@@ -479,7 +454,7 @@ def discover_plugins(plugins_dir: Path) -> PluginDiscoveryResult:
if not entry.is_dir() or entry.name.startswith("."):
continue
manifest_path = entry / PLUGIN_MANIFEST_FILE
if not manifest_path.exists() and not _looks_like_legacy_plugin(entry):
if not manifest_path.exists() and not looks_like_legacy_plugin(entry):
continue
plugin: PluginSpec | None = None
try:
@@ -629,7 +604,7 @@ def load_plugin(plugin: PluginSpec) -> LoadedPlugin:
plugin_config = _load_plugin_config(plugin)
for component_cls in _plugin_component_classes(plugin):
legacy_context = None
if _is_new_star_component(component_cls):
if is_new_star_component(component_cls):
instance = component_cls()
else:
construction = plan_legacy_component_construction(
+109 -6
View File
@@ -85,7 +85,7 @@ from __future__ import annotations
import asyncio
import inspect
from collections.abc import AsyncIterator, Awaitable, Callable
from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
from typing import Any
from ..context import CancelToken
@@ -109,6 +109,61 @@ InvokeHandler = Callable[
]
CancelHandler = Callable[[str], Awaitable[None]]
SUPPORTED_PROTOCOL_VERSIONS_METADATA_KEY = "supported_protocol_versions"
NEGOTIATED_PROTOCOL_VERSION_METADATA_KEY = "negotiated_protocol_version"
def _dedupe_protocol_versions(
versions: Sequence[str] | None, *, preferred_version: str
) -> list[str]:
ordered_versions: list[str] = [preferred_version]
if versions is not None:
ordered_versions.extend(versions)
deduped: list[str] = []
for version in ordered_versions:
if not isinstance(version, str) or not version:
continue
if version not in deduped:
deduped.append(version)
return deduped
def _parse_protocol_version(version: str) -> tuple[int, int] | None:
major, dot, minor = version.partition(".")
if not dot or not major.isdigit() or not minor.isdigit():
return None
return int(major), int(minor)
def _select_negotiated_protocol_version(
requested_version: str,
remote_metadata: dict[str, Any],
local_supported_versions: Sequence[str],
) -> str | None:
if requested_version in local_supported_versions:
return requested_version
requested_key = _parse_protocol_version(requested_version)
if requested_key is None:
return None
remote_supported = remote_metadata.get(SUPPORTED_PROTOCOL_VERSIONS_METADATA_KEY)
if not isinstance(remote_supported, (list, tuple)):
return None
local_supported_set = set(local_supported_versions)
compatible_versions: list[tuple[tuple[int, int], str]] = []
for version in remote_supported:
if not isinstance(version, str) or version not in local_supported_set:
continue
parsed_version = _parse_protocol_version(version)
if parsed_version is None:
continue
if parsed_version[0] != requested_key[0] or parsed_version > requested_key:
continue
compatible_versions.append((parsed_version, version))
if not compatible_versions:
return None
compatible_versions.sort(reverse=True)
return compatible_versions[0][1]
class Peer:
"""表示协议连接中的一个对等端。
@@ -124,17 +179,24 @@ class Peer:
transport,
peer_info: PeerInfo,
protocol_version: str = "1.0",
supported_protocol_versions: Sequence[str] | None = None,
) -> None:
"""创建一个协议对等端实例。
Args:
transport: 底层传输实现,负责发送字符串消息并回调入站消息。
peer_info: 当前端点对外声明的身份信息。
protocol_version: 当前端点支持的协议版本,用于初始化握手校验。
protocol_version: 当前端点首选的协议版本,用于初始化握手。
supported_protocol_versions: 当前端点可接受的协议版本列表。
"""
self.transport = transport
self.peer_info = peer_info
self.protocol_version = protocol_version
self.supported_protocol_versions = _dedupe_protocol_versions(
supported_protocol_versions,
preferred_version=protocol_version,
)
self.negotiated_protocol_version: str | None = None
self.remote_peer: PeerInfo | None = None
self.remote_handlers = []
self.remote_provided_capabilities = []
@@ -175,6 +237,7 @@ class Peer:
self._closed.clear()
self._unusable = False
self._stopping = False
self.negotiated_protocol_version = None
self._remote_initialized.clear()
self.transport.set_message_handler(self._handle_raw_message)
await self.transport.start()
@@ -273,6 +336,10 @@ class Peer:
"""
self._ensure_usable()
request_id = self._next_id()
handshake_metadata = dict(metadata or {})
handshake_metadata[SUPPORTED_PROTOCOL_VERSIONS_METADATA_KEY] = list(
self.supported_protocol_versions
)
future: asyncio.Future[ResultMessage] = (
asyncio.get_running_loop().create_future()
)
@@ -284,7 +351,7 @@ class Peer:
peer=self.peer_info,
handlers=list(handlers),
provided_capabilities=list(provided_capabilities or []),
metadata=metadata or {},
metadata=handshake_metadata,
)
)
result = await future
@@ -297,10 +364,25 @@ class Peer:
result.error.model_dump() if result.error else {}
)
output = InitializeOutput.model_validate(result.output)
negotiated_protocol_version = (
output.protocol_version
or output.metadata.get(NEGOTIATED_PROTOCOL_VERSION_METADATA_KEY)
or self.protocol_version
)
if (
not isinstance(negotiated_protocol_version, str)
or negotiated_protocol_version not in self.supported_protocol_versions
):
self._unusable = True
await self.stop()
raise AstrBotError.protocol_version_mismatch(
f"对端返回了当前端点不支持的协商协议版本:{negotiated_protocol_version}"
)
self.remote_peer = output.peer
self.remote_capabilities = output.capabilities
self.remote_capability_map = {item.name: item for item in output.capabilities}
self.remote_metadata = output.metadata
self.negotiated_protocol_version = negotiated_protocol_version
self._remote_initialized.set()
return output
@@ -458,7 +540,7 @@ class Peer:
self.remote_provided_capability_map = {
item.name: item for item in message.provided_capabilities
}
self.remote_metadata = message.metadata
self.remote_metadata = dict(message.metadata)
if self._initialize_handler is None:
await self._reject_initialize(
message,
@@ -466,16 +548,37 @@ class Peer:
)
return
if message.protocol_version != self.protocol_version:
negotiated_protocol_version = _select_negotiated_protocol_version(
message.protocol_version,
self.remote_metadata,
self.supported_protocol_versions,
)
if negotiated_protocol_version is None:
supported_versions = ", ".join(self.supported_protocol_versions)
await self._reject_initialize(
message,
AstrBotError.protocol_version_mismatch(
f"服务端支持协议版本 {self.protocol_version},客户端请求版本 {message.protocol_version}"
"服务端支持协议版本 "
f"{supported_versions},客户端请求版本 {message.protocol_version}"
),
)
return
self.negotiated_protocol_version = negotiated_protocol_version
self.remote_metadata[NEGOTIATED_PROTOCOL_VERSION_METADATA_KEY] = (
negotiated_protocol_version
)
output = await self._initialize_handler(message)
response_metadata = dict(output.metadata)
response_metadata[NEGOTIATED_PROTOCOL_VERSION_METADATA_KEY] = (
negotiated_protocol_version
)
output = output.model_copy(
update={
"protocol_version": negotiated_protocol_version,
"metadata": response_metadata,
}
)
await self._send(
ResultMessage(
id=message.id,
+704
View File
@@ -0,0 +1,704 @@
"""Supervisor 端运行时:SupervisorRuntime 管理多个 Worker 进程,WorkerSession 封装与单个 Worker 的通信。
架构层次:
AstrBot Core (Python)
|
v
SupervisorRuntime (管理多插件)
|
+-- WorkerSession (插件 A) -- StdioTransport -- PluginWorkerRuntime (子进程)
|
+-- WorkerSession (插件 B) -- StdioTransport -- PluginWorkerRuntime (子进程)
|
+-- WorkerSession (插件 C) -- StdioTransport -- PluginWorkerRuntime (子进程)
核心类:
SupervisorRuntime: 监管者运行时
- 发现并加载所有插件
- 为每个插件启动 Worker 进程
- 聚合所有 handler 并向 Core 注册
- 路由 Core 的调用请求到对应 Worker
- 处理 Worker 进程崩溃和重连
- handler ID 冲突检测和警告
WorkerSession: Worker 会话
- 管理单个插件 Worker 进程
- 通过 Peer 与 Worker 通信
- 提供 invoke_handler 和 cancel 方法
- 处理连接关闭回调
- 自动清理已注册的 handlers
信号处理:
- SIGTERM: 设置 stop_event,触发优雅关闭
- SIGINT: 设置 stop_event,触发优雅关闭
"""
from __future__ import annotations
import asyncio
import os
import signal
import sys
from collections.abc import Callable
from pathlib import Path
from typing import IO, Any
from loguru import logger
from ..errors import AstrBotError
from ..protocol.descriptors import CapabilityDescriptor
from ..protocol.messages import EventMessage, InitializeOutput, PeerInfo
from .capability_router import CapabilityRouter, StreamExecution
from .environment_groups import EnvironmentGroup
from .loader import (
PluginEnvironmentManager,
PluginSpec,
discover_plugins,
)
from .peer import Peer
from .transport import StdioTransport
__all__ = [
"SupervisorRuntime",
"WorkerSession",
"_install_signal_handlers",
"_prepare_stdio_transport",
"_sdk_source_dir",
"_wait_for_shutdown",
]
def _install_signal_handlers(stop_event: asyncio.Event) -> None:
loop = asyncio.get_running_loop()
for sig in (signal.SIGTERM, signal.SIGINT):
try:
loop.add_signal_handler(sig, stop_event.set)
except NotImplementedError:
logger.debug("Signal handlers are not supported for {}", sig)
def _prepare_stdio_transport(
stdin: IO[str] | None,
stdout: IO[str] | None,
) -> tuple[IO[str], IO[str], IO[str] | None]:
if stdin is not None and stdout is not None:
return stdin, stdout, None
transport_stdin = stdin or sys.stdin
transport_stdout = stdout or sys.stdout
original_stdout = sys.stdout
sys.stdout = sys.stderr
return transport_stdin, transport_stdout, original_stdout
def _sdk_source_dir(repo_root: Path) -> Path:
candidate = repo_root.resolve() / "src-new"
if (candidate / "astrbot_sdk").exists():
return candidate
return Path(__file__).resolve().parents[2]
async def _wait_for_shutdown(peer: Peer, stop_event: asyncio.Event) -> None:
stop_waiter = asyncio.create_task(stop_event.wait())
transport_waiter = asyncio.create_task(peer.wait_closed())
done, pending = await asyncio.wait(
{stop_waiter, transport_waiter},
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
for task in done:
if not task.cancelled():
task.result()
def _plugin_name_from_handler_id(handler_id: str) -> str:
if ":" in handler_id:
return handler_id.split(":", 1)[0]
return handler_id
class WorkerSession:
def __init__(
self,
*,
plugin: PluginSpec | None = None,
group: EnvironmentGroup | None = None,
repo_root: Path,
env_manager: PluginEnvironmentManager,
capability_router: CapabilityRouter,
on_closed: Callable[[], None] | None = None,
) -> None:
if plugin is None and group is None:
raise ValueError("WorkerSession requires either plugin or group")
self.group = group
self.plugins = list(group.plugins) if group is not None else [plugin]
self.plugin = plugin or self.plugins[0]
self.group_id = group.id if group is not None else self.plugin.name
self.repo_root = repo_root.resolve()
self.env_manager = env_manager
self.capability_router = capability_router
self.on_closed = on_closed
self.peer: Peer | None = None
self.handlers = []
self.provided_capabilities: list[CapabilityDescriptor] = []
self.loaded_plugins: list[str] = []
self.skipped_plugins: dict[str, str] = {}
self.capability_sources: dict[str, str] = {}
self._connection_watch_task: asyncio.Task[None] | None = None
async def start(self) -> None:
python_path, command, cwd = self._worker_command()
repo_src_dir = str(_sdk_source_dir(self.repo_root))
env = os.environ.copy()
existing_pythonpath = env.get("PYTHONPATH")
env["PYTHONPATH"] = (
f"{repo_src_dir}{os.pathsep}{existing_pythonpath}"
if existing_pythonpath
else repo_src_dir
)
env.setdefault("PYTHONIOENCODING", "utf-8")
env.setdefault("PYTHONUTF8", "1")
transport = StdioTransport(
command=command,
cwd=cwd,
env=env,
)
self.peer = Peer(
transport=transport,
peer_info=PeerInfo(name="astrbot-core", role="core", version="v4"),
)
self.peer.set_initialize_handler(self._handle_initialize)
self.peer.set_invoke_handler(self._handle_capability_invoke)
try:
await self.peer.start()
# 同时监听初始化完成和连接关闭,避免 worker 崩溃时等满超时
init_task = asyncio.create_task(
self.peer.wait_until_remote_initialized(timeout=None)
)
closed_task = asyncio.create_task(self.peer.wait_closed())
done, pending = await asyncio.wait(
{init_task, closed_task},
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
if closed_task in done:
raise RuntimeError(f"worker 组 {self.group_id} 在初始化阶段退出")
self.handlers = list(self.peer.remote_handlers)
self.provided_capabilities = list(self.peer.remote_provided_capabilities)
metadata = dict(self.peer.remote_metadata)
remote_loaded_plugins = metadata.get("loaded_plugins")
if isinstance(remote_loaded_plugins, list):
self.loaded_plugins = [
plugin_name
for plugin_name in remote_loaded_plugins
if isinstance(plugin_name, str)
]
else:
self.loaded_plugins = [plugin.name for plugin in self.plugins]
remote_skipped_plugins = metadata.get("skipped_plugins")
if isinstance(remote_skipped_plugins, dict):
self.skipped_plugins = {
str(plugin_name): str(reason)
for plugin_name, reason in remote_skipped_plugins.items()
}
remote_capability_sources = metadata.get("capability_sources")
if isinstance(remote_capability_sources, dict):
self.capability_sources = {
str(capability_name): str(plugin_name)
for capability_name, plugin_name in remote_capability_sources.items()
}
except Exception:
await self.stop()
raise
def _worker_command(self) -> tuple[Path, list[str], str]:
if self.group is not None:
prepare_group = getattr(self.env_manager, "prepare_group_environment", None)
if callable(prepare_group):
python_path = prepare_group(self.group)
else:
python_path = self.env_manager.prepare_environment(self.plugins[0])
return (
python_path,
[
str(python_path),
"-m",
"astrbot_sdk",
"worker",
"--group-metadata",
str(self.group.metadata_path),
],
str(self.repo_root),
)
python_path = self.env_manager.prepare_environment(self.plugin)
return (
python_path,
[
str(python_path),
"-m",
"astrbot_sdk",
"worker",
"--plugin-dir",
str(self.plugin.plugin_dir),
],
str(self.plugin.plugin_dir),
)
def start_close_watch(self) -> None:
if (
self.on_closed is None
or self.peer is None
or self._connection_watch_task is not None
):
return
self._connection_watch_task = asyncio.create_task(self._watch_connection())
async def _watch_connection(self) -> None:
"""监听 Worker 连接关闭,触发清理回调"""
try:
if self.peer is not None:
await self.peer.wait_closed()
if self.on_closed is not None:
try:
self.on_closed()
except Exception:
logger.exception(
"on_closed callback failed for worker group {}", self.group_id
)
finally:
current_task = asyncio.current_task()
if self._connection_watch_task is current_task:
self._connection_watch_task = None
async def stop(self) -> None:
if self.peer is not None:
await self.peer.stop()
async def invoke_handler(
self,
handler_id: str,
event_payload: dict[str, Any],
*,
request_id: str,
) -> dict[str, Any]:
if self.peer is None:
raise RuntimeError("worker session is not running")
return await self.peer.invoke(
"handler.invoke",
{
"handler_id": handler_id,
"event": event_payload,
},
request_id=request_id,
)
async def invoke_capability(
self,
capability_name: str,
payload: dict[str, Any],
*,
request_id: str,
) -> dict[str, Any]:
if self.peer is None:
raise RuntimeError("worker session is not running")
return await self.peer.invoke(
capability_name,
payload,
request_id=request_id,
)
async def invoke_capability_stream(
self,
capability_name: str,
payload: dict[str, Any],
*,
request_id: str,
):
if self.peer is None:
raise RuntimeError("worker session is not running")
event_stream = await self.peer.invoke_stream(
capability_name,
payload,
request_id=request_id,
include_completed=True,
)
async for event in event_stream:
yield event
async def cancel(self, request_id: str) -> None:
if self.peer is None:
return
await self.peer.cancel(request_id)
async def _handle_initialize(self, _message) -> InitializeOutput:
return InitializeOutput(
peer=PeerInfo(name="astrbot-supervisor", role="core", version="v4"),
capabilities=self.capability_router.descriptors(),
metadata={
"group_id": self.group_id,
"plugins": [plugin.name for plugin in self.plugins],
},
)
async def _handle_capability_invoke(self, message, cancel_token):
return await self.capability_router.execute(
message.capability,
message.input,
stream=message.stream,
cancel_token=cancel_token,
request_id=message.id,
)
def describe(self) -> dict[str, Any]:
return {
"group_id": self.group_id,
"plugins": [plugin.name for plugin in self.plugins],
"loaded_plugins": list(self.loaded_plugins),
"skipped_plugins": dict(self.skipped_plugins),
}
class SupervisorRuntime:
def __init__(
self,
*,
transport,
plugins_dir: Path,
env_manager: PluginEnvironmentManager | None = None,
) -> None:
self.transport = transport
self.plugins_dir = plugins_dir.resolve()
self.repo_root = Path(__file__).resolve().parents[3]
self.env_manager = env_manager or PluginEnvironmentManager(self.repo_root)
self.capability_router = CapabilityRouter()
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name="astrbot-supervisor", role="plugin", version="v4"),
)
self.peer.set_invoke_handler(self._handle_upstream_invoke)
self.peer.set_cancel_handler(self._handle_upstream_cancel)
self.worker_sessions: dict[str, WorkerSession] = {}
self.handler_to_worker: dict[str, WorkerSession] = {}
self.capability_to_worker: dict[str, WorkerSession] = {}
self.plugin_to_worker_session: dict[str, WorkerSession] = {}
self._handler_sources: dict[str, str] = {} # handler_id -> plugin_name
self._capability_sources: dict[str, str] = {} # capability_name -> plugin_name
self.active_requests: dict[str, WorkerSession] = {}
self.loaded_plugins: list[str] = []
self.skipped_plugins: dict[str, str] = {}
self._register_internal_capabilities()
def _register_internal_capabilities(self) -> None:
self.capability_router.register(
CapabilityDescriptor(
name="handler.invoke",
description="框架内部:转发到插件 handler",
input_schema={
"type": "object",
"properties": {
"handler_id": {"type": "string"},
"event": {"type": "object"},
},
"required": ["handler_id", "event"],
},
output_schema={
"type": "object",
"properties": {},
"required": [],
},
cancelable=True,
),
call_handler=self._route_handler_invoke,
exposed=False,
)
def _register_handler(
self, handler, session: WorkerSession, plugin_name: str
) -> None:
"""注册 handler,处理冲突时输出警告。
Args:
handler: Handler 描述符
session: Worker 会话
plugin_name: 插件名称
"""
handler_id = handler.id
existing_plugin = self._handler_sources.get(handler_id)
if existing_plugin is not None:
logger.warning(
f"Handler ID 冲突:'{handler_id}' 已被插件 '{existing_plugin}' 注册,"
f"现在被插件 '{plugin_name}' 覆盖。"
)
self.handler_to_worker[handler_id] = session
self._handler_sources[handler_id] = plugin_name
def _register_plugin_capability(
self,
descriptor: CapabilityDescriptor,
session: WorkerSession,
plugin_name: str,
) -> None:
capability_name = descriptor.name
if self.capability_router.contains(capability_name):
logger.warning(
"Capability 名称冲突:'{}' 已存在,跳过插件 '{}' 的注册。",
capability_name,
plugin_name,
# TODO: 更好的解决方案?
)
return
self.capability_router.register(
descriptor.model_copy(deep=True),
call_handler=self._make_plugin_capability_caller(session, capability_name),
stream_handler=(
self._make_plugin_capability_streamer(session, capability_name)
if descriptor.supports_stream
else None
),
)
self.capability_to_worker[capability_name] = session
self._capability_sources[capability_name] = plugin_name
def _make_plugin_capability_caller(
self,
session: WorkerSession,
capability_name: str,
):
async def call_handler(
request_id: str,
payload: dict[str, Any],
_cancel_token,
) -> dict[str, Any]:
self.active_requests[request_id] = session
try:
return await session.invoke_capability(
capability_name,
payload,
request_id=request_id,
)
finally:
self.active_requests.pop(request_id, None)
return call_handler
def _make_plugin_capability_streamer(
self,
session: WorkerSession,
capability_name: str,
):
async def stream_handler(
request_id: str,
payload: dict[str, Any],
_cancel_token,
):
completed_output: dict[str, Any] = {}
async def iterator():
self.active_requests[request_id] = session
try:
async for event in session.invoke_capability_stream(
capability_name,
payload,
request_id=request_id,
):
if not isinstance(event, EventMessage):
raise AstrBotError.protocol_error(
"插件 worker 返回了非法的流式事件"
)
if event.phase == "delta":
yield event.data or {}
continue
if event.phase == "completed":
completed_output.clear()
completed_output.update(event.output or {})
finally:
self.active_requests.pop(request_id, None)
return StreamExecution(
iterator=iterator(),
finalize=lambda chunks: completed_output or {"items": chunks},
)
return stream_handler
async def start(self) -> None:
discovery = discover_plugins(self.plugins_dir)
self.skipped_plugins = dict(discovery.skipped_plugins)
plan_result = self.env_manager.plan(discovery.plugins)
self.skipped_plugins.update(plan_result.skipped_plugins)
try:
planned_sessions: list[WorkerSession] = []
if plan_result.groups:
for group in plan_result.groups:
planned_sessions.append(
WorkerSession(
group=group,
repo_root=self.repo_root,
env_manager=self.env_manager,
capability_router=self.capability_router,
on_closed=lambda group_id=group.id: (
self._handle_worker_closed(group_id)
),
)
)
else:
for plugin in plan_result.plugins:
planned_sessions.append(
WorkerSession(
plugin=plugin,
repo_root=self.repo_root,
env_manager=self.env_manager,
capability_router=self.capability_router,
on_closed=lambda plugin_name=plugin.name: (
self._handle_worker_closed(plugin_name)
),
)
)
for session in planned_sessions:
try:
await session.start()
except Exception as exc:
for plugin in session.plugins:
self.skipped_plugins[plugin.name] = str(exc)
await session.stop()
continue
self.worker_sessions[session.group_id] = session
self.skipped_plugins.update(session.skipped_plugins)
for plugin_name in session.loaded_plugins:
self.plugin_to_worker_session[plugin_name] = session
if plugin_name not in self.loaded_plugins:
self.loaded_plugins.append(plugin_name)
for handler in session.handlers:
self._register_handler(
handler,
session,
_plugin_name_from_handler_id(handler.id),
)
for descriptor in session.provided_capabilities:
plugin_name = session.capability_sources.get(descriptor.name)
if plugin_name is None and len(session.loaded_plugins) == 1:
plugin_name = session.loaded_plugins[0]
if plugin_name is None:
plugin_name = session.group_id
self._register_plugin_capability(descriptor, session, plugin_name)
session.start_close_watch()
aggregated_handlers = list(self.handler_to_worker.keys())
logger.info(
"Loaded plugins: {}", ", ".join(sorted(self.loaded_plugins)) or "none"
)
await self.peer.start()
await self.peer.initialize(
[
handler
for session in self.worker_sessions.values()
for handler in session.handlers
],
provided_capabilities=self.capability_router.descriptors(),
metadata={
"plugins": sorted(self.loaded_plugins),
"skipped_plugins": self.skipped_plugins,
"aggregated_handler_ids": aggregated_handlers,
"worker_groups": [
session.describe() for session in self.worker_sessions.values()
],
"worker_group_count": len(self.worker_sessions),
},
)
except Exception:
await self.stop()
raise
def _handle_worker_closed(self, group_id: str) -> None:
"""Worker 连接关闭时的清理回调"""
session = self.worker_sessions.pop(group_id, None)
if session is None:
return
# 从 handler_to_worker 中移除该插件注册的 handlers(仅当来源仍为此插件时)
for handler in session.handlers:
source_plugin = self._handler_sources.get(handler.id)
if source_plugin == _plugin_name_from_handler_id(handler.id) or (
source_plugin == group_id
):
self.handler_to_worker.pop(handler.id, None)
self._handler_sources.pop(handler.id, None)
for descriptor in session.provided_capabilities:
source_plugin = self._capability_sources.get(descriptor.name)
capability_plugin = session.capability_sources.get(descriptor.name)
if source_plugin == capability_plugin or (
capability_plugin is None
and (
source_plugin == group_id or source_plugin in session.loaded_plugins
)
):
self.capability_to_worker.pop(descriptor.name, None)
self._capability_sources.pop(descriptor.name, None)
self.capability_router.unregister(descriptor.name)
session_loaded_plugins = getattr(session, "loaded_plugins", None)
if not isinstance(session_loaded_plugins, list):
session_loaded_plugins = [group_id]
for plugin_name in session_loaded_plugins:
if plugin_name in self.loaded_plugins:
self.loaded_plugins.remove(plugin_name)
self.plugin_to_worker_session.pop(plugin_name, None)
stale_requests = [
request_id
for request_id, active_session in self.active_requests.items()
if active_session is session
]
for request_id in stale_requests:
self.active_requests.pop(request_id, None)
logger.warning("worker 组 {} 连接已关闭,已清理相关 handlers", group_id)
async def stop(self) -> None:
for session in list(self.worker_sessions.values()):
await session.stop()
await self.peer.stop()
async def _handle_upstream_invoke(self, message, cancel_token):
return await self.capability_router.execute(
message.capability,
message.input,
stream=message.stream,
cancel_token=cancel_token,
request_id=message.id,
)
async def _route_handler_invoke(
self,
request_id: str,
payload: dict[str, Any],
_cancel_token,
) -> dict[str, Any]:
handler_id = str(payload.get("handler_id", ""))
session = self.handler_to_worker.get(handler_id)
if session is None:
raise AstrBotError.invalid_input(f"handler not found: {handler_id}")
self.active_requests[request_id] = session
try:
return await session.invoke_handler(
handler_id,
payload.get("event", {}),
request_id=request_id,
)
finally:
self.active_requests.pop(request_id, None)
async def _handle_upstream_cancel(self, request_id: str) -> None:
session = self.active_requests.get(request_id)
if session is not None:
await session.cancel(request_id)
+12 -1
View File
@@ -86,6 +86,17 @@ from loguru import logger
MessageHandler = Callable[[str], Awaitable[None]]
def _frame_stdio_payload(payload: str) -> str:
body = payload
if body.endswith("\r\n"):
body = body[:-2]
elif body.endswith(("\n", "\r")):
body = body[:-1]
if "\n" in body or "\r" in body:
raise ValueError("STDIO payload 不允许包含原始换行符")
return f"{body}\n"
class Transport(ABC):
def __init__(self) -> None:
self._handler: MessageHandler | None = None
@@ -175,7 +186,7 @@ class StdioTransport(Transport):
self._closed.set()
async def send(self, payload: str) -> None:
line = payload if payload.endswith("\n") else f"{payload}\n"
line = _frame_stdio_payload(payload)
if self._process is not None:
if self._process.stdin is None:
raise RuntimeError("STDIO subprocess stdin 不可用")
+436
View File
@@ -0,0 +1,436 @@
"""Worker 端运行时:PluginWorkerRuntime 运行单个插件,GroupWorkerRuntime 在同一进程中运行多个插件。
核心类:
GroupWorkerRuntime: 组 Worker 运行时
- 在同一进程中加载并运行多个插件
- 聚合所有插件的 handlers 和 capabilities
- 统一处理 invoke 和 cancel 请求
- 管理每个插件的生命周期回调
PluginWorkerRuntime: 单插件 Worker 运行时
- 加载单个插件
- 通过 Peer 与 Supervisor 通信
- 分发 handler 调用
- 处理生命周期回调 (on_start, on_stop)
启动流程:
Worker 启动:
1. load_plugin_spec() 加载插件规范
2. load_plugin() 加载插件组件
3. 创建 Peer 并设置处理器
4. 向 Supervisor 发送 initialize
5. 等待 Supervisor 的 initialize_result
6. 执行 on_start 生命周期回调
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from loguru import logger
from .._legacy_runtime import (
LegacyWorkerRuntimeBridge,
bind_legacy_runtime_contexts,
build_legacy_worker_runtime_bridge,
run_legacy_worker_shutdown_hooks,
run_legacy_worker_startup_hooks,
run_plugin_lifecycle,
)
from ..context import Context as RuntimeContext
from ..errors import AstrBotError
from ..protocol.messages import PeerInfo
from .handler_dispatcher import CapabilityDispatcher, HandlerDispatcher
from .loader import (
LoadedPlugin,
PluginSpec,
load_plugin,
load_plugin_spec,
)
from .peer import Peer
__all__ = [
"GroupPluginRuntimeState",
"GroupWorkerRuntime",
"PluginWorkerRuntime",
"_load_group_plugin_specs",
]
@dataclass(slots=True)
class GroupPluginRuntimeState:
plugin: PluginSpec
loaded_plugin: LoadedPlugin
lifecycle_context: RuntimeContext
def _load_group_plugin_specs(group_metadata_path: Path) -> tuple[str, list[PluginSpec]]:
try:
payload = json.loads(group_metadata_path.read_text(encoding="utf-8"))
except Exception as exc:
raise RuntimeError(
f"failed to read worker group metadata: {group_metadata_path}"
) from exc
if not isinstance(payload, dict):
raise RuntimeError(f"invalid worker group metadata: {group_metadata_path}")
entries = payload.get("plugin_entries")
if not isinstance(entries, list) or not entries:
raise RuntimeError(
f"worker group metadata missing plugin_entries: {group_metadata_path}"
)
plugins: list[PluginSpec] = []
for entry in entries:
if not isinstance(entry, dict):
raise RuntimeError(
f"worker group metadata contains invalid plugin entry: {group_metadata_path}"
)
plugin_dir = entry.get("plugin_dir")
if not isinstance(plugin_dir, str) or not plugin_dir:
raise RuntimeError(
f"worker group metadata contains invalid plugin_dir: {group_metadata_path}"
)
plugins.append(load_plugin_spec(Path(plugin_dir)))
group_id = payload.get("group_id")
if not isinstance(group_id, str) or not group_id:
group_id = group_metadata_path.stem
return group_id, plugins
class GroupWorkerRuntime:
def __init__(self, *, group_metadata_path: Path, transport) -> None:
self.group_metadata_path = group_metadata_path.resolve()
self.group_id, self.plugins = _load_group_plugin_specs(self.group_metadata_path)
self.transport = transport
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name=self.group_id, role="plugin", version="v4"),
)
self.skipped_plugins: dict[str, str] = {}
self._plugin_states: list[GroupPluginRuntimeState] = []
self._active_plugin_states: list[GroupPluginRuntimeState] = []
self._load_plugins()
self._refresh_dispatchers()
self.peer.set_invoke_handler(self._handle_invoke)
self.peer.set_cancel_handler(self._handle_cancel)
def _load_plugins(self) -> None:
for plugin in self.plugins:
try:
loaded_plugin = load_plugin(plugin)
except Exception as exc:
self.skipped_plugins[plugin.name] = str(exc)
logger.exception(
"组 {} 中插件 {} 加载失败,启动时将跳过",
self.group_id,
plugin.name,
)
continue
lifecycle_context = RuntimeContext(peer=self.peer, plugin_id=plugin.name)
bind_legacy_runtime_contexts(
[*loaded_plugin.handlers, *loaded_plugin.capabilities],
lifecycle_context,
)
self._plugin_states.append(
GroupPluginRuntimeState(
plugin=plugin,
loaded_plugin=loaded_plugin,
lifecycle_context=lifecycle_context,
)
)
self._active_plugin_states = list(self._plugin_states)
def _refresh_dispatchers(self) -> None:
handlers = [
handler
for state in self._active_plugin_states
for handler in state.loaded_plugin.handlers
]
capabilities = [
capability
for state in self._active_plugin_states
for capability in state.loaded_plugin.capabilities
]
self.dispatcher = HandlerDispatcher(
plugin_id=self.group_id,
peer=self.peer,
handlers=handlers,
)
self.capability_dispatcher = CapabilityDispatcher(
plugin_id=self.group_id,
peer=self.peer,
capabilities=capabilities,
)
async def start(self) -> None:
await self.peer.start()
started_states: list[GroupPluginRuntimeState] = []
try:
active_states: list[GroupPluginRuntimeState] = []
for state in self._plugin_states:
try:
await self._run_lifecycle(state, "on_start")
except Exception as exc:
self.skipped_plugins[state.plugin.name] = str(exc)
logger.exception(
"组 {} 中插件 {} on_start 失败,启动时将跳过",
self.group_id,
state.plugin.name,
)
continue
active_states.append(state)
started_states.append(state)
self._active_plugin_states = active_states
self._refresh_dispatchers()
if not self._active_plugin_states:
raise RuntimeError(
f"worker group {self.group_id} has no active plugins"
)
await self.peer.initialize(
[
handler.descriptor
for state in self._active_plugin_states
for handler in state.loaded_plugin.handlers
],
provided_capabilities=[
capability.descriptor
for state in self._active_plugin_states
for capability in state.loaded_plugin.capabilities
],
metadata=self._initialize_metadata(),
)
for state in self._active_plugin_states:
await self._run_legacy_worker_startup_hooks(
state,
metadata=dict(state.plugin.manifest_data),
)
except Exception:
for state in reversed(started_states):
try:
await self._run_lifecycle(state, "on_stop")
except Exception:
logger.exception(
"组 {} 在启动失败清理插件 {} on_stop 时发生异常",
self.group_id,
state.plugin.name,
)
await self.peer.stop()
raise
async def stop(self) -> None:
first_error: Exception | None = None
try:
for state in reversed(self._active_plugin_states):
try:
await self._run_legacy_worker_shutdown_hooks(
state,
metadata=dict(state.plugin.manifest_data),
)
await self._run_lifecycle(state, "on_stop")
except Exception as exc:
if first_error is None:
first_error = exc
logger.exception(
"组 {} 停止插件 {} 时发生异常",
self.group_id,
state.plugin.name,
)
finally:
await self.peer.stop()
if first_error is not None:
raise first_error
async def _handle_invoke(self, message, cancel_token):
if message.capability == "handler.invoke":
return await self.dispatcher.invoke(message, cancel_token)
try:
return await self.capability_dispatcher.invoke(message, cancel_token)
except LookupError as exc:
raise AstrBotError.capability_not_found(message.capability) from exc
async def _handle_cancel(self, request_id: str) -> None:
await self.dispatcher.cancel(request_id)
await self.capability_dispatcher.cancel(request_id)
def _initialize_metadata(self) -> dict[str, Any]:
return {
"group_id": self.group_id,
"plugins": [plugin.name for plugin in self.plugins],
"loaded_plugins": [
state.plugin.name for state in self._active_plugin_states
],
"skipped_plugins": dict(self.skipped_plugins),
"capability_sources": {
capability.descriptor.name: state.plugin.name
for state in self._active_plugin_states
for capability in state.loaded_plugin.capabilities
},
}
async def _run_lifecycle(
self,
state: GroupPluginRuntimeState,
method_name: str,
) -> None:
await run_plugin_lifecycle(
state.loaded_plugin.instances, method_name, state.lifecycle_context
)
async def _run_legacy_worker_startup_hooks(
self,
state: GroupPluginRuntimeState,
*,
metadata: dict[str, Any],
) -> None:
await run_legacy_worker_startup_hooks(
[
*state.loaded_plugin.handlers,
*state.loaded_plugin.capabilities,
],
context=state.lifecycle_context,
metadata=metadata,
)
async def _run_legacy_worker_shutdown_hooks(
self,
state: GroupPluginRuntimeState,
*,
metadata: dict[str, Any],
) -> None:
await run_legacy_worker_shutdown_hooks(
[
*state.loaded_plugin.handlers,
*state.loaded_plugin.capabilities,
],
context=state.lifecycle_context,
metadata=metadata,
)
class PluginWorkerRuntime:
def __init__(self, *, plugin_dir: Path, transport) -> None:
self.plugin = load_plugin_spec(plugin_dir)
self.transport = transport
self.loaded_plugin = load_plugin(self.plugin)
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name=self.plugin.name, role="plugin", version="v4"),
)
self.dispatcher = HandlerDispatcher(
plugin_id=self.plugin.name,
peer=self.peer,
handlers=self.loaded_plugin.handlers,
)
self.capability_dispatcher = CapabilityDispatcher(
plugin_id=self.plugin.name,
peer=self.peer,
capabilities=self.loaded_plugin.capabilities,
)
self._lifecycle_context = RuntimeContext(
peer=self.peer, plugin_id=self.plugin.name
)
self._legacy_worker_runtime: LegacyWorkerRuntimeBridge = (
build_legacy_worker_runtime_bridge(
lambda: [
*self.loaded_plugin.handlers,
*self.loaded_plugin.capabilities,
]
)
)
self._bind_legacy_runtime_contexts(self._lifecycle_context)
self.peer.set_invoke_handler(self._handle_invoke)
self.peer.set_cancel_handler(self._handle_cancel)
async def start(self) -> None:
await self.peer.start()
lifecycle_started = False
try:
await self._run_lifecycle("on_start")
lifecycle_started = True
await self.peer.initialize(
[item.descriptor for item in self.loaded_plugin.handlers],
provided_capabilities=[
item.descriptor for item in self.loaded_plugin.capabilities
],
metadata={
"plugin_id": self.plugin.name,
"plugins": [self.plugin.name],
"loaded_plugins": [self.plugin.name],
"skipped_plugins": {},
"capability_sources": {
item.descriptor.name: self.plugin.name
for item in self.loaded_plugin.capabilities
},
},
)
await self._run_legacy_worker_startup_hooks(
metadata=dict(self.plugin.manifest_data),
)
except Exception:
if lifecycle_started:
try:
await self._run_lifecycle("on_stop")
except Exception:
logger.exception(
"插件 {} 在启动失败清理 on_stop 时发生异常",
self.plugin.name,
)
await self.peer.stop()
raise
async def stop(self) -> None:
try:
await self._run_legacy_worker_shutdown_hooks(
metadata=dict(self.plugin.manifest_data),
)
await self._run_lifecycle("on_stop")
finally:
await self.peer.stop()
async def _handle_invoke(self, message, cancel_token):
if message.capability == "handler.invoke":
return await self.dispatcher.invoke(message, cancel_token)
try:
return await self.capability_dispatcher.invoke(message, cancel_token)
except LookupError as exc:
raise AstrBotError.capability_not_found(message.capability) from exc
async def _handle_cancel(self, request_id: str) -> None:
await self.dispatcher.cancel(request_id)
await self.capability_dispatcher.cancel(request_id)
async def _run_lifecycle(self, method_name: str) -> None:
await run_plugin_lifecycle(
self.loaded_plugin.instances, method_name, self._lifecycle_context
)
def _bind_legacy_runtime_contexts(self, runtime_context: RuntimeContext) -> None:
self._legacy_worker_runtime.bind_runtime_contexts(runtime_context)
async def _run_legacy_worker_startup_hooks(
self, *, metadata: dict[str, Any]
) -> None:
await self._legacy_worker_runtime.run_startup_hooks(
context=self._lifecycle_context,
metadata=metadata,
)
async def _run_legacy_worker_shutdown_hooks(
self,
*,
metadata: dict[str, Any],
) -> None:
await self._legacy_worker_runtime.run_shutdown_hooks(
context=self._lifecycle_context,
metadata=metadata,
)
+32
View File
@@ -88,6 +88,38 @@ def test_load_legacy_main_component_classes_supports_relative_imports():
assert classes[0].helper_value == "legacy-ok"
def test_load_legacy_main_component_classes_preserves_definition_order():
with tempfile.TemporaryDirectory() as temp_dir:
plugin_dir = Path(temp_dir) / "legacy_plugin"
plugin_dir.mkdir()
(plugin_dir / "main.py").write_text(
textwrap.dedent(
"""\
from astrbot_sdk.api.star import Star
class ZebraComponent(Star):
pass
class AlphaComponent(Star):
pass
"""
),
encoding="utf-8",
)
classes = load_legacy_main_component_classes(
plugin_name="legacy-plugin",
plugin_dir=plugin_dir,
)
assert [cls.__name__ for cls in classes] == [
"ZebraComponent",
"AlphaComponent",
]
def test_load_plugin_manifest_payload_prefers_plugin_yaml_when_present():
with tempfile.TemporaryDirectory() as temp_dir:
plugin_dir = Path(temp_dir) / "plugin"
+12 -6
View File
@@ -13,10 +13,16 @@ from unittest.mock import MagicMock, patch
import pytest
import yaml
from astrbot_sdk._legacy_api import LegacyContext
from astrbot_sdk._legacy_runtime import LegacyRuntimeAdapter
from astrbot_sdk.api.event.filter import CustomFilter, custom_filter
from astrbot_sdk._legacy_runtime import (
LegacyRuntimeAdapter,
)
from astrbot_sdk._legacy_runtime import (
create_legacy_component_context as _create_legacy_context,
)
from astrbot_sdk._legacy_runtime import (
is_new_star_component as _is_new_star_component,
)
from astrbot_sdk.protocol.descriptors import CommandTrigger, HandlerDescriptor
from astrbot_sdk.runtime.environment_groups import (
GROUP_STATE_FILE_NAME,
@@ -24,14 +30,12 @@ from astrbot_sdk.runtime.environment_groups import (
GroupEnvironmentManager,
)
from astrbot_sdk.runtime.loader import (
STATE_FILE_NAME,
LoadedHandler,
LoadedPlugin,
PluginDiscoveryResult,
PluginEnvironmentManager,
PluginSpec,
STATE_FILE_NAME,
_create_legacy_context,
_is_new_star_component,
_iter_handler_names,
_venv_python_path,
discover_plugins,
@@ -40,6 +44,8 @@ from astrbot_sdk.runtime.loader import (
load_plugin_spec,
)
from astrbot_sdk.api.event.filter import CustomFilter, custom_filter
def write_test_plugin(
plugins_dir: Path,
+41
View File
@@ -345,6 +345,7 @@ class PeerRuntimeTest(unittest.IsolatedAsyncioTestCase):
transport=self.right,
peer_info=PeerInfo(name="plugin", role="plugin", version="v4"),
protocol_version="2.0",
supported_protocol_versions=["1.0", "2.0"],
)
await core.start()
@@ -359,6 +360,46 @@ class PeerRuntimeTest(unittest.IsolatedAsyncioTestCase):
self.assertTrue(core._closed)
self.assertTrue(plugin._closed)
async def test_initialize_negotiates_lower_minor_protocol_version(self) -> None:
core = Peer(
transport=self.left,
peer_info=PeerInfo(name="core", role="core", version="v4"),
protocol_version="1.0",
supported_protocol_versions=["1.0"],
)
core.set_initialize_handler(
lambda _message: asyncio.sleep(
0,
result=InitializeOutput(
peer=PeerInfo(name="core", role="core", version="v4"),
capabilities=[],
metadata={},
),
)
)
plugin = Peer(
transport=self.right,
peer_info=PeerInfo(name="plugin", role="plugin", version="v4"),
protocol_version="1.1",
supported_protocol_versions=["1.0", "1.1"],
)
await core.start()
await plugin.start()
output = await plugin.initialize([])
self.assertEqual(output.protocol_version, "1.0")
self.assertEqual(plugin.negotiated_protocol_version, "1.0")
self.assertEqual(core.negotiated_protocol_version, "1.0")
self.assertEqual(
core.remote_metadata["supported_protocol_versions"], ["1.1", "1.0"]
)
self.assertEqual(plugin.remote_metadata["negotiated_protocol_version"], "1.0")
await plugin.stop()
await core.stop()
async def test_wait_until_remote_initialized_raises_if_connection_closes_first(
self,
) -> None:
+6
View File
@@ -233,6 +233,12 @@ class TestInitializeOutput:
output = InitializeOutput(peer=peer, metadata={"session": "abc"})
assert output.metadata["session"] == "abc"
def test_with_protocol_version(self):
"""InitializeOutput should accept negotiated protocol_version."""
peer = PeerInfo(name="core", role="core")
output = InitializeOutput(peer=peer, protocol_version="1.0")
assert output.protocol_version == "1.0"
class TestResultMessage:
"""Tests for ResultMessage model."""
+21
View File
@@ -239,6 +239,27 @@ class TestStdioTransportFileMode:
await transport.stop()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"payload",
["first\nsecond", "first\rsecond", "first\r\nsecond"],
)
async def test_send_rejects_embedded_newlines(self, payload):
"""send() should reject payloads containing raw embedded newlines."""
stdout = MagicMock()
stdout.write = MagicMock()
stdout.flush = MagicMock()
transport = StdioTransport(stdout=stdout)
with patch("sys.stdin"):
await transport.start()
with pytest.raises(ValueError, match="原始换行符"):
await transport.send(payload)
stdout.write.assert_not_called()
await transport.stop()
@pytest.mark.asyncio
async def test_send_raises_without_stdout(self):
"""send() should raise if stdout is None."""