mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -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()` 仍然不够。
|
||||
|
||||
@@ -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
+29
-1161
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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 不可用")
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user