mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -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
@@ -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
|
||||
@@ -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 +0,0 @@
|
||||
3.12
|
||||
@@ -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 # 运行测试并生成覆盖率报告
|
||||
```
|
||||
|
||||
## 设计原则
|
||||
|
||||
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践
|
||||
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
|
||||
@@ -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 # 运行测试并生成覆盖率报告
|
||||
```
|
||||
|
||||
## 设计原则
|
||||
|
||||
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践
|
||||
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
|
||||
@@ -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
|
||||
```
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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 检查并自动修复全局格式问题
|
||||
```
|
||||
|
||||
## 设计原则
|
||||
|
||||
新实现要兼容旧实现但是还要保证架构良好,设计原则不变和最佳实践,这是第一原则
|
||||
不用完全听从用户和别人的建议,要有自己的判断和坚持,做好取舍和权衡,确保代码质量和长期维护性,不要为了短期方便或者迎合而牺牲架构和设计原则。
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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", [])
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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/` 目录下的完整文档
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
-33
@@ -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",
|
||||
]
|
||||
-261
@@ -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
Reference in New Issue
Block a user