diff --git a/AGENTS.md b/AGENTS.md index 7d13fa3e9..91b4cf3a9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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()` 仍然不够。 diff --git a/CLAUDE.md b/CLAUDE.md index fcee5e9a1..bed121ef8 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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. # 开发命令 diff --git a/docs/v4/architecture-analysis.md b/docs/v4/architecture-analysis.md new file mode 100644 index 000000000..a2e701d1d --- /dev/null +++ b/docs/v4/architecture-analysis.md @@ -0,0 +1,1304 @@ +# AstrBot SDK v4 架构分析报告 + +> 版本:0.1.0 +> 生成日期:2026-03-14 +> 分析范围:`src-new/astrbot_sdk` 及相关测试 + +--- + +## 目录 + +1. [概述](#1-概述) +2. [优点](#2-优点) +3. [缺点](#3-缺点) +4. [设计理念](#4-设计理念) +5. [核心架构](#5-核心架构) +6. [实现思路](#6-实现思路) +7. [技术亮点](#7-技术亮点) +8. [演进规划](#8-演进规划) +9. [总结](#9-总结) + +--- + +## 1. 概述 + +AstrBot SDK v4 是一个**插件化机器人框架 SDK**,实现了从旧版 JSON-RPC 协议到新一代 v4 协议的架构重构。其核心特点包括: + +- **双层目标**:提供原生 v4 插件模型 + 维持旧版插件兼容 +- **协议优先**:设计清晰的 v4 线协议,兼容层作为过渡 +- **分层清晰**:插件作者、客户端、运行时、协议层职责明确 +- **进程隔离**:Supervisor-Worker 架构,每插件独立进程 +- **能力路由**:基于命名空间的 Capability 系统 + +### 项目结构概览 + +``` +astrbot-sdk/ +├── src-new/astrbot_sdk/ # v4 原生实现(主源码) +│ ├── protocol/ # v4 协议层(消息、描述符) +│ ├── runtime/ # 运行时核心(peer、transport、router、loader) +│ ├── clients/ # 能力客户端(llm、memory、db、platform) +│ ├── api/ # 旧 API 兼容层门面 +│ ├── _legacy_*.py # 私有兼容实现(收口边界) +│ └── astrbot/ # 旧包名 facade(受控兼容面) +├── src/ # 旧版代码(遗留) +├── tests_v4/ # v4 测试套件 +├── test_plugin/ # 测试插件示例(old/new 分离) +└── docs/ # 文档目录 +``` + +--- + +## 2. 优点 + +### 2.1 架构设计层面 + +#### 清晰的分层架构 + +``` +┌─────────────────────────────────────────┐ +│ 插件作者层 │ +│ Star / Context / MessageEvent │ +└─────────────────┬───────────────────┘ + │ +┌─────────────────▼───────────────────┐ +│ 客户端层 │ +│ LLMClient / DBClient / ... │ +│ CapabilityProxy │ +└─────────────────┬───────────────────┘ + │ +┌─────────────────▼───────────────────┐ +│ 运行时层 │ +│ Peer / Transport │ +│ CapabilityRouter / HandlerDispatcher│ +│ loader / bootstrap │ +└─────────────────┬───────────────────┘ + │ +┌─────────────────▼───────────────────┐ +│ 协议层 │ +│ messages / descriptors │ +│ legacy_adapter │ +└───────────────────────────────────┘ +``` + +每层职责单一,边界清晰,降低了理解和维护成本。 + +#### 协议优先的设计 + +v4 协议层(`protocol/messages.py`、`protocol/descriptors.py`)定义了清晰的线协议契约: + +- 5 种消息类型:`InitializeMessage`、`InvokeMessage`、`ResultMessage`、`EventMessage`、`CancelMessage` +- 强类型约束:使用 Pydantic 模型进行严格验证 +- 版本协商:支持 `protocol_version` 协商机制 +- 流式支持:统一的 `EventMessage` 处理流式调用 + +这种设计使得协议与实现解耦,便于跨语言实现和协议演进。 + +#### 窄导出的稳定 API + +顶层 `astrbot_sdk.__init__.py` 只导出 7 个核心类: + +```python +from .context import Context +from .decorators import (on_command, on_event, on_message, + on_schedule, provide_capability, require_admin) +from .errors import AstrBotError +from .events import MessageEvent +from .star import Star +``` + +这种"最小稳定面"设计减少了 API 变更的影响范围,有利于长期维护。 + +### 2.2 兼容性设计层面 + +#### 三级兼容策略 + +| 级别 | 路径 | 策略 | +|------|------|------| +| 一级 | `astrbot.api.*` | 优先做真实兼容 | +| 二级 | `astrbot.core.*` | 按需补薄 shim | +| 三级 | 旧应用内部系统 | 不做树级复刻 | + +这种分层策略避免了"全盘照搬旧架构"的陷阱,只保证真实插件使用的路径可用。 + +#### 私有边界收口 + +兼容逻辑集中在 `_legacy_api.py`、`_legacy_runtime.py`、`_legacy_loader.py` 等私有模块: + +- `LegacyContext`:旧版上下文适配 +- `LegacyRuntimeAdapter`:运行时执行适配 +- `SessionWaiterManager`:会话等待机制 + +这种收口设计让兼容层可被独立演进和最终移除。 + +### 2.3 运行时设计层面 + +#### Capability 模式 + +基于命名空间的能力系统: + +```python +# 注册能力 +router.register( + CapabilityDescriptor( + name="my_plugin.calculate", + description="执行计算", + input_schema={"type": "object", ...}, + output_schema={"type": "object", ...}, + ), + call_handler=my_calculate, +) + +# 调用能力 +result = await ctx.llm.chat(prompt="hello") +# 实际调用 peer.invoke("llm.chat", {"prompt": "hello"}) +``` + +优势: +- JSON Schema 输入输出验证 +- 支持同步和流式两种模式 +- 统一的错误处理 +- 命名空间避免冲突 + +#### Peer 模式 + +统一的对等端抽象,既是客户端也是服务端: + +```python +# 作为客户端 +peer = Peer(transport, PeerInfo(...)) +await peer.start() +output = await peer.initialize(handlers) +result = await peer.invoke("llm.chat", {"prompt": "hello"}) + +# 作为服务端 +peer.set_invoke_handler(my_handler) +await peer.start() +``` + +优势: +- 双向通信对称 +- 统一的初始化握手 +- 请求 ID 关联 +- 取消传播机制 + +#### Supervisor-Worker 架构 + +``` +AstrBot Core (Python) + | + v + SupervisorRuntime (管理多插件) + | + +-- WorkerSession (插件 A) -- StdioTransport -- PluginWorkerRuntime + | + +-- WorkerSession (插件 B) -- StdioTransport -- PluginWorkerRuntime + | + +-- WorkerSession (插件 C) -- StdioTransport -- PluginWorkerRuntime +``` + +优势: +- 进程隔离,单个插件崩溃不影响其他 +- 独立 Python 环境,依赖隔离 +- 支持 Worker 崩溃检测和清理 +- 支持分组 Worker 共享环境 + +### 2.4 开发体验层面 + +#### 完整的测试体系 + +``` +tests_v4/ +├── test_protocol.py # 协议模型测试 +├── test_peer.py # Peer 通信测试 +├── test_transport.py # 传输层测试 +├── test_loader.py # 插件加载测试 +├── test_capability_router.py # 能力路由测试 +├── test_handler_dispatcher.py # 处理器分发测试 +├── test_legacy_runtime.py # Legacy 运行时测试 +├── test_legacy_loader.py # Legacy 加载器测试 +├── test_api_*.py # API 兼容性测试 +├── test_new_plugin_integration.py # v4 插件集成测试 +├── test_legacy_plugin_integration.py # 旧插件集成测试 +└── test_grouped_environment_smoke.py # 分组环境测试 +``` + +#### 本地开发支持 + +`astrbot_sdk.testing` 提供本地开发 harness: + +```python +from astrbot_sdk.testing import PluginHarness, LocalRuntimeConfig + +harness = PluginHarness(config=LocalRuntimeConfig(...)) +await harness.start() + +# 测试插件 +result = await harness.invoke_handler("my_command", event) +``` + +优势: +- 无需启动完整 Core 即可测试 +- 复用真实 loader、dispatcher +- 支持交互式开发 + +--- + +## 3. 缺点 + +### 3.1 架构复杂度 + +#### 兼容层带来的认知负担 + +虽然兼容逻辑被收口到私有模块,但仍需维护: + +- `_legacy_api.py`:600+ 行 +- `_legacy_runtime.py`:500+ 行 +- `_legacy_loader.py`:400+ 行 +- `_session_waiter.py`:300+ 行 + +对于新开发者来说,理解"为什么要这些文件"需要额外学习成本。 + +#### 多层抽象的调用链 + +一个简单的 LLM 调用需要经过: + +``` +ctx.llm.chat(prompt) + -> LLMClient.chat() + -> CapabilityProxy.call("llm.chat") + -> Peer.invoke("llm.chat") + -> StdioTransport.send() + [跨进程] + -> Peer._handle_invoke() + -> CapabilityRouter.execute("llm.chat") + -> Supervisor 提供的实际实现 +``` + +这种多层调用链在调试时需要追踪多个文件。 + +### 3.2 兼容性限制 + +#### 降级兼容部分 + +某些能力只能"降级"实现: + +- `command_group`:旧版支持树状命令帮助,新版展平成普通命令名 +- legacy handshake 转 v4:只能近似恢复触发信息,原始 payload 保留在 metadata + +#### 明确不支持的部分 + +某些旧功能完全不支持: + +- `astrbot.api.agent()`:显式 `NotImplementedError` +- `register_platform_adapter`:不提供 +- 旧 LLM hook / plugin hook 的完整执行链:部分实现 + +### 3.3 测试覆盖的挑战 + +#### Legacy 插件矩阵维护 + +`tests_v4/external_plugin_matrix.json` 维护真实插件兼容矩阵: + +```json +{ + "plugins": [ + "astrbot_plugin_hapi_connector", + "astrbot_plugin_endfield" + ] +} +``` + +需要持续跟踪外部插件变更,维护成本较高。 + +#### 集成测试的依赖 + +真实集成测试需要: +- 克隆外部插件仓库 +- 运行完整的 Supervisor-Worker 链路 +- 处理网络和进程管理 + +这些测试执行较慢且容易受环境影响。 + +### 3.4 文档与代码的漂移 + +#### `refactor.md` 不再准确 + +架构文档明确指出: + +> `refactor.md` 仅保留历史设计意图和演进说明,不再描述现状。 + +这意味着: +- 新开发者可能被旧文档误导 +- 需要同时阅读 ARCHITECTURE.md 和 refactor.md +- 维护两份文档的成本 + +#### CLAUDE.md 中的 70+ 条备注 + +`CLAUDE.md` 记录了大量架构细节和陷阱,例如: + +- 2026-03-12: Legacy handshake payloads only contain `event_type` / `handler_full_name` metadata +- 2026-03-13: Keep `astrbot_sdk.runtime` root exports narrow +- 2026-03-14: `test_plugin/old/` and `test_plugin/new/` may contain checked-in `__pycache__` artifacts + +这些备注有价值但分散,不利于新人学习。 + +### 3.5 进程模型的开销 + +#### 一插件一进程 + +每个插件独立运行在子进程中,带来: + +- 启动延迟:插件数量多时启动时间长 +- 资源开销:Python 解释器和依赖的重复加载 +- 调试复杂:跨进程调试不如单进程方便 + +虽然有共享环境分组机制(`environment_groups.py`),但仍然无法完全消除进程开销。 + +--- + +## 4. 设计理念 + +### 4.1 协议优先 + +> v4 协议层是核心,兼容层是过渡 + +**体现**: + +- `protocol/` 目录独立设计,不依赖旧版代码 +- 协议消息使用强类型 Pydantic 模型 +- 协议版本协商机制 +- `legacy_adapter.py` 作为协议适配层,不污染核心 + +**好处**: + +- 协议可独立演进 +- 支持跨语言实现(未来 Go/Rust 版) +- 兼容层可最终移除 + +### 4.2 分层清晰 + +> 每层有明确职责,避免耦合 + +**体现**: + +- 插件作者层:`Star`、`Context`、`MessageEvent` +- 客户端层:`LLMClient`、`DBClient` 等 +- 运行时层:`Peer`、`Transport`、`CapabilityRouter` +- 协议层:`messages`、`descriptors` + +**好处**: + +- 各层可独立测试 +- 修改影响范围可控 +- 新人容易定位问题 + +### 4.3 窄导出 + +> 顶层只暴露稳定 API + +**体现**: + +- `astrbot_sdk.__init__` 只导出 7 个核心类 +- `astrbot_sdk.runtime.__init__` 不导出 loader/bootstrap +- `astrbot_sdk.protocol.__init__` 只导出 v4 原生模型 + +**好处**: + +- 减少变更影响面 +- 避免"意外公开内部实现" +- 长期兼容性更易保证 + +### 4.4 私有收口 + +> 兼容逻辑在私有模块 + +**体现**: + +- `_legacy_api.py`:私有兼容 API +- `_legacy_runtime.py`:私有运行时适配 +- `_legacy_loader.py`:私有加载器逻辑 + +**好处**: + +- 兼容层可独立演进 +- 不污染主代码库 +- 未来可整体移除 + +### 4.5 受控兼容 + +> 不是全盘复制旧架构 + +**体现**: + +- 三级兼容策略 +- 不支持的路径显式 `NotImplementedError` +- 外部插件矩阵作为真实标准 + +**好处**: + +- 避免维护负担无限增长 +- 清晰的兼容边界 +- 鼓励迁移到新 API + +--- + +## 5. 核心架构 + +### 5.1 协议层(Protocol) + +#### 消息类型 + +```python +# 1. InitializeMessage - 初始化握手 +{ + "type": "initialize", + "id": "msg_001", + "protocol_version": "1.0", + "peer": {"name": "plugin", "role": "plugin", "version": "v4"}, + "handlers": [...], + "provided_capabilities": [...], + "metadata": {} +} + +# 2. InvokeMessage - 能力调用 +{ + "type": "invoke", + "id": "msg_002", + "capability": "llm.chat", + "input": {"prompt": "hello"}, + "stream": false +} + +# 3. ResultMessage - 调用结果 +{ + "type": "result", + "id": "msg_002", + "success": true, + "output": {"text": "response"}, + "error": null +} + +# 4. EventMessage - 流式事件 +{ + "type": "event", + "id": "msg_003", + "phase": "delta", # started/delta/completed/failed + "data": {}, + "output": {}, + "error": null +} + +# 5. CancelMessage - 取消请求 +{ + "type": "cancel", + "id": "msg_003", + "reason": "user_cancelled" +} +``` + +#### 版本协商 + +```python +# PeerInfo.version: 软件版本标识("v4") +# protocol_version: 线协议版本("1.0") + +# 协商过程: +# 1. 发起方发送首选 protocol_version +# 2. 响应方检查支持列表,选择最佳版本 +# 3. 双方使用协商后的版本通信 +``` + +#### 描述符系统 + +```python +# HandlerDescriptor - 处理器描述 +@dataclass +class HandlerDescriptor: + id: str + trigger: Trigger # CommandTrigger | MessageTrigger | EventTrigger | ScheduleTrigger + permissions: Permissions + metadata: dict[str, Any] + +# CapabilityDescriptor - 能力描述 +@dataclass +class CapabilityDescriptor: + name: str # "llm.chat" + description: str + input_schema: dict # JSON Schema + output_schema: dict # JSON Schema + supports_stream: bool + cancelable: bool +``` + +### 5.2 运行时层(Runtime) + +#### Peer + +核心职责: + +```python +class Peer: + # 握手 + async def initialize(self, handlers, ...) -> InitializeOutput + + # 调用 + async def invoke(self, capability, payload) -> dict + async def invoke_stream(self, capability, payload) -> AsyncIterator[EventMessage] + + # 取消 + async def cancel(self, request_id, reason) + + # 生命周期 + async def start() + async def stop() +``` + +消息处理流程: + +``` +入站消息: + ResultMessage -> 唤醒 Future + EventMessage -> 投递到流式队列 + InitializeMessage -> 调用 initialize_handler + InvokeMessage -> 创建任务调用 invoke_handler + CancelMessage -> 取消对应任务 + +出站消息: + initialize() -> InitializeMessage + invoke() -> InvokeMessage(stream=False) + invoke_stream() -> InvokeMessage(stream=True) + cancel() -> CancelMessage +``` + +#### Transport + +抽象传输层: + +```python +class Transport(ABC): + @abstractmethod + async def start() + @abstractmethod + async def stop() + @abstractmethod + async def send(self, message: str) + @abstractmethod + def set_message_handler(self, handler) +``` + +实现: + +- `StdioTransport`:标准输入输出(支持子进程和文件模式) +- `WebSocketServerTransport`:WebSocket 服务端 +- `WebSocketClientTransport`:WebSocket 客户端 + +#### CapabilityRouter + +能力注册与执行: + +```python +class CapabilityRouter: + # 注册 + def register(self, descriptor, *, call_handler, stream_handler, finalize) + + # 执行 + async def execute(self, capability, payload, *, stream, cancel_token) + + # 18 个内建能力 + # llm: chat, chat_raw, stream_chat + # memory: search, save, get, delete + # db: get, set, delete, list, get_many, set_many, watch + # platform: send, send_image, send_chain, get_members +``` + +#### HandlerDispatcher + +处理器分发与参数注入: + +```python +class HandlerDispatcher: + async def invoke(self, message, cancel_token): + # 1. 检查 session_waiter + # 2. 准备 legacy 运行时(过滤器) + # 3. 构建参数(类型注入) + # 4. 执行 handler + # 5. 处理结果(legacy 结果兼容) + # 6. 错误处理 +``` + +#### Loader + +插件发现与加载: + +```python +def discover_plugins(plugins_dir) -> list[PluginSpec] + +def load_plugin(spec) -> LoadedPlugin + +# PluginSpec +@dataclass +class PluginSpec: + name: str + plugin_dir: Path + manifest_path: Path + requirements_path: Path + python_version: str + manifest_data: dict + +# LoadedPlugin +@dataclass +class LoadedPlugin: + plugin: PluginSpec + instances: list[Any] + handlers: list[HandlerWrapper] +``` + +### 5.3 客户端层(Clients) + +```python +class Context: + llm: LLMClient + memory: MemoryClient + db: DBClient + platform: PlatformClient + http: HTTPClient + metadata: MetadataClient + logger: Logger + cancel_token: CancelToken +``` + +每个客户端通过 `CapabilityProxy` 调用对应能力: + +```python +class LLMClient: + async def chat(self, prompt) -> str: + return await self._proxy.call("llm.chat", {"prompt": prompt}) + + async def chat_raw(self, prompt) -> LLMResponse: + return await self._proxy.call("llm.chat_raw", {"prompt": prompt}) + + async def stream_chat(self, prompt) -> AsyncIterator[str]: + async for event in self._proxy.stream("llm.stream_chat", {"prompt": prompt}): + yield event["data"]["text"] +``` + +### 5.4 兼容层(Compat) + +#### LegacyContext + +旧版上下文适配: + +```python +class LegacyContext: + def __init__(self, new_context: Context): + self._new_context = new_context + self.conversation_manager = LegacyConversationManager(self) + self.llm = ... + + def llm_generate(self, prompt) -> str: + return self._new_context.llm.chat(prompt) + + def put_kv_data(self, key, value): + asyncio.create_task(self._new_context.db.set(key, value)) + + def get_kv_data(self, key) -> Any: + return await self._new_context.db.get(key) +``` + +#### LegacyStar + +旧版 Star 基类: + +```python +class LegacyStar: + def __init__(self, context: LegacyContext): + self.context = context + + # 旧版方法 + async def initialize(self): + pass + + def register_component(self, component): + # 通过 _legacy_runtime 注册 + pass +``` + +#### LegacyRuntimeAdapter + +运行时执行适配: + +```python +class LegacyWorkerRuntimeBridge: + async def execute_legacy_handler(self, handler, event): + # 1. 应用自定义过滤器 + # 2. 执行 handler + # 3. 结果装饰(on_decorating_result) + # 4. 发送后 hook(after_message_sent) + # 5. 错误处理(on_plugin_error) +``` + +--- + +## 6. 实现思路 + +### 6.1 插件发现与加载 + +#### v4 插件(`plugin.yaml`) + +```yaml +name: my_plugin +version: "0.1.0" +description: My awesome plugin +runtime: + python: "3.12" +components: + - path: my_plugin/main.py + entry: MyComponent +permissions: + - type: admin + commands: [secure] +``` + +```python +# my_plugin/main.py +from astrbot_sdk import Star, Context, MessageEvent +from astrbot_sdk.decorators import on_command + +class MyComponent(Star): + @on_command("hello") + async def hello_cmd(self, event: MessageEvent): + await event.reply("Hello, world!") +``` + +#### Legacy 插件(`main.py`) + +```python +# main.py +from astrbot_sdk.api.star import Star +from astrbot_sdk.api.event import AstrMessageEvent + +class MyOldStar(Star): + async def initialize(self): + pass + + @filter.command("old_hello") + async def old_hello(self, event: AstrMessageEvent): + await event.reply("Old hello!") +``` + +发现流程: + +```python +def discover_plugins(plugins_dir): + for subdir in plugins_dir.iterdir(): + # 检查 plugin.yaml + yaml_path = subdir / "plugin.yaml" + if yaml_path.exists(): + return load_plugin_spec(subdir) + + # 检查 legacy main.py + main_path = subdir / "main.py" + if main_path.exists(): + return synthesize_legacy_spec(subdir) +``` + +### 6.2 环境管理与分组 + +```python +class PluginEnvironmentManager: + def plan(self, plugins: list[PluginSpec]) -> list[EnvironmentGroup]: + # 基于 runtime.python 和 requirements.txt 分组 + # 依赖兼容性分析 + # 返回共享环境规划 + + def prepare_environment(self, spec: PluginSpec): + # 创建虚拟环境 + # 安装依赖 + # 返回环境路径 + +class EnvironmentGroup: + def __init__(self, plugins: list[PluginSpec]): + self.plugins = plugins + self.env_path = self._create_shared_env() + self.lock_path = self._create_lock() + + def lock(self): + # 获取环境锁 + + def unlock(self): + # 释放环境锁 +``` + +### 6.3 消息处理流程 + +#### Handler 调用链 + +``` +Core 消息 + ↓ +Supervisor.handler_to_worker[handler_id] + ↓ +WorkerSession.invoke_handler(handler_id, event) + ↓ +Peer.invoke("handler.invoke", {handler_id, event}) + ↓ +HandlerDispatcher.invoke(message, cancel_token) + ↓ +1. 检查 session_waiter +2. 准备 legacy 运行时(过滤器) +3. 构建参数(类型注入) +4. 执行 handler +5. 处理结果(legacy 结果兼容) +6. 错误处理 +``` + +#### Capability 调用链 + +``` +插件代码调用 + ↓ +LLMClient.chat() → CapabilityProxy.call("llm.chat") + ↓ +Peer.invoke("llm.chat", payload) + ↓ +Supervisor.capability_to_worker[capability] + ↓ +WorkerSession.invoke_capability() + ↓ +CapabilityRouter.execute() + ↓ +内建或插件自定义 handler +``` + +### 6.4 Session Waiter 实现 + +```python +class SessionWaiterManager: + def __init__(self): + self._waiters: dict[str, deque[SessionWaiter]] = defaultdict(deque) + + def register(self, event: MessageEvent) -> SessionWaiter: + key = self._make_waiter_key(event) + waiter = SessionWaiter(event) + self._waiters[key].append(waiter) + return waiter + + async def dispatch(self, event: MessageEvent): + key = self._make_waiter_key(event) + queue = self._waiters.get(key) + if not queue: + return + + waiter = queue[0] + if waiter.match(event): + await waiter.resume(event) + queue.popleft() + +@dataclass +class SessionWaiter: + event: MessageEvent + future: asyncio.Future + condition: Callable[[MessageEvent], bool] + + async def wait(self, timeout: float): + return await asyncio.wait_for(self.future, timeout) +``` + +--- + +## 7. 技术亮点 + +### 7.1 取消机制 + +```python +class CancelToken: + def __init__(self): + self._cancelled = asyncio.Event() + + def cancel(self): + self._cancelled.set() + + def raise_if_cancelled(self): + if self.cancelled: + raise asyncio.CancelledError +``` + +调用链: + +``` +用户取消 + ↓ +peer.cancel(request_id) + ↓ +CancelMessage 发送 + ↓ +远端收到 CancelMessage + ↓ +CancelToken.cancel() + ↓ +asyncio.create_task().cancel() + ↓ +asyncio.CancelledError +``` + +早到取消避免: + +```python +async def _handle_invoke(self, message, token, started): + started.set() + token.raise_if_cancelled() # 早到取消检查 + # 执行逻辑... +``` + +### 7.2 JSON Schema 验证 + +```python +def _validate_schema(self, schema: dict, payload: dict): + properties = schema.get("properties", {}) + for field_name in schema.get("required", []): + if field_name not in payload: + raise AstrBotError.invalid_input(f"缺少必填字段:{field_name}") +``` + +能力注册时声明 Schema: + +```python +router.register( + CapabilityDescriptor( + name="my_plugin.calculate", + input_schema={ + "type": "object", + "properties": { + "x": {"type": "number"}, + "y": {"type": "number"}, + }, + "required": ["x", "y"], + }, + output_schema={ + "type": "object", + "properties": { + "result": {"type": "number"}, + }, + }, + ), + call_handler=my_calculate, +) +``` + +### 7.3 流式执行 + +```python +@dataclass(slots=True) +class StreamExecution: + iterator: AsyncIterator[dict[str, Any]] + finalize: FinalizeHandler # (chunks) -> dict + collect_chunks: bool = True + +# 注册流式能力 +async def stream_numbers(request_id, payload, token): + for i in range(10): + token.raise_if_cancelled() + yield {"number": i} + +router.register( + CapabilityDescriptor( + name="my_plugin.stream", + supports_stream=True, + cancelable=True, + ), + stream_handler=stream_numbers, + finalize=lambda chunks: {"count": len(chunks)}, +) + +# 调用流式能力 +async for event in peer.invoke_stream("my_plugin.stream", {}): + print(event["data"]["number"]) +``` + +### 7.4 参数注入 + +```python +class HandlerDispatcher: + async def invoke(self, message, cancel_token): + handler = self._handlers[message["handler_id"]] + ctx = Context(peer=..., plugin_id=...) + event = MessageEvent.from_dict(message["event"]) + + # 参数注入 + kwargs = {} + sig = inspect.signature(handler.method) + for param_name, param in sig.parameters.items(): + if param_name == "self": + continue + if param.annotation == Context: + kwargs[param_name] = ctx + elif param.annotation == MessageEvent: + kwargs[param_name] = event + elif param_name == "cancel_token": + kwargs[param_name] = cancel_token + else: + # 从 event 中获取 + kwargs[param_name] = getattr(event, param_name) + + return await handler.method(**kwargs) +``` + +### 7.5 传输抽象 + +```python +class StdioTransport: + def __init__(self, stdin, stdout): + self.stdin = stdin + self.stdout = stdout + + async def start(self): + self._read_task = asyncio.create_task(self._read_loop()) + + async def _read_loop(self): + while True: + line = await self.stdin.readline() + if not line: + break + self._message_handler(line.rstrip("\n")) + + async def send(self, message: str): + self.stdout.write(message + "\n") + await self.stdout.drain() +``` + +支持三种模式: + +1. **子进程模式**:`PluginWorkerRuntime` 通过子进程的 stdin/stdout 通信 +2. **文件模式**:通过临时文件交换消息(测试用) +3. **WebSocket 模式**:网络远程调用 + +--- + +## 8. 演进规划 + +### 8.1 当前规划(来自 ARCHITECTURE.md) + +1. **继续收口 runtime 对 compat 的认知** + - 统一通过 `_legacy_runtime.py` 与 `_legacy_loader.py` + - 避免直接展开更多 legacy 细节 + +2. **拆薄 `_legacy_api.py`** + - 让 `LegacyContext` 更偏向 facade 和 orchestration + - 减少直接适配逻辑 + +3. **保持 `src-new/astrbot` 为受控 facade** + - 不把旧应用整棵树重新复制进来 + - 只覆盖真实插件命中的路径 + +4. **契约测试保护** + - capability 注册表契约测试 + - compat hook 执行契约测试 + - facade 导入矩阵契约测试 + +### 8.2 建议的长期方向 + +#### 8.2.1 兼容层逐步淘汰 + +阶段 1(当前):兼容层完整功能 + +- 所有旧插件可运行 +- 文档明确兼容级别 + +阶段 2(中期):兼容层标记 deprecated + +- 新项目不再使用旧 API +- 迁移工具完善 +- 旧 API 发出警告 + +阶段 3(长期):兼容层移除 + +- 移除 `_legacy_*.py` +- 移除 `src-new/astrbot` facade +- 清理 `astrbot_sdk.api` + +#### 8.2.2 协议演进 + +v4.1:增强能力 + +- 更细粒度的权限控制 +- 插件间直接通信能力 +- 热更新支持 + +v5.0:可能的重大变更 + +- 二进制协议支持(性能优化) +- 更灵活的流式模型 +- 插件依赖管理 + +#### 8.2.3 运行时优化 + +当前痛点:一插件一进程的开销 + +可能优化方向: + +1. **共享 Python 进程**:多个插件在同一进程(需要更严格的隔离) +2. **轻量级进程**:使用 uvloop 或其他优化 +3. **预加载机制**:常用插件预加载,减少启动延迟 + +#### 8.2.4 工具链完善 + +1. **插件脚手架**: + +```bash +astrbot-sdk init my_plugin +# 生成项目结构 +# 添加示例代码 +# 配置 pyproject.toml +``` + +2. **迁移助手**: + +```bash +astrbot-sdk migrate old_plugin +# 自动转换旧 API 到新 API +# 生成迁移报告 +``` + +3. **调试工具**: + +```bash +astrbot-sdk debug plugin_dir +# 本地运行插件 +# 交互式测试 +# 查看调用链 +``` + +### 8.3 文档改进建议 + +#### 8.3.1 统一文档结构 + +``` +docs/ +├── v4/ +│ ├── README.md # v4 总览 +│ ├── architecture.md # 架构说明 +│ ├── getting-started.md # 快速开始 +│ ├── api/ # API 文档 +│ │ ├── star.md +│ │ ├── context.md +│ │ ├── events.md +│ │ └── decorators.md +│ ├── runtime/ # 运行时文档 +│ │ ├── peer.md +│ │ ├── transport.md +│ │ └── capabilities.md +│ └── migration.md # 迁移指南 +└── legacy/ # 兼容文档(逐步废弃) + ├── overview.md + ├── compatibility.md + └── migration-guide.md +``` + +#### 8.3.2 代码示例中心化 + +创建统一的示例仓库: + +```bash +astrbot-sdk-examples/ +├── 01-basic-command/ # 基础命令 +├── 02-message-filter/ # 消息过滤 +├── 03-llm-integration/ # LLM 集成 +├── 04-database/ # 数据库使用 +├── 05-stream-capability/ # 流式能力 +├── 06-session-management/ # 会话管理 +└── legacy-examples/ # 旧版示例 +``` + +#### 8.3.3 自动化文档生成 + +使用工具从 docstring 生成 API 文档: + +```bash +# 生成 API 文档 +astrbot-sdk docs generate --output docs/api/ + +# 检查文档覆盖 +astrbot-sdk docs check +``` + +--- + +## 9. 总结 + +### 9.1 整体评价 + +AstrBot SDK v4 是一个**设计良好、架构清晰、兼容性考虑周全**的插件框架。其核心优势在于: + +1. **协议优先**:清晰的 v4 协议设计,为长期演进打下基础 +2. **分层合理**:插件、客户端、运行时、协议四层职责明确 +3. **兼容务实**:三级兼容策略在维护成本和兼容性之间取得平衡 +4. **测试完善**:单元测试、集成测试、契约测试覆盖全面 +5. **开发友好**:本地开发 harness、CLI 工具、完整文档 + +主要挑战在于: + +1. **复杂度较高**:多层抽象和兼容层带来认知负担 +2. **进程开销**:一插件一进程模型的启动和资源成本 +3. **维护负担**:兼容层和外部插件矩阵的持续维护 +4. **文档漂移**:多份文档和大量 CLAUDE.md 备注不利于学习 + +### 9.2 适用场景 + +**非常适合**: + +- 需要插件化架构的机器人系统 +- 需要进程隔离的高可靠性场景 +- 有大量旧插件需要兼容的迁移项目 +- 需要 LLM 集成的智能对话系统 + +**需要权衡**: + +- 资源受限的嵌入式环境(进程开销) +- 单机小规模项目(复杂度收益不大) +- 需要极低延迟的场景(跨进程通信) + +### 9.3 与竞品对比 + +| 特性 | AstrBot SDK v4 | Plugin A | Plugin B | +|------|----------------|-----------|----------| +| 协议设计 | 自研 v4 协议 | JSON-RPC 2.0 | HTTP REST | +| 进程模型 | Supervisor-Worker | 单进程 | 单进程 | +| 类型安全 | Pydantic 模型 | 动态类型 | 无验证 | +| 流式支持 | 原生支持 | 不支持 | SSE | +| 兼容性 | 三级兼容策略 | 无 | 无 | +| 测试覆盖 | 完善 | 基础 | 不足 | +| 学习曲线 | 中等 | 低 | 高 | + +### 9.4 最终建议 + +**对于 SDK 维护者**: + +1. 继续推进兼容层收口和简化 +2. 完善自动化测试和 CI/CD +3. 统一文档结构,减少 CLAUDE.md 依赖 +4. 评估进程模型的优化可能性 + +**对于插件开发者**: + +1. 新项目直接使用 v4 API +2. 旧项目逐步迁移到新 API +3. 充分利用本地开发 harness +4. 参考官方示例项目 + +**对于 Core 开发者**: + +1. 理解 v4 协议规范 +2. 实现全部 18 个内建 capability +3. 提供可靠的 Supervisor 实现 +4. 支持 Worker 进程管理和监控 + +--- + +**文档结束** + +如有疑问或建议,请参考: +- ARCHITECTURE.md - 当前架构文档 +- COMPATIBILITY_MATRIX.md - 兼容矩阵 +- CLAUDE.md - 开发者注意事项 +- tests_v4/README.md - 测试指南 diff --git a/src-new/astrbot_sdk/_legacy_api.py b/src-new/astrbot_sdk/_legacy_api.py index 52c28aaf5..4eb4a0e3f 100644 --- a/src-new/astrbot_sdk/_legacy_api.py +++ b/src-new/astrbot_sdk/_legacy_api.py @@ -1,1175 +1,38 @@ -"""旧版 API 的兼容实现。 +"""旧版 API 兼容层聚合入口。 -这个模块承接旧 ``Context`` / ``CommandComponent`` 的运行时行为, -把仍然可映射到 v4 的能力落到 ``Context`` 客户端上, -无法等价支持的旧接口则显式给出迁移错误,而不是静默降级。 +这个模块重导出来自 ``_legacy_context`` 和 ``_legacy_star`` 的所有公开符号, +供 ``compat.py``、``api/star/``、``api/components/`` 等外部导入路径使用。 + +不要在这里添加新的运行时逻辑;业务实现分别在 ``_legacy_context.py`` 和 +``_legacy_star.py`` 中维护。 + +注意:``logger`` 在此显式导入,以保持向后兼容性——部分测试通过 +``patch("astrbot_sdk._legacy_api.logger.warning")`` 路径拦截日志调用。 +由于 loguru 的 ``logger`` 是全局单例,这里的引用与 ``_legacy_context`` +内部使用的是同一个对象。 """ from __future__ import annotations -import inspect -import json -from collections import defaultdict -from collections.abc import Callable -from dataclasses import dataclass -from pathlib import Path -from typing import TYPE_CHECKING, Any +from loguru import logger as logger # noqa: PLC0414 — re-exported for patch compat -from loguru import logger - -from ._legacy_llm import ( - CompatLLMToolManager, - _CompatProviderRequest, - _legacy_llm_response, - _tool_parameters_from_legacy_args, +from ._legacy_context import ( + COMPAT_CONVERSATIONS_KEY, + MIGRATION_DOC_URL, + LegacyContext, + LegacyConversationManager, + _CompatHookEntry, + _iter_registered_component_methods, + _warn_once, + _warned_methods, ) -from .context import Context as NewContext -from .star import Star - -if TYPE_CHECKING: - from .api.provider.entities import LLMResponse -MIGRATION_DOC_URL = "https://docs.astrbot.app/migration/v3" -COMPAT_CONVERSATIONS_KEY = "__compat_conversations__" -_warned_methods: set[str] = set() - - -def _warn_once(old_name: str, replacement: str) -> None: - if old_name in _warned_methods: - return - _warned_methods.add(old_name) - logger.warning( - "[AstrBot] 警告:{} 已过时。请替换为:{}\n迁移文档:{}", - old_name, - replacement, - MIGRATION_DOC_URL, - ) - - -def _iter_registered_component_methods( - component: Any, -) -> list[tuple[str, Callable[..., Any]]]: - methods: list[tuple[str, Callable[..., Any]]] = [] - for attr_name, static_attr in inspect.getmembers_static(component): - if attr_name.startswith("_") or isinstance(static_attr, property): - continue - if not callable(static_attr) and not isinstance( - static_attr, (staticmethod, classmethod) - ): - continue - try: - bound_attr = getattr(component, attr_name) - except Exception: - continue - if callable(bound_attr): - methods.append((attr_name, bound_attr)) - return methods - - -@dataclass(slots=True) -class _CompatHookEntry: - name: str - priority: int - handler: Callable[..., Any] - - -class LegacyConversationManager: - """旧版会话管理器的兼容实现。 - - 会话数据通过 ``ctx.db`` 存在统一 key 下。 - 数据是否持久化取决于当前 db capability 的后端实现,而不是 compat 层本身。 - """ - - __compat_component_name__ = "ConversationManager" - - def __init__(self, parent: "LegacyContext") -> None: - self._parent = parent - self._counters: defaultdict[str, int] = defaultdict(int) - # 记录每个 unified_msg_origin 的当前会话 ID - self._current_conversations: dict[str, str] = {} - - def _ctx(self) -> NewContext: - return self._parent.require_runtime_context() - - async def _get_stored(self) -> dict[str, dict[str, Any]]: - """获取存储的所有会话数据。""" - ctx = self._ctx() - stored = await ctx.db.get(COMPAT_CONVERSATIONS_KEY) - return stored if isinstance(stored, dict) else {} - - async def _set_stored(self, stored: dict[str, dict[str, Any]]) -> None: - """保存会话数据。""" - ctx = self._ctx() - await ctx.db.set(COMPAT_CONVERSATIONS_KEY, stored) - - async def new_conversation( - self, - unified_msg_origin: str, - platform_id: str | None = None, - content: list[dict] | None = None, - title: str | None = None, - persona_id: str | None = None, - ) -> str: - """创建新会话并返回会话 ID。""" - ctx = self._ctx() - stored = await self._get_stored() - next_counter = self._counters[unified_msg_origin] - while True: - next_counter += 1 - conversation_id = f"{ctx.plugin_id}-conv-{next_counter}" - if conversation_id not in stored: - break - self._counters[unified_msg_origin] = next_counter - stored[conversation_id] = { - "unified_msg_origin": unified_msg_origin, - "platform_id": platform_id, - "content": content or [], - "title": title, - "persona_id": persona_id, - } - await self._set_stored(stored) - # 设置为当前会话 - self._current_conversations[unified_msg_origin] = conversation_id - return conversation_id - - async def switch_conversation( - self, unified_msg_origin: str, conversation_id: str - ) -> None: - """切换到指定会话。 - - Args: - unified_msg_origin: 统一消息来源 - conversation_id: 要切换到的会话 ID - """ - stored = await self._get_stored() - if conversation_id not in stored: - return - # 验证会话属于该 unified_msg_origin - conv_data = stored[conversation_id] - if conv_data.get("unified_msg_origin") != unified_msg_origin: - return - self._current_conversations[unified_msg_origin] = conversation_id - - async def delete_conversation( - self, - unified_msg_origin: str, - conversation_id: str | None = None, - ) -> None: - """删除指定会话。 - - 当 conversation_id 为 None 时,删除当前会话。 - - Args: - unified_msg_origin: 统一消息来源 - conversation_id: 要删除的会话 ID,为 None 时删除当前会话 - """ - # 如果 conversation_id 为 None,使用当前会话 - if conversation_id is None: - conversation_id = self._current_conversations.get(unified_msg_origin) - if conversation_id is None: - return - - stored = await self._get_stored() - if conversation_id not in stored: - return - conv_data = stored[conversation_id] - if conv_data.get("unified_msg_origin") != unified_msg_origin: - return - del stored[conversation_id] - await self._set_stored(stored) - # 如果删除的是当前会话,清除当前会话记录 - if self._current_conversations.get(unified_msg_origin) == conversation_id: - del self._current_conversations[unified_msg_origin] - - async def get_curr_conversation_id(self, unified_msg_origin: str) -> str | None: - """获取当前会话 ID。 - - Args: - unified_msg_origin: 统一消息来源 - - Returns: - 当前会话 ID,若无则返回 None - """ - return self._current_conversations.get(unified_msg_origin) - - async def get_conversation( - self, - unified_msg_origin: str, - conversation_id: str, - create_if_not_exists: bool = False, - ) -> dict[str, Any] | None: - """获取指定会话的数据。 - - Args: - unified_msg_origin: 统一消息来源 - conversation_id: 会话 ID - create_if_not_exists: 如果会话不存在,是否创建新会话 - - Returns: - 会话数据字典,不存在则返回 None - """ - stored = await self._get_stored() - conv = stored.get(conversation_id) - if conv is None and create_if_not_exists: - # 创建新会话 - conv = { - "unified_msg_origin": unified_msg_origin, - "platform_id": None, - "content": [], - "title": None, - "persona_id": None, - } - stored[conversation_id] = conv - await self._set_stored(stored) - self._current_conversations[unified_msg_origin] = conversation_id - return conv - - async def get_conversations( - self, - unified_msg_origin: str | None = None, - platform_id: str | None = None, - ) -> list[dict[str, Any]]: - """获取会话列表。 - - Args: - unified_msg_origin: 统一消息来源,可选 - platform_id: 平台 ID,可选 - - Returns: - 会话列表,每个元素包含 conversation_id 和会话数据 - """ - stored = await self._get_stored() - result = [] - for conv_id, conv_data in stored.items(): - # 按 unified_msg_origin 过滤 - if unified_msg_origin is not None: - if conv_data.get("unified_msg_origin") != unified_msg_origin: - continue - # 按 platform_id 过滤 - if platform_id is not None: - if conv_data.get("platform_id") != platform_id: - continue - result.append({"conversation_id": conv_id, **conv_data}) - return result - - async def update_conversation( - self, - unified_msg_origin: str, - conversation_id: str | None = None, - history: list[dict] | None = None, - title: str | None = None, - persona_id: str | None = None, - ) -> None: - """更新会话数据。 - - Args: - unified_msg_origin: 统一消息来源 - conversation_id: 会话 ID,为 None 时更新当前会话 - history: 对话历史记录 - title: 会话标题 - persona_id: Persona ID - """ - # 如果 conversation_id 为 None,使用当前会话 - if conversation_id is None: - conversation_id = self._current_conversations.get(unified_msg_origin) - if conversation_id is None: - return - - stored = await self._get_stored() - if conversation_id not in stored: - return - - updates: dict[str, Any] = {} - if history is not None: - updates["content"] = history - if title is not None: - updates["title"] = title - if persona_id is not None: - updates["persona_id"] = persona_id - - stored[conversation_id].update(updates) - await self._set_stored(stored) - - async def delete_conversations_by_user_id(self, unified_msg_origin: str) -> None: - """删除指定用户的所有会话。 - - Args: - unified_msg_origin: 统一消息来源 - """ - stored = await self._get_stored() - to_delete = [ - conv_id - for conv_id, conv_data in stored.items() - if conv_data.get("unified_msg_origin") == unified_msg_origin - ] - for conv_id in to_delete: - del stored[conv_id] - await self._set_stored(stored) - # 清除当前会话记录 - if unified_msg_origin in self._current_conversations: - del self._current_conversations[unified_msg_origin] - - async def add_message_pair( - self, - cid: str, - user_message: str | dict, - assistant_message: str | dict, - ) -> None: - """向会话添加消息对。 - - Args: - cid: 会话 ID - user_message: 用户消息 - assistant_message: 助手消息 - """ - stored = await self._get_stored() - if cid not in stored: - return - content = stored[cid].get("content", []) - # 处理消息格式 - user_msg = ( - user_message - if isinstance(user_message, dict) - else {"role": "user", "content": user_message} - ) - assistant_msg = ( - assistant_message - if isinstance(assistant_message, dict) - else {"role": "assistant", "content": assistant_message} - ) - content.append(user_msg) - content.append(assistant_msg) - stored[cid]["content"] = content - await self._set_stored(stored) - - async def update_conversation_title( - self, - unified_msg_origin: str, - title: str, - conversation_id: str | None = None, - ) -> None: - """更新会话标题。 - - Args: - unified_msg_origin: 统一消息来源 - title: 会话标题 - conversation_id: 会话 ID,为 None 时更新当前会话 - - Deprecated: - 请使用 update_conversation() 的 title 参数。 - """ - await self.update_conversation(unified_msg_origin, conversation_id, title=title) - - async def update_conversation_persona_id( - self, - unified_msg_origin: str, - persona_id: str, - conversation_id: str | None = None, - ) -> None: - """更新会话 Persona ID。 - - Args: - unified_msg_origin: 统一消息来源 - persona_id: Persona ID - conversation_id: 会话 ID,为 None 时更新当前会话 - - Deprecated: - 请使用 update_conversation() 的 persona_id 参数。 - """ - await self.update_conversation( - unified_msg_origin, conversation_id, persona_id=persona_id - ) - - async def get_filtered_conversations(self, *args: Any, **kwargs: Any) -> Any: - """兼容旧版会话过滤接口。""" - unified_msg_origin = kwargs.get("unified_msg_origin") - platform_id = kwargs.get("platform_id") - keyword = kwargs.get("keyword") or kwargs.get("query") - conversations = await self.get_conversations( - unified_msg_origin=unified_msg_origin, - platform_id=platform_id, - ) - if not isinstance(keyword, str) or not keyword: - return conversations - filtered: list[dict[str, Any]] = [] - for conversation in conversations: - haystack = json.dumps(conversation, ensure_ascii=False) - if keyword in haystack: - filtered.append(conversation) - return filtered - - async def get_human_readable_context(self, *args: Any, **kwargs: Any) -> Any: - """把兼容会话内容格式化为可读文本。""" - unified_msg_origin = kwargs.get("unified_msg_origin") - conversation_id = kwargs.get("conversation_id") - if conversation_id is None and isinstance(unified_msg_origin, str): - conversation_id = await self.get_curr_conversation_id(unified_msg_origin) - if not isinstance(conversation_id, str) or not conversation_id: - return "" - conversation = await self.get_conversation( - unified_msg_origin or "", - conversation_id, - create_if_not_exists=False, - ) - if not isinstance(conversation, dict): - return "" - lines: list[str] = [] - for item in conversation.get("content", []): - if not isinstance(item, dict): - continue - role = str(item.get("role") or "unknown") - content = item.get("content") - if isinstance(content, list): - rendered = json.dumps(content, ensure_ascii=False) - else: - rendered = str(content or "") - lines.append(f"{role}: {rendered}".rstrip()) - return "\n".join(lines) - - -class LegacyContext: - """旧版 ``Context`` 的兼容外观。""" - - def __init__(self, plugin_id: str) -> None: - self.plugin_id = plugin_id - self._runtime_context: NewContext | None = None - self._registered_managers: dict[str, Any] = {} - self._registered_functions: dict[str, Callable[..., Any]] = {} - self._compat_hooks: defaultdict[str, list[_CompatHookEntry]] = defaultdict(list) - self._llm_tools = CompatLLMToolManager() - self.conversation_manager = LegacyConversationManager(self) - self._register_component(self.conversation_manager) - - def bind_runtime_context(self, runtime_context: NewContext) -> None: - self._runtime_context = runtime_context - - def require_runtime_context(self) -> NewContext: - if self._runtime_context is None: - raise RuntimeError("LegacyContext 尚未绑定运行时 Context") - return self._runtime_context - - def get_llm_tool_manager(self) -> CompatLLMToolManager: - return self._llm_tools - - def activate_llm_tool(self, name: str) -> bool: - return self._llm_tools.activate_llm_tool(name) - - def deactivate_llm_tool(self, name: str) -> bool: - return self._llm_tools.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._llm_tools.add_func(name, func_args, desc, func_obj) - - def unregister_llm_tool(self, name: str) -> None: - self._llm_tools.remove_func(name) - - def get_config(self) -> dict[str, Any]: - runtime_context = self._runtime_context - if runtime_context is None: - return {} - config = getattr(runtime_context, "_astrbot_config", None) - return dict(config) if isinstance(config, dict) else {} - - def _runtime_config(self) -> Any: - from .api.basic.astrbot_config import AstrBotConfig - - runtime_context = self._runtime_context - config = ( - getattr(runtime_context, "_astrbot_config", None) - if runtime_context - else None - ) - if isinstance(config, AstrBotConfig): - return config - if isinstance(config, dict): - return AstrBotConfig(dict(config)) - return AstrBotConfig({}) - - @staticmethod - def _merge_llm_kwargs( - *, - chat_provider_id: str, - kwargs: dict[str, Any], - ) -> dict[str, Any]: - merged = dict(kwargs) - if chat_provider_id: - merged.setdefault("provider_id", chat_provider_id) - return merged - - @staticmethod - def _apply_request_overrides( - call_kwargs: dict[str, Any], - request: _CompatProviderRequest, - ) -> dict[str, Any]: - updated = dict(call_kwargs) - if request.model: - updated["model"] = request.model - return updated - - @staticmethod - def _component_names(component: Any) -> list[str]: - names = [component.__class__.__name__] - compat_name = getattr(component, "__compat_component_name__", None) - if isinstance(compat_name, str) and compat_name and compat_name not in names: - names.insert(0, compat_name) - return names - - def _register_hook( - self, - name: str, - handler: Callable[..., Any], - *, - priority: int = 0, - ) -> None: - self._compat_hooks[name].append( - _CompatHookEntry(name=name, priority=priority, handler=handler) - ) - self._compat_hooks[name].sort(key=lambda item: item.priority, reverse=True) - - def _register_compat_component(self, component: Any) -> None: - from .api.event.filter import ( - get_compat_hook_metas, - get_compat_llm_tool_meta, - ) - - for _attr_name, attr in _iter_registered_component_methods(component): - tool_meta = get_compat_llm_tool_meta(attr) - if tool_meta is not None: - self._llm_tools.add_tool( - name=tool_meta.name, - description=tool_meta.description, - parameters=_tool_parameters_from_legacy_args(tool_meta.parameters), - handler=attr, - ) - for hook_meta in get_compat_hook_metas(attr): - self._register_hook( - hook_meta.name, - attr, - priority=hook_meta.priority, - ) - - @staticmethod - def _legacy_event(event: Any | None): - if event is None: - return None - from .api.event.astr_message_event import AstrMessageEvent - - if isinstance(event, AstrMessageEvent): - return event - return AstrMessageEvent.from_message_event(event) - - @staticmethod - def _hook_type_injection( - annotation: Any, - available: dict[str, Any], - ) -> Any: - from .api.event.astr_message_event import AstrMessageEvent - from .api.provider.entities import LLMResponse - from .context import Context as RuntimeContext - - if annotation is Any or annotation is inspect.Signature.empty: - return None - if annotation is AstrMessageEvent: - return available.get("event") - if annotation is RuntimeContext or annotation is NewContext: - return available.get("context") - if annotation is LegacyContext: - return available.get("legacy_context") - if annotation is LLMResponse: - return available.get("response") - return None - - async def _call_with_available( - self, - handler: Callable[..., Any], - available: dict[str, Any], - ) -> Any: - signature = inspect.signature(handler) - args: list[Any] = [] - kwargs: dict[str, Any] = {} - for parameter in signature.parameters.values(): - injected = None - if parameter.name in available: - injected = available[parameter.name] - else: - injected = self._hook_type_injection(parameter.annotation, available) - if injected is None: - if parameter.default is not parameter.empty: - continue - continue - if parameter.kind in ( - inspect.Parameter.POSITIONAL_ONLY, - inspect.Parameter.POSITIONAL_OR_KEYWORD, - ): - args.append(injected) - elif parameter.kind == inspect.Parameter.KEYWORD_ONLY: - kwargs[parameter.name] = injected - result = handler(*args, **kwargs) - if inspect.isasyncgen(result): - final_value = None - async for item in result: - final_value = item - await self._consume_tool_result( - available.get("event"), - available.get("context"), - item, - ) - return final_value - if inspect.isawaitable(result): - return await result - return result - - async def _run_compat_hook( - self, - name: str, - **available: Any, - ) -> list[Any]: - hook_results: list[Any] = [] - for entry in self._compat_hooks.get(name, []): - hook_results.append( - await self._call_with_available(entry.handler, available) - ) - return hook_results - - async def _consume_tool_result( - self, - event: Any | None, - runtime_context: NewContext | None, - item: Any, - ) -> None: - if event is None: - return - from .api.event.event_result import MessageEventResult - from .api.message.chain import MessageChain - - legacy_event = self._legacy_event(event) - if legacy_event is None: - return - if isinstance(item, MessageEventResult): - if ( - item.chain - and runtime_context is not None - and not item.is_plain_text_only() - ): - await runtime_context.platform.send_chain( - legacy_event.session_ref or legacy_event.session_id, - item.to_payload(), - ) - return - plain_text = item.get_plain_text() - if plain_text: - await legacy_event.reply(plain_text) - return - if isinstance(item, MessageChain): - if ( - item.chain - and runtime_context is not None - and not item.is_plain_text_only() - ): - await runtime_context.platform.send_chain( - legacy_event.session_ref or legacy_event.session_id, - item.to_payload(), - ) - return - plain_text = item.get_plain_text() - if plain_text: - await legacy_event.reply(plain_text) - return - if isinstance(item, str): - await legacy_event.reply(item) - - async def _invoke_llm_tool( - self, - *, - tool_name: str, - tool_args: dict[str, Any], - event: Any | None, - ) -> str: - tool = self._llm_tools.get_func(tool_name) - if tool is None or not tool.active: - return f"tool '{tool_name}' not found" - legacy_event = self._legacy_event(event) - runtime_context = self.require_runtime_context() - await self._run_compat_hook( - "on_using_llm_tool", - event=legacy_event, - context=runtime_context, - legacy_context=self, - tool=tool, - tool_args=tool_args, - ) - tool_result = await self._call_with_available( - tool.handler, - { - **tool_args, - "event": legacy_event, - "context": runtime_context, - "ctx": runtime_context, - "legacy_context": self, - }, - ) - if isinstance(tool_result, str): - normalized = tool_result - elif tool_result is None: - normalized = "" - else: - normalized = str(tool_result) - await self._run_compat_hook( - "on_llm_tool_respond", - event=legacy_event, - context=runtime_context, - legacy_context=self, - tool=tool, - tool_args=tool_args, - tool_result=normalized, - ) - return normalized - - def _register_component(self, *components: Any) -> None: - """保留旧版按名称暴露组件方法的兼容链路。""" - for component in components: - for class_name in self._component_names(component): - self._registered_managers[class_name] = component - for attr_name, attr in _iter_registered_component_methods(component): - self._registered_functions[f"{class_name}.{attr_name}"] = attr - self._register_compat_component(component) - - async def execute_registered_function( - self, - func_full_name: str, - args: dict[str, Any] | None = None, - ) -> Any: - if args is None: - call_args: dict[str, Any] = {} - elif isinstance(args, dict): - call_args = args - else: - raise TypeError("LegacyContext 调用参数必须是 dict") - - func = self._registered_functions.get(func_full_name) - if func is None: - raise ValueError(f"Function not found: {func_full_name}") - - result = func(**call_args) - if inspect.isawaitable(result): - return await result - return result - - async def call_context_function( - self, - func_full_name: str, - args: dict[str, Any] | None = None, - ) -> dict[str, Any]: - return { - "data": await self.execute_registered_function(func_full_name, args), - } - - async def llm_generate( - self, - chat_provider_id: str, - prompt: str | None = None, - image_urls: list[str] | None = None, - tools: Any | None = None, - system_prompt: str | None = None, - contexts: list[dict] | None = None, - event: Any | None = None, - **kwargs: Any, - ) -> LLMResponse: - _warn_once("context.llm_generate()", "ctx.llm.chat_raw(...)") - ctx = self.require_runtime_context() - call_kwargs = self._merge_llm_kwargs( - chat_provider_id=chat_provider_id, - kwargs=kwargs, - ) - legacy_event = self._legacy_event(event) - request = _CompatProviderRequest( - prompt=prompt or "", - session_id=legacy_event.session_id if legacy_event is not None else "", - image_urls=list(image_urls or []), - contexts=list(contexts or []), - system_prompt=system_prompt or "", - model=call_kwargs.get("model"), - ) - await self._run_compat_hook( - "on_waiting_llm_request", - event=legacy_event, - context=ctx, - legacy_context=self, - ) - await self._run_compat_hook( - "on_llm_request", - event=legacy_event, - context=ctx, - legacy_context=self, - request=request, - ) - call_kwargs = self._apply_request_overrides(call_kwargs, request) - response = await ctx.llm.chat_raw( - request.prompt or "", - system=request.system_prompt or None, - history=request.contexts or [], - image_urls=request.image_urls or [], - tools=tools, - **call_kwargs, - ) - legacy_response = _legacy_llm_response(response) - await self._run_compat_hook( - "on_llm_response", - event=legacy_event, - context=ctx, - legacy_context=self, - response=legacy_response, - ) - return legacy_response - - async def tool_loop_agent( - self, - chat_provider_id: str, - prompt: str | None = None, - image_urls: list[str] | None = None, - tools: Any | None = None, - system_prompt: str | None = None, - contexts: list[dict] | None = None, - max_steps: int = 30, - event: Any | None = None, - **kwargs: Any, - ) -> LLMResponse: - from .api.provider.entities import LLMResponse - - _warn_once("context.tool_loop_agent()", "compat local tool loop") - ctx = self.require_runtime_context() - call_kwargs = self._merge_llm_kwargs( - chat_provider_id=chat_provider_id, - kwargs=kwargs, - ) - legacy_event = self._legacy_event(event) - history = list(contexts or []) - request_prompt = prompt or "" - combined_tools = list(self._llm_tools.get_func_desc_openai_style()) - if isinstance(tools, list): - combined_tools.extend(item for item in tools if isinstance(item, dict)) - elif tools is not None: - openai_schema = getattr(tools, "openai_schema", None) - if callable(openai_schema): - extra_tools = openai_schema() - if isinstance(extra_tools, list): - combined_tools.extend( - item for item in extra_tools if isinstance(item, dict) - ) - - final_response = LLMResponse(role="assistant") - for _step in range(max_steps): - request = _CompatProviderRequest( - prompt=request_prompt, - session_id=legacy_event.session_id if legacy_event is not None else "", - image_urls=list(image_urls or []), - contexts=list(history), - system_prompt=system_prompt or "", - model=call_kwargs.get("model"), - ) - await self._run_compat_hook( - "on_waiting_llm_request", - event=legacy_event, - context=ctx, - legacy_context=self, - ) - await self._run_compat_hook( - "on_llm_request", - event=legacy_event, - context=ctx, - legacy_context=self, - request=request, - ) - call_kwargs = self._apply_request_overrides(call_kwargs, request) - response = await ctx.llm.chat_raw( - request.prompt or "", - system=request.system_prompt or None, - history=request.contexts or [], - image_urls=request.image_urls or [], - tools=combined_tools or None, - max_steps=max_steps, - **call_kwargs, - ) - final_response = _legacy_llm_response(response) - await self._run_compat_hook( - "on_llm_response", - event=legacy_event, - context=ctx, - legacy_context=self, - response=final_response, - ) - if not final_response.tools_call_name: - return final_response - - history.append( - { - "role": "assistant", - "content": final_response.completion_text, - "tool_calls": final_response.to_openai_tool_calls(), - } - ) - for tool_name, tool_args, tool_call_id in zip( - final_response.tools_call_name, - final_response.tools_call_args, - final_response.tools_call_ids, - strict=False, - ): - tool_result = await self._invoke_llm_tool( - tool_name=tool_name, - tool_args=tool_args, - event=legacy_event, - ) - history.append( - { - "role": "tool", - "tool_call_id": tool_call_id, - "name": tool_name, - "content": tool_result, - } - ) - request_prompt = "" - - return final_response - - async def send_message(self, session: str, message_chain: Any) -> None: - _warn_once( - "context.send_message()", - "ctx.platform.send(...) / ctx.platform.send_chain(...)", - ) - ctx = self.require_runtime_context() - chain = getattr(message_chain, "chain", None) - to_payload = getattr(message_chain, "to_payload", None) - is_plain_text_only = getattr(message_chain, "is_plain_text_only", None) - if ( - isinstance(chain, list) - and callable(to_payload) - and not (callable(is_plain_text_only) and is_plain_text_only()) - ): - await ctx.platform.send_chain(session, to_payload()) - return - - # 旧版插件也可能传纯文本对象,compat 层保留文本兜底。 - if hasattr(message_chain, "get_plain_text") and callable( - message_chain.get_plain_text - ): - text = message_chain.get_plain_text() - elif hasattr(message_chain, "to_text") and callable(message_chain.to_text): - text = message_chain.to_text() - else: - text = str(message_chain) - await ctx.platform.send(session, text) - - async def add_llm_tools(self, *tools: Any) -> None: - for tool in tools: - name = getattr(tool, "name", None) - if not isinstance(name, str) or not name: - raise TypeError("add_llm_tools() 需要带 name 的工具对象") - handler = getattr(tool, "handler", None) - if not callable(handler): - raise TypeError("add_llm_tools() 需要工具对象提供可调用的 handler") - parameters = getattr(tool, "parameters", None) - if not isinstance(parameters, dict): - func_args = getattr(tool, "func_args", None) - if isinstance(func_args, list): - parameters = _tool_parameters_from_legacy_args(func_args) - else: - parameters = {"type": "object", "properties": {}, "required": []} - description = str(getattr(tool, "description", "") or "") - self._llm_tools.add_tool( - name=name, - description=description, - parameters=parameters, - handler=handler, - ) - - async def put_kv_data(self, key: str, value: Any) -> None: - _warn_once("context.put_kv_data()", "ctx.db.set(key, value)") - ctx = self.require_runtime_context() - await ctx.db.set(key, value) - - async def get_kv_data(self, key: str, default: Any = None) -> Any: - _warn_once("context.get_kv_data()", "ctx.db.get(key)") - ctx = self.require_runtime_context() - value = await ctx.db.get(key) - return default if value is None else value - - async def delete_kv_data(self, key: str) -> None: - _warn_once("context.delete_kv_data()", "ctx.db.delete(key)") - ctx = self.require_runtime_context() - await ctx.db.delete(key) - - async def get_registered_star(self, star_name: str) -> Any: - ctx = self.require_runtime_context() - return await ctx.metadata.get_plugin(star_name) - - async def get_all_stars(self) -> list[Any]: - ctx = self.require_runtime_context() - return await ctx.metadata.list_plugins() - - -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 - +from ._legacy_star import CommandComponent, LegacyStar, StarTools, register +# Historical alias: ``Context`` was the original public name for ``LegacyContext``. Context = LegacyContext __all__ = [ + "COMPAT_CONVERSATIONS_KEY", "CommandComponent", "Context", "LegacyContext", @@ -1177,5 +40,10 @@ __all__ = [ "LegacyStar", "MIGRATION_DOC_URL", "StarTools", + "_CompatHookEntry", + "_iter_registered_component_methods", + "_warn_once", + "_warned_methods", + "logger", "register", ] diff --git a/src-new/astrbot_sdk/_legacy_context.py b/src-new/astrbot_sdk/_legacy_context.py new file mode 100644 index 000000000..3f2330317 --- /dev/null +++ b/src-new/astrbot_sdk/_legacy_context.py @@ -0,0 +1,1014 @@ +"""旧版 API 兼容层 — 会话管理与 Context 实现。 + +这个模块承接旧 ``Context`` / ``ConversationManager`` 的运行时行为, +把仍然可映射到 v4 的能力落到 v4 ``Context`` 客户端上, +无法等价支持的旧接口则显式给出迁移错误,而不是静默降级。 + +不要从 ``_legacy_star`` 导入任何符号(避免循环依赖)。 +外部代码应通过 ``_legacy_api`` 聚合入口导入。 +""" + +from __future__ import annotations + +import inspect +import json +from collections import defaultdict +from collections.abc import Callable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from loguru import logger + +from ._legacy_llm import ( + CompatLLMToolManager, + _CompatProviderRequest, + _legacy_llm_response, + _tool_parameters_from_legacy_args, +) +from .context import Context as NewContext + +if TYPE_CHECKING: + from .api.provider.entities import LLMResponse + +MIGRATION_DOC_URL = "https://docs.astrbot.app/migration/v3" +COMPAT_CONVERSATIONS_KEY = "__compat_conversations__" +_warned_methods: set[str] = set() + + +def _warn_once(old_name: str, replacement: str) -> None: + if old_name in _warned_methods: + return + _warned_methods.add(old_name) + logger.warning( + "[AstrBot] 警告:{} 已过时。请替换为:{}\n迁移文档:{}", + old_name, + replacement, + MIGRATION_DOC_URL, + ) + + +def _iter_registered_component_methods( + component: Any, +) -> list[tuple[str, Callable[..., Any]]]: + methods: list[tuple[str, Callable[..., Any]]] = [] + for attr_name, static_attr in inspect.getmembers_static(component): + if attr_name.startswith("_") or isinstance(static_attr, property): + continue + if not callable(static_attr) and not isinstance( + static_attr, (staticmethod, classmethod) + ): + continue + try: + bound_attr = getattr(component, attr_name) + except Exception: + continue + if callable(bound_attr): + methods.append((attr_name, bound_attr)) + return methods + + +@dataclass(slots=True) +class _CompatHookEntry: + name: str + priority: int + handler: Callable[..., Any] + + +class LegacyConversationManager: + """旧版会话管理器的兼容实现。 + + 会话数据通过 ``ctx.db`` 存在统一 key 下。 + 数据是否持久化取决于当前 db capability 的后端实现,而不是 compat 层本身。 + """ + + __compat_component_name__ = "ConversationManager" + + def __init__(self, parent: "LegacyContext") -> None: + self._parent = parent + self._counters: defaultdict[str, int] = defaultdict(int) + # 记录每个 unified_msg_origin 的当前会话 ID + self._current_conversations: dict[str, str] = {} + + def _ctx(self) -> NewContext: + return self._parent.require_runtime_context() + + async def _get_stored(self) -> dict[str, dict[str, Any]]: + """获取存储的所有会话数据。""" + ctx = self._ctx() + stored = await ctx.db.get(COMPAT_CONVERSATIONS_KEY) + return stored if isinstance(stored, dict) else {} + + async def _set_stored(self, stored: dict[str, dict[str, Any]]) -> None: + """保存会话数据。""" + ctx = self._ctx() + await ctx.db.set(COMPAT_CONVERSATIONS_KEY, stored) + + async def new_conversation( + self, + unified_msg_origin: str, + platform_id: str | None = None, + content: list[dict] | None = None, + title: str | None = None, + persona_id: str | None = None, + ) -> str: + """创建新会话并返回会话 ID。""" + ctx = self._ctx() + stored = await self._get_stored() + next_counter = self._counters[unified_msg_origin] + while True: + next_counter += 1 + conversation_id = f"{ctx.plugin_id}-conv-{next_counter}" + if conversation_id not in stored: + break + self._counters[unified_msg_origin] = next_counter + stored[conversation_id] = { + "unified_msg_origin": unified_msg_origin, + "platform_id": platform_id, + "content": content or [], + "title": title, + "persona_id": persona_id, + } + await self._set_stored(stored) + # 设置为当前会话 + self._current_conversations[unified_msg_origin] = conversation_id + return conversation_id + + async def switch_conversation( + self, unified_msg_origin: str, conversation_id: str + ) -> None: + """切换到指定会话。 + + Args: + unified_msg_origin: 统一消息来源 + conversation_id: 要切换到的会话 ID + """ + stored = await self._get_stored() + if conversation_id not in stored: + return + # 验证会话属于该 unified_msg_origin + conv_data = stored[conversation_id] + if conv_data.get("unified_msg_origin") != unified_msg_origin: + return + self._current_conversations[unified_msg_origin] = conversation_id + + async def delete_conversation( + self, + unified_msg_origin: str, + conversation_id: str | None = None, + ) -> None: + """删除指定会话。 + + 当 conversation_id 为 None 时,删除当前会话。 + + Args: + unified_msg_origin: 统一消息来源 + conversation_id: 要删除的会话 ID,为 None 时删除当前会话 + """ + # 如果 conversation_id 为 None,使用当前会话 + if conversation_id is None: + conversation_id = self._current_conversations.get(unified_msg_origin) + if conversation_id is None: + return + + stored = await self._get_stored() + if conversation_id not in stored: + return + conv_data = stored[conversation_id] + if conv_data.get("unified_msg_origin") != unified_msg_origin: + return + del stored[conversation_id] + await self._set_stored(stored) + # 如果删除的是当前会话,清除当前会话记录 + if self._current_conversations.get(unified_msg_origin) == conversation_id: + del self._current_conversations[unified_msg_origin] + + async def get_curr_conversation_id(self, unified_msg_origin: str) -> str | None: + """获取当前会话 ID。 + + Args: + unified_msg_origin: 统一消息来源 + + Returns: + 当前会话 ID,若无则返回 None + """ + return self._current_conversations.get(unified_msg_origin) + + async def get_conversation( + self, + unified_msg_origin: str, + conversation_id: str, + create_if_not_exists: bool = False, + ) -> dict[str, Any] | None: + """获取指定会话的数据。 + + Args: + unified_msg_origin: 统一消息来源 + conversation_id: 会话 ID + create_if_not_exists: 如果会话不存在,是否创建新会话 + + Returns: + 会话数据字典,不存在则返回 None + """ + stored = await self._get_stored() + conv = stored.get(conversation_id) + if conv is None and create_if_not_exists: + # 创建新会话 + conv = { + "unified_msg_origin": unified_msg_origin, + "platform_id": None, + "content": [], + "title": None, + "persona_id": None, + } + stored[conversation_id] = conv + await self._set_stored(stored) + self._current_conversations[unified_msg_origin] = conversation_id + return conv + + async def get_conversations( + self, + unified_msg_origin: str | None = None, + platform_id: str | None = None, + ) -> list[dict[str, Any]]: + """获取会话列表。 + + Args: + unified_msg_origin: 统一消息来源,可选 + platform_id: 平台 ID,可选 + + Returns: + 会话列表,每个元素包含 conversation_id 和会话数据 + """ + stored = await self._get_stored() + result = [] + for conv_id, conv_data in stored.items(): + # 按 unified_msg_origin 过滤 + if unified_msg_origin is not None: + if conv_data.get("unified_msg_origin") != unified_msg_origin: + continue + # 按 platform_id 过滤 + if platform_id is not None: + if conv_data.get("platform_id") != platform_id: + continue + result.append({"conversation_id": conv_id, **conv_data}) + return result + + async def update_conversation( + self, + unified_msg_origin: str, + conversation_id: str | None = None, + history: list[dict] | None = None, + title: str | None = None, + persona_id: str | None = None, + ) -> None: + """更新会话数据。 + + Args: + unified_msg_origin: 统一消息来源 + conversation_id: 会话 ID,为 None 时更新当前会话 + history: 对话历史记录 + title: 会话标题 + persona_id: Persona ID + """ + # 如果 conversation_id 为 None,使用当前会话 + if conversation_id is None: + conversation_id = self._current_conversations.get(unified_msg_origin) + if conversation_id is None: + return + + stored = await self._get_stored() + if conversation_id not in stored: + return + + updates: dict[str, Any] = {} + if history is not None: + updates["content"] = history + if title is not None: + updates["title"] = title + if persona_id is not None: + updates["persona_id"] = persona_id + + stored[conversation_id].update(updates) + await self._set_stored(stored) + + async def delete_conversations_by_user_id(self, unified_msg_origin: str) -> None: + """删除指定用户的所有会话。 + + Args: + unified_msg_origin: 统一消息来源 + """ + stored = await self._get_stored() + to_delete = [ + conv_id + for conv_id, conv_data in stored.items() + if conv_data.get("unified_msg_origin") == unified_msg_origin + ] + for conv_id in to_delete: + del stored[conv_id] + await self._set_stored(stored) + # 清除当前会话记录 + if unified_msg_origin in self._current_conversations: + del self._current_conversations[unified_msg_origin] + + async def add_message_pair( + self, + cid: str, + user_message: str | dict, + assistant_message: str | dict, + ) -> None: + """向会话添加消息对。 + + Args: + cid: 会话 ID + user_message: 用户消息 + assistant_message: 助手消息 + """ + stored = await self._get_stored() + if cid not in stored: + return + content = stored[cid].get("content", []) + # 处理消息格式 + user_msg = ( + user_message + if isinstance(user_message, dict) + else {"role": "user", "content": user_message} + ) + assistant_msg = ( + assistant_message + if isinstance(assistant_message, dict) + else {"role": "assistant", "content": assistant_message} + ) + content.append(user_msg) + content.append(assistant_msg) + stored[cid]["content"] = content + await self._set_stored(stored) + + async def update_conversation_title( + self, + unified_msg_origin: str, + title: str, + conversation_id: str | None = None, + ) -> None: + """更新会话标题。 + + Args: + unified_msg_origin: 统一消息来源 + title: 会话标题 + conversation_id: 会话 ID,为 None 时更新当前会话 + + Deprecated: + 请使用 update_conversation() 的 title 参数。 + """ + await self.update_conversation(unified_msg_origin, conversation_id, title=title) + + async def update_conversation_persona_id( + self, + unified_msg_origin: str, + persona_id: str, + conversation_id: str | None = None, + ) -> None: + """更新会话 Persona ID。 + + Args: + unified_msg_origin: 统一消息来源 + persona_id: Persona ID + conversation_id: 会话 ID,为 None 时更新当前会话 + + Deprecated: + 请使用 update_conversation() 的 persona_id 参数。 + """ + await self.update_conversation( + unified_msg_origin, conversation_id, persona_id=persona_id + ) + + async def get_filtered_conversations(self, *args: Any, **kwargs: Any) -> Any: + """兼容旧版会话过滤接口。""" + unified_msg_origin = kwargs.get("unified_msg_origin") + platform_id = kwargs.get("platform_id") + keyword = kwargs.get("keyword") or kwargs.get("query") + conversations = await self.get_conversations( + unified_msg_origin=unified_msg_origin, + platform_id=platform_id, + ) + if not isinstance(keyword, str) or not keyword: + return conversations + filtered: list[dict[str, Any]] = [] + for conversation in conversations: + haystack = json.dumps(conversation, ensure_ascii=False) + if keyword in haystack: + filtered.append(conversation) + return filtered + + async def get_human_readable_context(self, *args: Any, **kwargs: Any) -> Any: + """把兼容会话内容格式化为可读文本。""" + unified_msg_origin = kwargs.get("unified_msg_origin") + conversation_id = kwargs.get("conversation_id") + if conversation_id is None and isinstance(unified_msg_origin, str): + conversation_id = await self.get_curr_conversation_id(unified_msg_origin) + if not isinstance(conversation_id, str) or not conversation_id: + return "" + conversation = await self.get_conversation( + unified_msg_origin or "", + conversation_id, + create_if_not_exists=False, + ) + if not isinstance(conversation, dict): + return "" + lines: list[str] = [] + for item in conversation.get("content", []): + if not isinstance(item, dict): + continue + role = str(item.get("role") or "unknown") + content = item.get("content") + if isinstance(content, list): + rendered = json.dumps(content, ensure_ascii=False) + else: + rendered = str(content or "") + lines.append(f"{role}: {rendered}".rstrip()) + return "\n".join(lines) + + +class LegacyContext: + """旧版 ``Context`` 的兼容外观。""" + + def __init__(self, plugin_id: str) -> None: + self.plugin_id = plugin_id + self._runtime_context: NewContext | None = None + self._registered_managers: dict[str, Any] = {} + self._registered_functions: dict[str, Callable[..., Any]] = {} + self._compat_hooks: defaultdict[str, list[_CompatHookEntry]] = defaultdict(list) + self._llm_tools = CompatLLMToolManager() + self.conversation_manager = LegacyConversationManager(self) + self._register_component(self.conversation_manager) + + def bind_runtime_context(self, runtime_context: NewContext) -> None: + self._runtime_context = runtime_context + + def require_runtime_context(self) -> NewContext: + if self._runtime_context is None: + raise RuntimeError("LegacyContext 尚未绑定运行时 Context") + return self._runtime_context + + def get_llm_tool_manager(self) -> CompatLLMToolManager: + return self._llm_tools + + def activate_llm_tool(self, name: str) -> bool: + return self._llm_tools.activate_llm_tool(name) + + def deactivate_llm_tool(self, name: str) -> bool: + return self._llm_tools.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._llm_tools.add_func(name, func_args, desc, func_obj) + + def unregister_llm_tool(self, name: str) -> None: + self._llm_tools.remove_func(name) + + def get_config(self) -> dict[str, Any]: + runtime_context = self._runtime_context + if runtime_context is None: + return {} + config = getattr(runtime_context, "_astrbot_config", None) + return dict(config) if isinstance(config, dict) else {} + + def _runtime_config(self) -> Any: + from .api.basic.astrbot_config import AstrBotConfig + + runtime_context = self._runtime_context + config = ( + getattr(runtime_context, "_astrbot_config", None) + if runtime_context + else None + ) + if isinstance(config, AstrBotConfig): + return config + if isinstance(config, dict): + return AstrBotConfig(dict(config)) + return AstrBotConfig({}) + + @staticmethod + def _merge_llm_kwargs( + *, + chat_provider_id: str, + kwargs: dict[str, Any], + ) -> dict[str, Any]: + merged = dict(kwargs) + if chat_provider_id: + merged.setdefault("provider_id", chat_provider_id) + return merged + + @staticmethod + def _apply_request_overrides( + call_kwargs: dict[str, Any], + request: _CompatProviderRequest, + ) -> dict[str, Any]: + updated = dict(call_kwargs) + if request.model: + updated["model"] = request.model + return updated + + @staticmethod + def _component_names(component: Any) -> list[str]: + names = [component.__class__.__name__] + compat_name = getattr(component, "__compat_component_name__", None) + if isinstance(compat_name, str) and compat_name and compat_name not in names: + names.insert(0, compat_name) + return names + + def _register_hook( + self, + name: str, + handler: Callable[..., Any], + *, + priority: int = 0, + ) -> None: + self._compat_hooks[name].append( + _CompatHookEntry(name=name, priority=priority, handler=handler) + ) + self._compat_hooks[name].sort(key=lambda item: item.priority, reverse=True) + + def _register_compat_component(self, component: Any) -> None: + from .api.event.filter import ( + get_compat_hook_metas, + get_compat_llm_tool_meta, + ) + + for _attr_name, attr in _iter_registered_component_methods(component): + tool_meta = get_compat_llm_tool_meta(attr) + if tool_meta is not None: + self._llm_tools.add_tool( + name=tool_meta.name, + description=tool_meta.description, + parameters=_tool_parameters_from_legacy_args(tool_meta.parameters), + handler=attr, + ) + for hook_meta in get_compat_hook_metas(attr): + self._register_hook( + hook_meta.name, + attr, + priority=hook_meta.priority, + ) + + @staticmethod + def _legacy_event(event: Any | None): + if event is None: + return None + from .api.event.astr_message_event import AstrMessageEvent + + if isinstance(event, AstrMessageEvent): + return event + return AstrMessageEvent.from_message_event(event) + + @staticmethod + def _hook_type_injection( + annotation: Any, + available: dict[str, Any], + ) -> Any: + from .api.event.astr_message_event import AstrMessageEvent + from .api.provider.entities import LLMResponse + from .context import Context as RuntimeContext + + if annotation is Any or annotation is inspect.Signature.empty: + return None + if annotation is AstrMessageEvent: + return available.get("event") + if annotation is RuntimeContext or annotation is NewContext: + return available.get("context") + if annotation is LegacyContext: + return available.get("legacy_context") + if annotation is LLMResponse: + return available.get("response") + return None + + async def _call_with_available( + self, + handler: Callable[..., Any], + available: dict[str, Any], + ) -> Any: + signature = inspect.signature(handler) + args: list[Any] = [] + kwargs: dict[str, Any] = {} + for parameter in signature.parameters.values(): + injected = None + if parameter.name in available: + injected = available[parameter.name] + else: + injected = self._hook_type_injection(parameter.annotation, available) + if injected is None: + if parameter.default is not parameter.empty: + continue + continue + if parameter.kind in ( + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + ): + args.append(injected) + elif parameter.kind == inspect.Parameter.KEYWORD_ONLY: + kwargs[parameter.name] = injected + result = handler(*args, **kwargs) + if inspect.isasyncgen(result): + final_value = None + async for item in result: + final_value = item + await self._consume_tool_result( + available.get("event"), + available.get("context"), + item, + ) + return final_value + if inspect.isawaitable(result): + return await result + return result + + async def _run_compat_hook( + self, + name: str, + **available: Any, + ) -> list[Any]: + hook_results: list[Any] = [] + for entry in self._compat_hooks.get(name, []): + hook_results.append( + await self._call_with_available(entry.handler, available) + ) + return hook_results + + async def _consume_tool_result( + self, + event: Any | None, + runtime_context: NewContext | None, + item: Any, + ) -> None: + if event is None: + return + from .api.event.event_result import MessageEventResult + from .api.message.chain import MessageChain + + legacy_event = self._legacy_event(event) + if legacy_event is None: + return + if isinstance(item, MessageEventResult): + if ( + item.chain + and runtime_context is not None + and not item.is_plain_text_only() + ): + await runtime_context.platform.send_chain( + legacy_event.session_ref or legacy_event.session_id, + item.to_payload(), + ) + return + plain_text = item.get_plain_text() + if plain_text: + await legacy_event.reply(plain_text) + return + if isinstance(item, MessageChain): + if ( + item.chain + and runtime_context is not None + and not item.is_plain_text_only() + ): + await runtime_context.platform.send_chain( + legacy_event.session_ref or legacy_event.session_id, + item.to_payload(), + ) + return + plain_text = item.get_plain_text() + if plain_text: + await legacy_event.reply(plain_text) + return + if isinstance(item, str): + await legacy_event.reply(item) + + async def _invoke_llm_tool( + self, + *, + tool_name: str, + tool_args: dict[str, Any], + event: Any | None, + ) -> str: + tool = self._llm_tools.get_func(tool_name) + if tool is None or not tool.active: + return f"tool '{tool_name}' not found" + legacy_event = self._legacy_event(event) + runtime_context = self.require_runtime_context() + await self._run_compat_hook( + "on_using_llm_tool", + event=legacy_event, + context=runtime_context, + legacy_context=self, + tool=tool, + tool_args=tool_args, + ) + tool_result = await self._call_with_available( + tool.handler, + { + **tool_args, + "event": legacy_event, + "context": runtime_context, + "ctx": runtime_context, + "legacy_context": self, + }, + ) + if isinstance(tool_result, str): + normalized = tool_result + elif tool_result is None: + normalized = "" + else: + normalized = str(tool_result) + await self._run_compat_hook( + "on_llm_tool_respond", + event=legacy_event, + context=runtime_context, + legacy_context=self, + tool=tool, + tool_args=tool_args, + tool_result=normalized, + ) + return normalized + + def _register_component(self, *components: Any) -> None: + """保留旧版按名称暴露组件方法的兼容链路。""" + for component in components: + for class_name in self._component_names(component): + self._registered_managers[class_name] = component + for attr_name, attr in _iter_registered_component_methods(component): + self._registered_functions[f"{class_name}.{attr_name}"] = attr + self._register_compat_component(component) + + async def execute_registered_function( + self, + func_full_name: str, + args: dict[str, Any] | None = None, + ) -> Any: + if args is None: + call_args: dict[str, Any] = {} + elif isinstance(args, dict): + call_args = args + else: + raise TypeError("LegacyContext 调用参数必须是 dict") + + func = self._registered_functions.get(func_full_name) + if func is None: + raise ValueError(f"Function not found: {func_full_name}") + + result = func(**call_args) + if inspect.isawaitable(result): + return await result + return result + + async def call_context_function( + self, + func_full_name: str, + args: dict[str, Any] | None = None, + ) -> dict[str, Any]: + return { + "data": await self.execute_registered_function(func_full_name, args), + } + + async def llm_generate( + self, + chat_provider_id: str, + prompt: str | None = None, + image_urls: list[str] | None = None, + tools: Any | None = None, + system_prompt: str | None = None, + contexts: list[dict] | None = None, + event: Any | None = None, + **kwargs: Any, + ) -> LLMResponse: + _warn_once("context.llm_generate()", "ctx.llm.chat_raw(...)") + ctx = self.require_runtime_context() + call_kwargs = self._merge_llm_kwargs( + chat_provider_id=chat_provider_id, + kwargs=kwargs, + ) + legacy_event = self._legacy_event(event) + request = _CompatProviderRequest( + prompt=prompt or "", + session_id=legacy_event.session_id if legacy_event is not None else "", + image_urls=list(image_urls or []), + contexts=list(contexts or []), + system_prompt=system_prompt or "", + model=call_kwargs.get("model"), + ) + await self._run_compat_hook( + "on_waiting_llm_request", + event=legacy_event, + context=ctx, + legacy_context=self, + ) + await self._run_compat_hook( + "on_llm_request", + event=legacy_event, + context=ctx, + legacy_context=self, + request=request, + ) + call_kwargs = self._apply_request_overrides(call_kwargs, request) + response = await ctx.llm.chat_raw( + request.prompt or "", + system=request.system_prompt or None, + history=request.contexts or [], + image_urls=request.image_urls or [], + tools=tools, + **call_kwargs, + ) + legacy_response = _legacy_llm_response(response) + await self._run_compat_hook( + "on_llm_response", + event=legacy_event, + context=ctx, + legacy_context=self, + response=legacy_response, + ) + return legacy_response + + async def tool_loop_agent( + self, + chat_provider_id: str, + prompt: str | None = None, + image_urls: list[str] | None = None, + tools: Any | None = None, + system_prompt: str | None = None, + contexts: list[dict] | None = None, + max_steps: int = 30, + event: Any | None = None, + **kwargs: Any, + ) -> LLMResponse: + from .api.provider.entities import LLMResponse + + _warn_once("context.tool_loop_agent()", "compat local tool loop") + ctx = self.require_runtime_context() + call_kwargs = self._merge_llm_kwargs( + chat_provider_id=chat_provider_id, + kwargs=kwargs, + ) + legacy_event = self._legacy_event(event) + history = list(contexts or []) + request_prompt = prompt or "" + combined_tools = list(self._llm_tools.get_func_desc_openai_style()) + if isinstance(tools, list): + combined_tools.extend(item for item in tools if isinstance(item, dict)) + elif tools is not None: + openai_schema = getattr(tools, "openai_schema", None) + if callable(openai_schema): + extra_tools = openai_schema() + if isinstance(extra_tools, list): + combined_tools.extend( + item for item in extra_tools if isinstance(item, dict) + ) + + final_response = LLMResponse(role="assistant") + for _step in range(max_steps): + request = _CompatProviderRequest( + prompt=request_prompt, + session_id=legacy_event.session_id if legacy_event is not None else "", + image_urls=list(image_urls or []), + contexts=list(history), + system_prompt=system_prompt or "", + model=call_kwargs.get("model"), + ) + await self._run_compat_hook( + "on_waiting_llm_request", + event=legacy_event, + context=ctx, + legacy_context=self, + ) + await self._run_compat_hook( + "on_llm_request", + event=legacy_event, + context=ctx, + legacy_context=self, + request=request, + ) + call_kwargs = self._apply_request_overrides(call_kwargs, request) + response = await ctx.llm.chat_raw( + request.prompt or "", + system=request.system_prompt or None, + history=request.contexts or [], + image_urls=request.image_urls or [], + tools=combined_tools or None, + max_steps=max_steps, + **call_kwargs, + ) + final_response = _legacy_llm_response(response) + await self._run_compat_hook( + "on_llm_response", + event=legacy_event, + context=ctx, + legacy_context=self, + response=final_response, + ) + if not final_response.tools_call_name: + return final_response + + history.append( + { + "role": "assistant", + "content": final_response.completion_text, + "tool_calls": final_response.to_openai_tool_calls(), + } + ) + for tool_name, tool_args, tool_call_id in zip( + final_response.tools_call_name, + final_response.tools_call_args, + final_response.tools_call_ids, + strict=False, + ): + tool_result = await self._invoke_llm_tool( + tool_name=tool_name, + tool_args=tool_args, + event=legacy_event, + ) + history.append( + { + "role": "tool", + "tool_call_id": tool_call_id, + "name": tool_name, + "content": tool_result, + } + ) + request_prompt = "" + + return final_response + + async def send_message(self, session: str, message_chain: Any) -> None: + _warn_once( + "context.send_message()", + "ctx.platform.send(...) / ctx.platform.send_chain(...)", + ) + ctx = self.require_runtime_context() + chain = getattr(message_chain, "chain", None) + to_payload = getattr(message_chain, "to_payload", None) + is_plain_text_only = getattr(message_chain, "is_plain_text_only", None) + if ( + isinstance(chain, list) + and callable(to_payload) + and not (callable(is_plain_text_only) and is_plain_text_only()) + ): + await ctx.platform.send_chain(session, to_payload()) + return + + # 旧版插件也可能传纯文本对象,compat 层保留文本兜底。 + if hasattr(message_chain, "get_plain_text") and callable( + message_chain.get_plain_text + ): + text = message_chain.get_plain_text() + elif hasattr(message_chain, "to_text") and callable(message_chain.to_text): + text = message_chain.to_text() + else: + text = str(message_chain) + await ctx.platform.send(session, text) + + async def add_llm_tools(self, *tools: Any) -> None: + for tool in tools: + name = getattr(tool, "name", None) + if not isinstance(name, str) or not name: + raise TypeError("add_llm_tools() 需要带 name 的工具对象") + handler = getattr(tool, "handler", None) + if not callable(handler): + raise TypeError("add_llm_tools() 需要工具对象提供可调用的 handler") + parameters = getattr(tool, "parameters", None) + if not isinstance(parameters, dict): + func_args = getattr(tool, "func_args", None) + if isinstance(func_args, list): + parameters = _tool_parameters_from_legacy_args(func_args) + else: + parameters = {"type": "object", "properties": {}, "required": []} + description = str(getattr(tool, "description", "") or "") + self._llm_tools.add_tool( + name=name, + description=description, + parameters=parameters, + handler=handler, + ) + + async def put_kv_data(self, key: str, value: Any) -> None: + _warn_once("context.put_kv_data()", "ctx.db.set(key, value)") + ctx = self.require_runtime_context() + await ctx.db.set(key, value) + + async def get_kv_data(self, key: str, default: Any = None) -> Any: + _warn_once("context.get_kv_data()", "ctx.db.get(key)") + ctx = self.require_runtime_context() + value = await ctx.db.get(key) + return default if value is None else value + + async def delete_kv_data(self, key: str) -> None: + _warn_once("context.delete_kv_data()", "ctx.db.delete(key)") + ctx = self.require_runtime_context() + await ctx.db.delete(key) + + async def get_registered_star(self, star_name: str) -> Any: + ctx = self.require_runtime_context() + return await ctx.metadata.get_plugin(star_name) + + async def get_all_stars(self) -> list[Any]: + ctx = self.require_runtime_context() + return await ctx.metadata.list_plugins() diff --git a/src-new/astrbot_sdk/_legacy_loader.py b/src-new/astrbot_sdk/_legacy_loader.py index f737766eb..c6a233792 100644 --- a/src-new/astrbot_sdk/_legacy_loader.py +++ b/src-new/astrbot_sdk/_legacy_loader.py @@ -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( diff --git a/src-new/astrbot_sdk/_legacy_runtime.py b/src-new/astrbot_sdk/_legacy_runtime.py index a3c9f2caf..49ac2c898 100644 --- a/src-new/astrbot_sdk/_legacy_runtime.py +++ b/src-new/astrbot_sdk/_legacy_runtime.py @@ -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 diff --git a/src-new/astrbot_sdk/_legacy_star.py b/src-new/astrbot_sdk/_legacy_star.py new file mode 100644 index 000000000..7a2554be8 --- /dev/null +++ b/src-new/astrbot_sdk/_legacy_star.py @@ -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 diff --git a/src-new/astrbot_sdk/protocol/messages.py b/src-new/astrbot_sdk/protocol/messages.py index 8c226c6c9..45764d5fb 100644 --- a/src-new/astrbot_sdk/protocol/messages.py +++ b/src-new/astrbot_sdk/protocol/messages.py @@ -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) diff --git a/src-new/astrbot_sdk/runtime/bootstrap.py b/src-new/astrbot_sdk/runtime/bootstrap.py index be79fdd9e..7a8706965 100644 --- a/src-new/astrbot_sdk/runtime/bootstrap.py +++ b/src-new/astrbot_sdk/runtime/bootstrap.py @@ -1,1182 +1,49 @@ -"""启动引导模块。 +"""启动引导入口。 -定义 SupervisorRuntime 和 PluginWorkerRuntime 的启动逻辑。 -Supervisor 管理多个 Worker 进程,Worker 运行单个插件。 +对外提供三个顶层启动函数: -架构层次: - AstrBot Core (Python) - | - v - SupervisorRuntime (管理多插件) - | - +-- WorkerSession (插件 A) -- StdioTransport -- PluginWorkerRuntime (子进程) - | - +-- WorkerSession (插件 B) -- StdioTransport -- PluginWorkerRuntime (子进程) - | - +-- WorkerSession (插件 C) -- StdioTransport -- PluginWorkerRuntime (子进程) +- ``run_supervisor``: 启动 Supervisor 进程 +- ``run_plugin_worker``: 启动单插件或组 Worker 进程 +- ``run_websocket_server``: 以 WebSocket 方式启动 Worker -核心类: - SupervisorRuntime: 监管者运行时 - - 发现并加载所有插件 - - 为每个插件启动 Worker 进程 - - 聚合所有 handler 并向 Core 注册 - - 路由 Core 的调用请求到对应 Worker - - 处理 Worker 进程崩溃和重连 - - handler ID 冲突检测和警告 +运行时核心类分布在同目录的子模块: - WorkerSession: Worker 会话 - - 管理单个插件 Worker 进程 - - 通过 Peer 与 Worker 通信 - - 提供 invoke_handler 和 cancel 方法 - - 处理连接关闭回调 - - 自动清理已注册的 handlers - - PluginWorkerRuntime: 插件 Worker 运行时 - - 加载单个插件 - - 通过 Peer 与 Supervisor 通信 - - 分发 handler 调用 - - 处理生命周期回调 (on_start, on_stop) - -与旧版对比: - 旧版 supervisor.py: - - WorkerRuntime 管理单个插件进程 - - SupervisorRuntime 管理所有 Worker - - 使用 JSON-RPC 协议通信 - - call_context_function 调用核心功能 - - 使用 RPCRequestHelper 管理请求 - - 新版 bootstrap.py: - - WorkerSession 封装 Worker 会话 - - SupervisorRuntime 使用 Peer 通信 - - 使用新协议 (initialize/invoke/event/cancel) - - 通过 CapabilityRouter 路由能力调用 - - 支持 Worker 连接关闭回调 - - 支持 handler 冲突检测和警告 - -启动流程: - Supervisor 启动: - 1. discover_plugins() 发现所有插件 - 2. 为每个插件创建 WorkerSession - 3. 调用 session.start() 启动 Worker 进程 - 4. 等待 Worker 初始化完成或连接关闭 - 5. 聚合所有 handler 并向 Core 发送 initialize - 6. 等待 Core 的 initialize_result - - Worker 启动: - 1. load_plugin_spec() 加载插件规范 - 2. load_plugin() 加载插件组件 - 3. 创建 Peer 并设置处理器 - 4. 向 Supervisor 发送 initialize - 5. 等待 Supervisor 的 initialize_result - 6. 执行 on_start 生命周期回调 - -信号处理: - - SIGTERM: 设置 stop_event,触发优雅关闭 - - SIGINT: 设置 stop_event,触发优雅关闭 - -这层负责把 `loader`、`Peer`、`CapabilityRouter` 和 `HandlerDispatcher` 串起来: - -- `SupervisorRuntime`: 启动多个插件 Worker,并把所有 handler 暴露给上游 Core -- `WorkerSession`: Supervisor 侧对单个 Worker 的会话包装 -- `PluginWorkerRuntime`: Worker 进程内的插件加载与 handler 执行 - -当前实现会在 Worker 连接关闭时清理对应 handler,但不会自动重启或重连。 +- ``runtime.supervisor``: ``SupervisorRuntime`` / ``WorkerSession`` +- ``runtime.worker``: ``PluginWorkerRuntime`` / ``GroupWorkerRuntime`` """ from __future__ import annotations import asyncio -import inspect -import json -import os -import signal import sys -from collections.abc import Callable -from dataclasses import dataclass from pathlib import Path -from typing import IO, Any +from typing import IO -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, - resolve_plugin_lifecycle_hook, +from .loader import PluginEnvironmentManager +from .supervisor import ( + SupervisorRuntime, + WorkerSession, + _install_signal_handlers, + _prepare_stdio_transport, + _sdk_source_dir, + _wait_for_shutdown, ) -from ..context import Context as RuntimeContext -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 .handler_dispatcher import CapabilityDispatcher, HandlerDispatcher -from .loader import ( - LoadedPlugin, - PluginEnvironmentManager, - PluginSpec, - discover_plugins, - load_plugin, - load_plugin_spec, -) -from .peer import Peer from .transport import StdioTransport, WebSocketServerTransport - - -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() - - -@dataclass(slots=True) -class GroupPluginRuntimeState: - plugin: PluginSpec - loaded_plugin: LoadedPlugin - lifecycle_context: RuntimeContext - - -def _plugin_name_from_handler_id(handler_id: str) -> str: - if ":" in handler_id: - return handler_id.split(":", 1)[0] - return handler_id - - -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 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) - - -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: - for instance in state.loaded_plugin.instances: - hook = resolve_plugin_lifecycle_hook(instance, method_name) - if hook is None: - continue - args = [] - 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(state.lifecycle_context) - result = hook(*args) - if inspect.isawaitable(result): - await result - - 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: - for instance in self.loaded_plugin.instances: - hook = resolve_plugin_lifecycle_hook(instance, method_name) - if hook is None: - continue - args = [] - 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(self._lifecycle_context) - result = hook(*args) - if inspect.isawaitable(result): - await result - - 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, - ) +from .worker import GroupWorkerRuntime, PluginWorkerRuntime + +__all__ = [ + "GroupWorkerRuntime", + "PluginWorkerRuntime", + "SupervisorRuntime", + "WorkerSession", + "_install_signal_handlers", + "_prepare_stdio_transport", + "_sdk_source_dir", + "_wait_for_shutdown", + "run_supervisor", + "run_plugin_worker", + "run_websocket_server", +] async def run_supervisor( diff --git a/src-new/astrbot_sdk/runtime/capability_router.py b/src-new/astrbot_sdk/runtime/capability_router.py index 921504503..148d6c028 100644 --- a/src-new/astrbot_sdk/runtime/capability_router.py +++ b/src-new/astrbot_sdk/runtime/capability_router.py @@ -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, diff --git a/src-new/astrbot_sdk/runtime/loader.py b/src-new/astrbot_sdk/runtime/loader.py index da4ea4dab..c56eb83f3 100644 --- a/src-new/astrbot_sdk/runtime/loader.py +++ b/src-new/astrbot_sdk/runtime/loader.py @@ -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( diff --git a/src-new/astrbot_sdk/runtime/peer.py b/src-new/astrbot_sdk/runtime/peer.py index f501c420d..3501fbf26 100644 --- a/src-new/astrbot_sdk/runtime/peer.py +++ b/src-new/astrbot_sdk/runtime/peer.py @@ -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, diff --git a/src-new/astrbot_sdk/runtime/supervisor.py b/src-new/astrbot_sdk/runtime/supervisor.py new file mode 100644 index 000000000..f9f82331a --- /dev/null +++ b/src-new/astrbot_sdk/runtime/supervisor.py @@ -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) diff --git a/src-new/astrbot_sdk/runtime/transport.py b/src-new/astrbot_sdk/runtime/transport.py index c401d4c35..e8f7298a1 100644 --- a/src-new/astrbot_sdk/runtime/transport.py +++ b/src-new/astrbot_sdk/runtime/transport.py @@ -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 不可用") diff --git a/src-new/astrbot_sdk/runtime/worker.py b/src-new/astrbot_sdk/runtime/worker.py new file mode 100644 index 000000000..445aa69ca --- /dev/null +++ b/src-new/astrbot_sdk/runtime/worker.py @@ -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, + ) diff --git a/tests_v4/test_legacy_loader.py b/tests_v4/test_legacy_loader.py index 7f61161ab..8fab4f35c 100644 --- a/tests_v4/test_legacy_loader.py +++ b/tests_v4/test_legacy_loader.py @@ -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" diff --git a/tests_v4/test_loader.py b/tests_v4/test_loader.py index a01131bb6..e89231b12 100644 --- a/tests_v4/test_loader.py +++ b/tests_v4/test_loader.py @@ -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, diff --git a/tests_v4/test_peer.py b/tests_v4/test_peer.py index 70053a5f0..24ed5c7fd 100644 --- a/tests_v4/test_peer.py +++ b/tests_v4/test_peer.py @@ -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: diff --git a/tests_v4/test_protocol_messages.py b/tests_v4/test_protocol_messages.py index 4ed6b95fa..3d251b458 100644 --- a/tests_v4/test_protocol_messages.py +++ b/tests_v4/test_protocol_messages.py @@ -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.""" diff --git a/tests_v4/test_transport.py b/tests_v4/test_transport.py index 56dad555a..10880fcf5 100644 --- a/tests_v4/test_transport.py +++ b/tests_v4/test_transport.py @@ -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."""