diff --git a/PROJECT_ARCHITECTURE.md b/PROJECT_ARCHITECTURE.md index 1c3413112..8a534e6d0 100644 --- a/PROJECT_ARCHITECTURE.md +++ b/PROJECT_ARCHITECTURE.md @@ -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 相关测试 - ---- ## 目录 diff --git a/src-new/astrbot_sdk/__init__.py b/src-new/astrbot_sdk/__init__.py index 462e607bb..2b89ccc47 100644 --- a/src-new/astrbot_sdk/__init__.py +++ b/src-new/astrbot_sdk/__init__.py @@ -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", ] diff --git a/src-new/astrbot_sdk/_testing_support.py b/src-new/astrbot_sdk/_testing_support.py new file mode 100644 index 000000000..e78087582 --- /dev/null +++ b/src-new/astrbot_sdk/_testing_support.py @@ -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", +] diff --git a/src-new/astrbot_sdk/cli.py b/src-new/astrbot_sdk/cli.py index c00be793f..7c80bbaac 100644 --- a/src-new/astrbot_sdk/cli.py +++ b/src-new/astrbot_sdk/cli.py @@ -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 ""}, + 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}, ) diff --git a/src-new/astrbot_sdk/clients/http.py b/src-new/astrbot_sdk/clients/http.py index ec798ef2a..efec135e8 100644 --- a/src-new/astrbot_sdk/clients/http.py +++ b/src-new/astrbot_sdk/clients/http.py @@ -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( diff --git a/src-new/astrbot_sdk/clients/llm.py b/src-new/astrbot_sdk/clients/llm.py index fe637eb99..14d7393fd 100644 --- a/src-new/astrbot_sdk/clients/llm.py +++ b/src-new/astrbot_sdk/clients/llm.py @@ -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, diff --git a/src-new/astrbot_sdk/clients/memory.py b/src-new/astrbot_sdk/clients/memory.py index 3edcaad7f..98cfcf8b9 100644 --- a/src-new/astrbot_sdk/clients/memory.py +++ b/src-new/astrbot_sdk/clients/memory.py @@ -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 diff --git a/src-new/astrbot_sdk/clients/metadata.py b/src-new/astrbot_sdk/clients/metadata.py index dbe80af7d..c2f9ab65f 100644 --- a/src-new/astrbot_sdk/clients/metadata.py +++ b/src-new/astrbot_sdk/clients/metadata.py @@ -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", ""), diff --git a/src-new/astrbot_sdk/clients/platform.py b/src-new/astrbot_sdk/clients/platform.py index 2329fc245..3c9b4a914 100644 --- a/src-new/astrbot_sdk/clients/platform.py +++ b/src-new/astrbot_sdk/clients/platform.py @@ -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]]: """获取群组成员列表。 获取指定群组的成员信息列表。注意仅对群组会话有效。 diff --git a/src-new/astrbot_sdk/commands.py b/src-new/astrbot_sdk/commands.py new file mode 100644 index 000000000..0e90ab830 --- /dev/null +++ b/src-new/astrbot_sdk/commands.py @@ -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"] diff --git a/src-new/astrbot_sdk/context.py b/src-new/astrbot_sdk/context.py index fede4c642..873f71b0a 100644 --- a/src-new/astrbot_sdk/context.py +++ b/src-new/astrbot_sdk/context.py @@ -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", "")) diff --git a/src-new/astrbot_sdk/decorators.py b/src-new/astrbot_sdk/decorators.py index c94020ce7..199153ec0 100644 --- a/src-new/astrbot_sdk/decorators.py +++ b/src-new/astrbot_sdk/decorators.py @@ -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]: """注册事件处理方法。 diff --git a/src-new/astrbot_sdk/docs/PROJECT_ARCHITECTURE.md b/src-new/astrbot_sdk/docs/PROJECT_ARCHITECTURE.md new file mode 100644 index 000000000..3ab7e90b4 --- /dev/null +++ b/src-new/astrbot_sdk/docs/PROJECT_ARCHITECTURE.md @@ -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 diff --git a/src-new/astrbot_sdk/errors.py b/src-new/astrbot_sdk/errors.py index 7b3e90dc5..cd615631f 100644 --- a/src-new/astrbot_sdk/errors.py +++ b/src-new/astrbot_sdk/errors.py @@ -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: diff --git a/src-new/astrbot_sdk/events.py b/src-new/astrbot_sdk/events.py index 297444ce9..0188e7c80 100644 --- a/src-new/astrbot_sdk/events.py +++ b/src-new/astrbot_sdk/events.py @@ -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" + ) diff --git a/src-new/astrbot_sdk/filters.py b/src-new/astrbot_sdk/filters.py new file mode 100644 index 000000000..e0635adf3 --- /dev/null +++ b/src-new/astrbot_sdk/filters.py @@ -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", +] diff --git a/src-new/astrbot_sdk/message_components.py b/src-new/astrbot_sdk/message_components.py new file mode 100644 index 000000000..dfb17c7f0 --- /dev/null +++ b/src-new/astrbot_sdk/message_components.py @@ -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", +] diff --git a/src-new/astrbot_sdk/message_result.py b/src-new/astrbot_sdk/message_result.py new file mode 100644 index 000000000..3c593e537 --- /dev/null +++ b/src-new/astrbot_sdk/message_result.py @@ -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", +] diff --git a/src-new/astrbot_sdk/message_session.py b/src-new/astrbot_sdk/message_session.py new file mode 100644 index 000000000..a011f8dcc --- /dev/null +++ b/src-new/astrbot_sdk/message_session.py @@ -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, + ) diff --git a/src-new/astrbot_sdk/protocol/__init__.py b/src-new/astrbot_sdk/protocol/__init__.py index 4a98b52c6..9fd2fd0d4 100644 --- a/src-new/astrbot_sdk/protocol/__init__.py +++ b/src-new/astrbot_sdk/protocol/__init__.py @@ -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", diff --git a/src-new/astrbot_sdk/protocol/_builtin_schemas.py b/src-new/astrbot_sdk/protocol/_builtin_schemas.py new file mode 100644 index 000000000..302425f9d --- /dev/null +++ b/src-new/astrbot_sdk/protocol/_builtin_schemas.py @@ -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", +] diff --git a/src-new/astrbot_sdk/protocol/descriptors.py b/src-new/astrbot_sdk/protocol/descriptors.py index 697e088bb..eca4d4267 100644 --- a/src-new/astrbot_sdk/protocol/descriptors.py +++ b/src-new/astrbot_sdk/protocol/descriptors.py @@ -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", ] diff --git a/src-new/astrbot_sdk/protocol/messages.py b/src-new/astrbot_sdk/protocol/messages.py index 70e42a862..399051c96 100644 --- a/src-new/astrbot_sdk/protocol/messages.py +++ b/src-new/astrbot_sdk/protocol/messages.py @@ -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: 所有字段必须为空 diff --git a/src-new/astrbot_sdk/protocol/wire_codecs.py b/src-new/astrbot_sdk/protocol/wire_codecs.py deleted file mode 100644 index 494df4435..000000000 --- a/src-new/astrbot_sdk/protocol/wire_codecs.py +++ /dev/null @@ -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", -] diff --git a/src-new/astrbot_sdk/runtime/__init__.py b/src-new/astrbot_sdk/runtime/__init__.py index ef17b3fce..7601f745c 100644 --- a/src-new/astrbot_sdk/runtime/__init__.py +++ b/src-new/astrbot_sdk/runtime/__init__.py @@ -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) diff --git a/src-new/astrbot_sdk/runtime/_capability_router_builtins.py b/src-new/astrbot_sdk/runtime/_capability_router_builtins.py new file mode 100644 index 000000000..011286e4c --- /dev/null +++ b/src-new/astrbot_sdk/runtime/_capability_router_builtins.py @@ -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"{text}"} + + 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"] diff --git a/src-new/astrbot_sdk/runtime/_loader_support.py b/src-new/astrbot_sdk/runtime/_loader_support.py new file mode 100644 index 000000000..8575ff115 --- /dev/null +++ b/src-new/astrbot_sdk/runtime/_loader_support.py @@ -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", +] diff --git a/src-new/astrbot_sdk/runtime/_streaming.py b/src-new/astrbot_sdk/runtime/_streaming.py new file mode 100644 index 000000000..29d2671ca --- /dev/null +++ b/src-new/astrbot_sdk/runtime/_streaming.py @@ -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"] diff --git a/src-new/astrbot_sdk/runtime/bootstrap.py b/src-new/astrbot_sdk/runtime/bootstrap.py index 7d025b27f..7a8706965 100644 --- a/src-new/astrbot_sdk/runtime/bootstrap.py +++ b/src-new/astrbot_sdk/runtime/bootstrap.py @@ -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() diff --git a/src-new/astrbot_sdk/runtime/capability_dispatcher.py b/src-new/astrbot_sdk/runtime/capability_dispatcher.py new file mode 100644 index 000000000..652fceed5 --- /dev/null +++ b/src-new/astrbot_sdk/runtime/capability_dispatcher.py @@ -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__", "") + payload_keys = sorted(str(key) for key in payload.keys()) + payload_keys_text = ", ".join(payload_keys) if payload_keys else "" + return ( + f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:" + f"必填参数 '{parameter_name}' 无法注入。" + f"签名: {getattr(handler, '__name__', '')}" + 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"] diff --git a/src-new/astrbot_sdk/runtime/capability_router.py b/src-new/astrbot_sdk/runtime/capability_router.py index 7c0a974c0..3e689f4a3 100644 --- a/src-new/astrbot_sdk/runtime/capability_router.py +++ b/src-new/astrbot_sdk/runtime/capability_router.py @@ -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 # ------------------------------------------------------------------ diff --git a/src-new/astrbot_sdk/runtime/environment_groups.py b/src-new/astrbot_sdk/runtime/environment_groups.py index fe4c76af1..b742d66ec 100644 --- a/src-new/astrbot_sdk/runtime/environment_groups.py +++ b/src-new/astrbot_sdk/runtime/environment_groups.py @@ -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]: diff --git a/src-new/astrbot_sdk/runtime/handler_dispatcher.py b/src-new/astrbot_sdk/runtime/handler_dispatcher.py index fa5a8029a..56ed0bf5f 100644 --- a/src-new/astrbot_sdk/runtime/handler_dispatcher.py +++ b/src-new/astrbot_sdk/runtime/handler_dispatcher.py @@ -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__", "") - payload_keys = sorted(str(key) for key in payload.keys()) - payload_keys_text = ", ".join(payload_keys) if payload_keys else "" - return ( - f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:" - f"必填参数 '{parameter_name}' 无法注入。" - f"签名: {getattr(handler, '__name__', '')}" - 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"] diff --git a/src-new/astrbot_sdk/runtime/loader.py b/src-new/astrbot_sdk/runtime/loader.py index 1a4f89469..ad9aae4c0 100644 --- a/src-new/astrbot_sdk/runtime/loader.py +++ b/src-new/astrbot_sdk/runtime/loader.py @@ -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), ) ) diff --git a/src-new/astrbot_sdk/runtime/peer.py b/src-new/astrbot_sdk/runtime/peer.py index 26dc0e198..a84093e65 100644 --- a/src-new/astrbot_sdk/runtime/peer.py +++ b/src-new/astrbot_sdk/runtime/peer.py @@ -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)) diff --git a/src-new/astrbot_sdk/runtime/supervisor.py b/src-new/astrbot_sdk/runtime/supervisor.py index 5877c21e3..014e08441 100644 --- a/src-new/astrbot_sdk/runtime/supervisor.py +++ b/src-new/astrbot_sdk/runtime/supervisor.py @@ -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() diff --git a/src-new/astrbot_sdk/runtime/transport.py b/src-new/astrbot_sdk/runtime/transport.py index a5549f67d..22724c5cf 100644 --- a/src-new/astrbot_sdk/runtime/transport.py +++ b/src-new/astrbot_sdk/runtime/transport.py @@ -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 diff --git a/src-new/astrbot_sdk/runtime/worker.py b/src-new/astrbot_sdk/runtime/worker.py index 7b1b6cc55..2d3b4626f 100644 --- a/src-new/astrbot_sdk/runtime/worker.py +++ b/src-new/astrbot_sdk/runtime/worker.py @@ -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, diff --git a/src-new/astrbot_sdk/schedule.py b/src-new/astrbot_sdk/schedule.py new file mode 100644 index 000000000..e0aa20c7a --- /dev/null +++ b/src-new/astrbot_sdk/schedule.py @@ -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"] diff --git a/src-new/astrbot_sdk/session_waiter.py b/src-new/astrbot_sdk/session_waiter.py new file mode 100644 index 000000000..4813c7d9e --- /dev/null +++ b/src-new/astrbot_sdk/session_waiter.py @@ -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 diff --git a/src-new/astrbot_sdk/testing.py b/src-new/astrbot_sdk/testing.py index adbbb8dc4..cdec45ada 100644 --- a/src-new/astrbot_sdk/testing.py +++ b/src-new/astrbot_sdk/testing.py @@ -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): diff --git a/src-new/astrbot_sdk/types.py b/src-new/astrbot_sdk/types.py new file mode 100644 index 000000000..c2bc911ec --- /dev/null +++ b/src-new/astrbot_sdk/types.py @@ -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"]