Revert SDK integration, fix test paths, restore original code

- Delete astrbot-sdk package entirely
- Remove SDK dependencies from pyproject.toml
- Restore core modules to pre-sdk versions
- Fix test monkeypatch paths to use correct module paths
- Fix bool type checking bug in command filter
- Fix tools=None preservation in subagent orchestrator
This commit is contained in:
LIghtJUNction
2026-03-23 01:56:39 +08:00
parent 4db5063b77
commit 05053c221d
222 changed files with 115 additions and 72824 deletions
@@ -1,58 +0,0 @@
## Overview
我正在设计一个新的架构,以实现插件与核心系统的运行时环境隔离,以换取更佳的安全性和兼容性。这个架构将会形成一个 SDK,供插件开发者使用。以下是我目前的设计思路和功能规划:
这个 SDK 主要用于新的插件的 CLI bootstrap、Plugin Runtime 以及开发平台。
## 功能规划
### 对插件端:插件脚手架
1. CLI 指令: 初始化插件模版、指令等组件,作为 bootstraper
```bash
# === Scaffold ===
astr init # 新的插件模版
astr add command # 注册一个指令 / 指令组 handler 类
astr add listener # 注册一个监听器
astr add llmtool # 注册一个 LLM Tool
# 交互式创建,参考 Vue 脚手架,如:
# Is command group: [Y]es / [N]o
# Command Name: calc
# Description: xxxxxx
# === Deployment ===
astr tree # 解析 filters 已注册的 handlers,按类型列出
astr sync # 解析 filters 已注册的 handlers,并刷写到 plugin.yaml / metadata.yaml
astr dev # 启动开发环境(WebSockets 自动连接到 AstrBot Core)
astr build # 打包并构建资产
astr publish # 发布到 GitHub Issue / 插件市场!
```
2. 抽象 - 提供完整的插件开发时要用到的类和类方法的抽象
3. 注册器 - 接受插件注册的所有 Handlers
4. 通信 - 与 AstrBot Core 通信
### 对核心系统端:插件运行时环境
- 通信 - 与插件端的双向通信
- 插件管理 - 封装通信方法(如获取一个事件激活的 star handlers / 调用某个 Star Handler / 禁用某个插件)
## 架构
我们将旧插件命名为 LegacyStar,将新插件命名为 NewStar。LegacyStar 直接运行在 AstrBot Core 进程中,而 NewStar 则运行在一个独立的进程中,通过 IPC 与 AstrBot Core 通信。NewStar 进程将使用 astrbot-sdk 作为其运行时环境。
对于 NewStar 与 Core 之间的通信,我们将使用 stdio 或者 WebSockets 作为 IPC 的通信通道。
我们会设计一个 VirtualPluginLayer,以让 Core 端可以透明地调用 NewStar 的 Handlers,就像调用 LegacyStar 一样。
## 通信过程
通信过程应该是全双工的。
1. Core 调用 `VirtualPluginLayer.initialize`,启动插件进程。
2. 插件进程启动后,Core 调用 `VirtualPluginLayer.handshake()`,进行握手,获取插件的元数据,如支持的 Handlers 列表等。
3. 当消息平台有事件(AstrMessageEvent)下发时,Core 调用 `VirtualPluginLayer.get_triggered_handlers(event)`,获取需要处理该事件的 Handlers 列表。这一步不需要通信,因为 SDK 已经在上一步缓存了插件的 Handlers 列表元数据。
4. 如果有 handler 触发,Core 调用 `VirtualPluginLayer.call_handler(event, xxxx)`,调用某个 Handler 处理事件,等待结果返回。
5. 4 步骤期间,handler 可能会调用一些 Core 中的方法,如发送消息、获取对话历史等,这些调用通过 RPC 方式进行。
6. 插件 handler 处理完事件后,返回结果给 Core,Core 继续后续的事件处理流程。
7. 插件可能会主动向 Core 发送事件通知,Core 接收到后,进行相应的处理。
-35
View File
@@ -1,35 +0,0 @@
name: Code Quality Control
on:
push:
branches: [ "main", "dev" ]
pull_request:
branches: [ "main", "dev" ]
jobs:
lint-and-format:
runs-on: ubuntu-latest
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install tools
run: |
pip install pyclean ruff
- name: 1. Clean python bytecode
run: pyclean .
- name: 2. Ruff format
run: ruff format .
- name: 3. Ruff check and fix
run: ruff check . --fix
continue-on-error: true
env:
PYTHONIOENCODING: utf-8
-60
View File
@@ -1,60 +0,0 @@
# OS files
.DS_Store
# Python bytecode and caches
__pycache__/
*.py[cod]
*.pyd
*.so
.pytest_cache/
pytest-cache-files-*/
.mypy_cache/
.ruff_cache/
.coverage
.coverage.*
htmlcov/
# Build artifacts
build/
dist/
site/
wheels/
*.egg-info/
.eggs/
pip-wheel-metadata/
#
fork-docs/
tmp/
openspec/
scripts/
cs/
test_plugin/astrbot_plugin_interface_coverage
docs/
astrbot_sdk/
!src/astrbot_sdk/
!src/astrbot_sdk/**
src/astrbot_sdk/**/__pycache__/
src/astrbot_sdk/**/*.py[cod]
COMMAND_MATCH_REFACTOR_REPORT.md
# Virtual environments
.venv/
venv/
env/
ENV/
plugins/.venv/
# Tool caches
.uv-cache/
.astrbot/
.codex-local/
# IDE files
.idea/
.vscode/
*.iml
uv.lock
/astrBot/
plugins/
.serena/
-1
View File
@@ -1 +0,0 @@
3.12
-57
View File
@@ -1,57 +0,0 @@
# Notes
## v4 架构约束
### 运行时层
- `Peer` 必须将 transport EOF/连接断开视为一级失败路径。如果 transport 意外关闭而 `Peer` 没有主动失败 `_pending_results` / `_pending_streams`,supervisor 端对 worker 的调用可能永远挂起。
- `Peer.initialize()` 需要在发起端也标记远程已初始化。仅在被动接收 `InitializeMessage` 时设置 `_remote_initialized` 会导致 `wait_until_remote_initialized()` 单边 API 死锁。
- `Peer.invoke_stream()` 默认隐藏 `completed` 事件。需要保留最终结果的调用者必须显式启用 `include_completed=True`。
- `CapabilityRouter.register(..., stream_handler=...)` 使用 `(request_id, payload, cancel_token)` 签名,不是 peer 级别的 `(message, token)`。
### 模块导出约束
- 保持 `astrbot_sdk.runtime` 根导出狭窄。`Peer` / `Transport` / `CapabilityRouter` / `HandlerDispatcher` 是合理的高级运行时原语,但 `LoadedPlugin`、`PluginEnvironmentManager`、`WorkerSession`、`run_supervisor` 等应留在子模块中。
### 测试与 Mock 注意事项
- 当检查 peer 是否完成远程初始化时,避免对可能接收 `MagicMock` peer 的代码使用 `getattr(mock, "remote_peer")` 探测。`MagicMock` 会生成 truthy 子属性,`CapabilityProxy` 应从 `peer.__dict__` 或其他具体存储位置读取显式状态。
- `test_plugin/old/` 和 `test_plugin/new/` 可能包含已生成的 `__pycache__` / `*.pyc`。测试夹具复制示例插件时必须显式忽略这些缓存文件。
### 插件加载注意事项
- 本地 `dev --watch` 或同一路径插件重复加载场景,不能只依赖 `import_string()` 的跨插件模块根冲突清理。热重载前必须按插件目录清理模块缓存。
- `_prepare_plugin_import()` 不能只在插件目录"不在 `sys.path`"时才插入路径。像 `main.py` 这种通用模块名,如果插件目录已在 `sys.path` 但排在后面,`import main` 仍会先命中别处模块;导入前必须把目标插件目录提到 `sys.path[0]`。
- 示例/夹具测试如果直接用裸模块名导入插件入口(例如 `from main import HelloPlugin`),会污染 `sys.modules["main"]`,随后真实 loader 再按 `main:HelloPlugin` 加载时可能串到错误模块。
---
# 开发命令
## 格式化与检查
在提交代码前,请依次运行以下命令:
```bash
ruff format . # 使用 ruff 格式化全局代码
ruff check . --fix # 使用 ruff 检查并自动修复全局格式问题
```
## 测试
如果修改了内容可能影响现有功能,请运行测试以确保没有引入错误:
如果修改了bug或者更改了功能需要添加新的测试
当前仓库已统一使用 `tests/` 目录,`tests_v4/` 不再作为新增测试入口。
仓库当前没有 `run_tests.py`,请直接使用 `pytest`。
```bash
python -m pytest tests -q # 运行 tests 目录全部测试
python -m pytest tests -v # 详细输出
python -m pytest tests -k "test_context_register_task" # 运行匹配模式的测试
python -m pytest tests --cov=astrbot_sdk # 运行测试并生成覆盖率报告
```
## 设计原则
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
-57
View File
@@ -1,57 +0,0 @@
# CLAUDE Notes
## v4 架构约束
### 运行时层
- `Peer` 必须将 transport EOF/连接断开视为一级失败路径。如果 transport 意外关闭而 `Peer` 没有主动失败 `_pending_results` / `_pending_streams`,supervisor 端对 worker 的调用可能永远挂起。
- `Peer.initialize()` 需要在发起端也标记远程已初始化。仅在被动接收 `InitializeMessage` 时设置 `_remote_initialized` 会导致 `wait_until_remote_initialized()` 单边 API 死锁。
- `Peer.invoke_stream()` 默认隐藏 `completed` 事件。需要保留最终结果的调用者必须显式启用 `include_completed=True`。
- `CapabilityRouter.register(..., stream_handler=...)` 使用 `(request_id, payload, cancel_token)` 签名,不是 peer 级别的 `(message, token)`。
### 模块导出约束
- 保持 `astrbot_sdk.runtime` 根导出狭窄。`Peer` / `Transport` / `CapabilityRouter` / `HandlerDispatcher` 是合理的高级运行时原语,但 `LoadedPlugin`、`PluginEnvironmentManager`、`WorkerSession`、`run_supervisor` 等应留在子模块中。
### 测试与 Mock 注意事项
- 当检查 peer 是否完成远程初始化时,避免对可能接收 `MagicMock` peer 的代码使用 `getattr(mock, "remote_peer")` 探测。`MagicMock` 会生成 truthy 子属性,`CapabilityProxy` 应从 `peer.__dict__` 或其他具体存储位置读取显式状态。
- `test_plugin/old/` 和 `test_plugin/new/` 可能包含已生成的 `__pycache__` / `*.pyc`。测试夹具复制示例插件时必须显式忽略这些缓存文件。
### 插件加载注意事项
- 本地 `dev --watch` 或同一路径插件重复加载场景,不能只依赖 `import_string()` 的跨插件模块根冲突清理。热重载前必须按插件目录清理模块缓存。
- `_prepare_plugin_import()` 不能只在插件目录"不在 `sys.path`"时才插入路径。像 `main.py` 这种通用模块名,如果插件目录已在 `sys.path` 但排在后面,`import main` 仍会先命中别处模块;导入前必须把目标插件目录提到 `sys.path[0]`。
- 示例/夹具测试如果直接用裸模块名导入插件入口(例如 `from main import HelloPlugin`),会污染 `sys.modules["main"]`,随后真实 loader 再按 `main:HelloPlugin` 加载时可能串到错误模块。
---
# 开发命令
## 格式化与检查
在提交代码前,请依次运行以下命令:
```bash
ruff format . # 使用 ruff 格式化全局代码
ruff check . --fix # 使用 ruff 检查并自动修复全局格式问题
```
## 测试
如果修改了内容可能影响现有功能,请运行测试以确保没有引入错误:
如果修改了bug或者更改了功能需要添加新的测试
当前仓库已统一使用 `tests/` 目录,`tests_v4/` 不再作为新增测试入口。
仓库当前没有 `run_tests.py`,请直接使用 `pytest`。
```bash
python -m pytest tests -q # 运行 tests 目录全部测试
python -m pytest tests -v # 详细输出
python -m pytest tests -k "test_context_register_task" # 运行匹配模式的测试
python -m pytest tests --cov=astrbot_sdk # 运行测试并生成覆盖率报告
```
## 设计原则
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
-29
View File
@@ -1,29 +0,0 @@
# AstrBot SDK
AstrBot 插件开发 SDK,提供 v4 runtime、worker protocol 和插件工具链。
## 安装
```bash
pip install astrbot-sdk
```
## 开发安装
```bash
# 克隆仓库后
pip install -e .
# 或使用 uv
uv sync
```
## 目录结构
```
astrbot-sdk/
├── src/
│ └── astrbot_sdk/ # SDK 主包
├── pyproject.toml
└── README.md
```
-50
View File
@@ -1,50 +0,0 @@
[build-system]
requires = ["setuptools>=80", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "astrbot-sdk"
version = "0.1.0"
description = "AstrBot SDK with v4 runtime, worker protocol, and plugin tooling"
readme = "README.md"
requires-python = ">=3.12"
dependencies = [
"aiohttp>=3.13.2",
"anthropic>=0.72.1",
"certifi>=2025.10.5",
"click>=8.3.0",
"docstring-parser>=0.17.0",
"google-genai>=1.50.0",
"loguru>=0.7.3",
"msgpack>=1.1.1",
"openai>=2.7.2",
"pydantic>=2.12.3",
"pyyaml>=6.0.3",
"uv>=0.9.17",
]
[project.scripts]
astr = "astrbot_sdk.cli:cli"
astrbot-sdk = "astrbot_sdk.cli:cli"
[tool.pytest.ini_options]
markers = [
"unit: unit tests",
]
# ============================================================
# Package Discovery (src layout)
# ============================================================
[tool.setuptools.packages.find]
where = ["src"]
# ============================================================
# Optional Dependencies
# ============================================================
[project.optional-dependencies]
dev = [
"pytest>=8.0.0",
"pytest-asyncio>=0.24.0",
"pytest-cov>=5.0.0",
"ruff>=0.4.0",
]
-43
View File
@@ -1,43 +0,0 @@
# Notes
## v4 架构约束
### 运行时层
- `Peer` 必须将 transport EOF/连接断开视为一级失败路径。如果 transport 意外关闭而 `Peer` 没有主动失败 `_pending_results` / `_pending_streams`,supervisor 端对 worker 的调用可能永远挂起。
- `Peer.initialize()` 需要在发起端也标记远程已初始化。仅在被动接收 `InitializeMessage` 时设置 `_remote_initialized` 会导致 `wait_until_remote_initialized()` 单边 API 死锁。
- `Peer.invoke_stream()` 默认隐藏 `completed` 事件。需要保留最终结果的调用者必须显式启用 `include_completed=True`。
- `CapabilityRouter.register(..., stream_handler=...)` 使用 `(request_id, payload, cancel_token)` 签名,不是 peer 级别的 `(message, token)`。
### 模块导出约束
- 保持 `astrbot_sdk.runtime` 根导出狭窄。`Peer` / `Transport` / `CapabilityRouter` / `HandlerDispatcher` 是合理的高级运行时原语,但 `LoadedPlugin`、`PluginEnvironmentManager`、`WorkerSession`、`run_supervisor` 等应留在子模块中。
### 测试与 Mock 注意事项
- 当检查 peer 是否完成远程初始化时,避免对可能接收 `MagicMock` peer 的代码使用 `getattr(mock, "remote_peer")` 探测。`MagicMock` 会生成 truthy 子属性,`CapabilityProxy` 应从 `peer.__dict__` 或其他具体存储位置读取显式状态。
- `test_plugin/old/` 和 `test_plugin/new/` 可能包含已生成的 `__pycache__` / `*.pyc`。测试夹具复制示例插件时必须显式忽略这些缓存文件。
### 插件加载注意事项
- 本地 `dev --watch` 或同一路径插件重复加载场景,不能只依赖 `import_string()` 的跨插件模块根冲突清理。热重载前必须按插件目录清理模块缓存。
- `_prepare_plugin_import()` 不能只在插件目录"不在 `sys.path`"时才插入路径。像 `main.py` 这种通用模块名,如果插件目录已在 `sys.path` 但排在后面,`import main` 仍会先命中别处模块;导入前必须把目标插件目录提到 `sys.path[0]`。
- 示例/夹具测试如果直接用裸模块名导入插件入口(例如 `from main import HelloPlugin`),会污染 `sys.modules["main"]`,随后真实 loader 再按 `main:HelloPlugin` 加载时可能串到错误模块。
---
# 开发命令
## 格式化与检查
在提交代码前,请依次运行以下命令:
```bash
ruff format . # 使用 ruff 格式化全局代码
ruff check . --fix # 使用 ruff 检查并自动修复全局格式问题
```
## 设计原则
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践,这是第一原则
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
-200
View File
@@ -1,200 +0,0 @@
"""AstrBot SDK 的顶层公共 API。
这里仅重新导出 v4 推荐直接导入的稳定入口。
新插件应直接使用此模块的导出:
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message
迁移期适配入口位于独立模块;此处只暴露 v4 原生主入口。
"""
from .clients.managers import (
ConversationCreateParams,
ConversationManagerClient,
ConversationRecord,
ConversationUpdateParams,
KnowledgeBaseCreateParams,
KnowledgeBaseDocumentRecord,
KnowledgeBaseDocumentUploadParams,
KnowledgeBaseManagerClient,
KnowledgeBaseRecord,
KnowledgeBaseRetrieveResult,
KnowledgeBaseRetrieveResultItem,
KnowledgeBaseUpdateParams,
MessageHistoryManagerClient,
MessageHistoryPage,
MessageHistoryRecord,
MessageHistorySender,
PersonaCreateParams,
PersonaManagerClient,
PersonaRecord,
PersonaUpdateParams,
)
from .clients.mcp import MCPManagerClient, MCPServerRecord, MCPServerScope, MCPSession
from .clients.metadata import PluginMetadata, StarMetadata
from .clients.platform import PlatformError, PlatformStats, PlatformStatus
from .clients.provider import (
ManagedProviderRecord,
ProviderChangeEvent,
ProviderManagerClient,
)
from .clients.session import SessionPluginManager, SessionServiceManager
from .commands import CommandGroup, command_group, print_cmd_tree
from .context import Context
from .conversation import (
ConversationClosed,
ConversationReplaced,
ConversationSession,
ConversationState,
)
from .decorators import (
acknowledge_global_mcp_risk,
admin_only,
conversation_command,
cooldown,
group_only,
message_types,
on_command,
on_event,
on_message,
on_schedule,
platforms,
priority,
private_only,
provide_capability,
rate_limit,
require_admin,
)
from .errors import AstrBotError
from .events import MessageEvent
from .filters import (
CustomFilter,
MessageTypeFilter,
PlatformFilter,
all_of,
any_of,
custom_filter,
)
from .message.components import (
At,
AtAll,
BaseMessageComponent,
File,
Forward,
Image,
MediaHelper,
Plain,
Poke,
Record,
Reply,
UnknownComponent,
Video,
)
from .message.result import (
EventResultType,
MessageBuilder,
MessageChain,
MessageEventResult,
)
from .message.session import MessageSession
from .plugin_kv import PluginKVStoreMixin
from .schedule import ScheduleContext
from .session_waiter import SessionController, session_waiter
from .star import Star
from .star_tools import StarTools
from .types import GreedyStr
__all__ = [
"AstrBotError",
"At",
"AtAll",
"BaseMessageComponent",
"CommandGroup",
"ConversationClosed",
"ConversationCreateParams",
"ConversationManagerClient",
"ConversationReplaced",
"ConversationRecord",
"ConversationSession",
"ConversationState",
"ConversationUpdateParams",
"Context",
"CustomFilter",
"EventResultType",
"File",
"Forward",
"GreedyStr",
"Image",
"KnowledgeBaseCreateParams",
"KnowledgeBaseDocumentRecord",
"KnowledgeBaseDocumentUploadParams",
"KnowledgeBaseManagerClient",
"KnowledgeBaseRecord",
"KnowledgeBaseRetrieveResult",
"KnowledgeBaseRetrieveResultItem",
"KnowledgeBaseUpdateParams",
"ManagedProviderRecord",
"MCPManagerClient",
"MCPSession",
"MCPServerRecord",
"MCPServerScope",
"MediaHelper",
"MessageHistoryManagerClient",
"MessageHistoryPage",
"MessageHistoryRecord",
"MessageHistorySender",
"MessageEvent",
"MessageEventResult",
"MessageChain",
"MessageBuilder",
"MessageSession",
"MessageTypeFilter",
"Plain",
"PluginKVStoreMixin",
"PluginMetadata",
"PlatformFilter",
"PlatformError",
"PlatformStats",
"PlatformStatus",
"Poke",
"PersonaCreateParams",
"PersonaManagerClient",
"PersonaRecord",
"PersonaUpdateParams",
"ProviderChangeEvent",
"ProviderManagerClient",
"Record",
"Reply",
"ScheduleContext",
"SessionPluginManager",
"SessionServiceManager",
"SessionController",
"Star",
"StarMetadata",
"StarTools",
"UnknownComponent",
"Video",
"acknowledge_global_mcp_risk",
"admin_only",
"all_of",
"any_of",
"cooldown",
"conversation_command",
"command_group",
"custom_filter",
"group_only",
"message_types",
"on_command",
"on_event",
"on_message",
"on_schedule",
"platforms",
"print_cmd_tree",
"priority",
"provide_capability",
"private_only",
"rate_limit",
"require_admin",
"session_waiter",
]
-11
View File
@@ -1,11 +0,0 @@
"""`python -m astrbot_sdk` 的 CLI 入口。"""
from .cli import cli
def main() -> None:
cli()
if __name__ == "__main__":
main()
@@ -1,17 +0,0 @@
from ._internal.command_model import (
COMMAND_MODEL_DOCS_URL,
CommandModelParseResult,
ResolvedCommandModelParam,
format_command_model_help,
parse_command_model_remainder,
resolve_command_model_param,
)
__all__ = [
"COMMAND_MODEL_DOCS_URL",
"CommandModelParseResult",
"ResolvedCommandModelParam",
"format_command_model_help",
"parse_command_model_remainder",
"resolve_command_model_param",
]
@@ -1,7 +0,0 @@
"""Internal implementation modules for astrbot_sdk.
This package groups private helpers that are not part of the public SDK API.
Imports outside the SDK should avoid depending on these modules directly.
"""
__all__: list[str] = []
@@ -1,227 +0,0 @@
from __future__ import annotations
import inspect
from dataclasses import dataclass
from typing import Any
from pydantic import BaseModel
from ..errors import AstrBotError
from ..runtime._command_matching import split_command_remainder
from .injected_params import is_framework_injected_parameter
from .typing_utils import unwrap_optional
# TODO:文档内容喵
COMMAND_MODEL_DOCS_URL = "https://docs.astrbot.org/sdk/parameter-injection"
@dataclass(slots=True)
class ResolvedCommandModelParam:
name: str
model_cls: type[BaseModel]
@dataclass(slots=True)
class CommandModelParseResult:
model: BaseModel | None = None
help_text: str | None = None
def resolve_command_model_param(handler: Any) -> ResolvedCommandModelParam | None:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return None
try:
type_hints = inspect.get_annotations(handler, eval_str=True)
except Exception:
type_hints = {}
candidates: list[ResolvedCommandModelParam] = []
other_names: list[str] = []
for parameter in signature.parameters.values():
if parameter.kind not in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
continue
annotation = type_hints.get(parameter.name)
if _is_injected_parameter(parameter.name, annotation):
continue
normalized, _is_optional = unwrap_optional(annotation)
if isinstance(normalized, type) and issubclass(normalized, BaseModel):
candidates.append(
ResolvedCommandModelParam(
name=parameter.name,
model_cls=normalized,
)
)
continue
other_names.append(parameter.name)
if not candidates:
return None
if len(candidates) > 1 or other_names:
names = [item.name for item in candidates]
raise ValueError(
"Command BaseModel injection requires exactly one non-injected BaseModel "
f"parameter, got models={names!r} others={other_names!r}"
)
_validate_supported_model(candidates[0].model_cls)
return candidates[0]
def parse_command_model_remainder(
*,
remainder: str,
model_param: ResolvedCommandModelParam,
command_name: str,
) -> CommandModelParseResult:
tokens = split_command_remainder(remainder)
if any(token in {"-h", "--help"} for token in tokens):
return CommandModelParseResult(
help_text=format_command_model_help(command_name, model_param.model_cls)
)
fields = model_param.model_cls.model_fields
explicit_values: dict[str, Any] = {}
positional_values: dict[str, Any] = {}
positional_field_names = [
name
for name, field in fields.items()
if _supported_scalar_type(field.annotation)[0] is not bool
]
positional_index = 0
index = 0
while index < len(tokens):
token = tokens[index]
if not token.startswith("--"):
assigned = False
while positional_index < len(positional_field_names):
field_name = positional_field_names[positional_index]
positional_index += 1
if field_name in explicit_values or field_name in positional_values:
continue
positional_values[field_name] = token
assigned = True
break
if not assigned:
raise _command_parse_error("Too many positional arguments")
index += 1
continue
raw_name = token[2:]
if not raw_name:
raise _command_parse_error("Invalid option '--'")
explicit_value: str | None = None
if "=" in raw_name:
raw_name, explicit_value = raw_name.split("=", 1)
negated = raw_name.startswith("no-")
# 与 argparse/click 惯例一致:--foo-bar 自动映射为字段名 foo_bar
field_name = (raw_name[3:] if negated else raw_name).replace("-", "_")
field = fields.get(field_name)
if field is None:
raise _command_parse_error(f"Unknown field: {field_name}")
if field_name in explicit_values:
raise _command_parse_error(f"Duplicate field: {field_name}")
field_type, _is_optional = _supported_scalar_type(field.annotation)
if field_type is bool:
if explicit_value is not None:
raise _command_parse_error(
f"Boolean field '{field_name}' only supports --{field_name} or --no-{field_name}"
)
explicit_values[field_name] = not negated
index += 1
continue
if negated:
raise _command_parse_error(
f"Non-boolean field '{field_name}' does not support --no-{field_name}"
)
if explicit_value is None:
index += 1
if index >= len(tokens):
raise _command_parse_error(f"Missing value for field: {field_name}")
explicit_value = tokens[index]
explicit_values[field_name] = explicit_value
index += 1
values = {**positional_values, **explicit_values}
try:
model = model_param.model_cls.model_validate(values)
except Exception as exc:
raise AstrBotError.invalid_input(
"命令参数解析失败",
hint=str(exc),
docs_url=COMMAND_MODEL_DOCS_URL,
details={
"command": command_name,
"parameter": model_param.name,
"values": values,
},
) from exc
return CommandModelParseResult(model=model)
def format_command_model_help(command_name: str, model_cls: type[BaseModel]) -> str:
_validate_supported_model(model_cls)
lines = [f"用法: /{command_name} [options]"]
if model_cls.model_fields:
lines.append("参数:")
for name, field in model_cls.model_fields.items():
field_type, is_optional = _supported_scalar_type(field.annotation)
type_name = getattr(field_type, "__name__", str(field_type))
required = field.is_required()
default_text = ""
if not required:
default_text = f",默认 {field.default!r}"
elif is_optional:
default_text = ",默认 None"
description = str(field.description or "").strip()
detail = f"{name}: {type_name}"
if description:
detail += f" - {description}"
detail += ",必填" if required else ",可选"
detail += default_text
if field_type is bool:
detail += f",使用 --{name} / --no-{name}"
lines.append(detail)
return "\n".join(lines)
def _validate_supported_model(model_cls: type[BaseModel]) -> None:
for name, field in model_cls.model_fields.items():
try:
_supported_scalar_type(field.annotation)
except TypeError as exc:
raise ValueError(
f"Unsupported command model field '{name}': {exc}"
) from exc
def _supported_scalar_type(annotation: Any) -> tuple[type[Any], bool]:
normalized, is_optional = unwrap_optional(annotation)
if normalized in {str, int, float, bool}:
return normalized, is_optional
raise TypeError("only str/int/float/bool and Optional variants are supported")
def _command_parse_error(message: str) -> AstrBotError:
return AstrBotError.invalid_input(
message,
docs_url=COMMAND_MODEL_DOCS_URL,
)
def _is_injected_parameter(name: str, annotation: Any) -> bool:
return is_framework_injected_parameter(name, annotation)
__all__ = [
"COMMAND_MODEL_DOCS_URL",
"CommandModelParseResult",
"ResolvedCommandModelParam",
"format_command_model_help",
"parse_command_model_remainder",
"resolve_command_model_param",
]
@@ -1,91 +0,0 @@
from __future__ import annotations
import functools
import inspect
from typing import Any
try:
from typing import get_type_hints
except ImportError: # pragma: no cover
get_type_hints = None
from .typing_utils import unwrap_optional
_INJECTED_PARAMETER_NAMES = {
"event",
"ctx",
"context",
"sched",
"schedule",
"conversation",
"conv",
}
def is_framework_injected_parameter(name: str, annotation: Any) -> bool:
if name in _INJECTED_PARAMETER_NAMES:
return True
normalized, _is_optional = unwrap_optional(annotation)
if normalized is None:
return False
try:
injected_types = _framework_injected_types()
except Exception:
return False
if normalized in injected_types:
return True
if isinstance(normalized, type):
return issubclass(normalized, injected_types)
return False
def legacy_arg_parameter_names(handler: Any) -> list[str]:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return []
try:
if get_type_hints is None:
type_hints = {}
else:
type_hints = get_type_hints(handler)
except Exception:
type_hints = {}
names: list[str] = []
for parameter in signature.parameters.values():
if parameter.kind not in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
continue
if is_framework_injected_parameter(
parameter.name, type_hints.get(parameter.name)
):
continue
names.append(parameter.name)
return names
@functools.lru_cache(maxsize=1)
def _framework_injected_types() -> tuple[type[Any], ...]:
from ..clients.llm import LLMResponse
from ..context import Context
from ..conversation import ConversationSession
from ..events import MessageEvent
from ..llm.entities import ProviderRequest
from ..message.result import MessageEventResult
from ..schedule import ScheduleContext
return (
Context,
MessageEvent,
ScheduleContext,
ConversationSession,
ProviderRequest,
LLMResponse,
MessageEventResult,
)
__all__ = ["is_framework_injected_parameter", "legacy_arg_parameter_names"]
@@ -1,86 +0,0 @@
"""插件调用者身份上下文管理。
本模块使用 contextvars 实现跨异步任务传播插件身份,
用于在 capability 调用时自动识别调用者插件。
典型场景:
- http.register_api: 记录哪个插件注册了 API
- metadata.get_plugin_config: 只允许查询当前插件自己的配置
- 能力路由层权限校验
使用方式:
with caller_plugin_scope("my_plugin"):
# 在此作用域内,current_caller_plugin_id() 返回 "my_plugin"
await ctx.http.register_api(...)
注意:
contextvars 会自动传播到子任务(asyncio.create_task),
无需手动传递。
"""
from __future__ import annotations
from collections.abc import Iterator
from contextlib import contextmanager
from contextvars import ContextVar, Token
# 存储当前调用者插件 ID 的上下文变量
_CALLER_PLUGIN_ID: ContextVar[str | None] = ContextVar(
"astrbot_sdk_caller_plugin_id",
default=None,
)
def current_caller_plugin_id() -> str | None:
"""获取当前上下文中的调用者插件 ID。
Returns:
当前插件 ID,如果不在插件调用上下文中则返回 None
"""
return _CALLER_PLUGIN_ID.get()
def bind_caller_plugin_id(plugin_id: str | None) -> Token[str | None]:
"""绑定调用者插件 ID 到当前上下文。
Args:
plugin_id: 插件 ID,空字符串会被视为 None
Returns:
用于后续 reset 的 Token
Note:
通常使用 caller_plugin_scope 上下文管理器而非直接调用此函数
"""
normalized = plugin_id.strip() if isinstance(plugin_id, str) else ""
return _CALLER_PLUGIN_ID.set(normalized or None)
def reset_caller_plugin_id(token: Token[str | None]) -> None:
"""重置调用者插件 ID 到之前的状态。
Args:
token: bind_caller_plugin_id 返回的 Token
"""
_CALLER_PLUGIN_ID.reset(token)
@contextmanager
def caller_plugin_scope(plugin_id: str | None) -> Iterator[None]:
"""创建一个绑定插件身份的上下文作用域。
Args:
plugin_id: 要绑定的插件 ID
Yields:
None
示例:
with caller_plugin_scope("my_plugin"):
await some_capability_call()
"""
token = bind_caller_plugin_id(plugin_id)
try:
yield
finally:
reset_caller_plugin_id(token)
@@ -1,213 +0,0 @@
from __future__ import annotations
import json
import math
import re
from datetime import datetime, timedelta, timezone
from typing import Any
def is_ttl_memory_entry(value: Any) -> bool:
"""Return whether a stored memory payload uses the TTL wrapper shape."""
return isinstance(value, dict) and "value" in value and "ttl_seconds" in value
def memory_value_for_search(stored: Any) -> dict[str, Any] | None:
"""Unwrap the search payload from a stored memory record when possible."""
if not isinstance(stored, dict):
return None
if is_ttl_memory_entry(stored):
value = stored.get("value")
return value if isinstance(value, dict) else None
return stored
def extract_memory_text(stored: Any) -> str:
"""Pick the canonical text that keyword/vector search should index."""
value = memory_value_for_search(stored)
if not isinstance(value, dict):
return ""
for field_name in ("embedding_text", "content", "summary", "title", "text"):
item = value.get(field_name)
if isinstance(item, str) and item.strip():
return item.strip()
return json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)
def memory_expiration_from_ttl(ttl_seconds: Any) -> datetime | None:
"""Translate a TTL in seconds into an absolute UTC expiration timestamp."""
try:
ttl = int(ttl_seconds)
except (TypeError, ValueError):
return None
if ttl < 1:
return None
return datetime.now(timezone.utc) + timedelta(seconds=ttl)
def memory_expiration_from_stored_payload(stored: Any) -> datetime | None:
"""Recover an absolute expiration timestamp from a stored TTL payload."""
if not is_ttl_memory_entry(stored) or not isinstance(stored, dict):
return None
raw_expires_at = stored.get("expires_at")
if isinstance(raw_expires_at, (int, float)):
return datetime.fromtimestamp(float(raw_expires_at), tz=timezone.utc)
if not isinstance(raw_expires_at, str):
return None
normalized = raw_expires_at.strip()
if not normalized:
return None
if normalized.endswith("Z"):
normalized = f"{normalized[:-1]}+00:00"
try:
expires_at = datetime.fromisoformat(normalized)
except ValueError:
return None
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=timezone.utc)
return expires_at.astimezone(timezone.utc)
def normalize_memory_namespace(value: Any) -> str:
"""Normalize a namespace path into a stable slash-delimited string."""
if value is None:
return ""
if isinstance(value, (list, tuple)):
return join_memory_namespace(*value)
text = str(value).strip().replace("\\", "/")
if not text:
return ""
parts = [segment.strip() for segment in text.split("/") if segment.strip()]
return "/".join(parts)
def join_memory_namespace(*parts: Any) -> str:
"""Join namespace segments while preserving the root namespace as empty."""
normalized_parts: list[str] = []
for part in parts:
normalized = normalize_memory_namespace(part)
if not normalized:
continue
normalized_parts.extend(
segment for segment in normalized.split("/") if segment.strip()
)
return "/".join(normalized_parts)
def memory_namespace_matches(
candidate: str,
namespace: str | None,
*,
include_descendants: bool,
) -> bool:
"""Check whether a stored namespace belongs to the requested scope."""
if namespace is None:
return True
normalized_candidate = normalize_memory_namespace(candidate)
normalized_namespace = normalize_memory_namespace(namespace)
if not normalized_namespace:
return include_descendants or normalized_candidate == ""
if normalized_candidate == normalized_namespace:
return True
return include_descendants and normalized_candidate.startswith(
f"{normalized_namespace}/"
)
def display_memory_namespace(value: Any) -> str | None:
"""Return a user-facing namespace value."""
normalized = normalize_memory_namespace(value)
return normalized or None
def _memory_query_terms(value: str) -> list[str]:
normalized = re.sub(r"\s+", " ", str(value).strip().casefold())
if not normalized:
return []
terms = [item for item in re.findall(r"\w+", normalized, flags=re.UNICODE) if item]
if terms:
return terms
compact = normalized.replace(" ", "")
return [compact] if compact else []
def memory_keyword_score(query: str, key: str, text: str) -> float:
"""Score a keyword hit the same way across runtime and core bridge."""
normalized_query = str(query).casefold()
if not normalized_query:
return 1.0
normalized_key = str(key).casefold()
normalized_text = str(text).casefold()
best = 0.0
if normalized_query in normalized_key:
best = 1.0
if normalized_query in normalized_text:
best = max(best, 0.92)
terms = _memory_query_terms(normalized_query)
if not terms:
return best
key_hits = sum(1 for term in terms if term in normalized_key)
text_hits = sum(1 for term in terms if term in normalized_text)
if key_hits:
best = max(best, 0.5 + 0.5 * (key_hits / len(terms)))
if text_hits:
best = max(best, 0.35 + 0.55 * (text_hits / len(terms)))
return min(best, 1.0)
def cosine_similarity(left: list[float], right: list[float]) -> float:
"""Compute cosine similarity defensively for embedding vectors."""
if not left or not right or len(left) != len(right):
return 0.0
left_norm = math.sqrt(sum(value * value for value in left))
right_norm = math.sqrt(sum(value * value for value in right))
if left_norm <= 0 or right_norm <= 0:
return 0.0
return sum(a * b for a, b in zip(left, right, strict=False)) / (
left_norm * right_norm
)
def normalize_embedding(vector: list[float]) -> list[float]:
"""Normalize an embedding for cosine/inner-product search."""
if not vector:
return []
norm = math.sqrt(sum(value * value for value in vector))
if norm <= 0:
return [0.0 for _ in vector]
return [float(value) / norm for value in vector]
def memory_index_entry(entry: Any, *, text: str) -> dict[str, Any]:
"""Normalize cached sidecar data into a stable memory index record."""
if isinstance(entry, dict):
return {
"text": str(entry.get("text", text)),
"embedding": (
[float(item) for item in entry.get("embedding", [])]
if isinstance(entry.get("embedding"), list)
else None
),
"provider_id": (
str(entry.get("provider_id")).strip()
if entry.get("provider_id") is not None
else None
),
}
return {"text": text, "embedding": None, "provider_id": None}
@@ -1,54 +0,0 @@
from __future__ import annotations
import re
from pathlib import Path
PLUGIN_ID_PATTERN = re.compile(r"^[A-Za-z0-9_](?:[A-Za-z0-9._-]{0,126}[A-Za-z0-9_])?$")
_WINDOWS_RESERVED_PLUGIN_IDS = {
"CON",
"PRN",
"AUX",
"NUL",
"COM1",
"COM2",
"COM3",
"COM4",
"COM5",
"COM6",
"COM7",
"COM8",
"COM9",
"LPT1",
"LPT2",
"LPT3",
"LPT4",
"LPT5",
"LPT6",
"LPT7",
"LPT8",
"LPT9",
}
def validate_plugin_id(plugin_id: str) -> str:
normalized = str(plugin_id).strip()
if not normalized:
raise ValueError("plugin_id must not be empty")
if not PLUGIN_ID_PATTERN.fullmatch(normalized):
raise ValueError(
"plugin_id must use only letters, digits, dots, underscores, or hyphens"
)
if normalized.upper() in _WINDOWS_RESERVED_PLUGIN_IDS:
raise ValueError("plugin_id must not use a reserved Windows device name")
return normalized
def resolve_plugin_data_dir(root: Path, plugin_id: str) -> Path:
normalized = validate_plugin_id(plugin_id)
resolved_root = root.resolve()
candidate = (resolved_root / normalized).resolve()
try:
candidate.relative_to(resolved_root)
except ValueError as exc:
raise ValueError("plugin_id escapes the plugin data root") from exc
return candidate
@@ -1,313 +0,0 @@
from __future__ import annotations
import asyncio
import inspect
import os
import time
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
try:
from astrbot.core.config.default import VERSION as _ASTRBOT_VERSION
except Exception: # noqa: BLE001
_ASTRBOT_VERSION = ""
__all__ = ["PluginLogEntry", "PluginLogger"]
@dataclass(slots=True)
class PluginLogEntry:
level: str
time: float
message: str
plugin_id: str
context: dict[str, Any] = field(default_factory=dict)
class _PluginLogBroker:
def __init__(self, plugin_id: str) -> None:
self.plugin_id = plugin_id
self._subscribers: set[asyncio.Queue[PluginLogEntry]] = set()
def publish(self, entry: PluginLogEntry) -> None:
for queue in list(self._subscribers):
try:
queue.put_nowait(entry)
except asyncio.QueueFull:
continue
async def watch(self) -> AsyncIterator[PluginLogEntry]:
queue: asyncio.Queue[PluginLogEntry] = asyncio.Queue()
self._subscribers.add(queue)
try:
while True:
yield await queue.get()
finally:
self._subscribers.discard(queue)
_BROKERS: dict[str, _PluginLogBroker] = {}
_SHORT_LEVEL_NAMES = {
"DEBUG": "DBUG",
"INFO": "INFO",
"WARNING": "WARN",
"ERROR": "ERRO",
"CRITICAL": "CRIT",
}
_ANSI_RESET = "\u001b[0m"
_ANSI_GREEN = "\u001b[32m"
_ANSI_LEVEL_COLORS = {
"DEBUG": "\u001b[1;34m",
"INFO": "\u001b[1;36m",
"WARNING": "\u001b[1;33m",
"ERROR": "\u001b[31m",
"CRITICAL": "\u001b[1;31m",
}
def _get_short_level_name(level_name: str) -> str:
return _SHORT_LEVEL_NAMES.get(level_name.upper(), level_name[:4].upper())
def _build_source_file(pathname: str | None) -> str:
if not pathname:
return "unknown"
dirname = os.path.dirname(pathname)
return (
os.path.basename(dirname) + "." + os.path.basename(pathname).replace(".py", "")
)
def _plugin_tag_from_path(pathname: str | None) -> str:
if not pathname:
return "[Plug]"
norm_path = os.path.normpath(pathname)
if any(
marker in norm_path
for marker in (
os.path.normpath("data/plugins"),
os.path.normpath("data/sdk_plugins"),
os.path.normpath("astrbot/builtin_stars"),
)
):
return "[Plug]"
return "[Core]"
def _level_color(level: str) -> str:
return _ANSI_LEVEL_COLORS.get(level.upper(), _ANSI_RESET)
def _get_broker(plugin_id: str) -> _PluginLogBroker:
broker = _BROKERS.get(plugin_id)
if broker is None:
broker = _PluginLogBroker(plugin_id)
_BROKERS[plugin_id] = broker
return broker
class PluginLogger:
def __init__(
self,
*,
plugin_id: str,
logger: Any,
bound_context: dict[str, Any] | None = None,
) -> None:
self._plugin_id = plugin_id
self._logger = logger
self._broker = _get_broker(plugin_id)
self._bound_context = dict(bound_context or {})
@property
def plugin_id(self) -> str:
return self._plugin_id
def bind(self, **kwargs: Any) -> PluginLogger:
bind = getattr(self._logger, "bind", None)
next_logger = self._logger
if callable(bind):
try:
next_logger = bind(**kwargs)
except Exception:
next_logger = self._logger
return PluginLogger(
plugin_id=self._plugin_id,
logger=next_logger,
bound_context={**self._bound_context, **kwargs},
)
def opt(self, *args: Any, **kwargs: Any) -> PluginLogger:
opt = getattr(self._logger, "opt", None)
next_logger = self._logger
if callable(opt):
try:
next_logger = opt(*args, **kwargs)
except Exception:
next_logger = self._logger
return PluginLogger(
plugin_id=self._plugin_id,
logger=next_logger,
bound_context=self._bound_context,
)
async def watch(self) -> AsyncIterator[PluginLogEntry]:
async for entry in self._broker.watch():
yield entry
def log(self, level: str, message: Any, *args: Any, **kwargs: Any) -> None:
normalized_level = str(level).upper()
self._emit_console(normalized_level, message, *args, **kwargs)
self._publish(normalized_level, message, *args, **kwargs)
def debug(self, message: Any, *args: Any, **kwargs: Any) -> None:
self._emit_console("DEBUG", message, *args, **kwargs)
self._publish("DEBUG", message, *args, **kwargs)
def info(self, message: Any, *args: Any, **kwargs: Any) -> None:
self._emit_console("INFO", message, *args, **kwargs)
self._publish("INFO", message, *args, **kwargs)
def warning(self, message: Any, *args: Any, **kwargs: Any) -> None:
self._emit_console("WARNING", message, *args, **kwargs)
self._publish("WARNING", message, *args, **kwargs)
def error(self, message: Any, *args: Any, **kwargs: Any) -> None:
self._emit_console("ERROR", message, *args, **kwargs)
self._publish("ERROR", message, *args, **kwargs)
def exception(self, message: Any, *args: Any, **kwargs: Any) -> None:
self._emit_console("ERROR", message, *args, exception=True, **kwargs)
self._publish("ERROR", message, *args, **kwargs)
def _emit_console(
self,
level: str,
message: Any,
*args: Any,
exception: bool = False,
**kwargs: Any,
) -> None:
if self._emit_console_with_opt(
level,
message,
*args,
exception=exception,
**kwargs,
):
return
self._emit_console_fallback(
level,
message,
*args,
exception=exception,
**kwargs,
)
def _emit_console_with_opt(
self,
level: str,
message: Any,
*args: Any,
exception: bool = False,
**kwargs: Any,
) -> bool:
opt = getattr(self._logger, "opt", None)
if not callable(opt):
return False
formatted_message = self._format_message(message, *args, **kwargs)
pathname, source_line = self._caller_info()
plugin_tag = _plugin_tag_from_path(pathname)
source_file = _build_source_file(pathname)
version_tag = (
f" [v{_ASTRBOT_VERSION}]"
if _ASTRBOT_VERSION and level in {"WARNING", "ERROR", "CRITICAL"}
else ""
)
timestamp = datetime.now().strftime("%H:%M:%S.%f")[:-3]
level_text = _get_short_level_name(level)
level_color = _level_color(level)
line = (
f"{_ANSI_GREEN}[{timestamp}]{_ANSI_RESET} {plugin_tag} "
f"{level_color}[{level_text}]{_ANSI_RESET}{version_tag} "
f"[{source_file}:{source_line}]: {level_color}{formatted_message}{_ANSI_RESET}"
)
try:
emitter = opt(raw=True, exception=True) if exception else opt(raw=True)
log = getattr(emitter, "log", None)
if not callable(log):
return False
log(level, line + "\n")
return True
except Exception:
return False
def _emit_console_fallback(
self,
level: str,
message: Any,
*args: Any,
exception: bool = False,
**kwargs: Any,
) -> None:
method_names = []
if exception:
method_names.append("exception")
method_names.append(str(level).lower())
if exception:
method_names.append("error")
for method_name in method_names:
method = getattr(self._logger, method_name, None)
if not callable(method):
continue
try:
method(message, *args, **kwargs)
except Exception:
continue
return
log = getattr(self._logger, "log", None)
if callable(log):
try:
log(level, self._format_message(message, *args, **kwargs))
except Exception:
return
def _caller_info(self) -> tuple[str | None, int]:
frame = inspect.currentframe()
if frame is None:
return None, 0
frame = frame.f_back
while frame is not None and frame.f_globals.get("__name__") == __name__:
frame = frame.f_back
if frame is None:
return None, 0
return str(frame.f_code.co_filename), int(frame.f_lineno)
def _publish(self, level: str, message: Any, *args: Any, **kwargs: Any) -> None:
entry = PluginLogEntry(
level=level,
time=time.time(),
message=self._format_message(message, *args, **kwargs),
plugin_id=self._plugin_id,
context=dict(self._bound_context),
)
self._broker.publish(entry)
@staticmethod
def _format_message(message: Any, *args: Any, **kwargs: Any) -> str:
if not isinstance(message, str):
return str(message)
text = message
if not args and not kwargs:
return text
try:
return text.format(*args, **kwargs)
except Exception:
return text
def __getattr__(self, name: str) -> Any:
return getattr(self._logger, name)
@@ -1,46 +0,0 @@
from __future__ import annotations
from collections.abc import Iterator
from contextlib import contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from ..context import Context
from ..star import Star
_CURRENT_STAR_CONTEXT: ContextVar[Context | None] = ContextVar(
"astrbot_sdk_current_star_context",
default=None,
)
_CURRENT_STAR_INSTANCE: ContextVar[Star | None] = ContextVar(
"astrbot_sdk_current_star_instance",
default=None,
)
def current_star_context() -> Context | None:
return _CURRENT_STAR_CONTEXT.get()
def current_runtime_context() -> Context | None:
return _CURRENT_STAR_CONTEXT.get()
def current_star_instance() -> Star | None:
return _CURRENT_STAR_INSTANCE.get()
@contextmanager
def bind_star_runtime(star: Star | None, ctx: Context | None) -> Iterator[None]:
context_token = _CURRENT_STAR_CONTEXT.set(ctx)
star_token = _CURRENT_STAR_INSTANCE.set(star)
instance_token = star._bind_runtime_context(ctx) if star is not None else None
try:
yield
finally:
if star is not None and instance_token is not None:
star._reset_runtime_context(instance_token)
_CURRENT_STAR_INSTANCE.reset(star_token)
_CURRENT_STAR_CONTEXT.reset(context_token)
@@ -1,606 +0,0 @@
"""Shared support primitives for local SDK testing."""
from __future__ import annotations
import asyncio
import typing
from collections.abc import Mapping
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, TextIO
from ..context import CancelToken
from ..context import Context as RuntimeContext
from ..events import MessageEvent
from ..protocol.messages import EventMessage, PeerInfo
from ..runtime._streaming import StreamExecution
from ..runtime.capability_router import CapabilityRouter
def _clone_payload_mapping(value: Any) -> dict[str, Any] | None:
if not isinstance(value, dict):
return None
return {str(key): item for key, item in value.items()}
@dataclass(slots=True)
class RecordedSend:
kind: str
message_id: str
session_id: str
text: str | None = None
image_url: str | None = None
chain: list[dict[str, Any]] | None = None
target: dict[str, Any] | None = None
raw: dict[str, Any] = field(default_factory=dict)
@property
def session(self) -> str:
return self.session_id
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> RecordedSend:
if "text" in payload:
kind = "text"
elif "image_url" in payload:
kind = "image"
elif "chain" in payload:
kind = "chain"
else:
kind = "unknown"
return cls(
kind=kind,
message_id=str(payload.get("message_id", "")),
session_id=str(payload.get("session", "")),
text=payload.get("text") if isinstance(payload.get("text"), str) else None,
image_url=(
payload.get("image_url")
if isinstance(payload.get("image_url"), str)
else None
),
chain=(
[dict(item) for item in payload.get("chain", [])]
if isinstance(payload.get("chain"), list)
else None
),
target=_clone_payload_mapping(payload.get("target")),
raw=dict(payload),
)
class StdoutPlatformSink:
def __init__(self, stream: TextIO | None = None) -> None:
self._stream = stream
self.records: list[RecordedSend] = []
def record(self, item: RecordedSend) -> None:
self.records.append(item)
if self._stream is None:
return
self._stream.write(self._format(item) + "\n")
self._stream.flush()
def clear(self) -> None:
self.records.clear()
def _format(self, item: RecordedSend) -> str:
if item.kind == "text":
return f"[text][{item.session_id}] {item.text or ''}"
if item.kind == "image":
return f"[image][{item.session_id}] {item.image_url or ''}"
if item.kind == "chain":
count = len(item.chain or [])
return f"[chain][{item.session_id}] {count} components"
return f"[send][{item.session_id}] {item.raw}"
class InMemoryDB:
def __init__(self, store: dict[str, Any]) -> None:
self._store = store
def get(self, key: str, default: Any = None) -> Any:
return self._store.get(key, default)
def set(self, key: str, value: Any) -> None:
self._store[key] = value
def delete(self, key: str) -> None:
self._store.pop(key, None)
def list(self, prefix: str | None = None) -> list[str]:
keys = sorted(self._store.keys())
if prefix is None:
return keys
return [key for key in keys if key.startswith(prefix)]
def get_many(self, keys: list[str]) -> list[dict[str, Any]]:
return [{"key": key, "value": self._store.get(key)} for key in keys]
def set_many(self, items: list[dict[str, Any]]) -> None:
for item in items:
self.set(str(item.get("key", "")), item.get("value"))
class InMemoryMemory:
def __init__(
self,
store: dict[str, dict[str, Any]],
*,
expires_at: dict[str, datetime | None] | None = None,
) -> None:
self._store = store
self._expires_at = expires_at if expires_at is not None else {}
@staticmethod
def _is_ttl_entry(value: Any) -> bool:
"""判断测试 memory 值是否使用 TTL 包装结构。
Args:
value: 待检查的存储值。
Returns:
bool: 如果包含 ``value`` 和 ``ttl_seconds`` 字段则返回 ``True``。
"""
return isinstance(value, dict) and "value" in value and "ttl_seconds" in value
@classmethod
def _search_text(cls, value: Any) -> str:
"""提取测试用 memory.search 的匹配文本。
Args:
value: 当前存储的 memory 值。
Returns:
str: 用于本地测试搜索的文本内容。
"""
if cls._is_ttl_entry(value):
value = value.get("value")
if not isinstance(value, dict):
return ""
for field_name in ("embedding_text", "content", "summary", "title", "text"):
item = value.get(field_name)
if isinstance(item, str) and item.strip():
return item.strip()
return str(value)
def _is_expired(self, key: str) -> bool:
"""判断测试 memory 键是否已经过期。
Args:
key: memory 条目的键。
Returns:
bool: 如果当前时间已超过过期时间则返回 ``True``。
"""
expires_at = self._expires_at.get(key)
return expires_at is not None and expires_at <= datetime.now(timezone.utc)
def _purge_if_expired(self, key: str) -> bool:
"""在测试 helper 中清理已过期的 memory 条目。
Args:
key: memory 条目的键。
Returns:
bool: 如果条目已过期并被清理则返回 ``True``。
"""
if not self._is_expired(key):
return False
self._store.pop(key, None)
self._expires_at.pop(key, None)
return True
def get(self, key: str, default: Any = None) -> Any:
if self._purge_if_expired(key):
return default
return self._store.get(key, default)
def save(self, key: str, value: dict[str, Any]) -> None:
self._store[key] = dict(value)
def delete(self, key: str) -> None:
self._store.pop(key, None)
self._expires_at.pop(key, None)
def search(self, query: str) -> list[dict[str, Any]]:
results: list[dict[str, Any]] = []
for key, value in list(self._store.items()):
if self._purge_if_expired(key):
continue
if query in key or query in self._search_text(value):
results.append({"key": key, "value": value})
return results
class MockLLMClient:
def __init__(self, client: Any, router: MockCapabilityRouter) -> None:
self._client = client
self._router = router
def mock_response(self, text: str) -> None:
self._router.enqueue_llm_response(text)
def mock_stream_response(self, text: str) -> None:
self._router.enqueue_llm_stream_response(text)
def clear_mock_responses(self) -> None:
self._router.clear_llm_responses()
def __getattr__(self, name: str) -> Any:
return getattr(self._client, name)
class MockPlatformClient:
def __init__(self, client: Any, sink: StdoutPlatformSink) -> None:
self._client = client
self._sink = sink
@property
def records(self) -> list[RecordedSend]:
return list(self._sink.records)
def assert_sent(
self,
expected_text: str | None = None,
*,
kind: str = "text",
count: int | None = None,
) -> None:
matched = [item for item in self._sink.records if item.kind == kind]
if expected_text is not None:
matched = [item for item in matched if item.text == expected_text]
if count is not None:
if len(matched) != count:
raise AssertionError(
f"expected {count} sent records, got {len(matched)}: {matched}"
)
return
if not matched:
raise AssertionError(
f"expected sent record kind={kind!r} text={expected_text!r}, got {self._sink.records}"
)
def __getattr__(self, name: str) -> Any:
return getattr(self._client, name)
class MockCapabilityRouter(CapabilityRouter):
def __init__(self, *, platform_sink: StdoutPlatformSink | None = None) -> None:
self.platform_sink = platform_sink or StdoutPlatformSink()
self._llm_responses: list[str] = []
self._llm_stream_responses: list[str] = []
super().__init__()
self.db = InMemoryDB(self.db_store)
self.memory = InMemoryMemory(
self.memory_store,
expires_at=self._memory_expires_at,
)
def list_dynamic_command_routes(self, plugin_id: str) -> list[dict[str, Any]]:
return super().list_dynamic_command_routes(plugin_id)
def remove_dynamic_command_routes_for_plugin(self, plugin_id: str) -> None:
super().remove_dynamic_command_routes_for_plugin(plugin_id)
def emit_provider_change(
self,
provider_id: str,
provider_type: str,
umo: str | None = None,
) -> None:
super().emit_provider_change(provider_id, provider_type, umo)
def record_platform_error(
self,
platform_id: str,
message: str,
*,
traceback: str | None = None,
) -> None:
super().record_platform_error(platform_id, message, traceback=traceback)
def set_platform_stats(self, platform_id: str, stats: dict[str, Any]) -> None:
super().set_platform_stats(platform_id, stats)
def enqueue_llm_response(self, text: str) -> None:
self._llm_responses.append(text)
def enqueue_llm_stream_response(self, text: str) -> None:
self._llm_stream_responses.append(text)
def clear_llm_responses(self) -> None:
self._llm_responses.clear()
self._llm_stream_responses.clear()
async def execute(
self,
capability: str,
payload: dict[str, Any],
*,
stream: bool,
cancel_token,
request_id: str,
) -> dict[str, Any] | StreamExecution:
if capability == "llm.chat":
return {"text": self._take_llm_response(str(payload.get("prompt", "")))}
if capability == "llm.chat_raw":
text = self._take_llm_response(str(payload.get("prompt", "")))
return {
"text": text,
"usage": {
"input_tokens": len(str(payload.get("prompt", ""))),
"output_tokens": len(text),
},
"finish_reason": "stop",
"tool_calls": [],
"role": "assistant",
"reasoning_content": None,
"reasoning_signature": None,
}
if capability == "llm.stream_chat":
text = self._take_llm_stream_response(str(payload.get("prompt", "")))
async def iterator() -> typing.AsyncIterator[dict[str, Any]]:
for char in text:
cancel_token.raise_if_cancelled()
await asyncio.sleep(0)
yield {"text": char}
return StreamExecution(
iterator=iterator(),
finalize=lambda chunks: {
"text": "".join(item.get("text", "") for item in chunks)
},
)
before = len(self.sent_messages)
result = await super().execute(
capability,
payload,
stream=stream,
cancel_token=cancel_token,
request_id=request_id,
)
self._flush_platform_records(before)
return result
def _flush_platform_records(self, start_index: int) -> None:
for payload in self.sent_messages[start_index:]:
self.platform_sink.record(RecordedSend.from_payload(payload))
def _take_llm_response(self, prompt: str) -> str:
if self._llm_responses:
return self._llm_responses.pop(0)
return f"Echo: {prompt}"
def _take_llm_stream_response(self, prompt: str) -> str:
if self._llm_stream_responses:
return self._llm_stream_responses.pop(0)
if self._llm_responses:
return self._llm_responses.pop(0)
return f"Echo: {prompt}"
class MockPeer:
def __init__(self, router: MockCapabilityRouter) -> None:
self._router = router
self._counter = 0
self.remote_peer = PeerInfo(
name="astrbot-local-core",
role="core",
version="local",
)
self.remote_capabilities = list(router.descriptors())
self.remote_capability_map = {
item.name: item for item in self.remote_capabilities
}
self.remote_handlers: list[Any] = []
self.remote_provided_capabilities: list[Any] = []
self.remote_metadata = {"mode": "local"}
async def invoke(
self,
capability: str,
payload: dict[str, Any],
*,
stream: bool = False,
request_id: str | None = None,
) -> dict[str, Any]:
if stream:
raise ValueError("stream=True 请使用 invoke_stream()")
return typing.cast(
dict[str, Any],
await self._router.execute(
capability,
payload,
stream=False,
cancel_token=CancelToken(),
request_id=request_id or self._next_id(),
),
)
async def invoke_stream(
self,
capability: str,
payload: dict[str, Any],
*,
request_id: str | None = None,
include_completed: bool = False,
):
request_id = request_id or self._next_id()
execution = typing.cast(
StreamExecution,
await self._router.execute(
capability,
payload,
stream=True,
cancel_token=CancelToken(),
request_id=request_id,
),
)
async def iterator():
yield EventMessage.model_validate({"id": request_id, "phase": "started"})
chunks: list[dict[str, Any]] = []
async for chunk in execution.iterator:
if execution.collect_chunks:
chunks.append(chunk)
yield EventMessage.model_validate(
{"id": request_id, "phase": "delta", "data": chunk}
)
output = execution.finalize(chunks)
if include_completed:
yield EventMessage.model_validate(
{"id": request_id, "phase": "completed", "output": output}
)
return iterator()
def _next_id(self) -> str:
self._counter += 1
return f"local_{self._counter:04d}"
def _normalize_plugin_metadata(
plugin_id: str,
plugin_metadata: Mapping[str, Any] | None,
) -> dict[str, Any]:
if plugin_metadata is None:
plugin_metadata = {}
declared_name = plugin_metadata.get("name")
if declared_name is not None and str(declared_name) != plugin_id:
raise ValueError(
"MockContext.plugin_metadata['name'] 必须与 plugin_id 一致,"
f"当前收到 {declared_name!r} != {plugin_id!r}"
)
description = plugin_metadata.get("description")
if description is None:
description = plugin_metadata.get("desc", "")
return {
"name": plugin_id,
"display_name": str(plugin_metadata.get("display_name") or plugin_id),
"description": str(description or ""),
"author": str(plugin_metadata.get("author") or ""),
"version": str(plugin_metadata.get("version") or "0.0.0"),
"enabled": bool(plugin_metadata.get("enabled", True)),
"reserved": bool(plugin_metadata.get("reserved", False)),
"acknowledge_global_mcp_risk": bool(
plugin_metadata.get("acknowledge_global_mcp_risk", False)
),
"local_mcp_servers": (
{
str(server_name): dict(server_payload)
for server_name, server_payload in plugin_metadata.get(
"local_mcp_servers",
{},
).items()
if str(server_name).strip() and isinstance(server_payload, dict)
}
if isinstance(plugin_metadata.get("local_mcp_servers"), dict)
else {}
),
"support_platforms": [
str(item)
for item in plugin_metadata.get("support_platforms", [])
if isinstance(item, str)
]
if isinstance(plugin_metadata.get("support_platforms"), list)
else [],
"astrbot_version": (
str(plugin_metadata.get("astrbot_version"))
if plugin_metadata.get("astrbot_version") is not None
else None
),
}
class MockContext(RuntimeContext):
def __init__(
self,
*,
plugin_id: str = "test-plugin",
logger: Any | None = None,
cancel_token: CancelToken | None = None,
platform_sink: StdoutPlatformSink | None = None,
plugin_metadata: Mapping[str, Any] | None = None,
) -> None:
self.platform_sink = platform_sink or StdoutPlatformSink()
self.router = MockCapabilityRouter(platform_sink=self.platform_sink)
self.mock_peer = MockPeer(self.router)
super().__init__(
peer=self.mock_peer,
plugin_id=plugin_id,
cancel_token=cancel_token,
logger=logger,
)
self.router.upsert_plugin(
metadata=_normalize_plugin_metadata(plugin_id, plugin_metadata),
config={},
)
self.llm = MockLLMClient(self.llm, self.router)
self.platform = MockPlatformClient(self.platform, self.platform_sink)
@property
def sent_messages(self) -> list[RecordedSend]:
return list(self.platform_sink.records)
@property
def event_actions(self) -> list[dict[str, Any]]:
return list(self.router.event_actions)
class MockMessageEvent(MessageEvent):
def __init__(
self,
*,
text: str = "",
user_id: str | None = "test-user",
group_id: str | None = None,
platform: str | None = "test",
session_id: str | None = "test-session",
raw: dict[str, Any] | None = None,
context: MockContext | None = None,
) -> None:
self.replies: list[str] = []
super().__init__(
text=text,
user_id=user_id,
group_id=group_id,
platform=platform,
session_id=session_id,
raw=raw,
context=context,
)
if context is not None:
self.bind_runtime_reply(context)
elif self._reply_handler is None:
self.bind_reply_handler(self._capture_reply)
@property
def is_private(self) -> bool:
return self.group_id is None
def bind_runtime_reply(self, context: MockContext) -> None:
self._context = context
async def reply(text: str) -> None:
self.replies.append(text)
await context.platform.send(self.session_ref or self.session_id, text)
self.bind_reply_handler(reply)
async def _capture_reply(self, text: str) -> None:
self.replies.append(text)
__all__ = [
"InMemoryDB",
"InMemoryMemory",
"MockCapabilityRouter",
"MockContext",
"MockLLMClient",
"MockMessageEvent",
"MockPeer",
"MockPlatformClient",
"RecordedSend",
"StdoutPlatformSink",
]
@@ -1,17 +0,0 @@
from __future__ import annotations
import typing
from types import UnionType
from typing import Any
def unwrap_optional(annotation: Any) -> tuple[Any, bool]:
origin = typing.get_origin(annotation)
if origin in {typing.Union, UnionType}:
args = [item for item in typing.get_args(annotation) if item is not type(None)]
if len(args) == 1:
return args[0], True
return annotation, False
__all__ = ["unwrap_optional"]
File diff suppressed because it is too large Load Diff
@@ -1,3 +0,0 @@
from ._internal.plugin_logger import PluginLogEntry, PluginLogger
__all__ = ["PluginLogEntry", "PluginLogger"]
@@ -1,13 +0,0 @@
from ._internal.star_runtime import (
bind_star_runtime,
current_runtime_context,
current_star_context,
current_star_instance,
)
__all__ = [
"bind_star_runtime",
"current_runtime_context",
"current_star_context",
"current_star_instance",
]
@@ -1,25 +0,0 @@
from ._internal.testing_support import (
InMemoryDB,
InMemoryMemory,
MockCapabilityRouter,
MockContext,
MockLLMClient,
MockMessageEvent,
MockPeer,
MockPlatformClient,
RecordedSend,
StdoutPlatformSink,
)
__all__ = [
"InMemoryDB",
"InMemoryMemory",
"MockCapabilityRouter",
"MockContext",
"MockLLMClient",
"MockMessageEvent",
"MockPeer",
"MockPlatformClient",
"RecordedSend",
"StdoutPlatformSink",
]
File diff suppressed because it is too large Load Diff
@@ -1,105 +0,0 @@
"""Native v4 capability clients.
These clients provide the narrow, typed surface exposed by `Context` for
calling remote capabilities. They handle capability names, payload shaping,
and result decoding, without exposing protocol or transport details.
Migration shims and higher-level orchestration stay outside these native
capability clients so `Context` keeps a narrow, stable surface.
当前公开客户端:
- LLMClient: 文本/结构化/流式 LLM 调用
- MemoryClient: 记忆搜索、保存、读取、删除
- DBClient: 键值存储 get/set/delete/list
- FileServiceClient: 文件令牌注册与解析
- PlatformClient: 平台消息发送与成员查询
- ProviderClient: Provider 元信息与专用 provider proxy
- PersonaManagerClient: 人格管理
- ConversationManagerClient: 对话管理
- KnowledgeBaseManagerClient: 知识库管理
- HTTPClient: Web API 注册
- MetadataClient: 插件元数据查询
- SkillClient: 运行时注册插件 skill
"""
from .db import DBClient
from .files import FileRegistration, FileServiceClient
from .http import HTTPClient
from .llm import ChatMessage, LLMClient, LLMResponse
from .managers import (
ConversationCreateParams,
ConversationManagerClient,
ConversationRecord,
ConversationUpdateParams,
KnowledgeBaseCreateParams,
KnowledgeBaseManagerClient,
KnowledgeBaseRecord,
MessageHistoryManagerClient,
MessageHistoryPage,
MessageHistoryRecord,
MessageHistorySender,
PersonaCreateParams,
PersonaManagerClient,
PersonaRecord,
PersonaUpdateParams,
)
from .mcp import MCPManagerClient, MCPServerRecord, MCPServerScope, MCPSession
from .memory import MemoryClient
from .metadata import MetadataClient, PluginMetadata, StarMetadata
from .platform import PlatformClient, PlatformError, PlatformStats, PlatformStatus
from .provider import (
ManagedProviderRecord,
ProviderChangeEvent,
ProviderClient,
ProviderManagerClient,
)
from .registry import HandlerMetadata, RegistryClient
from .session import SessionPluginManager, SessionServiceManager
from .skills import SkillClient, SkillRegistration
__all__ = [
"ChatMessage",
"ConversationCreateParams",
"ConversationManagerClient",
"ConversationRecord",
"ConversationUpdateParams",
"DBClient",
"FileRegistration",
"FileServiceClient",
"HTTPClient",
"KnowledgeBaseCreateParams",
"KnowledgeBaseManagerClient",
"KnowledgeBaseRecord",
"MessageHistoryManagerClient",
"MessageHistoryPage",
"MessageHistoryRecord",
"MessageHistorySender",
"LLMClient",
"LLMResponse",
"MCPManagerClient",
"MCPSession",
"MCPServerRecord",
"MCPServerScope",
"MemoryClient",
"ManagedProviderRecord",
"MetadataClient",
"PlatformClient",
"PlatformError",
"PlatformStats",
"PlatformStatus",
"PersonaCreateParams",
"PersonaManagerClient",
"PersonaRecord",
"PersonaUpdateParams",
"ProviderChangeEvent",
"ProviderClient",
"ProviderManagerClient",
"PluginMetadata",
"StarMetadata",
"HandlerMetadata",
"RegistryClient",
"SessionPluginManager",
"SessionServiceManager",
"SkillClient",
"SkillRegistration",
]
@@ -1,188 +0,0 @@
"""能力代理模块。
提供 CapabilityProxy 类,作为客户端与 Peer 之间的中间层,负责:
- 检查远程能力是否可用
- 验证流式调用支持
- 统一封装 invoke 和 invoke_stream 调用
设计说明:
CapabilityProxy 是新版架构的核心组件。每个专用客户端 (LLMClient, DBClient 等)
都通过 CapabilityProxy 与远程通信,并在发起调用时绑定当前插件身份,
让运行时把调用者信息放进协议层而不是业务 payload。
使用示例:
proxy = CapabilityProxy(peer)
# 普通调用
result = await proxy.call("llm.chat", {"prompt": "hello"})
# 流式调用
async for delta in proxy.stream("llm.stream_chat", {"prompt": "hello"}):
print(delta["text"])
"""
from __future__ import annotations
from collections.abc import AsyncIterator, Mapping
from typing import Any, Protocol
from .._internal.invocation_context import caller_plugin_scope
from ..errors import AstrBotError
class _CapabilityDescriptorLike(Protocol):
supports_stream: bool | None
class _CapabilityPeerLike(Protocol):
remote_capability_map: Mapping[str, _CapabilityDescriptorLike]
remote_peer: Any | None
async def invoke(
self,
capability: str,
payload: dict[str, Any],
*,
stream: bool = False,
request_id: str | None = None,
) -> dict[str, Any]: ...
async def invoke_stream(
self,
capability: str,
payload: dict[str, Any],
*,
request_id: str | None = None,
) -> AsyncIterator[Any]: ...
class CapabilityProxy:
"""能力代理类,封装 Peer 的能力调用接口。
负责在调用前验证能力可用性和流式支持,提供统一的 call/stream 接口。
Attributes:
_peer: 底层 Peer 实例,负责实际的 RPC 通信
"""
def __init__(
self,
peer: _CapabilityPeerLike,
caller_plugin_id: str | None = None,
request_scope_id: str | None = None,
) -> None:
"""初始化能力代理。
Args:
peer: Peer 实例,提供 remote_capability_map 和 invoke/invoke_stream 方法
"""
self._peer = peer
self._caller_plugin_id = caller_plugin_id
self._request_scope_id = request_scope_id
def _get_descriptor(self, name: str):
"""获取能力描述符。
Args:
name: 能力名称,如 "llm.chat"
Returns:
能力描述符,若不存在则返回 None
"""
capability_map = getattr(self._peer, "remote_capability_map", {})
if not isinstance(capability_map, Mapping):
return None
return capability_map.get(name)
def _remote_initialized(self) -> bool:
peer_attrs = getattr(self._peer, "__dict__", None)
if not isinstance(peer_attrs, dict):
return False
# Avoid getattr() here: MagicMock synthesizes truthy child attributes and
# makes an uninitialized peer look ready.
remote_peer = peer_attrs.get("remote_peer")
capability_map = peer_attrs.get("remote_capability_map")
return bool(remote_peer) or (
isinstance(capability_map, Mapping) and bool(capability_map)
)
def _ensure_available(self, name: str, *, stream: bool) -> None:
"""确保能力可用且支持指定的调用模式。
Args:
name: 能力名称
stream: 是否需要流式支持
Raises:
AstrBotError: 能力不存在或流式不支持
"""
descriptor = self._get_descriptor(name)
if descriptor is None:
if self._remote_initialized():
raise AstrBotError.capability_not_found(name)
return
if stream and not descriptor.supports_stream:
raise AstrBotError.invalid_input(f"{name} 不支持 stream=true")
def _prepare_payload(self, name: str, payload: dict[str, Any]) -> dict[str, Any]:
if (
not isinstance(self._request_scope_id, str)
or not self._request_scope_id
or not name.startswith("system.event.")
):
return payload
scoped_payload = dict(payload)
scoped_payload.setdefault("_request_scope_id", self._request_scope_id)
return scoped_payload
async def call(self, name: str, payload: dict[str, Any]) -> dict[str, Any]:
"""执行普通能力调用(非流式)。
Args:
name: 能力名称,如 "llm.chat", "db.get"
payload: 调用参数字典
Returns:
调用结果字典
Raises:
AstrBotError: 能力不存在或调用失败
示例:
result = await proxy.call("llm.chat", {"prompt": "hello"})
print(result["text"])
"""
self._ensure_available(name, stream=False)
prepared_payload = self._prepare_payload(name, payload)
with caller_plugin_scope(self._caller_plugin_id):
return await self._peer.invoke(name, prepared_payload, stream=False)
async def stream(
self,
name: str,
payload: dict[str, Any],
) -> AsyncIterator[dict[str, Any]]:
"""执行流式能力调用。
Args:
name: 能力名称,如 "llm.stream_chat"
payload: 调用参数字典
Yields:
每个增量数据块(phase="delta" 时的 data 字段)
Raises:
AstrBotError: 能力不存在或不支持流式
示例:
async for delta in proxy.stream("llm.stream_chat", {"prompt": "hello"}):
print(delta["text"], end="")
"""
self._ensure_available(name, stream=True)
prepared_payload = self._prepare_payload(name, payload)
with caller_plugin_scope(self._caller_plugin_id):
event_stream = await self._peer.invoke_stream(name, prepared_payload)
async for event in event_stream:
if event.phase == "delta":
yield event.data
-161
View File
@@ -1,161 +0,0 @@
"""数据库客户端模块。
提供键值存储能力,用于持久化插件数据。
功能说明:
- 数据永久存储,除非用户显式删除
- 值类型支持任意 JSON 数据
- 支持前缀查询键列表
- 支持批量读写
- 支持订阅变更事件
"""
from __future__ import annotations
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import Any
from ._proxy import CapabilityProxy
class DBClient:
"""键值数据库客户端。
提供插件数据的持久化存储能力,数据永久保存直到显式删除。
Attributes:
_proxy: CapabilityProxy 实例,用于远程能力调用
"""
def __init__(self, proxy: CapabilityProxy) -> None:
"""初始化数据库客户端。
Args:
proxy: CapabilityProxy 实例
"""
self._proxy = proxy
async def get(self, key: str) -> Any | None:
"""获取指定键的值。
Args:
key: 数据键名
Returns:
存储的值,若键不存在则返回 None
示例:
data = await ctx.db.get("user_settings")
if data:
print(data["theme"])
"""
output = await self._proxy.call("db.get", {"key": key})
return output.get("value")
async def set(self, key: str, value: Any) -> None:
"""设置键值对。
Args:
key: 数据键名
value: 要存储的 JSON 值
示例:
await ctx.db.set("user_settings", {"theme": "dark", "lang": "zh"})
await ctx.db.set("greeted", True)
"""
await self._proxy.call("db.set", {"key": key, "value": value})
async def delete(self, key: str) -> None:
"""删除指定键的数据。
Args:
key: 要删除的数据键名
示例:
await ctx.db.delete("user_settings")
"""
await self._proxy.call("db.delete", {"key": key})
async def list(self, prefix: str | None = None) -> list[str]:
"""列出匹配前缀的所有键。
Args:
prefix: 键前缀过滤,None 表示列出所有键
Returns:
匹配的键名列表
示例:
# 列出所有用户设置相关的键
keys = await ctx.db.list("user_")
# ["user_settings", "user_profile", "user_history"]
"""
output = await self._proxy.call("db.list", {"prefix": prefix})
keys = output.get("keys")
if not isinstance(keys, (list, tuple)):
return []
return [str(item) for item in keys]
async def get_many(self, keys: Sequence[str]) -> dict[str, Any | None]:
"""批量获取多个键的值。
Args:
keys: 要读取的键列表
Returns:
一个 dict,key 为键名,value 为对应值(不存在则为 None)
示例:
values = await ctx.db.get_many(["user:1", "user:2"])
if values["user:1"] is None:
print("user:1 missing")
"""
output = await self._proxy.call("db.get_many", {"keys": list(keys)})
items = output.get("items")
if not isinstance(items, (list, tuple)):
return {}
result: dict[str, Any | None] = {}
for item in items:
if not isinstance(item, dict):
continue
key = item.get("key")
if not isinstance(key, str):
continue
result[key] = item.get("value")
return result
async def set_many(
self, items: Mapping[str, Any] | Sequence[tuple[str, Any]]
) -> None:
"""批量写入多个键值对。
Args:
items: 键值对集合(dict 或二元组序列)
示例:
await ctx.db.set_many({"user:1": {"name": "a"}, "user:2": {"name": "b"}})
"""
if isinstance(items, Mapping):
pairs = list(items.items())
else:
pairs = list(items)
payload_items: list[dict[str, Any]] = [
{"key": str(key), "value": value} for key, value in pairs
]
await self._proxy.call("db.set_many", {"items": payload_items})
def watch(self, prefix: str | None = None) -> AsyncIterator[dict[str, Any]]:
"""订阅 KV 变更事件(流式)。
Args:
prefix: 键前缀过滤;None 表示订阅所有键
Yields:
变更事件 dict:{"op": "set"|"delete", "key": str, "value": Any|None}
示例:
async for event in ctx.db.watch("user:"):
print(event["op"], event["key"])
"""
return self._proxy.stream("db.watch", {"prefix": prefix})
@@ -1,53 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from ._proxy import CapabilityProxy
@dataclass(slots=True)
class FileRegistration:
token: str
url: str
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> FileRegistration:
return cls(
token=str(payload.get("token", "")),
url=str(payload.get("url", "")),
)
class FileServiceClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def register_file(
self,
path: str,
timeout: float | None = None,
) -> str:
output = await self._proxy.call(
"system.file.register",
{"path": str(path), "timeout": timeout},
)
return FileRegistration.from_payload(output).token
async def handle_file(self, token: str) -> str:
output = await self._proxy.call(
"system.file.handle",
{"token": str(token)},
)
return str(output.get("path", ""))
async def _register_file_url(
self,
path: str,
timeout: float | None = None,
) -> str:
output = await self._proxy.call(
"system.file.register",
{"path": str(path), "timeout": timeout},
)
return FileRegistration.from_payload(output).url
-165
View File
@@ -1,165 +0,0 @@
"""HTTP 客户端模块。
提供 HTTP API 注册能力。
功能说明:
- 注册自定义 Web API 端点
- 支持异步请求处理
- 与宿主 Web 服务器集成
设计说明:
由于跨进程架构,handler 函数无法直接序列化传递。
插件需要先声明处理 HTTP 请求的 capability,然后注册路由到 capability 的映射。
当前插件身份由运行时在协议层透传,客户端 payload 不暴露 `plugin_id`。
调用流程:
HTTP 请求 → 宿主 Web 服务器 → 查找 route 映射 → invoke capability → Worker 执行 handler → 返回响应
示例:
# 插件声明处理 HTTP 请求的 capability
@provide_capability(
name="my_plugin.http_handler",
description="处理 /my-api 的 HTTP 请求",
input_schema={...},
output_schema={...}
)
async def handle_http_request(request_id: str, payload: dict, cancel_token):
return {"status": 200, "body": {"result": "ok"}}
# 注册路由 → capability 映射
await ctx.http.register_api(
route="/my-api",
methods=["GET", "POST"],
handler_capability="my_plugin.http_handler",
description="我的 API"
)
"""
from __future__ import annotations
from typing import Any
from ..decorators import get_capability_meta
from ..errors import AstrBotError
from ._proxy import CapabilityProxy
def _resolve_handler_capability(
handler_capability: str | None,
handler: Any | None,
) -> str:
if handler_capability and handler is not None:
raise AstrBotError.invalid_input(
"register_api 不能同时提供 handler_capability 和 handler",
hint="请二选一:传 capability 名称字符串,或传 @provide_capability 标记的方法",
)
if handler_capability:
return handler_capability
if handler is None:
raise AstrBotError.invalid_input(
"register_api 需要提供 handler_capability 或 handler",
hint="示例:handler_capability='demo.http_handler' 或 handler=self.http_handler_capability",
)
target = getattr(handler, "__func__", handler)
meta = get_capability_meta(target)
if meta is None:
raise AstrBotError.invalid_input(
"register_api(handler=...) 需要传入使用 @provide_capability 声明的方法",
hint="请先用 @provide_capability(name='demo.http_handler', ...) 标记该方法",
)
return meta.descriptor.name
class HTTPClient:
"""HTTP 能力客户端。
提供 Web API 注册能力,允许插件暴露自定义 HTTP 端点。
Attributes:
_proxy: CapabilityProxy 实例,用于远程能力调用
"""
def __init__(self, proxy: CapabilityProxy) -> None:
"""初始化 HTTP 客户端。
Args:
proxy: CapabilityProxy 实例
"""
self._proxy = proxy
async def register_api(
self,
route: str,
handler_capability: str | None = None,
*,
handler: Any | None = None,
methods: list[str] | None = None,
description: str = "",
) -> None:
"""注册 Web API 端点。
Args:
route: API 路由路径(如 "/my-api")
handler_capability: 处理此路由的 capability 名称
handler: 使用 @provide_capability 标记的方法引用
methods: HTTP 方法列表,默认 ["GET"]
description: API 描述
示例:
await ctx.http.register_api(
route="/my-api",
handler_capability="my_plugin.http_handler",
methods=["GET", "POST"],
description="我的 API"
)
"""
if methods is None:
methods = ["GET"]
resolved_handler = _resolve_handler_capability(handler_capability, handler)
await self._proxy.call(
"http.register_api",
{
"route": route,
"methods": methods,
"handler_capability": resolved_handler,
"description": description,
},
)
async def unregister_api(
self, route: str, methods: list[str] | None = None
) -> None:
"""注销 Web API 端点。
Args:
route: API 路由路径
methods: HTTP 方法列表,None 表示所有方法
示例:
await ctx.http.unregister_api("/my-api")
"""
if methods is None:
methods = []
await self._proxy.call(
"http.unregister_api",
{"route": route, "methods": methods},
)
async def list_apis(self) -> list[dict[str, Any]]:
"""列出当前插件注册的所有 API。
Returns:
API 列表,每项包含 route, methods, description
示例:
apis = await ctx.http.list_apis()
for api in apis:
print(f"{api['route']}: {api['methods']}")
"""
output = await self._proxy.call(
"http.list_apis",
{},
)
return output.get("apis", [])
-293
View File
@@ -1,293 +0,0 @@
"""大语言模型客户端模块。
提供 v4 原生的 LLM 能力调用接口。
设计边界:
- `chat()` 是便捷文本接口,返回最终文本
- `chat_raw()` 返回完整结构化响应
- `stream_chat()` 返回文本增量
- Agent 循环、动态工具注册等更高层 orchestration 不放在客户端内,
由上层运行时或独立迁移入口承接
"""
from __future__ import annotations
from collections.abc import AsyncGenerator, Mapping, Sequence
from typing import Any
from pydantic import BaseModel, Field
from ._proxy import CapabilityProxy
class ChatMessage(BaseModel):
"""聊天消息模型。
用于构建对话历史,传递给 LLM。
Attributes:
role: 消息角色,如 "user", "assistant", "system"
content: 消息内容
示例:
history = [
ChatMessage(role="user", content="你好"),
ChatMessage(role="assistant", content="你好!有什么可以帮助你的?"),
ChatMessage(role="user", content="今天天气怎么样?"),
]
"""
role: str
content: str
ChatHistoryItem = ChatMessage | Mapping[str, Any]
def _serialize_history(
history: Sequence[ChatHistoryItem] | None,
) -> list[dict[str, Any]]:
if history is None:
return []
serialized: list[dict[str, Any]] = []
for item in history:
if isinstance(item, ChatMessage):
serialized.append(item.model_dump())
continue
if isinstance(item, Mapping):
serialized.append(dict(item))
continue
raise TypeError("history 项必须是 ChatMessage 或 mapping")
return serialized
def _normalize_chat_context_payload(
*,
history: Sequence[ChatHistoryItem] | None = None,
contexts: Sequence[ChatHistoryItem] | None = None,
) -> dict[str, list[dict[str, Any]]]:
if contexts is not None:
return {"contexts": _serialize_history(contexts)}
if history is not None:
return {"contexts": _serialize_history(history)}
return {}
def _build_chat_payload(
prompt: str,
*,
system: str | None = None,
history: Sequence[ChatHistoryItem] | None = None,
contexts: Sequence[ChatHistoryItem] | None = None,
provider_id: str | None = None,
tool_calls_result: list[dict[str, Any]] | None = None,
model: str | None = None,
temperature: float | None = None,
extra: dict[str, Any] | None = None,
) -> dict[str, Any]:
payload: dict[str, Any] = {"prompt": prompt}
if system is not None:
payload["system"] = system
payload.update(_normalize_chat_context_payload(history=history, contexts=contexts))
if provider_id is not None:
payload["provider_id"] = provider_id
if tool_calls_result is not None:
payload["tool_calls_result"] = [dict(item) for item in tool_calls_result]
if model is not None:
payload["model"] = model
if temperature is not None:
payload["temperature"] = temperature
if extra:
payload.update(extra)
return payload
class LLMResponse(BaseModel):
"""LLM 响应模型。
包含完整的 LLM 响应信息,用于 chat_raw() 方法返回。
Attributes:
text: 生成的文本内容
usage: Token 使用统计,如 {"prompt_tokens": 10, "completion_tokens": 20}
finish_reason: 结束原因,如 "stop", "length", "tool_calls"
tool_calls: 工具调用列表(如果 LLM 决定调用工具)
"""
text: str
usage: dict[str, Any] | None = None
finish_reason: str | None = None
tool_calls: list[dict[str, Any]] = Field(default_factory=list)
role: str | None = None
reasoning_content: str | None = None
reasoning_signature: str | None = None
class LLMClient:
"""大语言模型客户端。
提供与 LLM 交互的能力,支持普通聊天和流式聊天。
Attributes:
_proxy: CapabilityProxy 实例,用于远程能力调用
"""
def __init__(self, proxy: CapabilityProxy) -> None:
"""初始化 LLM 客户端。
Args:
proxy: CapabilityProxy 实例
"""
self._proxy = proxy
async def chat(
self,
prompt: str,
*,
system: str | None = None,
history: Sequence[ChatHistoryItem] | None = None,
contexts: Sequence[ChatHistoryItem] | None = None,
provider_id: str | None = None,
tool_calls_result: list[dict[str, Any]] | None = None,
model: str | None = None,
temperature: float | None = None,
**kwargs: Any,
) -> str:
"""发送聊天请求并返回文本响应。
这是简化的聊天接口,仅返回生成的文本内容。
如需完整响应信息(包括 usage、tool_calls),请使用 chat_raw()。
Args:
prompt: 用户输入的提示文本
system: 系统提示词,用于指导 LLM 行为
history: 对话历史,用于保持上下文连续性
model: 指定使用的模型名称(可选,由核心自动选择)
temperature: 生成温度,控制随机性(0-1)
**kwargs: 额外透传参数,如 `image_urls`、`tools`
Returns:
LLM 生成的文本内容
示例:
# 简单对话
reply = await ctx.llm.chat("你好,介绍一下自己")
# 带历史的对话
history = [
ChatMessage(role="user", content="我叫小明"),
ChatMessage(role="assistant", content="你好小明!"),
]
reply = await ctx.llm.chat("你记得我的名字吗?", history=history)
"""
output = await self._proxy.call(
"llm.chat",
_build_chat_payload(
prompt,
system=system,
history=history,
contexts=contexts,
provider_id=provider_id,
tool_calls_result=tool_calls_result,
model=model,
temperature=temperature,
extra=kwargs,
),
)
return str(output.get("text", ""))
async def chat_raw(
self,
prompt: str,
*,
system: str | None = None,
history: Sequence[ChatHistoryItem] | None = None,
contexts: Sequence[ChatHistoryItem] | None = None,
provider_id: str | None = None,
tool_calls_result: list[dict[str, Any]] | None = None,
model: str | None = None,
temperature: float | None = None,
**kwargs: Any,
) -> LLMResponse:
"""发送聊天请求并返回完整响应。
与 chat() 不同,此方法返回完整的 LLMResponse 对象,
包含 usage、finish_reason、tool_calls 等信息。
Args:
prompt: 用户输入的提示文本
**kwargs: 额外参数,如 system, history, model, temperature 等
Returns:
LLMResponse 对象,包含完整响应信息
示例:
response = await ctx.llm.chat_raw("写一首诗", temperature=0.8)
print(f"生成文本: {response.text}")
print(f"Token 使用: {response.usage}")
"""
payload = _build_chat_payload(
prompt,
system=system,
history=history,
contexts=contexts,
provider_id=provider_id,
tool_calls_result=tool_calls_result,
model=model,
temperature=temperature,
extra=kwargs,
)
output = await self._proxy.call(
"llm.chat_raw",
payload,
)
return LLMResponse.model_validate(output)
async def stream_chat(
self,
prompt: str,
*,
system: str | None = None,
history: Sequence[ChatHistoryItem] | None = None,
contexts: Sequence[ChatHistoryItem] | None = None,
provider_id: str | None = None,
tool_calls_result: list[dict[str, Any]] | None = None,
model: str | None = None,
temperature: float | None = None,
**kwargs: Any,
) -> AsyncGenerator[str, None]:
"""流式聊天,逐块返回响应文本。
适用于需要实时显示生成内容的场景,如聊天界面。
Args:
prompt: 用户输入的提示文本
system: 系统提示词
history: 对话历史
model: 指定模型
temperature: 采样温度
**kwargs: 额外透传参数,如 `image_urls`、`tools`
Yields:
每个生成的文本块
示例:
async for chunk in ctx.llm.stream_chat("讲一个故事"):
print(chunk, end="", flush=True)
"""
async for data in self._proxy.stream(
"llm.stream_chat",
_build_chat_payload(
prompt,
system=system,
history=history,
contexts=contexts,
provider_id=provider_id,
tool_calls_result=tool_calls_result,
model=model,
temperature=temperature,
extra=kwargs,
),
):
yield str(data.get("text", ""))
@@ -1,875 +0,0 @@
"""Typed SDK manager clients for persona, conversation, and knowledge base."""
from __future__ import annotations
from datetime import datetime
from typing import Any
from pydantic import BaseModel, ConfigDict, Field, model_validator
from ..errors import AstrBotError, ErrorCodes
from ..message.components import (
BaseMessageComponent,
component_to_payload_sync,
payload_to_component,
)
from ..message.session import MessageSession
from ._proxy import CapabilityProxy
class _ManagerModel(BaseModel):
model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True)
def to_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_none=True)
def to_update_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_unset=True)
def _normalize_session(session: str | MessageSession) -> str:
if isinstance(session, MessageSession):
return str(session)
return str(session)
def _require_message_history_session(
session: MessageSession,
) -> dict[str, str]:
if not isinstance(session, MessageSession):
raise TypeError(
"message_history requires astrbot_sdk.message.session.MessageSession"
)
return {
"platform_id": str(session.platform_id),
"message_type": str(session.message_type),
"session_id": str(session.session_id),
}
def _normalize_message_history_parts(
parts: list[BaseMessageComponent],
) -> list[dict[str, Any]]:
normalized: list[dict[str, Any]] = []
for part in parts:
if not isinstance(part, BaseMessageComponent):
raise TypeError(
"message_history.append requires BaseMessageComponent items in parts"
)
normalized.append(component_to_payload_sync(part))
return normalized
class PersonaRecord(_ManagerModel):
persona_id: str
system_prompt: str
begin_dialogs: list[str] = Field(default_factory=list)
tools: list[str] | None = None
skills: list[str] | None = None
custom_error_message: str | None = None
folder_id: str | None = None
sort_order: int = 0
created_at: str | None = None
updated_at: str | None = None
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> PersonaRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class PersonaCreateParams(_ManagerModel):
persona_id: str
system_prompt: str
begin_dialogs: list[str] = Field(default_factory=list)
tools: list[str] | None = None
skills: list[str] | None = None
custom_error_message: str | None = None
folder_id: str | None = None
sort_order: int = 0
class PersonaUpdateParams(_ManagerModel):
system_prompt: str | None = None
begin_dialogs: list[str] | None = None
tools: list[str] | None = None
skills: list[str] | None = None
custom_error_message: str | None = None
class ConversationRecord(_ManagerModel):
conversation_id: str
session: str
platform_id: str
history: list[dict[str, Any]] = Field(default_factory=list)
title: str | None = None
persona_id: str | None = None
created_at: str | None = None
updated_at: str | None = None
token_usage: int | None = None
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> ConversationRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class ConversationCreateParams(_ManagerModel):
platform_id: str | None = None
history: list[dict[str, Any]] | None = None
title: str | None = None
persona_id: str | None = None
class ConversationUpdateParams(_ManagerModel):
history: list[dict[str, Any]] | None = None
title: str | None = None
persona_id: str | None = None
token_usage: int | None = None
class MessageHistorySender(_ManagerModel):
sender_id: str | None = None
sender_name: str | None = None
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> MessageHistorySender | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class MessageHistoryRecord(_ManagerModel):
id: int
session: MessageSession
sender: MessageHistorySender = Field(default_factory=MessageHistorySender)
parts: list[BaseMessageComponent] = Field(default_factory=list)
metadata: dict[str, Any] = Field(default_factory=dict)
created_at: datetime | None = None
updated_at: datetime | None = None
idempotency_key: str | None = None
@model_validator(mode="before")
@classmethod
def _normalize_payload(cls, value: Any) -> Any:
if not isinstance(value, dict):
return value
normalized = dict(value)
session_payload = normalized.get("session")
if isinstance(session_payload, dict):
normalized["session"] = MessageSession(
platform_id=str(session_payload.get("platform_id", "")),
message_type=str(session_payload.get("message_type", "")),
session_id=str(session_payload.get("session_id", "")),
)
sender_payload = normalized.get("sender")
if isinstance(sender_payload, dict):
normalized["sender"] = MessageHistorySender.model_validate(sender_payload)
elif sender_payload is None:
normalized["sender"] = MessageHistorySender()
parts_payload = normalized.get("parts")
if isinstance(parts_payload, list):
normalized["parts"] = [
payload_to_component(item)
for item in parts_payload
if isinstance(item, dict)
]
metadata_payload = normalized.get("metadata")
if not isinstance(metadata_payload, dict):
normalized["metadata"] = {}
return normalized
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> MessageHistoryRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class MessageHistoryPage(_ManagerModel):
records: list[MessageHistoryRecord] = Field(default_factory=list)
next_cursor: str | None = None
total: int | None = None
@model_validator(mode="before")
@classmethod
def _normalize_payload(cls, value: Any) -> Any:
if not isinstance(value, dict):
return value
normalized = dict(value)
records_payload = normalized.get("records")
if isinstance(records_payload, list):
normalized["records"] = [
record
for record in (
MessageHistoryRecord.from_payload(item)
if isinstance(item, dict)
else None
for item in records_payload
)
if record is not None
]
return normalized
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> MessageHistoryPage | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class KnowledgeBaseRecord(_ManagerModel):
kb_id: str
kb_name: str
description: str | None = None
emoji: str | None = None
embedding_provider_id: str
rerank_provider_id: str | None = None
chunk_size: int | None = None
chunk_overlap: int | None = None
top_k_dense: int | None = None
top_k_sparse: int | None = None
top_m_final: int | None = None
doc_count: int = 0
chunk_count: int = 0
created_at: str | None = None
updated_at: str | None = None
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> KnowledgeBaseRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class KnowledgeBaseCreateParams(_ManagerModel):
kb_name: str
embedding_provider_id: str
description: str | None = None
emoji: str | None = None
rerank_provider_id: str | None = None
chunk_size: int | None = None
chunk_overlap: int | None = None
top_k_dense: int | None = None
top_k_sparse: int | None = None
top_m_final: int | None = None
class KnowledgeBaseUpdateParams(_ManagerModel):
kb_name: str | None = None
embedding_provider_id: str | None = None
description: str | None = None
emoji: str | None = None
rerank_provider_id: str | None = None
chunk_size: int | None = None
chunk_overlap: int | None = None
top_k_dense: int | None = None
top_k_sparse: int | None = None
top_m_final: int | None = None
class KnowledgeBaseDocumentRecord(_ManagerModel):
doc_id: str
kb_id: str
doc_name: str
file_type: str
file_size: int
file_path: str = ""
chunk_count: int = 0
media_count: int = 0
created_at: str | None = None
updated_at: str | None = None
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> KnowledgeBaseDocumentRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class KnowledgeBaseRetrieveResultItem(_ManagerModel):
chunk_id: str
doc_id: str
kb_id: str
kb_name: str
doc_name: str
chunk_index: int
content: str
score: float
char_count: int
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> KnowledgeBaseRetrieveResultItem | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class KnowledgeBaseRetrieveResult(_ManagerModel):
context_text: str
results: list[KnowledgeBaseRetrieveResultItem] = Field(default_factory=list)
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> KnowledgeBaseRetrieveResult | None:
if not isinstance(payload, dict):
return None
items = payload.get("results")
normalized_items = (
[
item.model_dump()
for item in (
KnowledgeBaseRetrieveResultItem.from_payload(candidate)
if isinstance(candidate, dict)
else None
for candidate in items
)
if item is not None
]
if isinstance(items, list)
else []
)
return cls.model_validate(
{
"context_text": str(payload.get("context_text", "")),
"results": normalized_items,
}
)
class KnowledgeBaseDocumentUploadParams(_ManagerModel):
file_token: str | None = None
url: str | None = None
text: str | None = None
file_name: str | None = None
file_type: str | None = None
chunk_size: int | None = None
chunk_overlap: int | None = None
batch_size: int | None = None
tasks_limit: int | None = None
max_retries: int | None = None
enable_cleaning: bool | None = None
cleaning_provider_id: str | None = None
@model_validator(mode="after")
def _validate_source(self) -> KnowledgeBaseDocumentUploadParams:
if any(
isinstance(value, str) and value.strip()
for value in (self.file_token, self.url, self.text)
):
return self
raise ValueError(
"knowledge base document upload requires file_token, url, or text"
)
class PersonaManagerClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def get_persona(self, persona_id: str) -> PersonaRecord:
try:
output = await self._proxy.call(
"persona.get",
{"persona_id": str(persona_id)},
)
except AstrBotError as exc:
if exc.code == ErrorCodes.INVALID_INPUT:
raise ValueError(f"persona not found: {persona_id}") from exc
raise
persona = PersonaRecord.from_payload(output.get("persona"))
if persona is None:
raise ValueError(f"persona not found: {persona_id}")
return persona
async def get_all_personas(self) -> list[PersonaRecord]:
output = await self._proxy.call("persona.list", {})
items = output.get("personas")
if not isinstance(items, list):
return []
return [
persona
for persona in (
PersonaRecord.from_payload(item) if isinstance(item, dict) else None
for item in items
)
if persona is not None
]
async def create_persona(self, params: PersonaCreateParams) -> PersonaRecord:
output = await self._proxy.call(
"persona.create",
{"persona": params.to_payload()},
)
persona = PersonaRecord.from_payload(output.get("persona"))
if persona is None:
raise ValueError("persona.create returned no persona")
return persona
async def update_persona(
self,
persona_id: str,
params: PersonaUpdateParams,
) -> PersonaRecord | None:
output = await self._proxy.call(
"persona.update",
{"persona_id": str(persona_id), "persona": params.to_update_payload()},
)
return PersonaRecord.from_payload(output.get("persona"))
async def delete_persona(self, persona_id: str) -> None:
await self._proxy.call("persona.delete", {"persona_id": str(persona_id)})
class ConversationManagerClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def new_conversation(
self,
session: str | MessageSession,
params: ConversationCreateParams | None = None,
) -> str:
output = await self._proxy.call(
"conversation.new",
{
"session": _normalize_session(session),
"conversation": (params.to_payload() if params is not None else {}),
},
)
return str(output.get("conversation_id", ""))
async def switch_conversation(
self,
session: str | MessageSession,
conversation_id: str,
) -> None:
await self._proxy.call(
"conversation.switch",
{
"session": _normalize_session(session),
"conversation_id": str(conversation_id),
},
)
async def delete_conversation(
self,
session: str | MessageSession,
conversation_id: str | None = None,
) -> None:
"""Delete one conversation for the session.
When ``conversation_id`` is ``None``, this deletes the current selected
conversation for the session only. It does not delete all conversations
under the session.
"""
await self._proxy.call(
"conversation.delete",
{
"session": _normalize_session(session),
"conversation_id": conversation_id,
},
)
async def get_conversation(
self,
session: str | MessageSession,
conversation_id: str,
*,
create_if_not_exists: bool = False,
) -> ConversationRecord | None:
output = await self._proxy.call(
"conversation.get",
{
"session": _normalize_session(session),
"conversation_id": str(conversation_id),
"create_if_not_exists": bool(create_if_not_exists),
},
)
return ConversationRecord.from_payload(output.get("conversation"))
async def get_current_conversation(
self,
session: str | MessageSession,
*,
create_if_not_exists: bool = False,
) -> ConversationRecord | None:
output = await self._proxy.call(
"conversation.get_current",
{
"session": _normalize_session(session),
"create_if_not_exists": bool(create_if_not_exists),
},
)
return ConversationRecord.from_payload(output.get("conversation"))
async def get_conversations(
self,
session: str | MessageSession | None = None,
*,
platform_id: str | None = None,
) -> list[ConversationRecord]:
output = await self._proxy.call(
"conversation.list",
{
"session": (
_normalize_session(session) if session is not None else None
),
"platform_id": platform_id,
},
)
items = output.get("conversations")
if not isinstance(items, list):
return []
return [
conversation
for conversation in (
ConversationRecord.from_payload(item)
if isinstance(item, dict)
else None
for item in items
)
if conversation is not None
]
async def update_conversation(
self,
session: str | MessageSession,
conversation_id: str | None = None,
params: ConversationUpdateParams | None = None,
) -> None:
await self._proxy.call(
"conversation.update",
{
"session": _normalize_session(session),
"conversation_id": conversation_id,
"conversation": (
params.to_update_payload() if params is not None else {}
),
},
)
async def unset_persona(
self,
session: str | MessageSession,
conversation_id: str | None = None,
) -> None:
await self._proxy.call(
"conversation.unset_persona",
{
"session": _normalize_session(session),
"conversation_id": conversation_id,
},
)
class MessageHistoryManagerClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def list(
self,
session: MessageSession,
*,
cursor: str | None = None,
limit: int = 50,
) -> MessageHistoryPage:
output = await self._proxy.call(
"message_history.list",
{
"session": _require_message_history_session(session),
"cursor": str(cursor) if cursor is not None else None,
"limit": int(limit),
},
)
page = MessageHistoryPage.from_payload(output.get("page"))
if page is None:
raise ValueError("message_history.list returned no page")
return page
async def get(
self,
session: MessageSession,
record_id: int,
) -> MessageHistoryRecord | None:
output = await self._proxy.call(
"message_history.get_by_id",
{
"session": _require_message_history_session(session),
"record_id": int(record_id),
},
)
return MessageHistoryRecord.from_payload(output.get("record"))
async def get_by_id(
self,
session: MessageSession,
record_id: int,
) -> MessageHistoryRecord | None:
return await self.get(session, record_id)
async def append(
self,
session: MessageSession,
*,
parts: list[BaseMessageComponent],
sender: MessageHistorySender,
metadata: dict[str, Any] | None = None,
idempotency_key: str | None = None,
) -> MessageHistoryRecord:
if isinstance(sender, MessageHistorySender):
sender_payload = sender.to_payload()
elif isinstance(sender, dict):
sender_payload = MessageHistorySender.model_validate(sender).to_payload()
else:
raise TypeError(
"message_history.append requires MessageHistorySender for sender"
)
output = await self._proxy.call(
"message_history.append",
{
"session": _require_message_history_session(session),
"sender": sender_payload,
"parts": _normalize_message_history_parts(parts),
"metadata": dict(metadata or {}),
"idempotency_key": (
str(idempotency_key) if idempotency_key is not None else None
),
},
)
record = MessageHistoryRecord.from_payload(output.get("record"))
if record is None:
raise ValueError("message_history.append returned no record")
return record
async def delete_before(
self,
session: MessageSession,
*,
before: datetime,
) -> int:
output = await self._proxy.call(
"message_history.delete_before",
{
"session": _require_message_history_session(session),
"before": before.isoformat(),
},
)
return int(output.get("deleted_count", 0) or 0)
async def delete_after(
self,
session: MessageSession,
*,
after: datetime,
) -> int:
output = await self._proxy.call(
"message_history.delete_after",
{
"session": _require_message_history_session(session),
"after": after.isoformat(),
},
)
return int(output.get("deleted_count", 0) or 0)
async def delete_all(self, session: MessageSession) -> int:
output = await self._proxy.call(
"message_history.delete_all",
{"session": _require_message_history_session(session)},
)
return int(output.get("deleted_count", 0) or 0)
class KnowledgeBaseManagerClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def list_kbs(self) -> list[KnowledgeBaseRecord]:
output = await self._proxy.call("kb.list", {})
items = output.get("kbs")
if not isinstance(items, list):
return []
return [
kb
for kb in (
KnowledgeBaseRecord.from_payload(item)
if isinstance(item, dict)
else None
for item in items
)
if kb is not None
]
async def get_kb(self, kb_id: str) -> KnowledgeBaseRecord | None:
output = await self._proxy.call("kb.get", {"kb_id": str(kb_id)})
return KnowledgeBaseRecord.from_payload(output.get("kb"))
async def create_kb(
self,
params: KnowledgeBaseCreateParams,
) -> KnowledgeBaseRecord:
output = await self._proxy.call("kb.create", {"kb": params.to_payload()})
kb = KnowledgeBaseRecord.from_payload(output.get("kb"))
if kb is None:
raise ValueError("kb.create returned no knowledge base")
return kb
async def update_kb(
self,
kb_id: str,
params: KnowledgeBaseUpdateParams,
) -> KnowledgeBaseRecord | None:
output = await self._proxy.call(
"kb.update",
{"kb_id": str(kb_id), "kb": params.to_update_payload()},
)
return KnowledgeBaseRecord.from_payload(output.get("kb"))
async def delete_kb(self, kb_id: str) -> bool:
output = await self._proxy.call("kb.delete", {"kb_id": str(kb_id)})
return bool(output.get("deleted", False))
async def retrieve(
self,
query: str,
*,
kb_ids: list[str] | None = None,
kb_names: list[str] | None = None,
top_k_fusion: int | None = None,
top_m_final: int | None = None,
) -> KnowledgeBaseRetrieveResult | None:
request_payload: dict[str, Any] = {
"query": str(query),
"kb_ids": [str(item) for item in (kb_ids or [])],
"kb_names": [str(item) for item in (kb_names or [])],
}
if top_k_fusion is not None:
request_payload["top_k_fusion"] = int(top_k_fusion)
if top_m_final is not None:
request_payload["top_m_final"] = int(top_m_final)
output = await self._proxy.call(
"kb.retrieve",
request_payload,
)
return KnowledgeBaseRetrieveResult.from_payload(output.get("result"))
async def upload_document(
self,
kb_id: str,
params: KnowledgeBaseDocumentUploadParams,
) -> KnowledgeBaseDocumentRecord:
output = await self._proxy.call(
"kb.document.upload",
{"kb_id": str(kb_id), "document": params.to_payload()},
)
document = KnowledgeBaseDocumentRecord.from_payload(output.get("document"))
if document is None:
raise ValueError("kb.document.upload returned no document")
return document
async def list_documents(
self,
kb_id: str,
*,
offset: int = 0,
limit: int = 100,
) -> list[KnowledgeBaseDocumentRecord]:
output = await self._proxy.call(
"kb.document.list",
{"kb_id": str(kb_id), "offset": int(offset), "limit": int(limit)},
)
items = output.get("documents")
if not isinstance(items, list):
return []
return [
document
for document in (
KnowledgeBaseDocumentRecord.from_payload(item)
if isinstance(item, dict)
else None
for item in items
)
if document is not None
]
async def get_document(
self,
kb_id: str,
doc_id: str,
) -> KnowledgeBaseDocumentRecord | None:
output = await self._proxy.call(
"kb.document.get",
{"kb_id": str(kb_id), "doc_id": str(doc_id)},
)
return KnowledgeBaseDocumentRecord.from_payload(output.get("document"))
async def delete_document(
self,
kb_id: str,
doc_id: str,
) -> bool:
output = await self._proxy.call(
"kb.document.delete",
{"kb_id": str(kb_id), "doc_id": str(doc_id)},
)
return bool(output.get("deleted", False))
async def refresh_document(
self,
kb_id: str,
doc_id: str,
) -> KnowledgeBaseDocumentRecord | None:
output = await self._proxy.call(
"kb.document.refresh",
{"kb_id": str(kb_id), "doc_id": str(doc_id)},
)
return KnowledgeBaseDocumentRecord.from_payload(output.get("document"))
__all__ = [
"ConversationCreateParams",
"ConversationManagerClient",
"ConversationRecord",
"ConversationUpdateParams",
"KnowledgeBaseCreateParams",
"KnowledgeBaseDocumentRecord",
"KnowledgeBaseDocumentUploadParams",
"KnowledgeBaseManagerClient",
"KnowledgeBaseRecord",
"KnowledgeBaseRetrieveResult",
"KnowledgeBaseRetrieveResultItem",
"KnowledgeBaseUpdateParams",
"MessageHistoryManagerClient",
"MessageHistoryPage",
"MessageHistoryRecord",
"MessageHistorySender",
"PersonaCreateParams",
"PersonaManagerClient",
"PersonaRecord",
"PersonaUpdateParams",
]
-302
View File
@@ -1,302 +0,0 @@
from __future__ import annotations
from contextlib import AbstractAsyncContextManager
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
from ..errors import AstrBotError
from ._proxy import CapabilityProxy
class MCPServerScope(str, Enum):
local = "local"
global_ = "global"
@dataclass(slots=True)
class MCPServerRecord:
name: str
scope: MCPServerScope
active: bool
running: bool
config: dict[str, Any] = field(default_factory=dict)
tools: list[str] = field(default_factory=list)
errlogs: list[str] = field(default_factory=list)
last_error: str | None = None
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> MCPServerRecord | None:
if not isinstance(payload, dict):
return None
scope_value = str(payload.get("scope") or MCPServerScope.local.value).strip()
try:
scope = MCPServerScope(scope_value)
except ValueError:
scope = MCPServerScope.local
return cls(
name=str(payload.get("name", "")),
scope=scope,
active=bool(payload.get("active", False)),
running=bool(payload.get("running", False)),
config=(
dict(payload.get("config"))
if isinstance(payload.get("config"), dict)
else {}
),
tools=[
str(item)
for item in payload.get("tools", [])
if isinstance(item, str) and item
]
if isinstance(payload.get("tools"), list)
else [],
errlogs=[
str(item)
for item in payload.get("errlogs", [])
if isinstance(item, str)
]
if isinstance(payload.get("errlogs"), list)
else [],
last_error=(
str(payload.get("last_error"))
if payload.get("last_error") is not None
else None
),
)
class MCPSession(AbstractAsyncContextManager["MCPSession"]):
def __init__(
self,
proxy: CapabilityProxy,
*,
name: str,
config: dict[str, Any],
timeout: float,
) -> None:
self._proxy = proxy
self._name = str(name)
self._config = dict(config)
self._timeout = float(timeout)
self._session_id: str | None = None
self._tools: list[str] = []
async def __aenter__(self) -> MCPSession:
output = await self._proxy.call(
"mcp.session.open",
{
"name": self._name,
"config": dict(self._config),
"timeout": self._timeout,
},
)
session_id = str(output.get("session_id", "")).strip()
if not session_id:
raise ValueError("mcp.session.open returned no session_id")
self._session_id = session_id
tools = output.get("tools")
self._tools = (
[str(item) for item in tools if isinstance(item, str)]
if isinstance(tools, list)
else []
)
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
session_id = self._session_id
self._session_id = None
self._tools = []
if not session_id:
return
try:
await self._proxy.call("mcp.session.close", {"session_id": session_id})
except AstrBotError:
raise
except Exception:
# Session cleanup should not mask the original error raised inside the
# managed block.
if exc_type is None:
raise
async def call_tool(
self,
tool_name: str,
args: dict[str, Any] | None = None,
) -> dict[str, Any]:
session_id = self._require_session_id()
output = await self._proxy.call(
"mcp.session.call_tool",
{
"session_id": session_id,
"tool_name": str(tool_name),
"args": dict(args or {}),
},
)
result = output.get("result")
if not isinstance(result, dict):
raise ValueError("mcp.session.call_tool returned no result object")
return dict(result)
async def list_tools(self) -> list[str]:
session_id = self._require_session_id()
output = await self._proxy.call(
"mcp.session.list_tools",
{"session_id": session_id},
)
tools = output.get("tools")
self._tools = (
[str(item) for item in tools if isinstance(item, str)]
if isinstance(tools, list)
else []
)
return list(self._tools)
def _require_session_id(self) -> str:
if self._session_id is None:
raise RuntimeError("MCP session is not active; use 'async with'")
return self._session_id
class MCPManagerClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def get_server(self, name: str) -> MCPServerRecord | None:
output = await self._proxy.call("mcp.local.get", {"name": str(name)})
return MCPServerRecord.from_payload(output.get("server"))
async def list_servers(self) -> list[MCPServerRecord]:
output = await self._proxy.call("mcp.local.list", {})
items = output.get("servers")
if not isinstance(items, list):
return []
return [
record
for record in (
MCPServerRecord.from_payload(item) if isinstance(item, dict) else None
for item in items
)
if record is not None
]
async def enable_server(self, name: str) -> MCPServerRecord:
output = await self._proxy.call("mcp.local.enable", {"name": str(name)})
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.local.enable returned no server")
return record
async def disable_server(self, name: str) -> MCPServerRecord:
output = await self._proxy.call("mcp.local.disable", {"name": str(name)})
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.local.disable returned no server")
return record
async def wait_until_ready(
self,
name: str,
*,
timeout: float = 30.0,
) -> MCPServerRecord:
output = await self._proxy.call(
"mcp.local.wait_until_ready",
{"name": str(name), "timeout": float(timeout)},
)
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.local.wait_until_ready returned no server")
return record
def session(
self,
name: str,
config: dict[str, Any],
*,
timeout: float = 30.0,
) -> MCPSession:
return MCPSession(
self._proxy,
name=str(name),
config=dict(config),
timeout=float(timeout),
)
async def register_global_server(
self,
name: str,
config: dict[str, Any],
*,
timeout: float = 30.0,
) -> MCPServerRecord:
output = await self._proxy.call(
"mcp.global.register",
{
"name": str(name),
"config": dict(config),
"timeout": float(timeout),
},
)
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.global.register returned no server")
return record
async def get_global_server(self, name: str) -> MCPServerRecord | None:
output = await self._proxy.call("mcp.global.get", {"name": str(name)})
return MCPServerRecord.from_payload(output.get("server"))
async def list_global_servers(self) -> list[MCPServerRecord]:
output = await self._proxy.call("mcp.global.list", {})
items = output.get("servers")
if not isinstance(items, list):
return []
return [
record
for record in (
MCPServerRecord.from_payload(item) if isinstance(item, dict) else None
for item in items
)
if record is not None
]
async def enable_global_server(
self,
name: str,
*,
timeout: float = 30.0,
) -> MCPServerRecord:
output = await self._proxy.call(
"mcp.global.enable",
{"name": str(name), "timeout": float(timeout)},
)
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.global.enable returned no server")
return record
async def disable_global_server(self, name: str) -> MCPServerRecord:
output = await self._proxy.call("mcp.global.disable", {"name": str(name)})
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.global.disable returned no server")
return record
async def unregister_global_server(self, name: str) -> MCPServerRecord:
output = await self._proxy.call("mcp.global.unregister", {"name": str(name)})
record = MCPServerRecord.from_payload(output.get("server"))
if record is None:
raise ValueError("mcp.global.unregister returned no server")
return record
__all__ = [
"MCPManagerClient",
"MCPSession",
"MCPServerRecord",
"MCPServerScope",
]
@@ -1,373 +0,0 @@
"""记忆客户端模块。
提供 AI 记忆存储能力,用于存储和检索对话记忆、用户偏好等上下文数据。
设计说明:
MemoryClient 与 DBClient 的区别:
- DBClient: 简单的键值存储,精确匹配
- MemoryClient: 支持基于当前 bridge 行为的记忆检索,适合 AI 上下文管理
记忆系统可用于:
- 存储用户偏好和设置
- 记录对话摘要
- 缓存 AI 推理结果
"""
from __future__ import annotations
from typing import Any, Literal
from .._internal.memory_utils import join_memory_namespace
from ._proxy import CapabilityProxy
def _normalize_search_item(item: Any) -> dict[str, Any] | None:
if not isinstance(item, dict):
return None
normalized = dict(item)
value = normalized.get("value")
if isinstance(value, dict):
for key, payload_value in value.items():
normalized.setdefault(str(key), payload_value)
return normalized
class MemoryClient:
"""记忆客户端。
提供 AI 记忆的存储和检索能力。
Attributes:
_proxy: CapabilityProxy 实例,用于远程能力调用
"""
def __init__(
self,
proxy: CapabilityProxy,
*,
namespace: str | None = None,
) -> None:
"""初始化记忆客户端。
Args:
proxy: CapabilityProxy 实例
"""
self._proxy = proxy
self._namespace = join_memory_namespace(namespace)
def namespace(self, *parts: Any) -> MemoryClient:
"""Create a derived client that operates inside a child namespace."""
return MemoryClient(
self._proxy,
namespace=join_memory_namespace(self._namespace, *parts),
)
def _resolve_exact_namespace(self, namespace: str | None) -> str:
if namespace is None:
return self._namespace
return join_memory_namespace(self._namespace, namespace)
def _resolve_scope_namespace(self, namespace: str | None) -> tuple[bool, str]:
if namespace is None:
if self._namespace:
return True, self._namespace
return False, ""
return True, join_memory_namespace(self._namespace, namespace)
async def search(
self,
query: str,
*,
mode: Literal["auto", "keyword", "vector", "hybrid"] = "auto",
limit: int | None = None,
min_score: float | None = None,
provider_id: str | None = None,
namespace: str | None = None,
include_descendants: bool = True,
) -> list[dict[str, Any]]:
"""搜索记忆项。
默认会在有 embedding provider 时执行 hybrid 检索,
否则退化为关键词检索。返回结果包含 `score` 与 `match_type` 字段。
Args:
query: 搜索查询文本
mode: 搜索模式,支持 auto/keyword/vector/hybrid
limit: 最大返回条数
min_score: 最低分数阈值
provider_id: 指定 embedding provider,默认使用当前激活的 provider
Returns:
匹配的记忆项列表,按相关度排序
示例:
results = await ctx.memory.search(
"用户喜欢什么颜色",
mode="hybrid",
limit=5,
)
for item in results:
print(item["key"], item["score"], item["match_type"])
"""
payload: dict[str, Any] = {"query": query, "mode": mode}
if limit is not None:
payload["limit"] = limit
if min_score is not None:
payload["min_score"] = min_score
if provider_id is not None:
payload["provider_id"] = provider_id
has_namespace, resolved_namespace = self._resolve_scope_namespace(namespace)
if has_namespace:
payload["namespace"] = resolved_namespace
payload["include_descendants"] = bool(include_descendants)
output = await self._proxy.call("memory.search", payload)
items = output.get("items")
if not isinstance(items, (list, tuple)):
return []
normalized_items: list[dict[str, Any]] = []
for item in items:
normalized = _normalize_search_item(item)
if normalized is not None:
normalized_items.append(normalized)
return normalized_items
async def save(
self,
key: str,
value: dict[str, Any] | None = None,
namespace: str | None = None,
**extra: Any,
) -> None:
"""保存记忆项。
将数据存储到记忆系统,可通过 search() 检索或 get() 精确获取。
Args:
key: 记忆项的唯一标识键
value: 要存储的数据字典
**extra: 额外的键值对,会合并到 value 中
Raises:
TypeError: 如果 value 不是 dict 类型
示例:
保存用户偏好
await ctx.memory.save("user_pref", {"theme": "dark", "lang": "zh"})
使用关键字参数
await ctx.memory.save("note", None, content="重要笔记", tags=["work"])
使用 embedding_text 显式指定检索文本
await ctx.memory.save(
"profile",
{"name": "alice", "embedding_text": "Alice 喜欢蓝色和海边"},
)
"""
if value is not None and not isinstance(value, dict):
raise TypeError("memory.save 的 value 必须是 dict")
payload = dict(value or {})
if extra:
payload.update(extra)
request: dict[str, Any] = {"key": key, "value": payload}
request["namespace"] = self._resolve_exact_namespace(namespace)
await self._proxy.call("memory.save", request)
async def get(
self,
key: str,
*,
namespace: str | None = None,
) -> dict[str, Any] | None:
"""精确获取单个记忆项。
通过唯一键精确获取记忆内容,不经过搜索匹配。
Args:
key: 记忆项的唯一键
Returns:
记忆项内容字典,若不存在则返回 None
示例:
pref = await ctx.memory.get("user_pref")
if pref:
print(f"用户偏好主题: {pref.get('theme')}")
"""
payload: dict[str, Any] = {"key": key}
payload["namespace"] = self._resolve_exact_namespace(namespace)
output = await self._proxy.call("memory.get", payload)
value = output.get("value")
return value if isinstance(value, dict) else None
async def delete(
self,
key: str,
*,
namespace: str | None = None,
) -> None:
"""删除记忆项。
Args:
key: 要删除的记忆项键名
示例:
await ctx.memory.delete("old_note")
"""
payload: dict[str, Any] = {"key": key}
payload["namespace"] = self._resolve_exact_namespace(namespace)
await self._proxy.call("memory.delete", payload)
async def save_with_ttl(
self,
key: str,
value: dict[str, Any],
ttl_seconds: int,
*,
namespace: str | None = None,
) -> None:
"""保存带过期时间的记忆项。
与 save() 不同,此方法允许设置记忆项的存活时间(TTL),
过期后记忆项将自动删除。
Args:
key: 记忆项的唯一标识键
value: 要存储的数据字典
ttl_seconds: 存活时间(秒),必须大于 0
Raises:
TypeError: 如果 value 不是 dict 类型
ValueError: 如果 ttl_seconds 小于 1
示例:
# 保存临时会话状态,1小时后过期
await ctx.memory.save_with_ttl(
"session_temp",
{"state": "waiting"},
ttl_seconds=3600,
)
"""
if not isinstance(value, dict):
raise TypeError("memory.save_with_ttl 的 value 必须是 dict")
if ttl_seconds < 1:
raise ValueError("ttl_seconds 必须大于 0")
payload: dict[str, Any] = {
"key": key,
"value": value,
"ttl_seconds": ttl_seconds,
}
payload["namespace"] = self._resolve_exact_namespace(namespace)
await self._proxy.call("memory.save_with_ttl", payload)
async def get_many(
self,
keys: list[str],
*,
namespace: str | None = None,
) -> list[dict[str, Any]]:
"""批量获取多个记忆项。
一次性获取多个键对应的记忆内容,比多次调用 get() 更高效。
Args:
keys: 记忆项键名列表
Returns:
记忆项列表,每项包含 key 和 value 字段,
不存在的键返回 value 为 None
示例:
items = await ctx.memory.get_many(["pref1", "pref2", "pref3"])
for item in items:
if item["value"]:
print(f"{item['key']}: {item['value']}")
"""
payload: dict[str, Any] = {"keys": keys}
payload["namespace"] = self._resolve_exact_namespace(namespace)
output = await self._proxy.call("memory.get_many", payload)
items = output.get("items")
if not isinstance(items, (list, tuple)):
return []
return [dict(item) for item in items if isinstance(item, dict)]
async def delete_many(
self,
keys: list[str],
*,
namespace: str | None = None,
) -> int:
"""批量删除多个记忆项。
一次性删除多个键对应的记忆项,返回实际删除的数量。
Args:
keys: 要删除的记忆项键名列表
Returns:
实际删除的记忆项数量
示例:
deleted = await ctx.memory.delete_many(["old1", "old2", "old3"])
print(f"删除了 {deleted} 条记忆")
"""
payload: dict[str, Any] = {"keys": keys}
payload["namespace"] = self._resolve_exact_namespace(namespace)
output = await self._proxy.call("memory.delete_many", payload)
return int(output.get("deleted_count", 0))
async def stats(
self,
*,
namespace: str | None = None,
include_descendants: bool = True,
) -> dict[str, Any]:
"""获取记忆系统统计信息。
返回记忆系统的当前状态,包括条目数、索引状态和脏索引数量。
Returns:
统计信息字典,包含:
- total_items: 总记忆条目数
- total_bytes: 总占用字节数(可选)
- ttl_entries: 带过期时间的条目数(可选)
- indexed_items: 已建立检索索引的条目数(可选)
- embedded_items: 已生成向量的条目数(可选)
- dirty_items: 等待重建索引的条目数(可选)
示例:
stats = await ctx.memory.stats()
print(f"记忆库共有 {stats['total_items']} 条记录")
if "embedded_items" in stats:
print(f"其中 {stats['embedded_items']} 条已经向量化")
"""
payload: dict[str, Any] = {
"include_descendants": bool(include_descendants),
}
has_namespace, resolved_namespace = self._resolve_scope_namespace(namespace)
if has_namespace:
payload["namespace"] = resolved_namespace
output = await self._proxy.call("memory.stats", payload)
stats = {
"total_items": output.get("total_items", 0),
"total_bytes": output.get("total_bytes"),
}
if "namespace" in output:
stats["namespace"] = output.get("namespace")
if "namespace_count" in output:
stats["namespace_count"] = output.get("namespace_count")
if "fts_enabled" in output:
stats["fts_enabled"] = output.get("fts_enabled")
if "vector_backend" in output:
stats["vector_backend"] = output.get("vector_backend")
if "vector_indexes" in output:
stats["vector_indexes"] = output.get("vector_indexes")
if "plugin_id" in output:
stats["plugin_id"] = output.get("plugin_id")
if "ttl_entries" in output:
stats["ttl_entries"] = output.get("ttl_entries")
if "indexed_items" in output:
stats["indexed_items"] = output.get("indexed_items")
if "embedded_items" in output:
stats["embedded_items"] = output.get("embedded_items")
if "dirty_items" in output:
stats["dirty_items"] = output.get("dirty_items")
return stats
@@ -1,111 +0,0 @@
"""元数据客户端模块。
提供插件元数据查询能力。
功能说明:
- 查询已加载插件信息
- 获取插件列表
- 访问当前插件配置
安全边界:
插件身份由运行时透传到协议层;客户端只暴露业务参数,不接受外部指定调用者。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from ._proxy import CapabilityProxy
@dataclass
class StarMetadata:
"""插件元数据。"""
name: str
display_name: str
description: str
author: str
version: str
enabled: bool = True
support_platforms: list[str] = field(default_factory=list)
astrbot_version: str | None = None
@classmethod
def from_dict(cls, data: dict[str, Any]) -> StarMetadata:
raw_support_platforms = data.get("support_platforms")
support_platforms = (
[str(item) for item in raw_support_platforms if isinstance(item, str)]
if isinstance(raw_support_platforms, list)
else []
)
return cls(
name=str(data.get("name", "")),
display_name=str(data.get("display_name", data.get("name", ""))),
description=str(data.get("desc", data.get("description", ""))),
author=str(data.get("author", "")),
version=str(data.get("version", "0.0.0")),
enabled=bool(data.get("enabled", True)),
support_platforms=support_platforms,
astrbot_version=(
str(data.get("astrbot_version"))
if data.get("astrbot_version") is not None
else None
),
)
PluginMetadata = StarMetadata
class MetadataClient:
"""元数据能力客户端。"""
def __init__(self, proxy: CapabilityProxy, plugin_id: str) -> None:
self._proxy = proxy
self._plugin_id = plugin_id
async def get_plugin(self, name: str) -> StarMetadata | None:
output = await self._proxy.call(
"metadata.get_plugin",
{"name": name},
)
data = output.get("plugin")
if data is None:
return None
return StarMetadata.from_dict(data)
async def list_plugins(self) -> list[StarMetadata]:
output = await self._proxy.call("metadata.list_plugins", {})
items = output.get("plugins", [])
return [
StarMetadata.from_dict(item) for item in items if isinstance(item, dict)
]
async def get_current_plugin(self) -> StarMetadata | None:
return await self.get_plugin(self._plugin_id)
async def get_plugin_config(self, name: str | None = None) -> dict[str, Any] | None:
target = name or self._plugin_id
if target != self._plugin_id:
raise PermissionError(
"get_plugin_config 只允许访问当前插件自己的配置,"
f"请求的插件 '{target}' 不是当前插件 '{self._plugin_id}'"
)
output = await self._proxy.call(
"metadata.get_plugin_config",
{"name": target},
)
config = output.get("config")
return dict(config) if isinstance(config, dict) else None
async def save_plugin_config(self, config: dict[str, Any]) -> dict[str, Any]:
if not isinstance(config, dict):
raise TypeError("save_plugin_config requires a dict payload")
output = await self._proxy.call(
"metadata.save_plugin_config",
{"config": dict(config)},
)
saved = output.get("config")
return dict(saved) if isinstance(saved, dict) else {}
@@ -1,300 +0,0 @@
"""平台客户端模块。
提供 v4 原生的平台能力调用。
设计边界:
- `PlatformClient` 只负责直接的平台 capability
- 迁移期消息桥接由独立迁移入口承接,不放进原生客户端
- 富消息链通过 `platform.send_chain` 发送,链构建能力位于专门的消息模块
"""
from __future__ import annotations
from collections.abc import Sequence
from enum import Enum
from typing import Any, cast
from pydantic import BaseModel, ConfigDict, Field
from ..message.components import BaseMessageComponent, Plain
from ..message.result import MessageChain
from ..message.session import MessageSession
from ..protocol.descriptors import SessionRef
from ._proxy import CapabilityProxy
class _PlatformModel(BaseModel):
model_config = ConfigDict(extra="forbid")
class PlatformStatus(str, Enum):
PENDING = "pending"
RUNNING = "running"
ERROR = "error"
STOPPED = "stopped"
@classmethod
def from_value(cls, value: Any) -> PlatformStatus:
if isinstance(value, cls):
return value
try:
return cls(str(value).strip().lower())
except ValueError:
return cls.PENDING
class PlatformError(_PlatformModel):
message: str
timestamp: str
traceback: str | None = None
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> PlatformError | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class PlatformStats(_PlatformModel):
id: str
type: str
display_name: str
status: PlatformStatus
started_at: str | None = None
error_count: int
last_error: PlatformError | None = None
unified_webhook: bool
meta: dict[str, Any] = Field(default_factory=dict)
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> PlatformStats | None:
if not isinstance(payload, dict):
return None
normalized = dict(payload)
normalized["status"] = PlatformStatus.from_value(payload.get("status"))
normalized["last_error"] = PlatformError.from_payload(
payload.get("last_error") if isinstance(payload, dict) else None
)
meta = payload.get("meta")
normalized["meta"] = dict(meta) if isinstance(meta, dict) else {}
return cls.model_validate(normalized)
class PlatformClient:
"""平台消息客户端。
提供向聊天平台发送消息和获取信息的能力。
Attributes:
_proxy: CapabilityProxy 实例,用于远程能力调用
"""
def __init__(self, proxy: CapabilityProxy) -> None:
"""初始化平台客户端。
Args:
proxy: CapabilityProxy 实例
"""
self._proxy = proxy
def _build_target_payload(
self,
session: str | SessionRef | MessageSession,
) -> tuple[str, dict[str, Any]]:
if isinstance(session, SessionRef):
return session.session, {"target": session.to_payload()}
if isinstance(session, MessageSession):
return str(session), {}
return str(session), {}
async def _coerce_chain_payload(
self,
content: (
str
| MessageChain
| Sequence[BaseMessageComponent]
| Sequence[dict[str, Any]]
),
) -> list[dict[str, Any]]:
if isinstance(content, str):
return await MessageChain(
[Plain(content, convert=False)]
).to_payload_async()
if isinstance(content, MessageChain):
return await content.to_payload_async()
if (
isinstance(content, Sequence)
and not isinstance(content, (str, bytes))
and all(isinstance(item, BaseMessageComponent) for item in content)
):
components = cast(Sequence[BaseMessageComponent], content)
return await MessageChain(list(components)).to_payload_async()
if (
isinstance(content, Sequence)
and not isinstance(content, (str, bytes))
and all(isinstance(item, dict) for item in content)
):
payload_items = cast(Sequence[dict[str, Any]], content)
return [dict(item) for item in payload_items]
raise TypeError(
"content must be str, MessageChain, sequence of message components, "
"or sequence of platform.send_chain payload dicts"
)
async def send(
self,
session: str | SessionRef | MessageSession,
text: str,
) -> dict[str, Any]:
"""发送文本消息。
向指定的会话(用户或群组)发送文本消息。
Args:
session: 统一消息来源标识 (UMO),格式如 "platform:instance:user_id"
text: 要发送的文本内容
Returns:
发送结果,可能包含消息 ID 等信息
示例:
# 发送消息到当前会话
await ctx.platform.send(event.session_id, "收到您的消息!")
"""
session_id, extra = self._build_target_payload(session)
return await self._proxy.call(
"platform.send",
{"session": session_id, "text": text, **extra},
)
async def send_image(
self,
session: str | SessionRef | MessageSession,
image_url: str,
) -> dict[str, Any]:
"""发送图片消息。
向指定的会话发送图片,支持 URL 或本地路径。
Args:
session: 统一消息来源标识 (UMO)
image_url: 图片 URL 或本地文件路径
Returns:
发送结果
示例:
await ctx.platform.send_image(
event.session_id,
"https://example.com/image.png"
)
"""
session_id, extra = self._build_target_payload(session)
return await self._proxy.call(
"platform.send_image",
{"session": session_id, "image_url": image_url, **extra},
)
async def send_chain(
self,
session: str | SessionRef | MessageSession,
chain: MessageChain | Sequence[BaseMessageComponent] | Sequence[dict[str, Any]],
) -> dict[str, Any]:
"""发送富消息链。
Args:
session: 统一消息来源标识 (UMO)
chain: 序列化后的消息组件数组
Returns:
发送结果
"""
session_id, extra = self._build_target_payload(session)
chain_payload = await self._coerce_chain_payload(chain)
return await self._proxy.call(
"platform.send_chain",
{"session": session_id, "chain": chain_payload, **extra},
)
async def send_by_session(
self,
session: str | MessageSession,
content: (
str
| MessageChain
| Sequence[BaseMessageComponent]
| Sequence[dict[str, Any]]
),
) -> dict[str, Any]:
"""主动向指定会话发送消息链。
`Sequence[dict]` 的结构与 `platform.send_chain` 完全一致:
每一项都应是 `{"type": "...", "data": {...}}`。
"""
chain_payload = await self._coerce_chain_payload(content)
session_id = str(session)
return await self._proxy.call(
"platform.send_by_session",
{"session": session_id, "chain": chain_payload},
)
async def send_by_id(
self,
platform_id: str,
session_id: str,
content: (
str
| MessageChain
| Sequence[BaseMessageComponent]
| Sequence[dict[str, Any]]
),
*,
message_type: str = "private",
) -> dict[str, Any]:
"""主动向指定平台会话发送消息。"""
session = MessageSession(
platform_id=str(platform_id),
message_type=str(message_type),
session_id=str(session_id),
)
return await self.send_by_session(session, content)
async def get_members(
self,
session: str | SessionRef | MessageSession,
) -> list[dict[str, Any]]:
"""获取群组成员列表。
获取指定群组的成员信息列表。注意仅对群组会话有效。
Args:
session: 群组会话的统一消息来源标识 (UMO)
Returns:
成员信息列表,每个成员是一个字典,可能包含:
- user_id: 用户 ID
- nickname: 昵称
- role: 角色 (owner, admin, member)
示例:
members = await ctx.platform.get_members(event.session_id)
for member in members:
print(f"{member['nickname']} ({member['user_id']})")
"""
session_id, extra = self._build_target_payload(session)
output = await self._proxy.call(
"platform.get_members",
{"session": session_id, **extra},
)
members = output.get("members")
if not isinstance(members, (list, tuple)):
return []
return list(members)
__all__ = [
"PlatformClient",
"PlatformError",
"PlatformStats",
"PlatformStatus",
]
@@ -1,349 +0,0 @@
"""Provider discovery and provider-management clients."""
from __future__ import annotations
import asyncio
import contextlib
import inspect
from collections.abc import AsyncIterator, Awaitable, Callable
from typing import Any
from pydantic import BaseModel, ConfigDict
from ..llm.entities import ProviderMeta, ProviderType
from ..llm.providers import (
ProviderProxy,
STTProvider,
TTSProvider,
provider_proxy_from_meta,
)
from ._proxy import CapabilityProxy
class _ProviderModel(BaseModel):
model_config = ConfigDict(extra="forbid")
def to_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_none=True)
class ManagedProviderRecord(_ProviderModel):
id: str
model: str | None = None
type: str
provider_type: ProviderType
loaded: bool
enabled: bool
provider_source_id: str | None = None
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> ManagedProviderRecord | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class ProviderChangeEvent(_ProviderModel):
provider_id: str
provider_type: ProviderType
umo: str | None = None
@classmethod
def from_payload(
cls,
payload: dict[str, Any] | None,
) -> ProviderChangeEvent | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class ProviderClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
@staticmethod
def _provider_meta_list(items: Any) -> list[ProviderMeta]:
if not isinstance(items, list):
return []
providers: list[ProviderMeta] = []
for item in items:
if not isinstance(item, dict):
continue
provider = ProviderMeta.from_payload(item)
if provider is not None:
providers.append(provider)
return providers
async def list_all(self) -> list[ProviderMeta]:
output = await self._proxy.call("provider.list_all", {})
return self._provider_meta_list(output.get("providers"))
async def list_tts(self) -> list[ProviderMeta]:
output = await self._proxy.call("provider.list_all_tts", {})
return self._provider_meta_list(output.get("providers"))
async def list_stt(self) -> list[ProviderMeta]:
output = await self._proxy.call("provider.list_all_stt", {})
return self._provider_meta_list(output.get("providers"))
async def list_embedding(self) -> list[ProviderMeta]:
output = await self._proxy.call("provider.list_all_embedding", {})
return self._provider_meta_list(output.get("providers"))
async def list_rerank(self) -> list[ProviderMeta]:
output = await self._proxy.call("provider.list_all_rerank", {})
return self._provider_meta_list(output.get("providers"))
async def _get_tts_support_stream(self, provider_id: str) -> bool:
output = await self._proxy.call(
"provider.tts.support_stream",
{"provider_id": str(provider_id)},
)
return bool(output.get("supported", False))
async def _build_proxy(self, meta: ProviderMeta | None) -> ProviderProxy | None:
if meta is None:
return None
tts_supports_stream = None
if meta.provider_type == ProviderType.TEXT_TO_SPEECH:
tts_supports_stream = await self._get_tts_support_stream(meta.id)
return provider_proxy_from_meta(
self._proxy,
meta,
tts_supports_stream=tts_supports_stream,
)
async def get(self, provider_id: str) -> ProviderProxy | None:
output = await self._proxy.call(
"provider.get_by_id",
{"provider_id": str(provider_id)},
)
return await self._build_proxy(
ProviderMeta.from_payload(output.get("provider"))
)
async def get_using_chat(self, umo: str | None = None) -> ProviderMeta | None:
output = await self._proxy.call("provider.get_using", {"umo": umo})
return ProviderMeta.from_payload(output.get("provider"))
async def get_using_tts(self, umo: str | None = None) -> TTSProvider | None:
output = await self._proxy.call("provider.get_using_tts", {"umo": umo})
provider = await self._build_proxy(
ProviderMeta.from_payload(output.get("provider"))
)
return provider if isinstance(provider, TTSProvider) else None
async def get_using_stt(self, umo: str | None = None) -> STTProvider | None:
output = await self._proxy.call("provider.get_using_stt", {"umo": umo})
provider = await self._build_proxy(
ProviderMeta.from_payload(output.get("provider"))
)
return provider if isinstance(provider, STTProvider) else None
class ProviderManagerClient:
def __init__(
self,
proxy: CapabilityProxy,
*,
plugin_id: str | None = None,
logger: Any | None = None,
) -> None:
self._proxy = proxy
self._plugin_id = plugin_id
self._logger = logger
self._change_hook_tasks: set[asyncio.Task[None]] = set()
@staticmethod
def _provider_type_value(provider_type: ProviderType | str) -> str:
if isinstance(provider_type, ProviderType):
return provider_type.value
return str(provider_type).strip()
@staticmethod
def _record_from_output(output: dict[str, Any]) -> ManagedProviderRecord | None:
return ManagedProviderRecord.from_payload(output.get("provider"))
async def set_provider(
self,
provider_id: str,
provider_type: ProviderType | str,
umo: str | None = None,
) -> None:
await self._proxy.call(
"provider.manager.set",
{
"provider_id": str(provider_id),
"provider_type": self._provider_type_value(provider_type),
"umo": umo,
},
)
async def get_provider_by_id(
self,
provider_id: str,
) -> ManagedProviderRecord | None:
output = await self._proxy.call(
"provider.manager.get_by_id",
{"provider_id": str(provider_id)},
)
return self._record_from_output(output)
async def get_merged_provider_config(
self,
provider_id: str,
) -> dict[str, Any] | None:
output = await self._proxy.call(
"provider.manager.get_merged_provider_config",
{"provider_id": str(provider_id).strip()},
)
config = output.get("config")
return dict(config) if isinstance(config, dict) else None
async def load_provider(
self,
provider_config: dict[str, Any],
) -> ManagedProviderRecord | None:
output = await self._proxy.call(
"provider.manager.load",
{"provider_config": dict(provider_config)},
)
return self._record_from_output(output)
async def terminate_provider(self, provider_id: str) -> None:
await self._proxy.call(
"provider.manager.terminate",
{"provider_id": str(provider_id)},
)
async def create_provider(
self,
provider_config: dict[str, Any],
) -> ManagedProviderRecord | None:
output = await self._proxy.call(
"provider.manager.create",
{"provider_config": dict(provider_config)},
)
return self._record_from_output(output)
async def update_provider(
self,
origin_provider_id: str,
new_config: dict[str, Any],
) -> ManagedProviderRecord | None:
output = await self._proxy.call(
"provider.manager.update",
{
"origin_provider_id": str(origin_provider_id),
"new_config": dict(new_config),
},
)
return self._record_from_output(output)
async def delete_provider(
self,
provider_id: str | None = None,
provider_source_id: str | None = None,
) -> None:
await self._proxy.call(
"provider.manager.delete",
{
"provider_id": provider_id,
"provider_source_id": provider_source_id,
},
)
async def get_insts(self) -> list[ManagedProviderRecord]:
output = await self._proxy.call("provider.manager.get_insts", {})
items = output.get("providers")
if not isinstance(items, list):
return []
return [
record
for record in (
ManagedProviderRecord.from_payload(item)
if isinstance(item, dict)
else None
for item in items
)
if record is not None
]
async def watch_changes(self) -> AsyncIterator[ProviderChangeEvent]:
async for chunk in self._proxy.stream("provider.manager.watch_changes", {}):
event = ProviderChangeEvent.from_payload(chunk)
if event is not None:
yield event
async def register_provider_change_hook(
self,
callback: Callable[
[str, ProviderType, str | None],
Awaitable[None] | None,
],
) -> asyncio.Task[None]:
async def runner() -> None:
async for event in self.watch_changes():
result = callback(
event.provider_id,
event.provider_type,
event.umo,
)
if inspect.isawaitable(result):
await result
task = asyncio.create_task(runner())
self._change_hook_tasks.add(task)
task.add_done_callback(self._log_change_hook_result)
return task
async def unregister_provider_change_hook(
self,
task: asyncio.Task[None],
) -> None:
if task not in self._change_hook_tasks:
return
self._change_hook_tasks.discard(task)
if not task.done():
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
def _log_change_hook_result(self, task: asyncio.Task[None]) -> None:
self._change_hook_tasks.discard(task)
if task.cancelled():
debug_logger = getattr(self._logger, "debug", None)
if callable(debug_logger):
debug_logger(
"Provider change hook cancelled: plugin_id={}",
self._plugin_id,
)
return
try:
task.result()
except asyncio.CancelledError:
debug_logger = getattr(self._logger, "debug", None)
if callable(debug_logger):
debug_logger(
"Provider change hook cancelled: plugin_id={}",
self._plugin_id,
)
except Exception:
exception_logger = getattr(self._logger, "exception", None)
if callable(exception_logger):
exception_logger(
"Provider change hook failed: plugin_id={}",
self._plugin_id,
)
__all__ = [
"ManagedProviderRecord",
"ProviderChangeEvent",
"ProviderClient",
"ProviderManagerClient",
]
@@ -1,120 +0,0 @@
"""只读 handler 注册表客户端。"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from ._proxy import CapabilityProxy
def _coerce_int(value: Any, default: int = 0) -> int:
try:
return int(value)
except (TypeError, ValueError):
return default
@dataclass(slots=True)
class HandlerMetadata:
plugin_name: str
handler_full_name: str
trigger_type: str
description: str | None = None
event_types: list[str] = field(default_factory=list)
enabled: bool = True
group_path: list[str] = field(default_factory=list)
priority: int = 0
kind: str = "handler"
require_admin: bool = False
@classmethod
def from_dict(cls, data: dict[str, Any]) -> HandlerMetadata:
return cls(
plugin_name=str(data.get("plugin_name", "")),
handler_full_name=str(data.get("handler_full_name", "")),
trigger_type=str(data.get("trigger_type", "")),
description=(
None
if data.get("description") is None
else str(data.get("description", "")).strip() or None
),
event_types=[
str(item)
for item in data.get("event_types", [])
if isinstance(item, str)
],
enabled=bool(data.get("enabled", True)),
group_path=[
str(item)
for item in data.get("group_path", [])
if isinstance(item, str)
],
priority=_coerce_int(data.get("priority", 0), 0),
kind=str(data.get("kind", "handler") or "handler"),
require_admin=bool(data.get("require_admin", False)),
)
class RegistryClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def get_handlers_by_event_type(
self,
event_type: str,
) -> list[HandlerMetadata]:
output = await self._proxy.call(
"registry.get_handlers_by_event_type",
{"event_type": event_type},
)
return [
HandlerMetadata.from_dict(item)
for item in output.get("handlers", [])
if isinstance(item, dict)
]
async def get_handler_by_full_name(
self,
full_name: str,
) -> HandlerMetadata | None:
output = await self._proxy.call(
"registry.get_handler_by_full_name",
{"full_name": full_name},
)
handler = output.get("handler")
if not isinstance(handler, dict):
return None
return HandlerMetadata.from_dict(handler)
async def set_handler_whitelist(
self,
plugin_names: list[str] | set[str] | None,
) -> list[str] | None:
names = None
if plugin_names is not None:
names = sorted({str(item) for item in plugin_names if str(item).strip()})
output = await self._proxy.call(
"system.event.handler_whitelist.set",
{"plugin_names": names},
)
result = output.get("plugin_names")
if not isinstance(result, list):
return None
return [str(item) for item in result]
async def get_handler_whitelist(self) -> list[str] | None:
output = await self._proxy.call("system.event.handler_whitelist.get", {})
result = output.get("plugin_names")
if not isinstance(result, list):
return None
return [str(item) for item in result]
async def clear_handler_whitelist(self) -> None:
await self._proxy.call(
"system.event.handler_whitelist.set",
{"plugin_names": None},
)
__all__ = ["HandlerMetadata", "RegistryClient"]
@@ -1,135 +0,0 @@
"""Session-scoped SDK managers."""
from __future__ import annotations
from typing import Any
from ..events import MessageEvent
from ..message.session import MessageSession
from ._proxy import CapabilityProxy
from .registry import HandlerMetadata
def _normalize_session(session: str | MessageSession | MessageEvent) -> str:
if isinstance(session, MessageEvent):
return str(session.unified_msg_origin)
if isinstance(session, MessageSession):
return str(session)
return str(session)
def _handler_to_payload(handler: HandlerMetadata) -> dict[str, Any]:
return {
"plugin_name": handler.plugin_name,
"handler_full_name": handler.handler_full_name,
"trigger_type": handler.trigger_type,
"description": handler.description,
"event_types": list(handler.event_types),
"enabled": handler.enabled,
"group_path": list(handler.group_path),
"priority": handler.priority,
"kind": handler.kind,
"require_admin": handler.require_admin,
}
class SessionPluginManager:
"""Session-scoped plugin status manager."""
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def is_plugin_enabled_for_session(
self,
session: str | MessageSession | MessageEvent,
plugin_name: str,
) -> bool:
output = await self._proxy.call(
"session.plugin.is_enabled",
{
"session": _normalize_session(session),
"plugin_name": str(plugin_name),
},
)
return bool(output.get("enabled", False))
async def filter_handlers_by_session(
self,
session: str | MessageSession | MessageEvent,
handlers: list[HandlerMetadata],
) -> list[HandlerMetadata]:
output = await self._proxy.call(
"session.plugin.filter_handlers",
{
"session": _normalize_session(session),
"handlers": [_handler_to_payload(handler) for handler in handlers],
},
)
items = output.get("handlers")
if not isinstance(items, list):
return []
return [
HandlerMetadata.from_dict(item) for item in items if isinstance(item, dict)
]
class SessionServiceManager:
"""Session-scoped LLM/TTS service status manager."""
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def is_llm_enabled_for_session(
self,
session: str | MessageSession | MessageEvent,
) -> bool:
output = await self._proxy.call(
"session.service.is_llm_enabled",
{"session": _normalize_session(session)},
)
return bool(output.get("enabled", False))
async def set_llm_status_for_session(
self,
session: str | MessageSession | MessageEvent,
enabled: bool,
) -> None:
await self._proxy.call(
"session.service.set_llm_status",
{"session": _normalize_session(session), "enabled": bool(enabled)},
)
async def should_process_llm_request(
self,
event_or_session: str | MessageSession | MessageEvent,
) -> bool:
return await self.is_llm_enabled_for_session(event_or_session)
async def is_tts_enabled_for_session(
self,
session: str | MessageSession | MessageEvent,
) -> bool:
output = await self._proxy.call(
"session.service.is_tts_enabled",
{"session": _normalize_session(session)},
)
return bool(output.get("enabled", False))
async def set_tts_status_for_session(
self,
session: str | MessageSession | MessageEvent,
enabled: bool,
) -> None:
await self._proxy.call(
"session.service.set_tts_status",
{"session": _normalize_session(session), "enabled": bool(enabled)},
)
async def should_process_tts_request(
self,
event_or_session: str | MessageSession | MessageEvent,
) -> bool:
return await self.is_tts_enabled_for_session(event_or_session)
__all__ = ["SessionPluginManager", "SessionServiceManager"]
@@ -1,60 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from ._proxy import CapabilityProxy
@dataclass(slots=True)
class SkillRegistration:
name: str
description: str
path: str
skill_dir: str
@classmethod
def from_dict(cls, data: dict[str, Any]) -> SkillRegistration:
return cls(
name=str(data.get("name", "")),
description=str(data.get("description", "") or ""),
path=str(data.get("path", "")),
skill_dir=str(data.get("skill_dir", "")),
)
class SkillClient:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def register(
self,
*,
name: str,
path: str,
description: str = "",
) -> SkillRegistration:
output = await self._proxy.call(
"skill.register",
{
"name": name,
"path": path,
"description": description,
},
)
return SkillRegistration.from_dict(output)
async def unregister(self, name: str) -> bool:
output = await self._proxy.call("skill.unregister", {"name": name})
return bool(output.get("removed", False))
async def list(self) -> list[SkillRegistration]:
output = await self._proxy.call("skill.list", {})
return [
SkillRegistration.from_dict(item)
for item in output.get("skills", [])
if isinstance(item, dict)
]
__all__ = ["SkillClient", "SkillRegistration"]
-159
View File
@@ -1,159 +0,0 @@
"""SDK-native command group helpers.
本模块提供命令分组工具,用于组织具有层级关系的命令。
CommandGroup 允许以嵌套方式定义命令树,例如:
admin
├── user
│ ├── add
│ └── remove
└── config
├── get
└── set
特性:
- 支持命令别名,自动展开父级路径的所有别名组合
- 自动生成命令树的可视化输出 (print_cmd_tree)
- 与 @on_command 装饰器无缝集成
"""
from __future__ import annotations
from dataclasses import dataclass, field
from itertools import product
from .decorators import on_command, set_command_route_meta
from .protocol.descriptors import CommandRouteSpec
@dataclass(slots=True)
class _CommandNode:
name: str
aliases: list[str] = field(default_factory=list)
description: str | None = None
subgroups: list[CommandGroup] = field(default_factory=list)
commands: list[tuple[str, str | None]] = field(default_factory=list)
class CommandGroup:
def __init__(
self,
name: str,
*,
aliases: list[str] | None = None,
description: str | None = None,
parent: CommandGroup | None = None,
) -> None:
self.name = name
self.aliases = list(aliases or [])
self.description = description
self.parent = parent
self._tree = _CommandNode(
name=name, aliases=self.aliases, description=description
)
def group(
self,
name: str,
*,
aliases: list[str] | None = None,
description: str | None = None,
) -> CommandGroup:
child = CommandGroup(
name,
aliases=aliases,
description=description,
parent=self,
)
self._tree.subgroups.append(child)
return child
def command(
self,
name: str,
*,
aliases: list[str] | None = None,
description: str | None = None,
):
full_command = " ".join([*self.path, name])
full_aliases = self._expand_aliases(name=name, aliases=aliases or [])
display_command = full_command
route = CommandRouteSpec(
group_path=self.path,
display_command=display_command,
group_help=self.description,
)
def decorator(func):
decorated = on_command(
full_command,
aliases=full_aliases,
description=description,
)(func)
self._tree.commands.append((name, description))
set_command_route_meta(decorated, route)
return decorated
return decorator
@property
def path(self) -> list[str]:
if self.parent is None:
return [self.name]
return [*self.parent.path, self.name]
def print_cmd_tree(self) -> str:
lines: list[str] = []
self._append_tree_lines(lines, indent=0)
return "\n".join(lines)
def _append_tree_lines(self, lines: list[str], *, indent: int) -> None:
prefix = " " * indent
label = self.name
if self.aliases:
label += f" ({', '.join(self.aliases)})"
lines.append(f"{prefix}{label}")
for command_name, description in self._tree.commands:
command_label = f"{prefix} - {command_name}"
if description:
command_label += f": {description}"
lines.append(command_label)
for subgroup in self._tree.subgroups:
subgroup._append_tree_lines(lines, indent=indent + 1)
def _expand_aliases(self, *, name: str, aliases: list[str]) -> list[str]:
group_segments: list[list[str]] = []
cursor: CommandGroup | None = self
ancestry: list[CommandGroup] = []
while cursor is not None:
ancestry.append(cursor)
cursor = cursor.parent
for group in reversed(ancestry):
group_segments.append([group.name, *group.aliases])
leaf_segments = [name, *aliases]
expanded: set[str] = set()
for parts in product(*group_segments, leaf_segments):
route = " ".join(parts)
if route != " ".join([*self.path, name]):
expanded.add(route)
return sorted(expanded)
def command_group(
name: str,
*,
aliases: list[str] | None = None,
description: str | None = None,
) -> CommandGroup:
return CommandGroup(
name,
aliases=aliases,
description=description,
)
def print_cmd_tree(group: CommandGroup) -> str:
return group.print_cmd_tree()
__all__ = ["CommandGroup", "command_group", "print_cmd_tree"]
-729
View File
@@ -1,729 +0,0 @@
"""v4 原生运行时上下文。
`Context` 是插件与 AstrBot Core 交互的主要入口,
负责组合所有 capability 客户端并提供统一的访问接口。
每个 handler 调用都会创建一个新的 Context 实例,
绑定到当前的 Peer、插件 ID 和取消令牌。
Attributes:
llm: LLM 能力客户端,用于 AI 对话
memory: 记忆能力客户端,用于语义存储
db: 数据库客户端,用于 KV 持久化
files: 文件服务客户端,用于文件令牌注册与解析
platform: 平台客户端,用于发送消息
providers: Provider 客户端,用于查询和调用专用 Provider
provider_manager: Provider 管理客户端,用于 reserved/system 级操作
personas: 人格管理客户端
conversations: 对话管理客户端
kbs: 知识库管理客户端
message_history: 消息历史管理客户端
http: HTTP 客户端,用于注册 API 端点
metadata: 元数据客户端,用于查询插件信息
mcp: MCP 管理客户端,用于本地/全局 MCP 服务管理
skills: Skill 客户端,用于向 AstrBot 注册插件技能
plugin_id: 当前插件的唯一标识
logger: 绑定了插件 ID 的日志器
cancel_token: 取消令牌,用于处理请求取消
"""
from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from loguru import logger as base_logger
from ._internal.plugin_logger import PluginLogger
from ._internal.star_runtime import current_star_instance
from .clients import (
DBClient,
HTTPClient,
LLMClient,
MCPManagerClient,
MemoryClient,
MetadataClient,
PlatformClient,
PlatformError,
PlatformStats,
PlatformStatus,
RegistryClient,
SkillClient,
)
from .clients._proxy import CapabilityProxy
from .clients.files import FileServiceClient
from .clients.llm import LLMResponse
from .clients.managers import (
ConversationManagerClient,
KnowledgeBaseManagerClient,
MessageHistoryManagerClient,
PersonaManagerClient,
)
from .clients.provider import ProviderClient, ProviderManagerClient
from .clients.session import SessionPluginManager, SessionServiceManager
from .clients.skills import SkillRegistration
from .errors import AstrBotError
from .llm.entities import LLMToolSpec, ProviderMeta, ProviderRequest
from .llm.tools import LLMToolManager
from .message.components import BaseMessageComponent
from .message.result import MessageChain
from .message.session import MessageSession
from .session_waiter import (
_mark_session_waiter_background_task,
_unmark_session_waiter_background_task,
)
PlatformCompatContent = (
str | MessageChain | Sequence[BaseMessageComponent] | Sequence[dict[str, Any]]
)
@dataclass(slots=True)
class PlatformCompatFacade:
"""兼容层平台入口,仅暴露安全元信息和主动发送能力。"""
_ctx: Context
id: str
name: str
type: str
status: PlatformStatus = PlatformStatus.PENDING
errors: list[PlatformError] = field(default_factory=list)
last_error: PlatformError | None = None
unified_webhook: bool = False
_state_lock: asyncio.Lock = field(default_factory=asyncio.Lock, repr=False)
async def send_by_session(
self,
session: str | MessageSession,
content: PlatformCompatContent,
) -> dict[str, Any]:
return await self._ctx.platform.send_by_session(session, content)
async def send_by_id(
self,
session_id: str,
content: PlatformCompatContent,
*,
message_type: str = "private",
) -> dict[str, Any]:
return await self._ctx.platform.send_by_id(
self.id,
session_id,
content,
message_type=message_type,
)
async def send(
self,
session: str | MessageSession,
content: PlatformCompatContent,
*,
message_type: str = "private",
) -> dict[str, Any]:
if isinstance(session, MessageSession):
return await self.send_by_session(session, content)
session_text = str(session).strip()
if ":" in session_text:
return await self.send_by_session(session_text, content)
return await self.send_by_id(
session_text,
content,
message_type=message_type,
)
async def refresh(self) -> None:
async with self._state_lock:
await self._refresh_locked()
async def clear_errors(self) -> None:
async with self._state_lock:
await self._ctx._proxy.call(
"platform.manager.clear_errors",
{"platform_id": self.id},
)
await self._refresh_locked()
async def get_stats(self) -> PlatformStats | None:
output = await self._ctx._proxy.call(
"platform.manager.get_stats",
{"platform_id": self.id},
)
return PlatformStats.from_payload(output.get("stats"))
def _apply_snapshot(self, payload: Any) -> None:
if not isinstance(payload, dict):
return
self.name = str(payload.get("name", self.name))
self.type = str(payload.get("type", self.type))
self.status = PlatformStatus.from_value(payload.get("status"))
errors_payload = payload.get("errors")
if isinstance(errors_payload, list):
self.errors = [
error
for error in (
PlatformError.from_payload(item) if isinstance(item, dict) else None
for item in errors_payload
)
if error is not None
]
self.last_error = PlatformError.from_payload(payload.get("last_error"))
self.unified_webhook = bool(payload.get("unified_webhook", False))
async def _refresh_locked(self) -> None:
output = await self._ctx._proxy.call(
"platform.manager.get_by_id",
{"platform_id": self.id},
)
self._apply_snapshot(output.get("platform"))
@dataclass(slots=True)
class CancelToken:
"""请求取消令牌。
用于协调长时间运行操作的取消。当用户取消请求或
上游超时时,令牌会被触发,允许 handler 及时清理资源。
Example:
async def long_operation(ctx: Context):
for item in large_list:
ctx.cancel_token.raise_if_cancelled()
await process(item)
"""
_cancelled: asyncio.Event
def __init__(self) -> None:
self._cancelled = asyncio.Event()
def cancel(self) -> None:
"""触发取消信号。"""
self._cancelled.set()
@property
def cancelled(self) -> bool:
"""检查是否已被取消。"""
return self._cancelled.is_set()
async def wait(self) -> None:
"""等待取消信号。"""
await self._cancelled.wait()
def raise_if_cancelled(self) -> None:
"""如果已取消则抛出 CancelledError。
Raises:
asyncio.CancelledError: 如果令牌已被取消
"""
if self.cancelled:
raise asyncio.CancelledError
class Context:
"""插件运行时上下文。
组合所有 capability 客户端,提供统一的访问接口。
每个 handler 调用都会创建新的 Context 实例。
Attributes:
peer: 协议对等端,用于底层通信
llm: LLM 客户端
memory: 记忆客户端
db: 数据库客户端
platform: 平台客户端
providers: Provider 客户端
provider_manager: Provider 管理客户端
personas: 人格管理客户端
conversations: 对话管理客户端
kbs: 知识库管理客户端
message_history: 消息历史管理客户端
http: HTTP 客户端
metadata: 元数据客户端
mcp: MCP 管理客户端
plugin_id: 当前插件 ID
logger: 日志器
cancel_token: 取消令牌
"""
def __init__(
self,
*,
peer,
plugin_id: str,
request_id: str | None = None,
cancel_token: CancelToken | None = None,
logger: Any | None = None,
source_event_payload: dict[str, Any] | None = None,
) -> None:
"""初始化上下文。
Args:
peer: 协议对等端实例
plugin_id: 当前插件 ID
cancel_token: 取消令牌,None 时创建新令牌
logger: 日志器,None 时使用默认 logger 并绑定 plugin_id
"""
proxy = CapabilityProxy(
peer,
caller_plugin_id=plugin_id,
request_scope_id=request_id,
)
if isinstance(logger, PluginLogger):
bound_logger = logger
else:
bound_logger = logger or base_logger.bind(plugin_id=plugin_id)
self._proxy = proxy
self.peer = peer
self.llm = LLMClient(proxy)
self.memory = MemoryClient(proxy)
self.db = DBClient(proxy)
self.files = FileServiceClient(proxy)
self.platform = PlatformClient(proxy)
self.providers = ProviderClient(proxy)
self.provider_manager = ProviderManagerClient(
proxy,
plugin_id=plugin_id,
logger=bound_logger,
)
self.personas = PersonaManagerClient(proxy)
self.conversations = ConversationManagerClient(proxy)
self.kbs = KnowledgeBaseManagerClient(proxy)
self.message_history = MessageHistoryManagerClient(proxy)
self.http = HTTPClient(proxy)
self.metadata = MetadataClient(proxy, plugin_id)
self.mcp = MCPManagerClient(proxy)
self.registry = RegistryClient(proxy)
self.skills = SkillClient(proxy)
self.session_plugins = SessionPluginManager(proxy)
self.session_services = SessionServiceManager(proxy)
self.persona_manager = self.personas
self.conversation_manager = self.conversations
self.kb_manager = self.kbs
self.message_history_manager = self.message_history
self.mcp_manager = self.mcp
self._llm_tool_manager = LLMToolManager(proxy)
self.plugin_id = plugin_id
self.logger: PluginLogger = (
bound_logger
if isinstance(bound_logger, PluginLogger)
else PluginLogger(plugin_id=plugin_id, logger=bound_logger)
)
self.cancel_token = cancel_token or CancelToken()
self.request_id = request_id
self._source_event_payload = (
dict(source_event_payload) if isinstance(source_event_payload, dict) else {}
)
async def get_data_dir(self) -> Path:
"""Return the plugin-scoped data directory path."""
output = await self._proxy.call("system.get_data_dir", {})
return Path(str(output.get("path", "")))
async def _register_file_url(
self,
path: str,
timeout: float | None = None,
) -> str:
return await self.files._register_file_url(path, timeout=timeout)
async def text_to_image(
self,
text: str,
*,
return_url: bool = True,
) -> str:
"""Render plain text into an image using the host renderer."""
output = await self._proxy.call(
"system.text_to_image",
{"text": text, "return_url": return_url},
)
return str(output.get("result", ""))
async def html_render(
self,
tmpl: str,
data: dict[str, Any],
*,
return_url: bool = True,
options: dict[str, Any] | None = None,
) -> str:
"""Render an HTML template using the host renderer."""
output = await self._proxy.call(
"system.html_render",
{
"tmpl": tmpl,
"data": dict(data),
"return_url": return_url,
"options": options,
},
)
return str(output.get("result", ""))
async def get_using_provider(self, umo: str | None = None) -> ProviderMeta | None:
return await self.providers.get_using_chat(umo)
async def get_current_chat_provider_id(self, umo: str | None = None) -> str | None:
output = await self._proxy.call(
"provider.get_current_chat_provider_id",
{"umo": umo},
)
value = output.get("provider_id")
return str(value) if value else None
async def get_all_providers(self) -> list[ProviderMeta]:
return await self.providers.list_all()
async def get_all_tts_providers(self) -> list[ProviderMeta]:
return await self.providers.list_tts()
async def get_all_stt_providers(self) -> list[ProviderMeta]:
return await self.providers.list_stt()
async def get_all_embedding_providers(self) -> list[ProviderMeta]:
return await self.providers.list_embedding()
async def get_all_rerank_providers(self) -> list[ProviderMeta]:
return await self.providers.list_rerank()
async def get_using_tts_provider(
self, umo: str | None = None
) -> ProviderMeta | None:
provider = await self.providers.get_using_tts(umo)
return provider.meta() if provider is not None else None
async def get_using_stt_provider(
self, umo: str | None = None
) -> ProviderMeta | None:
provider = await self.providers.get_using_stt(umo)
return provider.meta() if provider is not None else None
async def send_message(
self,
session: str | MessageSession,
content: PlatformCompatContent,
) -> dict[str, Any]:
return await self.platform.send_by_session(session, content)
async def send_message_by_id(
self,
type: str,
id: str,
content: PlatformCompatContent,
*,
platform: str,
) -> dict[str, Any]:
platform_payload = await self._resolve_platform_target(platform)
return await self.platform.send_by_id(
str(platform_payload.get("id", "")),
str(id),
content,
message_type=self._normalize_compat_message_type(type),
)
@staticmethod
def _normalize_compat_message_type(value: str) -> str:
normalized = str(value).strip().lower()
if normalized in {"groupmessage", "group_message", "group"}:
return "group"
if normalized in {
"privatemessage",
"private_message",
"private",
"friendmessage",
"friend_message",
"friend",
}:
return "private"
if not normalized:
raise AstrBotError.invalid_input("send_message_by_id requires type")
return normalized
async def _resolve_platform_target(self, platform: str) -> dict[str, Any]:
target = str(platform).strip()
if not target:
raise AstrBotError.invalid_input(
"send_message_by_id requires explicit platform"
)
instances = await self._list_platform_instances()
id_matches = [
item for item in instances if str(item.get("id", "")).strip() == target
]
if len(id_matches) == 1:
return id_matches[0]
normalized_target = target.lower()
alias_matches = [
item
for item in instances
if str(item.get("type", "")).strip().lower() == normalized_target
or str(item.get("name", "")).strip().lower() == normalized_target
]
if len(alias_matches) == 1:
return alias_matches[0]
if len(alias_matches) > 1:
raise AstrBotError.invalid_input(
f"send_message_by_id platform '{target}' is ambiguous"
)
raise AstrBotError.invalid_input(
f"send_message_by_id cannot resolve platform '{target}'"
)
def get_llm_tool_manager(self) -> LLMToolManager:
return self._llm_tool_manager
async def activate_llm_tool(self, name: str) -> bool:
return await self._llm_tool_manager.activate(name)
async def deactivate_llm_tool(self, name: str) -> bool:
return await self._llm_tool_manager.deactivate(name)
async def add_llm_tools(self, *tools: LLMToolSpec) -> list[str]:
return await self._llm_tool_manager.add(*tools)
async def register_llm_tool(
self,
name: str,
parameters_schema: dict[str, Any],
desc: str,
func_obj: Callable[..., Any] | Callable[..., Awaitable[Any]],
*,
active: bool = True,
) -> list[str]:
if not callable(func_obj):
raise TypeError("register_llm_tool requires a callable func_obj")
tool_name = str(name).strip()
if not tool_name:
raise AstrBotError.invalid_input("register_llm_tool requires name")
if not isinstance(parameters_schema, dict):
raise TypeError("register_llm_tool requires parameters_schema dict")
handler_ref = f"__dynamic_llm_tool__:{tool_name}"
tool_spec = LLMToolSpec.create(
name=tool_name,
description=str(desc),
parameters_schema=dict(parameters_schema),
handler_ref=handler_ref,
active=bool(active),
)
owner = getattr(func_obj, "__self__", None) or current_star_instance()
dispatcher = getattr(self.peer, "_sdk_capability_dispatcher", None)
if dispatcher is not None and hasattr(dispatcher, "add_dynamic_llm_tool"):
dispatcher.add_dynamic_llm_tool(
plugin_id=self.plugin_id,
spec=tool_spec,
callable_obj=func_obj,
owner=owner,
)
try:
return await self._llm_tool_manager.add(tool_spec)
except Exception:
if dispatcher is not None and hasattr(dispatcher, "remove_llm_tool"):
dispatcher.remove_llm_tool(self.plugin_id, tool_name)
raise
async def unregister_llm_tool(self, name: str) -> bool:
removed = await self._llm_tool_manager.remove(str(name))
dispatcher = getattr(self.peer, "_sdk_capability_dispatcher", None)
if dispatcher is not None and hasattr(dispatcher, "remove_llm_tool"):
dispatcher.remove_llm_tool(self.plugin_id, str(name))
return removed
async def register_skill(
self,
*,
name: str,
path: str | Path,
description: str = "",
) -> SkillRegistration:
return await self.skills.register(
name=name,
path=str(path),
description=description,
)
async def unregister_skill(self, name: str) -> bool:
return await self.skills.unregister(name)
async def tool_loop_agent(
self,
request: ProviderRequest | None = None,
**kwargs: Any,
) -> LLMResponse:
provider_request = request or ProviderRequest()
if kwargs:
merged = provider_request.model_dump()
merged.update(kwargs)
provider_request = ProviderRequest.model_validate(merged)
payload = provider_request.to_payload()
target_payload = self._source_event_payload.get("target")
if isinstance(target_payload, dict):
# Preserve the original message target so core can recover the
# dispatch token for message-bound tool loop execution.
payload["target"] = dict(target_payload)
output = await self._proxy.call("agent.tool_loop.run", payload)
return LLMResponse.model_validate(output)
def _source_event_type(self) -> str:
event_type = self._source_event_payload.get("event_type")
if isinstance(event_type, str) and event_type.strip():
return event_type.strip()
fallback_type = self._source_event_payload.get("type")
if isinstance(fallback_type, str) and fallback_type.strip():
return fallback_type.strip()
raw_payload = self._source_event_payload.get("raw")
if isinstance(raw_payload, dict):
raw_event_type = raw_payload.get("event_type")
if isinstance(raw_event_type, str) and raw_event_type.strip():
return raw_event_type.strip()
return ""
async def register_commands(
self,
command_name: str,
handler_full_name: str,
*,
desc: str = "",
priority: int = 0,
use_regex: bool = False,
ignore_prefix: bool = False,
) -> None:
source_event_type = self._source_event_type()
if source_event_type not in {"astrbot_loaded", "platform_loaded"}:
raise AstrBotError.invalid_input(
"register_commands is only available in astrbot_loaded/platform_loaded events"
)
if ignore_prefix:
raise AstrBotError.invalid_input(
"register_commands(ignore_prefix=True) is unsupported in SDK runtime"
)
if isinstance(priority, bool) or not isinstance(priority, int):
raise AstrBotError.invalid_input(
"register_commands priority must be an integer"
)
await self._proxy.call(
"registry.command.register",
{
"command_name": str(command_name),
"handler_full_name": str(handler_full_name),
"source_event_type": source_event_type,
"desc": str(desc),
"priority": priority,
"use_regex": bool(use_regex),
"ignore_prefix": False,
},
)
async def register_task(
self,
task: Awaitable[Any],
desc: str,
) -> asyncio.Task[Any]:
"""Register a background task owned by the current SDK context.
This is the recommended way to launch follow-up work that should outlive
the current handler dispatch, including `session_waiter(...)` flows.
Directly awaiting a waiter inside the current handler keeps the original
dispatch open until the next message arrives.
Example:
await event.reply("请输入用户名:")
await ctx.register_task(
self.collect_username(event),
"waiter:collect_username",
)
"""
task_desc = str(desc)
async def _wrap_future(future: asyncio.Future[Any]) -> Any:
return await future
if isinstance(task, asyncio.Task):
background_task = task
elif asyncio.isfuture(task):
background_task = asyncio.create_task(_wrap_future(task))
elif asyncio.iscoroutine(task):
background_task = asyncio.create_task(task)
else:
raise TypeError("register_task requires an awaitable task object")
_mark_session_waiter_background_task(background_task)
def _on_done(done_task: asyncio.Task[Any]) -> None:
_unmark_session_waiter_background_task(done_task)
if done_task.cancelled():
debug_logger = getattr(self.logger, "debug", None)
if callable(debug_logger):
debug_logger(
"SDK background task cancelled: plugin_id={} desc={}",
self.plugin_id,
task_desc,
)
return
try:
done_task.result()
except Exception:
exception_logger = getattr(self.logger, "exception", None)
if callable(exception_logger):
exception_logger(
"SDK background task failed: plugin_id={} desc={}",
self.plugin_id,
task_desc,
)
background_task.add_done_callback(_on_done)
return background_task
async def _list_platform_instances(self) -> list[dict[str, Any]]:
output = await self._proxy.call("platform.list_instances", {})
items = output.get("platforms")
if not isinstance(items, list):
return []
normalized: list[dict[str, Any]] = []
for item in items:
if not isinstance(item, dict):
continue
platform_id = str(item.get("id", "")).strip()
platform_type = str(item.get("type", "")).strip()
if not platform_id or not platform_type:
continue
normalized.append(
{
"id": platform_id,
"name": str(item.get("name", platform_id)),
"type": platform_type,
"status": PlatformStatus.from_value(item.get("status")),
}
)
return normalized
def _build_platform_facade(
self,
platform_payload: dict[str, Any],
) -> PlatformCompatFacade:
return PlatformCompatFacade(
_ctx=self,
id=str(platform_payload.get("id", "")),
name=str(platform_payload.get("name", "")),
type=str(platform_payload.get("type", "")),
status=PlatformStatus.from_value(platform_payload.get("status")),
)
async def get_platform(self, platform_type: str) -> PlatformCompatFacade | None:
target_type = str(platform_type).strip().lower()
if not target_type:
return None
for item in await self._list_platform_instances():
if str(item.get("type", "")).strip().lower() == target_type:
return self._build_platform_facade(item)
return None
async def get_platform_inst(self, platform_id: str) -> PlatformCompatFacade | None:
target_id = str(platform_id).strip()
if not target_id:
return None
for item in await self._list_platform_instances():
if str(item.get("id", "")).strip() == target_id:
return self._build_platform_facade(item)
return None
-133
View File
@@ -1,133 +0,0 @@
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from enum import Enum
from typing import Any
from .context import Context
from .events import MessageEvent
from .message.components import BaseMessageComponent
from .message.result import MessageChain
from .session_waiter import SessionWaiterManager
DEFAULT_BUSY_MESSAGE = "当前会话已有进行中的交互,请先完成后再试。"
class ConversationState(str, Enum):
ACTIVE = "active"
REJECTED_BUSY = "rejected_busy"
REPLACED = "replaced"
TIMEOUT = "timeout"
COMPLETED = "completed"
CANCELLED = "cancelled"
class ConversationReplaced(RuntimeError):
pass
class ConversationClosed(RuntimeError):
pass
@dataclass(slots=True)
class ConversationSession:
ctx: Context
event: MessageEvent
waiter_manager: SessionWaiterManager
timeout: int
state: ConversationState = ConversationState.ACTIVE
_owner_task: asyncio.Task[Any] | None = None
def __post_init__(self) -> None:
if self.state != ConversationState.ACTIVE:
self.state = ConversationState.ACTIVE
def bind_owner_task(self, task: asyncio.Task[Any]) -> None:
self._owner_task = task
@property
def session_key(self) -> str:
return self.event.unified_msg_origin
@property
def active(self) -> bool:
return self.state == ConversationState.ACTIVE
async def ask(self, prompt: str, timeout: int | None = None) -> MessageEvent:
self._ensure_usable("ask")
if prompt:
await self.reply(prompt)
try:
return await self.waiter_manager.wait_for_event(
event=self.event,
timeout=timeout or self.timeout,
record_history_chains=False,
)
except asyncio.TimeoutError:
self.close(ConversationState.TIMEOUT)
raise
except asyncio.CancelledError as exc:
if self.state == ConversationState.REPLACED:
raise ConversationReplaced(
"conversation replaced by a newer session"
) from exc
self.close(ConversationState.CANCELLED)
raise
async def reply(self, text: str) -> None:
self._ensure_usable("reply")
await self.event.reply(text)
async def reply_chain(
self,
chain: MessageChain | list[BaseMessageComponent] | list[dict[str, Any]],
) -> None:
self._ensure_usable("reply_chain")
await self.event.reply_chain(chain)
async def send_message(
self,
content: str | MessageChain | list[BaseMessageComponent] | list[dict[str, Any]],
) -> dict[str, Any]:
self._ensure_usable("send_message")
return await self.ctx.platform.send_by_session(self.event.session_id, content)
def end(self) -> None:
self.close(ConversationState.COMPLETED)
def mark_replaced(self) -> None:
self.close(ConversationState.REPLACED)
def close(self, state: ConversationState) -> None:
if self.state != ConversationState.ACTIVE and state == self.state:
return
if (
self.state != ConversationState.ACTIVE
and state != ConversationState.REPLACED
):
return
self.state = state
def _ensure_usable(self, action: str) -> None:
if (
self._owner_task is not None
and asyncio.current_task() is not self._owner_task
):
raise ConversationClosed(
f"ConversationSession cannot be used outside its owner task during {action}"
)
if not self.active:
raise ConversationClosed(
f"ConversationSession is already closed ({self.state.value}) during {action}"
)
__all__ = [
"ConversationClosed",
"ConversationReplaced",
"ConversationSession",
"ConversationState",
"DEFAULT_BUSY_MESSAGE",
]
-928
View File
@@ -1,928 +0,0 @@
"""v4 原生装饰器。
提供声明式的方法来注册 handler 和 capability。
装饰器会在方法上附加元数据,由 Star.__init_subclass__ 自动收集。
触发器装饰器:
- @on_command: 命令触发器
- @on_message: 消息触发器(关键词/正则)
- @on_event: 事件触发器
- @on_schedule: 定时任务触发器
- @conversation_command: 带会话生命周期的命令触发器
权限与过滤装饰器:
- @require_admin / @admin_only: 管理员权限标记
- @platforms: 限定平台
- @group_only / @private_only: 群聊/私聊限定
- @message_types: 消息类型过滤
限流装饰器:
- @rate_limit: 滑动窗口限流
- @cooldown: 冷却时间
优先级装饰器:
- @priority: 设置执行优先级
能力导出装饰器:
- @provide_capability: 声明对外暴露的能力
- @register_llm_tool: 注册 LLM 工具
- @register_agent: 注册 Agent
Example:
class MyPlugin(Star):
@on_command("hello", aliases=["hi"])
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
@on_message(keywords=["help"])
async def help(self, event: MessageEvent, ctx: Context):
await event.reply("Help info...")
@provide_capability("my_plugin.calculate", description="计算")
async def calculate(self, payload: dict, ctx: Context):
return {"result": payload["x"] * 2}
"""
from __future__ import annotations
import inspect
import typing
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Literal, cast
from pydantic import BaseModel
from ._internal.typing_utils import unwrap_optional
from .llm.agents import AgentSpec, BaseAgentRunner
from .llm.entities import LLMToolSpec
from .protocol.descriptors import (
RESERVED_CAPABILITY_PREFIXES,
CapabilityDescriptor,
CommandRouteSpec,
CommandTrigger,
EventTrigger,
FilterSpec,
MessageTrigger,
MessageTypeFilterSpec,
Permissions,
PlatformFilterSpec,
ScheduleTrigger,
)
HandlerCallable = Callable[..., Any]
HANDLER_META_ATTR = "__astrbot_handler_meta__"
CAPABILITY_META_ATTR = "__astrbot_capability_meta__"
LLM_TOOL_META_ATTR = "__astrbot_llm_tool_meta__"
AGENT_META_ATTR = "__astrbot_agent_meta__"
LimiterScope = Literal["session", "user", "group", "global"]
LimiterBehavior = Literal["hint", "silent", "error"]
ConversationMode = Literal["replace", "reject"]
@dataclass(slots=True)
class LimiterMeta:
kind: Literal["rate_limit", "cooldown"]
limit: int
window: float
scope: LimiterScope = "session"
behavior: LimiterBehavior = "hint"
message: str | None = None
@dataclass(slots=True)
class ConversationMeta:
timeout: int = 60
mode: ConversationMode = "replace"
busy_message: str | None = None
grace_period: float = 1.0
@dataclass(slots=True)
class HandlerMeta:
"""Handler 元数据。
存储在方法上的 __astrbot_handler_meta__ 属性中。
Attributes:
trigger: 触发器(命令/消息/事件/定时)
kind: handler 类型标识
contract: 契约类型(可选)
priority: 执行优先级(数值越大越先执行)
permissions: 权限要求
"""
trigger: CommandTrigger | MessageTrigger | EventTrigger | ScheduleTrigger | None = (
None
)
kind: str = "handler"
contract: str | None = None
description: str | None = None
priority: int = 0
permissions: Permissions = field(default_factory=Permissions)
filters: list[FilterSpec] = field(default_factory=list)
local_filters: list[Any] = field(default_factory=list)
command_route: CommandRouteSpec | None = None
limiter: LimiterMeta | None = None
conversation: ConversationMeta | None = None
decorator_sources: dict[str, str] = field(default_factory=dict)
@dataclass(slots=True)
class CapabilityMeta:
"""Capability 元数据。
存储在方法上的 __astrbot_capability_meta__ 属性中。
Attributes:
descriptor: 能力描述符
"""
descriptor: CapabilityDescriptor
@dataclass(slots=True)
class LLMToolMeta:
spec: LLMToolSpec
@dataclass(slots=True)
class AgentMeta:
spec: AgentSpec
def _get_or_create_meta(func: HandlerCallable) -> HandlerMeta:
"""获取或创建 handler 元数据。"""
meta = getattr(func, HANDLER_META_ATTR, None)
if meta is None:
meta = HandlerMeta()
setattr(func, HANDLER_META_ATTR, meta)
return meta
def get_handler_meta(func: HandlerCallable) -> HandlerMeta | None:
"""获取方法的 handler 元数据。
Args:
func: 要检查的方法
Returns:
HandlerMeta 实例,如果没有则返回 None
"""
return getattr(func, HANDLER_META_ATTR, None)
def get_capability_meta(func: HandlerCallable) -> CapabilityMeta | None:
"""获取方法的 capability 元数据。
Args:
func: 要检查的方法
Returns:
CapabilityMeta 实例,如果没有则返回 None
"""
return getattr(func, CAPABILITY_META_ATTR, None)
def get_llm_tool_meta(func: HandlerCallable) -> LLMToolMeta | None:
return getattr(func, LLM_TOOL_META_ATTR, None)
def get_agent_meta(obj: Any) -> AgentMeta | None:
return getattr(obj, AGENT_META_ATTR, None)
def _replace_filter(meta: HandlerMeta, spec: FilterSpec) -> None:
kind = getattr(spec, "kind", None)
meta.filters = [
item for item in meta.filters if getattr(item, "kind", None) != kind
]
meta.filters.append(spec)
def _has_filter_kind(meta: HandlerMeta, kind: str) -> bool:
return any(getattr(item, "kind", None) == kind for item in meta.filters)
def _set_platform_filter(
meta: HandlerMeta,
values: list[str],
*,
source: str,
) -> None:
normalized = [
value for value in dict.fromkeys(str(item).strip() for item in values) if value
]
if not normalized:
return
existing = meta.decorator_sources.get("platforms")
if existing is not None and existing != source:
raise ValueError("platforms(...) 不能与 on_message(platforms=...) 混用")
if existing is None and _has_filter_kind(meta, "platform"):
raise ValueError("platforms(...) 不能与已有平台过滤器混用")
meta.decorator_sources["platforms"] = source
_replace_filter(meta, PlatformFilterSpec(platforms=normalized))
def _set_message_type_filter(
meta: HandlerMeta,
values: list[str],
*,
source: str,
) -> None:
normalized = [
value
for value in dict.fromkeys(str(item).strip().lower() for item in values)
if value
]
if not normalized:
return
existing = meta.decorator_sources.get("message_types")
if existing is not None and existing != source:
raise ValueError(
"group_only()/private_only()/message_types(...) 不能与已有消息类型约束混用"
)
if existing is None and _has_filter_kind(meta, "message_type"):
raise ValueError(
"group_only()/private_only()/message_types(...) 不能与已有消息类型过滤器混用"
)
meta.decorator_sources["message_types"] = source
_replace_filter(meta, MessageTypeFilterSpec(message_types=normalized))
def _validate_message_trigger_compatibility(meta: HandlerMeta) -> None:
if meta.limiter is None or meta.trigger is None:
return
trigger_type = getattr(meta.trigger, "type", None)
if trigger_type not in {"command", "message"}:
raise ValueError(
"rate_limit(...) 和 cooldown(...) 只适用于 on_command/on_message"
)
def _normalize_description(description: str | None) -> str | None:
if description is None:
return None
text = str(description).strip()
return text or None
def _validate_limiter_args(
*,
kind: str,
limit: int,
window: float,
scope: LimiterScope,
behavior: LimiterBehavior,
) -> None:
if isinstance(limit, bool) or int(limit) <= 0:
raise ValueError(f"{kind} requires a positive limit")
if float(window) <= 0:
raise ValueError(f"{kind} requires a positive window")
if scope not in {"session", "user", "group", "global"}:
raise ValueError(f"unsupported limiter scope: {scope}")
if behavior not in {"hint", "silent", "error"}:
raise ValueError(f"unsupported limiter behavior: {behavior}")
def _set_limiter(
func: HandlerCallable,
limiter: LimiterMeta,
) -> HandlerCallable:
meta = _get_or_create_meta(func)
if meta.limiter is not None:
raise ValueError("rate_limit(...) 和 cooldown(...) 不能叠加在同一个 handler 上")
meta.limiter = limiter
_validate_message_trigger_compatibility(meta)
return func
def _model_to_schema(
model: type[BaseModel] | None,
*,
label: str,
) -> dict[str, Any] | None:
"""将 pydantic 模型转换为 JSON Schema。
Args:
model: pydantic BaseModel 子类
label: 错误消息中的字段名
Returns:
JSON Schema 字典,如果 model 为 None 则返回 None
Raises:
TypeError: 如果 model 不是 BaseModel 子类
"""
if model is None:
return None
if not isinstance(model, type) or not issubclass(model, BaseModel):
raise TypeError(f"{label} 必须是 pydantic BaseModel 子类")
return cast(dict[str, Any], model.model_json_schema())
def on_command(
command: str | typing.Sequence[str],
*,
aliases: list[str] | None = None,
description: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册命令处理方法。
当用户发送指定命令时触发。命令格式为 `/{command}` 或直接 `{command}`,
取决于平台配置。
Args:
command: 命令名称(不包含前缀符)
aliases: 命令别名列表
description: 命令描述,用于帮助信息
Returns:
装饰器函数
Example:
@on_command("echo", aliases=["repeat"], description="重复消息")
async def echo(self, event: MessageEvent, ctx: Context):
await event.reply(event.text)
"""
commands = (
[str(command).strip()]
if isinstance(command, str)
else [str(item).strip() for item in command]
)
commands = [item for item in commands if item]
if not commands:
raise ValueError("on_command requires at least one non-empty command name")
canonical = commands[0]
merged_aliases: list[str] = [
item
for item in dict.fromkeys([*commands[1:], *(aliases or [])])
if isinstance(item, str) and item and item != canonical
]
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
normalized_description = _normalize_description(description)
meta.trigger = CommandTrigger(
command=canonical,
aliases=merged_aliases,
description=normalized_description,
)
meta.description = normalized_description
_validate_message_trigger_compatibility(meta)
return func
return decorator
def on_message(
*,
regex: str | None = None,
keywords: list[str] | None = None,
platforms: list[str] | None = None,
message_types: list[str] | None = None,
description: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册消息处理方法。
当消息匹配指定条件时触发。支持正则表达式或关键词匹配。
Args:
regex: 正则表达式模式
keywords: 关键词列表(任一匹配即可)
platforms: 限定平台列表(如 ["qq", "wechat"])
Returns:
装饰器函数
Note:
regex 和 keywords 至少提供一个
Example:
@on_message(keywords=["help", "帮助"])
async def help(self, event: MessageEvent, ctx: Context):
await event.reply("帮助信息")
@on_message(regex=r"\\d+") # 匹配数字
async def number_handler(self, event: MessageEvent, ctx: Context):
await event.reply("收到了数字")
"""
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
meta.trigger = MessageTrigger(
regex=regex,
keywords=keywords or [],
platforms=platforms or [],
message_types=message_types or [],
)
meta.description = _normalize_description(description)
if platforms:
_set_platform_filter(meta, list(platforms), source="trigger.platforms")
if message_types:
_set_message_type_filter(
meta,
list(message_types),
source="trigger.message_types",
)
_validate_message_trigger_compatibility(meta)
return func
return decorator
def append_filter_meta(
func: HandlerCallable,
*,
specs: list[FilterSpec] | None = None,
local_bindings: list[Any] | None = None,
) -> HandlerCallable:
"""追加过滤器元数据。"""
meta = _get_or_create_meta(func)
if specs:
meta.filters.extend(specs)
if local_bindings:
meta.local_filters.extend(local_bindings)
return func
def set_command_route_meta(
func: HandlerCallable,
route: CommandRouteSpec,
) -> HandlerCallable:
"""设置命令路由元数据。"""
meta = _get_or_create_meta(func)
meta.command_route = route
return func
def on_event(
event_type: str,
*,
description: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册事件处理方法。
当特定类型的事件发生时触发。用于处理非消息类型的事件,
如群成员变动、好友请求等。
Args:
event_type: 事件类型标识
Returns:
装饰器函数
Example:
@on_event("group_member_join")
async def on_join(self, event, ctx):
await ctx.platform.send(event.group_id, "欢迎新人!")
"""
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
meta.trigger = EventTrigger(event_type=event_type)
meta.description = _normalize_description(description)
_validate_message_trigger_compatibility(meta)
return func
return decorator
def on_schedule(
*,
cron: str | None = None,
interval_seconds: int | None = None,
description: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册定时任务方法。
按指定的时间计划定期执行。
Args:
cron: cron 表达式(如 "0 8 * * *" 表示每天 8:00)
interval_seconds: 执行间隔(秒)
Returns:
装饰器函数
Note:
cron 和 interval_seconds 至少提供一个
Example:
@on_schedule(cron="0 8 * * *") # 每天 8:00
async def morning_greeting(self, ctx):
await ctx.platform.send("group_123", "早上好!")
@on_schedule(interval_seconds=3600) # 每小时
async def hourly_check(self, ctx):
pass
"""
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
meta.trigger = ScheduleTrigger(cron=cron, interval_seconds=interval_seconds)
meta.description = _normalize_description(description)
_validate_message_trigger_compatibility(meta)
return func
return decorator
def require_admin(func: HandlerCallable) -> HandlerCallable:
"""标记 handler 需要管理员权限。
当用户不是管理员时,handler 将不会被调用。
Args:
func: 要标记的方法
Returns:
标记后的方法
Example:
@on_command("admin")
@require_admin
async def admin_only(self, event: MessageEvent, ctx: Context):
await event.reply("管理员命令执行成功")
"""
meta = _get_or_create_meta(func)
meta.permissions.require_admin = True
return func
def admin_only(func: HandlerCallable) -> HandlerCallable:
return require_admin(func)
def platforms(*names: str) -> Callable[[HandlerCallable], HandlerCallable]:
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
_set_platform_filter(meta, list(names), source="decorator.platforms")
return func
return decorator
def message_types(*types: str) -> Callable[[HandlerCallable], HandlerCallable]:
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
_set_message_type_filter(
meta,
list(types),
source="decorator.message_types",
)
return func
return decorator
def group_only() -> Callable[[HandlerCallable], HandlerCallable]:
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
_set_message_type_filter(meta, ["group"], source="decorator.group_only")
return func
return decorator
def private_only() -> Callable[[HandlerCallable], HandlerCallable]:
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
_set_message_type_filter(meta, ["private"], source="decorator.private_only")
return func
return decorator
def priority(value: int) -> Callable[[HandlerCallable], HandlerCallable]:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError("priority(...) requires an integer")
def decorator(func: HandlerCallable) -> HandlerCallable:
meta = _get_or_create_meta(func)
meta.priority = value
return func
return decorator
def rate_limit(
limit: int,
window: float,
*,
scope: LimiterScope = "session",
behavior: LimiterBehavior = "hint",
message: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
_validate_limiter_args(
kind="rate_limit",
limit=limit,
window=window,
scope=scope,
behavior=behavior,
)
def decorator(func: HandlerCallable) -> HandlerCallable:
return _set_limiter(
func,
LimiterMeta(
kind="rate_limit",
limit=int(limit),
window=float(window),
scope=scope,
behavior=behavior,
message=message,
),
)
return decorator
def cooldown(
seconds: float,
*,
scope: LimiterScope = "session",
behavior: LimiterBehavior = "hint",
message: str | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
_validate_limiter_args(
kind="cooldown",
limit=1,
window=seconds,
scope=scope,
behavior=behavior,
)
def decorator(func: HandlerCallable) -> HandlerCallable:
return _set_limiter(
func,
LimiterMeta(
kind="cooldown",
limit=1,
window=float(seconds),
scope=scope,
behavior=behavior,
message=message,
),
)
return decorator
def conversation_command(
command: str | typing.Sequence[str],
*,
aliases: list[str] | None = None,
description: str | None = None,
timeout: int = 60,
mode: ConversationMode = "replace",
busy_message: str | None = None,
grace_period: float = 1.0,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册带会话生命周期的命令处理方法。
在 ``on_command`` 基础上附加会话元数据,支持超时、并发策略和宽限期控制。
Args:
command: 命令名称或序列(首项为正式名,其余视为别名)
aliases: 额外别名列表
description: 命令描述
timeout: 会话超时时间(秒),必须为正整数
mode: 会话冲突时的行为:
- ``"replace"``: 替换当前会话
- ``"reject"``: 拒绝新请求
busy_message: 拒绝新请求时的提示消息
grace_period: 宽限期(秒),用于会话生命周期处理
Returns:
装饰器函数
Raises:
ValueError: mode 不合法、timeout 非正整数或 grace_period 非正数
Example:
@conversation_command("chat", timeout=120, mode="reject", busy_message="请稍后再试")
async def chat(self, event: MessageEvent, ctx: Context):
await event.reply("开始对话...")
"""
if mode not in {"replace", "reject"}:
raise ValueError("conversation_command mode must be 'replace' or 'reject'")
# bool 是 int 子类,需单独排除
if isinstance(timeout, bool) or int(timeout) <= 0:
raise ValueError("conversation_command timeout must be a positive integer")
if float(grace_period) <= 0:
raise ValueError("conversation_command grace_period must be positive")
command_decorator = on_command(
command,
aliases=aliases,
description=description,
)
def decorator(func: HandlerCallable) -> HandlerCallable:
decorated = command_decorator(func)
meta = _get_or_create_meta(decorated)
meta.conversation = ConversationMeta(
timeout=int(timeout),
mode=mode,
busy_message=busy_message,
grace_period=float(grace_period),
)
return decorated
return decorator
def provide_capability(
name: str,
*,
description: str,
input_schema: dict[str, Any] | None = None,
output_schema: dict[str, Any] | None = None,
input_model: type[BaseModel] | None = None,
output_model: type[BaseModel] | None = None,
supports_stream: bool = False,
cancelable: bool = False,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""声明插件对外暴露的 capability。
允许其他插件或 Core 通过 capability 名称调用此方法。
支持使用 JSON Schema 或 pydantic 模型定义输入输出。
Args:
name: capability 名称(不能使用保留命名空间)
description: 能力描述
input_schema: 输入 JSON Schema
output_schema: 输出 JSON Schema
input_model: 输入 pydantic 模型(与 input_schema 二选一)
output_model: 输出 pydantic 模型(与 output_schema 二选一)
supports_stream: 是否支持流式输出
cancelable: 是否可取消
Returns:
装饰器函数
Raises:
ValueError: 如果使用保留命名空间,或同时提供 schema 和 model
Example:
@provide_capability(
"my_plugin.calculate",
description="执行计算",
input_model=CalculateInput,
output_model=CalculateOutput,
)
async def calculate(self, payload: dict, ctx: Context):
return {"result": payload["x"] * 2}
"""
def decorator(func: HandlerCallable) -> HandlerCallable:
if name.startswith(RESERVED_CAPABILITY_PREFIXES):
raise ValueError(f"保留 capability 命名空间不能用于插件导出:{name}")
if input_schema is not None and input_model is not None:
raise ValueError("input_schema 和 input_model 不能同时提供")
if output_schema is not None and output_model is not None:
raise ValueError("output_schema 和 output_model 不能同时提供")
descriptor = CapabilityDescriptor(
name=name,
description=description,
input_schema=(
input_schema
if input_schema is not None
else _model_to_schema(input_model, label="input_model")
),
output_schema=(
output_schema
if output_schema is not None
else _model_to_schema(output_model, label="output_model")
),
supports_stream=supports_stream,
cancelable=cancelable,
)
setattr(func, CAPABILITY_META_ATTR, CapabilityMeta(descriptor=descriptor))
return func
return decorator
def _annotation_to_schema(annotation: Any) -> dict[str, Any]:
normalized, _is_optional = unwrap_optional(annotation)
origin = typing.get_origin(normalized)
if normalized is str:
return {"type": "string"}
if normalized is int:
return {"type": "integer"}
if normalized is float:
return {"type": "number"}
if normalized is bool:
return {"type": "boolean"}
if normalized is dict or origin is dict:
return {"type": "object"}
if normalized is list or origin is list:
args = typing.get_args(normalized)
item_schema = _annotation_to_schema(args[0]) if args else {}
return {"type": "array", "items": item_schema}
return {"type": "string"}
def _callable_parameters_schema(func: HandlerCallable) -> dict[str, Any]:
signature = inspect.signature(func)
type_hints: dict[str, Any] = {}
try:
type_hints = typing.get_type_hints(func)
except Exception:
type_hints = {}
properties: dict[str, Any] = {}
required: list[str] = []
for parameter in signature.parameters.values():
if parameter.kind not in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
continue
if parameter.name == "self":
continue
annotation = type_hints.get(parameter.name)
normalized, _is_optional = unwrap_optional(annotation)
if parameter.name in {"event", "ctx", "context"}:
continue
properties[parameter.name] = _annotation_to_schema(normalized)
if parameter.default is inspect.Parameter.empty and not _is_optional:
required.append(parameter.name)
schema: dict[str, Any] = {"type": "object", "properties": properties}
if required:
schema["required"] = required
return schema
def register_llm_tool(
name: str | None = None,
*,
description: str | None = None,
parameters_schema: dict[str, Any] | None = None,
active: bool = True,
) -> Callable[[HandlerCallable], HandlerCallable]:
def decorator(func: HandlerCallable) -> HandlerCallable:
tool_name = str(name or func.__name__).strip()
if not tool_name:
raise ValueError("LLM tool name must not be empty")
setattr(
func,
LLM_TOOL_META_ATTR,
LLMToolMeta(
spec=LLMToolSpec.create(
name=tool_name,
description=description
or (inspect.getdoc(func) or "").splitlines()[0]
if inspect.getdoc(func)
else "",
parameters_schema=parameters_schema
or _callable_parameters_schema(func),
handler_ref=tool_name,
active=active,
)
),
)
return func
return decorator
def register_agent(
name: str,
*,
description: str = "",
tool_names: list[str] | None = None,
) -> Callable[[type[BaseAgentRunner]], type[BaseAgentRunner]]:
def decorator(cls: type[BaseAgentRunner]) -> type[BaseAgentRunner]:
if not inspect.isclass(cls) or not issubclass(cls, BaseAgentRunner):
raise TypeError("@register_agent() 只接受 BaseAgentRunner 子类")
setattr(
cls,
AGENT_META_ATTR,
AgentMeta(
spec=AgentSpec(
name=name,
description=description,
tool_names=list(tool_names or []),
runner_class=f"{cls.__module__}.{cls.__qualname__}",
)
),
)
return cls
return decorator
def acknowledge_global_mcp_risk(cls: type[Any]) -> type[Any]:
"""Mark an SDK plugin class as eligible to mutate global MCP state.
This is intentionally a coarse, class-level marker. Runtime enforcement lives
in the Core MCP capability bridge.
"""
setattr(cls, "__astrbot_acknowledge_global_mcp_risk__", True)
return cls
@@ -1,665 +0,0 @@
# AstrBot SDK Context API 参考文档
## 概述
`Context` 是插件与 AstrBot Core 交互的主要入口,每个 handler 调用都会创建一个新的 Context 实例。Context 组合了所有 capability 客户端,提供统一的访问接口。
## 目录
- [Context 类属性](#context-类属性)
- [核心客户端](#核心客户端)
- [LLM 客户端 (ctx.llm)](#llm-客户端)
- [Memory 客户端 (ctx.memory)](#memory-客户端)
- [Database 客户端 (ctx.db)](#database-客户端)
- [Files 客户端 (ctx.files)](#files-客户端)
- [Platform 客户端 (ctx.platform)](#platform-客户端)
- [Provider 客户端 (ctx.providers)](#provider-客户端)
- [HTTP 客户端 (ctx.http)](#http-客户端)
- [Metadata 客户端 (ctx.metadata)](#metadata-客户端)
- [LLM Tool 管理方法](#llm-tool-管理方法)
- [系统工具方法](#系统工具方法)
---
## Context 类属性
### 基本属性
```python
@dataclass
class Context:
peer: Any # 协议对等端,用于底层通信
plugin_id: str # 当前插件 ID
logger: PluginLogger # 绑定了插件 ID 的日志器
cancel_token: CancelToken # 取消令牌,用于处理请求取消
```
### 客户端属性
```python
ctx.llm: LLMClient # LLM 能力客户端
ctx.memory: MemoryClient # 记忆能力客户端
ctx.db: DBClient # 数据库客户端
ctx.files: FileServiceClient # 文件服务客户端
ctx.platform: PlatformClient # 平台客户端
ctx.providers: ProviderClient # Provider 客户端
ctx.provider_manager: ProviderManagerClient # Provider 管理客户端
ctx.personas: PersonaManagerClient # 人格管理客户端
ctx.conversations: ConversationManagerClient # 对话管理客户端
ctx.kbs: KnowledgeBaseManagerClient # 知识库管理客户端
ctx.http: HTTPClient # HTTP 客户端
ctx.metadata: MetadataClient # 元数据客户端
```
---
## 核心客户端
### logger
绑定了插件 ID 的日志器,自动添加插件上下文信息。
```python
# 不同级别的日志
ctx.logger.debug("调试信息")
ctx.logger.info("普通信息")
ctx.logger.warning("警告信息")
ctx.logger.error("错误信息")
# 绑定额外上下文
logger = ctx.logger.bind(user_id="12345")
logger.info("用户操作")
# 流式日志监听
async for entry in ctx.logger.watch():
print(f"[{entry.level}] {entry.message}")
```
### cancel_token
取消令牌,用于长时间运行的任务中检查是否需要取消。
```python
# 检查是否取消
ctx.cancel_token.raise_if_cancelled()
# 触发取消
ctx.cancel_token.cancel()
# 等待取消信号
await ctx.cancel_token.wait()
```
---
## LLM 客户端
### chat()
发送聊天请求并返回文本响应。
```python
async def chat(
prompt: str,
*,
system: str | None = None,
history: Sequence[ChatHistoryItem] | None = None,
provider_id: str | None = None,
model: str | None = None,
temperature: float | None = None,
**kwargs: Any,
) -> str
```
**使用示例:**
```python
# 简单对话
reply = await ctx.llm.chat("你好,介绍一下自己")
# 带系统提示词
reply = await ctx.llm.chat(
"用 Python 写一个快速排序",
system="你是一个专业的程序员助手"
)
# 带历史的对话
from astrbot_sdk.clients.llm import ChatMessage
history = [
ChatMessage(role="user", content="我叫小明"),
ChatMessage(role="assistant", content="你好小明!"),
]
reply = await ctx.llm.chat("你记得我的名字吗?", history=history)
```
### chat_raw()
发送聊天请求并返回完整响应对象。
```python
response = await ctx.llm.chat_raw("写一首诗", temperature=0.8)
print(f"生成文本: {response.text}")
print(f"Token 使用: {response.usage}")
print(f"结束原因: {response.finish_reason}")
```
### stream_chat()
流式聊天,逐块返回响应文本。
```python
async for chunk in ctx.llm.stream_chat("讲一个故事"):
print(chunk, end="", flush=True)
```
---
## Memory 客户端
### search()
搜索记忆项。默认在有 embedding provider 时执行 hybrid 检索。
```python
results = await ctx.memory.search("用户喜欢什么颜色", mode="hybrid", limit=5)
for item in results:
print(item["key"], item["score"], item["match_type"])
```
### save()
保存记忆项。
```python
# 保存用户偏好
await ctx.memory.save("user_pref", {"theme": "dark", "lang": "zh"})
# 使用关键字参数
await ctx.memory.save("note", None, content="重要笔记", tags=["work"])
# 显式指定检索文本
await ctx.memory.save(
"profile:alice",
{"name": "Alice", "embedding_text": "Alice 喜欢蓝色和海边"},
)
```
### get()
精确获取单个记忆项。
```python
pref = await ctx.memory.get("user_pref")
if pref:
print(f"用户偏好主题: {pref.get('theme')}")
```
### save_with_ttl()
保存带过期时间的记忆项。
```python
# 保存临时会话状态,1小时后过期
await ctx.memory.save_with_ttl(
"session_temp",
{"state": "waiting"},
ttl_seconds=3600
)
```
### stats()
查看记忆索引状态。
```python
stats = await ctx.memory.stats()
print(stats["total_items"], stats.get("embedded_items"), stats.get("dirty_items"))
```
---
## Database 客户端
### get()
获取指定键的值。
```python
data = await ctx.db.get("user_settings")
if data:
print(data["theme"])
```
### set()
设置键值对。
```python
await ctx.db.set("user_settings", {"theme": "dark", "lang": "zh"})
await ctx.db.set("greeted", True)
```
### delete()
删除指定键的数据。
```python
await ctx.db.delete("user_settings")
```
### list()
列出匹配前缀的所有键。
```python
keys = await ctx.db.list("user_")
# ["user_settings", "user_profile", "user_history"]
```
### get_many()
批量获取多个键的值。
```python
values = await ctx.db.get_many(["user:1", "user:2"])
```
### set_many()
批量写入多个键值对。
```python
await ctx.db.set_many({
"user:1": {"name": "Alice"},
"user:2": {"name": "Bob"}
})
```
### watch()
订阅 KV 变更事件(流式)。
```python
async for event in ctx.db.watch("user:"):
print(event["op"], event["key"])
```
---
## Files 客户端
### register_file()
注册文件并获取令牌。
```python
token = await ctx.files.register_file("/path/to/file.jpg", timeout=3600)
```
### handle_file()
通过令牌解析文件路径。
```python
path = await ctx.files.handle_file(token)
```
---
## Platform 客户端
### send()
发送文本消息。
```python
await ctx.platform.send(event.session_id, "收到您的消息!")
```
### send_image()
发送图片消息。
```python
await ctx.platform.send_image(
event.session_id,
"https://example.com/image.png"
)
```
### send_chain()
发送富消息链。
```python
from astrbot_sdk.message_components import Plain, Image
chain = [Plain("文字"), Image(url="https://example.com/img.jpg")]
await ctx.platform.send_chain(event.session_id, chain)
```
### send_by_id()
主动向指定平台会话发送消息。
```python
await ctx.platform.send_by_id(
platform_id="qq",
session_id="user123",
content="Hello",
message_type="private"
)
```
### get_members()
获取群组成员列表。
```python
members = await ctx.platform.get_members("qq:group:123456")
for member in members:
print(f"{member['nickname']} ({member['user_id']})")
```
---
## Provider 客户端
### list_all()
列出所有 Provider。
```python
providers = await ctx.providers.list_all()
for p in providers:
print(f"{p.id}: {p.model}")
```
### get_using_chat()
获取当前使用的聊天 Provider。
```python
provider = await ctx.providers.get_using_chat()
if provider:
print(f"当前使用: {provider.id}")
```
---
## HTTP 客户端
### register_api()
注册 Web API 端点。
```python
from astrbot_sdk.decorators import provide_capability
@provide_capability(
name="my_plugin.http_handler",
description="处理 HTTP 请求"
)
async def handle_http_request(request_id: str, payload: dict, cancel_token):
return {"status": 200, "body": {"result": "ok"}}
await ctx.http.register_api(
route="/my-api",
handler=handle_http_request,
methods=["GET", "POST"]
)
```
### unregister_api()
注销 Web API 端点。
```python
await ctx.http.unregister_api("/my-api")
```
### list_apis()
列出当前插件注册的所有 API。
```python
apis = await ctx.http.list_apis()
for api in apis:
print(f"{api['route']}: {api['methods']}")
```
---
## Metadata 客户端
### get_plugin()
获取指定插件信息。
```python
plugin = await ctx.metadata.get_plugin("another_plugin")
if plugin:
print(f"插件: {plugin.display_name}")
print(f"版本: {plugin.version}")
```
### list_plugins()
获取所有插件列表。
```python
plugins = await ctx.metadata.list_plugins()
for plugin in plugins:
print(f"{plugin.display_name} v{plugin.version}")
```
### get_current_plugin()
获取当前插件信息。
```python
current = await ctx.metadata.get_current_plugin()
if current:
print(f"当前插件: {current.name} v{current.version}")
```
### get_plugin_config()
获取插件配置。
```python
config = await ctx.metadata.get_plugin_config()
if config:
api_key = config.get("api_key")
```
---
## LLM Tool 管理方法
### register_llm_tool()
注册可执行的 LLM 工具。
```python
async def search_weather(location: str) -> str:
return f"{location} 今天晴天"
await ctx.register_llm_tool(
name="search_weather",
parameters_schema={
"type": "object",
"properties": {
"location": {"type": "string", "description": "城市名称"}
},
"required": ["location"]
},
desc="搜索天气信息",
func_obj=search_weather,
active=True
)
```
### add_llm_tools()
添加 LLM 工具规范。
```python
from astrbot_sdk.llm.tools import LLMToolSpec
tool_spec = LLMToolSpec(
name="my_tool",
description="我的工具",
parameters_schema={...}
)
await ctx.add_llm_tools(tool_spec)
```
### activate_llm_tool() / deactivate_llm_tool()
激活/停用 LLM 工具。
```python
await ctx.activate_llm_tool("my_tool")
await ctx.deactivate_llm_tool("my_tool")
```
---
## 系统工具方法
### get_data_dir()
获取插件数据目录路径。
```python
data_dir = await ctx.get_data_dir()
print(f"数据目录: {data_dir}")
```
### text_to_image()
将文本渲染为图片。
```python
url = await ctx.text_to_image("Hello World", return_url=True)
```
### html_render()
渲染 HTML 模板。
```python
url = await ctx.html_render(
tmpl="<h1>{{ title }}</h1>",
data={"title": "标题"}
)
```
### send_message()
向会话发送消息。
```python
await ctx.send_message(event.session_id, "消息内容")
```
### send_message_by_id()
通过 ID 向平台发送消息。
```python
await ctx.send_message_by_id(
type="private",
id="user123",
content="Hello",
platform="qq"
)
```
### register_task()
注册后台任务。
```python
async def background_work():
while True:
await asyncio.sleep(60)
ctx.logger.info("每分钟执行一次")
task = await ctx.register_task(background_work(), "定时任务")
```
---
## 常见使用模式
### 1. 基本对话流程
```python
from astrbot_sdk.decorators import on_message
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
reply = await ctx.llm.chat(event.message_content)
await ctx.platform.send(event.session_id, reply)
```
### 2. 带历史的对话
```python
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
# 从 memory 获取历史
history_data = await ctx.memory.get(f"history:{event.session_id}")
history = history_data.get("messages", []) if history_data else []
# 对话
reply = await ctx.llm.chat(event.message_content, history=history)
# 保存新消息到历史
history.append(ChatMessage(role="user", content=event.message_content))
history.append(ChatMessage(role="assistant", content=reply))
await ctx.memory.save(f"history:{event.session_id}", {"messages": history})
await ctx.platform.send(event.session_id, reply)
```
### 3. 使用数据库持久化
```python
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
# 获取用户配置
config = await ctx.db.get(f"user_config:{event.sender_id}")
if not config:
config = {"theme": "light", "lang": "zh"}
await ctx.db.set(f"user_config:{event.sender_id}", config)
# 使用配置
reply = f"你的主题设置是: {config['theme']}"
await ctx.platform.send(event.session_id, reply)
```
---
## 注意事项
1. **跨进程通信**:Context 通过 capability 协议与核心通信,所有方法调用都是异步的
2. **插件隔离**:每个插件有独立的 Context 实例,数据和配置是隔离的
3. **取消处理**:长时间运行的操作应定期检查 `ctx.cancel_token.raise_if_cancelled()`
4. **错误处理**:所有远程调用都可能失败,建议使用 try-except 处理
5. **Memory vs DB**:
- Memory: 语义搜索,适合 AI 上下文
- DB: 精确匹配,适合结构化数据
6. **文件操作**:使用 `ctx.files` 注册文件令牌,不要直接传递本地路径
7. **平台标识**:使用 UMO(统一消息来源标识)格式:`"platform:instance:session_id"`
@@ -1,593 +0,0 @@
# AstrBot SDK 消息事件与组件 API 参考文档
## 概述
本文档详细介绍 `astrbot_sdk` 中消息事件和消息组件的使用方法,包括 `MessageEvent` 类和所有消息组件类。
## 目录
- [MessageEvent - 消息事件对象](#messageevent---消息事件对象)
- [消息组件类](#消息组件类)
- [MessageChain - 消息链](#messagechain---消息链)
- [MessageBuilder - 消息构建器](#messagebuilder---消息构建器)
---
## MessageEvent - 消息事件对象
**模块路径**: `astrbot_sdk.events.MessageEvent`
### 核心属性
| 属性名 | 类型 | 说明 |
|--------|------|------|
| `text` | `str` | 消息文本内容 |
| `user_id` | `str \| None` | 发送者用户 ID |
| `group_id` | `str \| None` | 群组 ID(私聊时为 None) |
| `platform` | `str \| None` | 平台标识(如 "qq", "wechat") |
| `session_id` | `str` | 会话 ID |
| `self_id` | `str` | 机器人账号 ID |
| `platform_id` | `str` | 平台实例标识 |
| `message_type` | `str` | 消息类型("private" 或 "group") |
| `sender_name` | `str` | 发送者昵称 |
### 消息组件访问方法
#### `get_messages()`
获取当前事件的所有 SDK 消息组件。
```python
components = event.get_messages()
for comp in components:
print(f"组件类型: {comp.type}")
```
#### `has_component(type_)`
检查是否包含特定类型的组件。
```python
if event.has_component(Image):
print("消息包含图片")
```
#### `get_components(type_)`
获取特定类型的所有组件。
```python
at_comps = event.get_components(At)
for at in at_comps:
print(f"@了用户: {at.qq}")
```
#### `get_images()`
获取所有图片组件。
```python
images = event.get_images()
for img in images:
path = await img.convert_to_file_path()
print(f"图片路径: {path}")
```
#### `get_files()`
获取所有文件组件。
```python
files = event.get_files()
```
#### `extract_plain_text()`
提取所有纯文本内容。
```python
text = event.extract_plain_text()
```
#### `get_at_users()`
获取消息中所有被@的用户ID列表。
```python
at_users = event.get_at_users()
```
### 会话与平台信息方法
#### `is_private_chat()` / `is_group_chat()`
判断消息类型。
```python
if event.is_private_chat():
await event.reply("这是私聊")
elif event.is_group_chat():
await event.reply("这是群聊")
```
#### `is_admin()`
判断发送者是否有管理员权限。
```python
if event.is_admin():
await event.reply("你是管理员")
```
### 回复与发送方法
#### `reply(text)`
回复纯文本消息。
```python
await event.reply("Hello World!")
```
#### `reply_image(image_url)`
回复图片消息。
```python
await event.reply_image("https://example.com/image.jpg")
```
#### `reply_chain(chain)`
回复消息链。
```python
from astrbot_sdk.message_components import Plain, At
await event.reply_chain([
Plain("Hello "),
At("123456"),
Plain("!")
])
```
### 事件控制方法
#### `stop_event()`
标记事件为已停止,阻止后续处理器执行。
```python
event.stop_event()
```
### 结果构建方法
#### `plain_result(text)`
创建纯文本结果。
```python
return event.plain_result("回复内容")
```
#### `image_result(url_or_path)`
创建图片结果。
```python
return event.image_result("https://example.com/image.jpg")
```
#### `chain_result(chain)`
创建链结果。
```python
return event.chain_result([
Plain("Hello"),
At("123456")
])
```
---
## 消息组件类
### Plain - 纯文本组件
```python
from astrbot_sdk.message_components import Plain
text = Plain("Hello World")
```
### At - @某人组件
```python
from astrbot_sdk.message_components import At
at = At("123456", name="张三")
```
### AtAll - @全体成员组件
```python
from astrbot_sdk.message_components import AtAll
at_all = AtAll()
```
### Image - 图片组件
```python
from astrbot_sdk.message_components import Image
# URL 图片
img1 = Image.fromURL("https://example.com/image.jpg")
# 本地文件
img2 = Image.fromFileSystem("/path/to/image.jpg")
# Base64
img3 = Image.fromBase64("iVBORw0KGgo...")
```
### Record - 语音组件
```python
from astrbot_sdk.message_components import Record
# URL 音频
audio = Record.fromURL("https://example.com/audio.mp3")
# 本地文件
audio = Record.fromFileSystem("/path/to/audio.mp3")
```
### Video - 视频组件
```python
from astrbot_sdk.message_components import Video
video = Video.fromURL("https://example.com/video.mp4")
```
### File - 文件组件
```python
from astrbot_sdk.message_components import File
# URL 文件
file1 = File(name="document.pdf", url="https://example.com/doc.pdf")
# 本地文件
file2 = File(name="image.jpg", file="/path/to/image.jpg")
```
### Reply - 回复组件
```python
from astrbot_sdk.message_components import Reply, Plain
reply = Reply(
id="msg_123",
sender_id="789",
chain=[Plain("被回复的消息")]
)
```
---
## MessageChain - 消息链
### 构造方法
```python
from astrbot_sdk.message_result import MessageChain
from astrbot_sdk.message_components import Plain, At
# 空消息链
chain = MessageChain()
# 带初始组件
chain = MessageChain([Plain("Hello"), At("123456")])
```
### 实例方法
#### `append(component)`
追加单个组件。
```python
chain.append(Plain("More text"))
```
#### `extend(components)`
追加多个组件。
```python
chain.extend([Plain("A"), Plain("B")])
```
#### `to_payload()`
转换为协议 payload。
```python
payload = chain.to_payload()
```
#### `get_plain_text()`
提取纯文本内容。
```python
text = chain.get_plain_text()
```
---
## MessageBuilder - 消息构建器
### 使用示例
```python
from astrbot_sdk.message_result import MessageBuilder
chain = (MessageBuilder()
.text("Hello ")
.at("123456")
.text("!\n")
.image("https://example.com/img.jpg")
.build())
await event.reply_chain(chain)
```
### 可用方法
- `.text(content)` - 添加文本
- `.at(user_id)` - 添加@用户
- `.at_all()` - 添加@全体成员
- `.image(url)` - 添加图片
- `.record(url)` - 添加语音
- `.video(url)` - 添加视频
- `.file(name, url=...)` - 添加文件
- `.build()` - 构建消息链
---
## 使用示例
### 处理图片消息
```python
@on_message()
async def handle_image(event: MessageEvent):
images = event.get_images()
if not images:
await event.reply("消息中没有图片")
return
for img in images:
path = await img.convert_to_file_path()
await event.reply(f"收到图片: {path}")
```
### 检测@和群聊/私聊
```python
@on_command("check")
async def check_handler(event: MessageEvent):
if event.is_group_chat():
await event.reply("这是群聊消息")
elif event.is_private_chat():
await event.reply("这是私聊消息")
at_users = event.get_at_users()
if at_users:
await event.reply(f"你@了: {', '.join(at_users)}")
```
### 返回富文本结果
```python
@on_command("info")
async def info_handler(event: MessageEvent):
return event.chain_result([
Plain(f"用户: {event.sender_name}\n"),
Plain(f"ID: {event.user_id}\n"),
Plain(f"平台: {event.platform}"),
])
```
---
## 媒体辅助类
### MediaHelper
媒体辅助类,提供从 URL 检测媒体类型和下载功能。
```python
from astrbot_sdk.message_components import MediaHelper
```
#### from_url - 从 URL 创建组件
自动检测 URL 的媒体类型并创建对应的消息组件。
```python
from astrbot_sdk.message_components import MediaHelper
# 自动检测媒体类型
component = await MediaHelper.from_url("https://example.com/video.mp4")
# 返回 Video 组件
component = await MediaHelper.from_url("https://example.com/image.jpg")
# 返回 Image 组件
component = await MediaHelper.from_url("https://example.com/audio.mp3")
# 返回 Record 组件
```
**参数**:
- `url`: 媒体文件 URL
- `headers`: 可选的请求头
**返回值**:
- `Image` / `Video` / `Record` / `File` 组件实例
#### download - 下载媒体文件
下载媒体文件到本地。
```python
from astrbot_sdk.message_components import MediaHelper
from pathlib import Path
# 下载到指定目录
path = await MediaHelper.download(
url="https://example.com/video.mp4",
save_dir=Path("/tmp/downloads")
)
print(f"下载到: {path}") # /tmp/downloads/video.mp4
# 下载到当前目录
path = await MediaHelper.download(
url="https://example.com/image.png"
)
```
**参数**:
- `url`: 文件 URL
- `save_dir`: 保存目录(可选,默认为当前目录)
- `filename`: 指定文件名(可选,自动从 URL 或响应头推断)
- `headers`: 请求头(可选)
**返回值**:
- `Path`: 下载文件的本地路径
**示例:完整媒体处理流程**
```python
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command
from astrbot_sdk.message_components import MediaHelper, Plain
class MediaPlugin(Star):
@on_command("download")
async def download_media(self, event: MessageEvent, ctx: Context, url: str):
"""下载媒体文件"""
try:
# 发送下载中提示
await event.reply(f"正在下载: {url}")
# 下载文件
path = await MediaHelper.download(url)
# 创建对应组件并发送
component = await MediaHelper.from_url(url)
component.file = str(path) # 使用本地文件
await event.reply([Plain("下载完成!"), component])
except Exception as e:
await event.reply(f"下载失败: {e}")
@on_command("mirror")
async def mirror_media(self, event: MessageEvent, ctx: Context):
"""转发收到的媒体"""
images = event.get_images()
if images:
for img in images:
# 下载并重新发送
if img.url:
local_path = await MediaHelper.download(img.url)
await event.reply(f"已镜像保存: {local_path}")
```
---
## 未知组件
### UnknownComponent
未知消息组件,用于表示 SDK 无法识别的平台特定组件。
```python
from astrbot_sdk.message_components import UnknownComponent
```
**说明**:
- 当收到 SDK 不支持的消息类型时,会返回此组件
- 保留原始数据供插件自行处理
- 通常出现在新平台或平台新功能中
**属性**:
- `raw_data`: 原始组件数据(dict)
- `type`: 组件类型字符串
```python
@on_message()
async def handle_unknown(self, event: MessageEvent, ctx: Context):
components = event.get_messages()
for comp in components:
if isinstance(comp, UnknownComponent):
ctx.logger.warning(f"未知组件类型: {comp.type}")
ctx.logger.debug(f"原始数据: {comp.raw_data}")
# 插件可以尝试自行处理 raw_data
```
---
## 特殊消息组件
### Forward - 合并转发消息
合并转发消息组件(仅部分平台支持,如 QQ)。
```python
from astrbot_sdk.message_components import Forward, ForwardNode
# 创建转发消息(需要平台支持)
nodes = [
ForwardNode(
user_id="123456",
nickname="用户A",
content=[Plain("消息内容1")]
),
ForwardNode(
user_id="789012",
nickname="用户B",
content=[Plain("消息内容2")]
),
]
forward = Forward(nodes=nodes)
```
**注意**:Forward 组件的支持程度取决于具体平台适配器。
### Poke - 戳一戳/拍一拍
戳一戳消息组件(QQ 等平台支持)。
```python
from astrbot_sdk.message_components import Poke
# 发送戳一戳
poke = Poke(user_id="123456")
await event.reply(poke)
# 检测戳一戳
@on_message()
async def on_poke(self, event: MessageEvent, ctx: Context):
for comp in event.get_messages():
if isinstance(comp, Poke):
await event.reply(f"{event.sender_name} 戳了你一下!")
```
**属性**:
- `user_id`: 被戳的用户 ID
@@ -1,610 +0,0 @@
# AstrBot SDK 装饰器使用指南
## 概述
本文档详细介绍 `astrbot_sdk.decorators` 中所有装饰器的使用方法、参数说明和最佳实践。
## 目录
- [事件触发装饰器](#事件触发装饰器)
- [修饰器装饰器](#修饰器装饰器)
- [过滤器装饰器](#过滤器装饰器)
- [限制器装饰器](#限制器装饰器)
- [能力暴露装饰器](#能力暴露装饰器)
- [LLM 工具装饰器](#llm-工具装饰器)
- [最佳实践](#最佳实践)
---
## 事件触发装饰器
### @on_command
命令触发装饰器。
**签名:**
```python
def on_command(
command: str | Sequence[str],
*,
aliases: list[str] | None = None,
description: str | None = None,
) -> Callable
```
**参数:**
- `command`: 命令名称(不包含前缀符)
- `aliases`: 命令别名列表
- `description`: 命令描述
**示例:**
```python
from astrbot_sdk.decorators import on_command
@on_command("hello")
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
@on_command(["echo", "repeat"], aliases=["say", "speak"])
async def echo(self, event: MessageEvent, text: str):
await event.reply(text)
```
### @on_message
消息触发装饰器。
**签名:**
```python
def on_message(
*,
regex: str | None = None,
keywords: list[str] | None = None,
platforms: list[str] | None = None,
message_types: list[str] | None = None,
) -> Callable
```
**参数:**
- `regex`: 正则表达式模式
- `keywords`: 关键词列表(任一匹配即触发)
- `platforms`: 限定平台列表
- `message_types`: 限定消息类型("group", "private")
**示例:**
```python
# 关键词匹配
@on_message(keywords=["帮助", "help"])
async def help_handler(self, event: MessageEvent, ctx: Context):
await event.reply("可用命令: /hello")
# 正则匹配
@on_message(regex=r"\d{4,}")
async def number_handler(self, event: MessageEvent, ctx: Context):
await event.reply("检测到数字!")
# 多条件过滤
@on_message(
keywords=["天气"],
platforms=["qq"],
message_types=["private"]
)
async def weather_query(self, event: MessageEvent, ctx: Context):
await event.reply("请输入城市名称")
```
### @on_event
事件触发装饰器。
**签名:**
```python
def on_event(event_type: str) -> Callable
```
**示例:**
```python
@on_event("group_member_join")
async def welcome_new_member(self, event, ctx: Context):
await ctx.platform.send(event.group_id, "欢迎新成员!")
```
### @on_schedule
定时任务装饰器。
**签名:**
```python
def on_schedule(
*,
cron: str | None = None,
interval_seconds: int | None = None,
) -> Callable
```
**示例:**
```python
# 固定间隔
@on_schedule(interval_seconds=3600)
async def hourly_check(self, ctx: Context):
pass
# cron 表达式
@on_schedule(cron="0 8 * * *") # 每天 8:00
async def morning_greeting(self, ctx: Context):
await ctx.platform.send("group_123", "早上好!")
```
---
## 修饰器装饰器
### @require_admin
管理员权限装饰器。
**示例:**
```python
from astrbot_sdk.decorators import on_command, require_admin
@on_command("admin")
@require_admin
async def admin_cmd(self, event: MessageEvent, ctx: Context):
await event.reply("管理员命令")
```
---
## 过滤器装饰器
### @platforms
限定平台装饰器。
**签名:**
```python
def platforms(*names: str) -> Callable
```
**示例:**
```python
@on_command("qq_only")
@platforms("qq")
async def qq_only_command(self, event: MessageEvent, ctx: Context):
await event.reply("这是 QQ 专属命令")
```
### @message_types
限定消息类型装饰器。
**签名:**
```python
def message_types(*types: str) -> Callable
```
**示例:**
```python
@on_command("group_only")
@message_types("group")
async def group_command(self, event: MessageEvent, ctx: Context):
await event.reply("这是群聊命令")
```
### @group_only
仅群聊装饰器。
```python
@on_command("group_admin")
@group_only()
async def group_admin_command(self, event: MessageEvent, ctx: Context):
await event.reply("这是群聊管理命令")
```
### @private_only
仅私聊装饰器。
```python
@on_command("private_chat")
@private_only()
async def private_command(self, event: MessageEvent, ctx: Context):
await event.reply("这是私聊命令")
```
---
## 限制器装饰器
### @rate_limit
速率限制装饰器。
**签名:**
```python
def rate_limit(
limit: int,
window: float,
*,
scope: LimiterScope = "session",
behavior: LimiterBehavior = "hint",
message: str | None = None,
) -> Callable
```
**参数:**
- `limit`: 时间窗口内最大调用次数
- `window`: 时间窗口大小(秒)
- `scope`: 限制范围("session", "user", "group", "global")
- `behavior`: 触发限制后的行为("hint", "silent", "error")
**示例:**
```python
@on_command("search")
@rate_limit(5, 60) # 每分钟最多5次
async def search_command(self, event: MessageEvent, ctx: Context):
await event.reply("搜索结果...")
@on_command("draw")
@rate_limit(3, 3600, scope="user") # 每用户每小时3次
async def draw_command(self, event: MessageEvent, ctx: Context):
await event.reply("绘图结果...")
```
### @cooldown
冷却时间装饰器。
**签名:**
```python
def cooldown(
seconds: float,
*,
scope: LimiterScope = "session",
behavior: LimiterBehavior = "hint",
message: str | None = None,
) -> Callable
```
**示例:**
```python
@on_command("cast_skill")
@cooldown(30) # 30秒冷却
async def cast_skill_command(self, event: MessageEvent, ctx: Context):
await event.reply("技能施放成功!")
```
---
### @admin_only
管理员权限装饰器(`@require_admin` 的别名)。
**签名:**
```python
def admin_only(func: HandlerCallable) -> HandlerCallable
```
**示例:**
```python
from astrbot_sdk.decorators import on_command, admin_only
@on_command("admin")
@admin_only
async def admin_cmd(self, event: MessageEvent, ctx: Context):
await event.reply("管理员命令")
```
**说明:**
- 功能与 `@require_admin` 完全相同
- 更简洁的语法,无需括号
- 适合快速标记管理员命令
---
## 优先级装饰器
### @priority
设置 handler 执行优先级。
**签名:**
```python
def priority(value: int) -> Callable[[HandlerCallable], HandlerCallable]
```
**参数:**
- `value`: 优先级数值,**越大越先执行**
- 默认优先级为 0
**示例:**
```python
from astrbot_sdk.decorators import on_command, priority
@on_command("hello")
@priority(10) # 高优先级,先执行
async def hello_high(self, event: MessageEvent, ctx: Context):
await event.reply("高优先级处理器")
@on_command("hello")
@priority(5) # 较低优先级,后执行
async def hello_low(self, event: MessageEvent, ctx: Context):
await event.reply("低优先级处理器")
```
**使用场景:**
- 多个插件注册了相同命令时控制执行顺序
- 确保核心处理器先于扩展处理器执行
- 实现插件间的协作处理链
**注意事项:**
- 相同优先级的 handler 执行顺序不确定
- 高优先级 handler 不会阻止低优先级 handler 执行(除非显式阻止)
---
## 对话装饰器
### @conversation_command
对话命令装饰器,用于创建交互式对话流程。
**签名:**
```python
def conversation_command(
command: str,
*,
timeout: float = 300.0,
description: str | None = None,
) -> Callable
```
**参数:**
- `command`: 命令名称
- `timeout`: 对话超时时间(秒),默认 300
- `description`: 命令描述
**示例:**
```python
from astrbot_sdk.decorators import conversation_command
from astrbot_sdk.conversation import ConversationSession
@conversation_command("survey", timeout=600)
async def survey(self, event: MessageEvent, ctx: Context, session: ConversationSession):
"""交互式调查问卷"""
# 第一轮对话
await event.reply("请输入您的姓名:")
# 等待用户回复(在下一个处理器中处理)
session.state["step"] = "name"
@conversation_command("survey")
async def survey_step2(self, event: MessageEvent, ctx: Context, session: ConversationSession):
"""问卷第二步"""
step = session.state.get("step")
if step == "name":
session.state["name"] = event.text
session.state["step"] = "age"
await event.reply("请输入您的年龄:")
elif step == "age":
session.state["age"] = event.text
# 完成问卷
await event.reply(f"感谢您的参与!姓名:{session.state['name']}, 年龄:{event.text}")
session.close() # 关闭对话会话
```
**工作流程:**
1. 用户发送 `/survey` 触发第一个处理器
2. 处理器使用 `ConversationSession` 维护对话状态
3. 后续消息在同一会话中路由到相同命令的处理器
4. 超时或调用 `session.close()` 结束对话
**异常处理:**
```python
from astrbot_sdk.conversation import ConversationClosed, ConversationReplaced
@conversation_command("demo")
async def demo(self, event: MessageEvent, ctx: Context, session: ConversationSession):
try:
await event.reply("输入 'exit' 结束对话")
if event.text.lower() == "exit":
session.close()
except ConversationClosed:
# 会话被关闭
await event.reply("对话已结束")
except ConversationReplaced:
# 会话被新会话替换
await event.reply("开始新的对话")
```
---
## 能力暴露装饰器
### @provide_capability
暴露能力装饰器。
**签名:**
```python
def provide_capability(
name: str,
*,
description: str,
input_schema: dict[str, Any] | None = None,
output_schema: dict[str, Any] | None = None,
input_model: type[BaseModel] | None = None,
output_model: type[BaseModel] | None = None,
supports_stream: bool = False,
cancelable: bool = False,
) -> Callable
```
**示例:**
```python
from pydantic import BaseModel, Field
from astrbot_sdk.decorators import provide_capability
class CalculateInput(BaseModel):
x: int = Field(description="第一个数")
y: int = Field(description="第二个数")
@provide_capability(
"my_plugin.calculate",
description="执行加法计算",
input_model=CalculateInput
)
async def calculate(self, payload: dict, ctx: Context):
x = payload["x"]
y = payload["y"]
return {"result": x + y}
```
---
## LLM 工具装饰器
### @register_llm_tool
注册 LLM 工具装饰器。
**签名:**
```python
def register_llm_tool(
name: str | None = None,
*,
description: str | None = None,
parameters_schema: dict[str, Any] | None = None,
active: bool = True,
) -> Callable
```
**示例:**
```python
from astrbot_sdk.decorators import register_llm_tool
@register_llm_tool()
async def get_weather(self, city: str, unit: str = "celsius"):
"""获取指定城市的天气信息"""
return f"{city} 的天气: 25°C"
```
### @register_agent
注册 Agent 装饰器。
**签名:**
```python
def register_agent(
name: str,
*,
description: str = "",
tool_names: list[str] | None = None,
) -> Callable
```
**示例:**
```python
from astrbot_sdk.decorators import register_agent
from astrbot_sdk.llm.agents import BaseAgentRunner
@register_agent("my_agent", description="我的智能助手")
class MyAgent(BaseAgentRunner):
async def run(self, ctx: Context, request) -> Any:
return "agent result"
```
---
## 最佳实践
### 1. 装饰器顺序
正确的装饰器顺序很重要:
```python
@on_command("command") # 1. 事件触发装饰器
@platforms("qq") # 2. 过滤器装饰器
@rate_limit(5, 60) # 3. 限制器装饰器
@require_admin # 4. 修饰器装饰器
async def my_handler(self, event: MessageEvent, ctx: Context):
pass
```
### 2. 错误处理
始终实现错误处理:
```python
@on_command("risky_command")
async def risky_handler(self, event: MessageEvent, ctx: Context):
try:
result = await some_risky_operation()
await event.reply(f"成功: {result}")
except Exception as e:
ctx.logger.error(f"操作失败: {e}")
await event.reply("操作失败,请稍后重试")
```
### 3. 类型注解
使用类型注解提高代码可读性:
```python
from typing import Optional
@on_command("greet")
async def greet_handler(
self,
event: MessageEvent,
ctx: Context
) -> None:
await event.reply("Hello!")
```
### 4. 避免常见陷阱
**不要混用冲突的装饰器:**
```python
# 错误
@on_message(platforms=["qq"])
@platforms("wechat") # 冲突!
async def handler(...): pass
# 正确
@on_message(platforms=["qq", "wechat"])
async def handler(...): pass
```
**不要在非消息处理器使用限制器:**
```python
# 错误
@on_event("ready")
@rate_limit(5, 60) # 不支持!
async def handler(...): pass
# 正确
@on_command("cmd")
@rate_limit(5, 60)
async def handler(...): pass
```
@@ -1,528 +0,0 @@
# AstrBot SDK Star 类与生命周期指南
## 概述
`Star` 是 AstrBot v4 SDK 的原生插件基类,提供了完整的插件生命周期管理、上下文访问和能力集成。
## 目录
- [Star 类概述](#star-类概述)
- [生命周期流程](#生命周期流程)
- [生命周期钩子](#生命周期钩子)
- [Context 上下文使用](#context-上下文使用)
- [插件元数据访问](#插件元数据访问)
- [错误处理模式](#错误处理模式)
- [最佳实践](#最佳实践)
---
## Star 类概述
### 什么是 Star 类?
`Star` 是所有 v4 原生插件必须继承的基类,提供插件生命周期管理和能力集成。
### 核心特性
```python
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message
class MyPlugin(Star):
"""插件类示例"""
@on_command("hello")
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
```
---
## 生命周期流程
### 完整生命周期
```
┌─────────────────────────────────────────────────────────────────┐
│ 插件加载阶段 │
├─────────────────────────────────────────────────────────────────┤
│ 1. 插件发现 (discover_plugins) │
│ ├─ 扫描插件目录 │
│ ├─ 读取 plugin.yaml │
│ └─ 验证组件类 (main:MyPlugin) │
│ │
│ 2. 插件加载 │
│ ├─ 动态导入插件模块 │
│ ├─ 实例化 Star 子类 │
│ ├─ 收集 __handlers__ 元组 │
│ └─ 注册装饰器元数据 │
│ │
│ 3. Worker 启动 (PluginWorkerRuntime.start) │
│ ├─ 向 Core 注册 handlers/capabilities │
│ └─ 建立通信对等端 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ 插件运行阶段 │
├─────────────────────────────────────────────────────────────────┤
│ 4. on_start() 生命周期钩子 │
│ ├─ 绑定运行时上下文 │
│ ├─ 调用 on_start(ctx) │
│ └─ 内部调用 initialize() │
│ │
│ 5. Handler 事件循环 │
│ ├─ 等待事件触发 (命令/消息/事件/定时) │
│ ├─ HandlerDispatcher.invoke() │
│ ├─ 创建 Context 和 MessageEvent │
│ ├─ 执行用户 handler │
│ └─ 处理返回值/异常 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ 插件卸载阶段 │
├─────────────────────────────────────────────────────────────────┤
│ 6. on_stop() 生命周期钩子 │
│ ├─ 调用 on_stop(ctx) │
│ ├─ 内部调用 terminate() │
│ ├─ 清理资源 (数据库连接、文件句柄等) │
│ └─ 重置运行时上下文 │
│ │
│ 7. Worker 关闭 │
│ ├─ 发送 finalize 消息给 Core │
│ ├─ 关闭通信传输层 │
│ └─ 退出子进程 │
└─────────────────────────────────────────────────────────────────┘
```
---
## 生命周期钩子
### 1. on_start() - 插件启动钩子
**触发时机**:Worker 启动后,在开始处理事件之前调用
**参数:**
- `ctx: Any | None` - 运行时上下文(通常为 Context 实例)
**用途:**
- 初始化数据库连接
- 加载配置文件
- 注册 LLM 工具
- 启动后台任务
**最佳实践:**
- `on_start()` 里只做初始化、能力注册和轻量状态恢复
- 需要长期保存的应是配置值、句柄、任务引用,不要把 `ctx` 实例长期挂到 `self`
- 如果要和 AstrBot 原生 persona / conversation 协作,优先在这里校验或创建所需资源
**示例:**
```python
class MyPlugin(Star):
async def on_start(self, ctx: Any | None = None) -> None:
"""插件启动时调用"""
await super().on_start(ctx)
# 加载配置
config = await ctx.metadata.get_plugin_config()
self.api_key = config.get("api_key", "")
# 注册 LLM 工具
await ctx.register_llm_tool(
name="search",
parameters_schema={...},
desc="搜索信息",
func_obj=self.search_tool
)
# 启动后台任务
await ctx.register_task(
self.background_sync(),
desc="后台数据同步"
)
```
### 2. on_stop() - 插件停止钩子
**触发时机**:插件卸载或程序关闭前调用
**用途:**
- 关闭数据库连接
- 清理临时文件
- 注销 LLM 工具
- 保存状态数据
**最佳实践:**
- 在 `on_stop()` 中释放 `on_start()` 注册的任务、监听器和外部资源
- 把需要持久化的状态尽量提前落库,不要把关键保存逻辑完全依赖在进程退出瞬间
- 始终把收到的 `ctx` 继续传给 `super().on_stop(ctx)`,不要手动丢掉它
**示例:**
```python
class MyPlugin(Star):
async def on_stop(self, ctx: Any | None = None) -> None:
"""插件停止时调用"""
# 保存状态
await self.put_kv_data("last_shutdown", time.time())
# 确保 terminate 被调用
await super().on_stop(ctx)
```
### 3. initialize() - 初始化钩子
**触发时机**:`on_start()` 内部自动调用
**用途:**
- 插件级别的初始化逻辑
- 不依赖 Context 的初始化
**示例:**
```python
class MyPlugin(Star):
async def initialize(self) -> None:
"""初始化插件"""
self._cache = {}
self._counter = 0
```
### 4. terminate() - 终止钩子
**触发时机**:`on_stop()` 内部自动调用
**用途:**
- 插件级别的清理逻辑
- 不依赖 Context 的清理
**示例:**
```python
class MyPlugin(Star):
async def terminate(self) -> None:
"""清理插件资源"""
self._cache.clear()
self.state = "stopped"
```
### 5. on_error() - 错误处理钩子
**触发时机**:任何 Handler 执行抛出异常时
**参数:**
- `error: Exception` - 捕获的异常
- `event` - 事件对象
- `ctx` - 上下文对象
**示例:**
```python
class MyPlugin(Star):
async def on_error(self, error: Exception, event, ctx) -> None:
"""自定义错误处理"""
from astrbot_sdk.errors import AstrBotError
if isinstance(error, AstrBotError):
await event.reply(error.hint or error.message)
elif isinstance(error, ValueError):
await event.reply(f"参数错误:{error}")
else:
await event.reply(f"发生错误: {type(error).__name__}")
ctx.logger.error(f"Handler error: {error}", exc_info=error)
```
---
## Context 上下文使用
### 在 Handler 中访问
```python
class MyPlugin(Star):
@on_command("test")
async def test_handler(self, event: MessageEvent, ctx: Context):
# Context 通过参数注入
await ctx.db.set("key", "value")
await event.reply("Done")
```
### 在生命周期钩子中访问
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 生命周期钩子中的 Context
config = await ctx.metadata.get_plugin_config()
```
---
## 插件元数据访问
### plugin.yaml 配置
```yaml
_schema_version: 2
name: my_plugin
author: your_name
version: 1.0.0
desc: 我的插件描述
repo: https://github.com/user/repo
logo: logo.png
runtime:
python: "3.12"
components:
- class: main:MyPlugin
support_platforms:
- aiocqhttp
- telegram
astrbot_version: ">=4.13.0,<5.0.0"
```
### StarMetadata 类
插件元数据 dataclass,描述插件的基本信息。
```python
from astrbot_sdk import StarMetadata
@dataclass
class StarMetadata:
name: str # 插件名称(唯一标识)
display_name: str # 显示名称
description: str # 插件描述
author: str # 作者
version: str # 版本号
enabled: bool = True # 是否启用
support_platforms: list[str] # 支持的平台列表
astrbot_version: str | None # 兼容的 AstrBot 版本范围
```
**使用示例:**
```python
from astrbot_sdk import Star, StarMetadata
class MyPlugin(Star):
async def on_start(self, ctx):
# 获取当前插件元数据
metadata: StarMetadata = await ctx.metadata.get_current_plugin()
print(f"插件名称: {metadata.name}")
print(f"显示名称: {metadata.display_name}")
print(f"版本: {metadata.version}")
print(f"作者: {metadata.author}")
print(f"支持平台: {', '.join(metadata.support_platforms)}")
# 检查兼容性
if metadata.astrbot_version:
print(f"兼容版本: {metadata.astrbot_version}")
```
### PluginMetadata 类
`StarMetadata` 的别名,功能完全相同。
```python
from astrbot_sdk import PluginMetadata
# PluginMetadata 是 StarMetadata 的别名
# 两者可以互换使用
metadata: PluginMetadata = await ctx.metadata.get_current_plugin()
```
**建议**:使用 `StarMetadata` 以符合 v4 SDK 的命名规范。
### 访问元数据
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 获取当前插件元数据
my_metadata = await ctx.metadata.get_current_plugin()
print(f"Starting {my_metadata.name} v{my_metadata.version}")
# 获取其他插件元数据
other_metadata = await ctx.metadata.get_plugin("other_plugin")
if other_metadata:
print(f"依赖插件版本: {other_metadata.version}")
```
---
## 错误处理模式
### 标准错误类型
```python
from astrbot_sdk.errors import AstrBotError
# 1. 输入无效错误
raise AstrBotError.invalid_input(
"参数格式错误",
hint="请使用 JSON 格式"
)
# 2. 能力未找到错误
raise AstrBotError.capability_not_found("unknown_capability")
# 3. 网络错误
raise AstrBotError.network_error(
"连接超时",
hint="请检查网络连接"
)
```
### 在 Handler 中捕获错误
```python
class MyPlugin(Star):
@on_command("risky_operation")
async def risky(self, event: MessageEvent, ctx: Context):
try:
result = await self.risky_operation()
await event.reply(f"成功: {result}")
except ValueError as e:
await event.reply(f"参数错误: {e}")
except ConnectionError as e:
ctx.logger.error(f"Network error: {e}")
await event.reply("网络连接失败")
except Exception as e:
ctx.logger.exception("Unexpected error")
raise
```
---
## 最佳实践
### 1. 插件结构
```
my_plugin/
├── plugin.yaml # 插件配置
├── main.py # 主入口
├── handlers/ # 处理器模块
├── utils/ # 工具函数
├── requirements.txt # 可选的 Python 依赖
└── README.md # 说明文档
```
### 2. 插件模板
```python
"""
插件说明
"""
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message
class MyPlugin(Star):
"""插件类"""
async def initialize(self) -> None:
"""初始化"""
self._cache = {}
self._counter = 0
async def on_start(self, ctx) -> None:
"""启动时调用"""
await super().on_start(ctx)
# 加载配置
config = await ctx.metadata.get_plugin_config()
self.setting = config.get("setting", "default")
# 注册工具
await ctx.register_llm_tool(
name="my_tool",
parameters_schema={...},
desc="我的工具",
func_obj=self.my_tool
)
ctx.logger.info(f"{ctx.plugin_id} started")
async def on_stop(self, ctx) -> None:
"""停止时调用"""
# 保存状态
await self.put_kv_data("counter", self._counter)
await super().on_stop(ctx)
ctx.logger.info(f"{ctx.plugin_id} stopped")
@on_command("hello", aliases=["hi"])
async def hello(self, event: MessageEvent, ctx: Context) -> None:
"""打招呼命令"""
await event.reply(f"你好,{event.sender_name}!")
async def my_tool(self, param: str) -> str:
"""LLM 工具实现"""
return f"处理结果: {param}"
```
### 3. 配置管理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 获取配置
config = await ctx.metadata.get_plugin_config()
# 提供默认值
self.timeout = config.get("timeout", 30)
self.max_retries = config.get("max_retries", 3)
self.debug = config.get("debug", False)
# 验证必需配置
if "api_key" not in config:
raise ValueError("缺少必需配置: api_key")
self.api_key = config["api_key"]
```
### 4. 数据持久化
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 加载状态
self.last_update = await self.get_kv_data("last_update", 0)
self.user_data = await self.get_kv_data("users", {})
async def save_state(self):
# 保存状态
await self.put_kv_data("last_update", time.time())
await self.put_kv_data("users", self.user_data)
```
### 5. 资源清理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 创建需要清理的资源
self._session = aiohttp.ClientSession()
self._task = asyncio.create_task(self.background_task())
async def on_stop(self, ctx):
# 清理资源
if hasattr(self, '_task'):
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
if hasattr(self, '_session'):
await self._session.close()
```
@@ -1,437 +0,0 @@
# AstrBot SDK 客户端 API 参考文档
## 概述
本文档详细介绍 `astrbot_sdk/clients/` 目录下所有客户端的 API,包括方法签名、使用示例和注意事项。
## 目录
- [LLMClient - AI 对话客户端](#1-llmclient---ai-对话客户端)
- [MemoryClient - 记忆存储客户端](#2-memoryclient---记忆存储客户端)
- [DBClient - KV 数据库客户端](#3-dbclient---kv-数据库客户端)
- [PlatformClient - 平台消息客户端](#4-platformclient---平台消息客户端)
- [FileServiceClient - 文件服务客户端](#5-fileserviceclient---文件服务客户端)
- [HTTPClient - HTTP API 客户端](#6-httpclient---http-api-客户端)
- [MetadataClient - 插件元数据客户端](#7-metadataclient---插件元数据客户端)
---
## 1. LLMClient - AI 对话客户端
### 导入
```python
from astrbot_sdk.clients import LLMClient, ChatMessage, LLMResponse
```
### 方法
#### chat()
简单对话。
```python
reply = await ctx.llm.chat("你好,介绍一下自己")
```
#### chat_raw()
获取完整响应。
```python
response = await ctx.llm.chat_raw("写一首诗", temperature=0.8)
print(f"Token 使用: {response.usage}")
```
#### stream_chat()
流式对话。
```python
async for chunk in ctx.llm.stream_chat("讲一个故事"):
print(chunk, end="")
```
---
## 2. MemoryClient - 记忆存储客户端
### 导入
```python
from astrbot_sdk.clients import MemoryClient
```
### 方法
#### search()
搜索记忆。默认在有 embedding provider 时执行 hybrid 检索。
```python
results = await ctx.memory.search("用户喜欢什么颜色", mode="hybrid", limit=5)
for item in results:
print(item["key"], item["score"], item["match_type"])
```
#### save()
保存记忆。
```python
await ctx.memory.save("user_pref", {"theme": "dark", "lang": "zh"})
await ctx.memory.save(
"profile:alice",
{"name": "Alice", "embedding_text": "Alice 喜欢蓝色和海边"},
)
```
#### get()
获取记忆。
```python
pref = await ctx.memory.get("user_pref")
```
#### save_with_ttl()
保存带过期时间的记忆。
```python
await ctx.memory.save_with_ttl(
"session_temp",
{"state": "waiting"},
ttl_seconds=3600
)
```
#### delete()
删除记忆。
```python
await ctx.memory.delete("old_note")
```
#### stats()
查看记忆索引状态。
```python
stats = await ctx.memory.stats()
print(stats["total_items"], stats.get("embedded_items"), stats.get("dirty_items"))
```
---
## 3. DBClient - KV 数据库客户端
### 导入
```python
from astrbot_sdk.clients import DBClient
```
### 方法
#### get() / set()
基本读写。
```python
data = await ctx.db.get("user_settings")
await ctx.db.set("user_settings", {"theme": "dark"})
```
#### delete()
删除数据。
```python
await ctx.db.delete("user_settings")
```
#### list()
列出键。
```python
keys = await ctx.db.list("user_")
```
#### get_many() / set_many()
批量操作。
```python
values = await ctx.db.get_many(["user:1", "user:2"])
await ctx.db.set_many({"user:1": {"name": "Alice"}, "user:2": {"name": "Bob"}})
```
#### watch()
监听变更。
```python
async for event in ctx.db.watch("user:"):
print(event["op"], event["key"])
```
---
## 4. PlatformClient - 平台消息客户端
### 导入
```python
from astrbot_sdk.clients import PlatformClient
```
### 方法
#### send()
发送文本消息。
```python
await ctx.platform.send("qq:group:123456", "大家好!")
```
#### send_image()
发送图片。
```python
await ctx.platform.send_image(event.session_id, "https://example.com/image.png")
```
#### send_chain()
发送消息链。
```python
from astrbot_sdk.message_components import Plain, Image
chain = [Plain("文字"), Image(url="https://example.com/img.jpg")]
await ctx.platform.send_chain(event.session_id, chain)
```
#### send_by_id()
通过 ID 发送。
```python
await ctx.platform.send_by_id(
platform_id="qq",
session_id="user123",
content="Hello",
message_type="private"
)
```
#### get_members()
获取群成员。
```python
members = await ctx.platform.get_members("qq:group:123456")
```
---
## 5. FileServiceClient - 文件服务客户端
### 导入
```python
from astrbot_sdk.clients import FileServiceClient
```
### 方法
#### register_file()
注册文件。
```python
token = await ctx.files.register_file("/path/to/file.jpg", timeout=3600)
```
#### handle_file()
解析令牌。
```python
path = await ctx.files.handle_file(token)
```
---
## 6. HTTPClient - HTTP API 客户端
### 导入
```python
from astrbot_sdk.clients import HTTPClient
from astrbot_sdk.decorators import provide_capability
```
### 方法
#### register_api()
注册 API。
```python
@provide_capability(
name="my_plugin.http_handler",
description="处理 HTTP 请求"
)
async def handle_http_request(request_id: str, payload: dict, cancel_token):
return {"status": 200, "body": {"result": "ok"}}
await ctx.http.register_api(
route="/my-api",
handler=handle_http_request,
methods=["GET", "POST"]
)
```
#### unregister_api()
注销 API。
```python
await ctx.http.unregister_api("/my-api")
```
#### list_apis()
列出 API。
```python
apis = await ctx.http.list_apis()
```
---
## 7. MetadataClient - 插件元数据客户端
### 导入
```python
from astrbot_sdk.clients import MetadataClient
```
### 方法
#### get_plugin()
获取插件信息。
```python
plugin = await ctx.metadata.get_plugin("another_plugin")
if plugin:
print(f"插件: {plugin.display_name}")
```
#### list_plugins()
列出所有插件。
```python
plugins = await ctx.metadata.list_plugins()
```
#### get_current_plugin()
获取当前插件。
```python
current = await ctx.metadata.get_current_plugin()
```
#### get_plugin_config()
获取配置。
```python
config = await ctx.metadata.get_plugin_config()
api_key = config.get("api_key")
```
---
## 客户端使用示例
### 1. 基本对话流程
```python
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
reply = await ctx.llm.chat(event.message_content)
await ctx.platform.send(event.session_id, reply)
```
### 2. 带历史的对话
```python
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
history_data = await ctx.memory.get(f"history:{event.session_id}")
history = history_data.get("messages", []) if history_data else []
reply = await ctx.llm.chat(event.message_content, history=history)
history.append(ChatMessage(role="user", content=event.message_content))
history.append(ChatMessage(role="assistant", content=reply))
await ctx.memory.save(f"history:{event.session_id}", {"messages": history})
await ctx.platform.send(event.session_id, reply)
```
### 3. 使用数据库持久化
```python
@on_message()
async def handle_message(event: MessageEvent, ctx: Context):
config = await ctx.db.get(f"user_config:{event.sender_id}")
if not config:
config = {"theme": "light", "lang": "zh"}
await ctx.db.set(f"user_config:{event.sender_id}", config)
reply = f"你的主题设置是: {config['theme']}"
await ctx.platform.send(event.session_id, reply)
```
### 4. 注册 Web API
```python
@provide_capability(
name="my_plugin.get_status",
description="获取插件状态",
)
async def get_status(request_id: str, payload: dict, cancel_token):
return {"status": "running", "version": "1.0.0"}
@on_command("setup_api")
async def setup_api(event: MessageEvent, ctx: Context):
await ctx.http.register_api(
route="/status",
handler=get_status,
methods=["GET"]
)
await ctx.platform.send(event.session_id, "API 已注册")
```
---
## 注意事项
1. 所有客户端方法都是异步的
2. 远程调用可能失败,建议使用 try-except
3. Memory 适合语义搜索,DB 适合精确匹配
4. 文件操作使用 file service 注册令牌
5. 平台标识使用 UMO 格式:`"platform:instance:session_id"`
@@ -1,625 +0,0 @@
# AstrBot SDK 错误处理与调试指南
本文档详细介绍 SDK 中的错误处理机制、错误类型、调试技巧和常见问题解决方案。
## 目录
- [错误处理概述](#错误处理概述)
- [AstrBotError 错误体系](#astrboterror-错误体系)
- [错误码参考](#错误码参考)
- [错误处理模式](#错误处理模式)
- [调试技巧](#调试技巧)
- [常见问题](#常见问题)
---
## 错误处理概述
AstrBot SDK 使用统一的错误体系 `AstrBotError`,支持跨进程传递(通过 to_payload/from_payload 序列化)。
### 错误处理流程
```
1. 运行时抛出 AstrBotError 子类或实例
2. 错误被捕获并序列化为 payload
3. 跨进程传输后反序列化
4. 在 on_error 钩子中统一处理
```
### 基本使用
```python
from astrbot_sdk.errors import AstrBotError, ErrorCodes
# 抛出错误
raise AstrBotError.invalid_input("参数不能为空")
# 捕获并处理
try:
await some_operation()
except AstrBotError as e:
if e.retryable:
# 可重试的错误
await retry()
else:
# 不可重试的错误
await event.reply(e.hint or e.message)
```
---
## AstrBotError 错误体系
### AstrBotError 类
```python
@dataclass(slots=True)
class AstrBotError(Exception):
code: str # 错误码
message: str # 错误消息(面向开发者)
hint: str = "" # 用户提示(面向终端用户)
retryable: bool = False # 是否可重试
docs_url: str = "" # 文档链接
details: dict[str, Any] | None = None # 详细信息
```
### 工厂方法
#### 1. invalid_input - 输入无效错误
**场景**:参数格式错误、缺少必需参数等
```python
raise AstrBotError.invalid_input(
message="参数格式错误",
hint="请使用 JSON 格式",
docs_url="https://docs.example.com/api"
)
```
**属性**:
- `retryable`: False
- 应该在修复输入后重试
#### 2. capability_not_found - 能力未找到
**场景**:调用的 capability 不存在或未注册
```python
raise AstrBotError.capability_not_found("unknown_capability")
```
**属性**:
- `retryable`: False
- 通常是配置或版本不匹配问题
#### 3. network_error - 网络错误
**场景**:连接超时、DNS 解析失败等
```python
raise AstrBotError.network_error(
message="连接超时",
hint="请检查网络连接后重试"
)
```
**属性**:
- `retryable`: True
- 通常可以重试
#### 4. internal_error - 内部错误
**场景**:SDK 或 Core 内部错误
```python
raise AstrBotError.internal_error(
message="数据库连接失败",
hint="请联系插件作者"
)
```
**属性**:
- `retryable`: False
- 需要开发者介入
#### 5. cancelled - 取消错误
**场景**:操作被取消
```python
raise AstrBotError.cancelled("用户取消了操作")
```
**属性**:
- `retryable`: False
#### 6. protocol_version_mismatch - 协议版本不匹配
**场景**:SDK 和 Core 协议版本不兼容
```python
raise AstrBotError.protocol_version_mismatch("协议版本不匹配: v4 vs v5")
```
**属性**:
- `retryable`: False
- 需要升级 SDK 或 Core
---
## 错误码参考
### 不可立即自动重试错误(retryable=False)
这些错误不适合框架做“立刻重试”的自动恢复;其中 `RATE_LIMITED` 和
`COOLDOWN_ACTIVE` 仍然可以在等待窗口结束后由用户或插件重新发起调用。
| 错误码 | 说明 | 处理方式 |
|--------|------|----------|
| `LLM_NOT_CONFIGURED` | LLM 未配置 | 配置 LLM Provider |
| `CAPABILITY_NOT_FOUND` | 能力未找到 | 检查 capability 名称 |
| `PERMISSION_DENIED` | 权限不足 | 检查用户权限 |
| `LLM_ERROR` | LLM 错误 | 查看详细错误信息 |
| `INVALID_INPUT` | 输入无效 | 修正输入参数 |
| `CANCELLED` | 操作被取消 | 无需处理 |
| `PROTOCOL_VERSION_MISMATCH` | 协议版本不匹配 | 升级 SDK |
| `PROTOCOL_ERROR` | 协议错误 | 检查实现 |
| `INTERNAL_ERROR` | 内部错误 | 联系开发者 |
| `RATE_LIMITED` | 速率限制 | 等待速率窗口结束后再重试 |
| `COOLDOWN_ACTIVE` | 冷却中 | 等待冷却结束后再重试 |
### 可重试错误(retryable=True)
| 错误码 | 说明 | 处理方式 |
|--------|------|----------|
| `CAPABILITY_TIMEOUT` | 能力调用超时 | 重试或增加超时时间 |
| `NETWORK_ERROR` | 网络错误 | 重试 |
| `LLM_TEMPORARY_ERROR` | LLM 临时错误 | 重试 |
---
## 对话相关异常
### ConversationClosed
对话已关闭异常。
**场景**:会话被显式关闭或超时时抛出
```python
from astrbot_sdk.conversation import ConversationClosed
@conversation_command("demo")
async def demo_handler(self, event, ctx, session):
try:
# 处理对话...
session.close() # 关闭会话
except ConversationClosed:
await event.reply("对话已结束")
```
**属性**:
- 继承自 `RuntimeError`
- 表示对话会话已结束,无法再接收消息
### ConversationReplaced
对话被替换异常。
**场景**:用户开始新对话,当前对话被替换时抛出
```python
from astrbot_sdk.conversation import ConversationReplaced
@conversation_command("survey")
async def survey_handler(self, event, ctx, session):
try:
# 处理对话...
pass
except ConversationReplaced:
# 用户开始了新对话
await event.reply("已切换到新对话")
```
**属性**:
- 继承自 `RuntimeError`
- 表示当前对话被新对话替换
---
## 错误处理模式
### 模式 1:基本错误处理
```python
@on_command("risky")
async def risky_handler(self, event: MessageEvent, ctx: Context):
try:
result = await risky_operation()
await event.reply(f"成功: {result}")
except AstrBotError as e:
# SDK 错误包含用户友好的提示
await event.reply(e.hint or e.message)
ctx.logger.error(f"操作失败: {e}")
except Exception as e:
# 未知错误
ctx.logger.exception("未知错误")
await event.reply("操作失败,请稍后重试")
```
### 模式 2:分层错误处理
```python
async def fetch_data(ctx: Context, url: str) -> dict:
"""获取数据,处理网络错误"""
try:
return await ctx.http.get(url)
except AstrBotError as e:
if e.code == ErrorCodes.NETWORK_ERROR:
# 网络错误可以重试
ctx.logger.warning(f"网络错误,重试: {e}")
await asyncio.sleep(1)
return await ctx.http.get(url)
raise
@on_command("data")
async def data_handler(self, event: MessageEvent, ctx: Context):
try:
data = await self.fetch_data(ctx, "https://api.example.com/data")
await event.reply(f"数据: {data}")
except AstrBotError as e:
if e.retryable:
await event.reply(f"暂时无法获取数据,请稍后重试")
else:
await event.reply(f"获取数据失败: {e.hint}")
```
### 模式 3:on_error 生命周期钩子
```python
class MyPlugin(Star):
async def on_error(self, error: Exception, event, ctx) -> None:
"""统一错误处理"""
from astrbot_sdk.errors import AstrBotError
if isinstance(error, AstrBotError):
# SDK 错误
if error.code == ErrorCodes.RATE_LIMITED:
await event.reply("操作过于频繁,请稍后再试")
elif error.code == ErrorCodes.PERMISSION_DENIED:
await event.reply("你没有权限执行此操作")
else:
await event.reply(error.hint or "操作失败")
elif isinstance(error, ValueError):
# 参数错误
await event.reply(f"参数错误: {error}")
else:
# 未知错误
ctx.logger.exception("未处理的错误")
await event.reply("发生未知错误,请联系管理员")
```
### 模式 4:重试机制
```python
from astrbot_sdk.errors import AstrBotError, ErrorCodes
async def with_retry(
operation,
max_retries: int = 3,
delay: float = 1.0
):
"""带重试的操作"""
last_error = None
for attempt in range(max_retries):
try:
return await operation()
except AstrBotError as e:
last_error = e
if not e.retryable:
raise # 不可重试错误直接抛出
ctx.logger.warning(f"第 {attempt + 1} 次尝试失败: {e}")
if attempt < max_retries - 1:
await asyncio.sleep(delay * (attempt + 1)) # 指数退避
raise last_error
# 使用
@on_command("fetch")
async def fetch_handler(self, event: MessageEvent, ctx: Context):
try:
result = await with_retry(
lambda: ctx.llm.chat("生成内容"),
max_retries=3
)
await event.reply(result)
except AstrBotError as e:
await event.reply(f"请求失败: {e.hint}")
```
### 模式 5:取消处理
```python
@on_command("long_task")
async def long_task_handler(self, event: MessageEvent, ctx: Context):
try:
for i in range(100):
# 检查是否取消
ctx.cancel_token.raise_if_cancelled()
await do_work(i)
await asyncio.sleep(0.1)
await event.reply("任务完成")
except asyncio.CancelledError:
await event.reply("任务已取消")
raise # 重新抛出以便框架处理
except AstrBotError as e:
if e.code == ErrorCodes.CANCELLED:
await event.reply("操作已取消")
else:
raise
```
---
## 调试技巧
### 1. 启用详细日志
```python
# 在插件中记录详细日志
@on_command("debug")
async def debug_handler(self, event: MessageEvent, ctx: Context):
ctx.logger.debug(f"收到消息: {event.text}")
ctx.logger.debug(f"用户ID: {event.user_id}")
ctx.logger.debug(f"会话ID: {event.session_id}")
ctx.logger.debug(f"平台: {event.platform}")
# 记录组件信息
components = event.get_messages()
for comp in components:
ctx.logger.debug(f"组件: {comp.type} - {comp}")
```
### 2. 使用测试框架调试
```python
from astrbot_sdk.testing import PluginTestHarness
async def test_with_debug():
harness = PluginTestHarness()
plugin = harness.load_plugin("my_plugin.main:MyPlugin")
# 启用详细日志
harness.enable_debug_logging()
# 模拟事件
result = await harness.simulate_command("/hello")
print(f"结果: {result}")
# 查看调用历史
for call in harness.get_call_history():
print(f"调用: {call}")
```
### 3. 使用 PDB 调试
```python
import pdb
@on_command("debug")
async def debug_handler(self, event: MessageEvent, ctx: Context):
# 设置断点
pdb.set_trace()
result = await ctx.llm.chat("测试")
await event.reply(result)
```
### 4. 记录完整错误信息
```python
import traceback
@on_command("risky")
async def risky_handler(self, event: MessageEvent, ctx: Context):
try:
result = await risky_operation()
await event.reply(f"成功: {result}")
except Exception as e:
# 记录完整堆栈
ctx.logger.error(f"错误: {e}")
ctx.logger.error(f"堆栈: {traceback.format_exc()}")
# 发送简化信息给用户
await event.reply("操作失败,请查看日志")
```
### 5. 使用 Context 的 cancel_token 调试
```python
@on_command("timeout_test")
async def timeout_test(self, event: MessageEvent, ctx: Context):
ctx.logger.info(f"取消状态: {ctx.cancel_token.cancelled}")
try:
# 长时间运行的操作
for i in range(10):
ctx.logger.debug(f"步骤 {i}, 取消状态: {ctx.cancel_token.cancelled}")
ctx.cancel_token.raise_if_cancelled()
await asyncio.sleep(1)
await event.reply("完成")
except asyncio.CancelledError:
ctx.logger.info("操作被取消")
raise
```
---
## 常见问题
### Q1: 如何处理 "CAPABILITY_NOT_FOUND" 错误?
**原因**:调用的 capability 不存在或未注册
**解决方案**:
```python
# 检查 Core 版本是否支持
# 确认 capability 名称正确
# 检查插件是否正确加载
try:
result = await ctx._proxy.call("unknown.capability", {})
except AstrBotError as e:
if e.code == ErrorCodes.CAPABILITY_NOT_FOUND:
ctx.logger.error("当前 AstrBot 版本不支持此功能")
await event.reply("请升级 AstrBot 到最新版本")
```
### Q2: 如何处理速率限制?
**解决方案**:
```python
from astrbot_sdk.errors import ErrorCodes
@on_command("api_call")
async def api_call_handler(self, event: MessageEvent, ctx: Context):
try:
result = await call_api()
await event.reply(result)
except AstrBotError as e:
if e.code == ErrorCodes.RATE_LIMITED:
# 获取重试时间(如果有)
retry_after = e.details.get("retry_after", 60)
await event.reply(f"操作过于频繁,请 {retry_after} 秒后再试")
else:
raise
```
### Q3: 如何区分用户错误和系统错误?
**解决方案**:
```python
@on_command("process")
async def process_handler(self, event: MessageEvent, ctx: Context):
try:
result = await process(event.text)
await event.reply(result)
except AstrBotError as e:
if e.code in {
ErrorCodes.INVALID_INPUT,
ErrorCodes.PERMISSION_DENIED
}:
# 用户错误,直接提示
await event.reply(e.hint or e.message)
else:
# 系统错误,记录并提示
ctx.logger.error(f"系统错误: {e}")
await event.reply("系统错误,请稍后重试")
```
### Q4: 如何在 on_error 中避免无限循环?
**注意**:如果 `on_error` 中抛出异常,会导致递归调用
**解决方案**:
```python
class MyPlugin(Star):
async def on_error(self, error: Exception, event, ctx) -> None:
try:
# 错误处理逻辑
await event.reply("发生错误")
except Exception as e:
# 避免递归,只记录不回复
ctx.logger.exception("on_error 失败")
```
### Q5: 如何调试跨进程通信问题?
**解决方案**:
```python
# 启用 SDK 调试日志
import logging
logging.getLogger("astrbot_sdk").setLevel(logging.DEBUG)
# 在关键位置添加日志
@on_command("debug_comm")
async def debug_comm_handler(self, event: MessageEvent, ctx: Context):
ctx.logger.debug("开始调用 capability")
try:
result = await ctx._proxy.call("test.capability", {"key": "value"})
ctx.logger.debug(f"调用成功: {result}")
except Exception as e:
ctx.logger.error(f"调用失败: {e}")
raise
```
---
## 最佳实践
### 1. 始终处理可重试错误
```python
# 好的做法
async def reliable_operation(ctx):
max_retries = 3
for i in range(max_retries):
try:
return await ctx.llm.chat("prompt")
except AstrBotError as e:
if e.retryable and i < max_retries - 1:
await asyncio.sleep(2 ** i) # 指数退避
else:
raise
```
### 2. 提供用户友好的错误提示
```python
# 好的做法
try:
result = await operation()
except AstrBotError as e:
# 使用 SDK 提供的 hint
await event.reply(e.hint or "操作失败,请稍后重试")
```
### 3. 区分日志级别
```python
# 好的做法
try:
result = await operation()
except AstrBotError as e:
if e.retryable:
ctx.logger.warning(f"临时错误: {e}")
else:
ctx.logger.error(f"严重错误: {e}")
```
### 4. 在 on_stop 中处理清理错误
```python
class MyPlugin(Star):
async def on_stop(self, ctx):
try:
await self.cleanup()
except Exception as e:
# 清理错误不应阻止停止流程
ctx.logger.error(f"清理失败: {e}")
```
---
## 相关文档
- [Context API 参考](./01_context_api.md)
- [Star 类与生命周期](./04_star_lifecycle.md)
- [高级主题](./07_advanced_topics.md)
@@ -1,575 +0,0 @@
# AstrBot SDK 高级主题
本文档介绍 AstrBot SDK 的高级用法,包括并发处理、性能优化、安全最佳实践和架构设计。
## 目录
- [并发处理](#并发处理)
- [性能优化](#性能优化)
- [安全最佳实践](#安全最佳实践)
- [架构设计模式](#架构设计模式)
- [高级客户端用法](#高级客户端用法)
---
## 并发处理
### asyncio 基础
SDK 完全基于 asyncio 构建,所有操作都是异步的。
```python
import asyncio
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("concurrent")
async def concurrent_handler(self, event: MessageEvent, ctx: Context):
# 并发执行多个操作
tasks = [
ctx.llm.chat("任务1"),
ctx.llm.chat("任务2"),
ctx.llm.chat("任务3"),
]
results = await asyncio.gather(*tasks, return_exceptions=True)
for i, result in enumerate(results):
if isinstance(result, Exception):
await event.reply(f"任务{i+1}失败: {result}")
else:
await event.reply(f"任务{i+1}结果: {result}")
```
### 并发限制
避免同时发起过多请求:
```python
import asyncio
from asyncio import Semaphore
class MyPlugin(Star):
def __init__(self):
# 限制并发数
self._semaphore = Semaphore(5)
async def limited_operation(self, ctx, prompt):
async with self._semaphore:
return await ctx.llm.chat(prompt)
@on_command("batch")
async def batch_handler(self, event: MessageEvent, ctx: Context):
prompts = ["任务1", "任务2", "任务3", "任务4", "任务5"]
# 使用 semaphore 限制并发
tasks = [self.limited_operation(ctx, p) for p in prompts]
results = await asyncio.gather(*tasks, return_exceptions=True)
await event.reply(f"完成 {len(results)} 个任务")
```
### 取消处理
正确处理操作取消:
```python
@on_command("cancelable")
async def cancelable_handler(self, event: MessageEvent, ctx: Context):
try:
# 长时间运行的操作
for i in range(100):
# 检查是否被取消
ctx.cancel_token.raise_if_cancelled()
await asyncio.sleep(0.1)
if i % 10 == 0:
await event.reply(f"进度: {i}%")
await event.reply("完成!")
except asyncio.CancelledError:
await event.reply("操作已取消")
raise # 重新抛出以便框架处理
```
### 锁和同步
保护共享资源:
```python
import asyncio
class MyPlugin(Star):
def __init__(self):
self._lock = asyncio.Lock()
self._counter = 0
async def increment(self):
async with self._lock:
# 临界区
current = self._counter
await asyncio.sleep(0.1) # 模拟操作
self._counter = current + 1
return self._counter
@on_command("count")
async def count_handler(self, event: MessageEvent, ctx: Context):
count = await self.increment()
await event.reply(f"当前计数: {count}")
```
---
## 性能优化
### 1. 连接池
复用 HTTP 连接:
```python
import aiohttp
class MyPlugin(Star):
async def on_start(self, ctx):
# 创建连接池
self._session = aiohttp.ClientSession(
connector=aiohttp.TCPConnector(limit=100, limit_per_host=20)
)
async def on_stop(self, ctx):
await self._session.close()
async def fetch_data(self, url):
# 复用连接
async with self._session.get(url) as response:
return await response.json()
```
### 2. 缓存策略
使用内存缓存减少重复计算:
```python
from functools import lru_cache
import asyncio
class MyPlugin(Star):
def __init__(self):
self._cache = {}
self._cache_lock = asyncio.Lock()
async def get_cached_data(self, key, ttl=300):
async with self._cache_lock:
if key in self._cache:
data, timestamp = self._cache[key]
if asyncio.get_event_loop().time() - timestamp < ttl:
return data
# 从数据库获取
data = await self.fetch_from_db(key)
async with self._cache_lock:
self._cache[key] = (data, asyncio.get_event_loop().time())
return data
async def invalidate_cache(self, key):
async with self._cache_lock:
self._cache.pop(key, None)
```
### 3. 批处理
批量操作减少网络往返:
```python
@on_command("batch_db")
async def batch_db_handler(self, event: MessageEvent, ctx: Context):
# 批量获取
keys = [f"user:{i}" for i in range(100)]
values = await ctx.db.get_many(keys)
# 批量设置
updates = {f"user:{i}": {"updated": True} for i in range(100)}
await ctx.db.set_many(updates)
await event.reply(f"更新了 {len(updates)} 条记录")
```
### 4. 流式处理
使用流式 API 处理大数据:
```python
@on_command("stream")
async def stream_handler(self, event: MessageEvent, ctx: Context):
# 流式 LLM 响应
message = await event.reply("正在生成...")
full_text = ""
async for chunk in ctx.llm.stream_chat("写一个很长的故事"):
full_text += chunk
# 每 100 个字符更新一次
if len(full_text) % 100 < 10:
await message.edit(full_text + "...")
await message.edit(full_text)
```
### 5. 懒加载
延迟初始化资源:
```python
class MyPlugin(Star):
def __init__(self):
self._expensive_resource = None
self._resource_lock = asyncio.Lock()
async def get_resource(self):
if self._expensive_resource is None:
async with self._resource_lock:
if self._expensive_resource is None:
# 昂贵的初始化
self._expensive_resource = await self.init_resource()
return self._expensive_resource
```
---
## 安全最佳实践
### 1. 输入验证
始终验证用户输入:
```python
import re
from astrbot_sdk.errors import AstrBotError
@on_command("search")
async def search_handler(self, event: MessageEvent, ctx: Context, query: str):
# 验证输入长度
if len(query) > 1000:
raise AstrBotError.invalid_input("查询过长,最多 1000 字符")
# 验证输入内容
if not re.match(r'^[\w\s\-]+$', query):
raise AstrBotError.invalid_input("查询包含非法字符")
# 执行搜索
result = await self.search(query)
await event.reply(result)
```
### 2. 防止注入攻击
```python
# 危险的代码
# await ctx.db.set(f"user:{event.user_id}", eval(user_input))
# 安全的代码
import json
@on_command("save")
async def save_handler(self, event: MessageEvent, ctx: Context, data: str):
try:
# 使用 JSON 解析而不是 eval
parsed = json.loads(data)
await ctx.db.set(f"user:{event.user_id}", parsed)
except json.JSONDecodeError:
raise AstrBotError.invalid_input("无效的 JSON 格式")
```
### 3. 敏感信息处理
```python
import os
class MyPlugin(Star):
async def on_start(self, ctx):
config = await ctx.metadata.get_plugin_config()
# 从配置或环境变量获取敏感信息
self.api_key = config.get("api_key") or os.getenv("MY_PLUGIN_API_KEY")
if not self.api_key:
raise ValueError("缺少 API Key")
# 不要在日志中打印敏感信息
ctx.logger.info("API Key 已配置")
# 不要: ctx.logger.info(f"API Key: {self.api_key}")
```
### 4. 权限检查
```python
from astrbot_sdk.decorators import require_admin
class MyPlugin(Star):
@on_command("admin_only")
@require_admin
async def admin_only(self, event: MessageEvent, ctx: Context):
await event.reply("管理员命令执行成功")
async def check_permission(self, event, required_role):
# 自定义权限检查
if not event.is_admin() and required_role == "admin":
raise AstrBotError.invalid_input("需要管理员权限")
```
### 5. 速率限制
```python
from astrbot_sdk.decorators import rate_limit
class MyPlugin(Star):
@on_command("expensive")
@rate_limit(
limit=5,
window=3600,
scope="user",
message="每小时只能调用 5 次"
)
async def expensive_operation(self, event: MessageEvent, ctx: Context):
# 昂贵的操作
result = await ctx.llm.chat("复杂任务", model="gpt-4")
await event.reply(result)
```
---
## 架构设计模式
### 1. 分层架构
```
my_plugin/
├── __init__.py
├── main.py # 插件入口
├── handlers/ # 处理器层
│ ├── __init__.py
│ ├── commands.py # 命令处理器
│ └── messages.py # 消息处理器
├── services/ # 业务逻辑层
│ ├── __init__.py
│ ├── user_service.py
│ └── data_service.py
├── models/ # 数据模型层
│ ├── __init__.py
│ └── user.py
└── utils/ # 工具层
├── __init__.py
└── helpers.py
```
### 2. 依赖注入
```python
class UserService:
def __init__(self, ctx: Context):
self._ctx = ctx
async def get_user(self, user_id: str):
return await self._ctx.db.get(f"user:{user_id}")
class MyPlugin(Star):
async def on_start(self, ctx):
# 注入依赖
self._user_service = UserService(ctx)
@on_command("profile")
async def profile_handler(self, event: MessageEvent, ctx: Context):
user = await self._user_service.get_user(event.user_id)
await event.reply(f"用户信息: {user}")
```
### 3. 事件驱动架构
```python
class MyPlugin(Star):
def __init__(self):
self._event_handlers = {}
def register_handler(self, event_type, handler):
if event_type not in self._event_handlers:
self._event_handlers[event_type] = []
self._event_handlers[event_type].append(handler)
async def emit_event(self, event_type, data):
handlers = self._event_handlers.get(event_type, [])
for handler in handlers:
try:
await handler(data)
except Exception as e:
self.logger.error(f"事件处理失败: {e}")
```
### 4. 状态机模式
```python
from enum import Enum, auto
class ConversationState(Enum):
IDLE = auto()
WAITING_INPUT = auto()
PROCESSING = auto()
class MyPlugin(Star):
def __init__(self):
self._states = {}
async def get_state(self, session_id):
return self._states.get(session_id, ConversationState.IDLE)
async def set_state(self, session_id, state):
self._states[session_id] = state
@on_message()
async def handle_message(self, event: MessageEvent, ctx: Context):
state = await self.get_state(event.session_id)
if state == ConversationState.IDLE:
await self.handle_idle(event, ctx)
elif state == ConversationState.WAITING_INPUT:
await self.handle_waiting(event, ctx)
```
---
## 高级客户端用法
### 1. ProviderManagerClient
```python
from astrbot_sdk import Star, Context
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("switch_provider")
async def switch_provider(self, event: MessageEvent, ctx: Context):
# 列出所有 Provider
providers = await ctx.provider_manager.get_insts()
# 切换 Provider
await ctx.provider_manager.set_provider(
provider_id="gpt-4",
provider_type="chat_completion"
)
# 监听 Provider 变更
async for change in ctx.provider_manager.watch_changes():
ctx.logger.info(f"Provider 变更: {change.provider_id}")
```
### 2. 平台管理
```python
@on_command("platform_info")
async def platform_info(self, event: MessageEvent, ctx: Context):
# 获取平台实例
platform = await ctx.get_platform_inst("qq:instance1")
if platform:
await platform.refresh()
await event.reply(
f"平台: {platform.name}\n"
f"状态: {platform.status}\n"
f"错误数: {len(platform.errors)}"
)
```
### 3. 高级 LLM 用法
```python
from astrbot_sdk.llm.entities import ProviderRequest
@on_command("advanced_llm")
async def advanced_llm(self, event: MessageEvent, ctx: Context):
# 使用 ProviderRequest 进行精细控制
request = ProviderRequest(
prompt="生成内容",
system_prompt="你是一个助手",
temperature=0.7,
max_tokens=2000
)
# 使用工具循环 Agent
response = await ctx.tool_loop_agent(
request=request,
tool_names=["search", "calculate"]
)
await event.reply(response.text)
```
### 4. 会话管理
```python
from astrbot_sdk.conversation import ConversationSession
@on_command("conversation")
async def conversation_handler(self, event: MessageEvent, ctx: Context):
# 创建会话
session = ConversationSession(
session_id=event.session_id,
conversation_id="conv_123"
)
# 使用会话上下文
async with session:
await session.send("开始对话")
response = await session.receive()
await session.send(f"收到: {response}")
```
---
## 性能监控
### 1. 添加性能指标
```python
import time
class MyPlugin(Star):
async def monitored_operation(self, operation, *args, **kwargs):
start = time.time()
try:
result = await operation(*args, **kwargs)
return result
finally:
duration = time.time() - start
self.logger.info(f"操作耗时: {duration:.2f}s")
@on_command("slow")
async def slow_handler(self, event: MessageEvent, ctx: Context):
result = await self.monitored_operation(
ctx.llm.chat,
"复杂查询"
)
await event.reply(result)
```
### 2. 内存监控
```python
import sys
import gc
class MyPlugin(Star):
def log_memory_usage(self):
# 获取内存使用
gc.collect()
objects = gc.get_objects()
self.logger.debug(f"当前对象数: {len(objects)}")
```
---
## 相关文档
- [错误处理与调试](./06_error_handling.md)
- [测试指南](./08_testing_guide.md)
- [安全检查清单](./11_security_checklist.md)
@@ -1,609 +0,0 @@
# AstrBot SDK 测试指南
本文档介绍如何测试 AstrBot SDK 插件,包括单元测试、集成测试和使用测试框架。
## 目录
- [测试概述](#测试概述)
- [测试框架](#测试框架)
- [单元测试](#单元测试)
- [集成测试](#集成测试)
- [Mock 使用](#mock-使用)
- [测试最佳实践](#测试最佳实践)
---
## 测试概述
### 为什么需要测试?
1. **确保功能正确性**:验证插件按预期工作
2. **防止回归**:修改代码时不破坏现有功能
3. **文档化**:测试用例展示了如何使用代码
4. **提高信心**:放心地重构和优化代码
### 测试类型
```
单元测试 ──→ 集成测试 ──→ 端到端测试
(最快) (中等) (最慢)
```
---
## 测试框架
### 安装测试依赖
```bash
pip install pytest pytest-asyncio pytest-cov
```
### 配置 pytest
```python
# conftest.py
import pytest
from astrbot_sdk.testing import PluginTestHarness
@pytest.fixture
async def harness():
"""提供测试 harness"""
h = PluginTestHarness()
yield h
await h.cleanup()
@pytest.fixture
async def plugin(harness):
"""加载插件"""
return await harness.load_plugin("my_plugin.main:MyPlugin")
```
---
## 单元测试
### 测试命令处理器
```python
import pytest
from astrbot_sdk.testing import PluginTestHarness
@pytest.mark.asyncio
async def test_hello_command():
"""测试 hello 命令"""
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 模拟命令调用
result = await harness.simulate_command("/hello")
# 验证结果
assert result.text == "Hello, World!"
await harness.cleanup()
```
### 测试消息处理器
```python
@pytest.mark.asyncio
async def test_message_handler():
"""测试消息处理器"""
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 模拟消息
result = await harness.simulate_message(
text="你好",
user_id="12345",
session_id="session_1"
)
# 验证响应
assert "你好" in result.text
await harness.cleanup()
```
### 测试装饰器
```python
@pytest.mark.asyncio
async def test_rate_limit():
"""测试速率限制"""
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 第一次调用应该成功
result1 = await harness.simulate_command("/limited")
assert result1.success
# 快速第二次调用应该被限制
result2 = await harness.simulate_command("/limited")
assert result2.error.code == "rate_limited"
await harness.cleanup()
```
---
## 集成测试
### 测试数据库操作
```python
@pytest.mark.asyncio
async def test_database_operations():
"""测试数据库操作"""
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 模拟事件以获取 ctx
event = harness.create_mock_event(text="test")
# 设置数据
await plugin.save_user_data(
event,
event.ctx,
user_id="123",
data={"name": "Alice"}
)
# 读取数据
data = await plugin.get_user_data(
event,
event.ctx,
user_id="123"
)
assert data["name"] == "Alice"
await harness.cleanup()
```
### 测试 LLM 调用
```python
@pytest.mark.asyncio
async def test_llm_integration():
"""测试 LLM 调用"""
harness = PluginTestHarness()
# 配置 mock LLM 响应
harness.mock_llm_response("模拟的 LLM 回复")
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 调用需要 LLM 的命令
result = await harness.simulate_command("/ask 问题")
assert "模拟的 LLM 回复" in result.text
await harness.cleanup()
```
### 测试平台发送
```python
@pytest.mark.asyncio
async def test_platform_send():
"""测试平台消息发送"""
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
# 模拟命令
await harness.simulate_command("/broadcast 大家好")
# 验证发送记录
sent_messages = harness.get_sent_messages()
assert len(sent_messages) >= 1
assert "大家好" in sent_messages[0].text
await harness.cleanup()
```
---
## Mock 使用
### Mock Context
```python
from unittest.mock import AsyncMock, MagicMock
from astrbot_sdk import Context
@pytest.fixture
def mock_ctx():
"""创建 mock Context"""
ctx = MagicMock(spec=Context)
# Mock LLM 客户端
ctx.llm = AsyncMock()
ctx.llm.chat.return_value = "Mocked response"
# Mock DB 客户端
ctx.db = AsyncMock()
ctx.db.get.return_value = {"key": "value"}
# Mock Logger
ctx.logger = MagicMock()
return ctx
@pytest.mark.asyncio
async def test_with_mock_ctx(mock_ctx):
"""使用 mock Context 测试"""
plugin = MyPlugin()
result = await plugin.some_method(mock_ctx)
# 验证调用
mock_ctx.llm.chat.assert_called_once()
assert result == "expected"
```
### Mock 事件
```python
from astrbot_sdk import MessageEvent
@pytest.fixture
def mock_event():
"""创建 mock 事件"""
event = MagicMock(spec=MessageEvent)
event.text = "测试消息"
event.user_id = "12345"
event.session_id = "session_1"
event.platform = "qq"
# Mock 回复方法
event.reply = AsyncMock()
return event
@pytest.mark.asyncio
async def test_with_mock_event(mock_event, mock_ctx):
"""使用 mock 事件测试"""
plugin = MyPlugin()
await plugin.handle_message(mock_event, mock_ctx)
# 验证回复
mock_event.reply.assert_called_once()
```
### Mock 时间
```python
import time
from unittest.mock import patch
@pytest.mark.asyncio
async def test_with_mock_time():
"""使用 mock 时间测试"""
with patch('time.time', return_value=1234567890):
result = await plugin.time_sensitive_operation()
assert result.timestamp == 1234567890
```
### Mock 外部 API
```python
import aiohttp
from aioresponses import aioresponses
@pytest.mark.asyncio
async def test_external_api():
"""测试外部 API 调用"""
with aioresponses() as mocked:
# Mock API 响应
mocked.get(
'https://api.example.com/data',
payload={'result': 'success'},
status=200
)
result = await plugin.fetch_external_data()
assert result['result'] == 'success'
```
---
## 测试最佳实践
### 1. 测试命名规范
```python
# 好的命名
def test_calculate_sum_with_positive_numbers():
"""测试正数相加"""
pass
def test_calculate_sum_with_negative_numbers():
"""测试负数相加"""
pass
# 不好的命名
def test1():
pass
def test_sum():
pass
```
### 2. 一个测试一个概念
```python
# 好的做法:每个测试一个断言
def test_user_creation():
user = create_user("alice")
assert user.name == "alice"
def test_user_creation_sets_default_role():
user = create_user("alice")
assert user.role == "user"
# 不好的做法:多个概念混在一起
def test_user():
user = create_user("alice")
assert user.name == "alice"
assert user.role == "user"
assert user.created_at is not None
```
### 3. 使用 Fixtures
```python
# conftest.py
import pytest
@pytest.fixture
def sample_user_data():
"""提供测试用户数据"""
return {
"user_id": "123",
"name": "Alice",
"email": "alice@example.com"
}
@pytest.fixture
async def initialized_plugin():
"""提供已初始化的插件"""
plugin = MyPlugin()
harness = PluginTestHarness()
await plugin.on_start(harness.create_mock_ctx())
yield plugin
await plugin.on_stop(None)
# 测试中使用
def test_with_fixture(sample_user_data, initialized_plugin):
result = initialized_plugin.process_user(sample_user_data)
assert result.success
```
### 4. 参数化测试
```python
import pytest
@pytest.mark.parametrize("input,expected", [
("hello", "Hello"),
("world", "World"),
("", ""),
])
def test_capitalize(input, expected):
assert input.capitalize() == expected
@pytest.mark.asyncio
@pytest.mark.parametrize("command,expected_response", [
("/help", "可用命令..."),
("/about", "关于信息..."),
("/version", "版本号..."),
])
async def test_commands(command, expected_response):
harness = PluginTestHarness()
plugin = await harness.load_plugin("my_plugin.main:MyPlugin")
result = await harness.simulate_command(command)
assert expected_response in result.text
```
### 5. 测试隔离
```python
# 每个测试使用独立的数据
@pytest.fixture(autouse=True)
def reset_state():
"""每个测试前重置状态"""
MyPlugin._instance_counter = 0
yield
# 测试后清理
MyPlugin._instance_counter = 0
@pytest.mark.asyncio
async def test_isolated():
# 这个测试不会受其他测试影响
plugin = MyPlugin()
assert plugin.id == 1
```
### 6. 异步测试模式
```python
import asyncio
import pytest
@pytest.mark.asyncio
async def test_async_operation():
"""测试异步操作"""
result = await async_function()
assert result == expected
@pytest.mark.asyncio
async def test_async_timeout():
"""测试超时"""
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(
slow_function(),
timeout=0.1
)
@pytest.mark.asyncio
async def test_async_exception():
"""测试异常"""
with pytest.raises(ValueError) as exc_info:
await function_that_raises()
assert "expected error" in str(exc_info.value)
```
### 7. 覆盖率检查
```bash
# 运行测试并生成覆盖率报告
pytest --cov=my_plugin --cov-report=html
# 检查覆盖率
pytest --cov=my_plugin --cov-fail-under=80
```
```ini
# .coveragerc
[run]
source = my_plugin
omit =
*/tests/*
*/venv/*
*/__pycache__/*
[report]
exclude_lines =
pragma: no cover
def __repr__
raise NotImplementedError
```
---
## 测试工具函数
### 常用测试辅助函数
```python
# test_utils.py
import asyncio
from contextlib import asynccontextmanager
async def run_with_timeout(coro, timeout=5):
"""带超时运行协程"""
return await asyncio.wait_for(coro, timeout=timeout)
@asynccontextmanager
async def temporary_database():
"""临时数据库上下文"""
db = await create_test_db()
try:
yield db
finally:
await db.cleanup()
def create_test_event(**kwargs):
"""创建测试事件"""
defaults = {
"text": "test",
"user_id": "12345",
"session_id": "test_session",
"platform": "qq",
}
defaults.update(kwargs)
return MockEvent(**defaults)
```
---
## 持续集成
### GitHub Actions 配置
```yaml
# .github/workflows/test.yml
name: Tests
on: [push, pull_request]
jobs:
test:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v3
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.12'
- name: Install dependencies
run: |
pip install -r requirements.txt
pip install -r requirements-dev.txt
- name: Run tests
run: |
pytest --cov=my_plugin --cov-report=xml
- name: Upload coverage
uses: codecov/codecov-action@v3
```
---
## 调试测试
### 使用 pdb
```python
import pytest
import pdb
def test_with_debug():
result = some_function()
# 设置断点
pdb.set_trace()
assert result.success
```
### 使用 pytest 的 --pdb
```bash
# 失败时自动进入 pdb
pytest --pdb
# 在第一个失败时停止
pytest -x --pdb
```
### 详细输出
```bash
# 详细输出
pytest -v
# 最详细输出
pytest -vv
# 显示 print 输出
pytest -s
```
---
## 相关文档
- [错误处理与调试](./06_error_handling.md)
- [高级主题](./07_advanced_topics.md)
@@ -1,34 +0,0 @@
# AstrBot SDK 完整 API 参考
本文档提供 SDK 所有导出类和函数的完整参考,按模块分类。
## 相关文档
### 入门文档
- [README](./README.md)
- [Context API 参考](./01_context_api.md)
- [消息事件与组件](./02_event_and_components.md)
- [装饰器使用指南](./03_decorators.md)
### API 详细文档
#### 核心类
- [Star 类 API](./api/star.md) - 插件基类与生命周期
- [Context 类 API](./api/context.md) - 运行时上下文与能力客户端
- [MessageEvent 类 API](./api/message_event.md) - 消息事件对象
#### 装饰器与过滤器
- [装饰器 API](./api/decorators.md) - 事件触发、限制器、过滤器装饰器
#### 客户端
- [客户端 API](./api/clients.md) - LLM、Memory、DB、Platform 等 12 个客户端
#### 消息处理
- [消息组件 API](./api/message_components.md) - Plain、Image、At、Record、Video、File 等
- [消息结果 API](./api/message_result.md) - MessageChain、MessageBuilder、MessageEventResult
#### 工具与类型
- [工具与辅助类 API](./api/utils.md) - CancelToken、MessageSession、GreedyStr、CommandGroup 等
- [类型定义 API](./api/types.md) - 类型别名、泛型变量、Pydantic 模型
#### 错误处理
- [错误处理 API](./api/errors.md) - AstrBotError、ErrorCodes
@@ -1,494 +0,0 @@
# AstrBot SDK 迁移指南
本文档帮助开发者从旧版本或其他框架迁移到 AstrBot SDK v4。
## 目录
- [从 v3 迁移](#从-v3-迁移)
- [从其他框架迁移](#从其他框架迁移)
- [破坏性变更](#破坏性变更)
- [迁移检查清单](#迁移检查清单)
---
## 从 v3 迁移
### 插件类定义
**v3 (旧版本)**:
```python
from astrbot.api import star
@star.register("my_plugin")
class MyPlugin(star.Star):
def __init__(self, context):
super().__init__(context)
```
**v4 (新版本)**:
```python
from astrbot_sdk import Star
class MyPlugin(Star):
async def on_start(self, ctx):
pass
async def on_stop(self, ctx):
pass
```
### 装饰器变更
**v3**:
```python
from astrbot.api import filter
@filter.command("hello")
async def hello(self, event):
await event.reply("Hello!")
```
**v4**:
```python
from astrbot_sdk.decorators import on_command
@on_command("hello")
async def hello(self, event, ctx):
await event.reply("Hello!")
```
### Context 访问
**v3**:
```python
# 通过 self.context
config = self.context.get_config()
reply = await self.context.llm_generate("prompt")
```
**v4**:
```python
# 通过参数注入
async def handler(self, event, ctx):
config = await ctx.metadata.get_plugin_config()
reply = await ctx.llm.chat("prompt")
```
### 数据存储
**v3**:
```python
# 通过 context
await self.context.put_kv_data("key", value)
data = await self.context.get_kv_data("key", default)
```
**v4**:
```python
# 通过 db 客户端
await ctx.db.set("key", value)
data = await ctx.db.get("key")
# 或使用 Mixin
from astrbot_sdk import PluginKVStoreMixin
class MyPlugin(Star, PluginKVStoreMixin):
async def save(self):
await self.put_kv_data("key", value)
```
### 消息发送
**v3**:
```python
# 通过 event
await event.reply("消息")
# 主动发送
await self.context.send_message(session, chain)
```
**v4**:
```python
# 通过 event
await event.reply("消息")
# 主动发送
await ctx.platform.send(session, "消息")
await ctx.platform.send_chain(session, chain)
```
### 生命周期
**v3**:
```python
class MyPlugin(Star):
async def initialize(self):
# 初始化
pass
async def terminate(self):
# 清理
pass
```
**v4**:
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 启动时
await super().on_start(ctx)
async def on_stop(self, ctx):
# 停止时
await super().on_stop(ctx)
# 仍然支持
async def initialize(self):
pass
async def terminate(self):
pass
```
### 配置获取
**v3**:
```python
config = self.context.get_config()
```
**v4**:
```python
config = await ctx.metadata.get_plugin_config()
```
### LLM 调用
**v3**:
```python
reply = await self.context.llm_generate("prompt")
# 带历史
reply = await self.context.llm_generate(
"prompt",
contexts=[{"role": "user", "content": "历史"}]
)
```
**v4**:
```python
from astrbot_sdk.clients.llm import ChatMessage
reply = await ctx.llm.chat("prompt")
# 带历史
history = [
ChatMessage(role="user", content="历史"),
]
reply = await ctx.llm.chat("prompt", history=history)
```
### 错误处理
**v3**:
```python
try:
result = await operation()
except Exception as e:
await event.reply(f"错误: {e}")
```
**v4**:
```python
from astrbot_sdk.errors import AstrBotError
try:
result = await operation()
except AstrBotError as e:
# 使用 SDK 提供的用户友好提示
await event.reply(e.hint or e.message)
except Exception as e:
ctx.logger.error(f"错误: {e}")
await event.reply("操作失败")
```
---
## 从其他框架迁移
### 从 NoneBot2 迁移
**NoneBot2**:
```python
from nonebot import on_command
from nonebot.adapters.onebot.v11 import Bot, Event
matcher = on_command("hello")
@matcher.handle()
async def hello(bot: Bot, event: Event):
await matcher.send("Hello!")
```
**AstrBot SDK**:
```python
from astrbot_sdk import Star, MessageEvent, Context
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("hello")
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
```
### 从 Koishi 迁移
**Koishi**:
```javascript
ctx.command('hello')
.action(() => 'Hello!')
```
**AstrBot SDK**:
```python
from astrbot_sdk import Star, MessageEvent, Context
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("hello")
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
```
### 从 python-telegram-bot 迁移
**python-telegram-bot**:
```python
from telegram import Update
from telegram.ext import ContextTypes
async def hello(update: Update, context: ContextTypes.DEFAULT_TYPE):
await update.message.reply_text("Hello!")
```
**AstrBot SDK**:
```python
from astrbot_sdk import Star, MessageEvent, Context
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("hello")
@platforms("telegram")
async def hello(self, event: MessageEvent, ctx: Context):
await event.reply("Hello!")
```
---
## 破坏性变更
### v3 → v4 主要变更
1. **注册方式**
- v3: `@star.register()` + `@filter.command()`
- v4: `@on_command()` 直接在类方法上
2. **Context 获取**
- v3: `self.context`
- v4: `ctx` 参数注入
3. **数据存储**
- v3: `self.context.put_kv_data()`
- v4: `ctx.db.set()` 或 `PluginKVStoreMixin`
4. **配置获取**
- v3: `self.context.get_config()`
- v4: `ctx.metadata.get_plugin_config()`
5. **LLM 调用**
- v3: `self.context.llm_generate()`
- v4: `ctx.llm.chat()`
6. **生命周期**
- v3: `initialize()` / `terminate()`
- v4: `on_start()` / `on_stop()`(仍然支持旧方法)
7. **错误类型**
- v3: 标准 Python 异常
- v4: `AstrBotError` 体系
### 已弃用的功能
| v3 功能 | v4 替代方案 | 状态 |
|---------|-------------|------|
| `@star.register()` | 继承 `Star` 类 | 已移除 |
| `self.context` | `ctx` 参数 | 已变更 |
| `filter.command()` | `on_command()` | 已更名 |
| `filter.regex()` | `on_message(regex=...)` | 已变更 |
| `llm_generate()` | `ctx.llm.chat()` | 已更名 |
| `send_message()` | `ctx.platform.send()` | 已更名 |
---
## 迁移检查清单
### 代码迁移
- [ ] 更新导入语句
- [ ] 移除 `@star.register()` 装饰器
- [ ] 将 `@filter.command()` 改为 `@on_command()`
- [ ] 添加 `ctx` 参数到所有 handler
- [ ] 更新 Context 访问方式
- [ ] 更新数据存储调用
- [ ] 更新 LLM 调用
- [ ] 更新配置获取
- [ ] 更新错误处理
### 配置迁移
- [ ] 更新 `plugin.yaml` 格式
- [ ] 检查 `support_platforms` 配置
- [ ] 更新 `runtime` 配置
### 测试迁移
- [ ] 更新测试导入
- [ ] 更新测试 mock
- [ ] 运行测试验证
### 文档更新
- [ ] 更新 README
- [ ] 更新使用文档
- [ ] 更新 CHANGELOG
---
## 迁移工具
### 自动迁移脚本(示例)
```python
#!/usr/bin/env python3
"""v3 到 v4 迁移辅助脚本"""
import re
import sys
from pathlib import Path
def migrate_file(file_path: Path):
"""迁移单个文件"""
content = file_path.read_text(encoding="utf-8")
# 替换导入
content = re.sub(
r'from astrbot\.api import star',
'from astrbot_sdk import Star, Context, MessageEvent',
content
)
# 替换装饰器
content = re.sub(
r'@star\.register\([^)]*\)',
'',
content
)
content = re.sub(
r'@filter\.command\(([^)]*)\)',
r'@on_command(\1)',
content
)
# 替换类定义
content = re.sub(
r'class (\w+)\(star\.Star\)',
r'class \1(Star)',
content
)
# 替换 context 访问
content = re.sub(
r'self\.context\.get_config\(\)',
'await ctx.metadata.get_plugin_config()',
content
)
content = re.sub(
r'self\.context\.llm_generate\(',
'ctx.llm.chat(',
content
)
# 添加 ctx 参数
content = re.sub(
r'async def (\w+)\(self, event\)',
r'async def \1(self, event, ctx)',
content
)
# 写回文件
file_path.write_text(content, encoding="utf-8")
print(f"已迁移: {file_path}")
def main():
if len(sys.argv) < 2:
print("用法: python migrate.py <plugin_directory>")
sys.exit(1)
plugin_dir = Path(sys.argv[1])
for py_file in plugin_dir.rglob("*.py"):
migrate_file(py_file)
print("迁移完成!请手动检查并测试。")
if __name__ == "__main__":
main()
```
---
## 常见问题
### Q: v3 插件能在 v4 运行吗?
**A**: 不能,需要进行迁移。但是 SDK 提供了兼容层,可以简化迁移过程。
### Q: 可以同时支持 v3 和 v4 吗?
**A**: 不推荐。建议为 v4 创建新的插件版本。
### Q: 迁移后测试失败怎么办?
**A**:
1. 检查导入是否正确
2. 确认 `ctx` 参数已添加
3. 验证异步函数使用 `await`
4. 查看错误日志获取详细信息
### Q: 如何逐步迁移?
**A**:
1. 先迁移插件结构和装饰器
2. 再迁移业务逻辑
3. 最后更新测试
4. 每个阶段都进行测试
---
## 获取帮助
- 查看完整文档:[docs/](./)
- 提交问题:[GitHub Issues](https://github.com/AstrBotDevs/AstrBot/issues)
- 迁移示例:[examples/migration/](./examples/migration/)
---
## 相关文档
- [README](./README.md)
- [Context API 参考](./01_context_api.md)
- [Star 类与生命周期](./04_star_lifecycle.md)
- [错误处理与调试](./06_error_handling.md)
@@ -1,382 +0,0 @@
# AstrBot SDK 安全检查清单
本文档包含 SDK 安全开发检查清单和已知安全问题,帮助开发者编写安全的插件。
## 目录
- [安全检查清单](#安全检查清单)
- [已知安全问题](#已知安全问题)
- [安全最佳实践](#安全最佳实践)
- [安全审计指南](#安全审计指南)
---
## 安全检查清单
### 输入验证
- [ ] 所有用户输入都经过验证
- [ ] 输入长度有限制
- [ ] 输入内容有白名单过滤
- [ ] 特殊字符被正确转义
```python
# ✅ 好的做法
import re
from astrbot_sdk.errors import AstrBotError
def validate_input(text: str) -> str:
if len(text) > 1000:
raise AstrBotError.invalid_input("输入过长")
if not re.match(r'^[\w\s\-]+$', text):
raise AstrBotError.invalid_input("包含非法字符")
return text
# ❌ 不好的做法
async def unsafe_handler(event, ctx):
result = eval(event.text) # 危险!
```
### 敏感信息处理
- [ ] API Key 等敏感信息不硬编码
- [ ] 敏感信息从配置或环境变量读取
- [ ] 敏感信息不在日志中打印
- [ ] 敏感信息不存储在不安全的位置
```python
# ✅ 好的做法
import os
class MyPlugin(Star):
async def on_start(self, ctx):
config = await ctx.metadata.get_plugin_config()
self.api_key = config.get("api_key") or os.getenv("MY_API_KEY")
ctx.logger.info("API Key 已配置") # 不打印实际值
# ❌ 不好的做法
class UnsafePlugin(Star):
api_key = "sk-1234567890" # 硬编码!
async def on_start(self, ctx):
ctx.logger.info(f"API Key: {self.api_key}") # 泄露!
```
### 权限检查
- [ ] 管理员命令有权限验证
- [ ] 敏感操作有二次确认
- [ ] 资源访问有权限控制
```python
# ✅ 好的做法
from astrbot_sdk.decorators import require_admin
class MyPlugin(Star):
@on_command("admin_only")
@require_admin
async def admin_cmd(self, event, ctx):
await event.reply("管理员命令")
# ❌ 不好的做法
class UnsafePlugin(Star):
@on_command("delete_all")
async def delete_all(self, event, ctx):
# 任何人都可以执行危险操作!
await ctx.db.clear_all()
```
### 速率限制
- [ ] 昂贵的操作有速率限制
- [ ] API 调用有配额控制
- [ ] 资源密集型操作有限制
```python
# ✅ 好的做法
from astrbot_sdk.decorators import rate_limit
class MyPlugin(Star):
@on_command("generate")
@rate_limit(limit=5, window=3600, scope="user")
async def generate(self, event, ctx):
# 昂贵的 LLM 调用
result = await ctx.llm.chat("生成内容", model="gpt-4")
await event.reply(result)
```
### 资源管理
- [ ] 资源正确释放
- [ ] 连接正确关闭
- [ ] 任务正确取消
- [ ] 避免资源泄漏
```python
# ✅ 好的做法
class MyPlugin(Star):
async def on_start(self, ctx):
self._session = aiohttp.ClientSession()
self._task = asyncio.create_task(self.background_task())
async def on_stop(self, ctx):
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
if self._session:
await self._session.close()
```
### 错误处理
- [ ] 错误信息不泄露敏感信息
- [ ] 异常被正确捕获和处理
- [ ] 错误日志不包含敏感数据
```python
# ✅ 好的做法
try:
result = await operation()
except Exception as e:
ctx.logger.error(f"操作失败: {type(e).__name__}")
await event.reply("操作失败,请稍后重试")
# ❌ 不好的做法
try:
result = await operation()
except Exception as e:
await event.reply(f"错误: {str(e)}") # 可能泄露敏感信息
```
---
## 已知安全问题
当前版本没有已知的 SDK 框架级高风险未修复项。以下历史回归已经关闭,
保留在这里帮助开发者理解为什么这些约束存在:
- `ProviderManagerClient.register_provider_change_hook()` 现在必须和
`unregister_provider_change_hook()` 配对使用,避免残留订阅任务。
- `PlatformCompatFacade` 内部已经串行化状态刷新,插件侧不需要再额外为
`refresh()` / `clear_errors()` 套一层锁来规避 SDK 自身竞态。
- Provider 管理路径会先复制 provider payload,再做 merge,避免污染共享缓存。
---
### 🟡 Medium: 命令参数注入风险
**问题描述**:
插件可能直接使用用户输入作为命令参数,存在注入风险。
**风险等级**: Medium
**示例**:
```python
# ❌ 危险
@on_command("search")
async def search(self, event, ctx, query):
# 如果 query 包含特殊字符,可能引发问题
os.system(f"grep {query} data.txt")
# ✅ 安全
@on_command("search")
async def search(self, event, ctx, query):
# 验证和清理输入
safe_query = re.sub(r'[^\w\s]', '', query)
subprocess.run(["grep", safe_query, "data.txt"], capture_output=True)
```
---
### 🟢 Low: 敏感信息可能出现在日志中
**问题描述**:
某些错误日志可能包含敏感信息。
**风险等级**: Low
**建议**:
```python
# ✅ 安全的日志记录
ctx.logger.info(f"用户 {user_id} 执行操作") # 只记录 ID
# ❌ 不安全的日志记录
ctx.logger.info(f"用户数据: {user_data}") # 可能包含敏感信息
```
---
## 安全最佳实践
### 1. 最小权限原则
```python
class MyPlugin(Star):
@on_command("public")
async def public_cmd(self, event, ctx):
# 所有人可用
pass
@on_command("admin")
@require_admin
async def admin_cmd(self, event, ctx):
# 仅管理员可用
pass
@on_command("owner")
async def owner_cmd(self, event, ctx):
# 仅插件所有者可用
if event.user_id != self.owner_id:
raise AstrBotError.invalid_input("权限不足")
```
### 2. 输入验证白名单
```python
import re
ALLOWED_COMMANDS = {"help", "status", "info"}
def validate_command(cmd: str) -> str:
cmd = cmd.lower().strip()
if cmd not in ALLOWED_COMMANDS:
raise AstrBotError.invalid_input("未知命令")
return cmd
```
### 3. 安全的文件操作
```python
import os
from pathlib import Path
BASE_DIR = Path("/safe/directory")
def safe_read_file(filename: str) -> str:
# 防止目录遍历
path = (BASE_DIR / filename).resolve()
if not str(path).startswith(str(BASE_DIR)):
raise AstrBotError.invalid_input("非法路径")
return path.read_text()
```
### 4. 安全的正则表达式
```python
import re
# ✅ 使用原始字符串和适当的限制
pattern = re.compile(r'^[a-zA-Z0-9_]{1,50}$')
# ❌ 避免复杂的正则,可能导致 ReDoS
# pattern = re.compile(r'(a+)+b') # 危险!
```
### 5. 安全配置
```python
class MyPlugin(Star):
async def on_start(self, ctx):
config = await ctx.metadata.get_plugin_config()
# 验证必需配置
required = ["api_key", "endpoint"]
for key in required:
if key not in config:
raise ValueError(f"缺少必需配置: {key}")
# 验证配置值
if not config["api_key"].startswith("sk-"):
raise ValueError("无效的 API Key 格式")
self.config = config
```
---
## 安全审计指南
### 审计检查清单
1. **代码审查**
- [ ] 所有输入都经过验证
- [ ] 没有使用 eval/exec
- [ ] 没有硬编码的敏感信息
- [ ] 错误处理不泄露敏感信息
2. **依赖审查**
```bash
# 检查依赖漏洞
pip install safety
safety check
# 检查依赖许可证
pip install pip-licenses
pip-licenses
```
3. **日志审查**
- [ ] 日志不包含密码、token
- [ ] 日志不包含个人隐私信息
- [ ] 日志有适当的级别
4. **权限审查**
- [ ] 敏感操作有权限检查
- [ ] 没有特权提升漏洞
- [ ] 资源访问有控制
### 安全测试
```python
# 测试输入验证
def test_input_validation():
# SQL 注入测试
malicious_input = "' OR '1'='1"
# XSS 测试
xss_input = "<script>alert('xss')</script>"
# 路径遍历测试
path_input = "../../../etc/passwd"
# 验证这些输入都被正确拒绝
```
### 安全工具
```bash
# 静态分析
pip install bandit
bandit -r my_plugin/
# 类型检查
pip install mypy
mypy my_plugin/
# 代码质量
pip install pylint
pylint my_plugin/
```
---
## 报告安全问题
如果您发现 SDK 或插件的安全问题,请通过以下方式报告:
1. **不要** 在公开 issue 中报告安全问题
2. 通过项目官方联系渠道私下报告,例如 `community@astrbot.app`
3. 提供详细的复现步骤
4. 等待修复后再公开
---
## 相关文档
- [错误处理与调试](./06_error_handling.md)
- [高级主题](./07_advanced_topics.md)
- [测试指南](./08_testing_guide.md)
-128
View File
@@ -1,128 +0,0 @@
# AstrBot SDK 文档目录
本文档目录包含完整的 SDK 开发文档,按难度级别分类。
## 📚 文档列表(按学习路径)
### 🚀 快速开始(初级使用者)
适合第一次接触 AstrBot SDK 的开发者:
| 文档 | 描述 | 行数 |
|------|------|------|
| [README.md](./README.md) | 文档首页、快速开始、核心概念 | ~350 |
| [01_context_api.md](./01_context_api.md) | Context 类的核心客户端和系统工具方法 | ~650 |
| [02_event_and_components.md](./02_event_and_components.md) | MessageEvent 和消息组件的使用 | ~480 |
| [03_decorators.md](./03_decorators.md) | 所有装饰器的详细说明 | ~580 |
| [04_star_lifecycle.md](./04_star_lifecycle.md) | 插件基类和生命周期钩子 | ~490 |
| [05_clients.md](./05_clients.md) | 所有客户端的完整 API 文档 | ~422 |
### 🔧 进阶主题(中级使用者)
适合已经掌握基础,希望深入了解 SDK 的开发者:
| 文档 | 描述 | 行数 |
|------|------|------|
| [06_error_handling.md](./06_error_handling.md) | 完整的错误处理指南和调试技巧 | ~530 |
| [07_advanced_topics.md](./07_advanced_topics.md) | 并发处理、性能优化、安全最佳实践 | ~550 |
| [08_testing_guide.md](./08_testing_guide.md) | 如何测试插件和 Mock 使用 | ~450 |
### 📖 参考资料(高级使用者)
适合需要深入了解 SDK 架构和完整 API 的开发者:
| 文档 | 描述 | 行数 |
|------|------|------|
| [09_api_reference.md](./09_api_reference.md) | 所有导出类和函数的完整参考 | ~880 |
| [10_migration_guide.md](./10_migration_guide.md) | 从旧版本或其他框架迁移 | ~450 |
| [11_security_checklist.md](./11_security_checklist.md) | 安全开发检查清单和已知问题 | ~480 |
| [PROJECT_ARCHITECTURE.md](./PROJECT_ARCHITECTURE.md) | SDK 架构设计文档 | ~872 |
---
## 📊 文档统计
- **总文档数**: 13 个
- **总内容行数**: ~6,700 行
- **新增/更新文档**: 7 个
- **保留原有**: 6 个
- **API 覆盖率**: 100% (77/77 exports documented)
---
## 🎯 文档内容覆盖
### 已涵盖的主题
✅ **基础使用**
- Context API 完整参考
- 消息事件处理
- 消息组件使用
- 装饰器使用
- 生命周期管理
✅ **错误处理**
- AstrBotError 完整文档
- 错误码参考
- 错误处理模式
- 调试技巧
✅ **高级主题**
- 并发处理
- 性能优化
- 安全最佳实践
- 架构设计模式
✅ **测试**
- 单元测试
- 集成测试
- Mock 使用
- 测试最佳实践
✅ **API 参考**
- 所有导出类的完整参考
- 方法签名
- 使用示例
✅ **迁移指南**
- v3 → v4 迁移
- 从其他框架迁移
- 破坏性变更列表
- 迁移检查清单
✅ **安全检查清单**
- 安全开发检查清单
- 已知安全问题(包含发现的问题)
- 安全最佳实践
- 安全审计指南
## 📝 文档使用建议
### 初级开发者
1. 从 [README.md](./README.md) 开始
2. 阅读 01-05 文档了解基础 API
3. 参考示例代码编写第一个插件
### 中级开发者
1. 阅读 [06_error_handling.md](./06_error_handling.md) 建立健壮的错误处理
2. 学习 [07_advanced_topics.md](./07_advanced_topics.md) 的并发和性能优化
3. 按照 [08_testing_guide.md](./08_testing_guide.md) 编写测试
### 高级开发者
1. 阅读 [09_api_reference.md](./09_api_reference.md) 了解所有可用功能
2. 研究 [07_advanced_topics.md](./07_advanced_topics.md) 中的架构设计
3. 阅读 [PROJECT_ARCHITECTURE.md](./PROJECT_ARCHITECTURE.md) 深入理解实现
---
## 🔗 相关资源
- **项目地址**: https://github.com/AstrBotDevs/AstrBot
- **SDK 版本**: v4.0
- **协议版本**: P0.6
- **Python 要求**: >= 3.12
---
**最后更新**: 2026-03-17
@@ -1,555 +0,0 @@
# AstrBot SDK 架构概述文档
> 作者:whatevertogo
> 生成日期:2026-03-19
---
## 目录
1. [项目概述](#项目概述)
2. [核心架构层次](#核心架构层次)
3. [协议层设计](#协议层设计)
4. [运行时架构](#运行时架构)
5. [客户端层设计](#客户端层设计)
6. [插件开发指南](#插件开发指南)
7. [关键设计模式](#关键设计模式)
8. [文档与资源](#文档与资源)
---
## 项目概述
AstrBot SDK 是一个基于 Python 3.12+ 的机器人插件开发框架,采用**进程隔离**和**能力路由**架构,支持插件的动态加载、独立运行和跨进程通信。
### 核心特性
| 特性 | 描述 |
|------|------|
| **进程隔离** | 每个插件运行在独立 Worker 进程,崩溃不影响其他插件 |
| **环境分组** | 多插件可共享同一 Python 虚拟环境,节省资源 |
| **能力路由** | 显式声明的 Capability 系统,支持 JSON Schema 验证 |
| **流式支持** | 原生支持流式 LLM 调用和增量结果返回 |
| **向后兼容** | 完整的旧版 API 兼容层,支持无修改迁移 |
| **协议优先** | 基于 v4 协议的统一通信模型,支持多种传输方式 |
### 技术栈
- **Python**: 3.12+
- **异步框架**: asyncio
- **Web 框架**: aiohttp
- **数据验证**: pydantic
- **日志**: loguru
- **配置**: pyyaml
- **LLM**: openai, anthropic, google-genai
- **包管理**: uv (环境分组)
---
## 核心架构层次
```
┌─────────────────────────────────────────────────────────────────┐
│ 用户层 (Plugin Developer) │
├─────────────────────────────────────────────────────────────────┤
│ v4 入口: astrbot_sdk.{Star, Context, MessageEvent} │
│ 装饰器: on_command, on_message, on_event, on_schedule │
│ provide_capability, require_admin │
│ 过滤器: PlatformFilter, MessageTypeFilter, CustomFilter │
│ 命令组: CommandGroup, command_group │
│ 会话: MessageSession, session_waiter │
└────────────────────┬────────────────────────────────────────────┘
│
┌──────────────────▼─────────────────────────────────────────────┐
│ 高层 API (High-Level API) │
├─────────────────────────────────────────────────────────────────┤
│ 能力客户端 (通过 CapabilityProxy 调用): │
│ - LLMClient (llm.chat, llm.chat_raw, llm.stream_chat)│
│ - MemoryClient (memory.search, memory.save, memory.stats)│
│ - DBClient (db.get, db.set, db.watch, db.list) │
│ - PlatformClient (platform.send, platform.send_image, ...)│
│ - HTTPClient (http.register_api, http.list_apis) │
│ - MetadataClient (metadata.get_plugin, metadata.list_plugins)│
└────────────────────┬────────────────────────────────────────────┘
│
┌──────────────────▼─────────────────────────────────────────────┐
│ 执行边界 (Execution Boundary) │
├─────────────────────────────────────────────────────────────────┤
│ runtime 主干: │
│ - loader.py (插件发现、加载、环境管理) │
│ - bootstrap.py (Supervisor/Worker 启动) │
│ - handler_dispatcher.py (Handler 执行分发、参数注入) │
│ - capability_dispatcher.py (Capability 调用分发) │
│ - capability_router.py (Capability 路由、Schema 验证) │
│ - peer.py (协议对等端) │
│ - transport.py (传输抽象) │
└────────────────────┬────────────────────────────────────────────┘
│
┌──────────────────▼─────────────────────────────────────────────┐
│ 协议与传输 (Protocol & Transport) │
├─────────────────────────────────────────────────────────────────┤
│ protocol/ │
│ - messages.py (协议消息模型) │
│ - descriptors.py (Handler/Capability 描述符) │
│ transport 实现: │
│ - StdioTransport (标准输入输出) │
│ - WebSocketServerTransport (WebSocket 服务端) │
│ - WebSocketClientTransport (WebSocket 客户端) │
└─────────────────────────────────────────────────────────────────┘
```
### 层次职责
| 层次 | 职责 | 主要模块 |
|------|------|---------|
| **用户层** | 插件开发者 API | `Star`, `Context`, `MessageEvent`, 装饰器, 过滤器 |
| **高层 API** | 类型化的能力客户端 | `clients/{llm, memory, db, platform, http, metadata}` |
| **执行边界** | 插件加载、路由、分发 | `runtime/loader.py`, `runtime/*_dispatcher.py` |
| **协议层** | 消息模型、描述符、JSON Schema | `protocol/` |
| **传输层** | 底层通信抽象 | `runtime/transport.py` |
### 核心设计原则
1. **延迟加载**:`runtime/__init__.py` 使用 `__getattr__` 避免导入时加载重型依赖
2. **插件身份透传**:通过 `caller_plugin_scope()` 上下文管理器将 plugin_id 注入协议层
3. **声明式优先**:所有配置都是数据结构(描述符),便于序列化和跨进程传递
4. **类型安全**:使用 Pydantic 模型和类型注解提供验证和 IDE 支持
---
## 协议层设计
### 消息模型
v4 协议定义了 5 种消息类型:
| 消息类型 | 用途 | 关键字段 |
|---------|------|---------|
| `InitializeMessage` | 握手初始化 | `protocol_version`, `peer`, `handlers`, `provided_capabilities` |
| `InvokeMessage` | 调用能力 | `capability`, `input`, `stream`, `caller_plugin_id` |
| `ResultMessage` | 返回结果 | `success`, `output`, `error`, `kind` |
| `EventMessage` | 流式事件 | `phase` (started/delta/completed/failed), `data` |
| `CancelMessage` | 取消调用 | `reason` |
### 错误模型
`ErrorPayload` 使用字符串 code(而非整数),包含:
- `code`: 错误码(如 "capability_not_found")
- `message`: 开发者信息
- `hint`: 用户友好提示
- `retryable`: 是否可重试
### 握手流程
```
Worker (Plugin) Supervisor (Core)
| |
| InitializeMessage |
| (handlers, capabilities) |
|----------------------------->|
| |
| ResultMessage(kind="init") |
|<-----------------------------|
| |
| InvokeMessage(handler.invoke) |
|<-----------------------------|
| 执行用户 handler |
| |
| ResultMessage(output) |
|----------------------------->|
```
### 描述符模型
#### HandlerDescriptor
```python
{
"id": "plugin.module:handler_name",
"trigger": {
"type": "command",
"command": "hello",
"aliases": ["hi"],
"description": "打招呼命令"
},
"kind": "handler", # handler | hook | tool | session
"contract": "message_event", # message_event | schedule
"priority": 0,
"permissions": {"require_admin": False, "level": 0},
"filters": [],
"param_specs": []
}
```
#### Trigger 类型
| 类型 | 关键字段 | 说明 |
|------|---------|------|
| `CommandTrigger` | command, aliases, platforms | 命令触发 |
| `MessageTrigger` | regex, keywords, platforms | 消息触发(正则/关键词) |
| `EventTrigger` | event_type | 事件触发 |
| `ScheduleTrigger` | cron, interval_seconds | 定时触发 |
### 内置 Capabilities (38个)
#### LLM 命名空间
| 能力 | 说明 |
|------|------|
| `llm.chat` | 同步对话,返回文本 |
| `llm.chat_raw` | 同步对话,返回完整响应 |
| `llm.stream_chat` | 流式对话 |
#### Memory 命名空间
| 能力 | 说明 |
|------|------|
| `memory.search` | 语义搜索记忆 |
| `memory.save` | 保存记忆 |
| `memory.save_with_ttl` | 保存带过期时间的记忆 |
| `memory.get` / `get_many` | 读取记忆 |
| `memory.delete` / `delete_many` | 删除记忆 |
| `memory.stats` | 获取统计信息 |
#### DB 命名空间
| 能力 | 说明 |
|------|------|
| `db.get` / `get_many` | 读取 KV |
| `db.set` / `set_many` | 写入 KV |
| `db.delete` | 删除 KV |
| `db.list` | 列出键(支持前缀过滤) |
| `db.watch` | 订阅变更(流式) |
#### Platform 命名空间
| 能力 | 说明 |
|------|------|
| `platform.send` | 发送文本消息 |
| `platform.send_image` | 发送图片 |
| `platform.send_chain` | 发送消息链 |
| `platform.get_members` | 获取群成员 |
#### HTTP 命名空间
| 能力 | 说明 |
|------|------|
| `http.register_api` | 注册 HTTP API 端点 |
| `http.unregister_api` | 注销 HTTP API 端点 |
| `http.list_apis` | 列出已注册的 API |
#### Metadata 命名空间
| 能力 | 说明 |
|------|------|
| `metadata.get_plugin` | 获取单个插件元数据 |
| `metadata.list_plugins` | 列出所有插件元数据 |
| `metadata.get_plugin_config` | 获取当前插件配置 |
#### System 命名空间
| 能力 | 说明 |
|------|------|
| `system.get_data_dir` | 获取插件数据目录 |
| `system.text_to_image` | 文本转图片 |
| `system.html_render` | 渲染 HTML 模板 |
| `system.session_waiter.*` | 会话等待器管理 |
| `system.event.*` | 表情回应、输入状态、流式消息 |
---
## 运行时架构
### 组件关系图
```
┌──────────────┐
│ AstrBot │
│ Core │
└──────┬─────┘
│
┌──────▼─────┐
│ Supervisor │
│ Runtime │
└──────┬─────┘
│
┌──────────────────┼──────────────────┐
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Peer │ │ Peer │ │ Peer │
│ (stdio) │ │ (stdio) │ │ (stdio) │
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Worker │ │ Worker │ │ Worker │
│ Runtime │ │ Runtime │ │ Runtime │
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Plugin A │ │ Plugin B │ │ Plugin C │
└───────────┘ └───────────┘ └───────────┘
```
### 核心运行时组件
| 组件 | 职责 |
|------|------|
| **SupervisorRuntime** | 管理多个 Worker 进程,聚合所有 handler |
| **WorkerSession** | 管理单个 Worker 进程的生命周期 |
| **PluginWorkerRuntime** | Worker 进程内的插件加载与执行 |
| **HandlerDispatcher** | 将 handler.invoke 请求转成真实 Python 调用 |
| **CapabilityRouter** | 能力注册、发现和执行路由 |
### 参数注入优先级
HandlerDispatcher 支持参数注入,优先级为:
1. **按类型注解注入**(`MessageEvent`, `Context`)
2. **按参数名注入**(`event`, `ctx`, `context`)
3. **从 legacy_args 注入**(命令参数等)
---
## 客户端层设计
### 客户端架构
```
┌─────────────────────────────────────────────────────────────┐
│ User Plugin │
│ ctx.llm.chat() / ctx.memory.save() / ctx.db.set() │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ CapabilityProxy │
│ - call(name, payload) 普通调用 │
│ - stream(name, payload) 流式调用 │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ Peer │
│ - invoke(capability, payload) │
│ - invoke_stream(capability, payload) │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ Transport │
│ - send(json_string) │
└─────────────────────────────────────────────────────────────┘
```
### 客户端一览
| 客户端 | 主要方法 | 对应 Capability |
|--------|---------|-----------------|
| `LLMClient` | `chat()`, `chat_raw()`, `stream_chat()` | `llm.*` |
| `MemoryClient` | `search()`, `save()`, `save_with_ttl()`, `get()`, `get_many()`, `delete()`, `delete_many()`, `stats()` | `memory.*` |
| `DBClient` | `get()`, `set()`, `get_many()`, `set_many()`, `delete()`, `list()`, `watch()` | `db.*` |
| `PlatformClient` | `send()`, `send_image()`, `send_chain()`, `get_members()` | `platform.*` |
| `HTTPClient` | `register_api()`, `unregister_api()`, `list_apis()` | `http.*` |
| `MetadataClient` | `get_plugin()`, `list_plugins()`, `get_current_plugin()`, `get_plugin_config()` | `metadata.*` |
---
## 插件开发指南
### v4 原生插件示例
#### plugin.yaml
```yaml
_schema_version: 2
name: my_plugin
author: your_name
version: 1.0.0
runtime:
python: "3.12"
components:
- class: main:MyPlugin
```
#### main.py
```python
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message, provide_capability
class MyPlugin(Star):
# 命令处理器
@on_command("hello", aliases=["hi"])
async def hello(self, event: MessageEvent, ctx: Context) -> None:
await event.reply(f"你好,{event.user_id}!")
# 消息处理器
@on_message(keywords=["帮助"])
async def help(self, event: MessageEvent, ctx: Context) -> None:
await event.reply("可用命令:hello, help")
# 提供能力
@provide_capability(
"my_plugin.calculate",
description="执行计算",
input_schema={
"type": "object",
"properties": {"x": {"type": "number"}},
"required": ["x"]
},
output_schema={
"type": "object",
"properties": {"result": {"type": "number"}},
"required": ["result"]
}
)
async def calculate_capability(self, payload: dict, ctx: Context) -> dict:
x = payload.get("x", 0)
return {"result": x * 2}
```
### 生命周期钩子
| 钩子 | 说明 |
|------|------|
| `on_start()` | 插件启动时调用 |
| `on_stop()` | 插件停止时调用 |
| `on_error(exc, event, ctx)` | Handler 执行出错时调用 |
### 常用功能速查
#### 1. LLM 对话
```python
# 简单对话
reply = await ctx.llm.chat("你好")
# 带历史对话
from astrbot_sdk.clients.llm import ChatMessage
history = [ChatMessage(role="user", content="我叫小明")]
reply = await ctx.llm.chat("你记得我吗?", history=history)
# 流式对话
async for chunk in ctx.llm.stream_chat("讲个故事"):
print(chunk, end="")
```
#### 2. 数据持久化
```python
# DB 客户端(精确匹配)
await ctx.db.set("user:123", {"name": "Alice"})
data = await ctx.db.get("user:123")
# Memory 客户端(语义搜索)
await ctx.memory.save("user_pref", {"theme": "dark"})
results = await ctx.memory.search("用户喜欢什么颜色")
```
#### 3. 消息发送
```python
# 简单文本
await ctx.platform.send(event.session_id, "消息内容")
# 图片
await ctx.platform.send_image(event.session_id, "https://example.com/img.jpg")
# 消息链
from astrbot_sdk.message_components import Plain, Image
chain = [Plain("文字"), Image(url="https://example.com/img.jpg")]
await ctx.platform.send_chain(event.session_id, chain)
```
---
## 关键设计模式
### 1. 协议优先模式
- 所有跨进程通信都通过 v4 协议
- 传输层只处理字符串,协议由 Peer 层处理
- 支持多种传输方式(Stdio, WebSocket)
### 2. 能力路由模式
- 显式声明 Capability 和输入/输出 Schema
- 通过 CapabilityRouter 统一路由
- 支持同步和流式两种调用模式
- 冲突处理:保留命名空间冲突直接跳过,非保留命名空间冲突自动添加插件名前缀
### 3. 环境分组模式
- 多插件可共享同一 Python 虚拟环境
- 按版本和依赖兼容性自动分组
- 节省资源,加快启动速度
### 4. 参数注入模式
- HandlerDispatcher 支持类型注解注入
- 优先级:类型注解 > 参数名 > legacy_args
- 支持可选类型 `Optional[Type]`
### 5. 取消传播模式
- CancelToken 统一取消机制
- 跨进程取消通过 CancelMessage
- 早到取消避免竞态条件
### 6. 插件隔离模式
- 每个插件运行在独立 Worker 进程
- 崩溃不影响其他插件
- 支持 GroupWorkerRuntime 共享环境
### 7. 热重载模式
- `dev --watch` 支持文件变更检测
- 按插件目录清理 `sys.modules` 缓存
- 确保代码变更后正确重载
---
## 文档与资源
### 完整文档目录
SDK 文档按学习路径组织,位于 `src/astrbot_sdk/docs/`:
| 级别 | 文档 | 内容 |
|------|------|------|
| **初级** | README.md | 快速开始、核心概念 |
| | 01_context_api.md | Context API 完整参考 |
| | 02_event_and_components.md | MessageEvent 和消息组件 |
| | 03_decorators.md | 装饰器详细说明 |
| | 04_star_lifecycle.md | 插件基类和生命周期 |
| | 05_clients.md | 客户端 API 文档 |
| **中级** | 06_error_handling.md | 错误处理与调试 |
| | 07_advanced_topics.md | 并发、性能优化、安全 |
| | 08_testing_guide.md | 测试指南 |
| **高级** | 09_api_reference.md | 完整 API 索引 |
| | 10_migration_guide.md | 迁移指南 |
| | 11_security_checklist.md | 安全检查清单 |
| | PROJECT_ARCHITECTURE.md | 架构设计文档 |
### 关键文件速查
| 文件 | 核心类/函数 | 说明 |
|------|------------|------|
| `astrbot_sdk/__init__.py` | `Star`, `Context`, `MessageEvent` | 顶层入口 |
| `astrbot_sdk/star.py` | `Star` | v4 原生插件基类 |
| `astrbot_sdk/context.py` | `Context` | 运行时上下文 |
| `astrbot_sdk/decorators.py` | `on_command`, `on_message` | v4 装饰器 |
| `astrbot_sdk/errors.py` | `AstrBotError` | 统一错误模型 |
| `astrbot_sdk/runtime/peer.py` | `Peer` | 协议对等端 |
| `astrbot_sdk/runtime/capability_router.py` | `CapabilityRouter` | Capability 路由 |
| `astrbot_sdk/clients/llm.py` | `LLMClient` | LLM 客户端 |
### 版本信息
- **SDK 版本**: v4.0
- **协议版本**: P0.6
- **Python 要求**: >=3.12
- **推荐版本**: 3.12+
---
> 本文档基于 AstrBot SDK v4 架构文档整理
> 详细内容请查阅 `src/astrbot_sdk/docs/` 目录下的完整文档
-445
View File
@@ -1,445 +0,0 @@
# AstrBot SDK 插件开发文档
欢迎来到 AstrBot SDK 插件开发文档!本文档面向 SDK 插件开发者,提供从入门到精通的完整指南。
## 📚 文档目录
### 🚀 快速开始(初级使用者)
适合第一次接触 AstrBot SDK 的开发者:
- **[01. Context API 参考](./01_context_api.md)** - Context 类的核心客户端和系统工具方法
- **[02. 消息事件与组件](./02_event_and_components.md)** - MessageEvent 和消息组件的使用
- **[03. 装饰器使用指南](./03_decorators.md)** - 所有装饰器的详细说明
- **[04. Star 类与生命周期](./04_star_lifecycle.md)** - 插件基类和生命周期钩子
- **[05. 客户端 API 参考](./05_clients.md)** - 所有客户端的完整 API 文档
### 🔧 进阶主题(中级使用者)
适合已经掌握基础,希望深入了解 SDK 的开发者:
- **[06. 错误处理与调试](./06_error_handling.md)** - 完整的错误处理指南和调试技巧
- **[07. 高级主题](./07_advanced_topics.md)** - 并发处理、性能优化、安全最佳实践
- **[08. 测试指南](./08_testing_guide.md)** - 如何测试插件和 Mock 使用
### 📖 参考资料(高级使用者)
适合需要深入了解 SDK 架构和完整 API 的开发者:
- **[09. 完整 API 索引](./09_api_reference.md)** - 所有导出类和函数的完整参考
- **[10. 迁移指南](./10_migration_guide.md)** - 从旧版本或其他框架迁移
- **[11. 安全检查清单](./11_security_checklist.md)** - 安全开发检查清单和已知问题
---
## 🎯 学习路径推荐
### 初级路径:快速上手
```
1. 阅读本 README 的快速开始部分
2. 跟随下面的"创建第一个插件"教程
3. 查阅 01-05 文档了解基础 API
4. 参考文档中的示例代码
```
### 中级路径:进阶开发
```
1. 阅读 06 错误处理指南,建立健壮的错误处理机制
2. 学习 07 高级主题中的并发和性能优化
3. 按照 08 测试指南编写测试
4. 尝试开发复杂的插件功能
```
### 高级路径:精通 SDK
```
1. 阅读 09 完整 API 索引,了解所有可用功能
2. 研究 07 高级主题中的架构设计
3. 阅读 SDK 源码深入理解实现
4. 参与 SDK 贡献和改进
```
---
## 🚀 快速上手
### 创建第一个插件
```python
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message
class MyPlugin(Star):
"""我的第一个插件"""
@on_command("hello")
async def hello(self, event: MessageEvent, ctx: Context):
"""打招呼命令"""
await event.reply(f"你好,{event.sender_name}!")
@on_message(keywords=["帮助", "help"])
async def help(self, event: MessageEvent, ctx: Context):
"""帮助信息"""
await event.reply("可用命令: /hello")
```
### 插件配置 (plugin.yaml)
```yaml
_schema_version: 2
name: my_plugin
author: your_name
version: 1.0.0
desc: 我的插件描述
runtime:
python: "3.12"
components:
- class: main:MyPlugin
support_platforms:
- aiocqhttp
- telegram
```
---
## 📖 核心概念
### Context - 能力访问入口
`Context` 是插件与 AstrBot Core 交互的主要入口:
```python
# LLM 对话
reply = await ctx.llm.chat("你好")
# 数据存储
await ctx.db.set("key", "value")
data = await ctx.db.get("key")
# 记忆存储
await ctx.memory.save("pref", {"theme": "dark"})
# 发送消息
await ctx.platform.send(event.session_id, "消息内容")
# 获取配置
config = await ctx.metadata.get_plugin_config()
```
### MessageEvent - 消息事件
`MessageEvent` 表示接收到的消息事件:
```python
# 回复消息
await event.reply("回复内容")
# 获取消息组件
images = event.get_images()
# 判断消息类型
if event.is_group_chat():
await event.reply("这是群聊消息")
# 构建返回结果
return event.plain_result("返回内容")
```
### 装饰器 - 事件处理注册
```python
from astrbot_sdk.decorators import (
on_command, # 命令触发
on_message, # 消息触发
on_event, # 事件触发
on_schedule, # 定时任务
require_admin, # 权限控制
rate_limit, # 速率限制
)
@on_command("test")
@rate_limit(5, 60)
async def test_handler(self, event: MessageEvent, ctx: Context):
await event.reply("测试")
```
---
## 🔧 常用功能速查
### 1. LLM 对话
```python
# 简单对话
reply = await ctx.llm.chat("你好")
# 带历史对话
from astrbot_sdk.clients.llm import ChatMessage
history = [
ChatMessage(role="user", content="我叫小明"),
ChatMessage(role="assistant", content="你好小明!"),
]
reply = await ctx.llm.chat("你记得我吗?", history=history)
# 流式对话
async for chunk in ctx.llm.stream_chat("讲个故事"):
print(chunk, end="")
```
### 2. 数据持久化
```python
# DB 客户端(精确匹配)
await ctx.db.set("user:123", {"name": "Alice"})
data = await ctx.db.get("user:123")
# Memory 客户端(语义搜索)
await ctx.memory.save("user_pref", {"theme": "dark"})
results = await ctx.memory.search("用户喜欢什么颜色")
```
### 3. 消息发送
```python
# 简单文本
await ctx.platform.send(event.session_id, "消息内容")
# 图片
await ctx.platform.send_image(event.session_id, "https://example.com/img.jpg")
# 消息链
from astrbot_sdk.message_components import Plain, Image
chain = [Plain("文字"), Image(url="https://example.com/img.jpg")]
await ctx.platform.send_chain(event.session_id, chain)
```
### 4. 文件处理
```python
from astrbot_sdk.message_components import Image
# 注册文件到文件服务
img = Image.fromFileSystem("/path/to/image.jpg")
public_url = await img.register_to_file_service()
```
---
## 🛠️ 高级功能
### 1. LLM 工具注册
```python
async def search_weather(location: str) -> str:
return f"{location} 今天晴天"
await ctx.register_llm_tool(
name="search_weather",
parameters_schema={
"type": "object",
"properties": {
"location": {"type": "string", "description": "城市名称"}
},
"required": ["location"]
},
desc="搜索天气信息",
func_obj=search_weather
)
```
### 2. Web API 注册
```python
from astrbot_sdk.decorators import provide_capability
@provide_capability(
name="my_plugin.api",
description="处理 HTTP 请求"
)
async def handle_api(request_id: str, payload: dict, cancel_token):
return {"status": 200, "body": {"result": "ok"}}
await ctx.http.register_api(
route="/my-api",
handler=handle_api,
methods=["GET", "POST"]
)
```
### 3. 后台任务
```python
async def background_work():
while True:
await asyncio.sleep(60)
ctx.logger.info("每分钟执行一次")
task = await ctx.register_task(background_work(), "定时任务")
```
---
## 📋 最佳实践
### 1. 错误处理
```python
from astrbot_sdk.errors import AstrBotError
@on_command("risky")
async def risky_handler(self, event: MessageEvent, ctx: Context):
try:
result = await risky_operation()
await event.reply(f"成功: {result}")
except AstrBotError as e:
# SDK 错误包含用户友好的提示
await event.reply(e.hint or e.message)
except ValueError as e:
await event.reply(f"参数错误: {e}")
except Exception as e:
ctx.logger.error(f"操作失败: {e}", exc_info=e)
raise
```
### 2. 日志记录
```python
# 不同级别的日志
ctx.logger.debug("调试信息")
ctx.logger.info("普通信息")
ctx.logger.warning("警告信息")
ctx.logger.error("错误信息")
# 绑定上下文
logger = ctx.logger.bind(user_id=event.user_id)
logger.info("用户操作")
```
### 3. 配置管理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
config = await ctx.metadata.get_plugin_config()
# 提供默认值
self.timeout = config.get("timeout", 30)
# 验证必需配置
if "api_key" not in config:
raise ValueError("缺少必需配置: api_key")
self.api_key = config["api_key"]
```
### 4. 资源清理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
self._session = aiohttp.ClientSession()
self._task = asyncio.create_task(self.background_task())
async def on_stop(self, ctx):
if hasattr(self, '_task'):
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
if hasattr(self, '_session'):
await self._session.close()
```
---
## 🔍 注意事项
1. **异步操作**:所有客户端方法都是异步的,需要使用 `await`
2. **插件隔离**:每个插件有独立的 Context 实例
3. **错误处理**:所有远程调用都可能失败,建议使用 try-except
4. **Memory vs DB**:
- Memory: 语义搜索,适合 AI 上下文
- DB: 精确匹配,适合结构化数据
5. **平台标识**:使用 UMO 格式 `"platform:instance:session_id"`
6. **装饰器顺序**:事件触发 → 过滤器 → 限制器 → 修饰器
7. **安全提示**:
- 不要在插件中存储敏感信息(API Key 等应使用配置)
- 验证所有用户输入
- 注意资源泄漏(任务、连接等需要正确清理)
- 遵循最小权限原则
---
## 🐛 调试技巧
### 启用调试日志
```python
# 在插件中获取 logger
logger = ctx.logger
# 记录详细信息
logger.debug(f"收到消息: {event.text}")
logger.debug(f"用户ID: {event.user_id}")
```
### 使用测试框架
```python
from astrbot_sdk.testing import PluginTestHarness
async def test_my_plugin():
harness = PluginTestHarness()
plugin = harness.load_plugin("my_plugin.main:MyPlugin")
# 模拟事件
result = await harness.simulate_command("/hello")
assert result.text == "Hello!"
```
---
## 📞 获取帮助
- **查看详细文档**:[docs/](./)
- **完整 API 索引**:[09_api_reference.md](./09_api_reference.md)
- **错误处理指南**:[06_error_handling.md](./06_error_handling.md)
- **安全检查清单**:[11_security_checklist.md](./11_security_checklist.md)
- **提交问题**:[GitHub Issues](https://github.com/AstrBotDevs/astrbot-sdk/issues)
- **参与讨论**:[GitHub Discussions](https://github.com/AstrBotDevs/astrbot-sdk/discussions)
---
## 📚 版本信息
- **SDK 版本**: v4.0
- **最后更新**: 2026-03-17
- **Python 要求**: >= 3.12
- **协议版本**: P0.6
---
## 📝 文档贡献
如果您发现文档中的错误或想改进文档,欢迎提交 PR!
**文档规范**:
- 使用清晰的代码示例
- 包含错误处理示例
- 标注 API 的稳定性和版本要求
- 提供初级和高级两种使用方式
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,651 +0,0 @@
# 错误处理 API 完整参考
## 概述
AstrBot SDK 提供了统一的错误处理机制,支持跨进程传递错误信息。所有可预期的错误都应使用 `AstrBotError` 类或其工厂方法创建。
**模块路径**: `astrbot_sdk.errors`
---
## 目录
- [错误处理流程](#错误处理流程)
- [导入方式](#导入方式)
- [ErrorCodes - 错误码常量](#errorcodes---错误码常量)
- [AstrBotError - 错误类](#astrboterror---错误类)
- [使用示例](#使用示例)
- [最佳实践](#最佳实践)
---
## 导入方式
```python
# 从主模块导入
from astrbot_sdk import AstrBotError
# 从 errors 模块导入
from astrbot_sdk.errors import AstrBotError, ErrorCodes
```
---
## 错误处理流程
```python
# 1. 抛出错误
raise AstrBotError.invalid_input("参数不能为空")
# 2. 错误被捕获并序列化为 payload
# 3. 跨进程传输后反序列化
# 4. 在 on_error 钩子中统一处理
```
```python
class MyPlugin(Star):
async def on_error(self, error: AstrBotError) -> None:
if error.retryable:
# 可重试的错误
ctx.logger.warning(f"可重试错误: {error.message}")
else:
# 不可重试的错误
ctx.logger.error(f"错误: {error.hint or error.message}")
```
---
## ErrorCodes - 错误码常量
稳定的错误码常量,用于标识不同类型的错误。
### 定义
```python
class ErrorCodes:
"""AstrBot v4 的稳定错误码常量。"""
```
### 错误码列表
#### 不可重试错误(retryable=False)
| 错误码 | 说明 | 默认提示 |
|--------|------|----------|
| `UNKNOWN_ERROR` | 未知错误 | - |
| `LLM_NOT_CONFIGURED` | LLM 未配置 | - |
| `CAPABILITY_NOT_FOUND` | 能力未找到 | 请确认 AstrBot Core 是否已注册该 capability |
| `PERMISSION_DENIED` | 权限被拒绝 | - |
| `LLM_ERROR` | LLM 错误 | - |
| `INVALID_INPUT` | 输入无效 | 请检查调用参数 |
| `CANCELLED` | 调用被取消 | - |
| `PROTOCOL_VERSION_MISMATCH` | 协议版本不匹配 | 请升级 astrbot_sdk 至最新版本 |
| `PROTOCOL_ERROR` | 协议错误 | 请检查通信双方的协议实现 |
| `INTERNAL_ERROR` | 内部错误 | 请联系插件作者 |
| `RATE_LIMITED` | 速率限制 | 操作过于频繁,请稍后再试 |
| `COOLDOWN_ACTIVE` | 冷却中 | - |
#### 可重试错误(retryable=True)
| 错误码 | 说明 | 默认提示 |
|--------|------|----------|
| `CAPABILITY_TIMEOUT` | 能力调用超时 | - |
| `NETWORK_ERROR` | 网络错误 | 网络请求失败,请稍后重试 |
| `LLM_TEMPORARY_ERROR` | LLM 临时错误 | - |
---
## AstrBotError - 错误类
AstrBot SDK 的标准错误类型,支持跨进程传递。
### 类定义
```python
@dataclass(slots=True)
class AstrBotError(Exception):
code: str
message: str
hint: str = ""
retryable: bool = False
docs_url: str = ""
details: dict[str, Any] | None = None
```
### 属性说明
| 属性 | 类型 | 说明 |
|------|------|------|
| `code` | `str` | 错误码,来自 ErrorCodes 常量 |
| `message` | `str` | 错误消息,面向开发者 |
| `hint` | `str` | 用户提示,面向终端用户 |
| `retryable` | `bool` | 是否可重试 |
| `docs_url` | `str` | 文档链接 |
| `details` | `dict[str, Any] \| None` | 详细信息 |
---
## 工厂方法
### `cancelled(message)`
创建取消错误。
```python
@classmethod
def cancelled(cls, message: str = "调用被取消") -> AstrBotError
```
**参数**:
- `message` (`str`): 错误消息
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.cancelled("用户取消操作")
```
---
### `capability_not_found(name)`
创建能力未找到错误。
```python
@classmethod
def capability_not_found(cls, name: str) -> AstrBotError
```
**参数**:
- `name` (`str`): 未找到的能力名称
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.capability_not_found("my_plugin.custom_capability")
```
---
### `invalid_input(message, *, hint, docs_url, details)`
创建输入无效错误。
```python
@classmethod
def invalid_input(
cls,
message: str,
*,
hint: str = "请检查调用参数",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError
```
**参数**:
- `message` (`str`): 详细错误消息
- `hint` (`str`): 用户提示,默认 "请检查调用参数"
- `docs_url` (`str`): 文档链接
- `details` (`dict[str, Any] | None`): 详细信息
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.invalid_input(
"参数格式错误",
hint="请使用 JSON 格式",
details={"expected": "json", "received": "text"}
)
```
---
### `protocol_version_mismatch(message)`
创建协议版本不匹配错误。
```python
@classmethod
def protocol_version_mismatch(cls, message: str) -> AstrBotError
```
**参数**:
- `message` (`str`): 详细错误消息
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.protocol_version_mismatch("SDK 版本 4.0 与 Core 版本 3.9 不兼容")
```
---
### `protocol_error(message)`
创建协议错误。
```python
@classmethod
def protocol_error(cls, message: str) -> AstrBotError
```
**参数**:
- `message` (`str`): 详细错误消息
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.protocol_error("无效的 payload 格式")
```
---
### `internal_error(message, *, hint, docs_url, details)`
创建内部错误。
```python
@classmethod
def internal_error(
cls,
message: str,
*,
hint: str = "请联系插件作者",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError
```
**参数**:
- `message` (`str`): 详细错误消息
- `hint` (`str`): 用户提示,默认 "请联系插件作者"
- `docs_url` (`str`): 文档链接
- `details` (`dict[str, Any] | None`): 详细信息
**返回**: `AstrBotError` 实例
**示例**:
```python
raise AstrBotError.internal_error(
"处理逻辑异常",
hint="请检查日志并联系插件作者",
details={"traceback": "..."}
)
```
---
### `network_error(message, *, hint, docs_url, details)`
创建网络错误。
```python
@classmethod
def network_error(
cls,
message: str,
*,
hint: str = "网络请求失败,请稍后重试",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError
```
**参数**:
- `message` (`str`): 详细错误消息
- `hint` (`str`): 用户提示,默认 "网络请求失败,请稍后重试"
- `docs_url` (`str`): 文档链接
- `details` (`dict[str, Any] | None`): 详细信息
**返回**: `AstrBotError` 实例
**特性**: `retryable=True`
**示例**:
```python
raise AstrBotError.network_error(
"连接超时",
hint="网络不稳定,请稍后重试",
details={"url": "...", "timeout": 30}
)
```
---
### `rate_limited(*, hint, details)`
创建速率限制错误。
```python
@classmethod
def rate_limited(
cls,
*,
hint: str = "操作过于频繁,请稍后再试。",
details: dict[str, Any] | None = None,
) -> AstrBotError
```
**参数**:
- `hint` (`str`): 用户提示,默认 "操作过于频繁,请稍后再试。"
- `details` (`dict[str, Any] | None`): 详细信息
**返回**: `AstrBotError` 实例
**特性**: `retryable=False`
**示例**:
```python
raise AstrBotError.rate_limited(
hint="每分钟最多调用 5 次",
details={"limit": 5, "window": 60, "remaining": 0}
)
```
---
### `cooldown_active(*, hint, details)`
创建冷却中错误。
```python
@classmethod
def cooldown_active(
cls,
*,
hint: str,
details: dict[str, Any] | None = None,
) -> AstrBotError
```
**参数**:
- `hint` (`str`): 用户提示
- `details` (`dict[str, Any] | None`): 详细信息
**返回**: `AstrBotError` 实例
**特性**: `retryable=False`
**示例**:
```python
raise AstrBotError.cooldown_active(
hint="技能冷却中,还需等待 25 秒",
details={"cooldown": 30, "remaining": 25}
)
```
---
## 实例方法
### `to_payload()`
序列化为可传输的字典格式,用于跨进程传递错误信息。
```python
def to_payload(self) -> dict[str, object]
```
**返回**: `dict[str, object]` - 包含错误信息的字典
**返回格式**:
```python
{
"code": "invalid_input",
"message": "参数格式错误",
"hint": "请使用 JSON 格式",
"retryable": False,
"docs_url": "",
"details": {"expected": "json", "received": "text"}
}
```
---
### `from_payload(payload)`
从字典反序列化错误实例。
```python
@classmethod
def from_payload(cls, payload: dict[str, object]) -> AstrBotError
```
**参数**:
- `payload` (`dict[str, object]`): 包含错误信息的字典
**返回**: `AstrBotError` 实例
**示例**:
```python
payload = error.to_payload()
restored_error = AstrBotError.from_payload(payload)
```
---
### `__str__()`
返回错误消息。
```python
def __str__(self) -> str
```
**返回**: `str` - `message` 属性的值
---
## 使用示例
### 基本错误处理
```python
from astrbot_sdk import AstrBotError
from astrbot_sdk.errors import ErrorCodes
@on_command("divide")
async def divide(self, event: MessageEvent, a: int, b: int):
if b == 0:
raise AstrBotError.invalid_input(
"除数不能为零",
hint="请输入非零的除数"
)
return event.plain_result(f"{a} / {b} = {a / b}")
```
### 带详细信息的错误
```python
@on_command("search")
async def search(self, event: MessageEvent, keyword: str):
if not keyword or len(keyword.strip()) == 0:
raise AstrBotError.invalid_input(
"搜索关键词不能为空",
hint="请输入要搜索的关键词",
details={
"field": "keyword",
"constraint": "non_empty",
"provided": keyword
}
)
# 执行搜索...
```
### 捕获和处理错误
```python
@on_command("risky")
async def risky_operation(self, event: MessageEvent):
try:
result = await some_network_request()
return event.plain_result(f"成功: {result}")
except AstrBotError as e:
ctx.logger.error(f"操作失败: {e.message}")
if e.retryable:
await event.reply(f"操作失败(可重试): {e.hint or e.message}")
else:
await event.reply(f"操作失败: {e.hint or e.message}")
```
### 在插件中处理错误
```python
class MyPlugin(Star):
async def on_error(self, error: AstrBotError) -> None:
"""统一处理插件中的所有错误"""
if error.code == ErrorCodes.CAPABILITY_NOT_FOUND:
self.logger.error(f"能力未找到: {error.message}")
elif error.code == ErrorCodes.NETWORK_ERROR:
self.logger.warning(f"网络错误: {error.message}")
elif error.retryable:
self.logger.warning(f"可重试错误: {error.code} - {error.message}")
else:
self.logger.error(f"错误: {error.code} - {error.message}")
```
### 检查特定错误码
```python
try:
await some_capability_call()
except AstrBotError as e:
if e.code == ErrorCodes.RATE_LIMITED:
remaining = e.details.get("remaining", 0)
await event.reply(f"请求过多,请稍后再试。剩余次数: {remaining}")
elif e.code == ErrorCodes.CAPABILITY_TIMEOUT:
await event.reply("请求超时,请稍后重试")
else:
await event.reply(f"错误: {e.hint or e.message}")
```
### 自定义错误(使用通用构造方法)
```python
# 使用通用构造方法创建自定义错误
error = AstrBotError(
code="custom_error_code",
message="自定义错误消息",
hint="这是给用户的提示",
retryable=False,
details={"custom_field": "custom_value"}
)
raise error
```
---
## 最佳实践
### 1. 使用工厂方法而非直接构造
```python
# 推荐
raise AstrBotError.invalid_input("参数错误")
# 不推荐(除非需要自定义错误码)
raise AstrBotError(
code=ErrorCodes.INVALID_INPUT,
message="参数错误",
hint="请检查调用参数"
)
```
### 2. 提供用户友好的提示
```python
# 推荐
raise AstrBotError.invalid_input(
"参数 'count' 必须为正整数",
hint="请输入大于 0 的数字"
)
# 不推荐
raise AstrBotError.invalid_input("参数错误")
```
### 3. 使用 details 提供调试信息
```python
raise AstrBotError.invalid_input(
"参数验证失败",
hint="请检查输入格式",
details={
"field": "email",
"pattern": "^[\\w\\.-]+@[\\w\\.-]+\\.\\w+$",
"provided": "invalid-email"
}
)
```
### 4. 区分可重试和不可重试错误
```python
# 网络错误 - 可重试
raise AstrBotError.network_error("连接失败")
# 参数错误 - 不可重试
raise AstrBotError.invalid_input("参数类型错误")
```
### 5. 在 on_error 中集中处理
```python
class MyPlugin(Star):
async def on_error(self, error: AstrBotError) -> None:
# 记录所有错误
self.logger.error(f"错误: [{error.code}] {error.message}")
# 可重试错误记录为警告级别
if error.retryable:
self.logger.warning(f"可重试错误,考虑实现重试逻辑")
# 特定错误码的特殊处理
if error.code == ErrorCodes.CAPABILITY_NOT_FOUND:
self.logger.critical("请检查 AstrBot Core 配置")
```
### 6. 向用户展示适当的错误信息
```python
try:
result = await operation()
except AstrBotError as e:
# 优先使用 hint(面向用户)
user_message = e.hint or e.message
await event.reply(user_message)
# 记录完整的错误信息(面向开发者)
ctx.logger.error(f"操作失败: {e.code} - {e.message}", extra=e.details)
```
---
## 相关模块
- **事件处理**: `astrbot_sdk.events.MessageEvent`
- **上下文**: `astrbot_sdk.context.Context`
- **插件基类**: `astrbot_sdk.star.Star`
---
**版本**: v4.0
**模块**: `astrbot_sdk.errors`
**最后更新**: 2026-03-17
@@ -1,948 +0,0 @@
# 消息组件 API 完整参考
## 概述
消息组件是用于构建聊天消息的各种元素。每个组件代表消息中的一种特定内容类型,可以单独使用或组合成消息链。
**模块路径**: `astrbot_sdk.message_components`
---
## 目录
- [BaseMessageComponent - 基类](#basemessagecomponent---基类)
- [Plain - 纯文本组件](#plain---纯文本组件)
- [At / AtAll - @组件](#at--atall---组件)
- [Image - 图片组件](#image---图片组件)
- [Record - 语音组件](#record---语音组件)
- [Video - 视频组件](#video---视频组件)
- [File - 文件组件](#file---文件组件)
- [Reply - 回复组件](#reply---回复组件)
- [Poke - 戳一戳组件](#poke---戳一戳组件)
- [Forward - 转发组件](#forward---转发组件)
- [MessageChain - 消息链](#messagechain---消息链)
- [辅助函数](#辅助函数)
---
## 导入方式
```python
# 从主模块导入(推荐)
from astrbot_sdk import (
Plain, At, AtAll, Image, Record, Video, File, Reply, Poke, Forward,
MessageChain, MessageBuilder
)
# 从子模块导入
from astrbot_sdk.message_components import (
Plain, At, AtAll, Image, Record, Video, File, Reply, Poke, Forward
)
from astrbot_sdk.message_result import MessageChain, MessageBuilder
# 辅助函数
from astrbot_sdk.message_components import (
payload_to_component,
component_to_payload_sync,
component_to_payload,
)
```
---
## BaseMessageComponent - 基类
所有消息组件的基类。
### 类定义
```python
class BaseMessageComponent:
type: str = "unknown"
def toDict(self) -> dict[str, Any]:
"""同步转换为字典 payload"""
async def to_dict(self) -> dict[str, Any]:
"""异步转换为字典 payload"""
```
---
## Plain - 纯文本组件
最简单的消息组件,只包含文本内容。
### 类定义
```python
class Plain(BaseMessageComponent):
type = "plain" # 序列化时为 "text"
def __init__(self, text: str, convert: bool = True, **_: Any) -> None:
self.text = text
self.convert = convert
```
### 构造方法
```python
from astrbot_sdk import Plain
# 基本用法
text = Plain("Hello World")
# 不自动 strip(保留首尾空格)
text = Plain(" Hello ", convert=False)
```
### 序列化格式
```python
# toDict() 会自动 strip 文本
{
"type": "text",
"data": {"text": "Hello World"}
}
# to_dict() 保留原始文本
{
"type": "text",
"data": {"text": " Hello "}
}
```
### 使用示例
```python
@on_command("echo")
async def echo(self, event: MessageEvent, text: str):
await event.reply_chain([Plain(f"你说: {text}")])
```
---
## At / AtAll - @组件
用于在消息中提及用户。
### At - @某人
#### 类定义
```python
class At(BaseMessageComponent):
type = "at"
def __init__(self, qq: int | str, name: str | None = "", **_: Any) -> None:
self.qq = qq
self.name = name or ""
```
#### 构造方法
```python
from astrbot_sdk import At
# @ 单个用户
at = At(123456)
at = At("123456", name="张三")
```
#### 序列化格式
```python
{
"type": "at",
"data": {"qq": "123456"}
}
```
---
### AtAll - @全体成员
#### 类定义
```python
class AtAll(At):
def __init__(self, **_: Any) -> None:
super().__init__(qq="all")
```
#### 构造方法
```python
from astrbot_sdk import AtAll
at_all = AtAll()
```
#### 序列化格式
```python
{
"type": "at",
"data": {"qq": "all"}
}
```
---
### 使用示例
```python
from astrbot_sdk import At, AtAll, Plain
@on_command("at_test")
async def at_test(self, event: MessageEvent):
await event.reply_chain([
Plain("你好 "),
At(event.user_id or "123456"),
Plain("!"),
AtAll(),
Plain("所有人请注意!")
])
```
---
## Image - 图片组件
用于在消息中发送图片。
### 类定义
```python
class Image(BaseMessageComponent):
type = "image"
def __init__(self, file: str | None, **kwargs: Any) -> None:
self.file = file or ""
self._type = kwargs.get("_type", "")
self.subType = kwargs.get("subType", 0)
self.url = kwargs.get("url", "")
self.cache = kwargs.get("cache", True)
self.id = kwargs.get("id", 40000)
self.c = kwargs.get("c", 2)
self.path = kwargs.get("path", "")
self.file_unique = kwargs.get("file_unique", "")
```
### 静态构造方法
#### `fromURL(url, **kwargs)`
从 URL 创建图片。
```python
from astrbot_sdk import Image
img = Image.fromURL("https://example.com/image.jpg")
```
#### `fromFileSystem(path, **kwargs)`
从本地文件系统创建图片。
```python
img = Image.fromFileSystem("/path/to/image.jpg")
```
#### `fromBase64(base64_data, **kwargs)`
从 Base64 数据创建图片。
```python
img = Image.fromBase64("iVBORw0KGgo...")
```
#### `fromBytes(data, **kwargs)`
从字节数据创建图片。
```python
img = Image.fromBytes(b"...")
```
### 实例方法
#### `convert_to_file_path()`
将图片转换为本地文件路径(下载或解码)。
```python
path = await img.convert_to_file_path()
```
#### `register_to_file_service()`
将图片注册到文件服务,返回可访问 URL。
```python
public_url = await img.register_to_file_service()
```
### 支持的格式
```python
# URL: "https://example.com/image.jpg"
# 本地文件: "file:///absolute/path/to/image.jpg"
# Base64: "base64://iVBORw0KGgo..."
```
### 使用示例
```python
from astrbot_sdk import Image
@on_command("cat")
async def cat(self, event: MessageEvent):
await event.reply_image("https://example.com/cat.jpg")
@on_command("local_img")
async def local_img(self, event: MessageEvent):
await event.reply_image("file:///path/to/image.jpg")
```
---
## Record - 语音组件
用于在消息中发送语音/音频。
### 类定义
```python
class Record(BaseMessageComponent):
type = "record"
def __init__(self, file: str | None, **kwargs: Any) -> None:
self.file = file or ""
self.magic = kwargs.get("magic", False)
self.url = kwargs.get("url", "")
self.cache = kwargs.get("cache", True)
self.proxy = kwargs.get("proxy", True)
self.timeout = kwargs.get("timeout", 0)
self.text = kwargs.get("text")
self.path = kwargs.get("path")
```
### 静态构造方法
#### `fromFileSystem(path, **kwargs)`
```python
from astrbot_sdk import Record
audio = Record.fromFileSystem("/path/to/audio.mp3")
```
#### `fromURL(url, **kwargs)`
```python
audio = Record.fromURL("https://example.com/audio.mp3")
```
### 实例方法
#### `convert_to_file_path()`
```python
path = await audio.convert_to_file_path()
```
#### `register_to_file_service()`
```python
public_url = await audio.register_to_file_service()
```
---
## Video - 视频组件
用于在消息中发送视频。
### 类定义
```python
class Video(BaseMessageComponent):
type = "video"
def __init__(self, file: str, **kwargs: Any) -> None:
self.file = file
self.cover = kwargs.get("cover", "")
self.c = kwargs.get("c", 2)
self.path = kwargs.get("path", "")
```
### 静态构造方法
#### `fromFileSystem(path, **kwargs)`
```python
from astrbot_sdk import Video
video = Video.fromFileSystem("/path/to/video.mp4")
```
#### `fromURL(url, **kwargs)`
```python
video = Video.fromURL("https://example.com/video.mp4")
```
---
## File - 文件组件
用于在消息中发送文件附件。
### 类定义
```python
class File(BaseMessageComponent):
type = "file"
def __init__(self, name: str, file: str = "", url: str = "") -> None:
self.name = name
self.file_ = file
self.url = url
```
### 属性
- `name` (`str`): 文件名
- `file_` (`str`): 本地文件路径(内部使用)
- `url` (`str`): 文件 URL
### file 属性 (getter/setter)
```python
@property
def file(self) -> str:
return self.file_
@file.setter
def file(self, value: str) -> None:
if value.startswith(("http://", "https://")):
self.url = value
else:
self.file_ = value
```
### 构造方法
```python
from astrbot_sdk import File
# URL 文件
file1 = File(name="document.pdf", url="https://example.com/doc.pdf")
# 本地文件
file2 = File(name="image.jpg", file="/path/to/image.jpg")
```
### 实例方法
#### `get_file(allow_return_url=False)`
获取文件路径或 URL。
```python
path = await file.get_file()
# 优先返回 URL
path = await file.get_file(allow_return_url=True)
```
#### `register_to_file_service()`
```python
public_url = await file.register_to_file_service()
```
### 序列化格式
```python
# toDict()
{
"type": "file",
"data": {
"name": "文件名.pdf",
"file": "本地路径或URL"
}
}
# to_dict()
{
"type": "file",
"data": {
"name": "文件名.pdf",
"file": "优先返回URL,否则本地路径"
}
}
```
---
## Reply - 回复组件
用于回复某条消息。
### 类定义
```python
class Reply(BaseMessageComponent):
type = "reply"
def __init__(self, **kwargs: Any) -> None:
self.id = kwargs.get("id", "")
self.chain = _coerce_reply_chain(kwargs.get("chain", []))
self.sender_id = kwargs.get("sender_id", 0)
self.sender_nickname = kwargs.get("sender_nickname", "")
self.time = kwargs.get("time", 0)
self.message_str = kwargs.get("message_str", "")
self.text = kwargs.get("text", "")
self.qq = kwargs.get("qq", 0)
self.seq = kwargs.get("seq", 0)
```
### 构造方法
```python
from astrbot_sdk import Reply, Plain
reply = Reply(
id="msg_123",
sender_id="789",
sender_nickname="张三",
chain=[Plain("被回复的消息")]
)
```
### 实例方法
#### `toDict()` / `to_dict()`
序列化为字典。
---
## Poke - 戳一戳组件
用于发送戳一戳操作。
### 类定义
```python
class Poke(BaseMessageComponent):
type = "poke"
def __init__(self, poke_type: str | int | None = None, **kwargs: Any) -> None:
self._type = str(poke_type)
self.id = kwargs.get("id")
self.qq = kwargs.get("qq", 0)
```
### 构造方法
```python
from astrbot_sdk import Poke
poke = Poke(poke_type="126", qq="123456")
```
---
## Forward - 转发组件
用于转发消息。
### 类定义
```python
class Forward(BaseMessageComponent):
type = "forward"
def __init__(self, id: str, **_: Any) -> None:
self.id = id
```
### 构造方法
```python
from astrbot_sdk import Forward
forward = Forward(id="forward_msg_123")
```
---
## UnknownComponent - 未知组件
用于表示无法识别的组件类型。
### 类定义
```python
class UnknownComponent(BaseMessageComponent):
type = "unknown"
def __init__(
self,
*,
raw_type: str = "unknown",
raw_data: dict[str, Any] | None = None,
) -> None:
self.raw_type = raw_type
self.raw_data = raw_data or {}
```
### 构造方法
```python
from astrbot_sdk import UnknownComponent
unknown = UnknownComponent(
raw_type="custom_type",
raw_data={"field": "value"}
)
```
### 说明
当 `payload_to_component()` 遇到无法识别的组件类型时,会返回 `UnknownComponent` 实例,保留原始数据以便调试。
---
## MessageChain - 消息链
用于组合多个消息组件。
### 类定义
```python
@dataclass(slots=True)
class MessageChain:
components: list[BaseMessageComponent] = field(default_factory=list)
```
### 构造方法
```python
from astrbot_sdk.message_result import MessageChain
from astrbot_sdk.message_components import Plain, At
# 空消息链
chain = MessageChain()
# 带初始组件
chain = MessageChain([Plain("Hello"), At("123456")])
```
### 实例方法
#### `append(component)`
追加单个组件,返回 self 支持链式调用。
```python
chain.append(Plain("More text"))
```
#### `extend(components)`
追加多个组件。
```python
chain.extend([Plain("A"), Plain("B")])
```
#### `to_payload()`
转换为协议 payload。
```python
payload = chain.to_payload()
```
#### `get_plain_text(with_other_comps_mark=False)`
提取纯文本内容。
```python
text = chain.get_plain_text()
```
---
## MessageBuilder - 消息构建器
流式构建消息链的工具类。
### 使用示例
```python
from astrbot_sdk.message_result import MessageBuilder
chain = (MessageBuilder()
.text("Hello ")
.at("123456")
.text("!\n")
.image("https://example.com/img.jpg")
.build())
await event.reply_chain(chain)
```
### 可用方法
- `.text(content)` - 添加文本
- `.at(user_id)` - 添加@用户
- `.at_all()` - 添加@全体成员
- `.image(url)` - 添加图片
- `.record(url)` - 添加语音
- `.video(url)` - 添加视频
- `.file(name, url=...)` - 添加文件
- `.build()` - 构建消息链
---
## 辅助函数
### `payload_to_component(payload)`
将协议 payload 转换为消息组件。
```python
from astrbot_sdk.message_components import payload_to_component
component = payload_to_component(payload)
```
### `component_to_payload_sync(component)`
将组件同步转换为 payload。
```python
from astrbot_sdk.message_components import component_to_payload_sync
payload = component_to_payload_sync(component)
```
### `component_to_payload(component)`
将组件异步转换为 payload。
```python
from astrbot_sdk.message_components import component_to_payload
payload = await component_to_payload(component)
```
---
### `is_message_component(value)`
检查值是否为消息组件。
```python
from astrbot_sdk.message_components import is_message_component
if is_message_component(value):
print("是消息组件")
```
---
### `payloads_to_components(payloads)`
批量将 payload 列表转换为组件列表。
```python
from astrbot_sdk.message_components import payloads_to_components
components = payloads_to_components(payload_list)
```
---
### `build_media_component_from_url(url, *, kind)`
从 URL 构建媒体组件。
```python
from astrbot_sdk.message_components import build_media_component_from_url
# 自动识别类型
component = build_media_component_from_url("https://example.com/image.jpg")
# 指定类型
component = build_media_component_from_url("https://example.com/file", kind="image")
```
---
## MediaHelper - 媒体辅助类
提供媒体处理的静态方法。
### `from_url(url, *, kind)`
从 URL 创建媒体组件。
**签名**:
```python
@staticmethod
async def from_url(
url: str,
*,
kind: str = "auto"
) -> BaseMessageComponent
```
**参数**:
- `url`: 媒体 URL
- `kind`: 媒体类型(`"auto"`, `"image"`, `"record"`, `"video"`, `"file"`)
**返回**: 对应的媒体组件
**示例**:
```python
from astrbot_sdk.message_components import MediaHelper
# 自动识别
img = await MediaHelper.from_url("https://example.com/photo.jpg")
# 指定类型
video = await MediaHelper.from_url("https://example.com/video.mp4", kind="video")
```
---
### `download(url, save_dir)`
下载媒体文件到指定目录。
**签名**:
```python
@staticmethod
async def download(url: str, save_dir: Path) -> Path
```
**参数**:
- `url`: 媒体 URL(仅支持 http/https)
- `save_dir`: 保存目录路径
**返回**: `Path` - 下载后的文件路径
**异常**:
- `AstrBotError`: 下载失败时抛出
**示例**:
```python
from pathlib import Path
from astrbot_sdk.message_components import MediaHelper
try:
path = await MediaHelper.download(
"https://example.com/image.jpg",
Path("./downloads")
)
print(f"下载到: {path}")
except AstrBotError as e:
print(f"下载失败: {e.message}")
```
---
## 使用示例
### 处理图片消息
```python
@on_message()
async def save_image(self, event: MessageEvent):
images = event.get_images()
if not images:
await event.reply("消息中没有图片")
return
for img in images:
try:
path = await img.convert_to_file_path()
# 保存图片...
await event.reply(f"已保存: {path}")
except Exception as e:
await event.reply(f"保存失败: {e}")
```
### 检测@和群聊/私聊
```python
@on_command("check")
async def check(self, event: MessageEvent):
# 检查是否群聊
if event.is_group_chat():
await event.reply("这是群聊消息")
elif event.is_private_chat():
await event.reply("这是私聊消息")
# 检查@的用户
at_users = event.get_at_users()
if at_users:
await event.reply(f"你@了: {', '.join(at_users)}")
```
### 返回富文本结果
```python
@on_command("info")
async def info(self, event: MessageEvent):
return event.chain_result([
Plain(f"用户: {event.sender_name}\n"),
Plain(f"ID: {event.user_id}\n"),
Plain(f"平台: {event.platform}"),
])
```
---
## 注意事项
1. **序列化差异**:
- `Plain.toDict()` 会 strip 文本
- `Plain.to_dict()` 保留原始文本
- `File.toDict()` 和 `to_dict()` 对 file 字段处理不同
2. **路径格式**:
- 本地文件: `file:///absolute/path` (Windows 下特殊处理)
- URL: `http://` 或 `https://`
- Base64: `base64://<data>`
3. **文件下载**:
- `convert_to_file_path()` 会下载网络文件到临时目录
- `register_to_file_service()` 需要运行时上下文
4. **兼容性**:
- `At` 和 `AtAll` 序列化后的 type 都是 "at"
- `Reply` 的 chain 字段在序列化时递归处理
---
## 相关模块
- **消息组件**: `astrbot_sdk.message_components`
- **消息链**: `astrbot_sdk.message_result.MessageChain`
- **消息构建器**: `astrbot_sdk.message_result.MessageBuilder`
- **协议描述符**: `astrbot_sdk.protocol.descriptors`
---
**版本**: v4.0
**模块**: `astrbot_sdk.message_components`
**最后更新**: 2026-03-17
File diff suppressed because it is too large Load Diff
@@ -1,728 +0,0 @@
# 消息结果 API 完整参考
## 概述
消息结果是用于构建和返回消息结果的类,包括消息链容器、流式构建器和事件结果包装器。
**模块路径**: `astrbot_sdk.message_result`
---
## 目录
- [EventResultType - 事件结果类型枚举](#eventresulttype---事件结果类型枚举)
- [MessageChain - 消息链](#messagechain---消息链)
- [MessageBuilder - 消息构建器](#messagebuilder---消息构建器)
- [MessageEventResult - 消息事件结果](#messageeventresult---消息事件结果)
---
## 导入方式
```python
# 从主模块导入
from astrbot_sdk import MessageChain, MessageBuilder, MessageEventResult
# 从子模块导入
from astrbot_sdk.message_result import (
MessageChain,
MessageBuilder,
MessageEventResult,
EventResultType,
)
# 消息组件(用于构建消息链)
from astrbot_sdk.message_components import Plain, At, Image, File
```
---
## EventResultType - 事件结果类型枚举
事件结果的类型枚举,定义消息结果的类型。
### 定义
```python
class EventResultType(str, Enum):
EMPTY = "empty" # 空结果
CHAIN = "chain" # 消息链结果
PLAIN = "plain" # 纯文本结果
```
### 值说明
| 值 | 说明 |
|------|------|
| `EventResultType.EMPTY` | 空结果,不返回任何内容 |
| `EventResultType.CHAIN` | 消息链结果,返回一个或多个消息组件 |
| `EventResultType.PLAIN` | 纯文本结果,返回文本内容 |
---
## MessageChain - 消息链
消息链是消息组件的容器,用于组合多个组件形成复杂的消息。
### 类定义
```python
@dataclass(slots=True)
class MessageChain:
components: list[BaseMessageComponent] = field(default_factory=list)
```
### 构造方法
#### 空消息链
```python
from astrbot_sdk.message_result import MessageChain
chain = MessageChain()
```
#### 带初始组件
```python
from astrbot_sdk.message_result import MessageChain
from astrbot_sdk.message_components import Plain, At
chain = MessageChain([
Plain("Hello"),
At("123456")
])
```
### 实例方法
#### `append(component)`
追加单个组件,返回 self 支持链式调用。
```python
def append(self, component: BaseMessageComponent) -> MessageChain:
"""追加单个组件,返回 self"""
self.components.append(component)
return self
```
**参数**:
- `component` (`BaseMessageComponent`): 要追加的组件
**返回**: `MessageChain` - self
**示例**:
```python
chain = MessageChain()
chain.append(Plain("Hello "))
.append(At("123456"))
.append(Plain("!"))
```
---
#### `extend(components)`
追加多个组件,返回 self。
```python
def extend(self, components: list[BaseMessageComponent]) -> MessageChain:
"""追加多个组件,返回 self"""
self.components.extend(components)
return self
```
**参数**:
- `components` (`list[BaseMessageComponent]`): 组件列表
**示例**:
```python
chain = MessageChain()
chain.extend([
Plain("A"),
Plain("B"),
Plain("C")
])
```
---
#### `to_payload()`
同步转换为协议 payload。
```python
def to_payload(self) -> list[dict[str, Any]]:
"""转换为协议 payload"""
return [component_to_payload_sync(c) for c in self.components]
```
**返回**: `list[dict]` - 可序列化的字典列表
---
#### `to_payload_async()`
异步转换为协议 payload。
```python
async def to_payload_async(self) -> list[dict[str, Any]]:
"""异步转换为协议 payload"""
return [await component_to_payload(c) for c in self.components]
```
**注意**: 某些组件(如 Reply)的异步序列化可能包含额外逻辑
---
#### `get_plain_text(with_other_comps_mark=False)`
提取纯文本内容。
```python
def get_plain_text(self, with_other_comps_mark: bool = False) -> str:
"""提取纯文本内容"""
texts: list[str] = []
for component in self.components:
if isinstance(component, Plain):
texts.append(component.text)
elif with_other_comps_mark:
texts.append(f"[{component.__class__.__name__}]")
return " ".join(texts)
```
**参数**:
- `with_other_comps_mark`: 是否为非文本组件显示类型标记
**返回**: `str` - 纯文本内容
**示例**:
```python
chain = MessageChain([
Plain("Hello "),
At("123456"),
Plain("!")
])
chain.get_plain_text() # "Hello !"
chain.get_plain_text(True) # "Hello [At] !"
```
---
#### `plain_text(with_other_comps_mark=False)`
`get_plain_text()` 的别名。
```python
def plain_text(self, with_other_comps_mark: bool = False) -> str:
return self.get_plain_text(with_other_comps_mark=with_other_comps_mark)
```
---
### 迭代与长度
```python
# 迭代
for component in chain:
print(f"组件: {component.__class__.__name__}")
# 长度
len(chain) # 组件数量
```
---
### 使用示例
```python
from astrbot_sdk.message_result import MessageChain
from astrbot_sdk.message_components import Plain, At, Image
# 创建并使用
chain = MessageChain([
Plain("Hello "),
At("123456"),
Plain("!"),
Image.fromURL("https://example.com/img.jpg")
])
# 转换为 payload
payload = chain.to_payload()
# 提取文本
text = chain.get_plain_text()
# 链式追加
chain.append(Plain("More text"))
```
---
## MessageBuilder - 消息构建器
流式构建消息链的工具类,提供流畅的 API。
### 类定义
```python
@dataclass(slots=True)
class MessageBuilder:
components: list[BaseMessageComponent] = field(default_factory=list)
```
### 链式方法
所有方法都返回 `self`,支持链式调用。
#### `text(content)`
添加文本组件。
```python
def text(self, content: str) -> MessageBuilder:
"""添加文本组件"""
self.components.append(Plain(content, convert=False))
return self
```
**示例**:
```python
builder = MessageBuilder()
builder.text("Hello ")
```
---
#### `at(user_id)`
添加@组件。
```python
def at(self, user_id: str) -> MessageBuilder:
"""添加@用户"""
self.components.append(At(user_id))
return self
```
---
#### `at_all()`
添加@全体成员。
```python
def at_all(self) -> MessageBuilder:
"""添加@全体成员"""
self.components.append(AtAll())
return self
```
---
#### `image(url)`
添加图片。
```python
def image(self, url: str) -> MessageBuilder:
"""添加图片"""
self.components.append(Image.fromURL(url))
return self
```
---
#### `record(url)`
添加语音。
```python
def record(self, url: str) -> MessageBuilder:
"""添加语音"""
self.components.append(Record.fromURL(url))
return self
```
---
#### `video(url)`
添加视频。
```python
def video(self, url: str) -> MessageBuilder:
"""添加视频"""
self.components.append(Video.fromURL(url))
return self
```
---
#### `file(name, *, file="", url="")`
添加文件。
```python
def file(self, name: str, *, file: str = "", url: str = "") -> MessageBuilder:
"""添加文件"""
self.components.append(File(name=name, file=file, url=url))
return self
```
---
#### `reply(**kwargs)`
添加回复组件。
```python
def reply(self, **kwargs: Any) -> MessageBuilder:
"""添加回复组件"""
self.components.append(Reply(**kwargs))
return self
```
---
#### `append(component)`
添加任意组件。
```python
def append(self, component: BaseMessageComponent) -> MessageBuilder:
"""添加任意组件"""
self.components.append(component)
return self
```
---
#### `extend(components)`
添加多个组件。
```python
def extend(self, components: list[BaseMessageComponent]) -> MessageBuilder:
"""添加多个组件"""
self.components.extend(components)
return self
```
---
#### `build()`
构建 MessageChain。
```python
def build(self) -> MessageChain:
"""构建消息链"""
return MessageChain(list(self.components))
```
**返回**: `MessageChain` - 包含所有组件的消息链对象
---
### 完整使用示例
```python
from astrbot_sdk.message_result import MessageBuilder
from astrbot_sdk.message_components import Plain, At, Image
# 链式构建
chain = (MessageBuilder()
.text("Hello ")
.at("123456")
.text("!\n")
.image("https://example.com/img.jpg")
.build())
# 使用 MessageChain
chain = MessageChain([
Plain("Hello "),
At("123456"),
Plain("!\n"),
Image.fromURL("https://example.com/img.jpg")
])
# 两种方式结果相同
```
---
## MessageEventResult - 消息事件结果
消息事件结果的包装类,用于 handler 返回值。
### 类定义
```python
@dataclass(slots=True)
class MessageEventResult:
type: EventResultType = EventResultType.EMPTY
chain: MessageChain = field(default_factory=MessageChain)
```
### 构造方法
#### 空结果
```python
from astrbot_sdk.message_result import MessageEventResult, EventResultType
result = MessageEventResult()
# 或
result = MessageEventResult(type=EventResultType.EMPTY)
```
---
#### 纯文本结果
```python
result = MessageEventResult(
type=EventResultType.PLAIN,
chain=MessageChain([Plain("返回内容")])
)
```
---
#### 消息链结果
```python
from astrbot_sdk.message_result import MessageEventResult, EventResultType, MessageChain
from astrbot_sdk.message_components import Plain, Image
result = MessageEventResult(
type=EventResultType.CHAIN,
chain=MessageChain([
Plain("文本"),
Image(url="https://example.com/a.png")
])
)
```
---
### 实例方法
#### `to_payload()`
转换为协议 payload。
```python
def to_payload(self) -> dict[str, Any]:
"""转换为协议 payload"""
return {
"type": self.type.value,
"chain": self.chain.to_payload(),
}
```
**返回格式**:
```python
# EMPTY
{"type": "empty", "chain": []}
# CHAIN
{
"type": "chain",
"chain": [
{"type": "text", "data": {"text": "内容"}},
{"type": "image", "data": {"url": "..."}}
]
}
# PLAIN
{
"type": "plain",
"chain": [{"type": "text", "data": {"text": "内容"}}]
}
```
---
#### `from_payload(payload)`
从协议 payload 创建实例。
```python
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> MessageEventResult:
result_type_raw = str(payload.get("type", EventResultType.EMPTY.value))
try:
result_type = EventResultType(result_type_raw)
except ValueError:
result_type = EventResultType.EMPTY
chain_payload = payload.get("chain")
components = (
payloads_to_components(chain_payload)
if isinstance(chain_payload, list)
else []
)
return cls(type=result_type, chain=MessageChain(components))
```
---
### 使用示例
```python
@on_command("return_text")
async def return_text(self, event: MessageEvent):
# 返回纯文本结果
return event.plain_result("返回内容")
@on_command("return_image")
async def return_image(self, event: MessageEvent):
# 返回图片结果
return event.image_result("https://example.com/image.jpg")
@on_command("return_chain")
async def return_chain(self, event: MessageEvent):
# 返回消息链结果
return event.chain_result([
Plain(f"用户: {event.sender_name}"),
Plain(f"ID: {event.user_id}"),
Plain(f"平台: {event.platform}"),
])
```
---
## 使用场景示例
### 场景1: 使用 MessageBuilder 构建复杂消息
```python
@on_command("rich")
async def rich_message(self, event: MessageEvent):
chain = (MessageBuilder()
.text("你好 ")
.at(event.user_id or "123456")
.text("!\n\n")
.image("https://example.com/welcome.jpg")
.text("这是欢迎图片")
.build())
await event.reply_chain(chain)
```
---
### 场景2: 使用 MessageChain 组合组件
```python
@on_command("multi")
async def multi_component(self, event: MessageEvent, count: int):
components = [Plain(f"发送 {count} 条消息:\n")]
for i in range(count):
components.append(Plain(f"{i+1}. "))
if i < count - 1:
components.append(Plain("\n"))
await event.reply_chain(components)
```
---
### 场景3: 返回结构化结果
```python
@on_command("user_info")
async def user_info(self, event: MessageEvent):
return event.chain_result([
Plain(f"用户: {event.sender_name}\n"),
Plain(f"ID: {event.user_id}\n"),
Plain(f"平台: {event.platform}\n"),
Plain(f"消息类型: {event.message_type}\n"),
])
```
---
## 辅助函数
### `coerce_message_chain(value)`
将多种输入格式统一转换为 MessageChain。
**签名**:
```python
def coerce_message_chain(value: Any) -> MessageChain | None
```
**参数**:
- `value`: 要转换的值,支持以下类型:
- `MessageEventResult`: 提取其中的 chain
- `MessageChain`: 直接返回
- `BaseMessageComponent`: 包装为单元素链
- `list[BaseMessageComponent]`: 包装为链
**返回**: `MessageChain | None` - 转换后的消息链,无法转换则返回 None
**示例**:
```python
from astrbot_sdk.message_result import coerce_message_chain, MessageChain
from astrbot_sdk.message_components import Plain, Image
# 从 MessageEventResult 提取
chain = coerce_message_chain(result)
# 从 MessageChain 返回
chain = coerce_message_chain(existing_chain)
# 从单个组件创建
chain = coerce_message_chain(Plain("文本"))
# 从组件列表创建
chain = coerce_message_chain([Plain("文本"), Image.fromURL("url")])
```
---
## 注意事项
1. **MessageChain 可变性**:
- `append()` 和 `extend()` 修改原链并返回 self
- 支持链式调用
- 注意:链式操作会修改原链
2. **异步序列化**:
- 大多数情况用 `to_payload()` 即可
- 包含 `Reply` 组件时建议用 `to_payload_async()`
3. **纯文本提取**:
- `get_plain_text()` 默认忽略非文本组件
- 设置 `with_other_comps_mark=True` 显示类型标记
4. **结果类型**:
- `EMPTY`: 不返回任何内容
- `CHAIN`: 返回一个或多个消息组件
- `PLAIN`: 返回文本内容
---
## 相关模块
- **消息组件**: `astrbot_sdk.message_components`
- **事件结果**: `astrbot_sdk.events.MessageEventResult`
- **事件类型**: `astrbot_sdk.events.EventResultType`
---
**版本**: v4.0
**模块**: `astrbot_sdk.message_result`
**最后更新**: 2026-03-17
@@ -1,740 +0,0 @@
# Star 类 - 插件基类完整参考
## 概述
`Star` 是 AstrBot SDK 的插件基类,所有 v4 原生插件都必须继承此类。它提供了完整的插件生命周期管理、上下文访问和能力集成。
**模块路径**: `astrbot_sdk.star.Star`
---
## 类定义
```python
class Star(PluginKVStoreMixin):
"""v4 原生插件基类"""
__handlers__: tuple[str, ...] # 自动收集的处理器列表
# 生命周期钩子
async def on_start(self, ctx: Any | None = None) -> None
async def on_stop(self, ctx: Any | None = None) -> None
async def initialize(self) -> None
async def terminate(self) -> None
async def on_error(self, error: Exception, event, ctx) -> None
# 便捷属性
@property
def context(self) -> Context | None
# 便捷方法
async def text_to_image(self, text: str, *, return_url: bool = True) -> str
async def html_render(self, tmpl: str, data: dict, *, return_url: bool = True) -> str
# KV 存储方法(继承自 PluginKVStoreMixin)
async def put_kv_data(self, key: str, value: Any) -> None
async def get_kv_data(self, key: str, default: _VT) -> _VT
async def delete_kv_data(self, key: str) -> None
```
---
## 导入方式
```python
# 从主模块导入(推荐)
from astrbot_sdk import Star
# 从子模块导入
from astrbot_sdk.star import Star
# 常用配套导入
from astrbot_sdk import Context, MessageEvent # 上下文和事件
from astrbot_sdk.decorators import on_command, on_message # 装饰器
from astrbot_sdk.errors import AstrBotError # 错误处理
```
---
## 核心属性
### `__handlers__`
自动收集的事件处理器元组。
```python
class MyPlugin(Star):
@on_command("cmd1")
async def cmd1_handler(self, event, ctx):
pass
# MyPlugin.__handlers__ == ("cmd1_handler",)
```
**说明**: 在子类创建时,`__init_subclass__()` 会自动扫描所有装饰了 `@on_command`、`@on_message` 等装饰器的方法,并将处理器名称收集到此元组中。
### `context`
获取当前运行时上下文的属性。
```python
class MyPlugin(Star):
async def some_method(self):
ctx = self.context
if ctx:
await ctx.db.set("key", "value")
```
**返回**: `Context | None` - 仅在生命周期钩子和 Handler 执行期间可用
**注意**: 不要存储此引用,它在插件停止后会被清除
---
## 生命周期钩子
### 1. `on_start(ctx)` - 插件启动钩子
**签名**:
```python
async def on_start(self, ctx: Any | None = None) -> None
```
**参数**:
- `ctx`: 运行时上下文(通常为 `Context` 实例)
**触发时机**: Worker 启动后,在开始处理事件之前调用
**用途**:
- 初始化数据库连接
- 加载配置文件
- 注册 LLM 工具
- 启动后台任务
- 验证外部依赖
**示例**:
```python
class MyPlugin(Star):
async def on_start(self, ctx) -> None:
# 确保 initialize 被调用
await super().on_start(ctx)
# 获取插件数据目录
data_dir = await ctx.get_data_dir()
# 加载配置
config = await ctx.metadata.get_plugin_config()
self.api_key = config.get("api_key", "")
# 注册 LLM 工具
await ctx.register_llm_tool(
name="search",
parameters_schema={...},
desc="搜索信息",
func_obj=self.search_tool
)
# 启动后台任务
await ctx.register_task(
self.background_sync(),
desc="后台数据同步"
)
ctx.logger.info(f"{ctx.plugin_id} 启动成功")
```
**注意事项**:
1. 始终调用 `await super().on_start(ctx)` 确保 `initialize()` 被调用
2. 在此方法中抛出的异常会导致插件加载失败
3. 此方法中 `ctx` 参数保证不为 `None`
---
### 2. `on_stop(ctx)` - 插件停止钩子
**签名**:
```python
async def on_stop(self, ctx: Any | None = None) -> None
```
**参数**:
- `ctx`: 运行时上下文
**触发时机**: 插件卸载或程序关闭前调用
**用途**:
- 关闭数据库连接
- 清理临时文件
- 注销 LLM 工具
- 保存状态数据
**示例**:
```python
class MyPlugin(Star):
async def on_stop(self, ctx) -> None:
# 保存状态
await self.put_kv_data("last_shutdown", time.time())
# 注销工具
if hasattr(self, '_tool_name'):
await ctx.unregister_llm_tool(self._tool_name)
# 确保 terminate 被调用
await super().on_stop(ctx)
ctx.logger.info(f"{ctx.plugin_id} 已停止")
```
**注意事项**:
1. 始终调用 `await super().on_stop(ctx)` 确保 `terminate()` 被调用
2. 此方法中的异常会被捕获并记录,不会阻止插件关闭
3. 此时可能没有活跃的事件处理,避免发送消息
---
### 3. `initialize()` - 初始化钩子
**签名**:
```python
async def initialize(self) -> None
```
**触发时机**: `on_start()` 内部自动调用
**用途**:
- 插件级别的初始化逻辑
- 不依赖 Context 的初始化
**示例**:
```python
class MyPlugin(Star):
async def initialize(self) -> None:
"""初始化插件"""
self._cache = {}
self._counter = 0
self.state = "ready"
```
**与 `on_start` 的区别**:
- `initialize()` 无 `Context` 参数,用于不依赖外部资源的初始化
- `on_start(ctx)` 有 `Context` 参数,用于需要访问 Core 的初始化
**调用顺序**:
```
插件实例化
↓
initialize() ← 先调用(无 Context)
↓
on_start(ctx) ← 后调用(有 Context)
```
---
### 4. `terminate()` - 终止钩子
**签名**:
```python
async def terminate(self) -> None
```
**触发时机**: `on_stop()` 内部自动调用
**用途**:
- 插件级别的清理逻辑
- 不依赖 Context 的清理
**示例**:
```python
class MyPlugin(Star):
async def terminate(self) -> None:
"""清理插件资源"""
self._cache.clear()
self.state = "stopped"
```
**与 `on_stop` 的区别**:
- `terminate()` 无 `Context` 参数,用于清理插件内部资源
- `on_stop(ctx)` 有 `Context` 参数,用于清理需要与 Core 交互的资源
**调用顺序**:
```
on_stop(ctx) ← 先调用(有 Context)
↓
terminate() ← 后调用(无 Context)
↓
插件卸载
```
---
### 5. `on_error(error, event, ctx)` - 错误处理钩子
**签名**:
```python
async def on_error(self, error: Exception, event, ctx) -> None
# 类方法
@classmethod
def __astrbot_is_new_star__(cls) -> bool
```
**参数**:
- `error`: 捕获的异常
- `event`: 事件对象(可能是 `MessageEvent` 或其他类型)
- `ctx`: 上下文对象
**触发时机**: 任何 Handler 执行抛出异常时
**默认行为**:
- `AstrBotError`:根据错误类型发送友好提示
- 其他异常:发送通用错误消息
- 记录错误日志
**示例**:
```python
from astrbot_sdk.errors import AstrBotError
class MyPlugin(Star):
async def on_error(self, error: Exception, event, ctx) -> None:
"""自定义错误处理"""
# SDK 标准错误
if isinstance(error, AstrBotError):
lines = []
if error.retryable:
lines.append("请求失败,请稍后重试")
elif error.hint:
lines.append(error.hint)
else:
lines.append(error.message)
if error.docs_url:
lines.append(f"文档:{error.docs_url}")
await event.reply("\n".join(lines))
# 业务逻辑错误
elif isinstance(error, ValueError):
await event.reply(f"参数错误:{error}")
# 网络错误
elif isinstance(error, ConnectionError):
await event.reply("网络连接失败,请检查网络设置")
# 未知错误
else:
await event.reply(f"出错了:{type(error).__name__}")
# 记录详细错误
ctx.logger.error(f"Handler failed: {error}", exc_info=error)
```
**覆盖建议**:
1. 始终记录错误日志
2. 向用户提供友好的错误提示
3. 调用 `await super().on_error(...)` 作为后备
---
## 类方法
### `__astrbot_is_new_star__()`
标识类为 v4 原生插件。
**签名**:
```python
@classmethod
def __astrbot_is_new_star__(cls) -> bool
```
**返回**: `bool` - 始终返回 `True`
**说明**: 此方法用于运行时识别插件类型,v4 原生插件返回 `True`,旧版插件无此方法。
---
## 便捷方法
### `text_to_image()`
将文本渲染为图片。
**签名**:
```python
async def text_to_image(
self,
text: str,
*,
return_url: bool = True
) -> str
```
**参数**:
- `text`: 要渲染的文本
- `return_url`: 是否返回 URL(False 则返回本地路径)
**返回**: 图片 URL 或路径
**示例**:
```python
class MyPlugin(Star):
@on_command("text_img")
async def text_to_image_cmd(self, event: MessageEvent):
url = await self.text_to_image("Hello World")
await event.reply_image(url)
```
**等价于**:
```python
url = await ctx.text_to_image("Hello World")
```
---
### `html_render()`
渲染 HTML 模板。
**签名**:
```python
async def html_render(
self,
tmpl: str,
data: dict,
*,
return_url: bool = True,
options: dict[str, Any] | None = None
) -> str
```
**参数**:
- `tmpl`: HTML 模板内容
- `data`: 模板数据
- `return_url`: 是否返回 URL
- `options`: 渲染选项
**返回**: 渲染结果 URL 或路径
**示例**:
```python
class MyPlugin(Star):
@on_command("card")
async def card_cmd(self, event: MessageEvent):
url = await self.html_render(
tmpl="<h1>{{ title }}</h1><p>{{ content }}</p>",
data={"title": "标题", "content": "内容"}
)
await event.reply_image(url)
```
**等价于**:
```python
url = await ctx.html_render(tmpl, data)
```
---
## KV 存储方法
这些方法继承自 `PluginKVStoreMixin`,提供简单的键值存储能力。
### `put_kv_data()`
存储数据。
**签名**:
```python
async def put_kv_data(self, key: str, value: Any) -> None
```
**示例**:
```python
await self.put_kv_data("last_run", time.time())
```
### `get_kv_data()`
获取数据。
**签名**:
```python
async def get_kv_data(self, key: str, default: _VT) -> _VT
```
**示例**:
```python
last_run = await self.get_kv_data("last_run", 0)
```
### `delete_kv_data()`
删除数据。
**签名**:
```python
async def delete_kv_data(self, key: str) -> None
```
**示例**:
```python
await self.delete_kv_data("temp_data")
```
---
## 完整插件示例
```python
"""
完整的插件示例
"""
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command, on_message, provide_capability
from astrbot_sdk.errors import AstrBotError
import asyncio
import time
class CompletePlugin(Star):
"""完整功能插件"""
async def initialize(self) -> None:
"""初始化"""
self._stats = {
"start_time": time.time(),
"command_count": 0
}
async def on_start(self, ctx) -> None:
"""启动"""
await super().on_start(ctx)
# 加载配置
config = await ctx.metadata.get_plugin_config()
self.greeting = config.get("greeting", "你好")
# 注册 LLM 工具
await ctx.register_llm_tool(
name="get_time",
parameters_schema={
"type": "object",
"properties": {},
"required": []
},
desc="获取当前时间",
func_obj=self.get_time_tool
)
# 启动后台任务
await ctx.register_task(
self.background_sync(),
desc="后台数据同步"
)
ctx.logger.info("Plugin started")
async def on_stop(self, ctx) -> None:
"""停止"""
# 保存统计
await self.put_kv_data("stats", self._stats)
await super().on_stop(ctx)
ctx.logger.info("Plugin stopped")
@on_command("hello", aliases=["hi", "greet"])
async def hello(self, event: MessageEvent, ctx: Context) -> None:
"""打招呼命令"""
self._stats["command_count"] += 1
await event.reply(f"{self.greeting},{event.sender_name}!")
@on_command("stats")
async def stats(self, event: MessageEvent, ctx: Context) -> None:
"""统计信息"""
uptime = time.time() - self._stats["start_time"]
await event.reply(f"""
运行时间: {uptime:.0f}秒
命令次数: {self._stats['command_count']}
""")
@on_message(keywords=["帮助"])
async def help(self, event: MessageEvent, ctx: Context) -> None:
"""帮助信息"""
await event.reply("""
可用命令:
/hello - 打招呼
/stats - 统计信息
/time - 当前时间
""")
@on_command("time")
async def time_cmd(self, event: MessageEvent, ctx: Context) -> None:
"""获取时间"""
result = await self.get_time_tool()
await event.reply(result)
async def get_time_tool(self) -> str:
"""LLM 工具实现"""
return f"当前时间: {time.strftime('%Y-%m-%d %H:%M:%S')}"
async def background_sync(self):
"""后台任务"""
while True:
await asyncio.sleep(3600)
# 执行同步逻辑
pass
async def on_error(self, error: Exception, event, ctx) -> None:
"""错误处理"""
if isinstance(error, AstrBotError):
await event.reply(error.hint or error.message)
else:
await event.reply(f"发生错误: {type(error).__name__}")
ctx.logger.error(f"Error: {error}", exc_info=error)
```
---
## plugin.yaml 配置
```yaml
_schema_version: 2
name: my_plugin
author: Your Name <email@example.com>
version: 1.0.0
desc: 我的插件描述
repo: https://github.com/user/repo
logo: assets/logo.png
runtime:
python: "3.12"
components:
- class: main:MyPlugin
support_platforms:
- aiocqhttp
- telegram
- discord
astrbot_version: ">=4.13.0,<5.0.0"
config:
timeout: 30
max_retries: 3
api_key: ""
```
---
## 最佳实践
### 1. 资源初始化与清理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 创建资源
self._session = aiohttp.ClientSession()
self._task = asyncio.create_task(self.background_task())
async def on_stop(self, ctx):
# 清理资源
if hasattr(self, '_task'):
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
if hasattr(self, '_session'):
await self._session.close()
```
### 2. 配置管理
```python
class MyPlugin(Star):
async def on_start(self, ctx):
config = await ctx.metadata.get_plugin_config()
# 提供默认值
self.timeout = config.get("timeout", 30)
# 验证必需配置
if "api_key" not in config:
raise ValueError("缺少必需配置: api_key")
self.api_key = config["api_key"]
```
### 3. 状态持久化
```python
class MyPlugin(Star):
async def on_start(self, ctx):
# 加载状态
self.last_update = await self.get_kv_data("last_update", 0)
self.user_data = await self.get_kv_data("users", {})
async def on_stop(self, ctx):
# 保存状态
await self.put_kv_data("last_update", time.time())
await self.put_kv_data("users", self.user_data)
```
### 4. 错误处理
```python
class MyPlugin(Star):
async def on_error(self, error, event, ctx):
# 根据错误类型发送不同的提示
if isinstance(error, ValueError):
await event.reply("参数错误")
elif isinstance(error, ConnectionError):
await event.reply("网络连接失败")
else:
# 使用默认处理
await super().on_error(error, event, ctx)
# 记录日志
ctx.logger.error(f"Handler error: {error}", exc_info=error)
```
---
## 注意事项
1. **异步方法**: 所有生命周期钩子都是异步方法,必须使用 `async def` 声明
2. **super() 调用**: 在 `on_start` 和 `on_stop` 中始终调用 `await super().xxx(ctx)` 确保 `initialize`/`terminate` 被调用
3. **context 属性**: 仅在生命周期钩子和 Handler 执行期间可用,不要存储此引用
4. **异常处理**: `on_start` 中的异常会导致插件加载失败,`on_stop` 中的异常会被捕获并记录
5. **资源清理**: 确保在 `on_stop` 或 `terminate` 中清理所有资源(连接、任务、文件等)
---
## 相关模块
- **装饰器**: `astrbot_sdk.decorators` - 事件处理装饰器
- **上下文**: `astrbot_sdk.context.Context` - 运行时上下文
- **事件**: `astrbot_sdk.events.MessageEvent` - 消息事件
- **错误**: `astrbot_sdk.errors.AstrBotError` - SDK 错误类
---
**版本**: v4.0
**模块**: `astrbot_sdk.star.Star`
**最后更新**: 2026-03-17
@@ -1,497 +0,0 @@
# 类型定义 API 完整参考
## 概述
本文档介绍 AstrBot SDK 中常用的类型定义,包括类型别名、泛型变量和类型注解。
**模块路径**: 分布在各个 SDK 模块中
---
## 目录
- [类型别名](#类型别名)
- [泛型变量](#泛型变量)
- [特殊类型](#特殊类型)
- [使用示例](#使用示例)
---
## 导入方式
```python
# 类型别名
from astrbot_sdk.context import PlatformCompatContent
from astrbot_sdk.clients.llm import ChatMessage, ChatHistoryItem, LLMResponse
# 泛型变量(通常不需要直接导入)
from astrbot_sdk.session_waiter import _P, _ResultT, _OwnerT
from astrbot_sdk.plugin_kv import _VT
# 通用类型
from typing import Callable, Awaitable, Any, Sequence, Mapping
HandlerType = Callable[..., Awaitable[Any]]
FilterType = Callable[..., Awaitable[bool]]
```
---
## 类型别名
### PlatformCompatContent
平台兼容的内容类型,用于表示可以发送到平台的各种消息格式。
**定义位置**: `astrbot_sdk.context`
**定义**:
```python
from collections.abc import Sequence
from typing import Any
PlatformCompatContent = (
str | MessageChain | Sequence[BaseMessageComponent] | Sequence[dict[str, Any]]
)
```
**说明**:
此类型别名表示可以用于平台发送方法的内容类型,支持以下四种格式:
| 格式 | 说明 | 示例 |
|------|------|------|
| `str` | 纯文本字符串 | `"Hello World"` |
| `MessageChain` | 消息链对象 | `MessageChain([Plain("Hi")])` |
| `Sequence[BaseMessageComponent]` | 消息组件列表 | `[Plain("Hi"), At("123")]` |
| `Sequence[dict[str, Any]]` | 序列化后的字典列表 | `[{"type": "text", "data": {"text": "Hi"}}]` |
**使用位置**:
- `Context.send_message()`
- `Context.send_message_by_id()`
- `PlatformClient.send_by_session()`
- `StarTools.send_message()`
**示例**:
```python
from astrbot_sdk import Plain, Image, MessageChain
# 纯文本
await ctx.platform.send_by_session("session_id", "Hello")
# 消息链
chain = MessageChain([Plain("Hello"), Image.fromURL("...")])
await ctx.platform.send_by_session("session_id", chain)
# 组件列表
await ctx.platform.send_by_session("session_id", [
Plain("Hello"),
At("123456")
])
# 字典列表
await ctx.platform.send_by_session("session_id", [
{"type": "text", "data": {"text": "Hello"}}
])
```
---
### ChatHistoryItem
聊天历史项类型,用于构建对话历史。
**定义位置**: `astrbot_sdk.clients.llm`
**定义**:
```python
from collections.abc import Mapping
from typing import Any
from pydantic import BaseModel
class ChatMessage(BaseModel):
role: str
content: str
ChatHistoryItem = ChatMessage | Mapping[str, Any]
```
**说明**:
此类型别名表示对话历史中的一项,可以是 `ChatMessage` 对象或任何字典类型的映射。
**支持格式**:
| 格式 | 说明 | 示例 |
|------|------|------|
| `ChatMessage` | Pydantic 模型对象 | `ChatMessage(role="user", content="Hi")` |
| `Mapping[str, Any]` | 字典类型 | `{"role": "user", "content": "Hi"}` |
**使用位置**:
- `LLMClient.chat()` - `history` 参数
- `LLMClient.chat_raw()` - `history` 参数
- `LLMClient.stream_chat()` - `history` 参数
**示例**:
```python
from astrbot_sdk.clients.llm import ChatMessage
# 使用 ChatMessage 对象
history = [
ChatMessage(role="user", content="你好"),
ChatMessage(role="assistant", content="你好!"),
]
# 使用字典
history = [
{"role": "user", "content": "你好"},
{"role": "assistant", "content": "你好!"},
]
# 混合使用
history = [
ChatMessage(role="user", content="你好"),
{"role": "assistant", "content": "你好!"},
{"role": "user", "content":今天天气怎么样?"},
]
```
---
## 泛型变量
SDK 内部使用的泛型类型变量,用于类型注解。
### `_P` - 参数规范
**定义位置**: `astrbot_sdk.session_waiter`
**定义**:
```python
from typing import ParamSpec
_P = ParamSpec("_P")
```
**说明**:
用于捕获可调用对象的参数签名,主要在装饰器中使用。
---
### `_ResultT` - 结果类型
**定义位置**: `astrbot_sdk.session_waiter`
**定义**:
```python
from typing import TypeVar
_ResultT = TypeVar("_ResultT")
```
**说明**:
表示异步函数的返回结果类型。
---
### `_OwnerT` - 所有者类型
**定义位置**: `astrbot_sdk.session_waiter`
**定义**:
```python
_OwnerT = TypeVar("_OwnerT")
```
**说明**:
表示类的所有者类型(通常是 `Star` 子类)。
---
### `_VT` - 值类型
**定义位置**: `astrbot_sdk.plugin_kv`
**定义**:
```python
_VT = TypeVar("_VT")
```
**说明**:
用于 KV 存储中默认值的类型。
**使用位置**:
- `PluginKVStoreMixin.get_kv_data()` - `default` 参数的类型注解
**示例**:
```python
# default 参数的类型会根据传入的值自动推断
value = await self.get_kv_data("key", default="default") # _VT 推断为 str
count = await self.get_kv_data("count", default=0) # _VT 推断为 int
```
---
## 特殊类型
### HandlerType
事件处理器函数类型。
**定义**:
```python
from typing import Callable, Awaitable, Any
HandlerType = Callable[..., Awaitable[Any]]
```
**说明**:
表示事件处理器的函数签名,接受任意参数并返回异步结果。
**特征**:
- 可变参数 (`...`)
- 异步返回 (`Awaitable[Any]`)
**示例**:
```python
async def my_handler(event: MessageEvent, ctx: Context) -> None:
pass
# 符合 HandlerType 类型
```
---
### FilterType
过滤器函数类型。
**定义**:
```python
FilterType = Callable[..., Awaitable[bool]]
```
**说明**:
表示过滤器函数的类型,返回布尔值。
**特征**:
- 可变参数 (`...`)
- 异步返回布尔值 (`Awaitable[bool]`)
**示例**:
```python
async def my_filter(event: MessageEvent, ctx: Context) -> bool:
return event.platform == "qq"
# 符合 FilterType 类型
```
---
## Pydantic 模型类型
### ChatMessage
聊天消息模型,用于构建对话历史。
**定义位置**: `astrbot_sdk.clients.llm`
**定义**:
```python
from pydantic import BaseModel
class ChatMessage(BaseModel):
"""聊天消息模型。"""
role: str
content: str
```
**属性**:
| 属性 | 类型 | 说明 |
|------|------|------|
| `role` | `str` | 消息角色,如 `"user"`, `"assistant"`, `"system"` |
| `content` | `str` | 消息内容 |
**示例**:
```python
from astrbot_sdk.clients.llm import ChatMessage
# 系统提示
system_msg = ChatMessage(
role="system",
content="你是一个友好的助手"
)
# 用户消息
user_msg = ChatMessage(
role="user",
content="你好"
)
# 助手回复
assistant_msg = ChatMessage(
role="assistant",
content="你好!有什么可以帮助你的?"
)
```
---
### LLMResponse
LLM 响应模型,包含完整的响应信息。
**定义位置**: `astrbot_sdk.clients.llm`
**定义**:
```python
from pydantic import BaseModel, Field
class LLMResponse(BaseModel):
"""LLM 响应模型。"""
text: str
usage: dict[str, Any] | None = None
finish_reason: str | None = None
tool_calls: list[dict[str, Any]] = Field(default_factory=list)
role: str | None = None
reasoning_content: str | None = None
reasoning_signature: str | None = None
```
**属性**:
| 属性 | 类型 | 说明 |
|------|------|------|
| `text` | `str` | 生成的文本内容 |
| `usage` | `dict[str, Any] \| None` | Token 使用统计 |
| `finish_reason` | `str \| None` | 结束原因(`"stop"`, `"length"`, `"tool_calls"`) |
| `tool_calls` | `list[dict[str, Any]]` | 工具调用列表 |
| `role` | `str \| None` | 响应角色 |
| `reasoning_content` | `str \| None` | 推理内容(用于推理模型) |
| `reasoning_signature` | `str \| None` | 推理签名 |
**示例**:
```python
from astrbot_sdk.clients.llm import LLMResponse
response = await ctx.llm.chat_raw("写一首诗")
print(f"生成内容: {response.text}")
print(f"Token 使用: {response.usage}")
print(f"结束原因: {response.finish_reason}")
if response.usage:
print(f"提示词 Token: {response.usage.get('prompt_tokens')}")
print(f"完成 Token: {response.usage.get('completion_tokens')}")
```
---
## 使用示例
### 类型注解在函数签名中的使用
```python
from typing import Sequence, Mapping, Any
from astrbot_sdk.clients.llm import ChatMessage, ChatHistoryItem
from astrbot_sdk import MessageChain, BaseMessageComponent, PlatformCompatContent
# 使用 ChatHistoryItem
async def chat_with_history(
prompt: str,
history: Sequence[ChatHistoryItem] | None = None
) -> str:
"""与 LLM 聊天的函数。"""
pass
# 使用 PlatformCompatContent
async def send_content(
session: str,
content: PlatformCompatContent
) -> dict[str, Any]:
"""发送内容的函数。"""
pass
```
### 类型检查和类型守卫
```python
from collections.abc import Mapping, Sequence
from astrbot_sdk.clients.llm import ChatMessage, ChatHistoryItem
def normalize_history_item(item: ChatHistoryItem) -> dict[str, Any]:
"""将聊天历史项规范化为字典。"""
if isinstance(item, ChatMessage):
return item.model_dump()
if isinstance(item, Mapping):
return dict(item)
raise TypeError("无效的聊天历史项类型")
# 使用
history: Sequence[ChatHistoryItem] = [
ChatMessage(role="user", content="Hi"),
{"role": "assistant", "content": "Hello"},
]
normalized = [normalize_history_item(item) for item in history]
```
### 泛型函数
```python
from typing import TypeVar, Generic
T = TypeVar("T")
class Container(Generic[T]):
def __init__(self, value: T) -> None:
self.value = value
def get(self) -> T:
return self.value
# 使用
int_container: Container[int] = Container(42)
str_container: Container[str] = Container("hello")
```
---
## 相关模块
- **LLM 客户端**: `astrbot_sdk.clients.LLMClient`
- **消息组件**: `astrbot_sdk.message_components`
- **消息链**: `astrbot_sdk.message_result.MessageChain`
- **上下文**: `astrbot_sdk.context.Context`
---
**版本**: v4.0
**最后更新**: 2026-03-17
File diff suppressed because it is too large Load Diff
-311
View File
@@ -1,311 +0,0 @@
"""跨运行时边界传递的统一错误模型。
AstrBotError 是 SDK 中所有可预期错误的标准格式,
支持跨进程传递(通过 to_payload/from_payload 序列化)。
错误处理流程:
1. 运行时抛出 AstrBotError 子类或实例
2. 错误被捕获并序列化为 payload
3. 跨进程传输后反序列化
4. 在 on_error 钩子中统一处理
Example:
# 抛出错误
raise AstrBotError.invalid_input("参数不能为空")
# 捕获并处理
try:
await some_operation()
except AstrBotError as e:
if e.retryable:
# 可重试的错误
await retry()
else:
# 不可重试的错误
await event.reply(e.hint or e.message)
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
class ErrorCodes:
"""AstrBot v4 的稳定错误码常量。
这些错误码在协议层稳定,不应随意更改。
新增错误码应放在对应分类的末尾。
分类:
- 不可重试错误(retryable=False):配置错误、权限错误等
- 可重试错误(retryable=True):网络超时、临时故障等
"""
UNKNOWN_ERROR = "unknown_error"
# 不可重试错误 - 配置或使用问题
LLM_NOT_CONFIGURED = "llm_not_configured"
CAPABILITY_NOT_FOUND = "capability_not_found"
PERMISSION_DENIED = "permission_denied"
LLM_ERROR = "llm_error"
INVALID_INPUT = "invalid_input"
CANCELLED = "cancelled"
PROTOCOL_VERSION_MISMATCH = "protocol_version_mismatch"
PROTOCOL_ERROR = "protocol_error"
INTERNAL_ERROR = "internal_error"
RATE_LIMITED = "rate_limited"
COOLDOWN_ACTIVE = "cooldown_active"
# 可重试错误 - 临时故障
CAPABILITY_TIMEOUT = "capability_timeout"
NETWORK_ERROR = "network_error"
LLM_TEMPORARY_ERROR = "llm_temporary_error"
@dataclass(slots=True)
class AstrBotError(Exception):
"""AstrBot SDK 的标准错误类型。
所有可预期的错误都应使用此类或其工厂方法创建。
支持跨进程传递,包含用户友好的提示信息。
Attributes:
code: 错误码,来自 ErrorCodes 常量
message: 错误消息,面向开发者
hint: 用户提示,面向终端用户
retryable: 是否可重试
Example:
# 使用工厂方法创建错误
raise AstrBotError.invalid_input("参数格式错误", hint="请使用 JSON 格式")
# 检查错误类型
try:
await operation()
except AstrBotError as e:
if e.code == ErrorCodes.CAPABILITY_NOT_FOUND:
logger.error(f"能力不存在: {e.message}")
"""
code: str
message: str
hint: str = ""
retryable: bool = False
docs_url: str = ""
details: dict[str, Any] | None = None
def __str__(self) -> str:
return self.message
@classmethod
def cancelled(cls, message: str = "调用被取消") -> AstrBotError:
"""创建取消错误。
Args:
message: 错误消息
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.CANCELLED,
message=message,
hint="",
retryable=False,
)
@classmethod
def capability_not_found(cls, name: str) -> AstrBotError:
"""创建能力未找到错误。
Args:
name: 未找到的能力名称
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.CAPABILITY_NOT_FOUND,
message=f"未找到能力:{name}",
hint="请确认 AstrBot Core 是否已注册该 capability",
retryable=False,
)
@classmethod
def invalid_input(
cls,
message: str,
*,
hint: str = "请检查调用参数",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError:
"""创建输入无效错误。
Args:
message: 详细错误消息
hint: 用户提示
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.INVALID_INPUT,
message=message,
hint=hint,
retryable=False,
docs_url=docs_url,
details=details,
)
@classmethod
def protocol_version_mismatch(cls, message: str) -> AstrBotError:
"""创建协议版本不匹配错误。
Args:
message: 详细错误消息
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.PROTOCOL_VERSION_MISMATCH,
message=message,
hint="请升级 astrbot_sdk 至最新版本",
retryable=False,
)
@classmethod
def protocol_error(cls, message: str) -> AstrBotError:
"""创建协议错误。
Args:
message: 详细错误消息
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.PROTOCOL_ERROR,
message=message,
hint="请检查通信双方的协议实现",
retryable=False,
)
@classmethod
def internal_error(
cls,
message: str,
*,
hint: str = "请联系插件作者",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError:
"""创建内部错误。
Args:
message: 详细错误消息
hint: 用户提示
Returns:
AstrBotError 实例
"""
return cls(
code=ErrorCodes.INTERNAL_ERROR,
message=message,
hint=hint,
retryable=False,
docs_url=docs_url,
details=details,
)
@classmethod
def network_error(
cls,
message: str,
*,
hint: str = "网络请求失败,请稍后重试",
docs_url: str = "",
details: dict[str, Any] | None = None,
) -> AstrBotError:
return cls(
code=ErrorCodes.NETWORK_ERROR,
message=message,
hint=hint,
retryable=True,
docs_url=docs_url,
details=details,
)
@classmethod
def rate_limited(
cls,
*,
hint: str = "操作过于频繁,请稍后再试。",
details: dict[str, Any] | None = None,
) -> AstrBotError:
return cls(
code=ErrorCodes.RATE_LIMITED,
message="handler invocation is rate limited",
hint=hint,
retryable=False,
details=details,
)
@classmethod
def cooldown_active(
cls,
*,
hint: str,
details: dict[str, Any] | None = None,
) -> AstrBotError:
return cls(
code=ErrorCodes.COOLDOWN_ACTIVE,
message="handler cooldown is active",
hint=hint,
retryable=False,
details=details,
)
def to_payload(self) -> dict[str, object]:
"""序列化为可传输的字典格式。
用于跨进程传递错误信息。
Returns:
包含错误信息的字典
"""
return {
"code": self.code,
"message": self.message,
"hint": self.hint,
"retryable": self.retryable,
"docs_url": self.docs_url,
"details": dict(self.details) if isinstance(self.details, dict) else None,
}
@classmethod
def from_payload(cls, payload: dict[str, object]) -> AstrBotError:
"""从字典反序列化错误实例。
Args:
payload: 包含错误信息的字典
Returns:
AstrBotError 实例
"""
details_payload = payload.get("details")
details = (
{str(key): value for key, value in details_payload.items()}
if isinstance(details_payload, dict)
else None
)
return cls(
code=str(payload.get("code", ErrorCodes.UNKNOWN_ERROR)),
message=str(payload.get("message", "未知错误")),
hint=str(payload.get("hint", "")),
retryable=bool(payload.get("retryable", False)),
docs_url=str(payload.get("docs_url", "")),
details=details,
)
-747
View File
@@ -1,747 +0,0 @@
"""v4 原生事件对象。
顶层 ``MessageEvent`` 保持精简,只承载 v4 运行时真正需要的基础能力。
迁移期扩展事件能力放在独立模块中,而不是继续塞回顶层事件类型。
MessageEvent 是 handler 接收的主要事件类型,封装了:
- 消息文本内容
- 发送者信息(user_id, group_id)
- 平台标识
- 回复能力(reply, reply_image, reply_chain)
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from .message.components import (
At,
BaseMessageComponent,
File,
Image,
Plain,
component_to_payload_sync,
payloads_to_components,
)
from .message.result import EventResultType, MessageChain, MessageEventResult
from .protocol.descriptors import SessionRef
if TYPE_CHECKING:
from .context import Context
@dataclass(slots=True)
class PlainTextResult:
"""纯文本结果。
用于 handler 返回简单的文本结果。
"""
text: str
ReplyHandler = Callable[[str], Awaitable[None]]
_JSON_DROP = object()
def _coerce_str(value: Any) -> str:
if value is None:
return ""
if isinstance(value, str):
return value
return str(value)
def _coerce_optional_str(value: Any) -> str | None:
if value is None:
return None
text = value if isinstance(value, str) else str(value)
return text or None
def _json_safe_value(value: Any) -> Any:
if value is None or isinstance(value, (str, int, float, bool)):
return value
if isinstance(value, (list, tuple)):
items = []
for item in value:
normalized = _json_safe_value(item)
if normalized is not _JSON_DROP:
items.append(normalized)
return items
if isinstance(value, dict):
normalized_dict: dict[str, Any] = {}
for key, item in value.items():
normalized = _json_safe_value(item)
if normalized is not _JSON_DROP:
normalized_dict[str(key)] = normalized
return normalized_dict
model_dump = getattr(value, "model_dump", None)
if callable(model_dump):
try:
return _json_safe_value(model_dump())
except Exception:
return _JSON_DROP
try:
json.dumps(value)
except (TypeError, ValueError):
return _JSON_DROP
return value
def _json_safe_mapping(value: Any) -> dict[str, Any]:
if not isinstance(value, dict):
return {}
normalized: dict[str, Any] = {}
for key, item in value.items():
safe_item = _json_safe_value(item)
if safe_item is not _JSON_DROP:
normalized[str(key)] = safe_item
return normalized
class MessageEvent:
"""消息事件对象。
封装收到的消息,提供便捷的回复方法。
每个 handler 调用都会创建新的 MessageEvent 实例。
Attributes:
text: 消息文本内容
user_id: 发送者用户 ID,缺失时为空字符串
group_id: 群组 ID(私聊时为 None)
platform: 平台标识(如 "qq", "wechat"),缺失时为空字符串
session_id: 会话 ID(通常是 group_id 或 user_id,缺失时为空字符串)
raw: 原始消息数据
Example:
@on_command("echo")
async def echo(self, event: MessageEvent, ctx: Context):
await event.reply(f"你说: {event.text}")
"""
text: str
user_id: str
group_id: str | None
platform: str
session_id: str
self_id: str
platform_id: str
message_type: str
sender_name: str
def __init__(
self,
*,
text: str = "",
user_id: str | None = None,
group_id: str | None = None,
platform: str | None = None,
session_id: str | None = None,
self_id: str | None = None,
platform_id: str | None = None,
message_type: str | None = None,
sender_name: str | None = None,
is_admin: bool = False,
raw: dict[str, Any] | None = None,
context: Context | None = None,
reply_handler: ReplyHandler | None = None,
) -> None:
"""初始化消息事件。
Args:
text: 消息文本
user_id: 用户 ID
group_id: 群组 ID
platform: 平台标识
session_id: 会话 ID,None 时自动从 group_id/user_id 推断
raw: 原始消息数据
context: 运行时上下文
reply_handler: 自定义回复处理器
"""
normalized_user_id = _coerce_str(user_id)
normalized_group_id = _coerce_optional_str(group_id)
normalized_platform = _coerce_str(platform)
normalized_session_id = _coerce_str(session_id)
self.text = text
self.user_id = normalized_user_id
self.group_id = normalized_group_id
self.platform = normalized_platform
self.session_id = (
normalized_session_id or normalized_group_id or normalized_user_id or ""
)
self.self_id = _coerce_str(self_id)
self.platform_id = _coerce_str(platform_id) or normalized_platform
self.message_type = _coerce_str(message_type).lower()
self.sender_name = _coerce_str(sender_name)
self._is_admin = bool(is_admin)
self.raw = raw or {}
self._stopped = False
host_extras = self.raw.get("host_extras")
raw_extras = self.raw.get("extras")
self._host_extras = _json_safe_mapping(
host_extras if isinstance(host_extras, dict) else raw_extras
)
self._host_extras_present = "host_extras" in self.raw or "extras" in self.raw
sdk_local_extras = self.raw.get("sdk_local_extras")
self._sdk_local_extras = _json_safe_mapping(sdk_local_extras)
self._sdk_local_extras_present = "sdk_local_extras" in self.raw
self._sdk_local_extras_dirty = False
messages_payload = self.raw.get("messages")
self._messages = (
payloads_to_components(messages_payload)
if isinstance(messages_payload, list)
else []
)
self._messages_present = "messages" in self.raw
self._message_outline = str(self.raw.get("message_outline", self.text))
sent_messages_payload = self.raw.get("sent_messages")
self._sent_messages = (
payloads_to_components(sent_messages_payload)
if isinstance(sent_messages_payload, list)
else []
)
self._sent_messages_present = "sent_messages" in self.raw
self._sent_message_outline = str(self.raw.get("sent_message_outline", ""))
self._sent_message_outline_present = "sent_message_outline" in self.raw
self._context = context
self._reply_handler = reply_handler
if self._reply_handler is None and context is not None:
self._reply_handler = lambda text: context.platform.send(
self.session_ref or self.session_id,
text,
)
def _require_runtime_context(self, action: str) -> Context:
"""获取运行时上下文,不存在则抛出异常。"""
if self._context is None:
raise RuntimeError(f"MessageEvent 未绑定运行时上下文,无法 {action}")
return self._context
def _reply_target(self) -> SessionRef | str:
"""获取回复目标。"""
return self.session_ref or self.session_id
@classmethod
def from_payload(
cls,
payload: dict[str, Any],
*,
context: Context | None = None,
reply_handler: ReplyHandler | None = None,
) -> MessageEvent:
"""从协议载荷创建事件实例。
Args:
payload: 协议层传递的消息数据
context: 运行时上下文
reply_handler: 自定义回复处理器
Returns:
新的 MessageEvent 实例
"""
target_payload = payload.get("target")
session_id = payload.get("session_id")
platform = payload.get("platform")
if isinstance(target_payload, dict):
target = SessionRef.model_validate(target_payload)
session_id = session_id or target.session
platform = platform or target.platform
return cls(
text=str(payload.get("text", "")),
user_id=payload.get("user_id"),
group_id=payload.get("group_id"),
platform=platform,
session_id=session_id,
self_id=payload.get("self_id"),
platform_id=payload.get("platform_id"),
message_type=payload.get("message_type"),
sender_name=payload.get("sender_name"),
is_admin=bool(payload.get("is_admin", False)),
raw=payload,
context=context,
reply_handler=reply_handler,
)
def to_payload(self) -> dict[str, Any]:
"""转换为协议载荷格式。
Returns:
可序列化的字典
"""
payload = dict(self.raw)
payload.update(
{
"text": self.text,
"user_id": self.user_id,
"group_id": self.group_id,
"platform": self.platform,
"session_id": self.session_id,
"self_id": self.self_id,
"platform_id": self.platform_id,
"message_type": self.message_type,
"sender_name": self.sender_name,
"is_admin": self._is_admin,
}
)
if self.session_ref is not None:
payload["target"] = self.session_ref.to_payload()
merged_extras = dict(self._host_extras)
merged_extras.update(self._sdk_local_extras_payload())
if merged_extras:
payload["extras"] = merged_extras
elif self._host_extras_present:
payload["extras"] = {}
else:
payload.pop("extras", None)
if self._host_extras or self._host_extras_present:
payload["host_extras"] = dict(self._host_extras)
else:
payload.pop("host_extras", None)
sdk_local_extras = self._sdk_local_extras_payload()
if sdk_local_extras or self._should_serialize_sdk_local_extras():
payload["sdk_local_extras"] = sdk_local_extras
else:
payload.pop("sdk_local_extras", None)
if self._messages or self._messages_present:
payload["messages"] = [
component_to_payload_sync(component) for component in self._messages
]
else:
payload.pop("messages", None)
payload["message_outline"] = self._message_outline
if self._sent_messages or self._sent_messages_present:
payload["sent_messages"] = [
component_to_payload_sync(component)
for component in self._sent_messages
]
else:
payload.pop("sent_messages", None)
if self._sent_message_outline or self._sent_message_outline_present:
payload["sent_message_outline"] = self._sent_message_outline
else:
payload.pop("sent_message_outline", None)
return payload
@property
def session_ref(self) -> SessionRef | None:
"""获取会话引用对象。
Returns:
SessionRef 实例,如果没有有效的 session_id 则返回 None
"""
if not self.session_id:
return None
return SessionRef(
conversation_id=self.session_id,
platform=self.platform,
raw=self.raw or None,
)
@property
def target(self) -> SessionRef | None:
"""session_ref 的别名。"""
return self.session_ref
@property
def unified_msg_origin(self) -> str:
"""Unified message origin string."""
return self.session_id
def is_private_chat(self) -> bool:
"""Whether the current event belongs to a private chat."""
if self.message_type:
return self.message_type == "private"
return not bool(self.group_id)
def is_group_chat(self) -> bool:
if self.message_type:
return self.message_type == "group"
return bool(self.group_id)
def get_platform_id(self) -> str:
"""Get the platform instance identifier."""
return self.platform_id
def get_message_type(self) -> str:
"""Get the normalized message type."""
return self.message_type
def get_session_id(self) -> str:
"""Get the current session identifier."""
return self.session_id
def is_admin(self) -> bool:
"""Whether the sender has admin permission."""
return self._is_admin
def get_messages(self) -> list[BaseMessageComponent]:
"""Return SDK message components for the current event."""
return list(self._messages)
def get_sent_messages(self) -> list[BaseMessageComponent]:
"""Return outbound SDK message components for after-send events."""
return list(self._sent_messages)
def has_component(self, type_: type[BaseMessageComponent]) -> bool:
return any(isinstance(component, type_) for component in self._messages)
def get_components(
self,
type_: type[BaseMessageComponent],
) -> list[BaseMessageComponent]:
return [
component for component in self._messages if isinstance(component, type_)
]
def get_images(self) -> list[Image]:
return [
component for component in self._messages if isinstance(component, Image)
]
def get_files(self) -> list[File]:
return [
component for component in self._messages if isinstance(component, File)
]
def extract_plain_text(self) -> str:
return " ".join(
component.text
for component in self._messages
if isinstance(component, Plain)
)
def get_at_users(self) -> list[str]:
return [
str(component.qq)
for component in self._messages
if isinstance(component, At) and str(component.qq).lower() != "all"
]
def get_message_outline(self) -> str:
"""Return the normalized message outline."""
return self._message_outline
def get_sent_message_outline(self) -> str:
"""Return the outbound message outline for after-send events."""
return self._sent_message_outline
async def get_group(self) -> dict[str, Any] | None:
"""Get current-group metadata for the bound message request."""
context = self._require_runtime_context("get_group")
output = await context._proxy.call( # noqa: SLF001
"platform.get_group",
{
"session": self.session_id,
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
payload = output.get("group")
if not isinstance(payload, dict):
return None
return dict(payload)
def set_extra(self, key: str, value: Any) -> None:
"""Store SDK-local transient event data."""
self._sdk_local_extras[key] = value
self._sdk_local_extras_dirty = True
def get_extra(self, key: str | None = None, default: Any = None) -> Any:
"""Read SDK-local transient event data."""
extras = dict(self._host_extras)
extras.update(self._sdk_local_extras)
if key is None:
return extras
return extras.get(key, default)
def clear_extra(self) -> None:
"""Clear SDK-local transient event data."""
self._sdk_local_extras.clear()
self._sdk_local_extras_dirty = True
def _sdk_local_extras_payload(self) -> dict[str, Any]:
return _json_safe_mapping(self._sdk_local_extras)
def _should_serialize_sdk_local_extras(self) -> bool:
return (
self._sdk_local_extras_present
or self._sdk_local_extras_dirty
or bool(self._sdk_local_extras)
)
async def request_llm(self) -> bool:
"""Request the default LLM chain for the current message request."""
context = self._require_runtime_context("request_llm")
output = await context._proxy.call( # noqa: SLF001
"system.event.llm.request",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
return bool(output.get("should_call_llm", False))
async def should_call_llm(self) -> bool:
"""Read the current default-LLM decision from the host bridge."""
context = self._require_runtime_context("should_call_llm")
output = await context._proxy.call( # noqa: SLF001
"system.event.llm.get_state",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
return bool(output.get("should_call_llm", False))
async def set_result(self, result: MessageEventResult) -> MessageEventResult:
"""Store a request-scoped SDK result in the host bridge."""
context = self._require_runtime_context("set_result")
await context._proxy.call( # noqa: SLF001
"system.event.result.set",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
"result": result.to_payload(),
},
)
return result
async def get_result(self) -> MessageEventResult | None:
"""Read the current request-scoped SDK result from the host bridge."""
context = self._require_runtime_context("get_result")
output = await context._proxy.call( # noqa: SLF001
"system.event.result.get",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
payload = output.get("result")
if not isinstance(payload, dict):
return None
return MessageEventResult.from_payload(payload)
async def clear_result(self) -> None:
"""Clear the current request-scoped SDK result."""
context = self._require_runtime_context("clear_result")
await context._proxy.call( # noqa: SLF001
"system.event.result.clear",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
def stop_event(self) -> None:
"""Mark the SDK-local event as stopped."""
self._stopped = True
def continue_event(self) -> None:
"""Clear the SDK-local stop flag."""
self._stopped = False
def is_stopped(self) -> bool:
"""Return whether the SDK-local event is stopped."""
return self._stopped
async def reply(self, text: str) -> None:
"""回复文本消息。
Args:
text: 要回复的文本内容
Raises:
RuntimeError: 如果未绑定 reply handler
"""
if self._reply_handler is None:
raise RuntimeError("MessageEvent 未绑定 reply handler,无法 reply")
await self._reply_handler(text)
async def reply_image(self, image_url: str) -> None:
"""回复图片消息。
Args:
image_url: 图片 URL
Raises:
RuntimeError: 如果未绑定运行时上下文
"""
context = self._require_runtime_context("reply_image")
await context.platform.send_image(self._reply_target(), image_url)
async def reply_chain(
self,
chain: MessageChain | list[BaseMessageComponent] | list[dict[str, Any]],
) -> None:
"""回复消息链(多类型消息组合)。
Args:
chain: 消息链组件列表
Raises:
RuntimeError: 如果未绑定运行时上下文
"""
context = self._require_runtime_context("reply_chain")
await context.platform.send_chain(self._reply_target(), chain)
async def react(self, emoji: str) -> bool:
"""Send a platform reaction when supported."""
context = self._require_runtime_context("react")
output = await context._proxy.call( # noqa: SLF001
"system.event.react",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
"emoji": emoji,
},
)
return bool(output.get("supported", False))
async def send_typing(self) -> bool:
"""Emit typing state when the host platform supports it."""
context = self._require_runtime_context("send_typing")
output = await context._proxy.call( # noqa: SLF001
"system.event.send_typing",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
},
)
return bool(output.get("supported", False))
async def send_streaming(
self,
generator,
use_fallback: bool = False,
) -> bool:
"""Replay normalized chunks through the host streaming pathway."""
context = self._require_runtime_context("send_streaming")
output = await context._proxy.call( # noqa: SLF001
"system.event.send_streaming",
{
"target": (
self.session_ref.to_payload()
if self.session_ref is not None
else None
),
"use_fallback": use_fallback,
},
)
if not bool(output.get("supported", False)):
return False
stream_id = str(output.get("stream_id", ""))
if not stream_id:
return False
try:
async for item in generator:
if isinstance(item, str):
chain = MessageChain([Plain(item, convert=False)])
else:
chain = self._coerce_chain_or_raise(item)
await context._proxy.call( # noqa: SLF001
"system.event.send_streaming_chunk",
{
"stream_id": stream_id,
"chain": await chain.to_payload_async(),
},
)
finally:
output = await context._proxy.call( # noqa: SLF001
"system.event.send_streaming_close",
{"stream_id": stream_id},
)
return bool(output.get("supported", False))
def bind_reply_handler(self, reply_handler: ReplyHandler) -> None:
"""绑定自定义回复处理器。
Args:
reply_handler: 回复处理函数
"""
self._reply_handler = reply_handler
def plain_result(self, text: str) -> PlainTextResult:
"""创建纯文本结果。
Args:
text: 结果文本
Returns:
PlainTextResult 实例
"""
return PlainTextResult(text=text)
def make_result(self) -> MessageEventResult:
"""Create an empty SDK-local result wrapper."""
return MessageEventResult(type=EventResultType.EMPTY)
def image_result(self, url_or_path: str) -> MessageEventResult:
"""Create a chain result that contains one image component."""
if url_or_path.startswith(("http://", "https://")):
image = Image.fromURL(url_or_path)
elif url_or_path.startswith("base64://"):
image = Image.fromBase64(url_or_path.removeprefix("base64://"))
else:
image = Image.fromFileSystem(url_or_path)
return MessageEventResult(
type=EventResultType.CHAIN,
chain=MessageChain([image]),
)
def chain_result(
self,
chain: MessageChain | list[BaseMessageComponent],
) -> MessageEventResult:
"""Create a chain result from SDK components."""
normalized = (
chain if isinstance(chain, MessageChain) else MessageChain(list(chain))
)
return MessageEventResult(type=EventResultType.CHAIN, chain=normalized)
@staticmethod
def _coerce_chain_or_raise(item: Any) -> MessageChain:
if isinstance(item, MessageEventResult):
return item.chain
if isinstance(item, MessageChain):
return item
if isinstance(item, BaseMessageComponent):
return MessageChain([item])
if isinstance(item, list) and all(
isinstance(component, BaseMessageComponent) for component in item
):
return MessageChain(list(item))
raise TypeError(
"send_streaming only accepts str, MessageChain, MessageEventResult or SDK message components"
)
-218
View File
@@ -1,218 +0,0 @@
"""SDK-native filter declarations.
本模块提供事件过滤器的声明式 API,用于在 handler 执行前进行条件判断。
内置过滤器类型:
- PlatformFilter: 按平台名称过滤(如 qq、wechat)
- MessageTypeFilter: 按消息类型过滤(如 group、private)
- CustomFilter: 用户自定义的同步布尔函数
组合操作:
- all_of(*filters): 所有过滤器都通过才执行(AND 逻辑)
- any_of(*filters): 任一过滤器通过即可执行(OR 逻辑)
- 支持 & 和 | 运算符进行链式组合
过滤器在本地(SDK worker 进程内)求值,避免不必要的跨进程调用。
"""
from __future__ import annotations
import inspect
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Literal, TypeAlias
from .decorators import append_filter_meta
from .protocol.descriptors import (
CompositeFilterSpec,
FilterSpec,
LocalFilterRefSpec,
MessageTypeFilterSpec,
PlatformFilterSpec,
)
FilterOperator: TypeAlias = Literal["and", "or"]
@dataclass(slots=True)
class LocalFilterBinding:
filter_id: str
callable: Callable[..., bool]
args: dict[str, Any] = field(default_factory=dict)
def evaluate(self, *, event=None, ctx=None) -> bool:
signature = inspect.signature(self.callable)
kwargs: dict[str, Any] = {}
if "event" in signature.parameters:
kwargs["event"] = event
if "ctx" in signature.parameters:
kwargs["ctx"] = ctx
result = self.callable(**kwargs)
if inspect.isawaitable(result):
raise TypeError("CustomFilter must return a synchronous bool")
if not isinstance(result, bool):
raise TypeError("CustomFilter must return bool")
return result
class FilterBinding:
def __and__(self, other: FilterBinding) -> CompositeFilter:
return CompositeFilter("and", [self, other])
def __or__(self, other: FilterBinding) -> CompositeFilter:
return CompositeFilter("or", [self, other])
def compile(self) -> tuple[FilterSpec, list[LocalFilterBinding]]:
raise NotImplementedError
@dataclass(slots=True)
class PlatformFilter(FilterBinding):
platforms: list[str]
def compile(self) -> tuple[FilterSpec, list[LocalFilterBinding]]:
return PlatformFilterSpec(platforms=list(self.platforms)), []
@dataclass(slots=True)
class MessageTypeFilter(FilterBinding):
message_types: list[str]
def compile(self) -> tuple[FilterSpec, list[LocalFilterBinding]]:
return MessageTypeFilterSpec(message_types=list(self.message_types)), []
@dataclass(slots=True)
class CustomFilter(FilterBinding):
callable: Callable[..., bool]
filter_id: str | None = None
def __post_init__(self) -> None:
if self.filter_id is None:
self.filter_id = f"{self.callable.__module__}.{getattr(self.callable, '__qualname__', self.callable.__name__)}"
def compile(self) -> tuple[FilterSpec, list[LocalFilterBinding]]:
assert self.filter_id is not None
return LocalFilterRefSpec(filter_id=self.filter_id), [
LocalFilterBinding(filter_id=self.filter_id, callable=self.callable),
]
@dataclass(slots=True)
class CompositeFilter(FilterBinding):
operator: FilterOperator
children: list[FilterBinding]
def compile(self) -> tuple[FilterSpec, list[LocalFilterBinding]]:
compiled_children: list[FilterSpec] = []
local_bindings: list[LocalFilterBinding] = []
for child in self.children:
spec, locals_for_child = child.compile()
compiled_children.append(spec)
local_bindings.extend(locals_for_child)
if local_bindings:
filter_id = (
"composite:"
+ ":".join(binding.filter_id for binding in local_bindings)
+ f":{self.operator}"
)
def _evaluate(*, event=None, ctx=None) -> bool:
results = [
_evaluate_filter_spec_locally(
spec, local_bindings, event=event, ctx=ctx
)
for spec in compiled_children
]
if self.operator == "and":
return all(results)
return any(results)
return (
LocalFilterRefSpec(filter_id=filter_id),
[LocalFilterBinding(filter_id=filter_id, callable=_evaluate)],
)
return CompositeFilterSpec(kind=self.operator, children=compiled_children), []
def _evaluate_filter_spec_locally(
spec: FilterSpec,
local_bindings: list[LocalFilterBinding],
*,
event=None,
ctx=None,
) -> bool:
if isinstance(spec, PlatformFilterSpec):
if event is None:
return True
platform = getattr(event, "platform", "") or ""
return platform in spec.platforms
if isinstance(spec, MessageTypeFilterSpec):
if event is None:
return True
message_type = getattr(event, "message_type", "") or ""
return message_type in spec.message_types
if isinstance(spec, LocalFilterRefSpec):
binding = next(
(item for item in local_bindings if item.filter_id == spec.filter_id),
None,
)
if binding is None:
# LocalFilterRefSpec 只在当前 worker 持有同名 local binding 时可真正执行。
# 缺失 binding 往往意味着描述符来自远端/测试快照,此时保持 fail-open,
# 避免因为无法调用进程内函数而把原本可执行的 handler 错误过滤掉。
return True
return binding.evaluate(event=event, ctx=ctx)
if isinstance(spec, CompositeFilterSpec):
results = [
_evaluate_filter_spec_locally(
child,
local_bindings,
event=event,
ctx=ctx,
)
for child in spec.children
]
if spec.kind == "and":
return all(results)
return any(results)
return True
def custom_filter(
binding: FilterBinding,
) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
"""Attach a filter declaration to a handler."""
def decorator(func: Callable[..., Any]) -> Callable[..., Any]:
spec, local_bindings = binding.compile()
append_filter_meta(
func,
specs=[spec],
local_bindings=local_bindings,
)
return func
return decorator
def all_of(*bindings: FilterBinding) -> CompositeFilter:
return CompositeFilter("and", list(bindings))
def any_of(*bindings: FilterBinding) -> CompositeFilter:
return CompositeFilter("or", list(bindings))
__all__ = [
"CustomFilter",
"FilterBinding",
"LocalFilterBinding",
"MessageTypeFilter",
"PlatformFilter",
"all_of",
"any_of",
"custom_filter",
]
-105
View File
@@ -1,105 +0,0 @@
"""Canonical SDK LLM/tool/provider entrypoints for P0.5."""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from .agents import AgentSpec, BaseAgentRunner
from .entities import (
LLMToolSpec,
ProviderMeta,
ProviderRequest,
ProviderType,
RerankResult,
ToolCallsResult,
)
from .providers import (
EmbeddingProvider,
ProviderProxy,
RerankProvider,
STTProvider,
TTSAudioChunk,
TTSProvider,
)
from .tools import LLMToolManager
__all__ = [
"AgentSpec",
"BaseAgentRunner",
"EmbeddingProvider",
"LLMToolManager",
"LLMToolSpec",
"ProviderMeta",
"ProviderProxy",
"ProviderRequest",
"ProviderType",
"RerankProvider",
"RerankResult",
"STTProvider",
"TTSAudioChunk",
"TTSProvider",
"ToolCallsResult",
]
def __getattr__(name: str) -> Any:
if name in {"AgentSpec", "BaseAgentRunner"}:
from .agents import AgentSpec, BaseAgentRunner
return {"AgentSpec": AgentSpec, "BaseAgentRunner": BaseAgentRunner}[name]
if name in {
"LLMToolSpec",
"ProviderMeta",
"ProviderRequest",
"ProviderType",
"RerankResult",
"ToolCallsResult",
}:
from .entities import (
LLMToolSpec,
ProviderMeta,
ProviderRequest,
ProviderType,
RerankResult,
ToolCallsResult,
)
return {
"LLMToolSpec": LLMToolSpec,
"ProviderMeta": ProviderMeta,
"ProviderRequest": ProviderRequest,
"ProviderType": ProviderType,
"RerankResult": RerankResult,
"ToolCallsResult": ToolCallsResult,
}[name]
if name in {
"EmbeddingProvider",
"ProviderProxy",
"RerankProvider",
"STTProvider",
"TTSAudioChunk",
"TTSProvider",
}:
from .providers import (
EmbeddingProvider,
ProviderProxy,
RerankProvider,
STTProvider,
TTSAudioChunk,
TTSProvider,
)
return {
"EmbeddingProvider": EmbeddingProvider,
"ProviderProxy": ProviderProxy,
"RerankProvider": RerankProvider,
"STTProvider": STTProvider,
"TTSAudioChunk": TTSAudioChunk,
"TTSProvider": TTSProvider,
}[name]
if name == "LLMToolManager":
from .tools import LLMToolManager
return LLMToolManager
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
-39
View File
@@ -1,39 +0,0 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any
from pydantic import BaseModel, ConfigDict, Field
from .entities import ProviderRequest
if TYPE_CHECKING:
from ..context import Context
class AgentSpec(BaseModel):
model_config = ConfigDict(extra="forbid")
name: str
description: str = ""
tool_names: list[str] = Field(default_factory=list)
runner_class: str
def to_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_none=True)
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> AgentSpec:
return cls.model_validate(payload)
class BaseAgentRunner(ABC):
"""P0.5 agent registration surface.
P0.5 only supports agent registration metadata. Actual execution remains
owned by the core tool loop and is not directly callable from SDK plugins.
"""
@abstractmethod
async def run(self, ctx: Context, request: ProviderRequest) -> Any:
raise NotImplementedError
-137
View File
@@ -1,137 +0,0 @@
from __future__ import annotations
import enum
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
class _EntityModel(BaseModel):
model_config = ConfigDict(extra="forbid")
def to_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_none=True)
class ProviderType(str, enum.Enum):
CHAT_COMPLETION = "chat_completion"
SPEECH_TO_TEXT = "speech_to_text"
TEXT_TO_SPEECH = "text_to_speech"
EMBEDDING = "embedding"
RERANK = "rerank"
class ProviderMeta(_EntityModel):
id: str
model: str | None = None
type: str
provider_type: ProviderType = ProviderType.CHAT_COMPLETION
@classmethod
def from_payload(cls, payload: dict[str, Any] | None) -> ProviderMeta | None:
if not isinstance(payload, dict):
return None
return cls.model_validate(payload)
class ToolCallsResult(_EntityModel):
tool_call_id: str | None = None
tool_name: str
content: str
success: bool = True
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> ToolCallsResult:
return cls.model_validate(payload)
class RerankResult(_EntityModel):
index: int
score: float
document: str
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> RerankResult:
return cls.model_validate(payload)
class LLMToolSpec(_EntityModel):
name: str
description: str = ""
parameters_schema: dict[str, Any] = Field(
default_factory=lambda: {"type": "object", "properties": {}}
)
handler_ref: str | None = Field(
default=None,
description="Worker-side handler reference used to resolve the tool callable.",
)
handler_capability: str | None = Field(
default=None,
description="Optional capability name override for executing this tool handler.",
)
active: bool = True
@classmethod
def create(
cls,
*,
name: str,
description: str = "",
parameters_schema: dict[str, Any] | None = None,
handler_ref: str | None = None,
handler_capability: str | None = None,
active: bool = True,
) -> LLMToolSpec:
# Keep an explicit factory signature so static analyzers do not depend on
# Pydantic's generated __init__ when SDK call sites construct tool specs.
payload: dict[str, Any] = {
"name": name,
"description": description,
"parameters_schema": parameters_schema
if parameters_schema is not None
else {"type": "object", "properties": {}},
"active": active,
}
if handler_ref is not None:
payload["handler_ref"] = handler_ref
if handler_capability is not None:
payload["handler_capability"] = handler_capability
return cls.from_payload(payload)
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> LLMToolSpec:
return cls.model_validate(payload)
class ProviderRequest(_EntityModel):
prompt: str | None = None
system_prompt: str | None = None
session_id: str | None = None
contexts: list[dict[str, Any]] = Field(default_factory=list)
image_urls: list[str] = Field(default_factory=list)
tool_names: list[str] | None = None
tool_calls_result: list[ToolCallsResult] = Field(default_factory=list)
provider_id: str | None = None
model: str | None = None
temperature: float | None = None
max_steps: int | None = None
tool_call_timeout: int | None = None
def to_payload(self) -> dict[str, Any]:
payload = super().to_payload()
payload["tool_calls_result"] = [
item.to_payload() for item in self.tool_calls_result
]
return payload
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> ProviderRequest:
normalized = dict(payload)
raw_results = normalized.get("tool_calls_result")
if isinstance(raw_results, list):
normalized["tool_calls_result"] = [
ToolCallsResult.from_payload(item)
for item in raw_results
if isinstance(item, dict)
]
return cls.model_validate(normalized)
@@ -1,199 +0,0 @@
"""Provider-facing SDK entities and typed proxy helpers."""
from __future__ import annotations
import base64
from collections.abc import AsyncIterable, AsyncIterator
from dataclasses import dataclass
from ..clients._proxy import CapabilityProxy
from .entities import ProviderMeta, ProviderType, RerankResult
@dataclass(slots=True)
class TTSAudioChunk:
audio: bytes
text: str | None = None
class _BaseProviderProxy:
def __init__(self, proxy: CapabilityProxy, meta: ProviderMeta) -> None:
self._proxy = proxy
self._meta = meta
@property
def id(self) -> str:
return self._meta.id
@property
def model(self) -> str | None:
return self._meta.model
@property
def type(self) -> str:
return self._meta.type
@property
def provider_type(self) -> ProviderType:
return self._meta.provider_type
def meta(self) -> ProviderMeta:
return self._meta
class STTProvider(_BaseProviderProxy):
async def get_text(self, audio_url: str) -> str:
output = await self._proxy.call(
"provider.stt.get_text",
{"provider_id": self.id, "audio_url": str(audio_url)},
)
return str(output.get("text", ""))
class TTSProvider(_BaseProviderProxy):
def __init__(
self,
proxy: CapabilityProxy,
meta: ProviderMeta,
*,
supports_stream: bool = False,
) -> None:
super().__init__(proxy, meta)
self._supports_stream = supports_stream
async def get_audio(self, text: str) -> str:
output = await self._proxy.call(
"provider.tts.get_audio",
{"provider_id": self.id, "text": str(text)},
)
return str(output.get("audio_path", ""))
def support_stream(self) -> bool:
return self._supports_stream
async def get_audio_stream(
self,
text: str | AsyncIterable[str],
) -> AsyncIterator[TTSAudioChunk]:
payload = await self._build_stream_payload(text)
async for chunk in self._proxy.stream("provider.tts.get_audio_stream", payload):
audio_base64 = str(chunk.get("audio_base64", ""))
yield TTSAudioChunk(
audio=base64.b64decode(audio_base64) if audio_base64 else b"",
text=(
str(chunk.get("text")) if chunk.get("text") is not None else None
),
)
async def _build_stream_payload(
self,
text: str | AsyncIterable[str],
) -> dict[str, object]:
payload: dict[str, object] = {"provider_id": self.id}
if isinstance(text, str):
payload["text"] = text
return payload
payload["text_chunks"] = [str(item) async for item in text]
return payload
class EmbeddingProvider(_BaseProviderProxy):
async def get_embedding(self, text: str) -> list[float]:
output = await self._proxy.call(
"provider.embedding.get_embedding",
{"provider_id": self.id, "text": str(text)},
)
embedding = output.get("embedding")
if not isinstance(embedding, list):
return []
return [float(item) for item in embedding]
async def get_embeddings(self, texts: list[str]) -> list[list[float]]:
output = await self._proxy.call(
"provider.embedding.get_embeddings",
{
"provider_id": self.id,
"texts": [str(item) for item in texts],
},
)
embeddings = output.get("embeddings")
if not isinstance(embeddings, list):
return []
return [
[float(value) for value in item]
for item in embeddings
if isinstance(item, list)
]
async def get_dim(self) -> int:
output = await self._proxy.call(
"provider.embedding.get_dim",
{"provider_id": self.id},
)
return int(output.get("dim", 0))
class RerankProvider(_BaseProviderProxy):
async def rerank(
self,
query: str,
documents: list[str],
top_n: int | None = None,
) -> list[RerankResult]:
output = await self._proxy.call(
"provider.rerank.rerank",
{
"provider_id": self.id,
"query": str(query),
"documents": [str(item) for item in documents],
"top_n": top_n,
},
)
results = output.get("results")
if not isinstance(results, list):
return []
return [
RerankResult.from_payload(item)
for item in results
if isinstance(item, dict)
]
ProviderProxy = STTProvider | TTSProvider | EmbeddingProvider | RerankProvider
def provider_proxy_from_meta(
proxy: CapabilityProxy,
meta: ProviderMeta | None,
*,
tts_supports_stream: bool | None = None,
) -> ProviderProxy | None:
if meta is None:
return None
if meta.provider_type == ProviderType.SPEECH_TO_TEXT:
return STTProvider(proxy, meta)
if meta.provider_type == ProviderType.TEXT_TO_SPEECH:
return TTSProvider(
proxy,
meta,
supports_stream=bool(tts_supports_stream),
)
if meta.provider_type == ProviderType.EMBEDDING:
return EmbeddingProvider(proxy, meta)
if meta.provider_type == ProviderType.RERANK:
return RerankProvider(proxy, meta)
return None
__all__ = [
"EmbeddingProvider",
"ProviderMeta",
"ProviderProxy",
"ProviderType",
"RerankProvider",
"RerankResult",
"STTProvider",
"TTSAudioChunk",
"TTSProvider",
"provider_proxy_from_meta",
]
-59
View File
@@ -1,59 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from .entities import LLMToolSpec
if TYPE_CHECKING:
from ..clients._proxy import CapabilityProxy
class LLMToolManager:
def __init__(self, proxy: CapabilityProxy) -> None:
self._proxy = proxy
async def list_registered(self) -> list[LLMToolSpec]:
output = await self._proxy.call("llm_tool.manager.get", {})
items = output.get("registered")
if not isinstance(items, list):
return []
return [
LLMToolSpec.from_payload(item) for item in items if isinstance(item, dict)
]
async def list_active(self) -> list[LLMToolSpec]:
output = await self._proxy.call("llm_tool.manager.get", {})
items = output.get("active")
if not isinstance(items, list):
return []
return [
LLMToolSpec.from_payload(item) for item in items if isinstance(item, dict)
]
async def activate(self, name: str) -> bool:
output = await self._proxy.call("llm_tool.manager.activate", {"name": name})
return bool(output.get("activated", False))
async def deactivate(self, name: str) -> bool:
output = await self._proxy.call("llm_tool.manager.deactivate", {"name": name})
return bool(output.get("deactivated", False))
async def add(self, *tools: LLMToolSpec) -> list[str]:
output = await self._proxy.call(
"llm_tool.manager.add",
{"tools": [tool.to_payload() for tool in tools]},
)
result = output.get("names")
if not isinstance(result, list):
return []
return [str(item) for item in result]
async def remove(self, name: str) -> bool:
output = await self._proxy.call("llm_tool.manager.remove", {"name": name})
return bool(output.get("removed", False))
async def get(self, name: str) -> LLMToolSpec | None:
for tool in await self.list_registered():
if tool.name == name:
return tool
return None
@@ -1,103 +0,0 @@
"""Message component, result, and session subpackage."""
from .components import (
At as At,
)
from .components import (
AtAll as AtAll,
)
from .components import (
BaseMessageComponent as BaseMessageComponent,
)
from .components import (
File as File,
)
from .components import (
Forward as Forward,
)
from .components import (
Image as Image,
)
from .components import (
MediaHelper as MediaHelper,
)
from .components import (
Plain as Plain,
)
from .components import (
Poke as Poke,
)
from .components import (
Record as Record,
)
from .components import (
Reply as Reply,
)
from .components import (
UnknownComponent as UnknownComponent,
)
from .components import (
Video as Video,
)
from .components import (
build_media_component_from_url as build_media_component_from_url,
)
from .components import (
component_to_payload as component_to_payload,
)
from .components import (
component_to_payload_sync as component_to_payload_sync,
)
from .components import (
is_message_component as is_message_component,
)
from .components import (
payload_to_component as payload_to_component,
)
from .components import (
payloads_to_components as payloads_to_components,
)
from .result import (
EventResultType as EventResultType,
)
from .result import (
MessageBuilder as MessageBuilder,
)
from .result import (
MessageChain as MessageChain,
)
from .result import (
MessageEventResult as MessageEventResult,
)
from .result import (
coerce_message_chain as coerce_message_chain,
)
from .session import MessageSession as MessageSession
__all__ = [
"At",
"AtAll",
"BaseMessageComponent",
"EventResultType",
"File",
"Forward",
"Image",
"MediaHelper",
"MessageBuilder",
"MessageChain",
"MessageEventResult",
"MessageSession",
"Plain",
"Poke",
"Record",
"Reply",
"UnknownComponent",
"Video",
"build_media_component_from_url",
"coerce_message_chain",
"component_to_payload",
"component_to_payload_sync",
"is_message_component",
"payload_to_component",
"payloads_to_components",
]
@@ -1,622 +0,0 @@
"""SDK message component compatibility layer.
该模块有意避免在导入时导入遗留核心组件模块。
SDK工作线程应该保持轻量级并且不能依赖于主机核心引导程序
仅用于构造消息对象的路径。
"""
from __future__ import annotations
import asyncio
import base64
import inspect
import os
import tempfile
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
from urllib.request import urlretrieve
from .._internal.star_runtime import current_runtime_context
from ..errors import AstrBotError
_IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp"}
_RECORD_SUFFIXES = {".mp3", ".wav", ".ogg", ".flac", ".aac", ".m4a"}
_VIDEO_SUFFIXES = {".mp4", ".webm", ".mov", ".mkv", ".avi"}
def _temp_path(prefix: str, suffix: str = "") -> Path:
return Path(tempfile.gettempdir()) / f"{prefix}_{uuid.uuid4().hex}{suffix}"
def _guess_suffix_from_url(url: str, fallback: str = "") -> str:
suffix = Path(urlparse(url).path).suffix
return suffix or fallback
def _download_to_temp(url: str, prefix: str, fallback_suffix: str = "") -> str:
target = _temp_path(prefix, _guess_suffix_from_url(url, fallback_suffix))
urlretrieve(url, target)
return str(target.resolve())
async def _download_to_temp_async(
url: str,
prefix: str,
fallback_suffix: str = "",
) -> str:
return await asyncio.to_thread(
_download_to_temp,
url,
prefix,
fallback_suffix,
)
def _stringify_mapping(mapping: Mapping[Any, Any]) -> dict[str, Any]:
return {str(key): value for key, value in mapping.items()}
async def _register_file_to_service(path: str) -> str:
context = current_runtime_context()
if context is None:
raise RuntimeError("message component file service requires runtime context")
return await context._register_file_url(path)
def _reply_chain_payloads_sync(value: Any) -> list[dict[str, Any]]:
if not isinstance(value, list):
return []
return [component_to_payload_sync(item) for item in value]
async def _reply_chain_payloads(value: Any) -> list[dict[str, Any]]:
if not isinstance(value, list):
return []
return [await component_to_payload(item) for item in value]
def _coerce_reply_chain(value: Any) -> list[BaseMessageComponent]:
if not isinstance(value, list):
return []
if value and all(isinstance(item, BaseMessageComponent) for item in value):
return list(value)
return payloads_to_components(value)
def _component_type_name(component: Any) -> str:
raw_type = getattr(component, "type", "unknown")
normalized = getattr(raw_type, "value", raw_type)
return str(normalized or "unknown").lower()
def _resolve_media_kind(url: str, kind: str = "auto") -> str:
normalized_kind = str(kind).strip().lower() or "auto"
if normalized_kind != "auto":
return normalized_kind
suffix = Path(urlparse(url).path).suffix.lower()
if suffix in _IMAGE_SUFFIXES:
return "image"
if suffix in _RECORD_SUFFIXES:
return "record"
if suffix in _VIDEO_SUFFIXES:
return "video"
return "file"
def build_media_component_from_url(
url: str,
*,
kind: str = "auto",
) -> BaseMessageComponent:
url_text = str(url).strip()
if not url_text:
raise AstrBotError.invalid_input(
"MediaHelper.from_url requires a non-empty url"
)
resolved_kind = _resolve_media_kind(url_text, kind=kind)
if resolved_kind == "image":
return Image.fromURL(url_text)
if resolved_kind in {"record", "audio"}:
return Record.fromURL(url_text)
if resolved_kind == "video":
return Video.fromURL(url_text)
if resolved_kind == "file":
return File(name=_filename_from_url(url_text), url=url_text)
raise AstrBotError.invalid_input(
f"Unsupported media kind: {kind}",
details={"kind": kind, "url": url_text},
)
def _filename_from_url(url: str) -> str:
name = Path(urlparse(url).path).name
return name or "download"
class BaseMessageComponent:
type: str = "unknown"
def toDict(self) -> dict[str, Any]:
data: dict[str, Any] = {}
for key, value in self.__dict__.items():
if key == "type" or value is None:
continue
data["type" if key == "_type" else key] = value
return {"type": str(self.type).lower(), "data": data}
async def to_dict(self) -> dict[str, Any]:
return self.toDict()
class Plain(BaseMessageComponent):
type = "plain"
def __init__(self, text: str, convert: bool = True, **_: Any) -> None:
self.text = text
self.convert = convert
def toDict(self) -> dict[str, Any]:
return {"type": "text", "data": {"text": self.text.strip()}}
async def to_dict(self) -> dict[str, Any]:
return {"type": "text", "data": {"text": self.text}}
class At(BaseMessageComponent):
type = "at"
def __init__(self, qq: int | str, name: str | None = "", **_: Any) -> None:
self.qq = qq
self.name = name or ""
def toDict(self) -> dict[str, Any]:
return {"type": "at", "data": {"qq": str(self.qq)}}
class AtAll(At):
def __init__(self, **_: Any) -> None:
super().__init__(qq="all")
class Reply(BaseMessageComponent):
type = "reply"
def __init__(self, **kwargs: Any) -> None:
self.id = kwargs.get("id", "")
self.chain = _coerce_reply_chain(kwargs.get("chain", []))
self.sender_id = kwargs.get("sender_id", 0)
self.sender_nickname = kwargs.get("sender_nickname", "")
self.time = kwargs.get("time", 0)
self.message_str = kwargs.get("message_str", "")
self.text = kwargs.get("text", "")
self.qq = kwargs.get("qq", 0)
self.seq = kwargs.get("seq", 0)
def toDict(self) -> dict[str, Any]:
return {
"type": "reply",
"data": {
"id": self.id,
"chain": _reply_chain_payloads_sync(self.chain),
"sender_id": self.sender_id,
"sender_nickname": self.sender_nickname,
"time": self.time,
"message_str": self.message_str,
"text": self.text,
"qq": self.qq,
"seq": self.seq,
},
}
async def to_dict(self) -> dict[str, Any]:
return {
"type": "reply",
"data": {
"id": self.id,
"chain": await _reply_chain_payloads(self.chain),
"sender_id": self.sender_id,
"sender_nickname": self.sender_nickname,
"time": self.time,
"message_str": self.message_str,
"text": self.text,
"qq": self.qq,
"seq": self.seq,
},
}
class Image(BaseMessageComponent):
type = "image"
def __init__(self, file: str | None, **kwargs: Any) -> None:
self.file = file or ""
self._type = kwargs.get("_type", "")
self.subType = kwargs.get("subType", 0)
self.url = kwargs.get("url", "")
self.cache = kwargs.get("cache", True)
self.id = kwargs.get("id", 40000)
self.c = kwargs.get("c", 2)
self.path = kwargs.get("path", "")
self.file_unique = kwargs.get("file_unique", "")
@staticmethod
def fromURL(url: str, **kwargs: Any) -> Image:
return Image(url, **kwargs)
@staticmethod
def fromFileSystem(path: str, **kwargs: Any) -> Image:
return Image(f"file:///{os.path.abspath(path)}", path=path, **kwargs)
@staticmethod
def fromBase64(base64_data: str, **kwargs: Any) -> Image:
return Image(f"base64://{base64_data}", **kwargs)
async def convert_to_file_path(self) -> str:
url = self.url or self.file
if not url:
raise ValueError("No valid file or URL provided")
if url.startswith("file:///"):
return os.path.abspath(url[8:])
if url.startswith(("http://", "https://")):
return await _download_to_temp_async(url, "imgseg", ".jpg")
if url.startswith("base64://"):
file_path = _temp_path("imgseg", ".jpg")
file_path.write_bytes(base64.b64decode(url.removeprefix("base64://")))
return str(file_path.resolve())
if os.path.exists(url):
return os.path.abspath(url)
raise ValueError(f"not a valid file: {url}")
async def register_to_file_service(self) -> str:
return await _register_file_to_service(await self.convert_to_file_path())
class Record(BaseMessageComponent):
type = "record"
def __init__(self, file: str | None, **kwargs: Any) -> None:
self.file = file or ""
self.magic = kwargs.get("magic", False)
self.url = kwargs.get("url", "")
self.cache = kwargs.get("cache", True)
self.proxy = kwargs.get("proxy", True)
self.timeout = kwargs.get("timeout", 0)
self.text = kwargs.get("text")
self.path = kwargs.get("path")
@staticmethod
def fromFileSystem(path: str, **kwargs: Any) -> Record:
return Record(f"file:///{os.path.abspath(path)}", path=path, **kwargs)
@staticmethod
def fromURL(url: str, **kwargs: Any) -> Record:
return Record(url, **kwargs)
async def convert_to_file_path(self) -> str:
if self.file.startswith("file:///"):
return os.path.abspath(self.file[8:])
if self.file.startswith(("http://", "https://")):
return await _download_to_temp_async(self.file, "recordseg", ".dat")
if self.file.startswith("base64://"):
file_path = _temp_path("recordseg", ".dat")
file_path.write_bytes(base64.b64decode(self.file.removeprefix("base64://")))
return str(file_path.resolve())
if os.path.exists(self.file):
return os.path.abspath(self.file)
raise ValueError(f"not a valid file: {self.file}")
async def register_to_file_service(self) -> str:
return await _register_file_to_service(await self.convert_to_file_path())
class Video(BaseMessageComponent):
type = "video"
def __init__(self, file: str, **kwargs: Any) -> None:
self.file = file
self.cover = kwargs.get("cover", "")
self.c = kwargs.get("c", 2)
self.path = kwargs.get("path", "")
@staticmethod
def fromFileSystem(path: str, **kwargs: Any) -> Video:
return Video(f"file:///{os.path.abspath(path)}", path=path, **kwargs)
@staticmethod
def fromURL(url: str, **kwargs: Any) -> Video:
return Video(url, **kwargs)
async def convert_to_file_path(self) -> str:
if self.file.startswith("file:///"):
return os.path.abspath(self.file[8:])
if self.file.startswith(("http://", "https://")):
return await _download_to_temp_async(self.file, "videoseg")
if os.path.exists(self.file):
return os.path.abspath(self.file)
raise ValueError(f"not a valid file: {self.file}")
async def register_to_file_service(self) -> str:
return await _register_file_to_service(await self.convert_to_file_path())
class File(BaseMessageComponent):
type = "file"
def __init__(self, name: str, file: str = "", url: str = "") -> None:
self.name = name
self.file_ = file
self.url = url
@property
def file(self) -> str:
return self.file_
@file.setter
def file(self, value: str) -> None:
if value.startswith(("http://", "https://")):
self.url = value
else:
self.file_ = value
async def get_file(self, allow_return_url: bool = False) -> str:
if allow_return_url and self.url:
return self.url
if self.file_:
path = self.file_
if path.startswith("file://"):
path = path[7:]
if (
os.name == "nt"
and len(path) > 2
and path[0] == "/"
and path[2] == ":"
):
path = path[1:]
if os.path.exists(path):
return os.path.abspath(path)
if self.url:
suffix = Path(urlparse(self.url).path).suffix
target = await _download_to_temp_async(self.url, "fileseg", suffix)
self.file_ = target
return target
return ""
async def register_to_file_service(self) -> str:
return await _register_file_to_service(await self.get_file())
def toDict(self) -> dict[str, Any]:
payload_file = self.url or self.file_
return {
"type": "file",
"data": {
"name": self.name,
"file": payload_file,
},
}
async def to_dict(self) -> dict[str, Any]:
payload_file = await self.get_file(allow_return_url=True)
return {
"type": "file",
"data": {
"name": self.name,
"file": payload_file,
},
}
class Poke(BaseMessageComponent):
type = "poke"
def __init__(self, poke_type: str | int | None = None, **kwargs: Any) -> None:
legacy_type = kwargs.pop("type", None)
if poke_type is None:
poke_type = legacy_type
if poke_type in (None, "", "poke", "Poke"):
poke_type = "126"
self._type = str(poke_type)
self.id = kwargs.get("id")
self.qq = kwargs.get("qq", 0)
def target_id(self) -> str | None:
for value in (self.id, self.qq):
if value is None:
continue
text = str(value).strip()
if text and text != "0":
return text
return None
def toDict(self) -> dict[str, Any]:
data = {"type": str(self._type or "126")}
target_id = self.target_id()
if target_id:
data["id"] = target_id
return {"type": "poke", "data": data}
class Forward(BaseMessageComponent):
type = "forward"
def __init__(self, id: str, **_: Any) -> None:
self.id = id
class UnknownComponent(BaseMessageComponent):
type = "unknown"
def __init__(
self,
*,
raw_type: str = "unknown",
raw_data: dict[str, Any] | None = None,
) -> None:
self.raw_type = raw_type
self.raw_data = raw_data or {}
def toDict(self) -> dict[str, Any]:
return {
"type": self.raw_type or "unknown",
"data": dict(self.raw_data),
}
def is_message_component(value: Any) -> bool:
return isinstance(value, BaseMessageComponent)
def payload_to_component(payload: Any) -> BaseMessageComponent:
if not isinstance(payload, dict):
return UnknownComponent(raw_data={"value": payload})
raw_type = str(payload.get("type", "unknown") or "unknown").lower()
data = payload.get("data")
if not isinstance(data, dict):
data = {}
if raw_type in {"text", "plain"}:
return Plain(str(data.get("text", "")), convert=False)
if raw_type == "image":
return Image(str(data.get("file") or data.get("url") or ""))
if raw_type == "at":
qq_value = data.get("qq")
if str(qq_value).lower() == "all":
return AtAll()
qq = "" if qq_value is None else str(qq_value)
return At(qq=qq, name=str(data.get("name", "")))
if raw_type == "reply":
return Reply(**data)
if raw_type == "record":
return Record(str(data.get("file") or data.get("url") or ""), **data)
if raw_type == "video":
return Video(str(data.get("file") or ""), **data)
if raw_type == "file":
file_value = str(data.get("file") or data.get("file_") or "")
if not file_value:
file_value = str(data.get("url") or "")
return File(
str(data.get("name", "")),
file="" if file_value.startswith(("http://", "https://")) else file_value,
url=file_value if file_value.startswith(("http://", "https://")) else "",
)
if raw_type == "poke":
return Poke(
poke_type=data.get("type"),
id=data.get("id"),
qq=data.get("qq"),
)
if raw_type == "forward":
return Forward(id=str(data.get("id", "")))
return UnknownComponent(raw_type=raw_type, raw_data=_stringify_mapping(data))
def payloads_to_components(payloads: list[Any]) -> list[BaseMessageComponent]:
return [payload_to_component(item) for item in payloads]
def component_to_payload_sync(component: Any) -> dict[str, Any]:
if isinstance(component, UnknownComponent):
return component.toDict()
if isinstance(component, Plain):
return {"type": "text", "data": {"text": component.text}}
if _component_type_name(component) == "reply":
return {
"type": "reply",
"data": {
"id": getattr(component, "id", ""),
"chain": _reply_chain_payloads_sync(getattr(component, "chain", [])),
"sender_id": getattr(component, "sender_id", 0),
"sender_nickname": getattr(component, "sender_nickname", ""),
"time": getattr(component, "time", 0),
"message_str": getattr(component, "message_str", ""),
"text": getattr(component, "text", ""),
"qq": getattr(component, "qq", 0),
"seq": getattr(component, "seq", 0),
},
}
to_dict = getattr(component, "toDict", None)
if callable(to_dict):
result = to_dict()
if isinstance(result, Mapping):
return _stringify_mapping(result)
return {"type": "unknown", "data": {"value": str(component)}}
async def component_to_payload(component: Any) -> dict[str, Any]:
if isinstance(component, (UnknownComponent, Plain)):
return component_to_payload_sync(component)
async_method = getattr(component, "to_dict", None)
if callable(async_method):
payload = async_method()
if inspect.isawaitable(payload):
result = await payload
if isinstance(result, dict):
return result
return component_to_payload_sync(component)
class MediaHelper:
@staticmethod
async def from_url(
url: str,
*,
kind: str = "auto",
) -> BaseMessageComponent:
return build_media_component_from_url(url, kind=kind)
@staticmethod
async def download(url: str, save_dir: Path) -> Path:
url_text = str(url).strip()
if not url_text:
raise AstrBotError.invalid_input(
"MediaHelper.download requires a non-empty url"
)
parsed = urlparse(url_text)
if parsed.scheme not in {"http", "https"}:
raise AstrBotError.invalid_input(
"MediaHelper.download only supports http/https urls",
details={"url": url_text},
)
target_dir = Path(save_dir)
try:
target_dir.mkdir(parents=True, exist_ok=True)
except OSError as exc:
raise AstrBotError.internal_error(
f"Failed to prepare download directory: {target_dir}",
details={"save_dir": str(target_dir)},
) from exc
target_path = target_dir / _filename_from_url(url_text)
try:
await asyncio.to_thread(urlretrieve, url_text, target_path)
except Exception as exc:
raise AstrBotError.network_error(
f"Failed to download media from '{url_text}'",
details={"url": url_text},
) from exc
return target_path.resolve()
__all__ = [
"At",
"AtAll",
"BaseMessageComponent",
"File",
"Forward",
"Image",
"MediaHelper",
"Plain",
"Poke",
"Record",
"Reply",
"UnknownComponent",
"Video",
"component_to_payload",
"component_to_payload_sync",
"is_message_component",
"payload_to_component",
"payloads_to_components",
]
@@ -1,173 +0,0 @@
"""SDK-local rich message result objects.
本模块定义消息事件的结果对象,用于构建和返回富文本/多媒体消息。
核心类:
- MessageChain: 消息组件列表,支持同步/异步序列化为协议 payload
- MessageEventResult: 事件处理结果,包含类型标记和消息链
- EventResultType: 结果类型枚举(EMPTY / CHAIN)
辅助函数:
- coerce_message_chain: 将多种输入格式统一转换为 MessageChain,
支持 MessageEventResult、MessageChain、单个组件或组件列表
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
from .components import (
At,
AtAll,
BaseMessageComponent,
File,
Plain,
Reply,
build_media_component_from_url,
component_to_payload,
component_to_payload_sync,
is_message_component,
payloads_to_components,
)
class EventResultType(str, Enum):
EMPTY = "empty"
CHAIN = "chain"
@dataclass(slots=True)
class MessageChain:
components: list[BaseMessageComponent] = field(default_factory=list)
def append(self, component: BaseMessageComponent) -> MessageChain:
self.components.append(component)
return self
def extend(self, components: list[BaseMessageComponent]) -> MessageChain:
self.components.extend(components)
return self
def __iter__(self):
return iter(self.components)
def __len__(self) -> int:
return len(self.components)
def to_payload(self) -> list[dict[str, Any]]:
return [component_to_payload_sync(component) for component in self.components]
async def to_payload_async(self) -> list[dict[str, Any]]:
return [await component_to_payload(component) for component in self.components]
def get_plain_text(self, with_other_comps_mark: bool = False) -> str:
texts: list[str] = []
for component in self.components:
if isinstance(component, Plain):
texts.append(component.text)
elif with_other_comps_mark:
texts.append(f"[{component.__class__.__name__}]")
return " ".join(texts)
def plain_text(self, with_other_comps_mark: bool = False) -> str:
return self.get_plain_text(with_other_comps_mark=with_other_comps_mark)
@dataclass(slots=True)
class MessageEventResult:
type: EventResultType = EventResultType.EMPTY
chain: MessageChain = field(default_factory=MessageChain)
def to_payload(self) -> dict[str, Any]:
return {
"type": self.type.value,
"chain": self.chain.to_payload(),
}
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> MessageEventResult:
result_type_raw = str(payload.get("type", EventResultType.EMPTY.value))
try:
result_type = EventResultType(result_type_raw)
except ValueError:
result_type = EventResultType.EMPTY
chain_payload = payload.get("chain")
components = (
payloads_to_components(chain_payload)
if isinstance(chain_payload, list)
else []
)
return cls(type=result_type, chain=MessageChain(components))
@dataclass(slots=True)
class MessageBuilder:
components: list[BaseMessageComponent] = field(default_factory=list)
def text(self, content: str) -> MessageBuilder:
self.components.append(Plain(content, convert=False))
return self
def at(self, user_id: str) -> MessageBuilder:
self.components.append(At(user_id))
return self
def at_all(self) -> MessageBuilder:
self.components.append(AtAll())
return self
def image(self, url: str) -> MessageBuilder:
self.components.append(build_media_component_from_url(url, kind="image"))
return self
def record(self, url: str) -> MessageBuilder:
self.components.append(build_media_component_from_url(url, kind="record"))
return self
def video(self, url: str) -> MessageBuilder:
self.components.append(build_media_component_from_url(url, kind="video"))
return self
def file(self, name: str, *, file: str = "", url: str = "") -> MessageBuilder:
self.components.append(File(name=name, file=file, url=url))
return self
def reply(self, **kwargs: Any) -> MessageBuilder:
self.components.append(Reply(**kwargs))
return self
def append(self, component: BaseMessageComponent) -> MessageBuilder:
self.components.append(component)
return self
def extend(self, components: list[BaseMessageComponent]) -> MessageBuilder:
self.components.extend(components)
return self
def build(self) -> MessageChain:
return MessageChain(list(self.components))
def coerce_message_chain(value: Any) -> MessageChain | None:
if isinstance(value, MessageEventResult):
return value.chain
if isinstance(value, MessageChain):
return value
if is_message_component(value):
return MessageChain([value])
if isinstance(value, (list, tuple)) and all(
is_message_component(item) for item in value
):
return MessageChain(list(value))
return None
__all__ = [
"EventResultType",
"MessageChain",
"MessageBuilder",
"MessageEventResult",
"coerce_message_chain",
]
@@ -1,46 +0,0 @@
"""SDK-visible message session identifier.
本模块定义 MessageSession 类,用于统一表示消息会话标识符。
会话标识符格式为:platform_id:message_type:session_id
例如:
- qq:group:123456 表示 QQ 群 123456
- wechat:private:user789 表示微信私聊用户 user789
该格式与 AstrBot 核心的 unified_msg_origin 保持兼容,
确保 SDK 与核心之间的会话信息能够正确传递。
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(slots=True)
class MessageSession:
"""SDK-visible message session identifier.
The string form stays compatible with AstrBot's unified message origin:
``platform_id:message_type:session_id``.
"""
platform_id: str
message_type: str
session_id: str
def __post_init__(self) -> None:
self.platform_id = str(self.platform_id)
self.message_type = str(self.message_type).lower()
self.session_id = str(self.session_id)
def __str__(self) -> str:
return f"{self.platform_id}:{self.message_type}:{self.session_id}"
@classmethod
def from_str(cls, session: str) -> MessageSession:
platform_id, message_type, session_id = str(session).split(":", 2)
return cls(
platform_id=platform_id,
message_type=message_type,
session_id=session_id,
)
@@ -1,13 +0,0 @@
"""Backward-compatible alias for ``astrbot_sdk.message.components``.
This module intentionally aliases the implementation module instead of re-exporting
names one by one so private helpers keep working with existing monkeypatch sites.
"""
from __future__ import annotations
import sys
from .message import components as _components_module
sys.modules[__name__] = _components_module
@@ -1,13 +0,0 @@
"""Backward-compatible alias for ``astrbot_sdk.message.result``.
Use a module alias so callers patching helper functions on the legacy module path
still affect ``MessageBuilder`` and other implementation globals.
"""
from __future__ import annotations
import sys
from .message import result as _result_module
sys.modules[__name__] = _result_module
@@ -1,9 +0,0 @@
"""Backward-compatible message session exports.
The canonical implementation moved to ``astrbot_sdk.message.session``. Preserve the
legacy import path to avoid breaking existing plugins.
"""
from .message.session import MessageSession
__all__ = ["MessageSession"]
-38
View File
@@ -1,38 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Protocol, TypeVar, cast
if TYPE_CHECKING:
from .context import Context
_VT = TypeVar("_VT")
class _HasRuntimeContext(Protocol):
def _require_runtime_context(self) -> Context: ...
class PluginKVStoreMixin:
"""Plugin-scoped KV helpers backed by the runtime db client."""
def _runtime_context(self) -> Context:
owner = cast(_HasRuntimeContext, self)
return owner._require_runtime_context()
@property
def plugin_id(self) -> str:
ctx = self._runtime_context()
return ctx.plugin_id
async def put_kv_data(self, key: str, value: Any) -> None:
ctx = self._runtime_context()
await ctx.db.set(str(key), value)
async def get_kv_data(self, key: str, default: _VT) -> _VT:
ctx = self._runtime_context()
value = await ctx.db.get(str(key))
return default if value is None else value
async def delete_kv_data(self, key: str) -> None:
ctx = self._runtime_context()
await ctx.db.delete(str(key))
@@ -1,160 +0,0 @@
"""AstrBot v4 协议公共入口。
这里暴露 v4 原生协议的消息模型、描述符和解析函数。
握手阶段由 `InitializeMessage` 发起,返回值不是另一条 initialize 消息,而是
`ResultMessage(kind="initialize_result")`,其 `output` 负载可解析为
`InitializeOutput`。
## 插件作者指南:什么时候用什么?
### CapabilityDescriptor vs BUILTIN_CAPABILITY_SCHEMAS
**CapabilityDescriptor** 用于**声明**能力:
- 当你的插件想**暴露**一个可被其他插件或核心调用的能力时
- 例如:你的插件提供了一个翻译功能,想让其他插件调用
```python
from astrbot_sdk.protocol import CapabilityDescriptor
descriptor = CapabilityDescriptor(
name="my_plugin.translate", # 格式: 插件名.能力名
description="翻译文本到指定语言",
input_schema={
"type": "object",
"properties": {
"text": {"type": "string", "description": "要翻译的文本"},
"target_lang": {"type": "string", "description": "目标语言"},
},
"required": ["text", "target_lang"],
},
output_schema={
"type": "object",
"properties": {
"translated": {"type": "string"},
},
},
)
```
**BUILTIN_CAPABILITY_SCHEMAS** 用于**查询**内置能力的参数格式:
- 当你想**调用**核心提供的内置能力时,用它了解参数结构
- 例如:你想调用 `llm.chat`,但不确定参数格式
```python
from astrbot_sdk.protocol import BUILTIN_CAPABILITY_SCHEMAS
# 查看 llm.chat 的输入参数格式
schema = BUILTIN_CAPABILITY_SCHEMAS["llm.chat"]
print(schema["input"]) # 输入参数的 JSON Schema
print(schema["output"]) # 输出结果的 JSON Schema
```
### 命名规范
能力名称必须遵循 `{namespace}.{action}` 或 `{namespace}.{sub_namespace}.{action}` 格式:
- `llm.chat` - LLM 对话
- `db.set` - 数据库写入
- `llm_tool.manager.activate` - LLM 工具管理
**保留命名空间**(插件不可使用):
- `handler.` - 处理器相关
- `system.` - 系统内部能力
- `internal.` - 内部实现细节
### 常用内置能力速查
| 能力名 | 用途 |
|-------|------|
| `llm.chat` | 同步 LLM 对话 |
| `llm.stream_chat` | 流式 LLM 对话 |
| `memory.save` / `memory.get` | 短期记忆存储 |
| `db.set` / `db.get` | 持久化键值存储 |
| `platform.send` | 发送消息 |
| `provider.get_using` | 获取当前 Provider |
"""
from __future__ import annotations
from typing import Any
from . import _builtin_schemas as builtin_schemas
from .descriptors import ( # noqa: F401
BUILTIN_CAPABILITY_SCHEMAS,
CapabilityDescriptor,
CommandRouteSpec,
CommandTrigger,
CompositeFilterSpec,
EventTrigger,
FilterSpec,
HandlerDescriptor,
LocalFilterRefSpec,
MessageTrigger,
MessageTypeFilterSpec,
ParamSpec,
Permissions,
PlatformFilterSpec,
ScheduleTrigger,
SessionRef,
Trigger,
)
from .messages import ( # noqa: F401
CancelMessage,
ErrorPayload,
EventMessage,
InitializeMessage,
InitializeOutput,
InvokeMessage,
PeerInfo,
ProtocolMessage,
ResultMessage,
parse_message,
)
_DIRECT_EXPORTS = [
"BUILTIN_CAPABILITY_SCHEMAS",
"CapabilityDescriptor",
"CommandRouteSpec",
"CommandTrigger",
"CancelMessage",
"builtin_schemas",
"CompositeFilterSpec",
"ErrorPayload",
"EventTrigger",
"EventMessage",
"FilterSpec",
"HandlerDescriptor",
"InitializeMessage",
"InitializeOutput",
"InvokeMessage",
"LocalFilterRefSpec",
"MessageTrigger",
"MessageTypeFilterSpec",
"ParamSpec",
"PeerInfo",
"PlatformFilterSpec",
"Permissions",
"ProtocolMessage",
"ResultMessage",
"ScheduleTrigger",
"SessionRef",
"Trigger",
"parse_message",
]
_BUILTIN_SCHEMA_EXPORTS = tuple(
name for name in builtin_schemas.__all__ if name != "BUILTIN_CAPABILITY_SCHEMAS"
)
def __getattr__(name: str) -> Any:
if name in _BUILTIN_SCHEMA_EXPORTS:
return getattr(builtin_schemas, name)
raise AttributeError(name)
def __dir__() -> list[str]:
return sorted(set(globals()) | set(_BUILTIN_SCHEMA_EXPORTS))
__all__ = list(dict.fromkeys([*_DIRECT_EXPORTS, *_BUILTIN_SCHEMA_EXPORTS]))
File diff suppressed because it is too large Load Diff
@@ -1,521 +0,0 @@
"""v4 协议描述符模型。
`protocol` 是 v4 新引入的协议层抽象,不对应旧树(圣诞树)中的一个同名目录。这里
定义的是跨进程握手和调度时使用的声明式元数据,而不是运行时的具体处理器/
能力实现。
"""
from __future__ import annotations
from typing import Annotated, Any, Literal
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, model_validator
from . import _builtin_schemas
from ._builtin_schemas import * # noqa: F403
JSONSchema = _builtin_schemas.JSONSchema
RESERVED_CAPABILITY_NAMESPACES = ("handler", "system", "internal")
RESERVED_CAPABILITY_PREFIXES = tuple(
f"{namespace}." for namespace in RESERVED_CAPABILITY_NAMESPACES
)
BUILTIN_CAPABILITY_SCHEMAS = _builtin_schemas.BUILTIN_CAPABILITY_SCHEMAS
_BUILTIN_SCHEMA_EXPORTS = frozenset(_builtin_schemas.__all__)
def __getattr__(name: str) -> Any:
if name in _BUILTIN_SCHEMA_EXPORTS:
return getattr(_builtin_schemas, name)
raise AttributeError(name)
def __dir__() -> list[str]:
return sorted(set(globals()) | _BUILTIN_SCHEMA_EXPORTS)
class _DescriptorBase(BaseModel):
model_config = ConfigDict(extra="forbid")
class Permissions(_DescriptorBase):
"""权限配置,控制处理器的访问权限。
Attributes:
require_admin: 是否需要管理员权限
level: 权限等级,数值越高权限越大
"""
require_admin: bool = False
level: int = 0
class SessionRef(_DescriptorBase):
"""结构化会话目标。
v4 运行时内部仍然保留 legacy `session` 字符串作为最低兼容层,
但对外模型允许同时携带平台与原始寻址信息,避免平台发送接口长期
只依赖一个不透明字符串。
"""
conversation_id: str = Field(
validation_alias=AliasChoices("conversation_id", "session"),
)
platform: str | None = None
raw: dict[str, Any] | None = None
@property
def session(self) -> str:
return self.conversation_id
def to_payload(self) -> dict[str, Any]:
return self.model_dump(exclude_none=True)
class CommandTrigger(_DescriptorBase):
"""命令触发器,响应特定命令。
Attributes:
type: 触发器类型,固定为 "command"
command: 命令名称(不含前缀,如 "help")
aliases: 命令别名列表
description: 命令描述,用于帮助文档
platforms: 允许的平台列表,为空表示所有平台
message_types: 限定的消息类型列表,为空表示不限
"""
type: Literal["command"] = "command"
command: str
aliases: list[str] = Field(default_factory=list)
description: str | None = None
platforms: list[str] = Field(default_factory=list)
message_types: list[str] = Field(default_factory=list)
class MessageTrigger(_DescriptorBase):
"""消息触发器,描述消息类处理器的订阅条件。
Attributes:
type: 触发器类型,固定为 "message"
regex: 正则表达式模式,匹配消息文本
keywords: 关键词列表,消息包含任一关键词即触发
platforms: 目标平台列表,为空表示所有平台
message_types: 限定的消息类型列表,为空表示不限
Note:
`regex` 和 `keywords` 可以同时为空,此时表示 "任意消息均可触发",
仅由平台过滤或上层运行时进一步筛选。
"""
type: Literal["message"] = "message"
regex: str | None = None
keywords: list[str] = Field(default_factory=list)
platforms: list[str] = Field(default_factory=list)
message_types: list[str] = Field(default_factory=list)
class EventTrigger(_DescriptorBase):
"""事件触发器,响应特定类型的事件。
Attributes:
type: 触发器类型,固定为 "event"
event_type: 事件类型,字符串形式(如 "message"、"notice")
"""
type: Literal["event"] = "event"
event_type: str
class ScheduleTrigger(_DescriptorBase):
"""定时触发器,按 cron 表达式或固定间隔执行。
Attributes:
type: 触发器类型,固定为 "schedule"
cron: cron 表达式(如 "0 9 * * *" 表示每天 9 点)
interval_seconds: 执行间隔(秒)
Note:
cron 和 interval_seconds 必须且只能有一个非空。
"""
type: Literal["schedule"] = "schedule"
cron: str | None = Field(
default=None,
validation_alias=AliasChoices("cron", "schedule"),
)
interval_seconds: int | None = None
@property
def schedule(self) -> str | None:
return self.cron
@model_validator(mode="after")
def validate_schedule(self) -> ScheduleTrigger:
has_cron = self.cron is not None
has_interval = self.interval_seconds is not None
if has_cron == has_interval:
raise ValueError("cron 和 interval_seconds 必须且只能有一个非 null")
return self
class PlatformFilterSpec(_DescriptorBase):
kind: Literal["platform"] = "platform"
platforms: list[str] = Field(default_factory=list)
class MessageTypeFilterSpec(_DescriptorBase):
kind: Literal["message_type"] = "message_type"
message_types: list[str] = Field(default_factory=list)
class LocalFilterRefSpec(_DescriptorBase):
kind: Literal["local"] = "local"
filter_id: str
args: dict[str, Any] = Field(default_factory=dict)
class CompositeFilterSpec(_DescriptorBase):
kind: Literal["and", "or"]
children: list[FilterSpec] = Field(default_factory=list)
FilterSpec = Annotated[
PlatformFilterSpec
| MessageTypeFilterSpec
| LocalFilterRefSpec
| CompositeFilterSpec,
Field(discriminator="kind"),
]
class ParamSpec(_DescriptorBase):
name: str
type: Literal["str", "int", "float", "bool", "optional", "greedy_str"]
required: bool = True
inner_type: Literal["str", "int", "float", "bool"] | None = None
class CommandRouteSpec(_DescriptorBase):
group_path: list[str] = Field(default_factory=list)
display_command: str
group_help: str | None = None
CompositeFilterSpec.model_rebuild()
Trigger = Annotated[
CommandTrigger | MessageTrigger | EventTrigger | ScheduleTrigger,
Field(discriminator="type"),
]
"""触发器联合类型,使用 type 字段作为判别器自动解析具体类型。"""
class HandlerDescriptor(_DescriptorBase):
"""处理器描述符,描述一个事件处理函数的元信息。
Attributes:
id: 处理器唯一标识,通常是 "模块.函数名" 格式
trigger: 触发器配置,决定何时执行该处理器
kind: 处理器类别,默认普通 handler
contract: 运行时契约名,描述入参/执行语义
priority: 优先级,数值越大越先执行
permissions: 权限配置,控制谁可以触发该处理器
使用场景:
HandlerDescriptor 通常由 `@on_command`、`@on_message` 等装饰器自动创建,
插件作者一般不需要手动实例化。但了解其结构有助于理解插件注册机制。
触发器类型:
- CommandTrigger: 响应特定命令,如 `/help`
- MessageTrigger: 响应消息(正则/关键词匹配)
- EventTrigger: 响应特定事件类型
- ScheduleTrigger: 定时触发
示例:
插件作者通常通过装饰器声明处理器,框架会自动生成 HandlerDescriptor:
```python
from astrbot_sdk.decorators import on_command, on_message
# 命令处理器
@on_command("hello")
async def hello_handler(ctx: Context):
await ctx.reply("Hello!")
# 消息处理器(正则匹配)
@on_message(regex=r"^test\\s+(.+)$")
async def test_handler(ctx: Context):
await ctx.reply(f"收到: {ctx.match.group(1)}")
```
See Also:
Trigger: 触发器联合类型
Permissions: 权限配置
"""
id: str
trigger: Trigger
kind: Literal["handler", "hook", "tool", "session"] = "handler"
contract: str | None = None
description: str | None = None
priority: int = 0
permissions: Permissions = Field(default_factory=Permissions)
filters: list[FilterSpec] = Field(default_factory=list)
param_specs: list[ParamSpec] = Field(default_factory=list)
command_route: CommandRouteSpec | None = None
@model_validator(mode="after")
def validate_contract_defaults(self) -> HandlerDescriptor:
if self.contract is None:
if isinstance(self.trigger, ScheduleTrigger):
self.contract = "schedule"
else:
self.contract = "message_event"
return self
class CapabilityDescriptor(_DescriptorBase):
"""能力描述符,描述一个可调用的远程能力。
能力命名规范:
- 使用 "namespace.action" 格式,如 "llm.chat"、"db.set"
- 支持多级命名空间,如 "llm_tool.manager.activate"
- 内置能力以 "internal." 开头,如 "internal.legacy.call_context_function"
保留命名空间(插件不可使用):
- `handler.` - 处理器相关
- `system.` - 系统内部能力
- `internal.` - 内部实现细节
Attributes:
name: 能力名称,格式为 "namespace.action"
description: 能力描述,用于文档和调试
input_schema: 输入参数的 JSON Schema,用于验证
output_schema: 输出结果的 JSON Schema,用于验证
supports_stream: 是否支持流式响应
cancelable: 是否支持取消
使用场景:
当你的插件需要**暴露**一个可被其他插件调用的能力时,使用此类声明。
示例:
```python
from astrbot_sdk.protocol import CapabilityDescriptor
# 声明一个翻译能力
translate_desc = CapabilityDescriptor(
name="my_plugin.translate",
description="翻译文本到指定语言",
input_schema={
"type": "object",
"properties": {
"text": {"type": "string", "description": "要翻译的文本"},
"target_lang": {"type": "string", "description": "目标语言"},
},
"required": ["text", "target_lang"],
},
output_schema={
"type": "object",
"properties": {
"translated": {"type": "string"},
},
},
)
# 声明一个流式数据能力
stream_desc = CapabilityDescriptor(
name="my_plugin.stream_data",
description="流式返回数据",
supports_stream=True,
cancelable=True,
input_schema={"type": "object", "properties": {"count": {"type": "integer"}}},
output_schema={"type": "object", "properties": {"items": {"type": "array"}}},
)
```
注意:
如果你要调用**内置能力**(如 `llm.chat`、`db.set`),不需要手动创建
CapabilityDescriptor,而是直接通过 `Context.invoke()` 调用,或查阅
`BUILTIN_CAPABILITY_SCHEMAS` 了解参数格式。
See Also:
BUILTIN_CAPABILITY_SCHEMAS: 内置能力的 schema 定义,用于查询参数格式
"""
name: str
description: str
input_schema: JSONSchema | None = None
output_schema: JSONSchema | None = None
supports_stream: bool = False
cancelable: bool = False
@model_validator(mode="after")
def validate_builtin_schema_governance(self) -> CapabilityDescriptor:
builtin_schema = BUILTIN_CAPABILITY_SCHEMAS.get(self.name)
if builtin_schema is None:
return self
if self.input_schema is None or self.output_schema is None:
raise ValueError(
f"内建 capability {self.name} 必须同时提供 input_schema 和 output_schema"
)
if (
self.input_schema != builtin_schema["input"]
or self.output_schema != builtin_schema["output"]
):
raise ValueError(
f"内建 capability {self.name} 的 schema 必须与协议注册表保持一致"
)
return self
__all__ = [
"AGENT_REGISTRY_GET_INPUT_SCHEMA",
"AGENT_REGISTRY_GET_OUTPUT_SCHEMA",
"AGENT_REGISTRY_LIST_INPUT_SCHEMA",
"AGENT_REGISTRY_LIST_OUTPUT_SCHEMA",
"AGENT_SPEC_SCHEMA",
"AGENT_TOOL_LOOP_RUN_INPUT_SCHEMA",
"AGENT_TOOL_LOOP_RUN_OUTPUT_SCHEMA",
"BUILTIN_CAPABILITY_SCHEMAS",
"CapabilityDescriptor",
"CommandRouteSpec",
"CommandTrigger",
"CompositeFilterSpec",
"DB_DELETE_INPUT_SCHEMA",
"DB_DELETE_OUTPUT_SCHEMA",
"DB_GET_INPUT_SCHEMA",
"DB_GET_MANY_INPUT_SCHEMA",
"DB_GET_MANY_OUTPUT_SCHEMA",
"DB_GET_OUTPUT_SCHEMA",
"DB_LIST_INPUT_SCHEMA",
"DB_LIST_OUTPUT_SCHEMA",
"DB_SET_INPUT_SCHEMA",
"DB_SET_MANY_INPUT_SCHEMA",
"DB_SET_MANY_OUTPUT_SCHEMA",
"DB_SET_OUTPUT_SCHEMA",
"DB_WATCH_INPUT_SCHEMA",
"DB_WATCH_OUTPUT_SCHEMA",
"EventTrigger",
"FilterSpec",
"HTTP_LIST_APIS_INPUT_SCHEMA",
"HTTP_LIST_APIS_OUTPUT_SCHEMA",
"HTTP_REGISTER_API_INPUT_SCHEMA",
"HTTP_REGISTER_API_OUTPUT_SCHEMA",
"HTTP_UNREGISTER_API_INPUT_SCHEMA",
"HTTP_UNREGISTER_API_OUTPUT_SCHEMA",
"HandlerDescriptor",
"JSONSchema",
"LLM_CHAT_INPUT_SCHEMA",
"LLM_CHAT_OUTPUT_SCHEMA",
"LLM_CHAT_RAW_INPUT_SCHEMA",
"LLM_CHAT_RAW_OUTPUT_SCHEMA",
"LLM_STREAM_CHAT_INPUT_SCHEMA",
"LLM_STREAM_CHAT_OUTPUT_SCHEMA",
"LLM_TOOL_MANAGER_ACTIVATE_INPUT_SCHEMA",
"LLM_TOOL_MANAGER_ACTIVATE_OUTPUT_SCHEMA",
"LLM_TOOL_MANAGER_ADD_INPUT_SCHEMA",
"LLM_TOOL_MANAGER_ADD_OUTPUT_SCHEMA",
"LLM_TOOL_MANAGER_DEACTIVATE_INPUT_SCHEMA",
"LLM_TOOL_MANAGER_DEACTIVATE_OUTPUT_SCHEMA",
"LLM_TOOL_MANAGER_GET_INPUT_SCHEMA",
"LLM_TOOL_MANAGER_GET_OUTPUT_SCHEMA",
"LLM_TOOL_SPEC_SCHEMA",
"LocalFilterRefSpec",
"MEMORY_DELETE_INPUT_SCHEMA",
"MEMORY_DELETE_MANY_INPUT_SCHEMA",
"MEMORY_DELETE_MANY_OUTPUT_SCHEMA",
"MEMORY_DELETE_OUTPUT_SCHEMA",
"MEMORY_GET_INPUT_SCHEMA",
"MEMORY_GET_MANY_INPUT_SCHEMA",
"MEMORY_GET_MANY_OUTPUT_SCHEMA",
"MEMORY_GET_OUTPUT_SCHEMA",
"MEMORY_SAVE_INPUT_SCHEMA",
"MEMORY_SAVE_OUTPUT_SCHEMA",
"MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA",
"MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA",
"MEMORY_SEARCH_INPUT_SCHEMA",
"MEMORY_SEARCH_OUTPUT_SCHEMA",
"MEMORY_STATS_INPUT_SCHEMA",
"MEMORY_STATS_OUTPUT_SCHEMA",
"METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA",
"METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA",
"METADATA_GET_PLUGIN_INPUT_SCHEMA",
"METADATA_GET_PLUGIN_OUTPUT_SCHEMA",
"METADATA_LIST_PLUGINS_INPUT_SCHEMA",
"METADATA_LIST_PLUGINS_OUTPUT_SCHEMA",
"MessageTrigger",
"MessageTypeFilterSpec",
"PROVIDER_GET_CURRENT_CHAT_PROVIDER_ID_INPUT_SCHEMA",
"PROVIDER_GET_CURRENT_CHAT_PROVIDER_ID_OUTPUT_SCHEMA",
"PROVIDER_GET_USING_INPUT_SCHEMA",
"PROVIDER_GET_USING_OUTPUT_SCHEMA",
"PROVIDER_LIST_ALL_INPUT_SCHEMA",
"PROVIDER_LIST_ALL_OUTPUT_SCHEMA",
"PROVIDER_META_SCHEMA",
"PLATFORM_GET_GROUP_INPUT_SCHEMA",
"PLATFORM_GET_GROUP_OUTPUT_SCHEMA",
"PLATFORM_GET_MEMBERS_INPUT_SCHEMA",
"PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA",
"PLATFORM_INSTANCE_SCHEMA",
"PLATFORM_LIST_INSTANCES_INPUT_SCHEMA",
"PLATFORM_LIST_INSTANCES_OUTPUT_SCHEMA",
"PLATFORM_SEND_BY_SESSION_INPUT_SCHEMA",
"PLATFORM_SEND_BY_SESSION_OUTPUT_SCHEMA",
"PLATFORM_SEND_CHAIN_INPUT_SCHEMA",
"PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA",
"PLATFORM_SEND_IMAGE_INPUT_SCHEMA",
"PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA",
"PLATFORM_SEND_INPUT_SCHEMA",
"PLATFORM_SEND_OUTPUT_SCHEMA",
"ParamSpec",
"Permissions",
"PlatformFilterSpec",
"REGISTRY_COMMAND_REGISTER_INPUT_SCHEMA",
"REGISTRY_COMMAND_REGISTER_OUTPUT_SCHEMA",
"REGISTRY_GET_HANDLER_BY_FULL_NAME_INPUT_SCHEMA",
"REGISTRY_GET_HANDLER_BY_FULL_NAME_OUTPUT_SCHEMA",
"REGISTRY_GET_HANDLERS_BY_EVENT_TYPE_INPUT_SCHEMA",
"REGISTRY_GET_HANDLERS_BY_EVENT_TYPE_OUTPUT_SCHEMA",
"RESERVED_CAPABILITY_NAMESPACES",
"RESERVED_CAPABILITY_PREFIXES",
"SESSION_PLUGIN_FILTER_HANDLERS_INPUT_SCHEMA",
"SESSION_PLUGIN_FILTER_HANDLERS_OUTPUT_SCHEMA",
"SESSION_PLUGIN_IS_ENABLED_INPUT_SCHEMA",
"SESSION_PLUGIN_IS_ENABLED_OUTPUT_SCHEMA",
"SESSION_REF_SCHEMA",
"SESSION_SERVICE_IS_LLM_ENABLED_INPUT_SCHEMA",
"SESSION_SERVICE_IS_LLM_ENABLED_OUTPUT_SCHEMA",
"SESSION_SERVICE_IS_TTS_ENABLED_INPUT_SCHEMA",
"SESSION_SERVICE_IS_TTS_ENABLED_OUTPUT_SCHEMA",
"SESSION_SERVICE_SET_LLM_STATUS_INPUT_SCHEMA",
"SESSION_SERVICE_SET_LLM_STATUS_OUTPUT_SCHEMA",
"SESSION_SERVICE_SET_TTS_STATUS_INPUT_SCHEMA",
"SESSION_SERVICE_SET_TTS_STATUS_OUTPUT_SCHEMA",
"ScheduleTrigger",
"SessionRef",
"SYSTEM_EVENT_HANDLER_WHITELIST_GET_INPUT_SCHEMA",
"SYSTEM_EVENT_HANDLER_WHITELIST_GET_OUTPUT_SCHEMA",
"SYSTEM_EVENT_HANDLER_WHITELIST_SET_INPUT_SCHEMA",
"SYSTEM_EVENT_HANDLER_WHITELIST_SET_OUTPUT_SCHEMA",
"SYSTEM_EVENT_LLM_GET_STATE_INPUT_SCHEMA",
"SYSTEM_EVENT_LLM_GET_STATE_OUTPUT_SCHEMA",
"SYSTEM_EVENT_LLM_REQUEST_INPUT_SCHEMA",
"SYSTEM_EVENT_LLM_REQUEST_OUTPUT_SCHEMA",
"SYSTEM_EVENT_REACT_INPUT_SCHEMA",
"SYSTEM_EVENT_REACT_OUTPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_CLEAR_INPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_CLEAR_OUTPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_GET_INPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_GET_OUTPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_SET_INPUT_SCHEMA",
"SYSTEM_EVENT_RESULT_SET_OUTPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA",
"SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA",
"SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA",
"SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA",
"Trigger",
]
@@ -1,289 +0,0 @@
"""v4 协议消息模型。
这些模型描述的是 `Peer` 与 `Peer` 之间的线协议。握手阶段通过
`InitializeMessage` 发起,再由 `ResultMessage(kind="initialize_result")`
返回 `InitializeOutput`;能力调用阶段则使用 `InvokeMessage` / `ResultMessage`
或 `EventMessage` 序列。
"""
from __future__ import annotations
import json
from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator
from .descriptors import CapabilityDescriptor, HandlerDescriptor
class _MessageBase(BaseModel):
model_config = ConfigDict(extra="forbid")
class ErrorPayload(_MessageBase):
"""错误载荷,用于 ResultMessage 和 EventMessage 中传递错误信息。
Attributes:
code: 错误码,字符串类型,便于语义化错误分类
message: 错误消息,人类可读的错误描述
hint: 错误提示,可选的解决方案或建议
retryable: 是否可重试,标识该错误是否可通过重试解决
docs_url: 可选的文档链接,帮助调用方定位更多说明
details: 可选的结构化细节,便于调试和日志展示
"""
code: str
message: str
hint: str = ""
retryable: bool = False
docs_url: str = ""
details: dict[str, Any] | None = None
class PeerInfo(_MessageBase):
"""对等节点信息,标识消息发送方的身份。
Attributes:
name: 节点名称,通常是插件 ID 或核心标识
role: 节点角色,"plugin" 或 "core"
version: 节点版本号,可选
"""
name: str
role: Literal["plugin", "core"]
version: str | None = None
class InitializeMessage(_MessageBase):
"""初始化消息,用于建立连接时交换信息。
Attributes:
type: 消息类型,固定为 "initialize"
id: 消息 ID,用于关联响应
protocol_version: 协议版本号
peer: 发送方节点信息
handlers: 注册的处理器描述符列表
provided_capabilities: 发送方对外暴露的能力描述符列表
metadata: 扩展元数据,可存储插件配置等信息
"""
type: Literal["initialize"] = "initialize"
id: str
protocol_version: str
peer: PeerInfo
handlers: list[HandlerDescriptor] = Field(default_factory=list)
provided_capabilities: list[CapabilityDescriptor] = Field(default_factory=list)
metadata: dict[str, Any] = Field(default_factory=dict)
class InitializeOutput(_MessageBase):
"""初始化输出,作为 InitializeMessage 的响应数据。
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)
class ResultMessage(_MessageBase):
"""结果消息,用于返回能力调用的结果。
Attributes:
type: 消息类型,固定为 "result"
id: 关联的请求 ID
kind: 结果类型,可选,如 "initialize_result" 标识初始化结果
success: 是否成功
output: 成功时的输出数据
error: 失败时的错误信息
"""
type: Literal["result"] = "result"
id: str
kind: str | None = None
success: bool
output: dict[str, Any] = Field(default_factory=dict)
error: ErrorPayload | None = None
@model_validator(mode="after")
def validate_result_state(self) -> ResultMessage:
"""约束 success / output / error 的组合状态。"""
if self.success:
if self.error is not None:
raise ValueError("success=true 时 error 必须为空")
return self
if self.error is None:
raise ValueError("success=false 时必须提供 error")
if self.output:
raise ValueError("success=false 时 output 必须为空")
return self
class InvokeMessage(_MessageBase):
"""调用消息,用于请求执行远程能力。
Attributes:
type: 消息类型,固定为 "invoke"
id: 请求 ID,用于关联响应
capability: 目标能力名称,格式为 "namespace.action"
input: 调用输入参数
stream: 是否期望流式响应,若为 True 将收到 EventMessage 序列
caller_plugin_id: 运行时透传的调用方插件 ID,不属于业务 payload
"""
type: Literal["invoke"] = "invoke"
id: str
capability: str
input: dict[str, Any] = Field(default_factory=dict)
stream: bool = False
caller_plugin_id: str | None = None
class EventMessage(_MessageBase):
"""事件消息,用于流式调用的状态通知。
流式调用生命周期:
1. started: 调用开始,所有字段为空
2. delta: 数据增量更新,包含 data 字段
3. completed: 调用完成,包含 output 字段
4. failed: 调用失败,包含 error 字段
Attributes:
type: 消息类型,固定为 "event"
id: 关联的请求 ID
phase: 事件阶段,started/delta/completed/failed
data: 增量数据,仅 delta 阶段有效
output: 最终输出,仅 completed 阶段有效
error: 错误信息,仅 failed 阶段有效
"""
type: Literal["event"] = "event"
id: str
phase: Literal["started", "delta", "completed", "failed"]
data: dict[str, Any] = Field(default_factory=dict)
output: dict[str, Any] = Field(default_factory=dict)
error: ErrorPayload | None = None
@model_validator(mode="after")
def validate_phase_constraints(self) -> EventMessage:
"""验证各 phase 的字段约束。
- started: 所有字段必须为空
- delta: 必须有 data,output/error 必须为空
- completed: 必须有 output,data/error 必须为空
- failed: 必须有 error,data/output 必须为空
"""
phase = self.phase
if phase == "started":
if self.data or self.output or self.error:
raise ValueError("started phase 必须所有字段为空")
elif phase == "delta":
if not self.data:
raise ValueError("delta phase 需要 data")
if self.output or self.error:
raise ValueError("delta phase 的 output/error 必须为空")
elif phase == "completed":
if not self.output:
raise ValueError("completed phase 需要 output")
if self.data or self.error:
raise ValueError("completed phase 的 data/error 必须为空")
elif phase == "failed":
if self.error is None:
raise ValueError("failed phase 需要 error")
if self.data or self.output:
raise ValueError("failed phase 的 data/output 必须为空")
return self
class CancelMessage(_MessageBase):
"""取消消息,用于取消正在进行的调用。
Attributes:
type: 消息类型,固定为 "cancel"
id: 要取消的请求 ID
reason: 取消原因,默认为 "user_cancelled"
"""
type: Literal["cancel"] = "cancel"
id: str
reason: str = "user_cancelled"
ProtocolMessage = (
InitializeMessage | ResultMessage | InvokeMessage | EventMessage | CancelMessage
)
"""协议消息联合类型,所有有效消息类型的联合。"""
_PROTOCOL_MESSAGE_MODELS = {
"initialize": InitializeMessage,
"result": ResultMessage,
"invoke": InvokeMessage,
"event": EventMessage,
"cancel": CancelMessage,
}
def parse_message(
payload: ProtocolMessage | str | bytes | dict[str, Any],
) -> ProtocolMessage:
"""解析协议消息。
从原始载荷(字符串、字节或字典)解析为对应的 ProtocolMessage 类型。
根据 "type" 字段自动识别消息类型并验证。
Args:
payload: 原始消息载荷,支持已解析模型、JSON 字符串、字节或字典
Returns:
解析后的协议消息对象
Raises:
ValueError: 未知的消息类型
Example:
>>> msg = parse_message('{"type": "invoke", "id": "1", "capability": "test"}')
>>> isinstance(msg, InvokeMessage)
True
"""
if isinstance(
payload,
(
InitializeMessage,
ResultMessage,
InvokeMessage,
EventMessage,
CancelMessage,
),
):
return payload
if isinstance(payload, bytes):
payload = payload.decode("utf-8")
if isinstance(payload, str):
payload = json.loads(payload)
if not isinstance(payload, dict):
raise ValueError("协议消息必须是 JSON object")
message_type = payload.get("type")
model = _PROTOCOL_MESSAGE_MODELS.get(str(message_type))
if model is not None:
return model.model_validate(payload)
raise ValueError(f"未知消息类型:{message_type}")
__all__ = [
"CancelMessage",
"ErrorPayload",
"EventMessage",
"InitializeMessage",
"InitializeOutput",
"InvokeMessage",
"PeerInfo",
"ProtocolMessage",
"ResultMessage",
"parse_message",
]
@@ -1,63 +0,0 @@
"""AstrBot SDK runtime public exports.
本模块提供运行时核心组件的公共导出,包括:
- CapabilityRouter: 能力路由器,处理能力调用的分发和路由
- HandlerDispatcher: 事件处理器分发器,将事件分发到注册的 handler
- Peer: 与 AstrBot 核心通信的对等端抽象
- Transport 系列: 进程间通信传输层实现(stdio/websocket)
延迟加载策略:
为避免导入时触发 websocket/aiohttp 等重型依赖,采用 __getattr__ 实现按需加载。
这样轻量级导入(如仅使用类型提示)不会产生不必要的依赖开销。
"""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from .capability_router import CapabilityRouter, StreamExecution
from .handler_dispatcher import HandlerDispatcher
from .peer import Peer
from .transport import (
MessageHandler,
StdioTransport,
Transport,
WebSocketClientTransport,
WebSocketServerTransport,
)
__all__ = [
"CapabilityRouter",
"HandlerDispatcher",
"MessageHandler",
"Peer",
"StdioTransport",
"StreamExecution",
"Transport",
"WebSocketClientTransport",
"WebSocketServerTransport",
]
def __getattr__(name: str) -> Any:
if name in {"CapabilityRouter", "StreamExecution"}:
module = import_module(".capability_router", __name__)
return getattr(module, name)
if name == "HandlerDispatcher":
module = import_module(".handler_dispatcher", __name__)
return getattr(module, name)
if name == "Peer":
module = import_module(".peer", __name__)
return getattr(module, name)
if name in {
"MessageHandler",
"StdioTransport",
"Transport",
"WebSocketClientTransport",
"WebSocketServerTransport",
}:
module = import_module(".transport", __name__)
return getattr(module, name)
raise AttributeError(name)
@@ -1,62 +0,0 @@
from __future__ import annotations
from .bridge_base import CapabilityRouterBridgeBase
from .capabilities import (
ConversationCapabilityMixin,
DBCapabilityMixin,
HttpCapabilityMixin,
KnowledgeBaseCapabilityMixin,
LLMCapabilityMixin,
McpCapabilityMixin,
MemoryCapabilityMixin,
MessageHistoryCapabilityMixin,
MetadataCapabilityMixin,
PersonaCapabilityMixin,
PlatformCapabilityMixin,
ProviderCapabilityMixin,
SessionCapabilityMixin,
SkillCapabilityMixin,
SystemCapabilityMixin,
)
class BuiltinCapabilityRouterMixin(
LLMCapabilityMixin,
MemoryCapabilityMixin,
DBCapabilityMixin,
PlatformCapabilityMixin,
HttpCapabilityMixin,
MetadataCapabilityMixin,
ProviderCapabilityMixin,
McpCapabilityMixin,
SessionCapabilityMixin,
SkillCapabilityMixin,
PersonaCapabilityMixin,
ConversationCapabilityMixin,
MessageHistoryCapabilityMixin,
KnowledgeBaseCapabilityMixin,
SystemCapabilityMixin,
CapabilityRouterBridgeBase,
):
def _register_builtin_capabilities(self) -> None:
self._register_llm_capabilities()
self._register_memory_capabilities()
self._register_db_capabilities()
self._register_platform_capabilities()
self._register_http_capabilities()
self._register_metadata_capabilities()
self._register_provider_capabilities()
self._register_agent_tool_capabilities()
self._register_mcp_capabilities()
self._register_session_capabilities()
self._register_skill_capabilities()
self._register_persona_capabilities()
self._register_conversation_capabilities()
self._register_message_history_capabilities()
self._register_kb_capabilities()
self._register_provider_manager_capabilities()
self._register_platform_manager_capabilities()
self._register_system_capabilities()
__all__ = ["BuiltinCapabilityRouterMixin"]
@@ -1,102 +0,0 @@
from __future__ import annotations
import asyncio
from datetime import datetime
from pathlib import Path
from typing import Any
from ...protocol.descriptors import CapabilityDescriptor
class CapabilityRouterHost:
memory_store: dict[str, dict[str, Any]]
_memory_backends: dict[str, Any]
_memory_index: dict[str, dict[str, Any]]
_memory_dirty_keys: set[str]
_memory_expires_at: dict[str, datetime | None]
db_store: dict[str, Any]
sent_messages: list[dict[str, Any]]
event_actions: list[dict[str, Any]]
http_api_store: list[dict[str, Any]]
_event_streams: dict[str, dict[str, Any]]
_plugins: dict[str, Any]
_request_overlays: dict[str, dict[str, Any]]
_provider_catalog: dict[str, list[dict[str, Any]]]
_provider_configs: dict[str, dict[str, Any]]
_active_provider_ids: dict[str, str | None]
_provider_change_subscriptions: dict[str, asyncio.Queue[dict[str, Any]]]
_system_data_root: Path
_session_waiters: dict[str, set[str]]
_session_plugin_configs: dict[str, dict[str, Any]]
_session_service_configs: dict[str, dict[str, Any]]
_db_watch_subscriptions: dict[str, tuple[str | None, asyncio.Queue[dict[str, Any]]]]
_dynamic_command_routes: dict[str, list[dict[str, Any]]]
_file_token_store: dict[str, str]
_platform_instances: list[dict[str, Any]]
_persona_store: dict[str, dict[str, Any]]
_conversation_store: dict[str, dict[str, Any]]
_session_current_conversation_ids: dict[str, str]
_kb_store: dict[str, dict[str, Any]]
_kb_document_store: dict[str, dict[str, dict[str, Any]]]
_kb_document_content_store: dict[str, str]
def register(
self,
descriptor: CapabilityDescriptor,
*,
call_handler=None,
stream_handler=None,
finalize=None,
exposed: bool = True,
) -> None:
raise NotImplementedError
def _emit_db_change(self, *, op: str, key: str, value: Any | None) -> None:
raise NotImplementedError
@staticmethod
def _require_caller_plugin_id(capability_name: str) -> str:
raise NotImplementedError
@staticmethod
def _validated_plugin_id(plugin_id: str, *, capability_name: str) -> str:
raise NotImplementedError
def _plugin_data_dir(self, plugin_id: str, *, capability_name: str) -> Path:
raise NotImplementedError
def register_dynamic_command_route(
self,
*,
plugin_id: str,
command_name: str,
handler_full_name: str,
desc: str = "",
priority: int = 0,
use_regex: bool = False,
) -> None:
raise NotImplementedError
def get_platform_instances(self) -> list[dict[str, Any]]:
raise NotImplementedError
def _register_agent_tool_capabilities(self) -> None:
raise NotImplementedError
def _provider_entry(
self,
payload: dict[str, Any],
capability_name: str,
expected_kind: str | None = None,
) -> dict[str, Any]:
raise NotImplementedError
async def _provider_embedding_get_embedding(
self, request_id: str, payload: dict[str, Any], token
) -> dict[str, Any]:
raise NotImplementedError
async def _provider_embedding_get_embeddings(
self, request_id: str, payload: dict[str, Any], token
) -> dict[str, Any]:
raise NotImplementedError
@@ -1,187 +0,0 @@
from __future__ import annotations
import copy
import hashlib
import math
import re
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from ..._internal.plugin_ids import resolve_plugin_data_dir, validate_plugin_id
from ...errors import AstrBotError
from ...protocol.descriptors import (
BUILTIN_CAPABILITY_SCHEMAS,
CapabilityDescriptor,
SessionRef,
)
from ._host import CapabilityRouterHost
def _clone_target_payload(value: Any) -> dict[str, Any] | None:
if not isinstance(value, dict):
return None
return {str(key): item for key, item in value.items()}
def _clone_chain_payload(value: Any) -> list[dict[str, Any]]:
if not isinstance(value, list):
return []
return [
{str(key): item for key, item in chunk.items()}
for chunk in value
if isinstance(chunk, dict)
]
_MOCK_EMBEDDING_DIM = 24
def _embedding_terms(text: str) -> list[str]:
"""Build stable tokens for the mock embedding implementation."""
normalized = re.sub(r"\s+", " ", str(text).strip().casefold())
compact = normalized.replace(" ", "")
if not normalized:
return []
terms = [word for word in re.findall(r"\w+", normalized, flags=re.UNICODE) if word]
if compact:
if len(compact) == 1:
terms.append(compact)
else:
terms.extend(
compact[index : index + 2] for index in range(len(compact) - 1)
)
terms.append(compact)
return terms or [normalized]
def _mock_embedding_vector(text: str, *, provider_id: str) -> list[float]:
"""Generate a deterministic normalized mock embedding vector."""
values = [0.0] * _MOCK_EMBEDDING_DIM
for term in _embedding_terms(text):
digest = hashlib.sha256(f"{provider_id}:{term}".encode()).digest()
index = int.from_bytes(digest[:2], "big") % _MOCK_EMBEDDING_DIM
values[index] += 1.0 + min(len(term), 8) * 0.05
norm = math.sqrt(sum(value * value for value in values))
if norm <= 0:
return values
return [value / norm for value in values]
class CapabilityRouterBridgeBase(CapabilityRouterHost):
_memory_backends: dict[str, Any]
@staticmethod
def _validated_plugin_id(plugin_id: str, *, capability_name: str) -> str:
try:
return validate_plugin_id(plugin_id)
except ValueError as exc:
raise AstrBotError.invalid_input(
f"{capability_name} requires a safe plugin_id: {exc}"
) from exc
def _plugin_data_dir(self, plugin_id: str, *, capability_name: str) -> Path:
try:
return resolve_plugin_data_dir(self._system_data_root, plugin_id)
except ValueError as exc:
raise AstrBotError.invalid_input(
f"{capability_name} requires a safe plugin_id: {exc}"
) from exc
def _builtin_descriptor(
self,
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 _resolve_target(
self, 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
@staticmethod
def _is_group_session(session: str) -> bool:
normalized = str(session).lower()
return ":group:" in normalized or ":groupmessage:" in normalized
@staticmethod
def _mock_group_payload(session: str) -> dict[str, Any] | None:
if not CapabilityRouterBridgeBase._is_group_session(session):
return None
members = [
{
"user_id": f"{session}:member-1",
"nickname": "Member 1",
"role": "member",
},
{
"user_id": f"{session}:member-2",
"nickname": "Member 2",
"role": "admin",
},
]
return {
"group_id": session.rsplit(":", maxsplit=1)[-1],
"group_name": f"Mock Group {session.rsplit(':', maxsplit=1)[-1]}",
"group_avatar": "",
"group_owner": members[0]["user_id"],
"group_admins": [members[1]["user_id"]],
"members": members,
}
def _session_plugin_config(self, session: str) -> dict[str, Any]:
config = self._session_plugin_configs.get(str(session), {})
return dict(config) if isinstance(config, dict) else {}
def _session_service_config(self, session: str) -> dict[str, Any]:
config = self._session_service_configs.get(str(session), {})
return dict(config) if isinstance(config, dict) else {}
@staticmethod
def _now_iso() -> str:
return datetime.now(timezone.utc).isoformat()
@staticmethod
def _session_platform_id(session: str) -> str:
parts = str(session).split(":", maxsplit=1)
if parts and parts[0].strip():
return parts[0].strip()
return "unknown"
@staticmethod
def _normalize_history_payload(value: Any) -> list[dict[str, Any]]:
if not isinstance(value, list):
return []
return [dict(item) for item in value if isinstance(item, dict)]
@staticmethod
def _normalize_persona_dialogs_payload(value: Any) -> list[str]:
if not isinstance(value, list):
return []
return [str(item) for item in value if isinstance(item, str)]
@staticmethod
def _optional_int(value: Any) -> int | None:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
@@ -1,33 +0,0 @@
from .conversation import ConversationCapabilityMixin
from .db import DBCapabilityMixin
from .http import HttpCapabilityMixin
from .kb import KnowledgeBaseCapabilityMixin
from .llm import LLMCapabilityMixin
from .mcp import McpCapabilityMixin
from .memory import MemoryCapabilityMixin
from .message_history import MessageHistoryCapabilityMixin
from .metadata import MetadataCapabilityMixin
from .persona import PersonaCapabilityMixin
from .platform import PlatformCapabilityMixin
from .provider import ProviderCapabilityMixin
from .session import SessionCapabilityMixin
from .skill import SkillCapabilityMixin
from .system import SystemCapabilityMixin
__all__ = [
"ConversationCapabilityMixin",
"DBCapabilityMixin",
"HttpCapabilityMixin",
"KnowledgeBaseCapabilityMixin",
"LLMCapabilityMixin",
"McpCapabilityMixin",
"MemoryCapabilityMixin",
"MessageHistoryCapabilityMixin",
"MetadataCapabilityMixin",
"PersonaCapabilityMixin",
"PlatformCapabilityMixin",
"ProviderCapabilityMixin",
"SessionCapabilityMixin",
"SkillCapabilityMixin",
"SystemCapabilityMixin",
]
@@ -1,261 +0,0 @@
from __future__ import annotations
import uuid
from typing import Any
from ....errors import AstrBotError
from ..bridge_base import CapabilityRouterBridgeBase
class ConversationCapabilityMixin(CapabilityRouterBridgeBase):
async def _conversation_new(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
if not session:
raise AstrBotError.invalid_input("conversation.new requires session")
raw_conversation = payload.get("conversation")
if raw_conversation is None:
raw_conversation = {}
if not isinstance(raw_conversation, dict):
raise AstrBotError.invalid_input(
"conversation.new requires conversation object"
)
conversation_id = uuid.uuid4().hex
now = self._now_iso()
record = {
"conversation_id": conversation_id,
"session": session,
"platform_id": (
str(raw_conversation.get("platform_id"))
if raw_conversation.get("platform_id") is not None
else self._session_platform_id(session)
),
"history": self._normalize_history_payload(raw_conversation.get("history")),
"title": (
str(raw_conversation.get("title"))
if raw_conversation.get("title") is not None
else None
),
"persona_id": (
str(raw_conversation.get("persona_id"))
if raw_conversation.get("persona_id") is not None
else None
),
"created_at": now,
"updated_at": now,
"token_usage": None,
}
self._conversation_store[conversation_id] = record
self._session_current_conversation_ids[session] = conversation_id
return {"conversation_id": conversation_id}
async def _conversation_switch(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = str(payload.get("conversation_id", "")).strip()
record = self._conversation_store.get(conversation_id)
if record is None or str(record.get("session", "")) != session:
raise AstrBotError.invalid_input(
"conversation.switch requires a conversation in the same session"
)
self._session_current_conversation_ids[session] = conversation_id
return {}
async def _conversation_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = payload.get("conversation_id")
normalized_conversation_id = (
str(conversation_id).strip() if conversation_id is not None else ""
)
if not normalized_conversation_id:
normalized_conversation_id = self._session_current_conversation_ids.get(
session, ""
)
if not normalized_conversation_id:
return {}
record = self._conversation_store.get(normalized_conversation_id)
if record is None:
return {}
if str(record.get("session", "")) != session:
raise AstrBotError.invalid_input(
"conversation.delete requires a conversation in the same session"
)
del self._conversation_store[normalized_conversation_id]
current_conversation_id = self._session_current_conversation_ids.get(session)
if current_conversation_id == normalized_conversation_id:
replacement = next(
(
conversation_id
for conversation_id, item in self._conversation_store.items()
if str(item.get("session", "")) == session
),
None,
)
if replacement is None:
self._session_current_conversation_ids.pop(session, None)
else:
self._session_current_conversation_ids[session] = replacement
return {}
async def _conversation_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = str(payload.get("conversation_id", "")).strip()
record = self._conversation_store.get(conversation_id)
if record is None and bool(payload.get("create_if_not_exists", False)):
created = await self._conversation_new(
_request_id,
{"session": session, "conversation": {}},
_token,
)
record = self._conversation_store.get(
str(created.get("conversation_id", "")).strip()
)
if record is None:
return {"conversation": None}
if str(record.get("session", "")) != session:
return {"conversation": None}
return {"conversation": dict(record)}
async def _conversation_get_current(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = self._session_current_conversation_ids.get(session, "")
if not conversation_id and bool(payload.get("create_if_not_exists", False)):
created = await self._conversation_new(
_request_id,
{"session": session, "conversation": {}},
_token,
)
conversation_id = str(created.get("conversation_id", "")).strip()
if not conversation_id:
return {"conversation": None}
record = self._conversation_store.get(conversation_id)
if record is None or str(record.get("session", "")) != session:
return {"conversation": None}
return {"conversation": dict(record)}
async def _conversation_list(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = payload.get("session")
platform_id = payload.get("platform_id")
conversations = []
for conversation_id in sorted(self._conversation_store.keys()):
item = self._conversation_store[conversation_id]
if session is not None and str(item.get("session", "")) != str(session):
continue
if platform_id is not None and str(item.get("platform_id", "")) != str(
platform_id
):
continue
conversations.append(dict(item))
return {"conversations": conversations}
async def _conversation_update(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = payload.get("conversation_id")
normalized_conversation_id = (
str(conversation_id).strip() if conversation_id is not None else ""
)
if not normalized_conversation_id:
normalized_conversation_id = self._session_current_conversation_ids.get(
session, ""
)
if not normalized_conversation_id:
return {}
record = self._conversation_store.get(normalized_conversation_id)
if record is None:
return {}
if str(record.get("session", "")) != session:
raise AstrBotError.invalid_input(
"conversation.update requires a conversation in the same session"
)
raw_conversation = payload.get("conversation")
if not isinstance(raw_conversation, dict):
raw_conversation = {}
if "history" in raw_conversation:
history = raw_conversation.get("history")
record["history"] = (
self._normalize_history_payload(history) if history is not None else []
)
if "title" in raw_conversation:
title = raw_conversation.get("title")
record["title"] = str(title) if title is not None else None
if "persona_id" in raw_conversation:
persona_id = raw_conversation.get("persona_id")
record["persona_id"] = str(persona_id) if persona_id is not None else None
if "token_usage" in raw_conversation:
token_usage = raw_conversation.get("token_usage")
record["token_usage"] = (
int(token_usage) if token_usage is not None else None
)
record["updated_at"] = self._now_iso()
return {}
async def _conversation_unset_persona(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session = str(payload.get("session", "")).strip()
conversation_id = payload.get("conversation_id")
normalized_conversation_id = (
str(conversation_id).strip() if conversation_id is not None else ""
)
if not normalized_conversation_id:
normalized_conversation_id = self._session_current_conversation_ids.get(
session, ""
)
if not normalized_conversation_id:
return {}
record = self._conversation_store.get(normalized_conversation_id)
if record is None:
return {}
if str(record.get("session", "")) != session:
raise AstrBotError.invalid_input(
"conversation.unset_persona requires a conversation in the same session"
)
record["persona_id"] = None
record["updated_at"] = self._now_iso()
return {}
def _register_conversation_capabilities(self) -> None:
self.register(
self._builtin_descriptor("conversation.new", "新建对话"),
call_handler=self._conversation_new,
)
self.register(
self._builtin_descriptor("conversation.switch", "切换对话"),
call_handler=self._conversation_switch,
)
self.register(
self._builtin_descriptor("conversation.delete", "删除对话"),
call_handler=self._conversation_delete,
)
self.register(
self._builtin_descriptor("conversation.get", "获取对话"),
call_handler=self._conversation_get,
)
self.register(
self._builtin_descriptor("conversation.get_current", "获取当前对话"),
call_handler=self._conversation_get_current,
)
self.register(
self._builtin_descriptor("conversation.list", "列出对话"),
call_handler=self._conversation_list,
)
self.register(
self._builtin_descriptor("conversation.update", "更新对话"),
call_handler=self._conversation_update,
)
self.register(
self._builtin_descriptor("conversation.unset_persona", "清空对话人格"),
call_handler=self._conversation_unset_persona,
)
@@ -1,170 +0,0 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from typing import Any
from ....errors import AstrBotError
from ..._streaming import StreamExecution
from ..bridge_base import CapabilityRouterBridgeBase
class DBCapabilityMixin(CapabilityRouterBridgeBase):
def _db_scoped_key(self, plugin_id: str, key: str) -> str:
"""将用户提供的 key 加上插件命名空间前缀,防止跨插件越权访问。"""
return f"{plugin_id}:{key}"
def _db_strip_scope(self, plugin_id: str, scoped_key: str) -> str:
"""去掉命名空间前缀,返回插件视角的原始 key。"""
prefix = f"{plugin_id}:"
return (
scoped_key[len(prefix) :] if scoped_key.startswith(prefix) else scoped_key
)
def _db_public_event(
self, plugin_id: str, raw_event: dict[str, Any]
) -> dict[str, Any]:
"""将内部事件转换回插件可见的 key 视图。"""
event = dict(raw_event)
key = event.get("key")
if isinstance(key, str):
event["key"] = self._db_strip_scope(plugin_id, key)
return event
async def _db_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("db.get")
key = self._db_scoped_key(plugin_id, str(payload.get("key", "")))
return {"value": self.db_store.get(key)}
async def _db_set(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("db.set")
key = self._db_scoped_key(plugin_id, 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]:
plugin_id = self._require_caller_plugin_id("db.delete")
key = self._db_scoped_key(plugin_id, 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]:
plugin_id = self._require_caller_plugin_id("db.list")
ns_prefix = f"{plugin_id}:"
# 只列出属于当前插件命名空间的 key,并去掉命名空间前缀返回给插件
user_prefix = payload.get("prefix")
all_keys = sorted(
key for key in self.db_store.keys() if key.startswith(ns_prefix)
)
stripped = [self._db_strip_scope(plugin_id, k) for k in all_keys]
if isinstance(user_prefix, str):
stripped = [k for k in stripped if k.startswith(user_prefix)]
return {"keys": stripped}
async def _db_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("db.get_many")
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
items = [
{
"key": str(k),
"value": self.db_store.get(self._db_scoped_key(plugin_id, str(k))),
}
for k in keys_payload
]
return {"items": items}
async def _db_set_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("db.set_many")
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 = self._db_scoped_key(plugin_id, 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:
plugin_id = self._require_caller_plugin_id("db.watch")
prefix = payload.get("prefix")
prefix_value: str | None
if isinstance(prefix, str):
# 将用户传入的前缀也加上命名空间,只监听本插件的 key 变更
prefix_value = self._db_scoped_key(plugin_id, prefix)
elif prefix is None:
# 无前缀时默认监听整个命名空间
prefix_value = f"{plugin_id}:"
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 self._db_public_event(plugin_id, 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(
self._builtin_descriptor("db.set", "写入 KV"), call_handler=self._db_set
)
self.register(
self._builtin_descriptor("db.delete", "删除 KV"),
call_handler=self._db_delete,
)
self.register(
self._builtin_descriptor("db.list", "列出 KV"), call_handler=self._db_list
)
self.register(
self._builtin_descriptor("db.get_many", "批量读取 KV"),
call_handler=self._db_get_many,
)
self.register(
self._builtin_descriptor("db.set_many", "批量写入 KV"),
call_handler=self._db_set_many,
)
self.register(
self._builtin_descriptor(
"db.watch",
"订阅 KV 变更",
supports_stream=True,
cancelable=True,
),
stream_handler=self._db_watch,
)
@@ -1,120 +0,0 @@
from __future__ import annotations
import re
from typing import Any
from ....errors import AstrBotError
from ..bridge_base import CapabilityRouterBridgeBase
# 路由只允许字母、数字、/, -, _, . 以及路径参数 {param},且必须以 / 开头
# 禁止 .. 防止路径遍历,禁止连续斜杠
_ROUTE_SAFE_RE = re.compile(r"^(/[\w\-._{}]*)+$")
def _validate_route(route: str, capability_name: str) -> None:
"""校验 HTTP 路由路径格式,阻止路径遍历和非法字符。"""
if ".." in route:
raise AstrBotError.invalid_input(f"{capability_name}: 路由路径不允许包含 '..'")
if not _ROUTE_SAFE_RE.match(route):
raise AstrBotError.invalid_input(
f"{capability_name}: 路由路径格式非法,只允许字母/数字/-/_/./{{param}} 段,"
"且必须以 / 开头,如 /foo/bar"
)
class HttpCapabilityMixin(CapabilityRouterBridgeBase):
async def _http_register_api(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
methods_payload = payload.get("methods")
if not isinstance(methods_payload, list) or not all(
isinstance(item, str) for item in methods_payload
):
raise AstrBotError.invalid_input(
"http.register_api 的 methods 必须是 string 数组"
)
route = str(payload.get("route", "")).strip()
handler_capability = str(payload.get("handler_capability", "")).strip()
if not route or not handler_capability:
raise AstrBotError.invalid_input(
"http.register_api 需要 route 和 handler_capability"
)
_validate_route(route, "http.register_api")
plugin_name = self._require_caller_plugin_id("http.register_api")
methods = sorted({method.upper() for method in methods_payload if method})
entry: dict[str, Any] = {
"route": route,
"methods": methods,
"handler_capability": handler_capability,
"description": str(payload.get("description", "")),
"plugin_id": plugin_name,
}
self.http_api_store = [
item
for item in self.http_api_store
if not (
item.get("route") == route
and item.get("plugin_id") == entry["plugin_id"]
and item.get("methods") == methods
)
]
self.http_api_store.append(entry)
return {}
async def _http_unregister_api(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
route = str(payload.get("route", "")).strip()
methods_payload = payload.get("methods")
if not isinstance(methods_payload, list) or not all(
isinstance(item, str) for item in methods_payload
):
raise AstrBotError.invalid_input(
"http.unregister_api 的 methods 必须是 string 数组"
)
plugin_name = self._require_caller_plugin_id("http.unregister_api")
methods = {method.upper() for method in methods_payload if method}
updated: list[dict[str, Any]] = []
for entry in self.http_api_store:
if entry.get("route") != route:
updated.append(entry)
continue
if entry.get("plugin_id") != plugin_name:
updated.append(entry)
continue
if not methods:
# `HTTPClient.unregister_api(methods=None)` 会归一化为空列表,
# 公开语义就是“移除当前插件在该 route 下注册的全部方法”。
continue
remaining_methods = [
method for method in entry.get("methods", []) if method not in methods
]
if remaining_methods:
updated.append({**entry, "methods": remaining_methods})
self.http_api_store = updated
return {}
async def _http_list_apis(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_name = self._require_caller_plugin_id("http.list_apis")
apis = [
dict(entry)
for entry in self.http_api_store
if entry.get("plugin_id") == plugin_name
]
return {"apis": apis}
def _register_http_capabilities(self) -> None:
self.register(
self._builtin_descriptor("http.register_api", "注册 HTTP 路由"),
call_handler=self._http_register_api,
)
self.register(
self._builtin_descriptor("http.unregister_api", "注销 HTTP 路由"),
call_handler=self._http_unregister_api,
)
self.register(
self._builtin_descriptor("http.list_apis", "列出 HTTP 路由"),
call_handler=self._http_list_apis,
)
@@ -1,427 +0,0 @@
from __future__ import annotations
import math
import uuid
from pathlib import Path
from typing import Any
from ....errors import AstrBotError
from ..bridge_base import CapabilityRouterBridgeBase
def _term_set(text: str) -> set[str]:
normalized = " ".join(str(text).strip().casefold().split())
compact = normalized.replace(" ", "")
if not normalized:
return set()
terms = {item for item in normalized.split(" ") if item}
if compact:
terms.add(compact)
if len(compact) > 1:
terms.update(
compact[index : index + 2] for index in range(len(compact) - 1)
)
return terms
class KnowledgeBaseCapabilityMixin(CapabilityRouterBridgeBase):
def _kb_documents(self, kb_id: str) -> dict[str, dict[str, Any]]:
return self._kb_document_store.setdefault(kb_id, {})
def _refresh_mock_kb_stats(self, kb_id: str) -> None:
kb = self._kb_store.get(kb_id)
if not isinstance(kb, dict):
return
documents = self._kb_documents(kb_id)
kb["doc_count"] = len(documents)
kb["chunk_count"] = sum(
int(document.get("chunk_count", 0) or 0) for document in documents.values()
)
kb["updated_at"] = self._now_iso()
def _resolve_mock_kb_ids(self, payload: dict[str, Any]) -> list[str]:
kb_ids = [
str(item).strip() for item in payload.get("kb_ids", []) if str(item).strip()
]
if kb_ids:
return [kb_id for kb_id in kb_ids if kb_id in self._kb_store]
kb_names = [
str(item).strip()
for item in payload.get("kb_names", [])
if str(item).strip()
]
if not kb_names:
return []
name_set = set(kb_names)
return [
kb_id
for kb_id, kb in self._kb_store.items()
if str(kb.get("kb_name", "")).strip() in name_set
]
@staticmethod
def _score_mock_document(query: str, content: str) -> float:
query_terms = _term_set(query)
content_terms = _term_set(content)
if not query_terms or not content_terms:
return 0.0
overlap = len(query_terms & content_terms)
if overlap <= 0:
return 0.0
score = overlap / len(query_terms)
if query.strip().casefold() in str(content).casefold():
score += 0.25
return min(score, 1.0)
@staticmethod
def _build_mock_context_text(results: list[dict[str, Any]]) -> str:
lines = ["以下是相关的知识库内容,请参考这些信息回答用户的问题:\n"]
for index, item in enumerate(results, start=1):
lines.append(f"【知识 {index}】")
lines.append(f"来源: {item['kb_name']} / {item['doc_name']}")
lines.append(f"内容: {item['content']}")
lines.append(f"相关度: {float(item['score']):.2f}")
lines.append("")
return "\n".join(lines)
async def _kb_list(
self,
_request_id: str,
_payload: dict[str, Any],
_token,
) -> dict[str, Any]:
return {
"kbs": [
dict(record)
for record in sorted(
self._kb_store.values(),
key=lambda item: str(item.get("created_at", "")),
)
]
}
async def _kb_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
record = self._kb_store.get(kb_id)
return {"kb": dict(record) if isinstance(record, dict) else None}
async def _kb_create(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
raw_kb = payload.get("kb")
if not isinstance(raw_kb, dict):
raise AstrBotError.invalid_input("kb.create requires kb object")
embedding_provider_id = str(raw_kb.get("embedding_provider_id", "")).strip()
if not embedding_provider_id:
raise AstrBotError.invalid_input("kb.create requires embedding_provider_id")
kb_id = uuid.uuid4().hex
now = self._now_iso()
record = {
"kb_id": kb_id,
"kb_name": str(raw_kb.get("kb_name", "")),
"description": (
str(raw_kb.get("description"))
if raw_kb.get("description") is not None
else None
),
"emoji": (
str(raw_kb.get("emoji")) if raw_kb.get("emoji") is not None else None
),
"embedding_provider_id": embedding_provider_id,
"rerank_provider_id": (
str(raw_kb.get("rerank_provider_id"))
if raw_kb.get("rerank_provider_id") is not None
else None
),
"chunk_size": self._optional_int(raw_kb.get("chunk_size")),
"chunk_overlap": self._optional_int(raw_kb.get("chunk_overlap")),
"top_k_dense": self._optional_int(raw_kb.get("top_k_dense")),
"top_k_sparse": self._optional_int(raw_kb.get("top_k_sparse")),
"top_m_final": self._optional_int(raw_kb.get("top_m_final")),
"doc_count": 0,
"chunk_count": 0,
"created_at": now,
"updated_at": now,
}
self._kb_store[kb_id] = record
self._kb_document_store[kb_id] = {}
return {"kb": dict(record)}
async def _kb_update(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
raw_kb = payload.get("kb")
if not isinstance(raw_kb, dict):
raise AstrBotError.invalid_input("kb.update requires kb object")
record = self._kb_store.get(kb_id)
if not isinstance(record, dict):
return {"kb": None}
for field_name in (
"kb_name",
"description",
"emoji",
"embedding_provider_id",
"rerank_provider_id",
):
if field_name in raw_kb:
value = raw_kb.get(field_name)
record[field_name] = str(value) if value is not None else None
for field_name in (
"chunk_size",
"chunk_overlap",
"top_k_dense",
"top_k_sparse",
"top_m_final",
):
if field_name in raw_kb:
record[field_name] = self._optional_int(raw_kb.get(field_name))
record["updated_at"] = self._now_iso()
return {"kb": dict(record)}
async def _kb_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
documents = self._kb_document_store.pop(kb_id, {})
for document in documents.values():
doc_id = str(document.get("doc_id", "")).strip()
if doc_id:
self._kb_document_content_store.pop(doc_id, None)
deleted = self._kb_store.pop(kb_id, None) is not None
return {"deleted": deleted}
async def _kb_retrieve(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
query = str(payload.get("query", "")).strip()
if not query:
raise AstrBotError.invalid_input("kb.retrieve requires query")
kb_ids = self._resolve_mock_kb_ids(payload)
if not kb_ids:
raise AstrBotError.invalid_input("kb.retrieve requires kb_ids or kb_names")
top_m_final = self._optional_int(payload.get("top_m_final")) or 5
results: list[dict[str, Any]] = []
for kb_id in kb_ids:
kb = self._kb_store.get(kb_id)
if not isinstance(kb, dict):
continue
for document in self._kb_documents(kb_id).values():
doc_id = str(document.get("doc_id", "")).strip()
if not doc_id:
continue
content = self._kb_document_content_store.get(doc_id, "")
score = self._score_mock_document(query, content)
if score <= 0:
continue
results.append(
{
"chunk_id": f"{doc_id}:0",
"doc_id": doc_id,
"kb_id": kb_id,
"kb_name": str(kb.get("kb_name", "")),
"doc_name": str(document.get("doc_name", "")),
"chunk_index": 0,
"content": content,
"score": score,
"char_count": len(content),
}
)
results.sort(key=lambda item: float(item["score"]), reverse=True)
results = results[:top_m_final]
if not results:
return {"result": None}
return {
"result": {
"context_text": self._build_mock_context_text(results),
"results": results,
}
}
async def _kb_document_upload(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
kb = self._kb_store.get(kb_id)
if not isinstance(kb, dict):
raise AstrBotError.invalid_input(f"Unknown knowledge base: {kb_id}")
raw_document = payload.get("document")
if not isinstance(raw_document, dict):
raise AstrBotError.invalid_input(
"kb.document.upload requires document object"
)
file_name = str(raw_document.get("file_name", "")).strip()
file_type = str(raw_document.get("file_type", "")).strip()
file_path = ""
content_text = ""
file_size = 0
text_value = raw_document.get("text")
url_value = raw_document.get("url")
file_token = str(raw_document.get("file_token", "")).strip()
if isinstance(text_value, str) and text_value.strip():
content_text = text_value
if not file_name:
file_name = "document.txt"
if not file_type:
file_type = "txt"
file_size = len(content_text.encode("utf-8"))
elif isinstance(url_value, str) and url_value.strip():
url_text = url_value.strip()
content_text = f"Imported from {url_text}"
if not file_name:
file_name = (
Path(url_text.split("?", maxsplit=1)[0]).name or "document.url"
)
if not file_type:
suffix = Path(file_name).suffix.lstrip(".")
file_type = suffix or "url"
file_path = url_text
file_size = len(content_text.encode("utf-8"))
elif file_token:
file_path = self._file_token_store.pop(file_token, "")
if not file_path:
raise AstrBotError.invalid_input(f"Unknown file token: {file_token}")
path = Path(file_path)
if not path.exists():
raise AstrBotError.invalid_input(f"File does not exist: {file_path}")
raw_bytes = path.read_bytes()
content_text = raw_bytes.decode("utf-8", errors="ignore")
if not file_name:
file_name = path.name
if not file_type:
file_type = path.suffix.lstrip(".")
if not file_type:
raise AstrBotError.invalid_input(
"kb.document.upload requires file_type when the file has no suffix"
)
file_size = len(raw_bytes)
else:
raise AstrBotError.invalid_input(
"kb.document.upload requires file_token, url, or text"
)
chunk_size = self._optional_int(raw_document.get("chunk_size"))
if chunk_size is None or chunk_size <= 0:
chunk_size = self._optional_int(kb.get("chunk_size")) or 512
chunk_count = max(1, math.ceil(max(len(content_text), 1) / chunk_size))
doc_id = uuid.uuid4().hex
now = self._now_iso()
document = {
"doc_id": doc_id,
"kb_id": kb_id,
"doc_name": file_name,
"file_type": file_type,
"file_size": file_size,
"file_path": file_path,
"chunk_count": chunk_count,
"media_count": 0,
"created_at": now,
"updated_at": now,
}
self._kb_documents(kb_id)[doc_id] = document
self._kb_document_content_store[doc_id] = content_text
self._refresh_mock_kb_stats(kb_id)
return {"document": dict(document)}
async def _kb_document_list(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
offset = max(self._optional_int(payload.get("offset")) or 0, 0)
limit = max(self._optional_int(payload.get("limit")) or 100, 0)
documents = list(self._kb_documents(kb_id).values())
documents.sort(key=lambda item: str(item.get("created_at", "")))
return {
"documents": [dict(item) for item in documents[offset : offset + limit]]
}
async def _kb_document_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
doc_id = str(payload.get("doc_id", "")).strip()
document = self._kb_documents(kb_id).get(doc_id)
return {"document": dict(document) if isinstance(document, dict) else None}
async def _kb_document_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
doc_id = str(payload.get("doc_id", "")).strip()
deleted = self._kb_documents(kb_id).pop(doc_id, None) is not None
if deleted:
self._kb_document_content_store.pop(doc_id, None)
self._refresh_mock_kb_stats(kb_id)
return {"deleted": deleted}
async def _kb_document_refresh(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
kb_id = str(payload.get("kb_id", "")).strip()
doc_id = str(payload.get("doc_id", "")).strip()
document = self._kb_documents(kb_id).get(doc_id)
if not isinstance(document, dict):
return {"document": None}
kb = self._kb_store.get(kb_id, {})
chunk_size = self._optional_int(kb.get("chunk_size")) or 512
content_text = self._kb_document_content_store.get(doc_id, "")
document["chunk_count"] = max(
1, math.ceil(max(len(content_text), 1) / chunk_size)
)
document["updated_at"] = self._now_iso()
self._refresh_mock_kb_stats(kb_id)
return {"document": dict(document)}
def _register_kb_capabilities(self) -> None:
self.register(
self._builtin_descriptor("kb.list", "列出知识库"),
call_handler=self._kb_list,
)
self.register(
self._builtin_descriptor("kb.get", "获取知识库"),
call_handler=self._kb_get,
)
self.register(
self._builtin_descriptor("kb.create", "创建知识库"),
call_handler=self._kb_create,
)
self.register(
self._builtin_descriptor("kb.update", "更新知识库"),
call_handler=self._kb_update,
)
self.register(
self._builtin_descriptor("kb.delete", "删除知识库"),
call_handler=self._kb_delete,
)
self.register(
self._builtin_descriptor("kb.retrieve", "检索知识库"),
call_handler=self._kb_retrieve,
)
self.register(
self._builtin_descriptor("kb.document.upload", "上传知识库文档"),
call_handler=self._kb_document_upload,
)
self.register(
self._builtin_descriptor("kb.document.list", "列出知识库文档"),
call_handler=self._kb_document_list,
)
self.register(
self._builtin_descriptor("kb.document.get", "获取知识库文档"),
call_handler=self._kb_document_get,
)
self.register(
self._builtin_descriptor("kb.document.delete", "删除知识库文档"),
call_handler=self._kb_document_delete,
)
self.register(
self._builtin_descriptor("kb.document.refresh", "刷新知识库文档"),
call_handler=self._kb_document_refresh,
)
@@ -1,64 +0,0 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from typing import Any
from ..bridge_base import CapabilityRouterBridgeBase
class LLMCapabilityMixin(CapabilityRouterBridgeBase):
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 _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 _llm_stream(
self,
_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}
def _register_llm_capabilities(self) -> None:
self.register(
self._builtin_descriptor("llm.chat", "发送对话请求,返回文本"),
call_handler=self._llm_chat,
)
self.register(
self._builtin_descriptor("llm.chat_raw", "发送对话请求,返回完整响应"),
call_handler=self._llm_chat_raw,
)
self.register(
self._builtin_descriptor(
"llm.stream_chat",
"流式对话",
supports_stream=True,
cancelable=True,
),
stream_handler=self._llm_stream,
finalize=lambda chunks: {
"text": "".join(item.get("text", "") for item in chunks)
},
)

Some files were not shown because too many files have changed in this diff Show More