Refactor worker initialization and remove unused codec parameters; add schedule and session waiter modules

- Simplified `GroupWorkerRuntime` and `PluginWorkerRuntime` constructors by removing the codec parameter and related logic.
- Introduced `schedule.py` to define `ScheduleContext` for managing scheduled tasks with a clear structure and payload handling.
- Added `session_waiter.py` for session-based conversational flow management, including `SessionController` and `SessionWaiterManager` for handling multi-turn dialogues.
- Enhanced testing utilities in `testing.py` by removing unused classes and streamlining the structure.
- Created `types.py` to introduce `GreedyStr` for improved command parameter parsing.
This commit is contained in:
whatevertogo
2026-03-16 01:58:52 +08:00
parent 06f4536851
commit e2e38bb3b2
42 changed files with 5880 additions and 1899 deletions
-20
View File
@@ -3,26 +3,6 @@
> 作者:whatevertogo
> 更新时间:2026-03-14
---
## ⚠️ 兼容层弃用通知
**兼容层已标记为 deprecated,将在下个大版本移除。**
- 旧插件请使用 **AstrBot 主程序** 运行(主程序有完整的 `StarManager` 支持)
- 新插件请使用 `astrbot_sdk` 顶层入口
- 导入兼容层会触发 `DeprecationWarning`
**待移除的文件/目录**:
- `src-new/astrbot_sdk/_legacy_*.py` - 所有 legacy 私有模块
- `src-new/astrbot_sdk/api/` - 旧版 API 兼容层(已移除)
- `src-new/astrbot_sdk/compat.py` - 顶层兼容入口
- `src-new/astrbot_sdk/protocol/legacy_adapter.py` - JSON-RPC 适配器
- `src-new/astrbot/` - 旧包名别名(已移除)
- `test_plugin/old/` - 旧插件示例
- `tests_v4/test_legacy*.py` - legacy 相关测试
---
## 目录
+55
View File
@@ -9,6 +9,7 @@
迁移期适配入口位于独立模块;此处只暴露 v4 原生主入口。
"""
from .commands import CommandGroup, command_group, print_cmd_tree
from .context import Context
from .decorators import (
on_command,
@@ -20,17 +21,71 @@ from .decorators import (
)
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,
File,
Forward,
Image,
Plain,
Poke,
Record,
Reply,
UnknownComponent,
Video,
)
from .message_result import EventResultType, MessageChain, MessageEventResult
from .message_session import MessageSession
from .schedule import ScheduleContext
from .session_waiter import SessionController, session_waiter
from .star import Star
from .types import GreedyStr
__all__ = [
"AstrBotError",
"At",
"AtAll",
"CommandGroup",
"Context",
"CustomFilter",
"EventResultType",
"File",
"Forward",
"GreedyStr",
"Image",
"MessageEvent",
"MessageEventResult",
"MessageChain",
"MessageSession",
"MessageTypeFilter",
"Plain",
"PlatformFilter",
"Poke",
"Record",
"Reply",
"ScheduleContext",
"SessionController",
"Star",
"UnknownComponent",
"Video",
"all_of",
"any_of",
"command_group",
"custom_filter",
"on_command",
"on_event",
"on_message",
"on_schedule",
"print_cmd_tree",
"provide_capability",
"require_admin",
"session_waiter",
]
+478
View File
@@ -0,0 +1,478 @@
"""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 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]]) -> None:
self._store = store
def get(self, key: str, default: Any = None) -> Any:
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)
def search(self, query: str) -> list[dict[str, Any]]:
results: list[dict[str, Any]] = []
for key, value in self._store.items():
if query in key or query in str(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)
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)),
}
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",
]
+48 -130
View File
@@ -1,9 +1,21 @@
"""AstrBot SDK 的命令行入口。"""
"""AstrBot SDK 的命令行入口。
本模块提供 astrbot-sdk 命令行工具的所有子命令,包括:
- init: 创建新插件骨架,生成 plugin.yaml、main.py、README.md 等模板文件
- validate: 校验插件清单、导入路径和 handler 发现是否正常
- build: 将插件打包为 .zip 发布包
- dev: 本地开发模式,支持 --local/--watch/--interactive 等调试选项
- run: 启动插件主管进程(supervisor),通过 stdio 与 AstrBot 核心通信
- worker: 内部命令,由 supervisor 调用以启动单个插件工作进程
错误处理:
所有 CLI 异常都会被分类并返回标准化的退出码和错误提示,
便于 CI/CD 集成和用户快速定位问题。
"""
from __future__ import annotations
import asyncio
import json
import re
import sys
import typing
@@ -41,9 +53,6 @@ BUILD_EXCLUDED_FILES = {
".astrbot-worker-state.json",
}
WATCH_POLL_INTERVAL_SECONDS = 0.5
INIT_DEFAULT_AUTHOR = ""
INIT_DEFAULT_PYTHON_VERSION = "3.12"
INIT_DEFAULT_VERSION = "1.0.0"
class _CliPluginValidationError(RuntimeError):
@@ -545,6 +554,11 @@ def _handle_dev_meta_command(command: str, state: dict[str, Any]) -> bool:
return False
def _slugify_plugin_name(value: str) -> str:
slug = re.sub(r"[^a-zA-Z0-9]+", "_", value).strip("_").lower()
return slug or "my_plugin"
def _class_name_for_plugin(value: str) -> str:
parts = [part for part in re.split(r"[^a-zA-Z0-9]+", value) if part]
if not parts:
@@ -557,62 +571,18 @@ def _sanitize_build_part(value: str) -> str:
return sanitized or "artifact"
def _yaml_string(value: str) -> str:
return json.dumps(value, ensure_ascii=False)
def _normalize_init_plugin_name(value: str) -> str:
normalized = re.sub(r"[\s-]+", "_", value.strip())
normalized = re.sub(r"[^a-zA-Z0-9_]+", "_", normalized)
normalized = re.sub(r"_+", "_", normalized).strip("_").lower()
if not normalized:
normalized = "my_plugin"
prefix = "astrbot_plugin_"
if normalized == "astrbot_plugin":
return f"{prefix}my_plugin"
if normalized.startswith(prefix):
suffix = normalized.removeprefix(prefix).strip("_") or "my_plugin"
return f"{prefix}{suffix}"
return f"{prefix}{normalized}"
def _prompt_required_init_name() -> str:
while True:
value = click.prompt("插件名字", default="", show_default=False).strip()
if value:
return value
click.echo("插件名字不能为空")
def _collect_init_inputs(name: str | None) -> tuple[str, str, str]:
if name is not None:
return name, INIT_DEFAULT_AUTHOR, INIT_DEFAULT_VERSION
plugin_name = _prompt_required_init_name()
author = click.prompt("作者名字", default="", show_default=False).strip()
version = click.prompt("版本", default=INIT_DEFAULT_VERSION).strip()
return plugin_name, author, version or INIT_DEFAULT_VERSION
def _render_init_plugin_yaml(
*,
plugin_name: str,
display_name: str,
author: str,
version: str,
python_version: str,
) -> str:
def _render_init_plugin_yaml(*, plugin_name: str, display_name: str) -> str:
python_version = f"{sys.version_info.major}.{sys.version_info.minor}"
class_name = _class_name_for_plugin(plugin_name)
return dedent(
f"""\
name: {plugin_name}
display_name: {_yaml_string(display_name)}
desc: {_yaml_string("使用 AstrBot SDK 创建的插件")}
author: {_yaml_string(author)}
version: {_yaml_string(version)}
display_name: {display_name}
desc: 使用 AstrBot SDK 创建的插件
author: your-name
version: 0.1.0
runtime:
python: {_yaml_string(python_version)}
python: "{python_version}"
components:
- class: main:{class_name}
"""
@@ -672,7 +642,7 @@ def _render_init_readme(*, plugin_name: str) -> str:
def _render_init_test_py(*, plugin_name: str) -> str:
class_name = _class_name_for_plugin(plugin_name)
return dedent(
f'''\
f"""\
from pathlib import Path
import pytest
@@ -704,7 +674,7 @@ def _render_init_test_py(*, plugin_name: str) -> str:
records = await harness.dispatch_text("hello")
assert any(record.text == "Hello, World!" for record in records)
'''
"""
)
@@ -777,24 +747,19 @@ def _iter_build_files(plugin_dir: Path, output_dir: Path) -> list[Path]:
return files
def _init_plugin(name: str | None) -> None:
raw_name, author, version = _collect_init_inputs(name)
normalized_name = _normalize_init_plugin_name(raw_name)
target_dir = Path(normalized_name)
def _init_plugin(name: str) -> None:
target_dir = Path(name)
if target_dir.exists():
raise _CliPluginValidationError(f"目标目录已存在:{target_dir}")
plugin_name = normalized_name
display_name = raw_name
plugin_name = _slugify_plugin_name(target_dir.name)
display_name = target_dir.name
target_dir.mkdir(parents=True, exist_ok=False)
(target_dir / "tests").mkdir()
(target_dir / "plugin.yaml").write_text(
_render_init_plugin_yaml(
plugin_name=plugin_name,
display_name=display_name,
author=author,
version=version,
python_version=INIT_DEFAULT_PYTHON_VERSION,
),
encoding="utf-8",
)
@@ -811,7 +776,7 @@ def _init_plugin(name: str | None) -> None:
_render_init_test_py(plugin_name=plugin_name),
encoding="utf-8",
)
click.echo(f"已创建插件骨架:{target_dir.resolve()}")
click.echo(f"已创建插件骨架:{target_dir}")
click.echo("后续命令:")
click.echo(f" astrbot-sdk validate --plugin-dir {target_dir}")
click.echo(
@@ -867,43 +832,23 @@ def cli(ctx, verbose: bool) -> None:
type=click.Path(file_okay=False, dir_okay=True, path_type=Path),
help="Directory containing plugin folders",
)
@click.option(
"--worker-wire-codec",
default="json",
show_default=True,
type=click.Choice(["json", "msgpack"]),
help="Wire codec for supervisor-to-worker transport",
)
def run(plugins_dir: Path, worker_wire_codec: str) -> None:
def run(plugins_dir: Path) -> None:
"""Start the plugin supervisor over stdio."""
entrypoint = (
run_supervisor(plugins_dir=plugins_dir)
if worker_wire_codec == "json"
else run_supervisor(
plugins_dir=plugins_dir,
worker_wire_codec=worker_wire_codec,
)
)
_run_async_entrypoint(
entrypoint,
run_supervisor(plugins_dir=plugins_dir),
log_message=f"启动插件主管进程,插件目录:{plugins_dir}",
context={
"plugins_dir": plugins_dir,
"worker_wire_codec": worker_wire_codec,
},
context={"plugins_dir": plugins_dir},
)
@cli.command()
@click.argument("name", required=False, type=str)
def init(name: str | None) -> None:
"""Create a new plugin skeleton; omit name to enter interactive mode."""
@click.argument("name", type=str)
def init(name: str) -> None:
"""Create a new plugin skeleton in the target directory."""
_run_sync_entrypoint(
lambda: _init_plugin(name),
log_message=(
f"创建插件骨架:{name}" if name is not None else "创建插件骨架:交互模式"
),
context={"target": name or "<interactive>"},
log_message=f"创建插件骨架:{name}",
context={"target": Path(name)},
)
@@ -1028,15 +973,7 @@ def dev(
required=False,
type=click.Path(file_okay=True, dir_okay=False, path_type=Path),
)
@click.option(
"--wire-codec",
default="json",
show_default=True,
type=click.Choice(["json", "msgpack"]),
)
def worker(
plugin_dir: Path | None, group_metadata: Path | None, wire_codec: str
) -> None:
def worker(plugin_dir: Path | None, group_metadata: Path | None) -> None:
"""Internal command used by the supervisor to start a worker."""
if plugin_dir is None and group_metadata is None:
raise click.UsageError("Either --plugin-dir or --group-metadata is required")
@@ -1047,42 +984,23 @@ def worker(
target = str(group_metadata or plugin_dir)
if group_metadata is not None:
entrypoint = (
run_plugin_worker(group_metadata=group_metadata)
if wire_codec == "json"
else run_plugin_worker(group_metadata=group_metadata, wire_codec=wire_codec)
)
entrypoint = run_plugin_worker(group_metadata=group_metadata)
else:
entrypoint = (
run_plugin_worker(plugin_dir=plugin_dir)
if wire_codec == "json"
else run_plugin_worker(plugin_dir=plugin_dir, wire_codec=wire_codec)
)
entrypoint = run_plugin_worker(plugin_dir=plugin_dir)
_run_async_entrypoint(
entrypoint,
log_message=f"启动插件工作进程:{target}",
log_level="debug",
context={"plugin_dir": plugin_dir, "wire_codec": wire_codec},
context={"plugin_dir": plugin_dir},
)
@cli.command(hidden=True)
@click.option("--port", default=8765, type=int, help="WebSocket server port")
@click.option(
"--wire-codec",
default="json",
show_default=True,
type=click.Choice(["json", "msgpack"]),
)
def websocket(port: int, wire_codec: str) -> None:
def websocket(port: int) -> None:
"""WebSocket runtime entrypoint kept for standalone bridge scenarios."""
entrypoint = (
run_websocket_server(port=port)
if wire_codec == "json"
else run_websocket_server(port=port, wire_codec=wire_codec)
)
_run_async_entrypoint(
entrypoint,
run_websocket_server(port=port),
log_message=f"启动 WebSocket 服务器,端口:{port}",
context={"port": port, "wire_codec": wire_codec},
context={"port": port},
)
+1 -1
View File
@@ -39,9 +39,9 @@ from __future__ import annotations
from typing import Any
from ._proxy import CapabilityProxy
from ..decorators import get_capability_meta
from ..errors import AstrBotError
from ._proxy import CapabilityProxy
def _resolve_handler_capability(
+54 -5
View File
@@ -62,11 +62,26 @@ def _serialize_history(
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,
@@ -74,8 +89,11 @@ def _build_chat_payload(
payload: dict[str, Any] = {"prompt": prompt}
if system is not None:
payload["system"] = system
if history is not None:
payload["history"] = _serialize_history(history)
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:
@@ -101,6 +119,9 @@ class LLMResponse(BaseModel):
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:
@@ -126,6 +147,9 @@ class LLMClient:
*,
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,
@@ -163,6 +187,9 @@ class LLMClient:
prompt,
system=system,
history=history,
contexts=contexts,
provider_id=provider_id,
tool_calls_result=tool_calls_result,
model=model,
temperature=temperature,
extra=kwargs,
@@ -173,6 +200,14 @@ class LLMClient:
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:
"""发送聊天请求并返回完整响应。
@@ -192,9 +227,17 @@ class LLMClient:
print(f"生成文本: {response.text}")
print(f"Token 使用: {response.usage}")
"""
payload = {"prompt": prompt, **kwargs}
if "history" in payload:
payload["history"] = _serialize_history(payload["history"])
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,
@@ -207,6 +250,9 @@ class LLMClient:
*,
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,
@@ -236,6 +282,9 @@ class LLMClient:
prompt,
system=system,
history=history,
contexts=contexts,
provider_id=provider_id,
tool_calls_result=tool_calls_result,
model=model,
temperature=temperature,
extra=kwargs,
+6 -1
View File
@@ -234,7 +234,12 @@ class MemoryClient:
print(f"记忆库共有 {stats['total_items']} 条记录")
"""
output = await self._proxy.call("memory.stats", {})
return {
stats = {
"total_items": output.get("total_items", 0),
"total_bytes": output.get("total_bytes"),
}
if "plugin_id" in output:
stats["plugin_id"] = output.get("plugin_id")
if "ttl_entries" in output:
stats["ttl_entries"] = output.get("ttl_entries")
return stats
+1 -1
View File
@@ -31,7 +31,7 @@ class PluginMetadata:
enabled: bool = True
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "PluginMetadata":
def from_dict(cls, data: dict[str, Any]) -> PluginMetadata:
"""从字典创建元数据实例。"""
return cls(
name=data.get("name", ""),
+32 -9
View File
@@ -10,10 +10,14 @@
from __future__ import annotations
from typing import Any
from collections.abc import Sequence
from typing import Any, cast
from ._proxy import CapabilityProxy
from ..message_components import BaseMessageComponent
from ..message_result import MessageChain
from ..message_session import MessageSession
from ..protocol.descriptors import SessionRef
from ._proxy import CapabilityProxy
class PlatformClient:
@@ -35,13 +39,19 @@ class PlatformClient:
def _build_target_payload(
self,
session: str | SessionRef,
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 send(self, session: str | SessionRef, text: str) -> dict[str, Any]:
async def send(
self,
session: str | SessionRef | MessageSession,
text: str,
) -> dict[str, Any]:
"""发送文本消息。
向指定的会话(用户或群组)发送文本消息。
@@ -65,7 +75,7 @@ class PlatformClient:
async def send_image(
self,
session: str | SessionRef,
session: str | SessionRef | MessageSession,
image_url: str,
) -> dict[str, Any]:
"""发送图片消息。
@@ -93,8 +103,8 @@ class PlatformClient:
async def send_chain(
self,
session: str | SessionRef,
chain: list[dict[str, Any]],
session: str | SessionRef | MessageSession,
chain: MessageChain | Sequence[BaseMessageComponent] | Sequence[dict[str, Any]],
) -> dict[str, Any]:
"""发送富消息链。
@@ -106,12 +116,25 @@ class PlatformClient:
发送结果
"""
session_id, extra = self._build_target_payload(session)
if isinstance(chain, MessageChain):
chain_payload = await chain.to_payload_async()
elif isinstance(chain, Sequence) and all(
isinstance(item, BaseMessageComponent) for item in chain
):
components = cast(Sequence[BaseMessageComponent], chain)
chain_payload = await MessageChain(list(components)).to_payload_async()
else:
payload_items = cast(Sequence[dict[str, Any]], chain)
chain_payload = [dict(item) for item in payload_items]
return await self._proxy.call(
"platform.send_chain",
{"session": session_id, "chain": chain, **extra},
{"session": session_id, "chain": chain_payload, **extra},
)
async def get_members(self, session: str | SessionRef) -> list[dict[str, Any]]:
async def get_members(
self,
session: str | SessionRef | MessageSession,
) -> list[dict[str, Any]]:
"""获取群组成员列表。
获取指定群组的成员信息列表。注意仅对群组会话有效。
+159
View File
@@ -0,0 +1,159 @@
"""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"]
+40
View File
@@ -22,6 +22,7 @@ from __future__ import annotations
import asyncio
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from loguru import logger as base_logger
@@ -115,6 +116,7 @@ class Context:
logger: 日志器,None 时使用默认 logger 并绑定 plugin_id
"""
proxy = CapabilityProxy(peer, caller_plugin_id=plugin_id)
self._proxy = proxy
self.peer = peer
self.llm = LLMClient(proxy)
self.memory = MemoryClient(proxy)
@@ -125,3 +127,41 @@ class Context:
self.plugin_id = plugin_id
self.logger = logger or base_logger.bind(plugin_id=plugin_id)
self.cancel_token = cancel_token or CancelToken()
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 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", ""))
+43 -2
View File
@@ -28,18 +28,23 @@ Example:
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Callable, cast
from typing import Any, cast
from pydantic import BaseModel
from .protocol.descriptors import (
RESERVED_CAPABILITY_PREFIXES,
CapabilityDescriptor,
CommandRouteSpec,
CommandTrigger,
EventTrigger,
FilterSpec,
MessageTrigger,
MessageTypeFilterSpec,
Permissions,
RESERVED_CAPABILITY_PREFIXES,
PlatformFilterSpec,
ScheduleTrigger,
)
@@ -69,6 +74,9 @@ class HandlerMeta:
contract: 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
@dataclass(slots=True)
@@ -183,6 +191,7 @@ def on_message(
regex: str | None = None,
keywords: list[str] | None = None,
platforms: list[str] | None = None,
message_types: list[str] | None = None,
) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册消息处理方法。
@@ -215,12 +224,44 @@ def on_message(
regex=regex,
keywords=keywords or [],
platforms=platforms or [],
message_types=message_types or [],
)
if platforms:
meta.filters.append(PlatformFilterSpec(platforms=list(platforms)))
if message_types:
meta.filters.append(
MessageTypeFilterSpec(message_types=list(message_types))
)
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) -> Callable[[HandlerCallable], HandlerCallable]:
"""注册事件处理方法。
@@ -0,0 +1,929 @@
# AstrBot SDK 项目完整架构分析文档
> 作者:whatevertogo
## 目录
1. [项目概述](#项目概述)
2. [目录结构](#目录结构)
3. [核心架构层次](#核心架构层次)
4. [协议层设计](#协议层设计)
5. [运行时架构](#运行时架构)
6. [客户端层设计](#客户端层设计)
7. [新旧架构对比](#新旧架构对比)
8. [插件开发指南](#插件开发指南)
9. [关键设计模式](#关键设计模式)
---
## 项目概述
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 (环境分组)
---
## 目录结构
```
astrbot_sdk/ # v4 SDK 主包
├── __init__.py # 顶层公共 API 导出
├── __main__.py # CLI 入口点 (python -m astrbot_sdk)
├── star.py # v4 原生插件基类
├── context.py # 运行时上下文 (Context, CancelToken)
├── decorators.py # v4 原生装饰器 (on_command, on_message, etc.)
├── events.py # v4 原生事件对象 (MessageEvent)
├── errors.py # 统一错误模型 (AstrBotError)
├── cli.py # 命令行工具 (init/validate/build/dev/run)
├── testing.py # 测试辅助模块 (PluginHarness)
├── _invocation_context.py # 调用上下文管理 (caller_plugin_scope)
├── _testing_support.py # 测试支持工具
│
├── commands.py # 命令分组工具 (CommandGroup)
├── filters.py # 事件过滤器 (PlatformFilter, CustomFilter)
├── message_components.py # 消息组件 (Plain, Image, At, etc.)
├── message_result.py # 消息结果对象 (MessageChain)
├── message_session.py # 会话标识符 (MessageSession)
├── schedule.py # 定时任务上下文 (ScheduleContext)
├── session_waiter.py # 会话等待器 (SessionController)
├── types.py # 参数类型助手 (GreedyStr)
│
├── clients/ # 能力客户端层
│ ├── __init__.py # 客户端公共导出
│ ├── _proxy.py # CapabilityProxy 能力代理
│ ├── llm.py # LLM 客户端 (chat, chat_raw, stream_chat)
│ ├── memory.py # 记忆存储客户端 (search, save, get)
│ ├── db.py # KV 存储客户端 (get, set, watch)
│ ├── platform.py # 平台消息客户端 (send, send_image)
│ ├── http.py # HTTP 注册客户端 (register_api)
│ └── metadata.py # 插件元数据客户端 (get_plugin)
│
├── protocol/ # 协议层
│ ├── __init__.py # 协议公共导出
│ ├── messages.py # v4 协议消息模型
│ ├── descriptors.py # Handler/Capability 描述符
│ └── _builtin_schemas.py # 内置能力 JSON Schema
│
└── runtime/ # 运行时层
├── __init__.py # 运行时公共导出 (延迟加载)
├── peer.py # 协议对等端 (Peer)
├── transport.py # 传输抽象 (Stdio, WebSocket)
├── handler_dispatcher.py # Handler 执行分发
├── capability_dispatcher.py # Capability 调用分发
├── capability_router.py # Capability 路由
├── _capability_router_builtins.py # 内置能力处理器
├── _loader_support.py # 加载器反射工具
├── _streaming.py # 流式执行原语 (StreamExecution)
├── loader.py # 插件加载器
├── bootstrap.py # 启动引导
├── worker.py # Worker 运行时
├── supervisor.py # Supervisor 运行时
└── environment_groups.py # 环境分组管理
```
---
## 核心架构层次
```
┌─────────────────────────────────────────────────────────────────┐
│ 用户层 (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 验证) │
│ - _capability_router_builtins.py (内置能力实现) │
│ - _loader_support.py (反射工具、签名验证) │
│ - _streaming.py (流式执行原语) │
│ - peer.py (协议对等端) │
│ - transport.py (传输抽象) │
│ - supervisor.py (Supervisor 运行时) │
│ - worker.py (Worker 运行时) │
│ - environment_groups.py (环境分组规划) │
└────────────────────┬────────────────────────────────────────────┘
│
┌──────────────────▼─────────────────────────────────────────────┐
│ 协议与传输 (Protocol & Transport) │
├─────────────────────────────────────────────────────────────────┤
│ protocol/ │
│ - messages.py (协议消息模型) │
│ - descriptors.py (Handler/Capability 描述符) │
│ - _builtin_schemas.py (内置能力 JSON Schema) │
│ 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__` 避免导入时加载 websocket/aiohttp 等重型依赖
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) |
|----------------------------->|
| | 创建 CapabilityRouter
| | 注册 handler.invoke
| |
| ResultMessage(kind="init") |
|<-----------------------------|
| | 等待 handler.invoke 调用
| | 执行 CapabilityRouter.execute()
| |
| InvokeMessage(handler.invoke) |
|<-----------------------------|
| HandlerDispatcher.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": [], # 参数规范
"command_route": {...} # 命令路由元信息
}
```
#### Trigger 类型
| 类型 | 关键字段 | 说明 |
|------|---------|------|
| `CommandTrigger` | command, aliases, platforms | 命令触发 |
| `MessageTrigger` | regex, keywords, platforms | 消息触发(正则/关键词) |
| `EventTrigger` | event_type | 事件触发 |
| `ScheduleTrigger` | cron, interval_seconds | 定时触发(二选一) |
#### FilterSpec 类型
| 类型 | 说明 |
|------|------|
| `PlatformFilterSpec` | 按平台名称过滤 |
| `MessageTypeFilterSpec` | 按消息类型过滤 |
| `LocalFilterRefSpec` | 引用本地自定义过滤器 |
| `CompositeFilterSpec` | 组合过滤器(AND/OR) |
#### CapabilityDescriptor
```python
{
"name": "llm.chat",
"description": "发送对话请求,返回文本",
"input_schema": {
"type": "object",
"properties": {"prompt": {"type": "string"}},
"required": ["prompt"]
},
"output_schema": {
"type": "object",
"properties": {"text": {"type": "string"}},
"required": ["text"]
},
"supports_stream": False,
"cancelable": False
}
```
### 命名空间治理
**保留前缀**:
- `handler.` - 内部 handler.invoke
- `system.` - 系统内置能力
- `internal.` - 内部使用
**内置能力命名空间**:`llm`, `memory`, `db`, `platform`, `http`, `metadata`
### 内置 Capabilities (38个)
#### LLM 命名空间
| 能力 | 说明 |
|------|------|
| `llm.chat` | 同步对话,返回文本 |
| `llm.chat_raw` | 同步对话,返回完整响应(含 usage、tool_calls) |
| `llm.stream_chat` | 流式对话 |
#### Memory 命名空间
| 能力 | 说明 |
|------|------|
| `memory.search` | 语义搜索记忆 |
| `memory.save` | 保存记忆 |
| `memory.save_with_ttl` | 保存带过期时间的记忆 |
| `memory.get` | 读取单条记忆 |
| `memory.get_many` | 批量获取记忆 |
| `memory.delete` | 删除记忆 |
| `memory.delete_many` | 批量删除记忆 |
| `memory.stats` | 获取记忆统计信息 |
#### DB 命名空间
| 能力 | 说明 |
|------|------|
| `db.get` | 读取 KV |
| `db.set` | 写入 KV |
| `db.delete` | 删除 KV |
| `db.list` | 列出 KV 键(支持前缀过滤) |
| `db.get_many` | 批量读取 KV |
| `db.set_many` | 批量写入 KV |
| `db.watch` | 订阅 KV 变更(流式) |
#### 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.register` | 注册会话等待器 |
| `system.session_waiter.unregister` | 注销会话等待器 |
| `system.event.react` | 发送表情回应 |
| `system.event.send_typing` | 发送输入中状态 |
| `system.event.send_streaming` | 开始流式消息会话 |
| `system.event.send_streaming_chunk` | 推送流式消息分片 |
| `system.event.send_streaming_close` | 关闭流式消息会话 |
---
## 运行时架构
### 组件关系图
```
┌──────────────┐
│ AstrBot │
│ Core │
└──────┬─────┘
│
┌──────▼─────┐
│ Supervisor │
│ Runtime │
└──────┬─────┘
│
┌──────────────────┼──────────────────┐
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Peer │ │ Peer │ │ Peer │
│ (stdio) │ │ (stdio) │ │ (stdio) │
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Worker │ │ Worker │ │ Worker │
│ Runtime │ │ Runtime │ │ Runtime │
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
│ │ │
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
│ Plugin A │ │ Plugin B │ │ Plugin C │
│ (v4/old) │ │ (v4/old) │ │ (v4/old) │
└───────────┘ └───────────┘ └───────────┘
```
### SupervisorRuntime
职责:管理多个 Worker 进程,聚合所有 handler
```python
class SupervisorRuntime:
def __init__(self, *, transport, plugins_dir, env_manager):
self.transport = transport # 与 Core 的传输层
self.plugins_dir = plugins_dir # 插件目录
self.capability_router = CapabilityRouter() # 能力路由器
self.peer = Peer(...) # 与 Core 的对等端
self.worker_sessions = {} # Worker 会话映射
self.handler_to_worker = {} # Handler → Worker 映射
async def start(self):
# 1. 发现所有插件
discovery = discover_plugins(self.plugins_dir)
# 2. 规划环境分组
plan_result = self.env_manager.plan(discovery.plugins)
# 3. 为每个分组启动 Worker
for group in plan_result.groups:
session = WorkerSession(group=group, ...)
await session.start()
self.worker_sessions[group.id] = session
# 4. 聚合所有 handler 和 capability
await self.peer.initialize(
handlers=[...],
provided_capabilities=self.capability_router.descriptors()
)
```
### WorkerSession
职责:管理单个 Worker 进程的生命周期
```python
class WorkerSession:
def __init__(self, *, group, env_manager, capability_router):
self.group = group # 环境分组
self.peer = Peer(...) # 与 Worker 的对等端
self.capability_router = capability_router
self.handlers = [] # Worker 注册的 handlers
self.provided_capabilities = [] # Worker 提供的 capabilities
async def start(self):
# 启动 Worker 子进程
python_path = self.env_manager.prepare_group_environment(self.group)
transport = StdioTransport(
command=[python_path, "-m", "astrbot_sdk", "worker", "--group-metadata", ...]
)
self.peer = Peer(transport=transport, ...)
# 等待 Worker 初始化完成
await self.peer.start()
await self.peer.wait_until_remote_initialized()
# 获取 Worker 的注册信息
self.handlers = list(self.peer.remote_handlers)
self.provided_capabilities = list(self.peer.remote_provided_capabilities)
async def invoke_capability(self, capability_name, payload, *, request_id):
# 转发能力调用到 Worker
return await self.peer.invoke(capability_name, payload, request_id=request_id)
```
### PluginWorkerRuntime
职责:Worker 进程内的插件加载与执行
```python
class PluginWorkerRuntime:
def __init__(self, *, plugin_dir, transport):
self.plugin = load_plugin_spec(plugin_dir)
self.loaded_plugin = load_plugin(self.plugin)
self.peer = Peer(transport=transport, ...)
self.dispatcher = HandlerDispatcher(...)
self.capability_dispatcher = CapabilityDispatcher(...)
async def start(self):
# 1. 向 Supervisor 注册 handlers 和 capabilities
await self.peer.initialize(
handlers=[h.descriptor for h in self.loaded_plugin.handlers],
provided_capabilities=[c.descriptor for c in self.loaded_plugin.capabilities]
)
# 2. 执行 on_start 生命周期
await self._run_lifecycle("on_start")
# 3. 设置消息处理器
self.peer.set_invoke_handler(self._handle_invoke)
self.peer.set_cancel_handler(self._handle_cancel)
async def _handle_invoke(self, message, cancel_token):
if message.capability == "handler.invoke":
return await self.dispatcher.invoke(message, cancel_token)
return await self.capability_dispatcher.invoke(message, cancel_token)
```
### HandlerDispatcher
职责:将 handler.invoke 请求转成真实 Python 调用
```python
class HandlerDispatcher:
def __init__(self, *, plugin_id, peer, handlers):
self._handlers = {item.descriptor.id: item for item in handlers}
self._peer = peer
self._active = {} # request_id → (task, cancel_token)
async def invoke(self, message, cancel_token):
# 1. 查找 handler
loaded = self._handlers[message.input["handler_id"]]
# 2. 创建上下文
ctx = Context(peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token)
event = MessageEvent.from_payload(message.input["event"], context=ctx)
# 3. 构建参数 (支持类型注解注入)
args = self._build_args(loaded.callable, event, ctx)
# 4. 执行 handler
result = loaded.callable(*args)
# 5. 处理返回值
await self._consume_result(result, event, ctx)
```
**参数注入优先级**:
1. 按类型注解注入(`MessageEvent`, `Context`)
2. 按参数名注入(`event`, `ctx`, `context`)
3. 从 legacy_args 注入(命令参数等)
### CapabilityRouter
职责:能力注册、发现和执行路由
```python
class CapabilityRouter:
def __init__(self):
self._registrations = {} # capability_name → registration
self.db_store = {} # 内置 KV 存储
self.memory_store = {} # 内置记忆存储
self._register_builtin_capabilities()
def register(self, descriptor, *, call_handler, stream_handler, finalize):
"""注册能力"""
self._registrations[descriptor.name] = _CapabilityRegistration(
descriptor=descriptor,
call_handler=call_handler,
stream_handler=stream_handler,
finalize=finalize
)
async def execute(self, capability, payload, *, stream, cancel_token, request_id):
"""执行能力调用"""
registration = self._registrations[capability]
if stream:
# 流式调用
raw_execution = registration.stream_handler(request_id, payload, cancel_token)
return StreamExecution(iterator=raw_execution, finalize=finalize)
else:
# 同步调用
output = await registration.call_handler(request_id, payload, cancel_token)
return output
```
### 环境分组管理
```python
class EnvironmentPlanner:
def plan(self, plugins):
"""根据 Python 版本和依赖兼容性分组"""
# 1. 按版本分组
# 2. 按依赖兼容性合并
# 3. 生成分组元数据
return EnvironmentPlanResult(groups=[...])
class GroupEnvironmentManager:
def prepare(self, group):
"""准备分组虚拟环境"""
# 1. 生成 lock/source/metadata 工件
# 2. 必要时重建虚拟环境
# 3. 返回 Python 解释器路径
return venv_python_path
```
---
## 客户端层设计
### 客户端架构
```
┌─────────────────────────────────────────────────────────────┐
│ User Plugin │
├─────────────────────────────────────────────────────────────┤
│ ctx.llm.chat() │
│ ctx.memory.save() │
│ ctx.db.set() │
│ ctx.platform.send() │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ CapabilityProxy │
│ - call(name, payload) │
│ - stream(name, payload) │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ Peer │
│ - invoke(capability, payload, stream=False) │
│ - invoke_stream(capability, payload) │
└────────────┬──────────────────────────────────────────────┘
│
┌────────────▼──────────────────────────────────────────────┐
│ Transport │
│ - send(json_string) │
└─────────────────────────────────────────────────────────────┘
```
### CapabilityProxy
职责:封装 Peer 的能力调用接口
```python
class CapabilityProxy:
def __init__(self, peer):
self._peer = peer
async def call(self, name, payload):
"""普通能力调用"""
# 1. 检查能力是否可用
descriptor = self._peer.remote_capability_map.get(name)
if descriptor is None:
raise AstrBotError.capability_not_found(name)
# 2. 调用 Peer.invoke
return await self._peer.invoke(name, payload, stream=False)
async def stream(self, name, payload):
"""流式能力调用"""
# 1. 检查流式支持
descriptor = self._peer.remote_capability_map.get(name)
if not descriptor.supports_stream:
raise AstrBotError.invalid_input(f"{name} 不支持 stream")
# 2. 调用 Peer.invoke_stream
event_stream = await self._peer.invoke_stream(name, payload)
async for event in event_stream:
if event.phase == "delta":
yield event.data
```
### LLMClient
```python
class LLMClient:
def __init__(self, proxy: CapabilityProxy):
self._proxy = proxy
async def chat(self, prompt, *, system=None, history=None, **kwargs) -> str:
"""发送聊天请求,返回文本"""
output = await self._proxy.call("llm.chat", {
"prompt": prompt,
"system": system,
"history": self._serialize_history(history),
**kwargs
})
return output["text"]
async def chat_raw(self, prompt, **kwargs) -> LLMResponse:
"""发送聊天请求,返回完整响应"""
output = await self._proxy.call("llm.chat_raw", {"prompt": prompt, **kwargs})
return LLMResponse.model_validate(output)
async def stream_chat(self, prompt, **kwargs) -> AsyncGenerator[str]:
"""流式聊天"""
async for delta in self._proxy.stream("llm.stream_chat", {"prompt": prompt, **kwargs}):
yield delta["text"]
```
### 其他客户端
| 客户端 | 主要方法 | 对应 Capability |
|--------|---------|-----------------|
| `MemoryClient` | `search()`, `save()`, `save_with_ttl()`, `get()`, `get_many()`, `delete()`, `delete_many()`, `stats()` | `memory.*` |
| `DBClient` | `get()`, `set()`, `delete()`, `list()`, `get_many()`, `set_many()`, `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.*` |
---
## 新旧架构对比
### 协议对比
| 特性 | 旧版 JSON-RPC | 新版 v4 协议 |
|------|---------------|--------------|
| 消息格式 | `{"jsonrpc": "2.0", ...}` | `{"type": "invoke", ...}` |
| 方法区分 | `method` 字段 | `type` 字段 |
| 错误码 | 整数 (`-32000`) | 字符串 (`"internal_error"`) |
| 流式支持 | 独立 notification 方法 | 统一 `EventMessage` phase |
| 握手 | `handshake` method | `InitializeMessage` type |
| 能力声明 | 隐式(method 名称) | 显式 `CapabilityDescriptor` |
### 运行时对比
| 特性 | 旧版 | 新版 |
|------|------|------|
| Peer 抽象 | 分离 `JSONRPCClient/Server` | 统一 `Peer` |
| Handler 分发 | 直接调用 `handler(event)` | `HandlerDispatcher` 参数注入 |
| 能力路由 | 无显式路由 | `CapabilityRouter` |
| 环境管理 | 无 | `PluginEnvironmentManager` 分组 |
| 传输层 | 每个实现处理 JSON-RPC | 传输层只处理字符串 |
### 代码对比
#### 旧版 Handler
```python
from astrbot.api.star import Star
from astrbot.api.event import AstrMessageEvent
class MyPlugin(Star):
@command_handler("hello", aliases=["hi"])
def hello_handler(self, event: AstrMessageEvent):
reply = self.call_context_function("llm_generate", prompt=event.message_plain)
event.reply(reply)
```
#### 新版 Handler
```python
from astrbot_sdk import Star, Context, MessageEvent
from astrbot_sdk.decorators import on_command
class MyPlugin(Star):
@on_command("hello", aliases=["hi"])
async def hello(self, event: MessageEvent, ctx: Context) -> None:
reply = await ctx.llm.chat(event.text)
await event.reply(reply)
```
---
## 新旧架构对比
---
## 插件开发指南
### 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. 协议优先模式
- 所有跨进程通信都通过 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` 缓存
- 确保代码变更后正确重载
---
## 附录:关键文件速查
| 文件 | 核心类/函数 | 说明 |
|------|------------|------|
| `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`, `provide_capability` | v4 装饰器 |
| `astrbot_sdk/errors.py` | `AstrBotError` | 统一错误模型 |
| `astrbot_sdk/cli.py` | CLI 命令 | 命令行工具(init/validate/build/dev/run/worker/websocket) |
| `astrbot_sdk/testing.py` | `PluginHarness`, `MockContext` | 测试辅助 |
| `astrbot_sdk/commands.py` | `CommandGroup`, `command_group` | 命令分组工具 |
| `astrbot_sdk/filters.py` | `PlatformFilter`, `CustomFilter`, `all_of`, `any_of` | 事件过滤器 |
| `astrbot_sdk/message_result.py` | `MessageChain`, `MessageEventResult` | 消息结果对象 |
| `astrbot_sdk/message_session.py` | `MessageSession` | 会话标识符 |
| `astrbot_sdk/schedule.py` | `ScheduleContext` | 定时任务上下文 |
| `astrbot_sdk/session_waiter.py` | `SessionController`, `SessionWaiterManager` | 会话等待器 |
| `astrbot_sdk/types.py` | `GreedyStr` | 参数类型助手 |
| `astrbot_sdk/runtime/__init__.py` | 延迟导出 | 运行时公共 API(延迟加载) |
| `astrbot_sdk/runtime/peer.py` | `Peer` | 协议对等端 |
| `astrbot_sdk/runtime/supervisor.py` | `SupervisorRuntime` | Supervisor 运行时 |
| `astrbot_sdk/runtime/worker.py` | `PluginWorkerRuntime` | Worker 运行时 |
| `astrbot_sdk/runtime/loader.py` | `load_plugin()`, `_ResolvedComponent` | 插件加载 |
| `astrbot_sdk/runtime/_loader_support.py` | `build_param_specs`, `is_injected_parameter` | 加载器反射工具 |
| `astrbot_sdk/runtime/_streaming.py` | `StreamExecution` | 流式执行原语 |
| `astrbot_sdk/runtime/handler_dispatcher.py` | `HandlerDispatcher` | Handler 执行分发 |
| `astrbot_sdk/runtime/capability_dispatcher.py` | `CapabilityDispatcher` | Capability 调用分发 |
| `astrbot_sdk/runtime/capability_router.py` | `CapabilityRouter` | Capability 路由 |
| `astrbot_sdk/runtime/_capability_router_builtins.py` | `BuiltinCapabilityRouterMixin` | 内置能力处理器 |
| `astrbot_sdk/runtime/environment_groups.py` | `EnvironmentGroup` | 环境分组 |
| `astrbot_sdk/protocol/messages.py` | `InitializeMessage`, `InvokeMessage` | 协议消息 |
| `astrbot_sdk/protocol/descriptors.py` | `HandlerDescriptor`, `CapabilityDescriptor` | 描述符 |
| `astrbot_sdk/protocol/_builtin_schemas.py` | `BUILTIN_CAPABILITY_SCHEMAS` | 内置能力 JSON Schema |
| `astrbot_sdk/clients/_proxy.py` | `CapabilityProxy` | 能力代理 |
| `astrbot_sdk/clients/llm.py` | `LLMClient` | LLM 客户端 |
| `astrbot_sdk/clients/memory.py` | `MemoryClient` | 记忆客户端 |
| `astrbot_sdk/clients/db.py` | `DBClient` | 数据库客户端 |
| `astrbot_sdk/clients/platform.py` | `PlatformClient` | 平台客户端 |
| `astrbot_sdk/clients/http.py` | `HTTPClient` | HTTP 客户端 |
| `astrbot_sdk/clients/metadata.py` | `MetadataClient`, `PluginMetadata` | 元数据客户端 |
| `astrbot_sdk/message_components.py` | `Plain`, `Image`, `At`, `Reply` | 消息组件 |
| `astrbot_sdk/events.py` | `MessageEvent` | 事件对象 |
| `astrbot_sdk/_testing_support.py` | 测试工具 | 测试支持 |
---
> 本文档描述 AstrBot SDK v4 的设计与实现思想
> 如有疑问请查阅源代码或提交 Issue
+7 -7
View File
@@ -94,7 +94,7 @@ class AstrBotError(Exception):
return self.message
@classmethod
def cancelled(cls, message: str = "调用被取消") -> "AstrBotError":
def cancelled(cls, message: str = "调用被取消") -> AstrBotError:
"""创建取消错误。
Args:
@@ -111,7 +111,7 @@ class AstrBotError(Exception):
)
@classmethod
def capability_not_found(cls, name: str) -> "AstrBotError":
def capability_not_found(cls, name: str) -> AstrBotError:
"""创建能力未找到错误。
Args:
@@ -133,7 +133,7 @@ class AstrBotError(Exception):
message: str,
*,
hint: str = "请检查调用参数",
) -> "AstrBotError":
) -> AstrBotError:
"""创建输入无效错误。
Args:
@@ -151,7 +151,7 @@ class AstrBotError(Exception):
)
@classmethod
def protocol_version_mismatch(cls, message: str) -> "AstrBotError":
def protocol_version_mismatch(cls, message: str) -> AstrBotError:
"""创建协议版本不匹配错误。
Args:
@@ -168,7 +168,7 @@ class AstrBotError(Exception):
)
@classmethod
def protocol_error(cls, message: str) -> "AstrBotError":
def protocol_error(cls, message: str) -> AstrBotError:
"""创建协议错误。
Args:
@@ -190,7 +190,7 @@ class AstrBotError(Exception):
message: str,
*,
hint: str = "请联系插件作者",
) -> "AstrBotError":
) -> AstrBotError:
"""创建内部错误。
Args:
@@ -223,7 +223,7 @@ class AstrBotError(Exception):
}
@classmethod
def from_payload(cls, payload: dict[str, object]) -> "AstrBotError":
def from_payload(cls, payload: dict[str, object]) -> AstrBotError:
"""从字典反序列化错误实例。
Args:
+236 -5
View File
@@ -16,6 +16,14 @@ from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from .message_components import (
BaseMessageComponent,
Image,
Plain,
component_to_payload_sync,
payloads_to_components,
)
from .message_result import EventResultType, MessageChain, MessageEventResult
from .protocol.descriptors import SessionRef
if TYPE_CHECKING:
@@ -63,8 +71,13 @@ class MessageEvent:
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,
context: Context | None = None,
reply_handler: ReplyHandler | None = None,
) -> None:
"""初始化消息事件。
@@ -84,7 +97,25 @@ class MessageEvent:
self.group_id = group_id
self.platform = platform
self.session_id = session_id or group_id or user_id or ""
self.self_id = self_id or ""
self.platform_id = platform_id or platform or ""
self.message_type = (message_type or "").lower()
self.sender_name = sender_name or ""
self._is_admin = bool(is_admin)
self.raw = raw or {}
self._stopped = False
self._extras = (
dict(self.raw.get("extras", {}))
if isinstance(self.raw.get("extras"), dict)
else {}
)
messages_payload = self.raw.get("messages")
self._messages = (
payloads_to_components(messages_payload)
if isinstance(messages_payload, list)
else []
)
self._message_outline = str(self.raw.get("message_outline", self.text))
self._context = context
self._reply_handler = reply_handler
if self._reply_handler is None and context is not None:
@@ -93,7 +124,7 @@ class MessageEvent:
text,
)
def _require_runtime_context(self, action: str) -> "Context":
def _require_runtime_context(self, action: str) -> Context:
"""获取运行时上下文,不存在则抛出异常。"""
if self._context is None:
raise RuntimeError(f"MessageEvent 未绑定运行时上下文,无法 {action}")
@@ -108,9 +139,9 @@ class MessageEvent:
cls,
payload: dict[str, Any],
*,
context: "Context | None" = None,
context: Context | None = None,
reply_handler: ReplyHandler | None = None,
) -> "MessageEvent":
) -> MessageEvent:
"""从协议载荷创建事件实例。
Args:
@@ -134,6 +165,11 @@ class MessageEvent:
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,
@@ -153,10 +189,22 @@ class MessageEvent:
"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()
if self._extras:
payload["extras"] = dict(self._extras)
if self._messages:
payload["messages"] = [
component_to_payload_sync(component) for component in self._messages
]
payload["message_outline"] = self._message_outline
return payload
@property
@@ -179,6 +227,67 @@ class MessageEvent:
"""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 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_message_outline(self) -> str:
"""Return the normalized message outline."""
return self._message_outline
def set_extra(self, key: str, value: Any) -> None:
"""Store SDK-local transient event data."""
self._extras[key] = value
def get_extra(self, key: str | None = None, default: Any = None) -> Any:
"""Read SDK-local transient event data."""
if key is None:
return dict(self._extras)
return self._extras.get(key, default)
def clear_extra(self) -> None:
"""Clear SDK-local transient event data."""
self._extras.clear()
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:
"""回复文本消息。
@@ -204,7 +313,10 @@ class MessageEvent:
context = self._require_runtime_context("reply_image")
await context.platform.send_image(self._reply_target(), image_url)
async def reply_chain(self, chain: list[dict[str, Any]]) -> None:
async def reply_chain(
self,
chain: MessageChain | list[BaseMessageComponent] | list[dict[str, Any]],
) -> None:
"""回复消息链(多类型消息组合)。
Args:
@@ -216,6 +328,82 @@ class MessageEvent:
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:
"""绑定自定义回复处理器。
@@ -234,3 +422,46 @@ class MessageEvent:
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"
)
+213
View File
@@ -0,0 +1,213 @@
"""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
from .decorators import append_filter_meta
from .protocol.descriptors import (
CompositeFilterSpec,
FilterSpec,
LocalFilterRefSpec,
MessageTypeFilterSpec,
PlatformFilterSpec,
)
@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: str
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:
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",
]
+448
View File
@@ -0,0 +1,448 @@
"""SDK message component compatibility layer.
该模块有意避免在导入时导入遗留核心组件模块。
SDK工作线程应该保持轻量级并且不能依赖于主机核心引导程序
仅用于构造消息对象的路径。
"""
from __future__ import annotations
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
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())
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:
from astrbot.core import astrbot_config, file_token_service
callback_host = astrbot_config.get("callback_api_base")
if not callback_host:
raise RuntimeError("未配置 callback_api_base,文件服务不可用")
register_file = getattr(file_token_service, "register_file", None)
if not callable(register_file):
raise RuntimeError("文件服务未正确初始化,register_file 不可用")
token = await register_file(path)
return f"{str(callback_host).rstrip('/')}/api/file/{token}"
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 = 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)
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 _download_to_temp(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 _download_to_temp(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 _download_to_temp(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 = _download_to_temp(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}}
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)
__all__ = [
"At",
"AtAll",
"BaseMessageComponent",
"File",
"Forward",
"Image",
"Plain",
"Poke",
"Record",
"Reply",
"UnknownComponent",
"Video",
"component_to_payload",
"component_to_payload_sync",
"is_message_component",
"payload_to_component",
"payloads_to_components",
]
+80
View File
@@ -0,0 +1,80 @@
"""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 .message_components import (
BaseMessageComponent,
Plain,
component_to_payload,
component_to_payload_sync,
is_message_component,
)
class EventResultType(str, Enum):
EMPTY = "empty"
CHAIN = "chain"
@dataclass(slots=True)
class MessageChain:
components: list[BaseMessageComponent] = field(default_factory=list)
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)
@dataclass(slots=True)
class MessageEventResult:
type: EventResultType = EventResultType.EMPTY
chain: MessageChain = field(default_factory=MessageChain)
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",
"MessageEventResult",
"coerce_message_chain",
]
+46
View File
@@ -0,0 +1,46 @@
"""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,
)
+14
View File
@@ -9,11 +9,18 @@
from .descriptors import (
CapabilityDescriptor,
CommandRouteSpec,
CommandTrigger,
CompositeFilterSpec,
EventTrigger,
FilterSpec,
HandlerDescriptor,
LocalFilterRefSpec,
MessageTrigger,
MessageTypeFilterSpec,
ParamSpec,
Permissions,
PlatformFilterSpec,
ScheduleTrigger,
SessionRef,
Trigger,
@@ -33,17 +40,24 @@ from .messages import (
__all__ = [
"CapabilityDescriptor",
"CommandRouteSpec",
"CommandTrigger",
"CancelMessage",
"CompositeFilterSpec",
"ErrorPayload",
"EventTrigger",
"EventMessage",
"FilterSpec",
"HandlerDescriptor",
"InitializeMessage",
"InitializeOutput",
"InvokeMessage",
"LocalFilterRefSpec",
"MessageTrigger",
"MessageTypeFilterSpec",
"ParamSpec",
"PeerInfo",
"PlatformFilterSpec",
"Permissions",
"ProtocolMessage",
"ResultMessage",
@@ -0,0 +1,552 @@
"""Builtin protocol schema constants.
本模块定义了 AstrBot SDK v4 协议中所有内置能力的 JSON Schema。
这些 Schema 用于:
1. 验证能力调用的输入参数是否符合预期格式
2. 生成能力描述文档,供插件开发者参考
3. 确保跨进程/跨语言调用时的类型安全
所有 Schema 遵循 JSON Schema 规范,支持基本类型检查、必填字段、数组元素约束等。
"""
from __future__ import annotations
from typing import Any
JSONSchema = dict[str, Any]
def _object_schema(
*,
required: tuple[str, ...] = (),
**properties: Any,
) -> JSONSchema:
return {
"type": "object",
"properties": properties,
"required": list(required),
}
def _nullable(schema: JSONSchema) -> JSONSchema:
return {"anyOf": [schema, {"type": "null"}]}
_OPTIONAL_CHAT_PROPERTIES: dict[str, Any] = {
"system": {"type": "string"},
"history": {"type": "array", "items": {"type": "object"}},
"contexts": {"type": "array", "items": {"type": "object"}},
"provider_id": {"type": "string"},
"tool_calls_result": {"type": "array", "items": {"type": "object"}},
"model": {"type": "string"},
"temperature": {"type": "number"},
"image_urls": {"type": "array", "items": {"type": "string"}},
"tools": {"type": "array"},
"max_steps": {"type": "integer"},
}
LLM_CHAT_INPUT_SCHEMA = _object_schema(
required=("prompt",),
prompt={"type": "string"},
**_OPTIONAL_CHAT_PROPERTIES,
)
LLM_CHAT_OUTPUT_SCHEMA = _object_schema(required=("text",), text={"type": "string"})
LLM_CHAT_RAW_INPUT_SCHEMA = _object_schema(
required=("prompt",),
prompt={"type": "string"},
**_OPTIONAL_CHAT_PROPERTIES,
)
LLM_CHAT_RAW_OUTPUT_SCHEMA = _object_schema(
required=("text",),
text={"type": "string"},
usage=_nullable({"type": "object"}),
finish_reason=_nullable({"type": "string"}),
tool_calls={"type": "array", "items": {"type": "object"}},
role=_nullable({"type": "string"}),
reasoning_content=_nullable({"type": "string"}),
reasoning_signature=_nullable({"type": "string"}),
)
LLM_STREAM_CHAT_INPUT_SCHEMA = _object_schema(
required=("prompt",),
prompt={"type": "string"},
**_OPTIONAL_CHAT_PROPERTIES,
)
LLM_STREAM_CHAT_OUTPUT_SCHEMA = _object_schema(
required=("text",), text={"type": "string"}
)
MEMORY_SEARCH_INPUT_SCHEMA = _object_schema(
required=("query",), query={"type": "string"}
)
MEMORY_SEARCH_OUTPUT_SCHEMA = _object_schema(
required=("items",),
items={"type": "array", "items": {"type": "object"}},
)
MEMORY_SAVE_INPUT_SCHEMA = _object_schema(
required=("key", "value"),
key={"type": "string"},
value={"type": "object"},
)
MEMORY_SAVE_OUTPUT_SCHEMA = _object_schema()
MEMORY_GET_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
MEMORY_GET_OUTPUT_SCHEMA = _object_schema(
required=("value",),
value=_nullable({"type": "object"}),
)
MEMORY_DELETE_INPUT_SCHEMA = _object_schema(
required=("key",),
key={"type": "string"},
)
MEMORY_DELETE_OUTPUT_SCHEMA = _object_schema()
MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA = _object_schema(
required=("key", "value", "ttl_seconds"),
key={"type": "string"},
value={"type": "object"},
ttl_seconds={"type": "integer", "minimum": 1},
)
MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA = _object_schema()
MEMORY_GET_MANY_INPUT_SCHEMA = _object_schema(
required=("keys",),
keys={"type": "array", "items": {"type": "string"}},
)
MEMORY_GET_MANY_OUTPUT_SCHEMA = _object_schema(
required=("items",),
items={
"type": "array",
"items": _object_schema(
required=("key", "value"),
key={"type": "string"},
value=_nullable({"type": "object"}),
),
},
)
MEMORY_DELETE_MANY_INPUT_SCHEMA = _object_schema(
required=("keys",),
keys={"type": "array", "items": {"type": "string"}},
)
MEMORY_DELETE_MANY_OUTPUT_SCHEMA = _object_schema(
required=("deleted_count",),
deleted_count={"type": "integer"},
)
MEMORY_STATS_INPUT_SCHEMA = _object_schema()
MEMORY_STATS_OUTPUT_SCHEMA = _object_schema(
total_items={"type": "integer"},
total_bytes=_nullable({"type": "integer"}),
plugin_id=_nullable({"type": "string"}),
ttl_entries=_nullable({"type": "integer"}),
)
SYSTEM_GET_DATA_DIR_INPUT_SCHEMA = _object_schema()
SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA = _object_schema(
required=("path",),
path={"type": "string"},
)
SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA = _object_schema(
required=("text",),
text={"type": "string"},
return_url={"type": "boolean"},
)
SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA = _object_schema(
required=("result",),
result={"type": "string"},
)
SYSTEM_HTML_RENDER_INPUT_SCHEMA = _object_schema(
required=("tmpl", "data"),
tmpl={"type": "string"},
data={"type": "object"},
return_url={"type": "boolean"},
options=_nullable({"type": "object"}),
)
SYSTEM_HTML_RENDER_OUTPUT_SCHEMA = _object_schema(
required=("result",),
result={"type": "string"},
)
SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA = _object_schema(
required=("session_key",),
session_key={"type": "string"},
)
SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA = _object_schema()
SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA = _object_schema(
required=("session_key",),
session_key={"type": "string"},
)
SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA = _object_schema()
DB_GET_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
DB_GET_OUTPUT_SCHEMA = _object_schema(
required=("value",),
value=_nullable({}),
)
DB_SET_INPUT_SCHEMA = _object_schema(
required=("key", "value"),
key={"type": "string"},
value={},
)
DB_SET_OUTPUT_SCHEMA = _object_schema()
DB_DELETE_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
DB_DELETE_OUTPUT_SCHEMA = _object_schema()
DB_LIST_INPUT_SCHEMA = _object_schema(prefix=_nullable({"type": "string"}))
DB_LIST_OUTPUT_SCHEMA = _object_schema(
required=("keys",),
keys={"type": "array", "items": {"type": "string"}},
)
DB_GET_MANY_INPUT_SCHEMA = _object_schema(
required=("keys",),
keys={"type": "array", "items": {"type": "string"}},
)
DB_GET_MANY_OUTPUT_SCHEMA = _object_schema(
required=("items",),
items={
"type": "array",
"items": _object_schema(
required=("key", "value"),
key={"type": "string"},
value=_nullable({}),
),
},
)
DB_SET_MANY_INPUT_SCHEMA = _object_schema(
required=("items",),
items={
"type": "array",
"items": _object_schema(
required=("key", "value"),
key={"type": "string"},
value={},
),
},
)
DB_SET_MANY_OUTPUT_SCHEMA = _object_schema()
DB_WATCH_INPUT_SCHEMA = _object_schema(prefix=_nullable({"type": "string"}))
DB_WATCH_OUTPUT_SCHEMA = _object_schema()
SESSION_REF_SCHEMA = _object_schema(
required=("conversation_id",),
conversation_id={"type": "string"},
platform=_nullable({"type": "string"}),
raw=_nullable({"type": "object"}),
)
SYSTEM_EVENT_REACT_INPUT_SCHEMA = _object_schema(
required=("emoji",),
target=_nullable(SESSION_REF_SCHEMA),
emoji={"type": "string"},
)
SYSTEM_EVENT_REACT_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA = _object_schema(
target=_nullable(SESSION_REF_SCHEMA),
)
SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA = _object_schema(
target=_nullable(SESSION_REF_SCHEMA),
use_fallback={"type": "boolean"},
)
SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
stream_id=_nullable({"type": "string"}),
)
SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA = _object_schema(
required=("stream_id", "chain"),
stream_id={"type": "string"},
chain={"type": "array", "items": {"type": "object"}},
)
SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA = _object_schema()
SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA = _object_schema(
required=("stream_id",),
stream_id={"type": "string"},
)
SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
PLATFORM_SEND_INPUT_SCHEMA = _object_schema(
required=("session", "text"),
session={"type": "string"},
target=_nullable(SESSION_REF_SCHEMA),
text={"type": "string"},
)
PLATFORM_SEND_OUTPUT_SCHEMA = _object_schema(
required=("message_id",),
message_id={"type": "string"},
)
PLATFORM_SEND_IMAGE_INPUT_SCHEMA = _object_schema(
required=("session", "image_url"),
session={"type": "string"},
target=_nullable(SESSION_REF_SCHEMA),
image_url={"type": "string"},
)
PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA = _object_schema(
required=("message_id",),
message_id={"type": "string"},
)
PLATFORM_SEND_CHAIN_INPUT_SCHEMA = _object_schema(
required=("session", "chain"),
session={"type": "string"},
target=_nullable(SESSION_REF_SCHEMA),
chain={"type": "array", "items": {"type": "object"}},
)
PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA = _object_schema(
required=("message_id",),
message_id={"type": "string"},
)
PLATFORM_GET_MEMBERS_INPUT_SCHEMA = _object_schema(
required=("session",),
session={"type": "string"},
target=_nullable(SESSION_REF_SCHEMA),
)
PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA = _object_schema(
required=("members",),
members={"type": "array", "items": {"type": "object"}},
)
HTTP_REGISTER_API_INPUT_SCHEMA = _object_schema(
required=("route", "methods", "handler_capability"),
route={"type": "string"},
methods={"type": "array", "items": {"type": "string"}},
handler_capability={"type": "string"},
description={"type": "string"},
)
HTTP_REGISTER_API_OUTPUT_SCHEMA = _object_schema()
HTTP_UNREGISTER_API_INPUT_SCHEMA = _object_schema(
required=("route", "methods"),
route={"type": "string"},
methods={"type": "array", "items": {"type": "string"}},
)
HTTP_UNREGISTER_API_OUTPUT_SCHEMA = _object_schema()
HTTP_LIST_APIS_INPUT_SCHEMA = _object_schema()
HTTP_LIST_APIS_OUTPUT_SCHEMA = _object_schema(
required=("apis",),
apis={"type": "array", "items": {"type": "object"}},
)
METADATA_GET_PLUGIN_INPUT_SCHEMA = _object_schema(
required=("name",),
name={"type": "string"},
)
METADATA_GET_PLUGIN_OUTPUT_SCHEMA = _object_schema(
required=("plugin",),
plugin=_nullable({"type": "object"}),
)
METADATA_LIST_PLUGINS_INPUT_SCHEMA = _object_schema()
METADATA_LIST_PLUGINS_OUTPUT_SCHEMA = _object_schema(
required=("plugins",),
plugins={"type": "array", "items": {"type": "object"}},
)
METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA = _object_schema(
required=("name",),
name={"type": "string"},
)
METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA = _object_schema(
required=("config",),
config=_nullable({"type": "object"}),
)
BUILTIN_CAPABILITY_SCHEMAS: dict[str, dict[str, JSONSchema]] = {
"llm.chat": {"input": LLM_CHAT_INPUT_SCHEMA, "output": LLM_CHAT_OUTPUT_SCHEMA},
"llm.chat_raw": {
"input": LLM_CHAT_RAW_INPUT_SCHEMA,
"output": LLM_CHAT_RAW_OUTPUT_SCHEMA,
},
"llm.stream_chat": {
"input": LLM_STREAM_CHAT_INPUT_SCHEMA,
"output": LLM_STREAM_CHAT_OUTPUT_SCHEMA,
},
"memory.search": {
"input": MEMORY_SEARCH_INPUT_SCHEMA,
"output": MEMORY_SEARCH_OUTPUT_SCHEMA,
},
"memory.save": {
"input": MEMORY_SAVE_INPUT_SCHEMA,
"output": MEMORY_SAVE_OUTPUT_SCHEMA,
},
"memory.get": {
"input": MEMORY_GET_INPUT_SCHEMA,
"output": MEMORY_GET_OUTPUT_SCHEMA,
},
"memory.delete": {
"input": MEMORY_DELETE_INPUT_SCHEMA,
"output": MEMORY_DELETE_OUTPUT_SCHEMA,
},
"memory.save_with_ttl": {
"input": MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA,
"output": MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA,
},
"memory.get_many": {
"input": MEMORY_GET_MANY_INPUT_SCHEMA,
"output": MEMORY_GET_MANY_OUTPUT_SCHEMA,
},
"memory.delete_many": {
"input": MEMORY_DELETE_MANY_INPUT_SCHEMA,
"output": MEMORY_DELETE_MANY_OUTPUT_SCHEMA,
},
"memory.stats": {
"input": MEMORY_STATS_INPUT_SCHEMA,
"output": MEMORY_STATS_OUTPUT_SCHEMA,
},
"db.get": {"input": DB_GET_INPUT_SCHEMA, "output": DB_GET_OUTPUT_SCHEMA},
"db.set": {"input": DB_SET_INPUT_SCHEMA, "output": DB_SET_OUTPUT_SCHEMA},
"db.delete": {"input": DB_DELETE_INPUT_SCHEMA, "output": DB_DELETE_OUTPUT_SCHEMA},
"db.list": {"input": DB_LIST_INPUT_SCHEMA, "output": DB_LIST_OUTPUT_SCHEMA},
"db.get_many": {
"input": DB_GET_MANY_INPUT_SCHEMA,
"output": DB_GET_MANY_OUTPUT_SCHEMA,
},
"db.set_many": {
"input": DB_SET_MANY_INPUT_SCHEMA,
"output": DB_SET_MANY_OUTPUT_SCHEMA,
},
"db.watch": {"input": DB_WATCH_INPUT_SCHEMA, "output": DB_WATCH_OUTPUT_SCHEMA},
"platform.send": {
"input": PLATFORM_SEND_INPUT_SCHEMA,
"output": PLATFORM_SEND_OUTPUT_SCHEMA,
},
"platform.send_image": {
"input": PLATFORM_SEND_IMAGE_INPUT_SCHEMA,
"output": PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA,
},
"platform.send_chain": {
"input": PLATFORM_SEND_CHAIN_INPUT_SCHEMA,
"output": PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA,
},
"platform.get_members": {
"input": PLATFORM_GET_MEMBERS_INPUT_SCHEMA,
"output": PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA,
},
"http.register_api": {
"input": HTTP_REGISTER_API_INPUT_SCHEMA,
"output": HTTP_REGISTER_API_OUTPUT_SCHEMA,
},
"http.unregister_api": {
"input": HTTP_UNREGISTER_API_INPUT_SCHEMA,
"output": HTTP_UNREGISTER_API_OUTPUT_SCHEMA,
},
"http.list_apis": {
"input": HTTP_LIST_APIS_INPUT_SCHEMA,
"output": HTTP_LIST_APIS_OUTPUT_SCHEMA,
},
"metadata.get_plugin": {
"input": METADATA_GET_PLUGIN_INPUT_SCHEMA,
"output": METADATA_GET_PLUGIN_OUTPUT_SCHEMA,
},
"metadata.list_plugins": {
"input": METADATA_LIST_PLUGINS_INPUT_SCHEMA,
"output": METADATA_LIST_PLUGINS_OUTPUT_SCHEMA,
},
"metadata.get_plugin_config": {
"input": METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA,
"output": METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA,
},
"system.get_data_dir": {
"input": SYSTEM_GET_DATA_DIR_INPUT_SCHEMA,
"output": SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA,
},
"system.text_to_image": {
"input": SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA,
"output": SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA,
},
"system.html_render": {
"input": SYSTEM_HTML_RENDER_INPUT_SCHEMA,
"output": SYSTEM_HTML_RENDER_OUTPUT_SCHEMA,
},
"system.session_waiter.register": {
"input": SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA,
"output": SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA,
},
"system.session_waiter.unregister": {
"input": SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA,
"output": SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA,
},
"system.event.react": {
"input": SYSTEM_EVENT_REACT_INPUT_SCHEMA,
"output": SYSTEM_EVENT_REACT_OUTPUT_SCHEMA,
},
"system.event.send_typing": {
"input": SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA,
},
"system.event.send_streaming": {
"input": SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA,
},
"system.event.send_streaming_chunk": {
"input": SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA,
},
"system.event.send_streaming_close": {
"input": SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA,
},
}
__all__ = [
"BUILTIN_CAPABILITY_SCHEMAS",
"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",
"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",
"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",
"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",
"PLATFORM_GET_MEMBERS_INPUT_SCHEMA",
"PLATFORM_GET_MEMBERS_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",
"SESSION_REF_SCHEMA",
"SYSTEM_EVENT_REACT_INPUT_SCHEMA",
"SYSTEM_EVENT_REACT_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",
]
+191 -3
View File
@@ -37,6 +37,9 @@ def _nullable(schema: JSONSchema) -> JSONSchema:
_OPTIONAL_CHAT_PROPERTIES: dict[str, Any] = {
"system": {"type": "string"},
"history": {"type": "array", "items": {"type": "object"}},
"contexts": {"type": "array", "items": {"type": "object"}},
"provider_id": {"type": "string"},
"tool_calls_result": {"type": "array", "items": {"type": "object"}},
"model": {"type": "string"},
"temperature": {"type": "number"},
"image_urls": {"type": "array", "items": {"type": "string"}},
@@ -64,6 +67,9 @@ LLM_CHAT_RAW_OUTPUT_SCHEMA = _object_schema(
usage=_nullable({"type": "object"}),
finish_reason=_nullable({"type": "string"}),
tool_calls={"type": "array", "items": {"type": "object"}},
role=_nullable({"type": "string"}),
reasoning_content=_nullable({"type": "string"}),
reasoning_signature=_nullable({"type": "string"}),
)
LLM_STREAM_CHAT_INPUT_SCHEMA = _object_schema(
required=("prompt",),
@@ -135,7 +141,44 @@ MEMORY_STATS_INPUT_SCHEMA = _object_schema()
MEMORY_STATS_OUTPUT_SCHEMA = _object_schema(
total_items={"type": "integer"},
total_bytes=_nullable({"type": "integer"}),
plugin_id=_nullable({"type": "string"}),
ttl_entries=_nullable({"type": "integer"}),
)
SYSTEM_GET_DATA_DIR_INPUT_SCHEMA = _object_schema()
SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA = _object_schema(
required=("path",),
path={"type": "string"},
)
SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA = _object_schema(
required=("text",),
text={"type": "string"},
return_url={"type": "boolean"},
)
SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA = _object_schema(
required=("result",),
result={"type": "string"},
)
SYSTEM_HTML_RENDER_INPUT_SCHEMA = _object_schema(
required=("tmpl", "data"),
tmpl={"type": "string"},
data={"type": "object"},
return_url={"type": "boolean"},
options=_nullable({"type": "object"}),
)
SYSTEM_HTML_RENDER_OUTPUT_SCHEMA = _object_schema(
required=("result",),
result={"type": "string"},
)
SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA = _object_schema(
required=("session_key",),
session_key={"type": "string"},
)
SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA = _object_schema()
SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA = _object_schema(
required=("session_key",),
session_key={"type": "string"},
)
SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA = _object_schema()
DB_GET_INPUT_SCHEMA = _object_schema(
required=("key",),
key={"type": "string"},
@@ -199,6 +242,45 @@ SESSION_REF_SCHEMA = _object_schema(
platform=_nullable({"type": "string"}),
raw=_nullable({"type": "object"}),
)
SYSTEM_EVENT_REACT_INPUT_SCHEMA = _object_schema(
required=("emoji",),
target=_nullable(SESSION_REF_SCHEMA),
emoji={"type": "string"},
)
SYSTEM_EVENT_REACT_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA = _object_schema(
target=_nullable(SESSION_REF_SCHEMA),
)
SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA = _object_schema(
target=_nullable(SESSION_REF_SCHEMA),
use_fallback={"type": "boolean"},
)
SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
stream_id=_nullable({"type": "string"}),
)
SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA = _object_schema(
required=("stream_id", "chain"),
stream_id={"type": "string"},
chain={"type": "array", "items": {"type": "object"}},
)
SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA = _object_schema()
SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA = _object_schema(
required=("stream_id",),
stream_id={"type": "string"},
)
SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA = _object_schema(
required=("supported",),
supported={"type": "boolean"},
)
PLATFORM_SEND_INPUT_SCHEMA = _object_schema(
required=("session", "text"),
session={"type": "string"},
@@ -392,6 +474,46 @@ BUILTIN_CAPABILITY_SCHEMAS: dict[str, dict[str, JSONSchema]] = {
"input": METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA,
"output": METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA,
},
"system.get_data_dir": {
"input": SYSTEM_GET_DATA_DIR_INPUT_SCHEMA,
"output": SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA,
},
"system.text_to_image": {
"input": SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA,
"output": SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA,
},
"system.html_render": {
"input": SYSTEM_HTML_RENDER_INPUT_SCHEMA,
"output": SYSTEM_HTML_RENDER_OUTPUT_SCHEMA,
},
"system.session_waiter.register": {
"input": SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA,
"output": SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA,
},
"system.session_waiter.unregister": {
"input": SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA,
"output": SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA,
},
"system.event.react": {
"input": SYSTEM_EVENT_REACT_INPUT_SCHEMA,
"output": SYSTEM_EVENT_REACT_OUTPUT_SCHEMA,
},
"system.event.send_typing": {
"input": SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA,
},
"system.event.send_streaming": {
"input": SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA,
},
"system.event.send_streaming_chunk": {
"input": SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA,
},
"system.event.send_streaming_close": {
"input": SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA,
"output": SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA,
},
}
@@ -532,7 +654,7 @@ class ScheduleTrigger(_DescriptorBase):
return self.cron
@model_validator(mode="after")
def validate_schedule(self) -> "ScheduleTrigger":
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:
@@ -540,6 +662,52 @@ class ScheduleTrigger(_DescriptorBase):
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"),
@@ -584,9 +752,12 @@ class HandlerDescriptor(_DescriptorBase):
contract: 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":
def validate_contract_defaults(self) -> HandlerDescriptor:
if self.contract is None:
if isinstance(self.trigger, ScheduleTrigger):
self.contract = "schedule"
@@ -623,7 +794,7 @@ class CapabilityDescriptor(_DescriptorBase):
cancelable: bool = False
@model_validator(mode="after")
def validate_builtin_schema_governance(self) -> "CapabilityDescriptor":
def validate_builtin_schema_governance(self) -> CapabilityDescriptor:
builtin_schema = BUILTIN_CAPABILITY_SCHEMAS.get(self.name)
if builtin_schema is None:
return self
@@ -644,7 +815,9 @@ class CapabilityDescriptor(_DescriptorBase):
__all__ = [
"BUILTIN_CAPABILITY_SCHEMAS",
"CapabilityDescriptor",
"CommandRouteSpec",
"CommandTrigger",
"CompositeFilterSpec",
"DB_DELETE_INPUT_SCHEMA",
"DB_DELETE_OUTPUT_SCHEMA",
"DB_GET_INPUT_SCHEMA",
@@ -660,6 +833,7 @@ __all__ = [
"DB_WATCH_INPUT_SCHEMA",
"DB_WATCH_OUTPUT_SCHEMA",
"EventTrigger",
"FilterSpec",
"HandlerDescriptor",
"HTTP_LIST_APIS_INPUT_SCHEMA",
"HTTP_LIST_APIS_OUTPUT_SCHEMA",
@@ -697,6 +871,8 @@ __all__ = [
"METADATA_LIST_PLUGINS_INPUT_SCHEMA",
"METADATA_LIST_PLUGINS_OUTPUT_SCHEMA",
"MessageTrigger",
"MessageTypeFilterSpec",
"ParamSpec",
"PLATFORM_GET_MEMBERS_INPUT_SCHEMA",
"PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA",
"PLATFORM_SEND_CHAIN_INPUT_SCHEMA",
@@ -711,5 +887,17 @@ __all__ = [
"ScheduleTrigger",
"SESSION_REF_SCHEMA",
"SessionRef",
"SYSTEM_EVENT_REACT_INPUT_SCHEMA",
"SYSTEM_EVENT_REACT_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",
"LocalFilterRefSpec",
"PlatformFilterSpec",
]
+2 -2
View File
@@ -152,7 +152,7 @@ class ResultMessage(_MessageBase):
error: ErrorPayload | None = None
@model_validator(mode="after")
def validate_result_state(self) -> "ResultMessage":
def validate_result_state(self) -> ResultMessage:
"""约束 success / output / error 的组合状态。"""
if self.success:
if self.error is not None:
@@ -238,7 +238,7 @@ class EventMessage(_MessageBase):
error: ErrorPayload | None = None
@model_validator(mode="after")
def validate_phase_constraints(self) -> "EventMessage":
def validate_phase_constraints(self) -> EventMessage:
"""验证各 phase 的字段约束。
- started: 所有字段必须为空
@@ -1,76 +0,0 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Literal
import msgpack
from .messages import ProtocolMessage, parse_message
StdioFraming = Literal["line", "length_prefixed"]
WebSocketFrameType = Literal["text", "binary"]
WireCodecName = Literal["json", "msgpack"]
class ProtocolCodec(ABC):
name: WireCodecName
stdio_framing: StdioFraming
websocket_frame_type: WebSocketFrameType
@abstractmethod
def encode_message(self, message: ProtocolMessage) -> bytes | str:
raise NotImplementedError
@abstractmethod
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
raise NotImplementedError
class JsonProtocolCodec(ProtocolCodec):
name: WireCodecName = "json"
stdio_framing: StdioFraming = "line"
websocket_frame_type: WebSocketFrameType = "text"
def encode_message(self, message: ProtocolMessage) -> str:
return message.model_dump_json(exclude_none=True)
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
if isinstance(payload, bytes):
return parse_message(payload.decode("utf-8"))
return parse_message(payload)
class MsgpackProtocolCodec(ProtocolCodec):
name: WireCodecName = "msgpack"
stdio_framing: StdioFraming = "length_prefixed"
websocket_frame_type: WebSocketFrameType = "binary"
def encode_message(self, message: ProtocolMessage) -> bytes:
return msgpack.packb(
message.model_dump(exclude_none=True),
use_bin_type=True,
)
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
if isinstance(payload, str):
return parse_message(payload)
return parse_message(msgpack.unpackb(payload, raw=False))
def make_protocol_codec(name: WireCodecName | str) -> ProtocolCodec:
if name == "json":
return JsonProtocolCodec()
if name == "msgpack":
return MsgpackProtocolCodec()
raise ValueError(f"未知 wire codec: {name}")
__all__ = [
"JsonProtocolCodec",
"MsgpackProtocolCodec",
"ProtocolCodec",
"StdioFraming",
"WebSocketFrameType",
"WireCodecName",
"make_protocol_codec",
]
+47 -14
View File
@@ -1,21 +1,32 @@
"""AstrBot SDK 的高级运行时原语。
"""AstrBot SDK runtime public exports.
这里仅暴露相对稳定的运行时构件:协议 `Peer`、传输抽象以及能力/处理器分发器。
大多数插件作者应优先使用顶层 `astrbot_sdk`。
本模块提供运行时核心组件的公共导出,包括:
- CapabilityRouter: 能力路由器,处理能力调用的分发和路由
- HandlerDispatcher: 事件处理器分发器,将事件分发到注册的 handler
- Peer: 与 AstrBot 核心通信的对等端抽象
- Transport 系列: 进程间通信传输层实现(stdio/websocket)
`loader` / `bootstrap` 等编排细节保留在各自子模块中,不作为根级稳定契约。
延迟加载策略:
为避免导入时触发 websocket/aiohttp 等重型依赖,采用 __getattr__ 实现按需加载。
这样轻量级导入(如仅使用类型提示)不会产生不必要的依赖开销。
"""
from .capability_router import CapabilityRouter, StreamExecution
from .handler_dispatcher import HandlerDispatcher
from .peer import Peer
from .transport import (
MessageHandler,
StdioTransport,
Transport,
WebSocketClientTransport,
WebSocketServerTransport,
)
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",
@@ -28,3 +39,25 @@ __all__ = [
"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)
@@ -0,0 +1,834 @@
"""Built-in capability registration and handlers for CapabilityRouter.
本模块为 CapabilityRouter 提供内置能力的注册逻辑和处理函数实现。
内置能力涵盖以下类别:
- LLM: 对话、流式对话等大语言模型能力
- Memory: 记忆存储、搜索、带 TTL 的键值对
- DB: 持久化键值存储及变更监听
- Platform: 跨平台消息发送、图片、消息链
- HTTP: 动态 API 路由注册与管理
- Metadata: 插件元数据查询
- System: 数据目录、文本转图片、HTML 渲染、会话等待器等
设计模式:
通过 Mixin 类 (BuiltinCapabilityRouterMixin) 将内置能力注入到 CapabilityRouter,
使其与用户自定义能力共享相同的注册和调用机制。
"""
from __future__ import annotations
import asyncio
import copy
import json
from collections.abc import AsyncIterator
from pathlib import Path
from typing import Any
from ..errors import AstrBotError
from ..protocol.descriptors import (
BUILTIN_CAPABILITY_SCHEMAS,
CapabilityDescriptor,
SessionRef,
)
from ._streaming import StreamExecution
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)
]
class _CapabilityRouterHost:
memory_store: dict[str, dict[str, Any]]
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]
_system_data_root: Path
_session_waiters: dict[str, set[str]]
_db_watch_subscriptions: dict[str, tuple[str | None, asyncio.Queue[dict[str, Any]]]]
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
class BuiltinCapabilityRouterMixin(_CapabilityRouterHost):
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_system_capabilities()
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
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)
},
)
async def _memory_search(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
query = str(payload.get("query", ""))
items = [
{"key": key, "value": value}
for key, value in self.memory_store.items()
if query in key or query in json.dumps(value, ensure_ascii=False)
]
return {"items": items}
async def _memory_save(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
if not isinstance(value, dict):
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
self.memory_store[key] = value
return {}
async def _memory_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.memory_store.get(str(payload.get("key", "")))}
async def _memory_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.memory_store.pop(str(payload.get("key", "")), None)
return {}
async def _memory_save_with_ttl(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
ttl_seconds = payload.get("ttl_seconds", 0)
if not isinstance(value, dict):
raise AstrBotError.invalid_input(
"memory.save_with_ttl 的 value 必须是 object"
)
self.memory_store[key] = {"value": value, "ttl_seconds": ttl_seconds}
return {}
async def _memory_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("memory.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = []
for key in keys:
stored = self.memory_store.get(key)
if (
isinstance(stored, dict)
and "value" in stored
and "ttl_seconds" in stored
):
value = stored["value"]
else:
value = stored
items.append({"key": key, "value": value})
return {"items": items}
async def _memory_delete_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("memory.delete_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
deleted_count = 0
for key in keys:
if key in self.memory_store:
del self.memory_store[key]
deleted_count += 1
return {"deleted_count": deleted_count}
async def _memory_stats(
self, _request_id: str, _payload: dict[str, Any], _token
) -> dict[str, Any]:
total_items = len(self.memory_store)
total_bytes = sum(
len(str(key)) + len(str(value)) for key, value in self.memory_store.items()
)
ttl_entries = sum(
1
for value in self.memory_store.values()
if isinstance(value, dict) and "value" in value and "ttl_seconds" in value
)
return {
"total_items": total_items,
"total_bytes": total_bytes,
"plugin_id": self._require_caller_plugin_id("memory.stats"),
"ttl_entries": ttl_entries,
}
def _register_memory_capabilities(self) -> None:
self.register(
self._builtin_descriptor("memory.search", "搜索记忆"),
call_handler=self._memory_search,
)
self.register(
self._builtin_descriptor("memory.save", "保存记忆"),
call_handler=self._memory_save,
)
self.register(
self._builtin_descriptor("memory.get", "读取单条记忆"),
call_handler=self._memory_get,
)
self.register(
self._builtin_descriptor("memory.delete", "删除记忆"),
call_handler=self._memory_delete,
)
self.register(
self._builtin_descriptor("memory.save_with_ttl", "保存带过期时间的记忆"),
call_handler=self._memory_save_with_ttl,
)
self.register(
self._builtin_descriptor("memory.get_many", "批量获取记忆"),
call_handler=self._memory_get_many,
)
self.register(
self._builtin_descriptor("memory.delete_many", "批量删除记忆"),
call_handler=self._memory_delete_many,
)
self.register(
self._builtin_descriptor("memory.stats", "获取记忆统计信息"),
call_handler=self._memory_stats,
)
async def _db_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.db_store.get(str(payload.get("key", "")))}
async def _db_set(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
self.db_store.pop(key, None)
self._emit_db_change(op="delete", key=key, value=None)
return {}
async def _db_list(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prefix = payload.get("prefix")
keys = sorted(self.db_store.keys())
if isinstance(prefix, str):
keys = [item for item in keys if item.startswith(prefix)]
return {"keys": keys}
async def _db_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
return {"items": items}
async def _db_set_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
items_payload = payload.get("items")
if not isinstance(items_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
for entry in items_payload:
if not isinstance(entry, dict):
raise AstrBotError.invalid_input(
"db.set_many 的 items 必须是 object 数组"
)
key = str(entry.get("key", ""))
value = entry.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_watch(
self, request_id: str, payload: dict[str, Any], _token
) -> StreamExecution:
prefix = payload.get("prefix")
prefix_value: str | None
if isinstance(prefix, str):
prefix_value = prefix
elif prefix is None:
prefix_value = None
else:
raise AstrBotError.invalid_input("db.watch 的 prefix 必须是 string 或 null")
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
async def iterator() -> AsyncIterator[dict[str, Any]]:
try:
while True:
yield await queue.get()
finally:
self._db_watch_subscriptions.pop(request_id, None)
return StreamExecution(
iterator=iterator(),
finalize=lambda _chunks: {},
collect_chunks=False,
)
def _register_db_capabilities(self) -> None:
self.register(
self._builtin_descriptor("db.get", "读取 KV"), call_handler=self._db_get
)
self.register(
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,
)
async def _platform_send(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
text = str(payload.get("text", ""))
message_id = f"msg_{len(self.sent_messages) + 1}"
sent: dict[str, Any] = {
"message_id": message_id,
"session": session,
"text": text,
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_image(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
image_url = str(payload.get("image_url", ""))
message_id = f"img_{len(self.sent_messages) + 1}"
sent: dict[str, Any] = {
"message_id": message_id,
"session": session,
"image_url": image_url,
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_chain(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
chain = payload.get("chain")
if not isinstance(chain, list) or not all(
isinstance(item, dict) for item in chain
):
raise AstrBotError.invalid_input(
"platform.send_chain 的 chain 必须是 object 数组"
)
message_id = f"chain_{len(self.sent_messages) + 1}"
sent: dict[str, Any] = {
"message_id": message_id,
"session": session,
"chain": [dict(item) for item in chain],
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_get_members(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, _target = self._resolve_target(payload)
return {
"members": [
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
]
}
def _register_platform_capabilities(self) -> None:
self.register(
self._builtin_descriptor("platform.send", "发送消息"),
call_handler=self._platform_send,
)
self.register(
self._builtin_descriptor("platform.send_image", "发送图片"),
call_handler=self._platform_send_image,
)
self.register(
self._builtin_descriptor("platform.send_chain", "发送消息链"),
call_handler=self._platform_send_chain,
)
self.register(
self._builtin_descriptor("platform.get_members", "获取群成员"),
call_handler=self._platform_get_members,
)
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"
)
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:
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,
)
async def _metadata_get_plugin(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
name = str(payload.get("name", "")).strip()
plugin = self._plugins.get(name)
if plugin is None:
return {"plugin": None}
return {"plugin": dict(plugin.metadata)}
async def _metadata_list_plugins(
self, _request_id: str, _payload: dict[str, Any], _token
) -> dict[str, Any]:
plugins = [
dict(self._plugins[name].metadata) for name in sorted(self._plugins.keys())
]
return {"plugins": plugins}
async def _metadata_get_plugin_config(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
name = str(payload.get("name", "")).strip()
caller_plugin_id = self._require_caller_plugin_id("metadata.get_plugin_config")
if name != caller_plugin_id:
return {"config": None}
plugin = self._plugins.get(name)
if plugin is None:
return {"config": None}
return {"config": dict(plugin.config)}
def _register_metadata_capabilities(self) -> None:
self.register(
self._builtin_descriptor("metadata.get_plugin", "获取单个插件元数据"),
call_handler=self._metadata_get_plugin,
)
self.register(
self._builtin_descriptor("metadata.list_plugins", "列出插件元数据"),
call_handler=self._metadata_list_plugins,
)
self.register(
self._builtin_descriptor(
"metadata.get_plugin_config",
"获取插件配置",
),
call_handler=self._metadata_get_plugin_config,
)
def _register_system_capabilities(self) -> None:
self.register(
self._builtin_descriptor("system.get_data_dir", "获取插件数据目录"),
call_handler=self._system_get_data_dir,
exposed=False,
)
self.register(
self._builtin_descriptor("system.text_to_image", "文本转图片"),
call_handler=self._system_text_to_image,
exposed=False,
)
self.register(
self._builtin_descriptor("system.html_render", "渲染 HTML 模板"),
call_handler=self._system_html_render,
exposed=False,
)
self.register(
self._builtin_descriptor(
"system.session_waiter.register",
"注册会话等待器",
),
call_handler=self._system_session_waiter_register,
exposed=False,
)
self.register(
self._builtin_descriptor(
"system.session_waiter.unregister",
"注销会话等待器",
),
call_handler=self._system_session_waiter_unregister,
exposed=False,
)
self.register(
self._builtin_descriptor("system.event.react", "发送事件表情回应"),
call_handler=self._system_event_react,
exposed=False,
)
self.register(
self._builtin_descriptor("system.event.send_typing", "发送输入中状态"),
call_handler=self._system_event_send_typing,
exposed=False,
)
self.register(
self._builtin_descriptor(
"system.event.send_streaming",
"发送事件流式消息",
),
call_handler=self._system_event_send_streaming,
exposed=False,
)
self.register(
self._builtin_descriptor(
"system.event.send_streaming_chunk",
"推送事件流式消息分片",
),
call_handler=self._system_event_send_streaming_chunk,
exposed=False,
)
self.register(
self._builtin_descriptor(
"system.event.send_streaming_close",
"关闭事件流式消息会话",
),
call_handler=self._system_event_send_streaming_close,
exposed=False,
)
async def _system_get_data_dir(
self, _request_id: str, _payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("system.get_data_dir")
data_dir = self._system_data_root / plugin_id
data_dir.mkdir(parents=True, exist_ok=True)
return {"path": str(data_dir)}
async def _system_text_to_image(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
text = str(payload.get("text", ""))
if bool(payload.get("return_url", True)):
return {"result": f"mock://text_to_image/{text}"}
return {"result": f"<image>{text}</image>"}
async def _system_html_render(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
tmpl = str(payload.get("tmpl", ""))
data = payload.get("data")
if not isinstance(data, dict):
raise AstrBotError.invalid_input("system.html_render requires object data")
if bool(payload.get("return_url", True)):
return {"result": f"mock://html_render/{tmpl}"}
return {"result": json.dumps({"tmpl": tmpl, "data": data}, ensure_ascii=False)}
async def _system_session_waiter_register(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("system.session_waiter.register")
session_key = str(payload.get("session_key", "")).strip()
if not session_key:
raise AstrBotError.invalid_input(
"system.session_waiter.register requires session_key"
)
self._session_waiters.setdefault(plugin_id, set()).add(session_key)
return {}
async def _system_session_waiter_unregister(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
plugin_id = self._require_caller_plugin_id("system.session_waiter.unregister")
session_key = str(payload.get("session_key", "")).strip()
plugin_waiters = self._session_waiters.get(plugin_id)
if plugin_waiters is None:
return {}
plugin_waiters.discard(session_key)
if not plugin_waiters:
self._session_waiters.pop(plugin_id, None)
return {}
async def _system_event_react(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.event_actions.append(
{
"action": "react",
"emoji": str(payload.get("emoji", "")),
"target": _clone_target_payload(payload.get("target")),
}
)
return {"supported": True}
async def _system_event_send_typing(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.event_actions.append(
{
"action": "send_typing",
"target": _clone_target_payload(payload.get("target")),
}
)
return {"supported": True}
async def _system_event_send_streaming(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
stream_id = f"mock-stream-{len(self._event_streams) + 1}"
stream_state: dict[str, Any] = {
"target": _clone_target_payload(payload.get("target")),
"chunks": [],
"use_fallback": bool(payload.get("use_fallback", False)),
}
self._event_streams[stream_id] = stream_state
return {"supported": True, "stream_id": stream_id}
async def _system_event_send_streaming_chunk(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
stream = self._event_streams.get(str(payload.get("stream_id", "")))
if stream is None:
raise AstrBotError.invalid_input("Unknown sdk event streaming session")
chain = payload.get("chain")
if not isinstance(chain, list):
raise AstrBotError.invalid_input(
"system.event.send_streaming_chunk requires a chain array"
)
stream["chunks"].append({"chain": _clone_chain_payload(chain)})
return {}
async def _system_event_send_streaming_close(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
stream_id = str(payload.get("stream_id", ""))
stream = self._event_streams.pop(stream_id, None)
if stream is None:
raise AstrBotError.invalid_input("Unknown sdk event streaming session")
self.event_actions.append(
{
"action": "send_streaming",
"target": stream["target"],
"chunks": list(stream["chunks"]),
"use_fallback": bool(stream["use_fallback"]),
}
)
return {"supported": True}
__all__ = ["BuiltinCapabilityRouterMixin"]
@@ -0,0 +1,173 @@
"""Support helpers for runtime loader reflection and signature validation.
本模块提供运行时加载器所需的反射和签名验证工具函数,主要用于:
1. 解析 handler/capability 函数签名,提取参数类型信息
2. 识别需要注入的框架对象(如 Context、MessageEvent、ScheduleContext)
3. 构建参数规格 (ParamSpec) 供协议层使用
4. 验证 schedule handler 的签名合法性
关键函数:
- build_param_specs: 从 handler 签名构建参数规格列表
- is_injected_parameter: 判断参数是否应由框架注入而非从命令行解析
- validate_schedule_signature: 确保 schedule handler 只接受允许的注入参数
"""
from __future__ import annotations
import inspect
import typing
from typing import Any, Literal, TypeAlias
from ..decorators import get_capability_meta, get_handler_meta
from ..protocol.descriptors import ParamSpec
from ..schedule import ScheduleContext
from ..types import GreedyStr
ParamTypeName: TypeAlias = Literal[
"str", "int", "float", "bool", "optional", "greedy_str"
]
OptionalInnerType: TypeAlias = Literal["str", "int", "float", "bool"] | None
def unwrap_optional(annotation: Any) -> tuple[Any, bool]:
origin = typing.get_origin(annotation)
if origin is typing.Union:
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
def is_injected_parameter(annotation: Any, parameter_name: str) -> bool:
if parameter_name in {"event", "ctx", "context", "sched", "schedule"}:
return True
normalized, _is_optional = unwrap_optional(annotation)
if normalized is None:
return False
if normalized in {ScheduleContext}:
return True
if isinstance(normalized, type):
from ..context import Context
from ..events import MessageEvent
return issubclass(normalized, (Context, MessageEvent, ScheduleContext))
return False
def param_type_name(annotation: Any) -> tuple[ParamTypeName, OptionalInnerType, bool]:
normalized, is_optional = unwrap_optional(annotation)
if normalized is GreedyStr:
return "greedy_str", None, False
if normalized in {int, float, bool, str}:
if is_optional:
return "optional", normalized.__name__, False
return normalized.__name__, None, True
if is_optional:
return "optional", "str", False
return "str", None, True
def build_param_specs(handler: Any) -> list[ParamSpec]:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return []
try:
type_hints = typing.get_type_hints(handler)
except Exception:
type_hints = {}
specs: list[ParamSpec] = []
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(annotation, parameter.name):
continue
param_type, inner_type, required = param_type_name(annotation)
if parameter.default is not inspect.Parameter.empty:
required = False
specs.append(
ParamSpec(
name=parameter.name,
type=param_type,
required=required,
inner_type=inner_type,
)
)
greedy_indexes = [
index for index, spec in enumerate(specs) if spec.type == "greedy_str"
]
if greedy_indexes and greedy_indexes[-1] != len(specs) - 1:
greedy_spec = specs[greedy_indexes[-1]]
raise ValueError(f"参数 '{greedy_spec.name}' (GreedyStr) 必须是最后一个参数。")
return specs
def validate_schedule_signature(handler: Any) -> None:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return
allowed_names = {"ctx", "context", "sched", "schedule"}
invalid = [
parameter.name
for parameter in signature.parameters.values()
if parameter.kind
in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
and parameter.name not in allowed_names
]
if invalid:
raise ValueError(
"Schedule handler 只允许注入 ctx/context 和 sched/schedule 参数。"
)
def resolve_handler_candidate(instance: Any, name: str) -> tuple[Any, Any] | None:
try:
raw = inspect.getattr_static(instance, name)
except AttributeError:
return None
candidates = [raw]
wrapped = getattr(raw, "__func__", None)
if wrapped is not None:
candidates.append(wrapped)
for candidate in candidates:
meta = get_handler_meta(candidate)
if meta is not None and meta.trigger is not None:
return getattr(instance, name), meta
return None
def resolve_capability_candidate(instance: Any, name: str) -> tuple[Any, Any] | None:
try:
raw = inspect.getattr_static(instance, name)
except AttributeError:
return None
candidates = [raw]
wrapped = getattr(raw, "__func__", None)
if wrapped is not None:
candidates.append(wrapped)
for candidate in candidates:
meta = get_capability_meta(candidate)
if meta is not None:
return getattr(instance, name), meta
return None
__all__ = [
"build_param_specs",
"is_injected_parameter",
"param_type_name",
"resolve_capability_candidate",
"resolve_handler_candidate",
"unwrap_optional",
"validate_schedule_signature",
]
+28
View File
@@ -0,0 +1,28 @@
"""Shared stream execution primitives for runtime internals.
本模块定义流式执行的通用数据结构 StreamExecution,用于:
1. 封装异步生成器迭代器,支持逐块返回数据
2. 提供收集完成后的聚合回调 (finalize)
3. 控制是否需要在内存中累积所有分块
使用场景:
- LLM 流式对话返回逐字输出
- DB watch 监听键值变更流
- 任何需要分块返回而非一次性返回的能力调用
"""
from __future__ import annotations
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from typing import Any
@dataclass(slots=True)
class StreamExecution:
iterator: AsyncIterator[dict[str, Any]]
finalize: Callable[[list[dict[str, Any]]], dict[str, Any]]
collect_chunks: bool = True
__all__ = ["StreamExecution"]
+7 -31
View File
@@ -19,7 +19,6 @@ import sys
from pathlib import Path
from typing import IO
from ..protocol.wire_codecs import make_protocol_codec
from .loader import PluginEnvironmentManager
from .supervisor import (
SupervisorRuntime,
@@ -50,10 +49,9 @@ __all__ = [
async def run_supervisor(
*,
plugins_dir: Path = Path("plugins"),
stdin: IO[str] | IO[bytes] | None = None,
stdout: IO[str] | IO[bytes] | None = None,
stdin: IO[str] | None = None,
stdout: IO[str] | None = None,
env_manager: PluginEnvironmentManager | None = None,
worker_wire_codec: str = "json",
) -> None:
transport_stdin, transport_stdout, original_stdout = _prepare_stdio_transport(
stdin,
@@ -64,7 +62,6 @@ async def run_supervisor(
transport=transport,
plugins_dir=plugins_dir,
env_manager=env_manager,
worker_wire_codec_name=worker_wire_codec,
)
try:
@@ -82,39 +79,26 @@ async def run_plugin_worker(
*,
plugin_dir: Path | None = None,
group_metadata: Path | None = None,
stdin: IO[str] | IO[bytes] | None = None,
stdout: IO[str] | IO[bytes] | None = None,
wire_codec: str = "json",
stdin: IO[str] | None = None,
stdout: IO[str] | None = None,
) -> None:
if plugin_dir is None and group_metadata is None:
raise ValueError("plugin_dir or group_metadata is required")
if plugin_dir is not None and group_metadata is not None:
raise ValueError("plugin_dir and group_metadata are mutually exclusive")
codec = make_protocol_codec(wire_codec)
transport_stdin, transport_stdout, original_stdout = _prepare_stdio_transport(
stdin,
stdout,
binary=codec.stdio_framing == "length_prefixed",
)
transport = StdioTransport(
stdin=transport_stdin,
stdout=transport_stdout,
framing=codec.stdio_framing,
)
transport = StdioTransport(stdin=transport_stdin, stdout=transport_stdout)
if group_metadata is not None:
runtime = GroupWorkerRuntime(
group_metadata_path=group_metadata,
transport=transport,
codec=codec,
)
else:
assert plugin_dir is not None
runtime = PluginWorkerRuntime(
plugin_dir=plugin_dir,
transport=transport,
codec=codec,
)
runtime = PluginWorkerRuntime(plugin_dir=plugin_dir, transport=transport)
try:
await runtime.start()
stop_event = asyncio.Event()
@@ -132,18 +116,10 @@ async def run_websocket_server(
port: int = 8765,
path: str = "/",
plugin_dir: Path | None = None,
wire_codec: str = "json",
) -> None:
codec = make_protocol_codec(wire_codec)
runtime = PluginWorkerRuntime(
plugin_dir=plugin_dir or Path.cwd(),
transport=WebSocketServerTransport(
host=host,
port=port,
path=path,
frame_type=codec.websocket_frame_type,
),
codec=codec,
transport=WebSocketServerTransport(host=host, port=port, path=path),
)
try:
await runtime.start()
@@ -0,0 +1,267 @@
"""Capability invocation dispatcher.
本模块实现能力调用的分发器,负责:
1. 接收能力调用请求,定位对应的已注册能力
2. 构建调用上下文 (Context),注入必要的依赖
3. 支持同步和流式两种调用模式
4. 管理活跃调用任务的生命周期和取消
参数注入策略:
按类型注入 Context / CancelToken / dict,或按参数名注入
ctx / context / payload / input / data / cancel_token / token。
若无法匹配则抛出详细的错误信息,帮助开发者定位问题。
"""
from __future__ import annotations
import asyncio
import inspect
import typing
from collections.abc import AsyncIterator
from typing import Any, get_type_hints
from .._invocation_context import caller_plugin_scope
from ..context import CancelToken, Context
from ..errors import AstrBotError
from ._streaming import StreamExecution
from .loader import LoadedCapability
class CapabilityDispatcher:
def __init__(
self,
*,
plugin_id: str,
peer,
capabilities: list[LoadedCapability],
) -> None:
self._plugin_id = plugin_id
self._peer = peer
self._capabilities = {item.descriptor.name: item for item in capabilities}
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
async def invoke(
self,
message,
cancel_token: CancelToken,
) -> dict[str, Any] | StreamExecution:
loaded = self._capabilities.get(message.capability)
if loaded is None:
raise LookupError(f"capability not found: {message.capability}")
plugin_id = self._resolve_plugin_id(loaded)
ctx = Context(
peer=self._peer,
plugin_id=plugin_id,
cancel_token=cancel_token,
)
with caller_plugin_scope(plugin_id):
task = asyncio.create_task(
self._run_capability(
loaded,
payload=dict(message.input),
ctx=ctx,
cancel_token=cancel_token,
stream=bool(message.stream),
)
)
self._active[message.id] = (task, cancel_token)
try:
return await task
finally:
self._active.pop(message.id, None)
def _resolve_plugin_id(self, loaded: LoadedCapability) -> str:
if loaded.plugin_id:
return loaded.plugin_id
return self._plugin_id
async def cancel(self, request_id: str) -> None:
active = self._active.get(request_id)
if active is None:
return
task, cancel_token = active
cancel_token.cancel()
task.cancel()
async def _run_capability(
self,
loaded: LoadedCapability,
*,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
stream: bool,
) -> dict[str, Any] | StreamExecution:
result = loaded.callable(
*self._build_args(
loaded.callable,
payload,
ctx,
cancel_token,
plugin_id=self._resolve_plugin_id(loaded),
capability_name=loaded.descriptor.name,
)
)
if stream:
if inspect.isasyncgen(result):
return StreamExecution(
iterator=self._iterate_generator(result),
finalize=lambda chunks: {"items": chunks},
)
if inspect.isawaitable(result):
result = await result
if isinstance(result, StreamExecution):
return result
raise AstrBotError.protocol_error(
"stream=true 的插件 capability 必须返回 async generator 或 StreamExecution"
)
if inspect.isasyncgen(result):
raise AstrBotError.protocol_error(
"stream=false 的插件 capability 不能返回 async generator"
)
if inspect.isawaitable(result):
result = await result
return self._normalize_output(result)
def _build_args(
self,
handler,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
*,
plugin_id: str | None = None,
capability_name: str | None = None,
) -> list[Any]:
signature = inspect.signature(handler)
args: list[Any] = []
type_hints: dict[str, Any] = {}
try:
type_hints = get_type_hints(handler)
except Exception:
pass
for parameter in signature.parameters.values():
if parameter.kind not in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
continue
injected = None
param_type = type_hints.get(parameter.name)
if param_type is not None:
injected = self._inject_by_type(param_type, payload, ctx, cancel_token)
if injected is None:
if parameter.name in {"ctx", "context"}:
injected = ctx
elif parameter.name in {"payload", "input", "data"}:
injected = payload
elif parameter.name in {"cancel_token", "token"}:
injected = cancel_token
if injected is None:
if parameter.default is not parameter.empty:
continue
raise TypeError(
self._format_capability_injection_error(
handler=handler,
parameter_name=parameter.name,
plugin_id=plugin_id,
capability_name=capability_name,
payload=payload,
)
)
args.append(injected)
return args
def _inject_by_type(
self,
param_type: Any,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
) -> Any:
origin = typing.get_origin(param_type)
if origin is typing.Union:
type_args = typing.get_args(param_type)
non_none_types = [item for item in type_args if item is not type(None)]
if len(non_none_types) == 1:
param_type = non_none_types[0]
origin = typing.get_origin(param_type)
if param_type is Context or (
isinstance(param_type, type) and issubclass(param_type, Context)
):
return ctx
if param_type is CancelToken or (
isinstance(param_type, type) and issubclass(param_type, CancelToken)
):
return cancel_token
if param_type is dict or origin is dict:
return payload
return None
def _format_capability_injection_error(
self,
*,
handler,
parameter_name: str,
plugin_id: str | None,
capability_name: str | None,
payload: dict[str, Any],
) -> str:
plugin_text = plugin_id or self._plugin_id
target = capability_name or getattr(handler, "__name__", "<anonymous>")
payload_keys = sorted(str(key) for key in payload.keys())
payload_keys_text = ", ".join(payload_keys) if payload_keys else "<none>"
return (
f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:"
f"必填参数 '{parameter_name}' 无法注入。"
f"签名: {getattr(handler, '__name__', '<anonymous>')}"
f"{self._callable_signature(handler)}。"
"当前支持按类型注入 Context / CancelToken / dict,"
"按参数名注入 ctx / context / payload / input / data / cancel_token / token,"
f"以及 payload 中现有键:{payload_keys_text}。"
)
async def _iterate_generator(
self,
generator: AsyncIterator[Any],
) -> AsyncIterator[dict[str, Any]]:
async for item in generator:
yield self._normalize_chunk(item)
def _normalize_chunk(self, item: Any) -> dict[str, Any]:
output = self._normalize_output(item)
if output:
return output
return {"ok": True}
def _normalize_output(self, result: Any) -> dict[str, Any]:
if result is None:
return {}
if isinstance(result, dict):
return result
model_dump = getattr(result, "model_dump", None)
if callable(model_dump):
dumped = model_dump()
if isinstance(dumped, dict):
return dumped
raise AstrBotError.invalid_input("插件 capability 必须返回 dict 或可序列化对象")
@staticmethod
def _callable_signature(handler) -> str:
try:
return str(inspect.signature(handler))
except (TypeError, ValueError):
return "(?)"
__all__ = ["CapabilityDispatcher"]
+16 -609
View File
@@ -97,35 +97,27 @@
from __future__ import annotations
import asyncio
import copy
import inspect
import json
import re
from collections.abc import AsyncIterator, Awaitable, Callable
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from .._invocation_context import current_caller_plugin_id
from ..errors import AstrBotError
from ..protocol.descriptors import (
BUILTIN_CAPABILITY_SCHEMAS,
RESERVED_CAPABILITY_PREFIXES,
CapabilityDescriptor,
SessionRef,
)
from ._capability_router_builtins import BuiltinCapabilityRouterMixin
from ._streaming import StreamExecution
CallHandler = Callable[[str, dict[str, Any], object], Awaitable[dict[str, Any]]]
FinalizeHandler = Callable[[list[dict[str, Any]]], dict[str, Any]]
CAPABILITY_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9_]*\.[a-z][a-z0-9_]*$")
@dataclass(slots=True)
class StreamExecution:
iterator: AsyncIterator[dict[str, Any]]
finalize: FinalizeHandler
collect_chunks: bool = True
StreamHandler = Callable[
[str, dict[str, Any], object],
AsyncIterator[dict[str, Any]]
@@ -149,14 +141,18 @@ class _RegisteredPlugin:
config: dict[str, Any]
class CapabilityRouter:
class CapabilityRouter(BuiltinCapabilityRouterMixin):
def __init__(self) -> None:
self._registrations: dict[str, _CapabilityRegistration] = {}
self.db_store: dict[str, Any] = {}
self.memory_store: dict[str, dict[str, Any]] = {}
self.sent_messages: list[dict[str, Any]] = []
self.event_actions: list[dict[str, Any]] = []
self._event_streams: dict[str, dict[str, Any]] = {}
self.http_api_store: list[dict[str, Any]] = []
self._plugins: dict[str, _RegisteredPlugin] = {}
self._system_data_root = Path.cwd() / ".astrbot_sdk_testing" / "plugin_data"
self._session_waiters: dict[str, set[str]] = {}
self._db_watch_subscriptions: dict[
str, tuple[str | None, asyncio.Queue[dict[str, Any]]]
] = {}
@@ -212,9 +208,7 @@ class CapabilityRouter:
queue.put_nowait(event)
def descriptors(self) -> list[CapabilityDescriptor]:
return [
entry.descriptor for entry in self._registrations.values() if entry.exposed
]
return [entry.descriptor for entry in self._registrations.values()]
def contains(self, name: str) -> bool:
return name in self._registrations
@@ -231,7 +225,13 @@ class CapabilityRouter:
finalize: FinalizeHandler | None = None,
exposed: bool = True,
) -> None:
if not CAPABILITY_NAME_PATTERN.fullmatch(descriptor.name):
is_internal_reserved = not exposed and descriptor.name.startswith(
RESERVED_CAPABILITY_PREFIXES
)
if (
not CAPABILITY_NAME_PATTERN.fullmatch(descriptor.name)
and not is_internal_reserved
):
raise ValueError(
f"capability 名称必须匹配 {{namespace}}.{{method}}:{descriptor.name}"
)
@@ -320,599 +320,6 @@ class CapabilityRouter:
collect_chunks=execution.collect_chunks,
)
# ------------------------------------------------------------------
# Built-in capability registration
# ------------------------------------------------------------------
def _register_builtin_capabilities(self) -> None:
"""注册全部内建 capability。"""
self._register_llm_capabilities()
self._register_memory_capabilities()
self._register_db_capabilities()
self._register_platform_capabilities()
self._register_http_capabilities()
self._register_metadata_capabilities()
def _builtin_descriptor(
self,
name: str,
description: str,
*,
supports_stream: bool = False,
cancelable: bool = False,
) -> CapabilityDescriptor:
"""构建内建 capability 描述符,schema 从注册表读取。"""
schema = BUILTIN_CAPABILITY_SCHEMAS[name]
return CapabilityDescriptor(
name=name,
description=description,
input_schema=copy.deepcopy(schema["input"]),
output_schema=copy.deepcopy(schema["output"]),
supports_stream=supports_stream,
cancelable=cancelable,
)
def _resolve_target(
self, payload: dict[str, Any]
) -> tuple[str, dict[str, Any] | None]:
"""从 payload 解析 session + target。"""
target_payload = payload.get("target")
if isinstance(target_payload, dict):
target = SessionRef.model_validate(target_payload)
return target.session, target.to_payload()
return str(payload.get("session", "")), None
# ------------------------------------------------------------------
# LLM handlers
# ------------------------------------------------------------------
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]]: # type: ignore[override]
text = f"Echo: {str(payload.get('prompt', ''))}"
for char in text:
token.raise_if_cancelled()
await asyncio.sleep(0)
yield {"text": char}
def _register_llm_capabilities(self) -> None:
self.register(
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)
},
)
# ------------------------------------------------------------------
# Memory handlers
# ------------------------------------------------------------------
async def _memory_search(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
query = str(payload.get("query", ""))
items = [
{"key": key, "value": value}
for key, value in self.memory_store.items()
if query in key or query in json.dumps(value, ensure_ascii=False)
]
return {"items": items}
async def _memory_save(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
if not isinstance(value, dict):
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
self.memory_store[key] = value
return {}
async def _memory_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.memory_store.get(str(payload.get("key", "")))}
async def _memory_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
self.memory_store.pop(str(payload.get("key", "")), None)
return {}
async def _memory_save_with_ttl(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
"""保存带 TTL 的记忆项(测试实现,TTL 仅记录但不实际过期)。"""
key = str(payload.get("key", ""))
value = payload.get("value")
ttl_seconds = payload.get("ttl_seconds", 0)
if not isinstance(value, dict):
raise AstrBotError.invalid_input(
"memory.save_with_ttl 的 value 必须是 object"
)
# 在测试实现中,我们只存储值,TTL 由实际后端实现
self.memory_store[key] = {"value": value, "ttl_seconds": ttl_seconds}
return {}
async def _memory_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
"""批量获取多个记忆项。"""
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("memory.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = []
for key in keys:
stored = self.memory_store.get(key)
# 如果存储的是带 TTL 的结构,提取实际值
if (
isinstance(stored, dict)
and "value" in stored
and "ttl_seconds" in stored
):
value = stored["value"]
else:
value = stored
items.append({"key": key, "value": value})
return {"items": items}
async def _memory_delete_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
"""批量删除多个记忆项。"""
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("memory.delete_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
deleted_count = 0
for key in keys:
if key in self.memory_store:
del self.memory_store[key]
deleted_count += 1
return {"deleted_count": deleted_count}
async def _memory_stats(
self, _request_id: str, _payload: dict[str, Any], _token
) -> dict[str, Any]:
"""获取记忆统计信息。"""
total_items = len(self.memory_store)
# 简单估算字节大小
total_bytes = sum(
len(str(key)) + len(str(value)) for key, value in self.memory_store.items()
)
return {"total_items": total_items, "total_bytes": total_bytes}
def _register_memory_capabilities(self) -> None:
self.register(
self._builtin_descriptor("memory.search", "搜索记忆"),
call_handler=self._memory_search,
)
self.register(
self._builtin_descriptor("memory.save", "保存记忆"),
call_handler=self._memory_save,
)
self.register(
self._builtin_descriptor("memory.get", "读取单条记忆"),
call_handler=self._memory_get,
)
self.register(
self._builtin_descriptor("memory.delete", "删除记忆"),
call_handler=self._memory_delete,
)
self.register(
self._builtin_descriptor("memory.save_with_ttl", "保存带过期时间的记忆"),
call_handler=self._memory_save_with_ttl,
)
self.register(
self._builtin_descriptor("memory.get_many", "批量获取记忆"),
call_handler=self._memory_get_many,
)
self.register(
self._builtin_descriptor("memory.delete_many", "批量删除记忆"),
call_handler=self._memory_delete_many,
)
self.register(
self._builtin_descriptor("memory.stats", "获取记忆统计信息"),
call_handler=self._memory_stats,
)
# ------------------------------------------------------------------
# DB handlers
# ------------------------------------------------------------------
async def _db_get(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
return {"value": self.db_store.get(str(payload.get("key", "")))}
async def _db_set(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
value = payload.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_delete(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
key = str(payload.get("key", ""))
self.db_store.pop(key, None)
self._emit_db_change(op="delete", key=key, value=None)
return {}
async def _db_list(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
prefix = payload.get("prefix")
keys = sorted(self.db_store.keys())
if isinstance(prefix, str):
keys = [item for item in keys if item.startswith(prefix)]
return {"keys": keys}
async def _db_get_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
keys_payload = payload.get("keys")
if not isinstance(keys_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
keys = [str(item) for item in keys_payload]
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
return {"items": items}
async def _db_set_many(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
items_payload = payload.get("items")
if not isinstance(items_payload, (list, tuple)):
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
for entry in items_payload:
if not isinstance(entry, dict):
raise AstrBotError.invalid_input(
"db.set_many 的 items 必须是 object 数组"
)
key = str(entry.get("key", ""))
value = entry.get("value")
self.db_store[key] = value
self._emit_db_change(op="set", key=key, value=value)
return {}
async def _db_watch(
self, request_id: str, payload: dict[str, Any], _token
) -> StreamExecution:
prefix = payload.get("prefix")
prefix_value: str | None
if isinstance(prefix, str):
prefix_value = prefix
elif prefix is None:
prefix_value = None
else:
raise AstrBotError.invalid_input("db.watch 的 prefix 必须是 string 或 null")
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
async def iterator() -> AsyncIterator[dict[str, Any]]:
try:
while True:
yield await queue.get()
finally:
self._db_watch_subscriptions.pop(request_id, None)
return StreamExecution(
iterator=iterator(),
finalize=lambda _chunks: {},
collect_chunks=False,
)
def _register_db_capabilities(self) -> None:
self.register(
self._builtin_descriptor("db.get", "读取 KV"),
call_handler=self._db_get,
)
self.register(
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,
)
# ------------------------------------------------------------------
# Platform handlers
# ------------------------------------------------------------------
async def _platform_send(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
text = str(payload.get("text", ""))
message_id = f"msg_{len(self.sent_messages) + 1}"
sent = {"message_id": message_id, "session": session, "text": text}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_image(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
image_url = str(payload.get("image_url", ""))
message_id = f"img_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"image_url": image_url,
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_send_chain(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, target = self._resolve_target(payload)
chain = payload.get("chain")
if not isinstance(chain, list) or not all(
isinstance(item, dict) for item in chain
):
raise AstrBotError.invalid_input(
"platform.send_chain 的 chain 必须是 object 数组"
)
message_id = f"chain_{len(self.sent_messages) + 1}"
sent = {
"message_id": message_id,
"session": session,
"chain": [dict(item) for item in chain],
}
if target is not None:
sent["target"] = target
self.sent_messages.append(sent)
return {"message_id": message_id}
async def _platform_get_members(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
session, _target = self._resolve_target(payload)
return {
"members": [
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
]
}
def _register_platform_capabilities(self) -> None:
self.register(
self._builtin_descriptor("platform.send", "发送消息"),
call_handler=self._platform_send,
)
self.register(
self._builtin_descriptor("platform.send_image", "发送图片"),
call_handler=self._platform_send_image,
)
self.register(
self._builtin_descriptor("platform.send_chain", "发送消息链"),
call_handler=self._platform_send_chain,
)
self.register(
self._builtin_descriptor("platform.get_members", "获取群成员"),
call_handler=self._platform_get_members,
)
# ------------------------------------------------------------------
# HTTP handlers
# ------------------------------------------------------------------
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"
)
plugin_name = self._require_caller_plugin_id("http.register_api")
methods = sorted({method.upper() for method in methods_payload if method})
entry = {
"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:
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,
)
# ------------------------------------------------------------------
# Metadata handlers
# ------------------------------------------------------------------
async def _metadata_get_plugin(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
name = str(payload.get("name", "")).strip()
plugin = self._plugins.get(name)
if plugin is None:
return {"plugin": None}
return {"plugin": dict(plugin.metadata)}
async def _metadata_list_plugins(
self, _request_id: str, _payload: dict[str, Any], _token
) -> dict[str, Any]:
plugins = [
dict(self._plugins[name].metadata) for name in sorted(self._plugins.keys())
]
return {"plugins": plugins}
async def _metadata_get_plugin_config(
self, _request_id: str, payload: dict[str, Any], _token
) -> dict[str, Any]:
name = str(payload.get("name", "")).strip()
caller_plugin_id = self._require_caller_plugin_id("metadata.get_plugin_config")
if name != caller_plugin_id:
return {"config": None}
plugin = self._plugins.get(name)
if plugin is None:
return {"config": None}
return {"config": dict(plugin.config)}
def _register_metadata_capabilities(self) -> None:
self.register(
self._builtin_descriptor("metadata.get_plugin", "获取单个插件元数据"),
call_handler=self._metadata_get_plugin,
)
self.register(
self._builtin_descriptor("metadata.list_plugins", "列出插件元数据"),
call_handler=self._metadata_list_plugins,
)
self.register(
self._builtin_descriptor(
"metadata.get_plugin_config",
"获取插件配置",
),
call_handler=self._metadata_get_plugin_config,
)
# ------------------------------------------------------------------
# Schema validation
# ------------------------------------------------------------------
@@ -37,6 +37,10 @@ GROUP_STATE_FILE_NAME = ".group-venv-state.json"
_EXACT_PIN_PATTERN = re.compile(r"^([A-Za-z0-9_.-]+)==([^\s;]+)$")
_NORMALIZE_PATTERN = re.compile(r"[-_.]+")
_PYVENV_VERSION_PATTERN = re.compile(
r"^(?:version|version_info)\s*=\s*(\d+\.\d+)(?:\.\d+)?\s*$",
re.IGNORECASE | re.MULTILINE,
)
def _venv_python_path(venv_path: Path) -> Path:
@@ -49,6 +53,19 @@ def _normalize_package_name(name: str) -> str:
return _NORMALIZE_PATTERN.sub("-", name).lower()
def _read_pyvenv_major_minor(pyvenv_cfg: Path) -> str | None:
if not pyvenv_cfg.exists():
return None
try:
content = pyvenv_cfg.read_text(encoding="utf-8")
except OSError:
return None
match = _PYVENV_VERSION_PATTERN.search(content)
if match is None:
return None
return match.group(1)
def _requirement_lines(plugin: PluginSpec) -> list[str]:
if not plugin.requirements_path.exists():
return []
@@ -612,15 +629,7 @@ class GroupEnvironmentManager:
@staticmethod
def _matches_python_version(venv_path: Path, version: str) -> bool:
pyvenv_cfg = venv_path / "pyvenv.cfg"
if not pyvenv_cfg.exists():
return False
try:
content = pyvenv_cfg.read_text(encoding="utf-8")
except OSError:
return False
match = re.search(r"version\s*=\s*(\d+\.\d+)\.\d+", content, re.IGNORECASE)
return match is not None and match.group(1) == version
return _read_pyvenv_major_minor(venv_path / "pyvenv.cfg") == version
@staticmethod
def _load_state(state_path: Path) -> dict[str, object]:
+200 -253
View File
@@ -27,17 +27,25 @@ import inspect
import re
import shlex
import typing
from collections.abc import AsyncIterator
from typing import Any, get_type_hints
from .._invocation_context import caller_plugin_scope
from ..context import CancelToken, Context
from ..errors import AstrBotError
from ..events import MessageEvent
from ..protocol.descriptors import CommandTrigger, MessageTrigger
from ..filters import LocalFilterBinding
from ..message_components import BaseMessageComponent
from ..message_result import MessageChain, MessageEventResult, coerce_message_chain
from ..protocol.descriptors import (
CommandTrigger,
MessageTrigger,
ParamSpec,
ScheduleTrigger,
)
from ..schedule import ScheduleContext
from ..session_waiter import SessionWaiterManager
from ..star import Star
from .capability_router import StreamExecution
from .loader import LoadedCapability, LoadedHandler
from .capability_dispatcher import CapabilityDispatcher
from .loader import LoadedHandler
class HandlerDispatcher:
@@ -46,9 +54,27 @@ class HandlerDispatcher:
self._peer = peer
self._handlers = {item.descriptor.id: item for item in handlers}
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
self._session_waiters = SessionWaiterManager(plugin_id=plugin_id, peer=peer)
setattr(peer, "_session_waiter_manager", self._session_waiters)
async def invoke(self, message, cancel_token: CancelToken) -> dict[str, Any]:
handler_id = str(message.input.get("handler_id", ""))
if handler_id == "__sdk_session_waiter__":
plugin_id = self._plugin_id
ctx = Context(
peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token
)
event = MessageEvent.from_payload(
message.input.get("event", {}), context=ctx
)
event.bind_reply_handler(self._create_reply_handler(ctx, event))
task = asyncio.create_task(self._session_waiters.dispatch(event))
self._active[message.id] = (task, cancel_token)
try:
return await task
finally:
self._active.pop(message.id, None)
loaded = self._handlers.get(handler_id)
if loaded is None:
raise LookupError(f"handler not found: {handler_id}")
@@ -57,6 +83,9 @@ class HandlerDispatcher:
ctx = Context(peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token)
event = MessageEvent.from_payload(message.input.get("event", {}), context=ctx)
event.bind_reply_handler(self._create_reply_handler(ctx, event))
schedule_context = self._build_schedule_context(
loaded, message.input.get("event", {})
)
# 提取 args 用于兼容 handler 签名
raw_args = message.input.get("args") or {}
@@ -65,7 +94,15 @@ class HandlerDispatcher:
args = self._derive_args(loaded, event)
with caller_plugin_scope(plugin_id):
task = asyncio.create_task(self._run_handler(loaded, event, ctx, args))
task = asyncio.create_task(
self._run_handler(
loaded,
event,
ctx,
args,
schedule_context=schedule_context,
)
)
self._active[message.id] = (task, cancel_token)
try:
return await task
@@ -108,17 +145,31 @@ class HandlerDispatcher:
event: MessageEvent,
ctx: Context,
args: dict[str, Any] | None = None,
*,
schedule_context: ScheduleContext | None = None,
) -> dict[str, Any]:
summary = {"sent_message": False, "stop": False, "call_llm": False}
try:
if not self._run_local_filters(
loaded.local_filters,
event=event,
ctx=ctx,
):
return summary
parsed_args = (
self._parse_handler_args(loaded.descriptor.param_specs, args or {})
if loaded.descriptor.param_specs
else dict(args or {})
)
result = loaded.callable(
*self._build_args(
loaded.callable,
event,
ctx,
args,
parsed_args,
plugin_id=self._resolve_plugin_id(loaded),
handler_ref=loaded.descriptor.id,
schedule_context=schedule_context,
)
)
if inspect.isasyncgen(result):
@@ -127,6 +178,7 @@ class HandlerDispatcher:
summary,
await self._handle_result_item(item, event, ctx),
)
summary["stop"] = bool(summary.get("stop")) or event.is_stopped()
return summary
if inspect.isawaitable(result):
result = await result
@@ -135,6 +187,7 @@ class HandlerDispatcher:
summary,
await self._handle_result_item(result, event, ctx),
)
summary["stop"] = bool(summary.get("stop")) or event.is_stopped()
return summary
except Exception as exc:
await self._handle_error(
@@ -154,16 +207,35 @@ class HandlerDispatcher:
) -> dict[str, Any]:
trigger = loaded.descriptor.trigger
if isinstance(trigger, CommandTrigger):
param_specs = loaded.descriptor.param_specs
for command_name in [trigger.command, *trigger.aliases]:
remainder = self._match_command_name(event.text, command_name)
if remainder is not None:
return self._build_command_args(loaded.callable, remainder)
if param_specs:
return self._build_command_args(param_specs, remainder)
return self._build_command_args(
[
ParamSpec(name=name, type="str")
for name in self._legacy_arg_parameter_names(
loaded.callable
)
],
remainder,
)
return {}
if isinstance(trigger, MessageTrigger) and trigger.regex:
match = re.search(trigger.regex, event.text)
if match is None:
return {}
return self._build_regex_args(loaded.callable, match)
if loaded.descriptor.param_specs:
return self._build_regex_args(loaded.descriptor.param_specs, match)
return self._build_regex_args(
[
ParamSpec(name=name, type="str")
for name in self._legacy_arg_parameter_names(loaded.callable)
],
match,
)
return {}
def _build_args(
@@ -175,6 +247,7 @@ class HandlerDispatcher:
*,
plugin_id: str | None = None,
handler_ref: str | None = None,
schedule_context: ScheduleContext | None = None,
) -> list[Any]:
"""构建 handler 参数列表。"""
from loguru import logger
@@ -201,7 +274,9 @@ class HandlerDispatcher:
# 1. 优先按类型注解注入
param_type = type_hints.get(parameter.name)
if param_type is not None:
injected = self._inject_by_type(param_type, event, ctx)
injected = self._inject_by_type(
param_type, event, ctx, schedule_context
)
# 2. Fallback 按名字注入
if injected is None:
@@ -209,6 +284,8 @@ class HandlerDispatcher:
injected = event
elif parameter.name in {"ctx", "context"}:
injected = ctx
elif parameter.name in {"sched", "schedule"}:
injected = schedule_context
elif parameter.name in args:
injected = args[parameter.name]
@@ -236,7 +313,11 @@ class HandlerDispatcher:
return injected_args
def _inject_by_type(
self, param_type: Any, event: MessageEvent, ctx: Context
self,
param_type: Any,
event: MessageEvent,
ctx: Context,
schedule_context: ScheduleContext | None,
) -> Any:
"""根据类型注解注入参数。"""
# 处理 Optional[Type] 情况
@@ -263,6 +344,10 @@ class HandlerDispatcher:
isinstance(param_type, type) and issubclass(param_type, Context)
):
return ctx
if param_type is ScheduleContext or (
isinstance(param_type, type) and issubclass(param_type, ScheduleContext)
):
return schedule_context
return None
@@ -341,6 +426,23 @@ class HandlerDispatcher:
if isinstance(item, dict) and "text" in item:
await event.reply(str(item["text"]))
return True
if isinstance(item, MessageEventResult):
chain = item.chain
if chain.components:
await event.reply_chain(chain)
return True
return False
chain = coerce_message_chain(item)
if chain is not None:
if chain.components:
await event.reply_chain(chain)
return True
return False
if isinstance(item, list) and all(
isinstance(component, BaseMessageComponent) for component in item
):
await event.reply_chain(MessageChain(list(item)))
return True
# 支持带 text 属性的对象
text = getattr(item, "text", None)
if isinstance(text, str):
@@ -358,27 +460,32 @@ class HandlerDispatcher:
return None
@classmethod
def _build_command_args(cls, handler, remainder: str) -> dict[str, Any]:
names = cls._legacy_arg_parameter_names(handler)
if not names or not remainder:
def _build_command_args(
cls, param_specs: list[ParamSpec], remainder: str
) -> dict[str, Any]:
if not param_specs or not remainder:
return {}
if len(names) == 1:
return {names[0]: remainder}
if len(param_specs) == 1:
return {param_specs[0].name: remainder}
parts = cls._split_command_remainder(remainder)
return {
name: parts[index] for index, name in enumerate(names) if index < len(parts)
}
values: dict[str, Any] = {}
for index, spec in enumerate(param_specs):
if index >= len(parts):
break
if spec.type == "greedy_str":
values[spec.name] = " ".join(parts[index:])
break
values[spec.name] = parts[index]
return values
@classmethod
def _build_regex_args(cls, handler, match: re.Match[str]) -> dict[str, Any]:
def _build_regex_args(
cls, param_specs: list[ParamSpec], match: re.Match[str]
) -> dict[str, Any]:
named = {
key: value for key, value in match.groupdict().items() if value is not None
}
names = [
name
for name in cls._legacy_arg_parameter_names(handler)
if name not in named
]
names = [spec.name for spec in param_specs if spec.name not in named]
positional = [value for value in match.groups() if value is not None]
for index, value in enumerate(positional):
if index >= len(names):
@@ -386,6 +493,73 @@ class HandlerDispatcher:
named[names[index]] = value
return named
@staticmethod
def _parse_handler_args(
param_specs: list[ParamSpec],
args: dict[str, Any],
) -> dict[str, Any]:
parsed: dict[str, Any] = {}
for spec in param_specs:
if spec.name not in args:
if spec.type == "optional":
parsed[spec.name] = None
continue
if spec.required:
raise TypeError(f"缺少参数: {spec.name}")
continue
parsed[spec.name] = HandlerDispatcher._convert_param(spec, args[spec.name])
return parsed
@staticmethod
def _convert_param(spec: ParamSpec, value: Any) -> Any:
if spec.type in {"str", "greedy_str"}:
return str(value)
if spec.type == "int":
return int(str(value))
if spec.type == "float":
return float(str(value))
if spec.type == "bool":
normalized = str(value).strip().lower()
if normalized in {"true", "1", "yes", "on"}:
return True
if normalized in {"false", "0", "no", "off"}:
return False
raise TypeError(f"无法解析布尔参数 {spec.name}: {value!r}")
if spec.type == "optional":
if value is None:
return None
inner = ParamSpec(
name=spec.name,
type=spec.inner_type or "str",
required=False,
)
return HandlerDispatcher._convert_param(inner, value)
return value
@staticmethod
def _run_local_filters(
bindings: list[LocalFilterBinding],
*,
event: MessageEvent,
ctx: Context,
) -> bool:
for binding in bindings:
if not binding.evaluate(event=event, ctx=ctx):
return False
return True
@staticmethod
def _build_schedule_context(
loaded: LoadedHandler,
event_payload: dict[str, Any],
) -> ScheduleContext | None:
if not isinstance(loaded.descriptor.trigger, ScheduleTrigger):
return None
try:
return ScheduleContext.from_payload(event_payload)
except Exception:
return None
@staticmethod
def _split_command_remainder(remainder: str) -> list[str]:
try:
@@ -464,231 +638,4 @@ class HandlerDispatcher:
await Star().on_error(exc, event, ctx)
class CapabilityDispatcher:
def __init__(
self,
*,
plugin_id: str,
peer,
capabilities: list[LoadedCapability],
) -> None:
self._plugin_id = plugin_id
self._peer = peer
self._capabilities = {item.descriptor.name: item for item in capabilities}
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
async def invoke(
self,
message,
cancel_token: CancelToken,
) -> dict[str, Any] | StreamExecution:
loaded = self._capabilities.get(message.capability)
if loaded is None:
raise LookupError(f"capability not found: {message.capability}")
plugin_id = self._resolve_plugin_id(loaded)
ctx = Context(
peer=self._peer,
plugin_id=plugin_id,
cancel_token=cancel_token,
)
with caller_plugin_scope(plugin_id):
task = asyncio.create_task(
self._run_capability(
loaded,
payload=dict(message.input),
ctx=ctx,
cancel_token=cancel_token,
stream=bool(message.stream),
)
)
self._active[message.id] = (task, cancel_token)
try:
return await task
finally:
self._active.pop(message.id, None)
def _resolve_plugin_id(self, loaded: LoadedCapability) -> str:
if loaded.plugin_id:
return loaded.plugin_id
return self._plugin_id
async def cancel(self, request_id: str) -> None:
active = self._active.get(request_id)
if active is None:
return
task, cancel_token = active
cancel_token.cancel()
task.cancel()
async def _run_capability(
self,
loaded: LoadedCapability,
*,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
stream: bool,
) -> dict[str, Any] | StreamExecution:
result = loaded.callable(
*self._build_args(
loaded.callable,
payload,
ctx,
cancel_token,
plugin_id=self._resolve_plugin_id(loaded),
capability_name=loaded.descriptor.name,
)
)
if stream:
if inspect.isasyncgen(result):
return StreamExecution(
iterator=self._iterate_generator(result),
finalize=lambda chunks: {"items": chunks},
)
if inspect.isawaitable(result):
result = await result
if isinstance(result, StreamExecution):
return result
raise AstrBotError.protocol_error(
"stream=true 的插件 capability 必须返回 async generator 或 StreamExecution"
)
if inspect.isasyncgen(result):
raise AstrBotError.protocol_error(
"stream=false 的插件 capability 不能返回 async generator"
)
if inspect.isawaitable(result):
result = await result
return self._normalize_output(result)
def _build_args(
self,
handler,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
*,
plugin_id: str | None = None,
capability_name: str | None = None,
) -> list[Any]:
signature = inspect.signature(handler)
args: list[Any] = []
type_hints: dict[str, Any] = {}
try:
type_hints = get_type_hints(handler)
except Exception:
pass
for parameter in signature.parameters.values():
if parameter.kind not in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
):
continue
injected = None
param_type = type_hints.get(parameter.name)
if param_type is not None:
injected = self._inject_by_type(param_type, payload, ctx, cancel_token)
if injected is None:
if parameter.name in {"ctx", "context"}:
injected = ctx
elif parameter.name in {"payload", "input", "data"}:
injected = payload
elif parameter.name in {"cancel_token", "token"}:
injected = cancel_token
if injected is None:
if parameter.default is not parameter.empty:
continue
raise TypeError(
self._format_capability_injection_error(
handler=handler,
parameter_name=parameter.name,
plugin_id=plugin_id,
capability_name=capability_name,
payload=payload,
)
)
args.append(injected)
return args
def _inject_by_type(
self,
param_type: Any,
payload: dict[str, Any],
ctx: Context,
cancel_token: CancelToken,
) -> Any:
origin = typing.get_origin(param_type)
if origin is typing.Union:
type_args = typing.get_args(param_type)
non_none_types = [item for item in type_args if item is not type(None)]
if len(non_none_types) == 1:
param_type = non_none_types[0]
origin = typing.get_origin(param_type)
if param_type is Context or (
isinstance(param_type, type) and issubclass(param_type, Context)
):
return ctx
if param_type is CancelToken or (
isinstance(param_type, type) and issubclass(param_type, CancelToken)
):
return cancel_token
if param_type is dict or origin is dict:
return payload
return None
def _format_capability_injection_error(
self,
*,
handler,
parameter_name: str,
plugin_id: str | None,
capability_name: str | None,
payload: dict[str, Any],
) -> str:
plugin_text = plugin_id or self._plugin_id
target = capability_name or getattr(handler, "__name__", "<anonymous>")
payload_keys = sorted(str(key) for key in payload.keys())
payload_keys_text = ", ".join(payload_keys) if payload_keys else "<none>"
return (
f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:"
f"必填参数 '{parameter_name}' 无法注入。"
f"签名: {getattr(handler, '__name__', '<anonymous>')}"
f"{HandlerDispatcher._callable_signature(handler)}。"
"当前支持按类型注入 Context / CancelToken / dict,"
"按参数名注入 ctx / context / payload / input / data / cancel_token / token,"
f"以及 payload 中现有键:{payload_keys_text}。"
)
async def _iterate_generator(
self,
generator: AsyncIterator[Any],
) -> AsyncIterator[dict[str, Any]]:
async for item in generator:
yield self._normalize_chunk(item)
def _normalize_chunk(self, item: Any) -> dict[str, Any]:
output = self._normalize_output(item)
if output:
return output
return {"ok": True}
def _normalize_output(self, result: Any) -> dict[str, Any]:
if result is None:
return {}
if isinstance(result, dict):
return result
model_dump = getattr(result, "model_dump", None)
if callable(model_dump):
dumped = model_dump()
if isinstance(dumped, dict):
return dumped
raise AstrBotError.invalid_input("插件 capability 必须返回 dict 或可序列化对象")
__all__ = ["CapabilityDispatcher", "HandlerDispatcher"]
+124 -1
View File
@@ -59,6 +59,7 @@ import os
import re
import shutil
import sys
import typing
from dataclasses import dataclass, field
from importlib import import_module
from pathlib import Path
@@ -67,7 +68,14 @@ from typing import Any
import yaml
from ..decorators import get_capability_meta, get_handler_meta
from ..protocol.descriptors import CapabilityDescriptor, HandlerDescriptor
from ..protocol.descriptors import (
CapabilityDescriptor,
HandlerDescriptor,
ParamSpec,
ScheduleTrigger,
)
from ..schedule import ScheduleContext
from ..types import GreedyStr
from .environment_groups import (
EnvironmentGroup,
EnvironmentPlanner,
@@ -113,6 +121,7 @@ class LoadedHandler:
callable: Any
owner: Any
plugin_id: str = ""
local_filters: list[Any] = field(default_factory=list)
@dataclass(slots=True)
@@ -152,6 +161,107 @@ def _iter_discoverable_names(instance: Any) -> list[str]:
return [*handler_names, *extra_names]
def _unwrap_optional(annotation: Any) -> tuple[Any, bool]:
origin = typing.get_origin(annotation)
if origin is typing.Union:
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
def _is_injected_parameter(annotation: Any, parameter_name: str) -> bool:
if parameter_name in {"event", "ctx", "context", "sched", "schedule"}:
return True
normalized, _is_optional = _unwrap_optional(annotation)
if normalized is None:
return False
if normalized in {ScheduleContext}:
return True
if isinstance(normalized, type):
from ..context import Context
from ..events import MessageEvent
return issubclass(normalized, (Context, MessageEvent, ScheduleContext))
return False
def _param_type_name(annotation: Any) -> tuple[str, str | None, bool]:
normalized, is_optional = _unwrap_optional(annotation)
if normalized is GreedyStr:
return "greedy_str", None, False
if normalized in {int, float, bool, str}:
if is_optional:
return "optional", normalized.__name__, False
return normalized.__name__, None, True
if is_optional:
return "optional", "str", False
return "str", None, True
def _build_param_specs(handler: Any) -> list[ParamSpec]:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return []
try:
type_hints = typing.get_type_hints(handler)
except Exception:
type_hints = {}
specs: list[ParamSpec] = []
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(annotation, parameter.name):
continue
param_type, inner_type, required = _param_type_name(annotation)
if parameter.default is not inspect.Parameter.empty:
required = False
specs.append(
ParamSpec(
name=parameter.name,
type=param_type,
required=required,
inner_type=inner_type,
)
)
greedy_indexes = [
index for index, spec in enumerate(specs) if spec.type == "greedy_str"
]
if greedy_indexes and greedy_indexes[-1] != len(specs) - 1:
greedy_spec = specs[greedy_indexes[-1]]
raise ValueError(f"参数 '{greedy_spec.name}' (GreedyStr) 必须是最后一个参数。")
return specs
def _validate_schedule_signature(handler: Any) -> None:
try:
signature = inspect.signature(handler)
except (TypeError, ValueError):
return
allowed_names = {"ctx", "context", "sched", "schedule"}
invalid = [
parameter.name
for parameter in signature.parameters.values()
if parameter.kind
in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
and parameter.name not in allowed_names
]
if invalid:
raise ValueError(
"Schedule handler 只允许注入 ctx/context 和 sched/schedule 参数。"
)
def _plugin_context(plugin: PluginSpec) -> str:
return f"插件 '{plugin.name}'({plugin.manifest_path})"
@@ -626,6 +736,9 @@ def load_plugin(plugin: PluginSpec) -> LoadedPlugin:
bound, meta = resolved
handler_id = f"{plugin.name}:{instance.__class__.__module__}.{instance.__class__.__name__}.{name}"
if isinstance(meta.trigger, ScheduleTrigger):
_validate_schedule_signature(bound)
param_specs = _build_param_specs(bound)
handlers.append(
LoadedHandler(
descriptor=HandlerDescriptor(
@@ -635,10 +748,20 @@ def load_plugin(plugin: PluginSpec) -> LoadedPlugin:
contract=meta.contract,
priority=meta.priority,
permissions=meta.permissions.model_copy(deep=True),
filters=[item.model_copy(deep=True) for item in meta.filters],
param_specs=[
item.model_copy(deep=True) for item in param_specs
],
command_route=(
meta.command_route.model_copy(deep=True)
if meta.command_route is not None
else None
),
),
callable=bound,
owner=instance,
plugin_id=plugin.name,
local_filters=list(meta.local_filters),
)
)
+6 -13
View File
@@ -5,7 +5,7 @@
而不是业务上的用户、群聊或会话对象。
核心职责:
- 通过可插拔 codec 做消息编解码
- 消息序列化/反序列化
- 初始化握手协议
- 能力调用(同步/流式)
- 取消处理
@@ -69,7 +69,7 @@
- 入站任务在收到 CancelMessage 时被取消
- 早到取消:在任务执行前检查 cancel_token,避免竞态条件
`Peer` 把 `Transport`、wire codec 和 v4 协议消息模型接起来,负责:
`Peer` 把 `Transport` 和 v4 协议消息模型接起来,负责:
- 握手与远端元数据缓存
- 请求 ID 关联
@@ -100,8 +100,8 @@ from ..protocol.messages import (
InvokeMessage,
PeerInfo,
ResultMessage,
parse_message,
)
from ..protocol.wire_codecs import JsonProtocolCodec, ProtocolCodec
from .capability_router import StreamExecution
InitializeHandler = Callable[[InitializeMessage], Awaitable[InitializeOutput]]
@@ -181,7 +181,6 @@ class Peer:
peer_info: PeerInfo,
protocol_version: str = "1.0",
supported_protocol_versions: Sequence[str] | None = None,
codec: ProtocolCodec | None = None,
) -> None:
"""创建一个协议对等端实例。
@@ -193,10 +192,6 @@ class Peer:
"""
self.transport = transport
self.peer_info = peer_info
self.codec = codec or JsonProtocolCodec()
configure_for_codec = getattr(self.transport, "configure_for_codec", None)
if callable(configure_for_codec):
configure_for_codec(self.codec)
self.protocol_version = protocol_version
self.supported_protocol_versions = _dedupe_protocol_versions(
supported_protocol_versions,
@@ -350,7 +345,6 @@ class Peer:
asyncio.get_running_loop().create_future()
)
self._pending_results[request_id] = future
# FIXME: 这里会输出乱七八糟的各种东西
await self._send(
InitializeMessage(
id=request_id,
@@ -361,7 +355,6 @@ class Peer:
metadata=handshake_metadata,
)
)
# FIXME: 👆会输出各种乱七八糟的东西
result = await future
if result.kind != "initialize_result":
raise AstrBotError.protocol_error("initialize 必须收到 initialize_result")
@@ -507,10 +500,10 @@ class Peer:
if self._unusable:
raise AstrBotError.protocol_error("连接已进入不可用状态")
async def _handle_raw_message(self, payload: bytes) -> None:
async def _handle_raw_message(self, payload: str) -> None:
"""解析原始消息并分发到对应的消息处理分支。"""
try:
message = self.codec.decode_message(payload)
message = parse_message(payload)
if isinstance(message, ResultMessage):
await self._handle_result(message)
return
@@ -759,4 +752,4 @@ class Peer:
async def _send(self, message) -> None:
"""序列化协议消息并通过底层传输发送出去。"""
await self.transport.send(self.codec.encode_message(message))
await self.transport.send(message.model_dump_json(exclude_none=True))
+12 -41
View File
@@ -41,14 +41,13 @@ import signal
import sys
from collections.abc import Callable
from pathlib import Path
from typing import IO, Any, cast
from typing import IO, Any
from loguru import logger
from ..errors import AstrBotError
from ..protocol.descriptors import CapabilityDescriptor
from ..protocol.messages import EventMessage, InitializeOutput, PeerInfo
from ..protocol.wire_codecs import ProtocolCodec, make_protocol_codec
from .capability_router import CapabilityRouter, StreamExecution
from .environment_groups import EnvironmentGroup
from .loader import (
@@ -80,15 +79,13 @@ def _install_signal_handlers(stop_event: asyncio.Event) -> None:
def _prepare_stdio_transport(
stdin: IO[str] | IO[bytes] | None,
stdout: IO[str] | IO[bytes] | None,
*,
binary: bool = False,
) -> tuple[IO[str] | IO[bytes], IO[str] | IO[bytes], IO[str] | None]:
stdin: IO[str] | None,
stdout: IO[str] | None,
) -> tuple[IO[str], IO[str], IO[str] | None]:
if stdin is not None and stdout is not None:
return stdin, stdout, None
transport_stdin = stdin or (sys.stdin.buffer if binary else sys.stdin)
transport_stdout = stdout or (sys.stdout.buffer if binary else sys.stdout)
transport_stdin = stdin or sys.stdin
transport_stdout = stdout or sys.stdout
original_stdout = sys.stdout
sys.stdout = sys.stderr
return transport_stdin, transport_stdout, original_stdout
@@ -131,25 +128,17 @@ class WorkerSession:
env_manager: PluginEnvironmentManager,
capability_router: CapabilityRouter,
on_closed: Callable[[], None] | None = None,
codec: ProtocolCodec | None = None,
wire_codec_name: str = "json",
) -> None:
if plugin is None and group is None:
raise ValueError("WorkerSession requires either plugin or group")
if group is None and plugin is None:
raise ValueError("WorkerSession requires a plugin when group is absent")
self.group = group
self.plugins = (
list(group.plugins) if group is not None else [cast(PluginSpec, plugin)]
)
self.plugins = list(group.plugins) if group is not None else [plugin]
self.plugin = plugin or self.plugins[0]
self.group_id = group.id if group is not None else self.plugin.name
self.repo_root = repo_root.resolve()
self.env_manager = env_manager
self.capability_router = capability_router
self.on_closed = on_closed
self.codec = codec or make_protocol_codec(wire_codec_name)
self.wire_codec_name = self.codec.name
self.peer: Peer | None = None
self.handlers = []
self.provided_capabilities: list[CapabilityDescriptor] = []
@@ -175,12 +164,10 @@ class WorkerSession:
command=command,
cwd=cwd,
env=env,
framing=self.codec.stdio_framing,
)
self.peer = Peer(
transport=transport,
peer_info=PeerInfo(name="astrbot-core", role="core", version="v4"),
codec=self.codec,
)
self.peer.set_initialize_handler(self._handle_initialize)
self.peer.set_invoke_handler(self._handle_capability_invoke)
@@ -238,7 +225,7 @@ class WorkerSession:
if self.group is not None:
prepare_group = getattr(self.env_manager, "prepare_group_environment", None)
if callable(prepare_group):
python_path = cast(Path, prepare_group(self.group))
python_path = prepare_group(self.group)
else:
python_path = self.env_manager.prepare_environment(self.plugins[0])
return (
@@ -248,16 +235,13 @@ class WorkerSession:
"-m",
"astrbot_sdk",
"worker",
"--wire-codec",
self.wire_codec_name,
"--group-metadata",
str(self.group.metadata_path),
],
str(self.repo_root),
)
plugin = self.plugin
python_path = self.env_manager.prepare_environment(plugin)
python_path = self.env_manager.prepare_environment(self.plugin)
return (
python_path,
[
@@ -265,12 +249,10 @@ class WorkerSession:
"-m",
"astrbot_sdk",
"worker",
"--wire-codec",
self.wire_codec_name,
"--plugin-dir",
str(plugin.plugin_dir),
str(self.plugin.plugin_dir),
],
str(plugin.plugin_dir),
str(self.plugin.plugin_dir),
)
def start_close_watch(self) -> None:
@@ -396,20 +378,15 @@ class SupervisorRuntime:
transport,
plugins_dir: Path,
env_manager: PluginEnvironmentManager | None = None,
codec: ProtocolCodec | None = None,
worker_wire_codec_name: str = "json",
) -> None:
self.transport = transport
self.plugins_dir = plugins_dir.resolve()
self.repo_root = Path(__file__).resolve().parents[3]
self.env_manager = env_manager or PluginEnvironmentManager(self.repo_root)
self.codec = codec or make_protocol_codec("json")
self.worker_wire_codec_name = worker_wire_codec_name
self.capability_router = CapabilityRouter()
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name="astrbot-supervisor", role="plugin", version="v4"),
codec=self.codec,
)
self.peer.set_invoke_handler(self._handle_upstream_invoke)
self.peer.set_cancel_handler(self._handle_upstream_cancel)
@@ -635,9 +612,6 @@ class SupervisorRuntime:
discovery = discover_plugins(self.plugins_dir)
self.skipped_plugins = dict(discovery.skipped_plugins)
plan_result = self.env_manager.plan(discovery.plugins)
logger.info(
f"发现 {len(discovery.plugins)} 个插件,{len(plan_result.groups)} 个环境组"
)
self.skipped_plugins.update(plan_result.skipped_plugins)
self._sync_plugin_registry(discovery.plugins)
try:
@@ -650,7 +624,6 @@ class SupervisorRuntime:
repo_root=self.repo_root,
env_manager=self.env_manager,
capability_router=self.capability_router,
wire_codec_name=self.worker_wire_codec_name,
on_closed=lambda group_id=group.id: (
self._handle_worker_closed(group_id)
),
@@ -664,7 +637,6 @@ class SupervisorRuntime:
repo_root=self.repo_root,
env_manager=self.env_manager,
capability_router=self.capability_router,
wire_codec_name=self.worker_wire_codec_name,
on_closed=lambda plugin_name=plugin.name: (
self._handle_worker_closed(plugin_name)
),
@@ -704,8 +676,7 @@ class SupervisorRuntime:
aggregated_handlers = list(self.handler_to_worker.keys())
logger.info(
"Loaded plugins: \n{}",
"\n ".join(sorted(self.loaded_plugins)) or "none",
"Loaded plugins: {}", ", ".join(sorted(self.loaded_plugins)) or "none"
)
await self.peer.start()
+102 -150
View File
@@ -1,7 +1,7 @@
"""传输层抽象模块。
定义 Transport 抽象基类及其实现,负责底层原始载荷的传输。
传输层只关心分帧后的 bytes 或 text frame,不处理协议细节。
定义 Transport 抽象基类及其实现,负责底层的消息传输。
传输层只关心"发送字符串"和"接收字符串",不处理协议细节。
传输实现:
Transport: 抽象基类,定义 start/stop/send/wait_closed 接口
StdioTransport: 标准输入输出传输
@@ -37,7 +37,7 @@
- 支持心跳配置
- WebSocketClientTransport:
- 自动重连需要外部实现
- 传输层只处理 framed payload,协议由 Peer 层处理
- 传输层只处理字符串,协议由 Peer 层处理
使用示例:
# 子进程模式
@@ -58,15 +58,15 @@
# 统一接口
transport.set_message_handler(my_handler)
await transport.start()
await transport.send(encoded_payload)
await transport.send(json_string)
await transport.stop()
`Transport` 只处理 framed payload,不做协议解析,也不关心能力、handler 或
legacy 兼容。当前实现包括:
`Transport` 只处理“字符串发出去 / 字符串收进来”这件事,不做协议解析,也不关心
能力、handler 或迁移适配策略。当前实现包括:
- `StdioTransport`: 子进程或文件对象上的按行或 length-prefixed 传输
- `WebSocketServerTransport`: 单连接 WebSocket 服务端,支持 text/binary frame
- `WebSocketClientTransport`: WebSocket 客户端,支持 text/binary frame
- `StdioTransport`: 子进程或文件对象上的按行文本传输
- `WebSocketServerTransport`: 单连接 WebSocket 服务端
- `WebSocketClientTransport`: WebSocket 客户端
自动重连、消息重放等策略不在这里实现,统一留给更上层编排。
"""
@@ -74,57 +74,45 @@ legacy 兼容。当前实现包括:
from __future__ import annotations
import asyncio
import io
import struct
import sys
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable, Sequence
from typing import IO, cast
from typing import IO, Any
import aiohttp
from aiohttp import web
from loguru import logger
from ..protocol.wire_codecs import StdioFraming, WebSocketFrameType
MessageHandler = Callable[[bytes], Awaitable[None]]
RawPayload = bytes | str
MessageHandler = Callable[[str], Awaitable[None]]
def _ensure_bytes(payload: RawPayload) -> bytes:
if isinstance(payload, bytes):
return payload
return payload.encode("utf-8")
def _get_aiohttp():
import aiohttp
return aiohttp
def _frame_stdio_line_payload(payload: bytes) -> bytes:
def _get_web():
from aiohttp import web
return web
def _frame_stdio_payload(payload: str) -> str:
body = payload
if body.endswith(b"\r\n"):
if body.endswith("\r\n"):
body = body[:-2]
elif body.endswith((b"\n", b"\r")):
elif body.endswith(("\n", "\r")):
body = body[:-1]
if b"\n" in body or b"\r" in body:
if "\n" in body or "\r" in body:
raise ValueError("STDIO payload 不允许包含原始换行符")
return body + b"\n"
return f"{body}\n"
def _frame_stdio_length_prefixed_payload(payload: bytes) -> bytes:
return struct.pack(">I", len(payload)) + payload
def _write_stdio_payload(stream: IO[str] | IO[bytes], payload: bytes) -> None:
if hasattr(stream, "buffer"):
stream.buffer.write(payload) # type: ignore[attr-defined]
stream.flush() # type: ignore[call-arg]
return
if isinstance(stream, io.TextIOBase):
text_stream = cast(IO[str], stream)
text_stream.write(payload.decode("utf-8"))
text_stream.flush()
return
binary_stream = cast(IO[bytes], stream)
binary_stream.write(payload)
binary_stream.flush()
#TODO 一个更好的解决方案?
def _is_windows_access_denied(error: BaseException) -> bool:
return (
sys.platform == "win32"
and isinstance(error, PermissionError)
and getattr(error, "winerror", None) == 5
)
class Transport(ABC):
@@ -136,10 +124,6 @@ class Transport(ABC):
"""注册收到原始字符串消息后的回调。"""
self._handler = handler
def configure_for_codec(self, codec) -> None:
"""Allow transports to align framing or frame type with the selected codec."""
return None
@abstractmethod
async def start(self) -> None:
raise NotImplementedError
@@ -149,14 +133,14 @@ class Transport(ABC):
raise NotImplementedError
@abstractmethod
async def send(self, payload: RawPayload) -> None:
async def send(self, payload: str) -> None:
raise NotImplementedError
async def wait_closed(self) -> None:
"""等待传输层进入关闭状态。"""
await self._closed.wait()
async def _dispatch(self, payload: bytes) -> None:
async def _dispatch(self, payload: str) -> None:
"""把收到的原始载荷转交给上层处理器。"""
if self._handler is not None:
await self._handler(payload)
@@ -166,12 +150,11 @@ class StdioTransport(Transport):
def __init__(
self,
*,
stdin: IO[str] | IO[bytes] | None = None,
stdout: IO[str] | IO[bytes] | None = None,
stdin: IO[str] | None = None,
stdout: IO[str] | None = None,
command: Sequence[str] | None = None,
cwd: str | None = None,
env: dict[str, str] | None = None,
framing: StdioFraming = "line",
) -> None:
super().__init__()
self._stdin = stdin
@@ -179,34 +162,48 @@ class StdioTransport(Transport):
self._command = list(command) if command is not None else None
self._cwd = cwd
self._env = env
self._framing = framing
self._process: asyncio.subprocess.Process | None = None
self._reader_task: asyncio.Task[None] | None = None
async def start(self) -> None:
self._closed.clear()
if self._command is not None:
self._process = await asyncio.create_subprocess_exec(
*self._command,
cwd=self._cwd,
env=self._env,
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.PIPE,
stderr=sys.stderr,
)
self._process = await self._start_subprocess_with_retry()
self._reader_task = asyncio.create_task(self._read_process_loop())
return
if self._framing == "length_prefixed":
self._stdin = self._stdin or sys.stdin.buffer
self._stdout = self._stdout or sys.stdout.buffer
else:
self._stdin = self._stdin or sys.stdin
self._stdout = self._stdout or sys.stdout
self._stdin = self._stdin or sys.stdin
self._stdout = self._stdout or sys.stdout
self._reader_task = asyncio.create_task(self._read_file_loop())
def configure_for_codec(self, codec) -> None:
self._framing = codec.stdio_framing
async def _start_subprocess_with_retry(self) -> asyncio.subprocess.Process:
delays = [0.15, 0.35, 0.75]
last_error: BaseException | None = None
for attempt, delay in enumerate([0.0, *delays], start=1):
if delay:
await asyncio.sleep(delay)
try:
return await asyncio.create_subprocess_exec(
*self._command,
cwd=self._cwd,
env=self._env,
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.PIPE,
stderr=sys.stderr,
)
except Exception as exc:
last_error = exc
if not _is_windows_access_denied(exc) or attempt == len(delays) + 1:
raise
logger.warning(
"Windows denied access while starting freshly prepared worker "
"interpreter, retrying attempt {}/{}: {}",
attempt,
len(delays) + 1,
exc,
)
assert last_error is not None
raise last_error
async def stop(self) -> None:
if self._reader_task is not None:
@@ -228,16 +225,12 @@ class StdioTransport(Transport):
self._process = None
self._closed.set()
async def send(self, payload: RawPayload) -> None:
encoded = _ensure_bytes(payload)
if self._framing == "line":
framed = _frame_stdio_line_payload(encoded)
else:
framed = _frame_stdio_length_prefixed_payload(encoded)
async def send(self, payload: str) -> None:
line = _frame_stdio_payload(payload)
if self._process is not None:
if self._process.stdin is None:
raise RuntimeError("STDIO subprocess stdin 不可用")
self._process.stdin.write(framed)
self._process.stdin.write(line.encode("utf-8"))
await self._process.stdin.drain()
return
@@ -246,7 +239,8 @@ class StdioTransport(Transport):
def _write() -> None:
assert self._stdout is not None
_write_stdio_payload(self._stdout, framed)
self._stdout.write(line)
self._stdout.flush()
await asyncio.to_thread(_write)
@@ -255,18 +249,10 @@ class StdioTransport(Transport):
assert self._process.stdout is not None
try:
while True:
if self._framing == "line":
raw = await self._process.stdout.readline()
if not raw:
break
await self._dispatch(raw.rstrip(b"\r\n"))
continue
header = await self._process.stdout.readexactly(4)
length = struct.unpack(">I", header)[0]
payload = await self._process.stdout.readexactly(length)
await self._dispatch(payload)
except asyncio.IncompleteReadError:
pass
raw = await self._process.stdout.readline()
if not raw:
break
await self._dispatch(raw.decode("utf-8").rstrip("\r\n"))
finally:
self._closed.set()
@@ -274,29 +260,10 @@ class StdioTransport(Transport):
assert self._stdin is not None
try:
while True:
if self._framing == "line":
raw = await asyncio.to_thread(self._stdin.readline)
if not raw:
break
if isinstance(raw, bytes):
await self._dispatch(raw.rstrip(b"\r\n"))
else:
await self._dispatch(raw.rstrip("\r\n").encode("utf-8"))
continue
header = await asyncio.to_thread(self._stdin.read, 4)
if not header:
raw = await asyncio.to_thread(self._stdin.readline)
if not raw:
break
if isinstance(header, str):
raise RuntimeError("length_prefixed STDIO 需要二进制 stdin")
if len(header) < 4:
break
length = struct.unpack(">I", header)[0]
payload = await asyncio.to_thread(self._stdin.read, length)
if isinstance(payload, str):
raise RuntimeError("length_prefixed STDIO 需要二进制 stdin")
if len(payload) < length:
break
await self._dispatch(payload)
await self._dispatch(raw.rstrip("\r\n"))
finally:
self._closed.set()
@@ -309,7 +276,6 @@ class WebSocketServerTransport(Transport):
port: int = 8765,
path: str = "/",
heartbeat: float = 30.0,
frame_type: WebSocketFrameType = "text",
) -> None:
super().__init__()
self._host = host
@@ -317,15 +283,15 @@ class WebSocketServerTransport(Transport):
self._actual_port: int | None = None
self._path = path
self._heartbeat = heartbeat
self._frame_type = frame_type
self._app: web.Application | None = None
self._runner: web.AppRunner | None = None
self._site: web.TCPSite | None = None
self._ws: web.WebSocketResponse | None = None
self._app: Any | None = None
self._runner: Any | None = None
self._site: Any | None = None
self._ws: Any | None = None
self._write_lock = asyncio.Lock()
self._connected = asyncio.Event()
async def start(self) -> None:
web = _get_web()
self._closed.clear()
self._connected.clear()
self._app = web.Application()
@@ -334,10 +300,8 @@ class WebSocketServerTransport(Transport):
await self._runner.setup()
self._site = web.TCPSite(self._runner, self._host, self._port)
await self._site.start()
server = getattr(self._site, "_server", None)
sockets = getattr(server, "sockets", None)
if sockets:
socket = sockets[0]
if self._site._server and getattr(self._site._server, "sockets", None):
socket = self._site._server.sockets[0]
self._actual_port = socket.getsockname()[1]
async def stop(self) -> None:
@@ -352,19 +316,17 @@ class WebSocketServerTransport(Transport):
self._runner = None
self._closed.set()
async def send(self, payload: RawPayload) -> None:
async def send(self, payload: str) -> None:
if self._ws is None or self._ws.closed:
await asyncio.wait_for(self._connected.wait(), timeout=30.0)
if self._ws is None or self._ws.closed:
raise RuntimeError("WebSocket 尚未连接")
async with self._write_lock:
encoded = _ensure_bytes(payload)
if self._frame_type == "text":
await self._ws.send_str(encoded.decode("utf-8"))
else:
await self._ws.send_bytes(encoded)
await self._ws.send_str(payload)
async def _handle_socket(self, request: web.Request) -> web.WebSocketResponse:
async def _handle_socket(self, request) -> Any:
web = _get_web()
aiohttp = _get_aiohttp()
if self._ws is not None and not self._ws.closed:
ws = web.WebSocketResponse()
await ws.prepare(request)
@@ -380,9 +342,9 @@ class WebSocketServerTransport(Transport):
try:
async for msg in ws:
if msg.type == aiohttp.WSMsgType.TEXT:
await self._dispatch(msg.data.encode("utf-8"))
elif msg.type == aiohttp.WSMsgType.BINARY:
await self._dispatch(msg.data)
elif msg.type == aiohttp.WSMsgType.BINARY:
await self._dispatch(msg.data.decode("utf-8"))
elif msg.type == aiohttp.WSMsgType.ERROR:
logger.error("websocket server error: {}", ws.exception())
break
@@ -400,9 +362,6 @@ class WebSocketServerTransport(Transport):
def url(self) -> str:
return f"ws://{self._host}:{self.port}{self._path}"
def configure_for_codec(self, codec) -> None:
self._frame_type = codec.websocket_frame_type
class WebSocketClientTransport(Transport):
def __init__(
@@ -410,17 +369,16 @@ class WebSocketClientTransport(Transport):
*,
url: str,
heartbeat: float = 30.0,
frame_type: WebSocketFrameType = "text",
) -> None:
super().__init__()
self._url = url
self._heartbeat = heartbeat
self._frame_type = frame_type
self._session: aiohttp.ClientSession | None = None
self._ws: aiohttp.ClientWebSocketResponse | None = None
self._session: Any | None = None
self._ws: Any | None = None
self._reader_task: asyncio.Task[None] | None = None
async def start(self) -> None:
aiohttp = _get_aiohttp()
self._closed.clear()
self._session = aiohttp.ClientSession()
self._ws = await self._session.ws_connect(
@@ -445,28 +403,22 @@ class WebSocketClientTransport(Transport):
self._session = None
self._closed.set()
async def send(self, payload: RawPayload) -> None:
async def send(self, payload: str) -> None:
if self._ws is None or self._ws.closed:
raise RuntimeError("WebSocket client 尚未连接")
encoded = _ensure_bytes(payload)
if self._frame_type == "text":
await self._ws.send_str(encoded.decode("utf-8"))
else:
await self._ws.send_bytes(encoded)
await self._ws.send_str(payload)
async def _read_loop(self) -> None:
assert self._ws is not None
aiohttp = _get_aiohttp()
try:
async for msg in self._ws:
if msg.type == aiohttp.WSMsgType.TEXT:
await self._dispatch(msg.data.encode("utf-8"))
elif msg.type == aiohttp.WSMsgType.BINARY:
await self._dispatch(msg.data)
elif msg.type == aiohttp.WSMsgType.BINARY:
await self._dispatch(msg.data.decode("utf-8"))
elif msg.type == aiohttp.WSMsgType.ERROR:
logger.error("websocket client error: {}", self._ws.exception())
break
finally:
self._closed.set()
def configure_for_codec(self, codec) -> None:
self._frame_type = codec.websocket_frame_type
+2 -21
View File
@@ -37,7 +37,6 @@ from .._invocation_context import caller_plugin_scope
from ..context import Context as RuntimeContext
from ..errors import AstrBotError
from ..protocol.messages import PeerInfo
from ..protocol.wire_codecs import ProtocolCodec, make_protocol_codec
from .handler_dispatcher import CapabilityDispatcher, HandlerDispatcher
from .loader import (
LoadedPlugin,
@@ -115,22 +114,13 @@ async def run_plugin_lifecycle(
class GroupWorkerRuntime:
def __init__(
self,
*,
group_metadata_path: Path,
transport,
codec: ProtocolCodec | None = None,
wire_codec_name: str = "json",
) -> None:
def __init__(self, *, group_metadata_path: Path, transport) -> None:
self.group_metadata_path = group_metadata_path.resolve()
self.group_id, self.plugins = _load_group_plugin_specs(self.group_metadata_path)
self.transport = transport
self.codec = codec or make_protocol_codec(wire_codec_name)
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name=self.group_id, role="plugin", version="v4"),
codec=self.codec,
)
self.skipped_plugins: dict[str, str] = {}
self._plugin_states: list[GroupPluginRuntimeState] = []
@@ -294,22 +284,13 @@ class GroupWorkerRuntime:
class PluginWorkerRuntime:
def __init__(
self,
*,
plugin_dir: Path,
transport,
codec: ProtocolCodec | None = None,
wire_codec_name: str = "json",
) -> None:
def __init__(self, *, plugin_dir: Path, transport) -> None:
self.plugin = load_plugin_spec(plugin_dir)
self.transport = transport
self.codec = codec or make_protocol_codec(wire_codec_name)
self.loaded_plugin = load_plugin(self.plugin)
self.peer = Peer(
transport=self.transport,
peer_info=PeerInfo(name=self.plugin.name, role="plugin", version="v4"),
codec=self.codec,
)
self.dispatcher = HandlerDispatcher(
plugin_id=self.plugin.name,
+60
View File
@@ -0,0 +1,60 @@
"""Schedule-specific SDK types.
本模块定义定时任务相关的 SDK 类型,主要为 ScheduleContext 提供数据结构。
ScheduleContext 包含:
- schedule_id: 调度任务唯一标识
- plugin_id: 所属插件 ID
- handler_id: 对应 handler 的标识
- trigger_kind: 触发类型(cron / interval / once)
- cron: cron 表达式(仅 cron 类型)
- interval_seconds: 间隔秒数(仅 interval 类型)
- scheduled_at: 计划执行时间(仅 once 类型)
使用方式:
通过 @on_schedule 装饰器注册的 handler 可通过参数注入获取 ScheduleContext。
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
@dataclass(slots=True)
class ScheduleContext:
schedule_id: str
plugin_id: str
handler_id: str
trigger_kind: str
cron: str | None = None
interval_seconds: int | None = None
scheduled_at: str | None = None
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> ScheduleContext:
schedule = payload.get("schedule")
if not isinstance(schedule, dict):
raise ValueError("schedule payload is required")
return cls(
schedule_id=str(schedule.get("schedule_id", "")),
plugin_id=str(schedule.get("plugin_id", "")),
handler_id=str(schedule.get("handler_id", "")),
trigger_kind=str(schedule.get("trigger_kind", "")),
cron=(
str(schedule["cron"]) if isinstance(schedule.get("cron"), str) else None
),
interval_seconds=(
int(schedule["interval_seconds"])
if isinstance(schedule.get("interval_seconds"), int)
else None
),
scheduled_at=(
str(schedule["scheduled_at"])
if isinstance(schedule.get("scheduled_at"), str)
else None
),
)
__all__ = ["ScheduleContext"]
+239
View File
@@ -0,0 +1,239 @@
"""Session-based conversational flow management.
本模块实现会话等待器 (session_waiter),用于构建多轮对话流程。
核心组件:
- SessionController: 控制会话生命周期,支持超时管理、会话保持、历史记录
- SessionWaiterManager: 管理活跃的会话等待器,处理事件分发和注册/注销
- @session_waiter 装饰器: 将普通 handler 转换为会话式 handler
使用场景:
当需要在用户首次触发后继续监听后续消息(如分步表单、问答游戏),
可使用 @session_waiter 装饰器自动管理会话状态和超时。
注意事项:
在当前桥接设计中,不应在普通 SDK handler 内直接 await session_waiter,
这会导致首次 dispatch 保持打开直到下一条消息到达。
如需非阻塞的会话等待,应从后台任务启动或添加显式的调度/恢复机制。
"""
from __future__ import annotations
import asyncio
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
from typing import Any
from loguru import logger
from .events import MessageEvent
@dataclass(slots=True)
class SessionController:
future: asyncio.Future[Any] = field(default_factory=asyncio.Future)
current_event: asyncio.Event | None = None
ts: float | None = None
timeout: float | None = None
history_chains: list[list[dict[str, Any]]] = field(default_factory=list)
def stop(self, error: Exception | None = None) -> None:
if self.future.done():
return
if error is not None:
self.future.set_exception(error)
else:
self.future.set_result(None)
def keep(self, timeout: float = 0, reset_timeout: bool = False) -> None:
new_ts = time.time()
if reset_timeout:
if timeout <= 0:
self.stop()
return
else:
assert self.timeout is not None
assert self.ts is not None
left_timeout = self.timeout - (new_ts - self.ts)
timeout = left_timeout + timeout
if timeout <= 0:
self.stop()
return
if self.current_event and not self.current_event.is_set():
self.current_event.set()
current_event = asyncio.Event()
self.current_event = current_event
self.ts = new_ts
self.timeout = timeout
asyncio.create_task(self._holding(current_event, timeout))
async def _holding(self, event: asyncio.Event, timeout: float) -> None:
try:
await asyncio.wait_for(event.wait(), timeout)
except asyncio.TimeoutError as exc:
self.stop(exc)
except asyncio.CancelledError:
return
def get_history_chains(self) -> list[list[dict[str, Any]]]:
return list(self.history_chains)
@dataclass(slots=True)
class _WaiterEntry:
session_key: str
handler: Callable[[SessionController, MessageEvent], Awaitable[Any]]
controller: SessionController
record_history_chains: bool
class SessionWaiterManager:
def __init__(self, *, plugin_id: str, peer) -> None:
self._plugin_id = plugin_id
self._peer = peer
self._entries: dict[str, _WaiterEntry] = {}
self._locks: dict[str, asyncio.Lock] = {}
async def register(
self,
*,
event: MessageEvent,
handler: Callable[[SessionController, MessageEvent], Awaitable[Any]],
timeout: int,
record_history_chains: bool,
) -> Any:
if event._context is None:
raise RuntimeError("session_waiter requires runtime context")
session_key = event.unified_msg_origin
entry = _WaiterEntry(
session_key=session_key,
handler=handler,
controller=SessionController(),
record_history_chains=record_history_chains,
)
replaced = session_key in self._entries
self._entries[session_key] = entry
self._locks.setdefault(session_key, asyncio.Lock())
if replaced:
logger.warning(
"Session waiter replaced: plugin_id=%s session_key=%s",
self._plugin_id,
session_key,
)
await self._peer.invoke(
"system.session_waiter.register",
{"session_key": session_key},
)
entry.controller.keep(timeout, reset_timeout=True)
try:
return await entry.controller.future
finally:
await self.unregister(session_key)
async def unregister(self, session_key: str) -> None:
self._entries.pop(session_key, None)
self._locks.pop(session_key, None)
try:
await self._peer.invoke(
"system.session_waiter.unregister",
{"session_key": session_key},
)
except Exception:
logger.debug(
"Failed to unregister session waiter: plugin_id=%s session_key=%s",
self._plugin_id,
session_key,
)
def has_waiter(self, event: MessageEvent) -> bool:
return event.unified_msg_origin in self._entries
async def dispatch(self, event: MessageEvent) -> dict[str, Any]:
session_key = event.unified_msg_origin
entry = self._entries.get(session_key)
if entry is None:
return {"sent_message": False, "stop": False, "call_llm": False}
lock = self._locks.setdefault(session_key, asyncio.Lock())
async with lock:
if entry.record_history_chains:
chain = []
raw_chain = (
event.raw.get("chain") if isinstance(event.raw, dict) else None
)
if isinstance(raw_chain, list):
chain = [dict(item) for item in raw_chain if isinstance(item, dict)]
entry.controller.history_chains.append(chain)
await entry.handler(entry.controller, event)
return {
"sent_message": False,
"stop": event.is_stopped(),
"call_llm": False,
}
def session_waiter(
timeout: int = 30,
*,
record_history_chains: bool = False,
):
def decorator(
func: Callable[[SessionController, MessageEvent], Awaitable[Any]],
):
async def wrapper(*args, **kwargs):
owner = None
event: MessageEvent | None = None
trailing_args = ()
if args and isinstance(args[0], MessageEvent):
event = args[0]
trailing_args = args[1:]
elif len(args) >= 2 and isinstance(args[1], MessageEvent):
owner = args[0]
event = args[1]
trailing_args = args[2:]
if event is None:
raise RuntimeError("session_waiter requires a MessageEvent argument")
if event._context is None:
raise RuntimeError("session_waiter requires runtime context")
manager = getattr(event._context.peer, "_session_waiter_manager", None)
if manager is None:
raise RuntimeError("session_waiter manager is unavailable")
if owner is None:
async def bound_handler(
controller: SessionController,
waiter_event: MessageEvent,
) -> Any:
return await func(
controller,
waiter_event,
*trailing_args,
**kwargs,
)
else:
async def bound_handler(
controller: SessionController,
waiter_event: MessageEvent,
) -> Any:
return await func(
owner,
controller,
waiter_event,
*trailing_args,
**kwargs,
)
return await manager.register(
event=event,
handler=bound_handler,
timeout=timeout,
record_history_chains=record_history_chains,
)
return wrapper
return decorator
+88 -495
View File
@@ -18,21 +18,38 @@ import inspect
import re
import shlex
import typing
from dataclasses import dataclass, field
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Mapping, TextIO, get_type_hints
from typing import Any, get_type_hints
from .context import CancelToken, Context as RuntimeContext
from ._testing_support import (
InMemoryDB,
InMemoryMemory,
MockCapabilityRouter,
MockContext,
MockLLMClient,
MockMessageEvent,
MockPeer,
MockPlatformClient,
RecordedSend,
StdoutPlatformSink,
)
from .context import CancelToken
from .context import Context as RuntimeContext
from .errors import AstrBotError
from .events import MessageEvent
from .protocol.descriptors import (
CommandTrigger,
CompositeFilterSpec,
EventTrigger,
LocalFilterRefSpec,
MessageTrigger,
MessageTypeFilterSpec,
PlatformFilterSpec,
ScheduleTrigger,
)
from .protocol.messages import EventMessage, InvokeMessage, PeerInfo
from .runtime.capability_router import CapabilityRouter, StreamExecution
from .protocol.messages import InvokeMessage
from .runtime._streaming import StreamExecution
from .runtime.handler_dispatcher import CapabilityDispatcher, HandlerDispatcher
from .runtime.loader import (
LoadedHandler,
@@ -54,357 +71,6 @@ class _PluginExecutionError(RuntimeError):
"""本地 harness 执行插件代码时的已知插件异常。"""
@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=(
dict(payload.get("target"))
if isinstance(payload.get("target"), dict)
else None
),
raw=dict(payload),
)
class StdoutPlatformSink:
"""把 platform.* 的发送结果同时写到终端与内存记录。"""
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:
"""测试友好的 KV 视图,直接绑定到 mock router 的内存存储。"""
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:
"""测试友好的 memory 视图,保持与 mock router 同步。"""
def __init__(self, store: dict[str, dict[str, Any]]) -> None:
self._store = store
def get(self, key: str, default: Any = None) -> Any:
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)
def search(self, query: str) -> list[dict[str, Any]]:
results: list[dict[str, Any]] = []
for key, value in self._store.items():
if query in key or query in str(value):
results.append({"key": key, "value": value})
return results
class MockLLMClient:
"""在真实 LLMClient 之上补一层测试控制能力。"""
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:
"""在真实 PlatformClient 之上补一层断言入口。"""
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):
"""本地 mock core,直接复用已有的内建 capability 实现。"""
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)
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": [],
}
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:
"""满足 `Context`/`CapabilityProxy` 需要的最小 peer。"""
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(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(id=request_id, phase="delta", data=chunk)
output = execution.finalize(chunks)
if include_completed:
yield EventMessage(id=request_id, phase="completed", output=output)
return iterator()
def _next_id(self) -> str:
self._counter += 1
return f"local_{self._counter:04d}"
def _plugin_metadata_from_spec(
plugin: PluginSpec,
*,
@@ -421,110 +87,6 @@ def _plugin_metadata_from_spec(
}
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)),
}
class MockContext(RuntimeContext):
"""直接用于 handler 单元测试的轻量运行时上下文。"""
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)
class MockMessageEvent(MessageEvent):
"""直接用于 handler 单元测试的轻量消息事件。"""
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)
@dataclass(slots=True)
class LocalRuntimeConfig:
"""本地 harness 的稳定配置对象。"""
@@ -575,7 +137,7 @@ class PluginHarness:
group_id: str | None = None,
event_type: str = "message",
platform_sink: StdoutPlatformSink | None = None,
) -> "PluginHarness":
) -> PluginHarness:
return cls(
LocalRuntimeConfig(
plugin_dir=Path(plugin_dir),
@@ -588,7 +150,7 @@ class PluginHarness:
platform_sink=platform_sink,
)
async def __aenter__(self) -> "PluginHarness":
async def __aenter__(self) -> PluginHarness:
await self.start()
return self
@@ -761,7 +323,11 @@ class PluginHarness:
"session_id": session_value,
"user_id": user_id or self.config.user_id,
"platform": platform or self.config.platform,
"platform_id": platform or self.config.platform,
"group_id": group_value,
"self_id": f"{platform or self.config.platform}-bot",
"sender_name": str(user_id or self.config.user_id or ""),
"is_admin": False,
"raw": {
"trace_id": request_id or self._next_request_id("trace"),
"event_type": event_type_value,
@@ -876,6 +442,11 @@ class PluginHarness:
return None
return {}
if isinstance(trigger, ScheduleTrigger):
if (
str(event_payload.get("event_type") or event_payload.get("type"))
== "schedule"
):
return {}
return None
return None
@@ -885,9 +456,7 @@ class PluginHarness:
trigger: CommandTrigger,
event_payload: dict[str, Any],
) -> dict[str, Any] | None:
if not self._passes_trigger_constraints(
trigger.platforms, trigger.message_types, event_payload
):
if not self._passes_filters(loaded, event_payload):
return None
text = str(event_payload.get("text", "")).strip()
for command_name in [trigger.command, *trigger.aliases]:
@@ -896,7 +465,7 @@ class PluginHarness:
match = self._match_command_name(text, command_name)
if match is None:
continue
return self._build_command_args(loaded.callable, match)
return self._build_command_args(loaded.descriptor.param_specs, match)
return None
def _match_message_trigger(
@@ -905,35 +474,64 @@ class PluginHarness:
trigger: MessageTrigger,
event_payload: dict[str, Any],
) -> dict[str, Any] | None:
if not self._passes_trigger_constraints(
trigger.platforms, trigger.message_types, event_payload
):
if not self._passes_filters(loaded, event_payload):
return None
text = str(event_payload.get("text", ""))
if trigger.regex:
match = re.search(trigger.regex, text)
if match is None:
return None
return self._build_regex_args(loaded.callable, match)
return self._build_regex_args(loaded.descriptor.param_specs, match)
if trigger.keywords and not any(
keyword in text for keyword in trigger.keywords
):
return None
return {}
def _passes_trigger_constraints(
def _passes_filters(
self,
platforms: list[str],
message_types: list[str],
loaded: LoadedHandler,
event_payload: dict[str, Any],
) -> bool:
platform = str(event_payload.get("platform", ""))
if platforms and platform not in platforms:
return False
if not message_types:
return True
current_message_type = self._message_type_name(event_payload)
return current_message_type in message_types
for filter_spec in loaded.descriptor.filters:
if isinstance(filter_spec, PlatformFilterSpec):
if str(event_payload.get("platform", "")) not in filter_spec.platforms:
return False
elif isinstance(filter_spec, MessageTypeFilterSpec):
if (
self._message_type_name(event_payload)
not in filter_spec.message_types
):
return False
elif isinstance(filter_spec, CompositeFilterSpec):
if not self._passes_composite_filter(filter_spec, event_payload):
return False
elif isinstance(filter_spec, LocalFilterRefSpec):
continue
return True
def _passes_composite_filter(
self,
filter_spec: CompositeFilterSpec,
event_payload: dict[str, Any],
) -> bool:
results: list[bool] = []
for child in filter_spec.children:
if isinstance(child, PlatformFilterSpec):
results.append(
str(event_payload.get("platform", "")) in child.platforms
)
elif isinstance(child, MessageTypeFilterSpec):
results.append(
self._message_type_name(event_payload) in child.message_types
)
elif isinstance(child, LocalFilterRefSpec):
results.append(True)
elif isinstance(child, CompositeFilterSpec):
results.append(self._passes_composite_filter(child, event_payload))
if filter_spec.kind == "and":
return all(results)
return any(results)
def _has_waiter_for_event(self, event_payload: dict[str, Any]) -> bool:
assert self.dispatcher is not None
@@ -973,34 +571,29 @@ class PluginHarness:
return text[len(command_name) :].strip()
return None
def _build_command_args(self, handler, remainder: str) -> dict[str, Any]:
names = self._legacy_arg_parameter_names(handler)
if not names or not remainder:
def _build_command_args(self, param_specs, remainder: str) -> dict[str, Any]:
if not param_specs or not remainder:
return {}
if len(names) == 1:
return {names[0]: remainder}
if len(param_specs) == 1:
return {param_specs[0].name: remainder}
tokens = self._split_command_remainder(remainder)
if not tokens:
return {}
values: dict[str, Any] = {}
for index, name in enumerate(names):
for index, spec in enumerate(param_specs):
if index >= len(tokens):
break
if index == len(names) - 1:
values[name] = " ".join(tokens[index:])
if spec.type == "greedy_str":
values[spec.name] = " ".join(tokens[index:])
break
values[name] = tokens[index]
values[spec.name] = tokens[index]
return values
def _build_regex_args(self, handler, match: re.Match[str]) -> dict[str, Any]:
def _build_regex_args(self, param_specs, match: re.Match[str]) -> dict[str, Any]:
named = {
key: value for key, value in match.groupdict().items() if value is not None
}
names = [
name
for name in self._legacy_arg_parameter_names(handler)
if name not in named
]
names = [spec.name for spec in param_specs if spec.name not in named]
positional = [value for value in match.groups() if value is not None]
for index, value in enumerate(positional):
if index >= len(names):
+22
View File
@@ -0,0 +1,22 @@
"""SDK parameter helper types.
本模块提供 SDK 参数类型助手,用于增强命令参数解析能力。
GreedyStr:
用于标记"贪婪字符串"参数,在命令解析时将剩余所有文本作为一个整体参数。
例如:/echo hello world this is a test
如果最后一个参数类型为 GreedyStr,将获取 "hello world this is a test" 而非仅 "hello"
使用方式:
在 handler 签名中将最后一个参数标注为 GreedyStr 类型,
_loader_support 会识别此类型并调整参数解析逻辑。
"""
from __future__ import annotations
class GreedyStr(str):
"""Consume the remaining command text as one argument."""
__all__ = ["GreedyStr"]