mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
Refactor worker initialization and remove unused codec parameters; add schedule and session waiter modules
- Simplified `GroupWorkerRuntime` and `PluginWorkerRuntime` constructors by removing the codec parameter and related logic. - Introduced `schedule.py` to define `ScheduleContext` for managing scheduled tasks with a clear structure and payload handling. - Added `session_waiter.py` for session-based conversational flow management, including `SessionController` and `SessionWaiterManager` for handling multi-turn dialogues. - Enhanced testing utilities in `testing.py` by removing unused classes and streamlining the structure. - Created `types.py` to introduce `GreedyStr` for improved command parameter parsing.
This commit is contained in:
@@ -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 相关测试
|
||||
|
||||
---
|
||||
|
||||
## 目录
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,478 @@
|
||||
"""Shared support primitives for local SDK testing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import typing
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, TextIO
|
||||
|
||||
from .context import CancelToken
|
||||
from .context import Context as RuntimeContext
|
||||
from .events import MessageEvent
|
||||
from .protocol.messages import EventMessage, PeerInfo
|
||||
from .runtime._streaming import StreamExecution
|
||||
from .runtime.capability_router import CapabilityRouter
|
||||
|
||||
|
||||
def _clone_payload_mapping(value: Any) -> dict[str, Any] | None:
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
return {str(key): item for key, item in value.items()}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RecordedSend:
|
||||
kind: str
|
||||
message_id: str
|
||||
session_id: str
|
||||
text: str | None = None
|
||||
image_url: str | None = None
|
||||
chain: list[dict[str, Any]] | None = None
|
||||
target: dict[str, Any] | None = None
|
||||
raw: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def session(self) -> str:
|
||||
return self.session_id
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: dict[str, Any]) -> RecordedSend:
|
||||
if "text" in payload:
|
||||
kind = "text"
|
||||
elif "image_url" in payload:
|
||||
kind = "image"
|
||||
elif "chain" in payload:
|
||||
kind = "chain"
|
||||
else:
|
||||
kind = "unknown"
|
||||
return cls(
|
||||
kind=kind,
|
||||
message_id=str(payload.get("message_id", "")),
|
||||
session_id=str(payload.get("session", "")),
|
||||
text=payload.get("text") if isinstance(payload.get("text"), str) else None,
|
||||
image_url=(
|
||||
payload.get("image_url")
|
||||
if isinstance(payload.get("image_url"), str)
|
||||
else None
|
||||
),
|
||||
chain=(
|
||||
[dict(item) for item in payload.get("chain", [])]
|
||||
if isinstance(payload.get("chain"), list)
|
||||
else None
|
||||
),
|
||||
target=_clone_payload_mapping(payload.get("target")),
|
||||
raw=dict(payload),
|
||||
)
|
||||
|
||||
|
||||
class StdoutPlatformSink:
|
||||
def __init__(self, stream: TextIO | None = None) -> None:
|
||||
self._stream = stream
|
||||
self.records: list[RecordedSend] = []
|
||||
|
||||
def record(self, item: RecordedSend) -> None:
|
||||
self.records.append(item)
|
||||
if self._stream is None:
|
||||
return
|
||||
self._stream.write(self._format(item) + "\n")
|
||||
self._stream.flush()
|
||||
|
||||
def clear(self) -> None:
|
||||
self.records.clear()
|
||||
|
||||
def _format(self, item: RecordedSend) -> str:
|
||||
if item.kind == "text":
|
||||
return f"[text][{item.session_id}] {item.text or ''}"
|
||||
if item.kind == "image":
|
||||
return f"[image][{item.session_id}] {item.image_url or ''}"
|
||||
if item.kind == "chain":
|
||||
count = len(item.chain or [])
|
||||
return f"[chain][{item.session_id}] {count} components"
|
||||
return f"[send][{item.session_id}] {item.raw}"
|
||||
|
||||
|
||||
class InMemoryDB:
|
||||
def __init__(self, store: dict[str, Any]) -> None:
|
||||
self._store = store
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
return self._store.get(key, default)
|
||||
|
||||
def set(self, key: str, value: Any) -> None:
|
||||
self._store[key] = value
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
self._store.pop(key, None)
|
||||
|
||||
def list(self, prefix: str | None = None) -> list[str]:
|
||||
keys = sorted(self._store.keys())
|
||||
if prefix is None:
|
||||
return keys
|
||||
return [key for key in keys if key.startswith(prefix)]
|
||||
|
||||
def get_many(self, keys: list[str]) -> list[dict[str, Any]]:
|
||||
return [{"key": key, "value": self._store.get(key)} for key in keys]
|
||||
|
||||
def set_many(self, items: list[dict[str, Any]]) -> None:
|
||||
for item in items:
|
||||
self.set(str(item.get("key", "")), item.get("value"))
|
||||
|
||||
|
||||
class InMemoryMemory:
|
||||
def __init__(self, store: dict[str, dict[str, Any]]) -> None:
|
||||
self._store = store
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
return self._store.get(key, default)
|
||||
|
||||
def save(self, key: str, value: dict[str, Any]) -> None:
|
||||
self._store[key] = dict(value)
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
self._store.pop(key, None)
|
||||
|
||||
def search(self, query: str) -> list[dict[str, Any]]:
|
||||
results: list[dict[str, Any]] = []
|
||||
for key, value in self._store.items():
|
||||
if query in key or query in str(value):
|
||||
results.append({"key": key, "value": value})
|
||||
return results
|
||||
|
||||
|
||||
class MockLLMClient:
|
||||
def __init__(self, client: Any, router: MockCapabilityRouter) -> None:
|
||||
self._client = client
|
||||
self._router = router
|
||||
|
||||
def mock_response(self, text: str) -> None:
|
||||
self._router.enqueue_llm_response(text)
|
||||
|
||||
def mock_stream_response(self, text: str) -> None:
|
||||
self._router.enqueue_llm_stream_response(text)
|
||||
|
||||
def clear_mock_responses(self) -> None:
|
||||
self._router.clear_llm_responses()
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._client, name)
|
||||
|
||||
|
||||
class MockPlatformClient:
|
||||
def __init__(self, client: Any, sink: StdoutPlatformSink) -> None:
|
||||
self._client = client
|
||||
self._sink = sink
|
||||
|
||||
@property
|
||||
def records(self) -> list[RecordedSend]:
|
||||
return list(self._sink.records)
|
||||
|
||||
def assert_sent(
|
||||
self,
|
||||
expected_text: str | None = None,
|
||||
*,
|
||||
kind: str = "text",
|
||||
count: int | None = None,
|
||||
) -> None:
|
||||
matched = [item for item in self._sink.records if item.kind == kind]
|
||||
if expected_text is not None:
|
||||
matched = [item for item in matched if item.text == expected_text]
|
||||
if count is not None:
|
||||
if len(matched) != count:
|
||||
raise AssertionError(
|
||||
f"expected {count} sent records, got {len(matched)}: {matched}"
|
||||
)
|
||||
return
|
||||
if not matched:
|
||||
raise AssertionError(
|
||||
f"expected sent record kind={kind!r} text={expected_text!r}, got {self._sink.records}"
|
||||
)
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._client, name)
|
||||
|
||||
|
||||
class MockCapabilityRouter(CapabilityRouter):
|
||||
def __init__(self, *, platform_sink: StdoutPlatformSink | None = None) -> None:
|
||||
self.platform_sink = platform_sink or StdoutPlatformSink()
|
||||
self._llm_responses: list[str] = []
|
||||
self._llm_stream_responses: list[str] = []
|
||||
super().__init__()
|
||||
self.db = InMemoryDB(self.db_store)
|
||||
self.memory = InMemoryMemory(self.memory_store)
|
||||
|
||||
def enqueue_llm_response(self, text: str) -> None:
|
||||
self._llm_responses.append(text)
|
||||
|
||||
def enqueue_llm_stream_response(self, text: str) -> None:
|
||||
self._llm_stream_responses.append(text)
|
||||
|
||||
def clear_llm_responses(self) -> None:
|
||||
self._llm_responses.clear()
|
||||
self._llm_stream_responses.clear()
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
capability: str,
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
stream: bool,
|
||||
cancel_token,
|
||||
request_id: str,
|
||||
) -> dict[str, Any] | StreamExecution:
|
||||
if capability == "llm.chat":
|
||||
return {"text": self._take_llm_response(str(payload.get("prompt", "")))}
|
||||
if capability == "llm.chat_raw":
|
||||
text = self._take_llm_response(str(payload.get("prompt", "")))
|
||||
return {
|
||||
"text": text,
|
||||
"usage": {
|
||||
"input_tokens": len(str(payload.get("prompt", ""))),
|
||||
"output_tokens": len(text),
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
"tool_calls": [],
|
||||
"role": "assistant",
|
||||
"reasoning_content": None,
|
||||
"reasoning_signature": None,
|
||||
}
|
||||
if capability == "llm.stream_chat":
|
||||
text = self._take_llm_stream_response(str(payload.get("prompt", "")))
|
||||
|
||||
async def iterator() -> typing.AsyncIterator[dict[str, Any]]:
|
||||
for char in text:
|
||||
cancel_token.raise_if_cancelled()
|
||||
await asyncio.sleep(0)
|
||||
yield {"text": char}
|
||||
|
||||
return StreamExecution(
|
||||
iterator=iterator(),
|
||||
finalize=lambda chunks: {
|
||||
"text": "".join(item.get("text", "") for item in chunks)
|
||||
},
|
||||
)
|
||||
before = len(self.sent_messages)
|
||||
result = await super().execute(
|
||||
capability,
|
||||
payload,
|
||||
stream=stream,
|
||||
cancel_token=cancel_token,
|
||||
request_id=request_id,
|
||||
)
|
||||
self._flush_platform_records(before)
|
||||
return result
|
||||
|
||||
def _flush_platform_records(self, start_index: int) -> None:
|
||||
for payload in self.sent_messages[start_index:]:
|
||||
self.platform_sink.record(RecordedSend.from_payload(payload))
|
||||
|
||||
def _take_llm_response(self, prompt: str) -> str:
|
||||
if self._llm_responses:
|
||||
return self._llm_responses.pop(0)
|
||||
return f"Echo: {prompt}"
|
||||
|
||||
def _take_llm_stream_response(self, prompt: str) -> str:
|
||||
if self._llm_stream_responses:
|
||||
return self._llm_stream_responses.pop(0)
|
||||
if self._llm_responses:
|
||||
return self._llm_responses.pop(0)
|
||||
return f"Echo: {prompt}"
|
||||
|
||||
|
||||
class MockPeer:
|
||||
def __init__(self, router: MockCapabilityRouter) -> None:
|
||||
self._router = router
|
||||
self._counter = 0
|
||||
self.remote_peer = PeerInfo(
|
||||
name="astrbot-local-core",
|
||||
role="core",
|
||||
version="local",
|
||||
)
|
||||
self.remote_capabilities = list(router.descriptors())
|
||||
self.remote_capability_map = {
|
||||
item.name: item for item in self.remote_capabilities
|
||||
}
|
||||
self.remote_handlers: list[Any] = []
|
||||
self.remote_provided_capabilities: list[Any] = []
|
||||
self.remote_metadata = {"mode": "local"}
|
||||
|
||||
async def invoke(
|
||||
self,
|
||||
capability: str,
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
stream: bool = False,
|
||||
request_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if stream:
|
||||
raise ValueError("stream=True 请使用 invoke_stream()")
|
||||
return typing.cast(
|
||||
dict[str, Any],
|
||||
await self._router.execute(
|
||||
capability,
|
||||
payload,
|
||||
stream=False,
|
||||
cancel_token=CancelToken(),
|
||||
request_id=request_id or self._next_id(),
|
||||
),
|
||||
)
|
||||
|
||||
async def invoke_stream(
|
||||
self,
|
||||
capability: str,
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
request_id: str | None = None,
|
||||
include_completed: bool = False,
|
||||
):
|
||||
request_id = request_id or self._next_id()
|
||||
execution = typing.cast(
|
||||
StreamExecution,
|
||||
await self._router.execute(
|
||||
capability,
|
||||
payload,
|
||||
stream=True,
|
||||
cancel_token=CancelToken(),
|
||||
request_id=request_id,
|
||||
),
|
||||
)
|
||||
|
||||
async def iterator():
|
||||
yield EventMessage.model_validate({"id": request_id, "phase": "started"})
|
||||
chunks: list[dict[str, Any]] = []
|
||||
async for chunk in execution.iterator:
|
||||
if execution.collect_chunks:
|
||||
chunks.append(chunk)
|
||||
yield EventMessage.model_validate(
|
||||
{"id": request_id, "phase": "delta", "data": chunk}
|
||||
)
|
||||
output = execution.finalize(chunks)
|
||||
if include_completed:
|
||||
yield EventMessage.model_validate(
|
||||
{"id": request_id, "phase": "completed", "output": output}
|
||||
)
|
||||
|
||||
return iterator()
|
||||
|
||||
def _next_id(self) -> str:
|
||||
self._counter += 1
|
||||
return f"local_{self._counter:04d}"
|
||||
|
||||
|
||||
def _normalize_plugin_metadata(
|
||||
plugin_id: str,
|
||||
plugin_metadata: Mapping[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
if plugin_metadata is None:
|
||||
plugin_metadata = {}
|
||||
declared_name = plugin_metadata.get("name")
|
||||
if declared_name is not None and str(declared_name) != plugin_id:
|
||||
raise ValueError(
|
||||
"MockContext.plugin_metadata['name'] 必须与 plugin_id 一致,"
|
||||
f"当前收到 {declared_name!r} != {plugin_id!r}"
|
||||
)
|
||||
description = plugin_metadata.get("description")
|
||||
if description is None:
|
||||
description = plugin_metadata.get("desc", "")
|
||||
return {
|
||||
"name": plugin_id,
|
||||
"display_name": str(plugin_metadata.get("display_name") or plugin_id),
|
||||
"description": str(description or ""),
|
||||
"author": str(plugin_metadata.get("author") or ""),
|
||||
"version": str(plugin_metadata.get("version") or "0.0.0"),
|
||||
"enabled": bool(plugin_metadata.get("enabled", True)),
|
||||
}
|
||||
|
||||
|
||||
class MockContext(RuntimeContext):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plugin_id: str = "test-plugin",
|
||||
logger: Any | None = None,
|
||||
cancel_token: CancelToken | None = None,
|
||||
platform_sink: StdoutPlatformSink | None = None,
|
||||
plugin_metadata: Mapping[str, Any] | None = None,
|
||||
) -> None:
|
||||
self.platform_sink = platform_sink or StdoutPlatformSink()
|
||||
self.router = MockCapabilityRouter(platform_sink=self.platform_sink)
|
||||
self.mock_peer = MockPeer(self.router)
|
||||
super().__init__(
|
||||
peer=self.mock_peer,
|
||||
plugin_id=plugin_id,
|
||||
cancel_token=cancel_token,
|
||||
logger=logger,
|
||||
)
|
||||
self.router.upsert_plugin(
|
||||
metadata=_normalize_plugin_metadata(plugin_id, plugin_metadata),
|
||||
config={},
|
||||
)
|
||||
self.llm = MockLLMClient(self.llm, self.router)
|
||||
self.platform = MockPlatformClient(self.platform, self.platform_sink)
|
||||
|
||||
@property
|
||||
def sent_messages(self) -> list[RecordedSend]:
|
||||
return list(self.platform_sink.records)
|
||||
|
||||
@property
|
||||
def event_actions(self) -> list[dict[str, Any]]:
|
||||
return list(self.router.event_actions)
|
||||
|
||||
|
||||
class MockMessageEvent(MessageEvent):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
text: str = "",
|
||||
user_id: str | None = "test-user",
|
||||
group_id: str | None = None,
|
||||
platform: str | None = "test",
|
||||
session_id: str | None = "test-session",
|
||||
raw: dict[str, Any] | None = None,
|
||||
context: MockContext | None = None,
|
||||
) -> None:
|
||||
self.replies: list[str] = []
|
||||
super().__init__(
|
||||
text=text,
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
platform=platform,
|
||||
session_id=session_id,
|
||||
raw=raw,
|
||||
context=context,
|
||||
)
|
||||
if context is not None:
|
||||
self.bind_runtime_reply(context)
|
||||
elif self._reply_handler is None:
|
||||
self.bind_reply_handler(self._capture_reply)
|
||||
|
||||
@property
|
||||
def is_private(self) -> bool:
|
||||
return self.group_id is None
|
||||
|
||||
def bind_runtime_reply(self, context: MockContext) -> None:
|
||||
self._context = context
|
||||
|
||||
async def reply(text: str) -> None:
|
||||
self.replies.append(text)
|
||||
await context.platform.send(self.session_ref or self.session_id, text)
|
||||
|
||||
self.bind_reply_handler(reply)
|
||||
|
||||
async def _capture_reply(self, text: str) -> None:
|
||||
self.replies.append(text)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InMemoryDB",
|
||||
"InMemoryMemory",
|
||||
"MockCapabilityRouter",
|
||||
"MockContext",
|
||||
"MockLLMClient",
|
||||
"MockMessageEvent",
|
||||
"MockPeer",
|
||||
"MockPlatformClient",
|
||||
"RecordedSend",
|
||||
"StdoutPlatformSink",
|
||||
]
|
||||
+48
-130
@@ -1,9 +1,21 @@
|
||||
"""AstrBot SDK 的命令行入口。"""
|
||||
"""AstrBot SDK 的命令行入口。
|
||||
|
||||
本模块提供 astrbot-sdk 命令行工具的所有子命令,包括:
|
||||
- init: 创建新插件骨架,生成 plugin.yaml、main.py、README.md 等模板文件
|
||||
- validate: 校验插件清单、导入路径和 handler 发现是否正常
|
||||
- build: 将插件打包为 .zip 发布包
|
||||
- dev: 本地开发模式,支持 --local/--watch/--interactive 等调试选项
|
||||
- run: 启动插件主管进程(supervisor),通过 stdio 与 AstrBot 核心通信
|
||||
- worker: 内部命令,由 supervisor 调用以启动单个插件工作进程
|
||||
|
||||
错误处理:
|
||||
所有 CLI 异常都会被分类并返回标准化的退出码和错误提示,
|
||||
便于 CI/CD 集成和用户快速定位问题。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import sys
|
||||
import typing
|
||||
@@ -41,9 +53,6 @@ BUILD_EXCLUDED_FILES = {
|
||||
".astrbot-worker-state.json",
|
||||
}
|
||||
WATCH_POLL_INTERVAL_SECONDS = 0.5
|
||||
INIT_DEFAULT_AUTHOR = ""
|
||||
INIT_DEFAULT_PYTHON_VERSION = "3.12"
|
||||
INIT_DEFAULT_VERSION = "1.0.0"
|
||||
|
||||
|
||||
class _CliPluginValidationError(RuntimeError):
|
||||
@@ -545,6 +554,11 @@ def _handle_dev_meta_command(command: str, state: dict[str, Any]) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _slugify_plugin_name(value: str) -> str:
|
||||
slug = re.sub(r"[^a-zA-Z0-9]+", "_", value).strip("_").lower()
|
||||
return slug or "my_plugin"
|
||||
|
||||
|
||||
def _class_name_for_plugin(value: str) -> str:
|
||||
parts = [part for part in re.split(r"[^a-zA-Z0-9]+", value) if part]
|
||||
if not parts:
|
||||
@@ -557,62 +571,18 @@ def _sanitize_build_part(value: str) -> str:
|
||||
return sanitized or "artifact"
|
||||
|
||||
|
||||
def _yaml_string(value: str) -> str:
|
||||
return json.dumps(value, ensure_ascii=False)
|
||||
|
||||
|
||||
def _normalize_init_plugin_name(value: str) -> str:
|
||||
normalized = re.sub(r"[\s-]+", "_", value.strip())
|
||||
normalized = re.sub(r"[^a-zA-Z0-9_]+", "_", normalized)
|
||||
normalized = re.sub(r"_+", "_", normalized).strip("_").lower()
|
||||
if not normalized:
|
||||
normalized = "my_plugin"
|
||||
|
||||
prefix = "astrbot_plugin_"
|
||||
if normalized == "astrbot_plugin":
|
||||
return f"{prefix}my_plugin"
|
||||
if normalized.startswith(prefix):
|
||||
suffix = normalized.removeprefix(prefix).strip("_") or "my_plugin"
|
||||
return f"{prefix}{suffix}"
|
||||
return f"{prefix}{normalized}"
|
||||
|
||||
|
||||
def _prompt_required_init_name() -> str:
|
||||
while True:
|
||||
value = click.prompt("插件名字", default="", show_default=False).strip()
|
||||
if value:
|
||||
return value
|
||||
click.echo("插件名字不能为空")
|
||||
|
||||
|
||||
def _collect_init_inputs(name: str | None) -> tuple[str, str, str]:
|
||||
if name is not None:
|
||||
return name, INIT_DEFAULT_AUTHOR, INIT_DEFAULT_VERSION
|
||||
|
||||
plugin_name = _prompt_required_init_name()
|
||||
author = click.prompt("作者名字", default="", show_default=False).strip()
|
||||
version = click.prompt("版本", default=INIT_DEFAULT_VERSION).strip()
|
||||
return plugin_name, author, version or INIT_DEFAULT_VERSION
|
||||
|
||||
|
||||
def _render_init_plugin_yaml(
|
||||
*,
|
||||
plugin_name: str,
|
||||
display_name: str,
|
||||
author: str,
|
||||
version: str,
|
||||
python_version: str,
|
||||
) -> str:
|
||||
def _render_init_plugin_yaml(*, plugin_name: str, display_name: str) -> str:
|
||||
python_version = f"{sys.version_info.major}.{sys.version_info.minor}"
|
||||
class_name = _class_name_for_plugin(plugin_name)
|
||||
return dedent(
|
||||
f"""\
|
||||
name: {plugin_name}
|
||||
display_name: {_yaml_string(display_name)}
|
||||
desc: {_yaml_string("使用 AstrBot SDK 创建的插件")}
|
||||
author: {_yaml_string(author)}
|
||||
version: {_yaml_string(version)}
|
||||
display_name: {display_name}
|
||||
desc: 使用 AstrBot SDK 创建的插件
|
||||
author: your-name
|
||||
version: 0.1.0
|
||||
runtime:
|
||||
python: {_yaml_string(python_version)}
|
||||
python: "{python_version}"
|
||||
components:
|
||||
- class: main:{class_name}
|
||||
"""
|
||||
@@ -672,7 +642,7 @@ def _render_init_readme(*, plugin_name: str) -> str:
|
||||
def _render_init_test_py(*, plugin_name: str) -> str:
|
||||
class_name = _class_name_for_plugin(plugin_name)
|
||||
return dedent(
|
||||
f'''\
|
||||
f"""\
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -704,7 +674,7 @@ def _render_init_test_py(*, plugin_name: str) -> str:
|
||||
records = await harness.dispatch_text("hello")
|
||||
|
||||
assert any(record.text == "Hello, World!" for record in records)
|
||||
'''
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
@@ -777,24 +747,19 @@ def _iter_build_files(plugin_dir: Path, output_dir: Path) -> list[Path]:
|
||||
return files
|
||||
|
||||
|
||||
def _init_plugin(name: str | None) -> None:
|
||||
raw_name, author, version = _collect_init_inputs(name)
|
||||
normalized_name = _normalize_init_plugin_name(raw_name)
|
||||
target_dir = Path(normalized_name)
|
||||
def _init_plugin(name: str) -> None:
|
||||
target_dir = Path(name)
|
||||
if target_dir.exists():
|
||||
raise _CliPluginValidationError(f"目标目录已存在:{target_dir}")
|
||||
|
||||
plugin_name = normalized_name
|
||||
display_name = raw_name
|
||||
plugin_name = _slugify_plugin_name(target_dir.name)
|
||||
display_name = target_dir.name
|
||||
target_dir.mkdir(parents=True, exist_ok=False)
|
||||
(target_dir / "tests").mkdir()
|
||||
(target_dir / "plugin.yaml").write_text(
|
||||
_render_init_plugin_yaml(
|
||||
plugin_name=plugin_name,
|
||||
display_name=display_name,
|
||||
author=author,
|
||||
version=version,
|
||||
python_version=INIT_DEFAULT_PYTHON_VERSION,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
@@ -811,7 +776,7 @@ def _init_plugin(name: str | None) -> None:
|
||||
_render_init_test_py(plugin_name=plugin_name),
|
||||
encoding="utf-8",
|
||||
)
|
||||
click.echo(f"已创建插件骨架:{target_dir.resolve()}")
|
||||
click.echo(f"已创建插件骨架:{target_dir}")
|
||||
click.echo("后续命令:")
|
||||
click.echo(f" astrbot-sdk validate --plugin-dir {target_dir}")
|
||||
click.echo(
|
||||
@@ -867,43 +832,23 @@ def cli(ctx, verbose: bool) -> None:
|
||||
type=click.Path(file_okay=False, dir_okay=True, path_type=Path),
|
||||
help="Directory containing plugin folders",
|
||||
)
|
||||
@click.option(
|
||||
"--worker-wire-codec",
|
||||
default="json",
|
||||
show_default=True,
|
||||
type=click.Choice(["json", "msgpack"]),
|
||||
help="Wire codec for supervisor-to-worker transport",
|
||||
)
|
||||
def run(plugins_dir: Path, worker_wire_codec: str) -> None:
|
||||
def run(plugins_dir: Path) -> None:
|
||||
"""Start the plugin supervisor over stdio."""
|
||||
entrypoint = (
|
||||
run_supervisor(plugins_dir=plugins_dir)
|
||||
if worker_wire_codec == "json"
|
||||
else run_supervisor(
|
||||
plugins_dir=plugins_dir,
|
||||
worker_wire_codec=worker_wire_codec,
|
||||
)
|
||||
)
|
||||
_run_async_entrypoint(
|
||||
entrypoint,
|
||||
run_supervisor(plugins_dir=plugins_dir),
|
||||
log_message=f"启动插件主管进程,插件目录:{plugins_dir}",
|
||||
context={
|
||||
"plugins_dir": plugins_dir,
|
||||
"worker_wire_codec": worker_wire_codec,
|
||||
},
|
||||
context={"plugins_dir": plugins_dir},
|
||||
)
|
||||
|
||||
|
||||
@cli.command()
|
||||
@click.argument("name", required=False, type=str)
|
||||
def init(name: str | None) -> None:
|
||||
"""Create a new plugin skeleton; omit name to enter interactive mode."""
|
||||
@click.argument("name", type=str)
|
||||
def init(name: str) -> None:
|
||||
"""Create a new plugin skeleton in the target directory."""
|
||||
_run_sync_entrypoint(
|
||||
lambda: _init_plugin(name),
|
||||
log_message=(
|
||||
f"创建插件骨架:{name}" if name is not None else "创建插件骨架:交互模式"
|
||||
),
|
||||
context={"target": name or "<interactive>"},
|
||||
log_message=f"创建插件骨架:{name}",
|
||||
context={"target": Path(name)},
|
||||
)
|
||||
|
||||
|
||||
@@ -1028,15 +973,7 @@ def dev(
|
||||
required=False,
|
||||
type=click.Path(file_okay=True, dir_okay=False, path_type=Path),
|
||||
)
|
||||
@click.option(
|
||||
"--wire-codec",
|
||||
default="json",
|
||||
show_default=True,
|
||||
type=click.Choice(["json", "msgpack"]),
|
||||
)
|
||||
def worker(
|
||||
plugin_dir: Path | None, group_metadata: Path | None, wire_codec: str
|
||||
) -> None:
|
||||
def worker(plugin_dir: Path | None, group_metadata: Path | None) -> None:
|
||||
"""Internal command used by the supervisor to start a worker."""
|
||||
if plugin_dir is None and group_metadata is None:
|
||||
raise click.UsageError("Either --plugin-dir or --group-metadata is required")
|
||||
@@ -1047,42 +984,23 @@ def worker(
|
||||
|
||||
target = str(group_metadata or plugin_dir)
|
||||
if group_metadata is not None:
|
||||
entrypoint = (
|
||||
run_plugin_worker(group_metadata=group_metadata)
|
||||
if wire_codec == "json"
|
||||
else run_plugin_worker(group_metadata=group_metadata, wire_codec=wire_codec)
|
||||
)
|
||||
entrypoint = run_plugin_worker(group_metadata=group_metadata)
|
||||
else:
|
||||
entrypoint = (
|
||||
run_plugin_worker(plugin_dir=plugin_dir)
|
||||
if wire_codec == "json"
|
||||
else run_plugin_worker(plugin_dir=plugin_dir, wire_codec=wire_codec)
|
||||
)
|
||||
entrypoint = run_plugin_worker(plugin_dir=plugin_dir)
|
||||
_run_async_entrypoint(
|
||||
entrypoint,
|
||||
log_message=f"启动插件工作进程:{target}",
|
||||
log_level="debug",
|
||||
context={"plugin_dir": plugin_dir, "wire_codec": wire_codec},
|
||||
context={"plugin_dir": plugin_dir},
|
||||
)
|
||||
|
||||
|
||||
@cli.command(hidden=True)
|
||||
@click.option("--port", default=8765, type=int, help="WebSocket server port")
|
||||
@click.option(
|
||||
"--wire-codec",
|
||||
default="json",
|
||||
show_default=True,
|
||||
type=click.Choice(["json", "msgpack"]),
|
||||
)
|
||||
def websocket(port: int, wire_codec: str) -> None:
|
||||
def websocket(port: int) -> None:
|
||||
"""WebSocket runtime entrypoint kept for standalone bridge scenarios."""
|
||||
entrypoint = (
|
||||
run_websocket_server(port=port)
|
||||
if wire_codec == "json"
|
||||
else run_websocket_server(port=port, wire_codec=wire_codec)
|
||||
)
|
||||
_run_async_entrypoint(
|
||||
entrypoint,
|
||||
run_websocket_server(port=port),
|
||||
log_message=f"启动 WebSocket 服务器,端口:{port}",
|
||||
context={"port": port, "wire_codec": wire_codec},
|
||||
context={"port": port},
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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", ""),
|
||||
|
||||
@@ -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]]:
|
||||
"""获取群组成员列表。
|
||||
|
||||
获取指定群组的成员信息列表。注意仅对群组会话有效。
|
||||
|
||||
@@ -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"]
|
||||
@@ -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", ""))
|
||||
|
||||
@@ -28,18 +28,23 @@ Example:
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, cast
|
||||
from typing import Any, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from .protocol.descriptors import (
|
||||
RESERVED_CAPABILITY_PREFIXES,
|
||||
CapabilityDescriptor,
|
||||
CommandRouteSpec,
|
||||
CommandTrigger,
|
||||
EventTrigger,
|
||||
FilterSpec,
|
||||
MessageTrigger,
|
||||
MessageTypeFilterSpec,
|
||||
Permissions,
|
||||
RESERVED_CAPABILITY_PREFIXES,
|
||||
PlatformFilterSpec,
|
||||
ScheduleTrigger,
|
||||
)
|
||||
|
||||
@@ -69,6 +74,9 @@ class HandlerMeta:
|
||||
contract: str | None = None
|
||||
priority: int = 0
|
||||
permissions: Permissions = field(default_factory=Permissions)
|
||||
filters: list[FilterSpec] = field(default_factory=list)
|
||||
local_filters: list[Any] = field(default_factory=list)
|
||||
command_route: CommandRouteSpec | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
@@ -183,6 +191,7 @@ def on_message(
|
||||
regex: str | None = None,
|
||||
keywords: list[str] | None = None,
|
||||
platforms: list[str] | None = None,
|
||||
message_types: list[str] | None = None,
|
||||
) -> Callable[[HandlerCallable], HandlerCallable]:
|
||||
"""注册消息处理方法。
|
||||
|
||||
@@ -215,12 +224,44 @@ def on_message(
|
||||
regex=regex,
|
||||
keywords=keywords or [],
|
||||
platforms=platforms or [],
|
||||
message_types=message_types or [],
|
||||
)
|
||||
if platforms:
|
||||
meta.filters.append(PlatformFilterSpec(platforms=list(platforms)))
|
||||
if message_types:
|
||||
meta.filters.append(
|
||||
MessageTypeFilterSpec(message_types=list(message_types))
|
||||
)
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def append_filter_meta(
|
||||
func: HandlerCallable,
|
||||
*,
|
||||
specs: list[FilterSpec] | None = None,
|
||||
local_bindings: list[Any] | None = None,
|
||||
) -> HandlerCallable:
|
||||
"""追加过滤器元数据。"""
|
||||
meta = _get_or_create_meta(func)
|
||||
if specs:
|
||||
meta.filters.extend(specs)
|
||||
if local_bindings:
|
||||
meta.local_filters.extend(local_bindings)
|
||||
return func
|
||||
|
||||
|
||||
def set_command_route_meta(
|
||||
func: HandlerCallable,
|
||||
route: CommandRouteSpec,
|
||||
) -> HandlerCallable:
|
||||
"""设置命令路由元数据。"""
|
||||
meta = _get_or_create_meta(func)
|
||||
meta.command_route = route
|
||||
return func
|
||||
|
||||
|
||||
def on_event(event_type: str) -> Callable[[HandlerCallable], HandlerCallable]:
|
||||
"""注册事件处理方法。
|
||||
|
||||
|
||||
@@ -0,0 +1,929 @@
|
||||
# AstrBot SDK 项目完整架构分析文档
|
||||
|
||||
> 作者:whatevertogo
|
||||
|
||||
## 目录
|
||||
|
||||
1. [项目概述](#项目概述)
|
||||
2. [目录结构](#目录结构)
|
||||
3. [核心架构层次](#核心架构层次)
|
||||
4. [协议层设计](#协议层设计)
|
||||
5. [运行时架构](#运行时架构)
|
||||
6. [客户端层设计](#客户端层设计)
|
||||
7. [新旧架构对比](#新旧架构对比)
|
||||
8. [插件开发指南](#插件开发指南)
|
||||
9. [关键设计模式](#关键设计模式)
|
||||
|
||||
---
|
||||
|
||||
## 项目概述
|
||||
|
||||
AstrBot SDK 是一个基于 Python 3.12+ 的机器人插件开发框架,采用**进程隔离**和**能力路由**架构,支持插件的动态加载、独立运行和跨进程通信。
|
||||
|
||||
### 核心特性
|
||||
|
||||
| 特性 | 描述 |
|
||||
|------|------|
|
||||
| 进程隔离 | 每个插件运行在独立 Worker 进程,崩溃不影响其他插件 |
|
||||
| 环境分组 | 多插件可共享同一 Python 虚拟环境,节省资源 |
|
||||
| 能力路由 | 显式声明的 Capability 系统,支持 JSON Schema 验证 |
|
||||
| 流式支持 | 原生支持流式 LLM 调用和增量结果返回 |
|
||||
| 向后兼容 | 完整的旧版 API 兼容层,支持无修改迁移 |
|
||||
| 协议优先 | 基于 v4 协议的统一通信模型,支持多种传输方式 |
|
||||
|
||||
### 技术栈
|
||||
|
||||
- **Python**: 3.12+
|
||||
- **异步框架**: asyncio
|
||||
- **Web 框架**: aiohttp
|
||||
- **数据验证**: pydantic
|
||||
- **日志**: loguru
|
||||
- **配置**: pyyaml
|
||||
- **LLM**: openai, anthropic, google-genai
|
||||
- **包管理**: uv (环境分组)
|
||||
|
||||
---
|
||||
|
||||
## 目录结构
|
||||
|
||||
```
|
||||
astrbot_sdk/ # v4 SDK 主包
|
||||
├── __init__.py # 顶层公共 API 导出
|
||||
├── __main__.py # CLI 入口点 (python -m astrbot_sdk)
|
||||
├── star.py # v4 原生插件基类
|
||||
├── context.py # 运行时上下文 (Context, CancelToken)
|
||||
├── decorators.py # v4 原生装饰器 (on_command, on_message, etc.)
|
||||
├── events.py # v4 原生事件对象 (MessageEvent)
|
||||
├── errors.py # 统一错误模型 (AstrBotError)
|
||||
├── cli.py # 命令行工具 (init/validate/build/dev/run)
|
||||
├── testing.py # 测试辅助模块 (PluginHarness)
|
||||
├── _invocation_context.py # 调用上下文管理 (caller_plugin_scope)
|
||||
├── _testing_support.py # 测试支持工具
|
||||
│
|
||||
├── commands.py # 命令分组工具 (CommandGroup)
|
||||
├── filters.py # 事件过滤器 (PlatformFilter, CustomFilter)
|
||||
├── message_components.py # 消息组件 (Plain, Image, At, etc.)
|
||||
├── message_result.py # 消息结果对象 (MessageChain)
|
||||
├── message_session.py # 会话标识符 (MessageSession)
|
||||
├── schedule.py # 定时任务上下文 (ScheduleContext)
|
||||
├── session_waiter.py # 会话等待器 (SessionController)
|
||||
├── types.py # 参数类型助手 (GreedyStr)
|
||||
│
|
||||
├── clients/ # 能力客户端层
|
||||
│ ├── __init__.py # 客户端公共导出
|
||||
│ ├── _proxy.py # CapabilityProxy 能力代理
|
||||
│ ├── llm.py # LLM 客户端 (chat, chat_raw, stream_chat)
|
||||
│ ├── memory.py # 记忆存储客户端 (search, save, get)
|
||||
│ ├── db.py # KV 存储客户端 (get, set, watch)
|
||||
│ ├── platform.py # 平台消息客户端 (send, send_image)
|
||||
│ ├── http.py # HTTP 注册客户端 (register_api)
|
||||
│ └── metadata.py # 插件元数据客户端 (get_plugin)
|
||||
│
|
||||
├── protocol/ # 协议层
|
||||
│ ├── __init__.py # 协议公共导出
|
||||
│ ├── messages.py # v4 协议消息模型
|
||||
│ ├── descriptors.py # Handler/Capability 描述符
|
||||
│ └── _builtin_schemas.py # 内置能力 JSON Schema
|
||||
│
|
||||
└── runtime/ # 运行时层
|
||||
├── __init__.py # 运行时公共导出 (延迟加载)
|
||||
├── peer.py # 协议对等端 (Peer)
|
||||
├── transport.py # 传输抽象 (Stdio, WebSocket)
|
||||
├── handler_dispatcher.py # Handler 执行分发
|
||||
├── capability_dispatcher.py # Capability 调用分发
|
||||
├── capability_router.py # Capability 路由
|
||||
├── _capability_router_builtins.py # 内置能力处理器
|
||||
├── _loader_support.py # 加载器反射工具
|
||||
├── _streaming.py # 流式执行原语 (StreamExecution)
|
||||
├── loader.py # 插件加载器
|
||||
├── bootstrap.py # 启动引导
|
||||
├── worker.py # Worker 运行时
|
||||
├── supervisor.py # Supervisor 运行时
|
||||
└── environment_groups.py # 环境分组管理
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 核心架构层次
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ 用户层 (Plugin Developer) │
|
||||
├─────────────────────────────────────────────────────────────────┤
|
||||
│ v4 入口: astrbot_sdk.{Star, Context, MessageEvent} │
|
||||
│ 装饰器: on_command, on_message, on_event, on_schedule │
|
||||
│ provide_capability, require_admin │
|
||||
│ 过滤器: PlatformFilter, MessageTypeFilter, CustomFilter │
|
||||
│ 命令组: CommandGroup, command_group │
|
||||
│ 会话: MessageSession, session_waiter │
|
||||
└────────────────────┬────────────────────────────────────────────┘
|
||||
│
|
||||
┌──────────────────▼─────────────────────────────────────────────┐
|
||||
│ 高层 API (High-Level API) │
|
||||
├─────────────────────────────────────────────────────────────────┤
|
||||
│ 能力客户端 (通过 CapabilityProxy 调用): │
|
||||
│ - LLMClient (llm.chat, llm.chat_raw, llm.stream_chat)│
|
||||
│ - MemoryClient (memory.search, memory.save, memory.stats)│
|
||||
│ - DBClient (db.get, db.set, db.watch, db.list) │
|
||||
│ - PlatformClient (platform.send, platform.send_image, ...)│
|
||||
│ - HTTPClient (http.register_api, http.list_apis) │
|
||||
│ - MetadataClient (metadata.get_plugin, metadata.list_plugins)│
|
||||
└────────────────────┬────────────────────────────────────────────┘
|
||||
│
|
||||
┌──────────────────▼─────────────────────────────────────────────┐
|
||||
│ 执行边界 (Execution Boundary) │
|
||||
├─────────────────────────────────────────────────────────────────┤
|
||||
│ runtime 主干: │
|
||||
│ - loader.py (插件发现、加载、环境管理) │
|
||||
│ - bootstrap.py (Supervisor/Worker 启动) │
|
||||
│ - handler_dispatcher.py (Handler 执行分发、参数注入) │
|
||||
│ - capability_dispatcher.py (Capability 调用分发) │
|
||||
│ - capability_router.py (Capability 路由、Schema 验证) │
|
||||
│ - _capability_router_builtins.py (内置能力实现) │
|
||||
│ - _loader_support.py (反射工具、签名验证) │
|
||||
│ - _streaming.py (流式执行原语) │
|
||||
│ - peer.py (协议对等端) │
|
||||
│ - transport.py (传输抽象) │
|
||||
│ - supervisor.py (Supervisor 运行时) │
|
||||
│ - worker.py (Worker 运行时) │
|
||||
│ - environment_groups.py (环境分组规划) │
|
||||
└────────────────────┬────────────────────────────────────────────┘
|
||||
│
|
||||
┌──────────────────▼─────────────────────────────────────────────┐
|
||||
│ 协议与传输 (Protocol & Transport) │
|
||||
├─────────────────────────────────────────────────────────────────┤
|
||||
│ protocol/ │
|
||||
│ - messages.py (协议消息模型) │
|
||||
│ - descriptors.py (Handler/Capability 描述符) │
|
||||
│ - _builtin_schemas.py (内置能力 JSON Schema) │
|
||||
│ transport 实现: │
|
||||
│ - StdioTransport (标准输入输出) │
|
||||
│ - WebSocketServerTransport (WebSocket 服务端) │
|
||||
│ - WebSocketClientTransport (WebSocket 客户端) │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
### 层次职责
|
||||
|
||||
| 层次 | 职责 | 主要模块 |
|
||||
|------|------|---------|
|
||||
| 用户层 | 插件开发者 API | `Star`, `Context`, `MessageEvent`, 装饰器, 过滤器, 命令组 |
|
||||
| 高层 API | 类型化的能力客户端 | `clients/{llm, memory, db, platform, http, metadata}` |
|
||||
| 执行边界 | 插件加载、路由、分发、参数注入 | `runtime/loader.py`, `runtime/*_dispatcher.py` |
|
||||
| 协议层 | 消息模型、描述符、JSON Schema | `protocol/` |
|
||||
| 传输层 | 底层通信抽象 | `runtime/transport.py` |
|
||||
|
||||
### 核心设计原则
|
||||
|
||||
1. **延迟加载**:`runtime/__init__.py` 使用 `__getattr__` 避免导入时加载 websocket/aiohttp 等重型依赖
|
||||
2. **插件身份透传**:通过 `caller_plugin_scope()` 上下文管理器将 plugin_id 注入协议层
|
||||
3. **声明式优先**:所有配置都是数据结构(描述符),便于序列化和跨进程传递
|
||||
4. **类型安全**:使用 Pydantic 模型和类型注解提供验证和 IDE 支持
|
||||
|
||||
---
|
||||
|
||||
## 协议层设计
|
||||
|
||||
### 消息模型
|
||||
|
||||
v4 协议定义了 5 种消息类型:
|
||||
|
||||
| 消息类型 | 用途 | 关键字段 |
|
||||
|---------|------|---------|
|
||||
| `InitializeMessage` | 握手初始化 | `protocol_version`, `peer`, `handlers`, `provided_capabilities` |
|
||||
| `InvokeMessage` | 调用能力 | `capability`, `input`, `stream`, `caller_plugin_id` |
|
||||
| `ResultMessage` | 返回结果 | `success`, `output`, `error`, `kind` |
|
||||
| `EventMessage` | 流式事件 | `phase` (started/delta/completed/failed), `data` |
|
||||
| `CancelMessage` | 取消调用 | `reason` |
|
||||
|
||||
### 错误模型
|
||||
|
||||
`ErrorPayload` 使用字符串 code(而非整数),包含:
|
||||
- `code`: 错误码(如 "capability_not_found")
|
||||
- `message`: 开发者信息
|
||||
- `hint`: 用户友好提示
|
||||
- `retryable`: 是否可重试
|
||||
|
||||
### 握手流程
|
||||
|
||||
```
|
||||
Worker (Plugin) Supervisor (Core)
|
||||
| |
|
||||
| InitializeMessage |
|
||||
| (handlers, capabilities) |
|
||||
|----------------------------->|
|
||||
| | 创建 CapabilityRouter
|
||||
| | 注册 handler.invoke
|
||||
| |
|
||||
| ResultMessage(kind="init") |
|
||||
|<-----------------------------|
|
||||
| | 等待 handler.invoke 调用
|
||||
| | 执行 CapabilityRouter.execute()
|
||||
| |
|
||||
| InvokeMessage(handler.invoke) |
|
||||
|<-----------------------------|
|
||||
| HandlerDispatcher.invoke() |
|
||||
| 执行用户 handler |
|
||||
| |
|
||||
| ResultMessage(output) |
|
||||
|----------------------------->|
|
||||
```
|
||||
|
||||
### 描述符模型
|
||||
|
||||
#### HandlerDescriptor
|
||||
|
||||
```python
|
||||
{
|
||||
"id": "plugin.module:handler_name",
|
||||
"trigger": {
|
||||
"type": "command",
|
||||
"command": "hello",
|
||||
"aliases": ["hi"],
|
||||
"description": "打招呼命令"
|
||||
},
|
||||
"kind": "handler", # handler | hook | tool | session
|
||||
"contract": "message_event", # message_event | schedule
|
||||
"priority": 0,
|
||||
"permissions": {"require_admin": False, "level": 0},
|
||||
"filters": [], # 高级过滤器列表
|
||||
"param_specs": [], # 参数规范
|
||||
"command_route": {...} # 命令路由元信息
|
||||
}
|
||||
```
|
||||
|
||||
#### Trigger 类型
|
||||
|
||||
| 类型 | 关键字段 | 说明 |
|
||||
|------|---------|------|
|
||||
| `CommandTrigger` | command, aliases, platforms | 命令触发 |
|
||||
| `MessageTrigger` | regex, keywords, platforms | 消息触发(正则/关键词) |
|
||||
| `EventTrigger` | event_type | 事件触发 |
|
||||
| `ScheduleTrigger` | cron, interval_seconds | 定时触发(二选一) |
|
||||
|
||||
#### FilterSpec 类型
|
||||
|
||||
| 类型 | 说明 |
|
||||
|------|------|
|
||||
| `PlatformFilterSpec` | 按平台名称过滤 |
|
||||
| `MessageTypeFilterSpec` | 按消息类型过滤 |
|
||||
| `LocalFilterRefSpec` | 引用本地自定义过滤器 |
|
||||
| `CompositeFilterSpec` | 组合过滤器(AND/OR) |
|
||||
|
||||
#### CapabilityDescriptor
|
||||
|
||||
```python
|
||||
{
|
||||
"name": "llm.chat",
|
||||
"description": "发送对话请求,返回文本",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"prompt": {"type": "string"}},
|
||||
"required": ["prompt"]
|
||||
},
|
||||
"output_schema": {
|
||||
"type": "object",
|
||||
"properties": {"text": {"type": "string"}},
|
||||
"required": ["text"]
|
||||
},
|
||||
"supports_stream": False,
|
||||
"cancelable": False
|
||||
}
|
||||
```
|
||||
|
||||
### 命名空间治理
|
||||
|
||||
**保留前缀**:
|
||||
- `handler.` - 内部 handler.invoke
|
||||
- `system.` - 系统内置能力
|
||||
- `internal.` - 内部使用
|
||||
|
||||
**内置能力命名空间**:`llm`, `memory`, `db`, `platform`, `http`, `metadata`
|
||||
|
||||
### 内置 Capabilities (38个)
|
||||
|
||||
#### LLM 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `llm.chat` | 同步对话,返回文本 |
|
||||
| `llm.chat_raw` | 同步对话,返回完整响应(含 usage、tool_calls) |
|
||||
| `llm.stream_chat` | 流式对话 |
|
||||
|
||||
#### Memory 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `memory.search` | 语义搜索记忆 |
|
||||
| `memory.save` | 保存记忆 |
|
||||
| `memory.save_with_ttl` | 保存带过期时间的记忆 |
|
||||
| `memory.get` | 读取单条记忆 |
|
||||
| `memory.get_many` | 批量获取记忆 |
|
||||
| `memory.delete` | 删除记忆 |
|
||||
| `memory.delete_many` | 批量删除记忆 |
|
||||
| `memory.stats` | 获取记忆统计信息 |
|
||||
|
||||
#### DB 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `db.get` | 读取 KV |
|
||||
| `db.set` | 写入 KV |
|
||||
| `db.delete` | 删除 KV |
|
||||
| `db.list` | 列出 KV 键(支持前缀过滤) |
|
||||
| `db.get_many` | 批量读取 KV |
|
||||
| `db.set_many` | 批量写入 KV |
|
||||
| `db.watch` | 订阅 KV 变更(流式) |
|
||||
|
||||
#### Platform 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `platform.send` | 发送文本消息 |
|
||||
| `platform.send_image` | 发送图片 |
|
||||
| `platform.send_chain` | 发送消息链 |
|
||||
| `platform.get_members` | 获取群成员 |
|
||||
|
||||
#### HTTP 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `http.register_api` | 注册 HTTP API 端点 |
|
||||
| `http.unregister_api` | 注销 HTTP API 端点 |
|
||||
| `http.list_apis` | 列出已注册的 API |
|
||||
|
||||
#### Metadata 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `metadata.get_plugin` | 获取单个插件元数据 |
|
||||
| `metadata.list_plugins` | 列出所有插件元数据 |
|
||||
| `metadata.get_plugin_config` | 获取当前插件配置 |
|
||||
|
||||
#### System 命名空间
|
||||
|
||||
| 能力 | 说明 |
|
||||
|------|------|
|
||||
| `system.get_data_dir` | 获取插件数据目录 |
|
||||
| `system.text_to_image` | 文本转图片 |
|
||||
| `system.html_render` | 渲染 HTML 模板 |
|
||||
| `system.session_waiter.register` | 注册会话等待器 |
|
||||
| `system.session_waiter.unregister` | 注销会话等待器 |
|
||||
| `system.event.react` | 发送表情回应 |
|
||||
| `system.event.send_typing` | 发送输入中状态 |
|
||||
| `system.event.send_streaming` | 开始流式消息会话 |
|
||||
| `system.event.send_streaming_chunk` | 推送流式消息分片 |
|
||||
| `system.event.send_streaming_close` | 关闭流式消息会话 |
|
||||
|
||||
---
|
||||
|
||||
## 运行时架构
|
||||
|
||||
### 组件关系图
|
||||
|
||||
```
|
||||
┌──────────────┐
|
||||
│ AstrBot │
|
||||
│ Core │
|
||||
└──────┬─────┘
|
||||
│
|
||||
┌──────▼─────┐
|
||||
│ Supervisor │
|
||||
│ Runtime │
|
||||
└──────┬─────┘
|
||||
│
|
||||
┌──────────────────┼──────────────────┐
|
||||
│ │ │
|
||||
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
|
||||
│ Peer │ │ Peer │ │ Peer │
|
||||
│ (stdio) │ │ (stdio) │ │ (stdio) │
|
||||
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
|
||||
│ │ │
|
||||
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
|
||||
│ Worker │ │ Worker │ │ Worker │
|
||||
│ Runtime │ │ Runtime │ │ Runtime │
|
||||
└─────┬─────┘ └─────┬─────┘ └─────┬─────┘
|
||||
│ │ │
|
||||
┌─────▼─────┐ ┌─────▼─────┐ ┌─────▼─────┐
|
||||
│ Plugin A │ │ Plugin B │ │ Plugin C │
|
||||
│ (v4/old) │ │ (v4/old) │ │ (v4/old) │
|
||||
└───────────┘ └───────────┘ └───────────┘
|
||||
```
|
||||
|
||||
### SupervisorRuntime
|
||||
|
||||
职责:管理多个 Worker 进程,聚合所有 handler
|
||||
|
||||
```python
|
||||
class SupervisorRuntime:
|
||||
def __init__(self, *, transport, plugins_dir, env_manager):
|
||||
self.transport = transport # 与 Core 的传输层
|
||||
self.plugins_dir = plugins_dir # 插件目录
|
||||
self.capability_router = CapabilityRouter() # 能力路由器
|
||||
self.peer = Peer(...) # 与 Core 的对等端
|
||||
self.worker_sessions = {} # Worker 会话映射
|
||||
self.handler_to_worker = {} # Handler → Worker 映射
|
||||
|
||||
async def start(self):
|
||||
# 1. 发现所有插件
|
||||
discovery = discover_plugins(self.plugins_dir)
|
||||
|
||||
# 2. 规划环境分组
|
||||
plan_result = self.env_manager.plan(discovery.plugins)
|
||||
|
||||
# 3. 为每个分组启动 Worker
|
||||
for group in plan_result.groups:
|
||||
session = WorkerSession(group=group, ...)
|
||||
await session.start()
|
||||
self.worker_sessions[group.id] = session
|
||||
|
||||
# 4. 聚合所有 handler 和 capability
|
||||
await self.peer.initialize(
|
||||
handlers=[...],
|
||||
provided_capabilities=self.capability_router.descriptors()
|
||||
)
|
||||
```
|
||||
|
||||
### WorkerSession
|
||||
|
||||
职责:管理单个 Worker 进程的生命周期
|
||||
|
||||
```python
|
||||
class WorkerSession:
|
||||
def __init__(self, *, group, env_manager, capability_router):
|
||||
self.group = group # 环境分组
|
||||
self.peer = Peer(...) # 与 Worker 的对等端
|
||||
self.capability_router = capability_router
|
||||
self.handlers = [] # Worker 注册的 handlers
|
||||
self.provided_capabilities = [] # Worker 提供的 capabilities
|
||||
|
||||
async def start(self):
|
||||
# 启动 Worker 子进程
|
||||
python_path = self.env_manager.prepare_group_environment(self.group)
|
||||
transport = StdioTransport(
|
||||
command=[python_path, "-m", "astrbot_sdk", "worker", "--group-metadata", ...]
|
||||
)
|
||||
self.peer = Peer(transport=transport, ...)
|
||||
|
||||
# 等待 Worker 初始化完成
|
||||
await self.peer.start()
|
||||
await self.peer.wait_until_remote_initialized()
|
||||
|
||||
# 获取 Worker 的注册信息
|
||||
self.handlers = list(self.peer.remote_handlers)
|
||||
self.provided_capabilities = list(self.peer.remote_provided_capabilities)
|
||||
|
||||
async def invoke_capability(self, capability_name, payload, *, request_id):
|
||||
# 转发能力调用到 Worker
|
||||
return await self.peer.invoke(capability_name, payload, request_id=request_id)
|
||||
```
|
||||
|
||||
### PluginWorkerRuntime
|
||||
|
||||
职责:Worker 进程内的插件加载与执行
|
||||
|
||||
```python
|
||||
class PluginWorkerRuntime:
|
||||
def __init__(self, *, plugin_dir, transport):
|
||||
self.plugin = load_plugin_spec(plugin_dir)
|
||||
self.loaded_plugin = load_plugin(self.plugin)
|
||||
self.peer = Peer(transport=transport, ...)
|
||||
self.dispatcher = HandlerDispatcher(...)
|
||||
self.capability_dispatcher = CapabilityDispatcher(...)
|
||||
|
||||
async def start(self):
|
||||
# 1. 向 Supervisor 注册 handlers 和 capabilities
|
||||
await self.peer.initialize(
|
||||
handlers=[h.descriptor for h in self.loaded_plugin.handlers],
|
||||
provided_capabilities=[c.descriptor for c in self.loaded_plugin.capabilities]
|
||||
)
|
||||
|
||||
# 2. 执行 on_start 生命周期
|
||||
await self._run_lifecycle("on_start")
|
||||
|
||||
# 3. 设置消息处理器
|
||||
self.peer.set_invoke_handler(self._handle_invoke)
|
||||
self.peer.set_cancel_handler(self._handle_cancel)
|
||||
|
||||
async def _handle_invoke(self, message, cancel_token):
|
||||
if message.capability == "handler.invoke":
|
||||
return await self.dispatcher.invoke(message, cancel_token)
|
||||
return await self.capability_dispatcher.invoke(message, cancel_token)
|
||||
```
|
||||
|
||||
### HandlerDispatcher
|
||||
|
||||
职责:将 handler.invoke 请求转成真实 Python 调用
|
||||
|
||||
```python
|
||||
class HandlerDispatcher:
|
||||
def __init__(self, *, plugin_id, peer, handlers):
|
||||
self._handlers = {item.descriptor.id: item for item in handlers}
|
||||
self._peer = peer
|
||||
self._active = {} # request_id → (task, cancel_token)
|
||||
|
||||
async def invoke(self, message, cancel_token):
|
||||
# 1. 查找 handler
|
||||
loaded = self._handlers[message.input["handler_id"]]
|
||||
|
||||
# 2. 创建上下文
|
||||
ctx = Context(peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token)
|
||||
event = MessageEvent.from_payload(message.input["event"], context=ctx)
|
||||
|
||||
# 3. 构建参数 (支持类型注解注入)
|
||||
args = self._build_args(loaded.callable, event, ctx)
|
||||
|
||||
# 4. 执行 handler
|
||||
result = loaded.callable(*args)
|
||||
|
||||
# 5. 处理返回值
|
||||
await self._consume_result(result, event, ctx)
|
||||
```
|
||||
|
||||
**参数注入优先级**:
|
||||
1. 按类型注解注入(`MessageEvent`, `Context`)
|
||||
2. 按参数名注入(`event`, `ctx`, `context`)
|
||||
3. 从 legacy_args 注入(命令参数等)
|
||||
|
||||
### CapabilityRouter
|
||||
|
||||
职责:能力注册、发现和执行路由
|
||||
|
||||
```python
|
||||
class CapabilityRouter:
|
||||
def __init__(self):
|
||||
self._registrations = {} # capability_name → registration
|
||||
self.db_store = {} # 内置 KV 存储
|
||||
self.memory_store = {} # 内置记忆存储
|
||||
self._register_builtin_capabilities()
|
||||
|
||||
def register(self, descriptor, *, call_handler, stream_handler, finalize):
|
||||
"""注册能力"""
|
||||
self._registrations[descriptor.name] = _CapabilityRegistration(
|
||||
descriptor=descriptor,
|
||||
call_handler=call_handler,
|
||||
stream_handler=stream_handler,
|
||||
finalize=finalize
|
||||
)
|
||||
|
||||
async def execute(self, capability, payload, *, stream, cancel_token, request_id):
|
||||
"""执行能力调用"""
|
||||
registration = self._registrations[capability]
|
||||
|
||||
if stream:
|
||||
# 流式调用
|
||||
raw_execution = registration.stream_handler(request_id, payload, cancel_token)
|
||||
return StreamExecution(iterator=raw_execution, finalize=finalize)
|
||||
else:
|
||||
# 同步调用
|
||||
output = await registration.call_handler(request_id, payload, cancel_token)
|
||||
return output
|
||||
```
|
||||
|
||||
### 环境分组管理
|
||||
|
||||
```python
|
||||
class EnvironmentPlanner:
|
||||
def plan(self, plugins):
|
||||
"""根据 Python 版本和依赖兼容性分组"""
|
||||
# 1. 按版本分组
|
||||
# 2. 按依赖兼容性合并
|
||||
# 3. 生成分组元数据
|
||||
return EnvironmentPlanResult(groups=[...])
|
||||
|
||||
class GroupEnvironmentManager:
|
||||
def prepare(self, group):
|
||||
"""准备分组虚拟环境"""
|
||||
# 1. 生成 lock/source/metadata 工件
|
||||
# 2. 必要时重建虚拟环境
|
||||
# 3. 返回 Python 解释器路径
|
||||
return venv_python_path
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 客户端层设计
|
||||
|
||||
### 客户端架构
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────────┐
|
||||
│ User Plugin │
|
||||
├─────────────────────────────────────────────────────────────┤
|
||||
│ ctx.llm.chat() │
|
||||
│ ctx.memory.save() │
|
||||
│ ctx.db.set() │
|
||||
│ ctx.platform.send() │
|
||||
└────────────┬──────────────────────────────────────────────┘
|
||||
│
|
||||
┌────────────▼──────────────────────────────────────────────┐
|
||||
│ CapabilityProxy │
|
||||
│ - call(name, payload) │
|
||||
│ - stream(name, payload) │
|
||||
└────────────┬──────────────────────────────────────────────┘
|
||||
│
|
||||
┌────────────▼──────────────────────────────────────────────┐
|
||||
│ Peer │
|
||||
│ - invoke(capability, payload, stream=False) │
|
||||
│ - invoke_stream(capability, payload) │
|
||||
└────────────┬──────────────────────────────────────────────┘
|
||||
│
|
||||
┌────────────▼──────────────────────────────────────────────┐
|
||||
│ Transport │
|
||||
│ - send(json_string) │
|
||||
└─────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
### CapabilityProxy
|
||||
|
||||
职责:封装 Peer 的能力调用接口
|
||||
|
||||
```python
|
||||
class CapabilityProxy:
|
||||
def __init__(self, peer):
|
||||
self._peer = peer
|
||||
|
||||
async def call(self, name, payload):
|
||||
"""普通能力调用"""
|
||||
# 1. 检查能力是否可用
|
||||
descriptor = self._peer.remote_capability_map.get(name)
|
||||
if descriptor is None:
|
||||
raise AstrBotError.capability_not_found(name)
|
||||
|
||||
# 2. 调用 Peer.invoke
|
||||
return await self._peer.invoke(name, payload, stream=False)
|
||||
|
||||
async def stream(self, name, payload):
|
||||
"""流式能力调用"""
|
||||
# 1. 检查流式支持
|
||||
descriptor = self._peer.remote_capability_map.get(name)
|
||||
if not descriptor.supports_stream:
|
||||
raise AstrBotError.invalid_input(f"{name} 不支持 stream")
|
||||
|
||||
# 2. 调用 Peer.invoke_stream
|
||||
event_stream = await self._peer.invoke_stream(name, payload)
|
||||
async for event in event_stream:
|
||||
if event.phase == "delta":
|
||||
yield event.data
|
||||
```
|
||||
|
||||
### LLMClient
|
||||
|
||||
```python
|
||||
class LLMClient:
|
||||
def __init__(self, proxy: CapabilityProxy):
|
||||
self._proxy = proxy
|
||||
|
||||
async def chat(self, prompt, *, system=None, history=None, **kwargs) -> str:
|
||||
"""发送聊天请求,返回文本"""
|
||||
output = await self._proxy.call("llm.chat", {
|
||||
"prompt": prompt,
|
||||
"system": system,
|
||||
"history": self._serialize_history(history),
|
||||
**kwargs
|
||||
})
|
||||
return output["text"]
|
||||
|
||||
async def chat_raw(self, prompt, **kwargs) -> LLMResponse:
|
||||
"""发送聊天请求,返回完整响应"""
|
||||
output = await self._proxy.call("llm.chat_raw", {"prompt": prompt, **kwargs})
|
||||
return LLMResponse.model_validate(output)
|
||||
|
||||
async def stream_chat(self, prompt, **kwargs) -> AsyncGenerator[str]:
|
||||
"""流式聊天"""
|
||||
async for delta in self._proxy.stream("llm.stream_chat", {"prompt": prompt, **kwargs}):
|
||||
yield delta["text"]
|
||||
```
|
||||
|
||||
### 其他客户端
|
||||
|
||||
| 客户端 | 主要方法 | 对应 Capability |
|
||||
|--------|---------|-----------------|
|
||||
| `MemoryClient` | `search()`, `save()`, `save_with_ttl()`, `get()`, `get_many()`, `delete()`, `delete_many()`, `stats()` | `memory.*` |
|
||||
| `DBClient` | `get()`, `set()`, `delete()`, `list()`, `get_many()`, `set_many()`, `watch()` | `db.*` |
|
||||
| `PlatformClient` | `send()`, `send_image()`, `send_chain()`, `get_members()` | `platform.*` |
|
||||
| `HTTPClient` | `register_api()`, `unregister_api()`, `list_apis()` | `http.*` |
|
||||
| `MetadataClient` | `get_plugin()`, `list_plugins()`, `get_current_plugin()`, `get_plugin_config()` | `metadata.*` |
|
||||
|
||||
---
|
||||
|
||||
## 新旧架构对比
|
||||
|
||||
### 协议对比
|
||||
|
||||
| 特性 | 旧版 JSON-RPC | 新版 v4 协议 |
|
||||
|------|---------------|--------------|
|
||||
| 消息格式 | `{"jsonrpc": "2.0", ...}` | `{"type": "invoke", ...}` |
|
||||
| 方法区分 | `method` 字段 | `type` 字段 |
|
||||
| 错误码 | 整数 (`-32000`) | 字符串 (`"internal_error"`) |
|
||||
| 流式支持 | 独立 notification 方法 | 统一 `EventMessage` phase |
|
||||
| 握手 | `handshake` method | `InitializeMessage` type |
|
||||
| 能力声明 | 隐式(method 名称) | 显式 `CapabilityDescriptor` |
|
||||
|
||||
### 运行时对比
|
||||
|
||||
| 特性 | 旧版 | 新版 |
|
||||
|------|------|------|
|
||||
| Peer 抽象 | 分离 `JSONRPCClient/Server` | 统一 `Peer` |
|
||||
| Handler 分发 | 直接调用 `handler(event)` | `HandlerDispatcher` 参数注入 |
|
||||
| 能力路由 | 无显式路由 | `CapabilityRouter` |
|
||||
| 环境管理 | 无 | `PluginEnvironmentManager` 分组 |
|
||||
| 传输层 | 每个实现处理 JSON-RPC | 传输层只处理字符串 |
|
||||
|
||||
### 代码对比
|
||||
|
||||
#### 旧版 Handler
|
||||
|
||||
```python
|
||||
from astrbot.api.star import Star
|
||||
from astrbot.api.event import AstrMessageEvent
|
||||
|
||||
class MyPlugin(Star):
|
||||
@command_handler("hello", aliases=["hi"])
|
||||
def hello_handler(self, event: AstrMessageEvent):
|
||||
reply = self.call_context_function("llm_generate", prompt=event.message_plain)
|
||||
event.reply(reply)
|
||||
```
|
||||
|
||||
#### 新版 Handler
|
||||
|
||||
```python
|
||||
from astrbot_sdk import Star, Context, MessageEvent
|
||||
from astrbot_sdk.decorators import on_command
|
||||
|
||||
class MyPlugin(Star):
|
||||
@on_command("hello", aliases=["hi"])
|
||||
async def hello(self, event: MessageEvent, ctx: Context) -> None:
|
||||
reply = await ctx.llm.chat(event.text)
|
||||
await event.reply(reply)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 新旧架构对比
|
||||
|
||||
---
|
||||
|
||||
## 插件开发指南
|
||||
|
||||
### v4 原生插件
|
||||
|
||||
#### plugin.yaml
|
||||
|
||||
```yaml
|
||||
_schema_version: 2
|
||||
name: my_plugin
|
||||
author: your_name
|
||||
version: 1.0.0
|
||||
runtime:
|
||||
python: "3.12"
|
||||
components:
|
||||
- class: main:MyPlugin
|
||||
```
|
||||
|
||||
#### main.py
|
||||
|
||||
```python
|
||||
from astrbot_sdk import Star, Context, MessageEvent
|
||||
from astrbot_sdk.decorators import on_command, on_message, provide_capability
|
||||
|
||||
class MyPlugin(Star):
|
||||
# 命令处理器
|
||||
@on_command("hello", aliases=["hi"])
|
||||
async def hello(self, event: MessageEvent, ctx: Context) -> None:
|
||||
await event.reply(f"你好,{event.user_id}!")
|
||||
|
||||
# 消息处理器
|
||||
@on_message(keywords=["帮助"])
|
||||
async def help(self, event: MessageEvent, ctx: Context) -> None:
|
||||
await event.reply("可用命令:hello, help")
|
||||
|
||||
# 提供能力
|
||||
@provide_capability(
|
||||
"my_plugin.calculate",
|
||||
description="执行计算",
|
||||
input_schema={
|
||||
"type": "object",
|
||||
"properties": {"x": {"type": "number"}},
|
||||
"required": ["x"]
|
||||
},
|
||||
output_schema={
|
||||
"type": "object",
|
||||
"properties": {"result": {"type": "number"}},
|
||||
"required": ["result"]
|
||||
}
|
||||
)
|
||||
async def calculate_capability(
|
||||
self,
|
||||
payload: dict,
|
||||
ctx: Context
|
||||
) -> dict:
|
||||
x = payload.get("x", 0)
|
||||
return {"result": x * 2}
|
||||
```
|
||||
|
||||
### 生命周期钩子
|
||||
|
||||
| 钩子 | 说明 |
|
||||
|------|------|
|
||||
| `on_start()` | 插件启动时调用 |
|
||||
| `on_stop()` | 插件停止时调用 |
|
||||
| `on_error(exc, event, ctx)` | Handler 执行出错时调用 |
|
||||
|
||||
---
|
||||
|
||||
## 关键设计模式
|
||||
|
||||
### 1. 协议优先模式
|
||||
|
||||
- 所有跨进程通信都通过 v4 协议
|
||||
- 传输层只处理字符串,协议由 Peer 层处理
|
||||
- 支持多种传输方式(Stdio, WebSocket)
|
||||
|
||||
### 2. 能力路由模式
|
||||
|
||||
- 显式声明 Capability 和输入/输出 Schema
|
||||
- 通过 CapabilityRouter 统一路由
|
||||
- 支持同步和流式两种调用模式
|
||||
- 冲突处理:保留命名空间冲突直接跳过,非保留命名空间冲突自动添加插件名前缀
|
||||
|
||||
### 3. 环境分组模式
|
||||
|
||||
- 多插件可共享同一 Python 虚拟环境
|
||||
- 按版本和依赖兼容性自动分组
|
||||
- 节省资源,加快启动速度
|
||||
|
||||
### 4. 参数注入模式
|
||||
|
||||
- HandlerDispatcher 支持类型注解注入
|
||||
- 优先级:类型注解 > 参数名 > legacy_args
|
||||
- 支持可选类型 `Optional[Type]`
|
||||
|
||||
### 5. 取消传播模式
|
||||
|
||||
- CancelToken 统一取消机制
|
||||
- 跨进程取消通过 CancelMessage
|
||||
- 早到取消避免竞态条件
|
||||
|
||||
### 6. 插件隔离模式
|
||||
|
||||
- 每个插件运行在独立 Worker 进程
|
||||
- 崩溃不影响其他插件
|
||||
- 支持 GroupWorkerRuntime 共享环境
|
||||
|
||||
### 7. 热重载模式
|
||||
|
||||
- `dev --watch` 支持文件变更检测
|
||||
- 按插件目录清理 `sys.modules` 缓存
|
||||
- 确保代码变更后正确重载
|
||||
|
||||
---
|
||||
|
||||
## 附录:关键文件速查
|
||||
|
||||
| 文件 | 核心类/函数 | 说明 |
|
||||
|------|------------|------|
|
||||
| `astrbot_sdk/__init__.py` | `Star`, `Context`, `MessageEvent` | 顶层入口 |
|
||||
| `astrbot_sdk/star.py` | `Star` | v4 原生插件基类 |
|
||||
| `astrbot_sdk/context.py` | `Context` | 运行时上下文 |
|
||||
| `astrbot_sdk/decorators.py` | `on_command`, `on_message`, `provide_capability` | v4 装饰器 |
|
||||
| `astrbot_sdk/errors.py` | `AstrBotError` | 统一错误模型 |
|
||||
| `astrbot_sdk/cli.py` | CLI 命令 | 命令行工具(init/validate/build/dev/run/worker/websocket) |
|
||||
| `astrbot_sdk/testing.py` | `PluginHarness`, `MockContext` | 测试辅助 |
|
||||
| `astrbot_sdk/commands.py` | `CommandGroup`, `command_group` | 命令分组工具 |
|
||||
| `astrbot_sdk/filters.py` | `PlatformFilter`, `CustomFilter`, `all_of`, `any_of` | 事件过滤器 |
|
||||
| `astrbot_sdk/message_result.py` | `MessageChain`, `MessageEventResult` | 消息结果对象 |
|
||||
| `astrbot_sdk/message_session.py` | `MessageSession` | 会话标识符 |
|
||||
| `astrbot_sdk/schedule.py` | `ScheduleContext` | 定时任务上下文 |
|
||||
| `astrbot_sdk/session_waiter.py` | `SessionController`, `SessionWaiterManager` | 会话等待器 |
|
||||
| `astrbot_sdk/types.py` | `GreedyStr` | 参数类型助手 |
|
||||
| `astrbot_sdk/runtime/__init__.py` | 延迟导出 | 运行时公共 API(延迟加载) |
|
||||
| `astrbot_sdk/runtime/peer.py` | `Peer` | 协议对等端 |
|
||||
| `astrbot_sdk/runtime/supervisor.py` | `SupervisorRuntime` | Supervisor 运行时 |
|
||||
| `astrbot_sdk/runtime/worker.py` | `PluginWorkerRuntime` | Worker 运行时 |
|
||||
| `astrbot_sdk/runtime/loader.py` | `load_plugin()`, `_ResolvedComponent` | 插件加载 |
|
||||
| `astrbot_sdk/runtime/_loader_support.py` | `build_param_specs`, `is_injected_parameter` | 加载器反射工具 |
|
||||
| `astrbot_sdk/runtime/_streaming.py` | `StreamExecution` | 流式执行原语 |
|
||||
| `astrbot_sdk/runtime/handler_dispatcher.py` | `HandlerDispatcher` | Handler 执行分发 |
|
||||
| `astrbot_sdk/runtime/capability_dispatcher.py` | `CapabilityDispatcher` | Capability 调用分发 |
|
||||
| `astrbot_sdk/runtime/capability_router.py` | `CapabilityRouter` | Capability 路由 |
|
||||
| `astrbot_sdk/runtime/_capability_router_builtins.py` | `BuiltinCapabilityRouterMixin` | 内置能力处理器 |
|
||||
| `astrbot_sdk/runtime/environment_groups.py` | `EnvironmentGroup` | 环境分组 |
|
||||
| `astrbot_sdk/protocol/messages.py` | `InitializeMessage`, `InvokeMessage` | 协议消息 |
|
||||
| `astrbot_sdk/protocol/descriptors.py` | `HandlerDescriptor`, `CapabilityDescriptor` | 描述符 |
|
||||
| `astrbot_sdk/protocol/_builtin_schemas.py` | `BUILTIN_CAPABILITY_SCHEMAS` | 内置能力 JSON Schema |
|
||||
| `astrbot_sdk/clients/_proxy.py` | `CapabilityProxy` | 能力代理 |
|
||||
| `astrbot_sdk/clients/llm.py` | `LLMClient` | LLM 客户端 |
|
||||
| `astrbot_sdk/clients/memory.py` | `MemoryClient` | 记忆客户端 |
|
||||
| `astrbot_sdk/clients/db.py` | `DBClient` | 数据库客户端 |
|
||||
| `astrbot_sdk/clients/platform.py` | `PlatformClient` | 平台客户端 |
|
||||
| `astrbot_sdk/clients/http.py` | `HTTPClient` | HTTP 客户端 |
|
||||
| `astrbot_sdk/clients/metadata.py` | `MetadataClient`, `PluginMetadata` | 元数据客户端 |
|
||||
| `astrbot_sdk/message_components.py` | `Plain`, `Image`, `At`, `Reply` | 消息组件 |
|
||||
| `astrbot_sdk/events.py` | `MessageEvent` | 事件对象 |
|
||||
| `astrbot_sdk/_testing_support.py` | 测试工具 | 测试支持 |
|
||||
|
||||
---
|
||||
|
||||
> 本文档描述 AstrBot SDK v4 的设计与实现思想
|
||||
> 如有疑问请查阅源代码或提交 Issue
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -9,11 +9,18 @@
|
||||
|
||||
from .descriptors import (
|
||||
CapabilityDescriptor,
|
||||
CommandRouteSpec,
|
||||
CommandTrigger,
|
||||
CompositeFilterSpec,
|
||||
EventTrigger,
|
||||
FilterSpec,
|
||||
HandlerDescriptor,
|
||||
LocalFilterRefSpec,
|
||||
MessageTrigger,
|
||||
MessageTypeFilterSpec,
|
||||
ParamSpec,
|
||||
Permissions,
|
||||
PlatformFilterSpec,
|
||||
ScheduleTrigger,
|
||||
SessionRef,
|
||||
Trigger,
|
||||
@@ -33,17 +40,24 @@ from .messages import (
|
||||
|
||||
__all__ = [
|
||||
"CapabilityDescriptor",
|
||||
"CommandRouteSpec",
|
||||
"CommandTrigger",
|
||||
"CancelMessage",
|
||||
"CompositeFilterSpec",
|
||||
"ErrorPayload",
|
||||
"EventTrigger",
|
||||
"EventMessage",
|
||||
"FilterSpec",
|
||||
"HandlerDescriptor",
|
||||
"InitializeMessage",
|
||||
"InitializeOutput",
|
||||
"InvokeMessage",
|
||||
"LocalFilterRefSpec",
|
||||
"MessageTrigger",
|
||||
"MessageTypeFilterSpec",
|
||||
"ParamSpec",
|
||||
"PeerInfo",
|
||||
"PlatformFilterSpec",
|
||||
"Permissions",
|
||||
"ProtocolMessage",
|
||||
"ResultMessage",
|
||||
|
||||
@@ -0,0 +1,552 @@
|
||||
"""Builtin protocol schema constants.
|
||||
|
||||
本模块定义了 AstrBot SDK v4 协议中所有内置能力的 JSON Schema。
|
||||
这些 Schema 用于:
|
||||
1. 验证能力调用的输入参数是否符合预期格式
|
||||
2. 生成能力描述文档,供插件开发者参考
|
||||
3. 确保跨进程/跨语言调用时的类型安全
|
||||
|
||||
所有 Schema 遵循 JSON Schema 规范,支持基本类型检查、必填字段、数组元素约束等。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
JSONSchema = dict[str, Any]
|
||||
|
||||
|
||||
def _object_schema(
|
||||
*,
|
||||
required: tuple[str, ...] = (),
|
||||
**properties: Any,
|
||||
) -> JSONSchema:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": properties,
|
||||
"required": list(required),
|
||||
}
|
||||
|
||||
|
||||
def _nullable(schema: JSONSchema) -> JSONSchema:
|
||||
return {"anyOf": [schema, {"type": "null"}]}
|
||||
|
||||
|
||||
_OPTIONAL_CHAT_PROPERTIES: dict[str, Any] = {
|
||||
"system": {"type": "string"},
|
||||
"history": {"type": "array", "items": {"type": "object"}},
|
||||
"contexts": {"type": "array", "items": {"type": "object"}},
|
||||
"provider_id": {"type": "string"},
|
||||
"tool_calls_result": {"type": "array", "items": {"type": "object"}},
|
||||
"model": {"type": "string"},
|
||||
"temperature": {"type": "number"},
|
||||
"image_urls": {"type": "array", "items": {"type": "string"}},
|
||||
"tools": {"type": "array"},
|
||||
"max_steps": {"type": "integer"},
|
||||
}
|
||||
|
||||
LLM_CHAT_INPUT_SCHEMA = _object_schema(
|
||||
required=("prompt",),
|
||||
prompt={"type": "string"},
|
||||
**_OPTIONAL_CHAT_PROPERTIES,
|
||||
)
|
||||
LLM_CHAT_OUTPUT_SCHEMA = _object_schema(required=("text",), text={"type": "string"})
|
||||
LLM_CHAT_RAW_INPUT_SCHEMA = _object_schema(
|
||||
required=("prompt",),
|
||||
prompt={"type": "string"},
|
||||
**_OPTIONAL_CHAT_PROPERTIES,
|
||||
)
|
||||
LLM_CHAT_RAW_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("text",),
|
||||
text={"type": "string"},
|
||||
usage=_nullable({"type": "object"}),
|
||||
finish_reason=_nullable({"type": "string"}),
|
||||
tool_calls={"type": "array", "items": {"type": "object"}},
|
||||
role=_nullable({"type": "string"}),
|
||||
reasoning_content=_nullable({"type": "string"}),
|
||||
reasoning_signature=_nullable({"type": "string"}),
|
||||
)
|
||||
LLM_STREAM_CHAT_INPUT_SCHEMA = _object_schema(
|
||||
required=("prompt",),
|
||||
prompt={"type": "string"},
|
||||
**_OPTIONAL_CHAT_PROPERTIES,
|
||||
)
|
||||
LLM_STREAM_CHAT_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("text",), text={"type": "string"}
|
||||
)
|
||||
MEMORY_SEARCH_INPUT_SCHEMA = _object_schema(
|
||||
required=("query",), query={"type": "string"}
|
||||
)
|
||||
MEMORY_SEARCH_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("items",),
|
||||
items={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
MEMORY_SAVE_INPUT_SCHEMA = _object_schema(
|
||||
required=("key", "value"),
|
||||
key={"type": "string"},
|
||||
value={"type": "object"},
|
||||
)
|
||||
MEMORY_SAVE_OUTPUT_SCHEMA = _object_schema()
|
||||
MEMORY_GET_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
|
||||
MEMORY_GET_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("value",),
|
||||
value=_nullable({"type": "object"}),
|
||||
)
|
||||
MEMORY_DELETE_INPUT_SCHEMA = _object_schema(
|
||||
required=("key",),
|
||||
key={"type": "string"},
|
||||
)
|
||||
MEMORY_DELETE_OUTPUT_SCHEMA = _object_schema()
|
||||
MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA = _object_schema(
|
||||
required=("key", "value", "ttl_seconds"),
|
||||
key={"type": "string"},
|
||||
value={"type": "object"},
|
||||
ttl_seconds={"type": "integer", "minimum": 1},
|
||||
)
|
||||
MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA = _object_schema()
|
||||
MEMORY_GET_MANY_INPUT_SCHEMA = _object_schema(
|
||||
required=("keys",),
|
||||
keys={"type": "array", "items": {"type": "string"}},
|
||||
)
|
||||
MEMORY_GET_MANY_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("items",),
|
||||
items={
|
||||
"type": "array",
|
||||
"items": _object_schema(
|
||||
required=("key", "value"),
|
||||
key={"type": "string"},
|
||||
value=_nullable({"type": "object"}),
|
||||
),
|
||||
},
|
||||
)
|
||||
MEMORY_DELETE_MANY_INPUT_SCHEMA = _object_schema(
|
||||
required=("keys",),
|
||||
keys={"type": "array", "items": {"type": "string"}},
|
||||
)
|
||||
MEMORY_DELETE_MANY_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("deleted_count",),
|
||||
deleted_count={"type": "integer"},
|
||||
)
|
||||
MEMORY_STATS_INPUT_SCHEMA = _object_schema()
|
||||
MEMORY_STATS_OUTPUT_SCHEMA = _object_schema(
|
||||
total_items={"type": "integer"},
|
||||
total_bytes=_nullable({"type": "integer"}),
|
||||
plugin_id=_nullable({"type": "string"}),
|
||||
ttl_entries=_nullable({"type": "integer"}),
|
||||
)
|
||||
SYSTEM_GET_DATA_DIR_INPUT_SCHEMA = _object_schema()
|
||||
SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("path",),
|
||||
path={"type": "string"},
|
||||
)
|
||||
SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA = _object_schema(
|
||||
required=("text",),
|
||||
text={"type": "string"},
|
||||
return_url={"type": "boolean"},
|
||||
)
|
||||
SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("result",),
|
||||
result={"type": "string"},
|
||||
)
|
||||
SYSTEM_HTML_RENDER_INPUT_SCHEMA = _object_schema(
|
||||
required=("tmpl", "data"),
|
||||
tmpl={"type": "string"},
|
||||
data={"type": "object"},
|
||||
return_url={"type": "boolean"},
|
||||
options=_nullable({"type": "object"}),
|
||||
)
|
||||
SYSTEM_HTML_RENDER_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("result",),
|
||||
result={"type": "string"},
|
||||
)
|
||||
SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA = _object_schema(
|
||||
required=("session_key",),
|
||||
session_key={"type": "string"},
|
||||
)
|
||||
SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA = _object_schema()
|
||||
SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA = _object_schema(
|
||||
required=("session_key",),
|
||||
session_key={"type": "string"},
|
||||
)
|
||||
SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA = _object_schema()
|
||||
DB_GET_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
|
||||
DB_GET_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("value",),
|
||||
value=_nullable({}),
|
||||
)
|
||||
DB_SET_INPUT_SCHEMA = _object_schema(
|
||||
required=("key", "value"),
|
||||
key={"type": "string"},
|
||||
value={},
|
||||
)
|
||||
DB_SET_OUTPUT_SCHEMA = _object_schema()
|
||||
DB_DELETE_INPUT_SCHEMA = _object_schema(required=("key",), key={"type": "string"})
|
||||
DB_DELETE_OUTPUT_SCHEMA = _object_schema()
|
||||
DB_LIST_INPUT_SCHEMA = _object_schema(prefix=_nullable({"type": "string"}))
|
||||
DB_LIST_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("keys",),
|
||||
keys={"type": "array", "items": {"type": "string"}},
|
||||
)
|
||||
DB_GET_MANY_INPUT_SCHEMA = _object_schema(
|
||||
required=("keys",),
|
||||
keys={"type": "array", "items": {"type": "string"}},
|
||||
)
|
||||
DB_GET_MANY_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("items",),
|
||||
items={
|
||||
"type": "array",
|
||||
"items": _object_schema(
|
||||
required=("key", "value"),
|
||||
key={"type": "string"},
|
||||
value=_nullable({}),
|
||||
),
|
||||
},
|
||||
)
|
||||
DB_SET_MANY_INPUT_SCHEMA = _object_schema(
|
||||
required=("items",),
|
||||
items={
|
||||
"type": "array",
|
||||
"items": _object_schema(
|
||||
required=("key", "value"),
|
||||
key={"type": "string"},
|
||||
value={},
|
||||
),
|
||||
},
|
||||
)
|
||||
DB_SET_MANY_OUTPUT_SCHEMA = _object_schema()
|
||||
DB_WATCH_INPUT_SCHEMA = _object_schema(prefix=_nullable({"type": "string"}))
|
||||
DB_WATCH_OUTPUT_SCHEMA = _object_schema()
|
||||
SESSION_REF_SCHEMA = _object_schema(
|
||||
required=("conversation_id",),
|
||||
conversation_id={"type": "string"},
|
||||
platform=_nullable({"type": "string"}),
|
||||
raw=_nullable({"type": "object"}),
|
||||
)
|
||||
SYSTEM_EVENT_REACT_INPUT_SCHEMA = _object_schema(
|
||||
required=("emoji",),
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
emoji={"type": "string"},
|
||||
)
|
||||
SYSTEM_EVENT_REACT_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("supported",),
|
||||
supported={"type": "boolean"},
|
||||
)
|
||||
SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA = _object_schema(
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
)
|
||||
SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("supported",),
|
||||
supported={"type": "boolean"},
|
||||
)
|
||||
SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA = _object_schema(
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
use_fallback={"type": "boolean"},
|
||||
)
|
||||
SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("supported",),
|
||||
supported={"type": "boolean"},
|
||||
stream_id=_nullable({"type": "string"}),
|
||||
)
|
||||
SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA = _object_schema(
|
||||
required=("stream_id", "chain"),
|
||||
stream_id={"type": "string"},
|
||||
chain={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA = _object_schema()
|
||||
SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA = _object_schema(
|
||||
required=("stream_id",),
|
||||
stream_id={"type": "string"},
|
||||
)
|
||||
SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("supported",),
|
||||
supported={"type": "boolean"},
|
||||
)
|
||||
PLATFORM_SEND_INPUT_SCHEMA = _object_schema(
|
||||
required=("session", "text"),
|
||||
session={"type": "string"},
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
text={"type": "string"},
|
||||
)
|
||||
PLATFORM_SEND_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("message_id",),
|
||||
message_id={"type": "string"},
|
||||
)
|
||||
PLATFORM_SEND_IMAGE_INPUT_SCHEMA = _object_schema(
|
||||
required=("session", "image_url"),
|
||||
session={"type": "string"},
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
image_url={"type": "string"},
|
||||
)
|
||||
PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("message_id",),
|
||||
message_id={"type": "string"},
|
||||
)
|
||||
PLATFORM_SEND_CHAIN_INPUT_SCHEMA = _object_schema(
|
||||
required=("session", "chain"),
|
||||
session={"type": "string"},
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
chain={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("message_id",),
|
||||
message_id={"type": "string"},
|
||||
)
|
||||
PLATFORM_GET_MEMBERS_INPUT_SCHEMA = _object_schema(
|
||||
required=("session",),
|
||||
session={"type": "string"},
|
||||
target=_nullable(SESSION_REF_SCHEMA),
|
||||
)
|
||||
PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("members",),
|
||||
members={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
HTTP_REGISTER_API_INPUT_SCHEMA = _object_schema(
|
||||
required=("route", "methods", "handler_capability"),
|
||||
route={"type": "string"},
|
||||
methods={"type": "array", "items": {"type": "string"}},
|
||||
handler_capability={"type": "string"},
|
||||
description={"type": "string"},
|
||||
)
|
||||
HTTP_REGISTER_API_OUTPUT_SCHEMA = _object_schema()
|
||||
HTTP_UNREGISTER_API_INPUT_SCHEMA = _object_schema(
|
||||
required=("route", "methods"),
|
||||
route={"type": "string"},
|
||||
methods={"type": "array", "items": {"type": "string"}},
|
||||
)
|
||||
HTTP_UNREGISTER_API_OUTPUT_SCHEMA = _object_schema()
|
||||
HTTP_LIST_APIS_INPUT_SCHEMA = _object_schema()
|
||||
HTTP_LIST_APIS_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("apis",),
|
||||
apis={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
METADATA_GET_PLUGIN_INPUT_SCHEMA = _object_schema(
|
||||
required=("name",),
|
||||
name={"type": "string"},
|
||||
)
|
||||
METADATA_GET_PLUGIN_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("plugin",),
|
||||
plugin=_nullable({"type": "object"}),
|
||||
)
|
||||
METADATA_LIST_PLUGINS_INPUT_SCHEMA = _object_schema()
|
||||
METADATA_LIST_PLUGINS_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("plugins",),
|
||||
plugins={"type": "array", "items": {"type": "object"}},
|
||||
)
|
||||
METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA = _object_schema(
|
||||
required=("name",),
|
||||
name={"type": "string"},
|
||||
)
|
||||
METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA = _object_schema(
|
||||
required=("config",),
|
||||
config=_nullable({"type": "object"}),
|
||||
)
|
||||
|
||||
BUILTIN_CAPABILITY_SCHEMAS: dict[str, dict[str, JSONSchema]] = {
|
||||
"llm.chat": {"input": LLM_CHAT_INPUT_SCHEMA, "output": LLM_CHAT_OUTPUT_SCHEMA},
|
||||
"llm.chat_raw": {
|
||||
"input": LLM_CHAT_RAW_INPUT_SCHEMA,
|
||||
"output": LLM_CHAT_RAW_OUTPUT_SCHEMA,
|
||||
},
|
||||
"llm.stream_chat": {
|
||||
"input": LLM_STREAM_CHAT_INPUT_SCHEMA,
|
||||
"output": LLM_STREAM_CHAT_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.search": {
|
||||
"input": MEMORY_SEARCH_INPUT_SCHEMA,
|
||||
"output": MEMORY_SEARCH_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.save": {
|
||||
"input": MEMORY_SAVE_INPUT_SCHEMA,
|
||||
"output": MEMORY_SAVE_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.get": {
|
||||
"input": MEMORY_GET_INPUT_SCHEMA,
|
||||
"output": MEMORY_GET_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.delete": {
|
||||
"input": MEMORY_DELETE_INPUT_SCHEMA,
|
||||
"output": MEMORY_DELETE_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.save_with_ttl": {
|
||||
"input": MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA,
|
||||
"output": MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.get_many": {
|
||||
"input": MEMORY_GET_MANY_INPUT_SCHEMA,
|
||||
"output": MEMORY_GET_MANY_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.delete_many": {
|
||||
"input": MEMORY_DELETE_MANY_INPUT_SCHEMA,
|
||||
"output": MEMORY_DELETE_MANY_OUTPUT_SCHEMA,
|
||||
},
|
||||
"memory.stats": {
|
||||
"input": MEMORY_STATS_INPUT_SCHEMA,
|
||||
"output": MEMORY_STATS_OUTPUT_SCHEMA,
|
||||
},
|
||||
"db.get": {"input": DB_GET_INPUT_SCHEMA, "output": DB_GET_OUTPUT_SCHEMA},
|
||||
"db.set": {"input": DB_SET_INPUT_SCHEMA, "output": DB_SET_OUTPUT_SCHEMA},
|
||||
"db.delete": {"input": DB_DELETE_INPUT_SCHEMA, "output": DB_DELETE_OUTPUT_SCHEMA},
|
||||
"db.list": {"input": DB_LIST_INPUT_SCHEMA, "output": DB_LIST_OUTPUT_SCHEMA},
|
||||
"db.get_many": {
|
||||
"input": DB_GET_MANY_INPUT_SCHEMA,
|
||||
"output": DB_GET_MANY_OUTPUT_SCHEMA,
|
||||
},
|
||||
"db.set_many": {
|
||||
"input": DB_SET_MANY_INPUT_SCHEMA,
|
||||
"output": DB_SET_MANY_OUTPUT_SCHEMA,
|
||||
},
|
||||
"db.watch": {"input": DB_WATCH_INPUT_SCHEMA, "output": DB_WATCH_OUTPUT_SCHEMA},
|
||||
"platform.send": {
|
||||
"input": PLATFORM_SEND_INPUT_SCHEMA,
|
||||
"output": PLATFORM_SEND_OUTPUT_SCHEMA,
|
||||
},
|
||||
"platform.send_image": {
|
||||
"input": PLATFORM_SEND_IMAGE_INPUT_SCHEMA,
|
||||
"output": PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA,
|
||||
},
|
||||
"platform.send_chain": {
|
||||
"input": PLATFORM_SEND_CHAIN_INPUT_SCHEMA,
|
||||
"output": PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA,
|
||||
},
|
||||
"platform.get_members": {
|
||||
"input": PLATFORM_GET_MEMBERS_INPUT_SCHEMA,
|
||||
"output": PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA,
|
||||
},
|
||||
"http.register_api": {
|
||||
"input": HTTP_REGISTER_API_INPUT_SCHEMA,
|
||||
"output": HTTP_REGISTER_API_OUTPUT_SCHEMA,
|
||||
},
|
||||
"http.unregister_api": {
|
||||
"input": HTTP_UNREGISTER_API_INPUT_SCHEMA,
|
||||
"output": HTTP_UNREGISTER_API_OUTPUT_SCHEMA,
|
||||
},
|
||||
"http.list_apis": {
|
||||
"input": HTTP_LIST_APIS_INPUT_SCHEMA,
|
||||
"output": HTTP_LIST_APIS_OUTPUT_SCHEMA,
|
||||
},
|
||||
"metadata.get_plugin": {
|
||||
"input": METADATA_GET_PLUGIN_INPUT_SCHEMA,
|
||||
"output": METADATA_GET_PLUGIN_OUTPUT_SCHEMA,
|
||||
},
|
||||
"metadata.list_plugins": {
|
||||
"input": METADATA_LIST_PLUGINS_INPUT_SCHEMA,
|
||||
"output": METADATA_LIST_PLUGINS_OUTPUT_SCHEMA,
|
||||
},
|
||||
"metadata.get_plugin_config": {
|
||||
"input": METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA,
|
||||
"output": METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.get_data_dir": {
|
||||
"input": SYSTEM_GET_DATA_DIR_INPUT_SCHEMA,
|
||||
"output": SYSTEM_GET_DATA_DIR_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.text_to_image": {
|
||||
"input": SYSTEM_TEXT_TO_IMAGE_INPUT_SCHEMA,
|
||||
"output": SYSTEM_TEXT_TO_IMAGE_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.html_render": {
|
||||
"input": SYSTEM_HTML_RENDER_INPUT_SCHEMA,
|
||||
"output": SYSTEM_HTML_RENDER_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.session_waiter.register": {
|
||||
"input": SYSTEM_SESSION_WAITER_REGISTER_INPUT_SCHEMA,
|
||||
"output": SYSTEM_SESSION_WAITER_REGISTER_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.session_waiter.unregister": {
|
||||
"input": SYSTEM_SESSION_WAITER_UNREGISTER_INPUT_SCHEMA,
|
||||
"output": SYSTEM_SESSION_WAITER_UNREGISTER_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.event.react": {
|
||||
"input": SYSTEM_EVENT_REACT_INPUT_SCHEMA,
|
||||
"output": SYSTEM_EVENT_REACT_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.event.send_typing": {
|
||||
"input": SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA,
|
||||
"output": SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.event.send_streaming": {
|
||||
"input": SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA,
|
||||
"output": SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.event.send_streaming_chunk": {
|
||||
"input": SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA,
|
||||
"output": SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA,
|
||||
},
|
||||
"system.event.send_streaming_close": {
|
||||
"input": SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA,
|
||||
"output": SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BUILTIN_CAPABILITY_SCHEMAS",
|
||||
"DB_DELETE_INPUT_SCHEMA",
|
||||
"DB_DELETE_OUTPUT_SCHEMA",
|
||||
"DB_GET_INPUT_SCHEMA",
|
||||
"DB_GET_MANY_INPUT_SCHEMA",
|
||||
"DB_GET_MANY_OUTPUT_SCHEMA",
|
||||
"DB_GET_OUTPUT_SCHEMA",
|
||||
"DB_LIST_INPUT_SCHEMA",
|
||||
"DB_LIST_OUTPUT_SCHEMA",
|
||||
"DB_SET_INPUT_SCHEMA",
|
||||
"DB_SET_MANY_INPUT_SCHEMA",
|
||||
"DB_SET_MANY_OUTPUT_SCHEMA",
|
||||
"DB_SET_OUTPUT_SCHEMA",
|
||||
"DB_WATCH_INPUT_SCHEMA",
|
||||
"DB_WATCH_OUTPUT_SCHEMA",
|
||||
"HTTP_LIST_APIS_INPUT_SCHEMA",
|
||||
"HTTP_LIST_APIS_OUTPUT_SCHEMA",
|
||||
"HTTP_REGISTER_API_INPUT_SCHEMA",
|
||||
"HTTP_REGISTER_API_OUTPUT_SCHEMA",
|
||||
"HTTP_UNREGISTER_API_INPUT_SCHEMA",
|
||||
"HTTP_UNREGISTER_API_OUTPUT_SCHEMA",
|
||||
"JSONSchema",
|
||||
"LLM_CHAT_INPUT_SCHEMA",
|
||||
"LLM_CHAT_OUTPUT_SCHEMA",
|
||||
"LLM_CHAT_RAW_INPUT_SCHEMA",
|
||||
"LLM_CHAT_RAW_OUTPUT_SCHEMA",
|
||||
"LLM_STREAM_CHAT_INPUT_SCHEMA",
|
||||
"LLM_STREAM_CHAT_OUTPUT_SCHEMA",
|
||||
"MEMORY_DELETE_INPUT_SCHEMA",
|
||||
"MEMORY_DELETE_MANY_INPUT_SCHEMA",
|
||||
"MEMORY_DELETE_MANY_OUTPUT_SCHEMA",
|
||||
"MEMORY_DELETE_OUTPUT_SCHEMA",
|
||||
"MEMORY_GET_INPUT_SCHEMA",
|
||||
"MEMORY_GET_MANY_INPUT_SCHEMA",
|
||||
"MEMORY_GET_MANY_OUTPUT_SCHEMA",
|
||||
"MEMORY_GET_OUTPUT_SCHEMA",
|
||||
"MEMORY_SAVE_INPUT_SCHEMA",
|
||||
"MEMORY_SAVE_OUTPUT_SCHEMA",
|
||||
"MEMORY_SAVE_WITH_TTL_INPUT_SCHEMA",
|
||||
"MEMORY_SAVE_WITH_TTL_OUTPUT_SCHEMA",
|
||||
"MEMORY_SEARCH_INPUT_SCHEMA",
|
||||
"MEMORY_SEARCH_OUTPUT_SCHEMA",
|
||||
"MEMORY_STATS_INPUT_SCHEMA",
|
||||
"MEMORY_STATS_OUTPUT_SCHEMA",
|
||||
"METADATA_GET_PLUGIN_CONFIG_INPUT_SCHEMA",
|
||||
"METADATA_GET_PLUGIN_CONFIG_OUTPUT_SCHEMA",
|
||||
"METADATA_GET_PLUGIN_INPUT_SCHEMA",
|
||||
"METADATA_GET_PLUGIN_OUTPUT_SCHEMA",
|
||||
"METADATA_LIST_PLUGINS_INPUT_SCHEMA",
|
||||
"METADATA_LIST_PLUGINS_OUTPUT_SCHEMA",
|
||||
"PLATFORM_GET_MEMBERS_INPUT_SCHEMA",
|
||||
"PLATFORM_GET_MEMBERS_OUTPUT_SCHEMA",
|
||||
"PLATFORM_SEND_CHAIN_INPUT_SCHEMA",
|
||||
"PLATFORM_SEND_CHAIN_OUTPUT_SCHEMA",
|
||||
"PLATFORM_SEND_IMAGE_INPUT_SCHEMA",
|
||||
"PLATFORM_SEND_IMAGE_OUTPUT_SCHEMA",
|
||||
"PLATFORM_SEND_INPUT_SCHEMA",
|
||||
"PLATFORM_SEND_OUTPUT_SCHEMA",
|
||||
"SESSION_REF_SCHEMA",
|
||||
"SYSTEM_EVENT_REACT_INPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_REACT_OUTPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_CHUNK_INPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_CHUNK_OUTPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_CLOSE_INPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_CLOSE_OUTPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_INPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_STREAMING_OUTPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_TYPING_INPUT_SCHEMA",
|
||||
"SYSTEM_EVENT_SEND_TYPING_OUTPUT_SCHEMA",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -152,7 +152,7 @@ class ResultMessage(_MessageBase):
|
||||
error: ErrorPayload | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_result_state(self) -> "ResultMessage":
|
||||
def validate_result_state(self) -> ResultMessage:
|
||||
"""约束 success / output / error 的组合状态。"""
|
||||
if self.success:
|
||||
if self.error is not None:
|
||||
@@ -238,7 +238,7 @@ class EventMessage(_MessageBase):
|
||||
error: ErrorPayload | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_phase_constraints(self) -> "EventMessage":
|
||||
def validate_phase_constraints(self) -> EventMessage:
|
||||
"""验证各 phase 的字段约束。
|
||||
|
||||
- started: 所有字段必须为空
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Literal
|
||||
|
||||
import msgpack
|
||||
|
||||
from .messages import ProtocolMessage, parse_message
|
||||
|
||||
StdioFraming = Literal["line", "length_prefixed"]
|
||||
WebSocketFrameType = Literal["text", "binary"]
|
||||
WireCodecName = Literal["json", "msgpack"]
|
||||
|
||||
|
||||
class ProtocolCodec(ABC):
|
||||
name: WireCodecName
|
||||
stdio_framing: StdioFraming
|
||||
websocket_frame_type: WebSocketFrameType
|
||||
|
||||
@abstractmethod
|
||||
def encode_message(self, message: ProtocolMessage) -> bytes | str:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class JsonProtocolCodec(ProtocolCodec):
|
||||
name: WireCodecName = "json"
|
||||
stdio_framing: StdioFraming = "line"
|
||||
websocket_frame_type: WebSocketFrameType = "text"
|
||||
|
||||
def encode_message(self, message: ProtocolMessage) -> str:
|
||||
return message.model_dump_json(exclude_none=True)
|
||||
|
||||
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
|
||||
if isinstance(payload, bytes):
|
||||
return parse_message(payload.decode("utf-8"))
|
||||
return parse_message(payload)
|
||||
|
||||
|
||||
class MsgpackProtocolCodec(ProtocolCodec):
|
||||
name: WireCodecName = "msgpack"
|
||||
stdio_framing: StdioFraming = "length_prefixed"
|
||||
websocket_frame_type: WebSocketFrameType = "binary"
|
||||
|
||||
def encode_message(self, message: ProtocolMessage) -> bytes:
|
||||
return msgpack.packb(
|
||||
message.model_dump(exclude_none=True),
|
||||
use_bin_type=True,
|
||||
)
|
||||
|
||||
def decode_message(self, payload: bytes | str) -> ProtocolMessage:
|
||||
if isinstance(payload, str):
|
||||
return parse_message(payload)
|
||||
return parse_message(msgpack.unpackb(payload, raw=False))
|
||||
|
||||
|
||||
def make_protocol_codec(name: WireCodecName | str) -> ProtocolCodec:
|
||||
if name == "json":
|
||||
return JsonProtocolCodec()
|
||||
if name == "msgpack":
|
||||
return MsgpackProtocolCodec()
|
||||
raise ValueError(f"未知 wire codec: {name}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"JsonProtocolCodec",
|
||||
"MsgpackProtocolCodec",
|
||||
"ProtocolCodec",
|
||||
"StdioFraming",
|
||||
"WebSocketFrameType",
|
||||
"WireCodecName",
|
||||
"make_protocol_codec",
|
||||
]
|
||||
@@ -1,21 +1,32 @@
|
||||
"""AstrBot SDK 的高级运行时原语。
|
||||
"""AstrBot SDK runtime public exports.
|
||||
|
||||
这里仅暴露相对稳定的运行时构件:协议 `Peer`、传输抽象以及能力/处理器分发器。
|
||||
大多数插件作者应优先使用顶层 `astrbot_sdk`。
|
||||
本模块提供运行时核心组件的公共导出,包括:
|
||||
- CapabilityRouter: 能力路由器,处理能力调用的分发和路由
|
||||
- HandlerDispatcher: 事件处理器分发器,将事件分发到注册的 handler
|
||||
- Peer: 与 AstrBot 核心通信的对等端抽象
|
||||
- Transport 系列: 进程间通信传输层实现(stdio/websocket)
|
||||
|
||||
`loader` / `bootstrap` 等编排细节保留在各自子模块中,不作为根级稳定契约。
|
||||
延迟加载策略:
|
||||
为避免导入时触发 websocket/aiohttp 等重型依赖,采用 __getattr__ 实现按需加载。
|
||||
这样轻量级导入(如仅使用类型提示)不会产生不必要的依赖开销。
|
||||
"""
|
||||
|
||||
from .capability_router import CapabilityRouter, StreamExecution
|
||||
from .handler_dispatcher import HandlerDispatcher
|
||||
from .peer import Peer
|
||||
from .transport import (
|
||||
MessageHandler,
|
||||
StdioTransport,
|
||||
Transport,
|
||||
WebSocketClientTransport,
|
||||
WebSocketServerTransport,
|
||||
)
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .capability_router import CapabilityRouter, StreamExecution
|
||||
from .handler_dispatcher import HandlerDispatcher
|
||||
from .peer import Peer
|
||||
from .transport import (
|
||||
MessageHandler,
|
||||
StdioTransport,
|
||||
Transport,
|
||||
WebSocketClientTransport,
|
||||
WebSocketServerTransport,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CapabilityRouter",
|
||||
@@ -28,3 +39,25 @@ __all__ = [
|
||||
"WebSocketClientTransport",
|
||||
"WebSocketServerTransport",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name in {"CapabilityRouter", "StreamExecution"}:
|
||||
module = import_module(".capability_router", __name__)
|
||||
return getattr(module, name)
|
||||
if name == "HandlerDispatcher":
|
||||
module = import_module(".handler_dispatcher", __name__)
|
||||
return getattr(module, name)
|
||||
if name == "Peer":
|
||||
module = import_module(".peer", __name__)
|
||||
return getattr(module, name)
|
||||
if name in {
|
||||
"MessageHandler",
|
||||
"StdioTransport",
|
||||
"Transport",
|
||||
"WebSocketClientTransport",
|
||||
"WebSocketServerTransport",
|
||||
}:
|
||||
module = import_module(".transport", __name__)
|
||||
return getattr(module, name)
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -0,0 +1,834 @@
|
||||
"""Built-in capability registration and handlers for CapabilityRouter.
|
||||
|
||||
本模块为 CapabilityRouter 提供内置能力的注册逻辑和处理函数实现。
|
||||
内置能力涵盖以下类别:
|
||||
- LLM: 对话、流式对话等大语言模型能力
|
||||
- Memory: 记忆存储、搜索、带 TTL 的键值对
|
||||
- DB: 持久化键值存储及变更监听
|
||||
- Platform: 跨平台消息发送、图片、消息链
|
||||
- HTTP: 动态 API 路由注册与管理
|
||||
- Metadata: 插件元数据查询
|
||||
- System: 数据目录、文本转图片、HTML 渲染、会话等待器等
|
||||
|
||||
设计模式:
|
||||
通过 Mixin 类 (BuiltinCapabilityRouterMixin) 将内置能力注入到 CapabilityRouter,
|
||||
使其与用户自定义能力共享相同的注册和调用机制。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ..errors import AstrBotError
|
||||
from ..protocol.descriptors import (
|
||||
BUILTIN_CAPABILITY_SCHEMAS,
|
||||
CapabilityDescriptor,
|
||||
SessionRef,
|
||||
)
|
||||
from ._streaming import StreamExecution
|
||||
|
||||
|
||||
def _clone_target_payload(value: Any) -> dict[str, Any] | None:
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
return {str(key): item for key, item in value.items()}
|
||||
|
||||
|
||||
def _clone_chain_payload(value: Any) -> list[dict[str, Any]]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [
|
||||
{str(key): item for key, item in chunk.items()}
|
||||
for chunk in value
|
||||
if isinstance(chunk, dict)
|
||||
]
|
||||
|
||||
|
||||
class _CapabilityRouterHost:
|
||||
memory_store: dict[str, dict[str, Any]]
|
||||
db_store: dict[str, Any]
|
||||
sent_messages: list[dict[str, Any]]
|
||||
event_actions: list[dict[str, Any]]
|
||||
http_api_store: list[dict[str, Any]]
|
||||
_event_streams: dict[str, dict[str, Any]]
|
||||
_plugins: dict[str, Any]
|
||||
_system_data_root: Path
|
||||
_session_waiters: dict[str, set[str]]
|
||||
_db_watch_subscriptions: dict[str, tuple[str | None, asyncio.Queue[dict[str, Any]]]]
|
||||
|
||||
def register(
|
||||
self,
|
||||
descriptor: CapabilityDescriptor,
|
||||
*,
|
||||
call_handler=None,
|
||||
stream_handler=None,
|
||||
finalize=None,
|
||||
exposed: bool = True,
|
||||
) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def _emit_db_change(self, *, op: str, key: str, value: Any | None) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def _require_caller_plugin_id(capability_name: str) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class BuiltinCapabilityRouterMixin(_CapabilityRouterHost):
|
||||
def _register_builtin_capabilities(self) -> None:
|
||||
self._register_llm_capabilities()
|
||||
self._register_memory_capabilities()
|
||||
self._register_db_capabilities()
|
||||
self._register_platform_capabilities()
|
||||
self._register_http_capabilities()
|
||||
self._register_metadata_capabilities()
|
||||
self._register_system_capabilities()
|
||||
|
||||
def _builtin_descriptor(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
*,
|
||||
supports_stream: bool = False,
|
||||
cancelable: bool = False,
|
||||
) -> CapabilityDescriptor:
|
||||
schema = BUILTIN_CAPABILITY_SCHEMAS[name]
|
||||
return CapabilityDescriptor(
|
||||
name=name,
|
||||
description=description,
|
||||
input_schema=copy.deepcopy(schema["input"]),
|
||||
output_schema=copy.deepcopy(schema["output"]),
|
||||
supports_stream=supports_stream,
|
||||
cancelable=cancelable,
|
||||
)
|
||||
|
||||
def _resolve_target(
|
||||
self, payload: dict[str, Any]
|
||||
) -> tuple[str, dict[str, Any] | None]:
|
||||
target_payload = payload.get("target")
|
||||
if isinstance(target_payload, dict):
|
||||
target = SessionRef.model_validate(target_payload)
|
||||
return target.session, target.to_payload()
|
||||
return str(payload.get("session", "")), None
|
||||
|
||||
async def _llm_chat(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prompt = str(payload.get("prompt", ""))
|
||||
return {"text": f"Echo: {prompt}"}
|
||||
|
||||
async def _llm_chat_raw(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prompt = str(payload.get("prompt", ""))
|
||||
text = f"Echo: {prompt}"
|
||||
return {
|
||||
"text": text,
|
||||
"usage": {
|
||||
"input_tokens": len(prompt),
|
||||
"output_tokens": len(text),
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
"tool_calls": [],
|
||||
}
|
||||
|
||||
async def _llm_stream(
|
||||
self,
|
||||
_request_id: str,
|
||||
payload: dict[str, Any],
|
||||
token,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
text = f"Echo: {str(payload.get('prompt', ''))}"
|
||||
for char in text:
|
||||
token.raise_if_cancelled()
|
||||
await asyncio.sleep(0)
|
||||
yield {"text": char}
|
||||
|
||||
def _register_llm_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("llm.chat", "发送对话请求,返回文本"),
|
||||
call_handler=self._llm_chat,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("llm.chat_raw", "发送对话请求,返回完整响应"),
|
||||
call_handler=self._llm_chat_raw,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"llm.stream_chat",
|
||||
"流式对话",
|
||||
supports_stream=True,
|
||||
cancelable=True,
|
||||
),
|
||||
stream_handler=self._llm_stream,
|
||||
finalize=lambda chunks: {
|
||||
"text": "".join(item.get("text", "") for item in chunks)
|
||||
},
|
||||
)
|
||||
|
||||
async def _memory_search(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
query = str(payload.get("query", ""))
|
||||
items = [
|
||||
{"key": key, "value": value}
|
||||
for key, value in self.memory_store.items()
|
||||
if query in key or query in json.dumps(value, ensure_ascii=False)
|
||||
]
|
||||
return {"items": items}
|
||||
|
||||
async def _memory_save(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
if not isinstance(value, dict):
|
||||
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
|
||||
self.memory_store[key] = value
|
||||
return {}
|
||||
|
||||
async def _memory_get(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
return {"value": self.memory_store.get(str(payload.get("key", "")))}
|
||||
|
||||
async def _memory_delete(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
self.memory_store.pop(str(payload.get("key", "")), None)
|
||||
return {}
|
||||
|
||||
async def _memory_save_with_ttl(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
ttl_seconds = payload.get("ttl_seconds", 0)
|
||||
if not isinstance(value, dict):
|
||||
raise AstrBotError.invalid_input(
|
||||
"memory.save_with_ttl 的 value 必须是 object"
|
||||
)
|
||||
self.memory_store[key] = {"value": value, "ttl_seconds": ttl_seconds}
|
||||
return {}
|
||||
|
||||
async def _memory_get_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("memory.get_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
items = []
|
||||
for key in keys:
|
||||
stored = self.memory_store.get(key)
|
||||
if (
|
||||
isinstance(stored, dict)
|
||||
and "value" in stored
|
||||
and "ttl_seconds" in stored
|
||||
):
|
||||
value = stored["value"]
|
||||
else:
|
||||
value = stored
|
||||
items.append({"key": key, "value": value})
|
||||
return {"items": items}
|
||||
|
||||
async def _memory_delete_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("memory.delete_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
deleted_count = 0
|
||||
for key in keys:
|
||||
if key in self.memory_store:
|
||||
del self.memory_store[key]
|
||||
deleted_count += 1
|
||||
return {"deleted_count": deleted_count}
|
||||
|
||||
async def _memory_stats(
|
||||
self, _request_id: str, _payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
total_items = len(self.memory_store)
|
||||
total_bytes = sum(
|
||||
len(str(key)) + len(str(value)) for key, value in self.memory_store.items()
|
||||
)
|
||||
ttl_entries = sum(
|
||||
1
|
||||
for value in self.memory_store.values()
|
||||
if isinstance(value, dict) and "value" in value and "ttl_seconds" in value
|
||||
)
|
||||
return {
|
||||
"total_items": total_items,
|
||||
"total_bytes": total_bytes,
|
||||
"plugin_id": self._require_caller_plugin_id("memory.stats"),
|
||||
"ttl_entries": ttl_entries,
|
||||
}
|
||||
|
||||
def _register_memory_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.search", "搜索记忆"),
|
||||
call_handler=self._memory_search,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.save", "保存记忆"),
|
||||
call_handler=self._memory_save,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.get", "读取单条记忆"),
|
||||
call_handler=self._memory_get,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.delete", "删除记忆"),
|
||||
call_handler=self._memory_delete,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.save_with_ttl", "保存带过期时间的记忆"),
|
||||
call_handler=self._memory_save_with_ttl,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.get_many", "批量获取记忆"),
|
||||
call_handler=self._memory_get_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.delete_many", "批量删除记忆"),
|
||||
call_handler=self._memory_delete_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.stats", "获取记忆统计信息"),
|
||||
call_handler=self._memory_stats,
|
||||
)
|
||||
|
||||
async def _db_get(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
return {"value": self.db_store.get(str(payload.get("key", "")))}
|
||||
|
||||
async def _db_set(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
self.db_store[key] = value
|
||||
self._emit_db_change(op="set", key=key, value=value)
|
||||
return {}
|
||||
|
||||
async def _db_delete(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
self.db_store.pop(key, None)
|
||||
self._emit_db_change(op="delete", key=key, value=None)
|
||||
return {}
|
||||
|
||||
async def _db_list(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prefix = payload.get("prefix")
|
||||
keys = sorted(self.db_store.keys())
|
||||
if isinstance(prefix, str):
|
||||
keys = [item for item in keys if item.startswith(prefix)]
|
||||
return {"keys": keys}
|
||||
|
||||
async def _db_get_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
|
||||
return {"items": items}
|
||||
|
||||
async def _db_set_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
items_payload = payload.get("items")
|
||||
if not isinstance(items_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
|
||||
for entry in items_payload:
|
||||
if not isinstance(entry, dict):
|
||||
raise AstrBotError.invalid_input(
|
||||
"db.set_many 的 items 必须是 object 数组"
|
||||
)
|
||||
key = str(entry.get("key", ""))
|
||||
value = entry.get("value")
|
||||
self.db_store[key] = value
|
||||
self._emit_db_change(op="set", key=key, value=value)
|
||||
return {}
|
||||
|
||||
async def _db_watch(
|
||||
self, request_id: str, payload: dict[str, Any], _token
|
||||
) -> StreamExecution:
|
||||
prefix = payload.get("prefix")
|
||||
prefix_value: str | None
|
||||
if isinstance(prefix, str):
|
||||
prefix_value = prefix
|
||||
elif prefix is None:
|
||||
prefix_value = None
|
||||
else:
|
||||
raise AstrBotError.invalid_input("db.watch 的 prefix 必须是 string 或 null")
|
||||
|
||||
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
|
||||
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
|
||||
|
||||
async def iterator() -> AsyncIterator[dict[str, Any]]:
|
||||
try:
|
||||
while True:
|
||||
yield await queue.get()
|
||||
finally:
|
||||
self._db_watch_subscriptions.pop(request_id, None)
|
||||
|
||||
return StreamExecution(
|
||||
iterator=iterator(),
|
||||
finalize=lambda _chunks: {},
|
||||
collect_chunks=False,
|
||||
)
|
||||
|
||||
def _register_db_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("db.get", "读取 KV"), call_handler=self._db_get
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.set", "写入 KV"), call_handler=self._db_set
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.delete", "删除 KV"),
|
||||
call_handler=self._db_delete,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.list", "列出 KV"), call_handler=self._db_list
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.get_many", "批量读取 KV"),
|
||||
call_handler=self._db_get_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.set_many", "批量写入 KV"),
|
||||
call_handler=self._db_set_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"db.watch",
|
||||
"订阅 KV 变更",
|
||||
supports_stream=True,
|
||||
cancelable=True,
|
||||
),
|
||||
stream_handler=self._db_watch,
|
||||
)
|
||||
|
||||
async def _platform_send(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
text = str(payload.get("text", ""))
|
||||
message_id = f"msg_{len(self.sent_messages) + 1}"
|
||||
sent: dict[str, Any] = {
|
||||
"message_id": message_id,
|
||||
"session": session,
|
||||
"text": text,
|
||||
}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_send_image(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
image_url = str(payload.get("image_url", ""))
|
||||
message_id = f"img_{len(self.sent_messages) + 1}"
|
||||
sent: dict[str, Any] = {
|
||||
"message_id": message_id,
|
||||
"session": session,
|
||||
"image_url": image_url,
|
||||
}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_send_chain(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
chain = payload.get("chain")
|
||||
if not isinstance(chain, list) or not all(
|
||||
isinstance(item, dict) for item in chain
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"platform.send_chain 的 chain 必须是 object 数组"
|
||||
)
|
||||
message_id = f"chain_{len(self.sent_messages) + 1}"
|
||||
sent: dict[str, Any] = {
|
||||
"message_id": message_id,
|
||||
"session": session,
|
||||
"chain": [dict(item) for item in chain],
|
||||
}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_get_members(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, _target = self._resolve_target(payload)
|
||||
return {
|
||||
"members": [
|
||||
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
|
||||
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
|
||||
]
|
||||
}
|
||||
|
||||
def _register_platform_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send", "发送消息"),
|
||||
call_handler=self._platform_send,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send_image", "发送图片"),
|
||||
call_handler=self._platform_send_image,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send_chain", "发送消息链"),
|
||||
call_handler=self._platform_send_chain,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.get_members", "获取群成员"),
|
||||
call_handler=self._platform_get_members,
|
||||
)
|
||||
|
||||
async def _http_register_api(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
methods_payload = payload.get("methods")
|
||||
if not isinstance(methods_payload, list) or not all(
|
||||
isinstance(item, str) for item in methods_payload
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.register_api 的 methods 必须是 string 数组"
|
||||
)
|
||||
route = str(payload.get("route", "")).strip()
|
||||
handler_capability = str(payload.get("handler_capability", "")).strip()
|
||||
if not route or not handler_capability:
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.register_api 需要 route 和 handler_capability"
|
||||
)
|
||||
plugin_name = self._require_caller_plugin_id("http.register_api")
|
||||
methods = sorted({method.upper() for method in methods_payload if method})
|
||||
entry: dict[str, Any] = {
|
||||
"route": route,
|
||||
"methods": methods,
|
||||
"handler_capability": handler_capability,
|
||||
"description": str(payload.get("description", "")),
|
||||
"plugin_id": plugin_name,
|
||||
}
|
||||
self.http_api_store = [
|
||||
item
|
||||
for item in self.http_api_store
|
||||
if not (
|
||||
item.get("route") == route
|
||||
and item.get("plugin_id") == entry["plugin_id"]
|
||||
and item.get("methods") == methods
|
||||
)
|
||||
]
|
||||
self.http_api_store.append(entry)
|
||||
return {}
|
||||
|
||||
async def _http_unregister_api(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
route = str(payload.get("route", "")).strip()
|
||||
methods_payload = payload.get("methods")
|
||||
if not isinstance(methods_payload, list) or not all(
|
||||
isinstance(item, str) for item in methods_payload
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.unregister_api 的 methods 必须是 string 数组"
|
||||
)
|
||||
plugin_name = self._require_caller_plugin_id("http.unregister_api")
|
||||
methods = {method.upper() for method in methods_payload if method}
|
||||
updated: list[dict[str, Any]] = []
|
||||
for entry in self.http_api_store:
|
||||
if entry.get("route") != route:
|
||||
updated.append(entry)
|
||||
continue
|
||||
if entry.get("plugin_id") != plugin_name:
|
||||
updated.append(entry)
|
||||
continue
|
||||
if not methods:
|
||||
continue
|
||||
remaining_methods = [
|
||||
method for method in entry.get("methods", []) if method not in methods
|
||||
]
|
||||
if remaining_methods:
|
||||
updated.append({**entry, "methods": remaining_methods})
|
||||
self.http_api_store = updated
|
||||
return {}
|
||||
|
||||
async def _http_list_apis(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugin_name = self._require_caller_plugin_id("http.list_apis")
|
||||
apis = [
|
||||
dict(entry)
|
||||
for entry in self.http_api_store
|
||||
if entry.get("plugin_id") == plugin_name
|
||||
]
|
||||
return {"apis": apis}
|
||||
|
||||
def _register_http_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("http.register_api", "注册 HTTP 路由"),
|
||||
call_handler=self._http_register_api,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("http.unregister_api", "注销 HTTP 路由"),
|
||||
call_handler=self._http_unregister_api,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("http.list_apis", "列出 HTTP 路由"),
|
||||
call_handler=self._http_list_apis,
|
||||
)
|
||||
|
||||
async def _metadata_get_plugin(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
name = str(payload.get("name", "")).strip()
|
||||
plugin = self._plugins.get(name)
|
||||
if plugin is None:
|
||||
return {"plugin": None}
|
||||
return {"plugin": dict(plugin.metadata)}
|
||||
|
||||
async def _metadata_list_plugins(
|
||||
self, _request_id: str, _payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugins = [
|
||||
dict(self._plugins[name].metadata) for name in sorted(self._plugins.keys())
|
||||
]
|
||||
return {"plugins": plugins}
|
||||
|
||||
async def _metadata_get_plugin_config(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
name = str(payload.get("name", "")).strip()
|
||||
caller_plugin_id = self._require_caller_plugin_id("metadata.get_plugin_config")
|
||||
if name != caller_plugin_id:
|
||||
return {"config": None}
|
||||
plugin = self._plugins.get(name)
|
||||
if plugin is None:
|
||||
return {"config": None}
|
||||
return {"config": dict(plugin.config)}
|
||||
|
||||
def _register_metadata_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("metadata.get_plugin", "获取单个插件元数据"),
|
||||
call_handler=self._metadata_get_plugin,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("metadata.list_plugins", "列出插件元数据"),
|
||||
call_handler=self._metadata_list_plugins,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"metadata.get_plugin_config",
|
||||
"获取插件配置",
|
||||
),
|
||||
call_handler=self._metadata_get_plugin_config,
|
||||
)
|
||||
|
||||
def _register_system_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("system.get_data_dir", "获取插件数据目录"),
|
||||
call_handler=self._system_get_data_dir,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("system.text_to_image", "文本转图片"),
|
||||
call_handler=self._system_text_to_image,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("system.html_render", "渲染 HTML 模板"),
|
||||
call_handler=self._system_html_render,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"system.session_waiter.register",
|
||||
"注册会话等待器",
|
||||
),
|
||||
call_handler=self._system_session_waiter_register,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"system.session_waiter.unregister",
|
||||
"注销会话等待器",
|
||||
),
|
||||
call_handler=self._system_session_waiter_unregister,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("system.event.react", "发送事件表情回应"),
|
||||
call_handler=self._system_event_react,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("system.event.send_typing", "发送输入中状态"),
|
||||
call_handler=self._system_event_send_typing,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"system.event.send_streaming",
|
||||
"发送事件流式消息",
|
||||
),
|
||||
call_handler=self._system_event_send_streaming,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"system.event.send_streaming_chunk",
|
||||
"推送事件流式消息分片",
|
||||
),
|
||||
call_handler=self._system_event_send_streaming_chunk,
|
||||
exposed=False,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"system.event.send_streaming_close",
|
||||
"关闭事件流式消息会话",
|
||||
),
|
||||
call_handler=self._system_event_send_streaming_close,
|
||||
exposed=False,
|
||||
)
|
||||
|
||||
async def _system_get_data_dir(
|
||||
self, _request_id: str, _payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugin_id = self._require_caller_plugin_id("system.get_data_dir")
|
||||
data_dir = self._system_data_root / plugin_id
|
||||
data_dir.mkdir(parents=True, exist_ok=True)
|
||||
return {"path": str(data_dir)}
|
||||
|
||||
async def _system_text_to_image(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
text = str(payload.get("text", ""))
|
||||
if bool(payload.get("return_url", True)):
|
||||
return {"result": f"mock://text_to_image/{text}"}
|
||||
return {"result": f"<image>{text}</image>"}
|
||||
|
||||
async def _system_html_render(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
tmpl = str(payload.get("tmpl", ""))
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
raise AstrBotError.invalid_input("system.html_render requires object data")
|
||||
if bool(payload.get("return_url", True)):
|
||||
return {"result": f"mock://html_render/{tmpl}"}
|
||||
return {"result": json.dumps({"tmpl": tmpl, "data": data}, ensure_ascii=False)}
|
||||
|
||||
async def _system_session_waiter_register(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugin_id = self._require_caller_plugin_id("system.session_waiter.register")
|
||||
session_key = str(payload.get("session_key", "")).strip()
|
||||
if not session_key:
|
||||
raise AstrBotError.invalid_input(
|
||||
"system.session_waiter.register requires session_key"
|
||||
)
|
||||
self._session_waiters.setdefault(plugin_id, set()).add(session_key)
|
||||
return {}
|
||||
|
||||
async def _system_session_waiter_unregister(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugin_id = self._require_caller_plugin_id("system.session_waiter.unregister")
|
||||
session_key = str(payload.get("session_key", "")).strip()
|
||||
plugin_waiters = self._session_waiters.get(plugin_id)
|
||||
if plugin_waiters is None:
|
||||
return {}
|
||||
plugin_waiters.discard(session_key)
|
||||
if not plugin_waiters:
|
||||
self._session_waiters.pop(plugin_id, None)
|
||||
return {}
|
||||
|
||||
async def _system_event_react(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
self.event_actions.append(
|
||||
{
|
||||
"action": "react",
|
||||
"emoji": str(payload.get("emoji", "")),
|
||||
"target": _clone_target_payload(payload.get("target")),
|
||||
}
|
||||
)
|
||||
return {"supported": True}
|
||||
|
||||
async def _system_event_send_typing(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
self.event_actions.append(
|
||||
{
|
||||
"action": "send_typing",
|
||||
"target": _clone_target_payload(payload.get("target")),
|
||||
}
|
||||
)
|
||||
return {"supported": True}
|
||||
|
||||
async def _system_event_send_streaming(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
stream_id = f"mock-stream-{len(self._event_streams) + 1}"
|
||||
stream_state: dict[str, Any] = {
|
||||
"target": _clone_target_payload(payload.get("target")),
|
||||
"chunks": [],
|
||||
"use_fallback": bool(payload.get("use_fallback", False)),
|
||||
}
|
||||
self._event_streams[stream_id] = stream_state
|
||||
return {"supported": True, "stream_id": stream_id}
|
||||
|
||||
async def _system_event_send_streaming_chunk(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
stream = self._event_streams.get(str(payload.get("stream_id", "")))
|
||||
if stream is None:
|
||||
raise AstrBotError.invalid_input("Unknown sdk event streaming session")
|
||||
chain = payload.get("chain")
|
||||
if not isinstance(chain, list):
|
||||
raise AstrBotError.invalid_input(
|
||||
"system.event.send_streaming_chunk requires a chain array"
|
||||
)
|
||||
stream["chunks"].append({"chain": _clone_chain_payload(chain)})
|
||||
return {}
|
||||
|
||||
async def _system_event_send_streaming_close(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
stream_id = str(payload.get("stream_id", ""))
|
||||
stream = self._event_streams.pop(stream_id, None)
|
||||
if stream is None:
|
||||
raise AstrBotError.invalid_input("Unknown sdk event streaming session")
|
||||
self.event_actions.append(
|
||||
{
|
||||
"action": "send_streaming",
|
||||
"target": stream["target"],
|
||||
"chunks": list(stream["chunks"]),
|
||||
"use_fallback": bool(stream["use_fallback"]),
|
||||
}
|
||||
)
|
||||
return {"supported": True}
|
||||
|
||||
|
||||
__all__ = ["BuiltinCapabilityRouterMixin"]
|
||||
@@ -0,0 +1,173 @@
|
||||
"""Support helpers for runtime loader reflection and signature validation.
|
||||
|
||||
本模块提供运行时加载器所需的反射和签名验证工具函数,主要用于:
|
||||
1. 解析 handler/capability 函数签名,提取参数类型信息
|
||||
2. 识别需要注入的框架对象(如 Context、MessageEvent、ScheduleContext)
|
||||
3. 构建参数规格 (ParamSpec) 供协议层使用
|
||||
4. 验证 schedule handler 的签名合法性
|
||||
|
||||
关键函数:
|
||||
- build_param_specs: 从 handler 签名构建参数规格列表
|
||||
- is_injected_parameter: 判断参数是否应由框架注入而非从命令行解析
|
||||
- validate_schedule_signature: 确保 schedule handler 只接受允许的注入参数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import typing
|
||||
from typing import Any, Literal, TypeAlias
|
||||
|
||||
from ..decorators import get_capability_meta, get_handler_meta
|
||||
from ..protocol.descriptors import ParamSpec
|
||||
from ..schedule import ScheduleContext
|
||||
from ..types import GreedyStr
|
||||
|
||||
ParamTypeName: TypeAlias = Literal[
|
||||
"str", "int", "float", "bool", "optional", "greedy_str"
|
||||
]
|
||||
OptionalInnerType: TypeAlias = Literal["str", "int", "float", "bool"] | None
|
||||
|
||||
|
||||
def unwrap_optional(annotation: Any) -> tuple[Any, bool]:
|
||||
origin = typing.get_origin(annotation)
|
||||
if origin is typing.Union:
|
||||
args = [item for item in typing.get_args(annotation) if item is not type(None)]
|
||||
if len(args) == 1:
|
||||
return args[0], True
|
||||
return annotation, False
|
||||
|
||||
|
||||
def is_injected_parameter(annotation: Any, parameter_name: str) -> bool:
|
||||
if parameter_name in {"event", "ctx", "context", "sched", "schedule"}:
|
||||
return True
|
||||
normalized, _is_optional = unwrap_optional(annotation)
|
||||
if normalized is None:
|
||||
return False
|
||||
if normalized in {ScheduleContext}:
|
||||
return True
|
||||
if isinstance(normalized, type):
|
||||
from ..context import Context
|
||||
from ..events import MessageEvent
|
||||
|
||||
return issubclass(normalized, (Context, MessageEvent, ScheduleContext))
|
||||
return False
|
||||
|
||||
|
||||
def param_type_name(annotation: Any) -> tuple[ParamTypeName, OptionalInnerType, bool]:
|
||||
normalized, is_optional = unwrap_optional(annotation)
|
||||
if normalized is GreedyStr:
|
||||
return "greedy_str", None, False
|
||||
if normalized in {int, float, bool, str}:
|
||||
if is_optional:
|
||||
return "optional", normalized.__name__, False
|
||||
return normalized.__name__, None, True
|
||||
if is_optional:
|
||||
return "optional", "str", False
|
||||
return "str", None, True
|
||||
|
||||
|
||||
def build_param_specs(handler: Any) -> list[ParamSpec]:
|
||||
try:
|
||||
signature = inspect.signature(handler)
|
||||
except (TypeError, ValueError):
|
||||
return []
|
||||
try:
|
||||
type_hints = typing.get_type_hints(handler)
|
||||
except Exception:
|
||||
type_hints = {}
|
||||
|
||||
specs: list[ParamSpec] = []
|
||||
for parameter in signature.parameters.values():
|
||||
if parameter.kind not in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
):
|
||||
continue
|
||||
annotation = type_hints.get(parameter.name)
|
||||
if is_injected_parameter(annotation, parameter.name):
|
||||
continue
|
||||
param_type, inner_type, required = param_type_name(annotation)
|
||||
if parameter.default is not inspect.Parameter.empty:
|
||||
required = False
|
||||
specs.append(
|
||||
ParamSpec(
|
||||
name=parameter.name,
|
||||
type=param_type,
|
||||
required=required,
|
||||
inner_type=inner_type,
|
||||
)
|
||||
)
|
||||
|
||||
greedy_indexes = [
|
||||
index for index, spec in enumerate(specs) if spec.type == "greedy_str"
|
||||
]
|
||||
if greedy_indexes and greedy_indexes[-1] != len(specs) - 1:
|
||||
greedy_spec = specs[greedy_indexes[-1]]
|
||||
raise ValueError(f"参数 '{greedy_spec.name}' (GreedyStr) 必须是最后一个参数。")
|
||||
return specs
|
||||
|
||||
|
||||
def validate_schedule_signature(handler: Any) -> None:
|
||||
try:
|
||||
signature = inspect.signature(handler)
|
||||
except (TypeError, ValueError):
|
||||
return
|
||||
allowed_names = {"ctx", "context", "sched", "schedule"}
|
||||
invalid = [
|
||||
parameter.name
|
||||
for parameter in signature.parameters.values()
|
||||
if parameter.kind
|
||||
in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
and parameter.name not in allowed_names
|
||||
]
|
||||
if invalid:
|
||||
raise ValueError(
|
||||
"Schedule handler 只允许注入 ctx/context 和 sched/schedule 参数。"
|
||||
)
|
||||
|
||||
|
||||
def resolve_handler_candidate(instance: Any, name: str) -> tuple[Any, Any] | None:
|
||||
try:
|
||||
raw = inspect.getattr_static(instance, name)
|
||||
except AttributeError:
|
||||
return None
|
||||
candidates = [raw]
|
||||
wrapped = getattr(raw, "__func__", None)
|
||||
if wrapped is not None:
|
||||
candidates.append(wrapped)
|
||||
for candidate in candidates:
|
||||
meta = get_handler_meta(candidate)
|
||||
if meta is not None and meta.trigger is not None:
|
||||
return getattr(instance, name), meta
|
||||
return None
|
||||
|
||||
|
||||
def resolve_capability_candidate(instance: Any, name: str) -> tuple[Any, Any] | None:
|
||||
try:
|
||||
raw = inspect.getattr_static(instance, name)
|
||||
except AttributeError:
|
||||
return None
|
||||
candidates = [raw]
|
||||
wrapped = getattr(raw, "__func__", None)
|
||||
if wrapped is not None:
|
||||
candidates.append(wrapped)
|
||||
for candidate in candidates:
|
||||
meta = get_capability_meta(candidate)
|
||||
if meta is not None:
|
||||
return getattr(instance, name), meta
|
||||
return None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_param_specs",
|
||||
"is_injected_parameter",
|
||||
"param_type_name",
|
||||
"resolve_capability_candidate",
|
||||
"resolve_handler_candidate",
|
||||
"unwrap_optional",
|
||||
"validate_schedule_signature",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -19,7 +19,6 @@ import sys
|
||||
from pathlib import Path
|
||||
from typing import IO
|
||||
|
||||
from ..protocol.wire_codecs import make_protocol_codec
|
||||
from .loader import PluginEnvironmentManager
|
||||
from .supervisor import (
|
||||
SupervisorRuntime,
|
||||
@@ -50,10 +49,9 @@ __all__ = [
|
||||
async def run_supervisor(
|
||||
*,
|
||||
plugins_dir: Path = Path("plugins"),
|
||||
stdin: IO[str] | IO[bytes] | None = None,
|
||||
stdout: IO[str] | IO[bytes] | None = None,
|
||||
stdin: IO[str] | None = None,
|
||||
stdout: IO[str] | None = None,
|
||||
env_manager: PluginEnvironmentManager | None = None,
|
||||
worker_wire_codec: str = "json",
|
||||
) -> None:
|
||||
transport_stdin, transport_stdout, original_stdout = _prepare_stdio_transport(
|
||||
stdin,
|
||||
@@ -64,7 +62,6 @@ async def run_supervisor(
|
||||
transport=transport,
|
||||
plugins_dir=plugins_dir,
|
||||
env_manager=env_manager,
|
||||
worker_wire_codec_name=worker_wire_codec,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -82,39 +79,26 @@ async def run_plugin_worker(
|
||||
*,
|
||||
plugin_dir: Path | None = None,
|
||||
group_metadata: Path | None = None,
|
||||
stdin: IO[str] | IO[bytes] | None = None,
|
||||
stdout: IO[str] | IO[bytes] | None = None,
|
||||
wire_codec: str = "json",
|
||||
stdin: IO[str] | None = None,
|
||||
stdout: IO[str] | None = None,
|
||||
) -> None:
|
||||
if plugin_dir is None and group_metadata is None:
|
||||
raise ValueError("plugin_dir or group_metadata is required")
|
||||
if plugin_dir is not None and group_metadata is not None:
|
||||
raise ValueError("plugin_dir and group_metadata are mutually exclusive")
|
||||
|
||||
codec = make_protocol_codec(wire_codec)
|
||||
transport_stdin, transport_stdout, original_stdout = _prepare_stdio_transport(
|
||||
stdin,
|
||||
stdout,
|
||||
binary=codec.stdio_framing == "length_prefixed",
|
||||
)
|
||||
transport = StdioTransport(
|
||||
stdin=transport_stdin,
|
||||
stdout=transport_stdout,
|
||||
framing=codec.stdio_framing,
|
||||
)
|
||||
transport = StdioTransport(stdin=transport_stdin, stdout=transport_stdout)
|
||||
if group_metadata is not None:
|
||||
runtime = GroupWorkerRuntime(
|
||||
group_metadata_path=group_metadata,
|
||||
transport=transport,
|
||||
codec=codec,
|
||||
)
|
||||
else:
|
||||
assert plugin_dir is not None
|
||||
runtime = PluginWorkerRuntime(
|
||||
plugin_dir=plugin_dir,
|
||||
transport=transport,
|
||||
codec=codec,
|
||||
)
|
||||
runtime = PluginWorkerRuntime(plugin_dir=plugin_dir, transport=transport)
|
||||
try:
|
||||
await runtime.start()
|
||||
stop_event = asyncio.Event()
|
||||
@@ -132,18 +116,10 @@ async def run_websocket_server(
|
||||
port: int = 8765,
|
||||
path: str = "/",
|
||||
plugin_dir: Path | None = None,
|
||||
wire_codec: str = "json",
|
||||
) -> None:
|
||||
codec = make_protocol_codec(wire_codec)
|
||||
runtime = PluginWorkerRuntime(
|
||||
plugin_dir=plugin_dir or Path.cwd(),
|
||||
transport=WebSocketServerTransport(
|
||||
host=host,
|
||||
port=port,
|
||||
path=path,
|
||||
frame_type=codec.websocket_frame_type,
|
||||
),
|
||||
codec=codec,
|
||||
transport=WebSocketServerTransport(host=host, port=port, path=path),
|
||||
)
|
||||
try:
|
||||
await runtime.start()
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
"""Capability invocation dispatcher.
|
||||
|
||||
本模块实现能力调用的分发器,负责:
|
||||
1. 接收能力调用请求,定位对应的已注册能力
|
||||
2. 构建调用上下文 (Context),注入必要的依赖
|
||||
3. 支持同步和流式两种调用模式
|
||||
4. 管理活跃调用任务的生命周期和取消
|
||||
|
||||
参数注入策略:
|
||||
按类型注入 Context / CancelToken / dict,或按参数名注入
|
||||
ctx / context / payload / input / data / cancel_token / token。
|
||||
若无法匹配则抛出详细的错误信息,帮助开发者定位问题。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import typing
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, get_type_hints
|
||||
|
||||
from .._invocation_context import caller_plugin_scope
|
||||
from ..context import CancelToken, Context
|
||||
from ..errors import AstrBotError
|
||||
from ._streaming import StreamExecution
|
||||
from .loader import LoadedCapability
|
||||
|
||||
|
||||
class CapabilityDispatcher:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plugin_id: str,
|
||||
peer,
|
||||
capabilities: list[LoadedCapability],
|
||||
) -> None:
|
||||
self._plugin_id = plugin_id
|
||||
self._peer = peer
|
||||
self._capabilities = {item.descriptor.name: item for item in capabilities}
|
||||
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
|
||||
|
||||
async def invoke(
|
||||
self,
|
||||
message,
|
||||
cancel_token: CancelToken,
|
||||
) -> dict[str, Any] | StreamExecution:
|
||||
loaded = self._capabilities.get(message.capability)
|
||||
if loaded is None:
|
||||
raise LookupError(f"capability not found: {message.capability}")
|
||||
|
||||
plugin_id = self._resolve_plugin_id(loaded)
|
||||
ctx = Context(
|
||||
peer=self._peer,
|
||||
plugin_id=plugin_id,
|
||||
cancel_token=cancel_token,
|
||||
)
|
||||
|
||||
with caller_plugin_scope(plugin_id):
|
||||
task = asyncio.create_task(
|
||||
self._run_capability(
|
||||
loaded,
|
||||
payload=dict(message.input),
|
||||
ctx=ctx,
|
||||
cancel_token=cancel_token,
|
||||
stream=bool(message.stream),
|
||||
)
|
||||
)
|
||||
self._active[message.id] = (task, cancel_token)
|
||||
try:
|
||||
return await task
|
||||
finally:
|
||||
self._active.pop(message.id, None)
|
||||
|
||||
def _resolve_plugin_id(self, loaded: LoadedCapability) -> str:
|
||||
if loaded.plugin_id:
|
||||
return loaded.plugin_id
|
||||
return self._plugin_id
|
||||
|
||||
async def cancel(self, request_id: str) -> None:
|
||||
active = self._active.get(request_id)
|
||||
if active is None:
|
||||
return
|
||||
task, cancel_token = active
|
||||
cancel_token.cancel()
|
||||
task.cancel()
|
||||
|
||||
async def _run_capability(
|
||||
self,
|
||||
loaded: LoadedCapability,
|
||||
*,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
stream: bool,
|
||||
) -> dict[str, Any] | StreamExecution:
|
||||
result = loaded.callable(
|
||||
*self._build_args(
|
||||
loaded.callable,
|
||||
payload,
|
||||
ctx,
|
||||
cancel_token,
|
||||
plugin_id=self._resolve_plugin_id(loaded),
|
||||
capability_name=loaded.descriptor.name,
|
||||
)
|
||||
)
|
||||
if stream:
|
||||
if inspect.isasyncgen(result):
|
||||
return StreamExecution(
|
||||
iterator=self._iterate_generator(result),
|
||||
finalize=lambda chunks: {"items": chunks},
|
||||
)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
if isinstance(result, StreamExecution):
|
||||
return result
|
||||
raise AstrBotError.protocol_error(
|
||||
"stream=true 的插件 capability 必须返回 async generator 或 StreamExecution"
|
||||
)
|
||||
|
||||
if inspect.isasyncgen(result):
|
||||
raise AstrBotError.protocol_error(
|
||||
"stream=false 的插件 capability 不能返回 async generator"
|
||||
)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
return self._normalize_output(result)
|
||||
|
||||
def _build_args(
|
||||
self,
|
||||
handler,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
*,
|
||||
plugin_id: str | None = None,
|
||||
capability_name: str | None = None,
|
||||
) -> list[Any]:
|
||||
signature = inspect.signature(handler)
|
||||
args: list[Any] = []
|
||||
|
||||
type_hints: dict[str, Any] = {}
|
||||
try:
|
||||
type_hints = get_type_hints(handler)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for parameter in signature.parameters.values():
|
||||
if parameter.kind not in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
):
|
||||
continue
|
||||
|
||||
injected = None
|
||||
param_type = type_hints.get(parameter.name)
|
||||
if param_type is not None:
|
||||
injected = self._inject_by_type(param_type, payload, ctx, cancel_token)
|
||||
|
||||
if injected is None:
|
||||
if parameter.name in {"ctx", "context"}:
|
||||
injected = ctx
|
||||
elif parameter.name in {"payload", "input", "data"}:
|
||||
injected = payload
|
||||
elif parameter.name in {"cancel_token", "token"}:
|
||||
injected = cancel_token
|
||||
|
||||
if injected is None:
|
||||
if parameter.default is not parameter.empty:
|
||||
continue
|
||||
raise TypeError(
|
||||
self._format_capability_injection_error(
|
||||
handler=handler,
|
||||
parameter_name=parameter.name,
|
||||
plugin_id=plugin_id,
|
||||
capability_name=capability_name,
|
||||
payload=payload,
|
||||
)
|
||||
)
|
||||
args.append(injected)
|
||||
|
||||
return args
|
||||
|
||||
def _inject_by_type(
|
||||
self,
|
||||
param_type: Any,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
) -> Any:
|
||||
origin = typing.get_origin(param_type)
|
||||
if origin is typing.Union:
|
||||
type_args = typing.get_args(param_type)
|
||||
non_none_types = [item for item in type_args if item is not type(None)]
|
||||
if len(non_none_types) == 1:
|
||||
param_type = non_none_types[0]
|
||||
origin = typing.get_origin(param_type)
|
||||
|
||||
if param_type is Context or (
|
||||
isinstance(param_type, type) and issubclass(param_type, Context)
|
||||
):
|
||||
return ctx
|
||||
if param_type is CancelToken or (
|
||||
isinstance(param_type, type) and issubclass(param_type, CancelToken)
|
||||
):
|
||||
return cancel_token
|
||||
if param_type is dict or origin is dict:
|
||||
return payload
|
||||
return None
|
||||
|
||||
def _format_capability_injection_error(
|
||||
self,
|
||||
*,
|
||||
handler,
|
||||
parameter_name: str,
|
||||
plugin_id: str | None,
|
||||
capability_name: str | None,
|
||||
payload: dict[str, Any],
|
||||
) -> str:
|
||||
plugin_text = plugin_id or self._plugin_id
|
||||
target = capability_name or getattr(handler, "__name__", "<anonymous>")
|
||||
payload_keys = sorted(str(key) for key in payload.keys())
|
||||
payload_keys_text = ", ".join(payload_keys) if payload_keys else "<none>"
|
||||
return (
|
||||
f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:"
|
||||
f"必填参数 '{parameter_name}' 无法注入。"
|
||||
f"签名: {getattr(handler, '__name__', '<anonymous>')}"
|
||||
f"{self._callable_signature(handler)}。"
|
||||
"当前支持按类型注入 Context / CancelToken / dict,"
|
||||
"按参数名注入 ctx / context / payload / input / data / cancel_token / token,"
|
||||
f"以及 payload 中现有键:{payload_keys_text}。"
|
||||
)
|
||||
|
||||
async def _iterate_generator(
|
||||
self,
|
||||
generator: AsyncIterator[Any],
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
async for item in generator:
|
||||
yield self._normalize_chunk(item)
|
||||
|
||||
def _normalize_chunk(self, item: Any) -> dict[str, Any]:
|
||||
output = self._normalize_output(item)
|
||||
if output:
|
||||
return output
|
||||
return {"ok": True}
|
||||
|
||||
def _normalize_output(self, result: Any) -> dict[str, Any]:
|
||||
if result is None:
|
||||
return {}
|
||||
if isinstance(result, dict):
|
||||
return result
|
||||
model_dump = getattr(result, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
dumped = model_dump()
|
||||
if isinstance(dumped, dict):
|
||||
return dumped
|
||||
raise AstrBotError.invalid_input("插件 capability 必须返回 dict 或可序列化对象")
|
||||
|
||||
@staticmethod
|
||||
def _callable_signature(handler) -> str:
|
||||
try:
|
||||
return str(inspect.signature(handler))
|
||||
except (TypeError, ValueError):
|
||||
return "(?)"
|
||||
|
||||
|
||||
__all__ = ["CapabilityDispatcher"]
|
||||
@@ -97,35 +97,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import inspect
|
||||
import json
|
||||
import re
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .._invocation_context import current_caller_plugin_id
|
||||
from ..errors import AstrBotError
|
||||
from ..protocol.descriptors import (
|
||||
BUILTIN_CAPABILITY_SCHEMAS,
|
||||
RESERVED_CAPABILITY_PREFIXES,
|
||||
CapabilityDescriptor,
|
||||
SessionRef,
|
||||
)
|
||||
from ._capability_router_builtins import BuiltinCapabilityRouterMixin
|
||||
from ._streaming import StreamExecution
|
||||
|
||||
CallHandler = Callable[[str, dict[str, Any], object], Awaitable[dict[str, Any]]]
|
||||
FinalizeHandler = Callable[[list[dict[str, Any]]], dict[str, Any]]
|
||||
CAPABILITY_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9_]*\.[a-z][a-z0-9_]*$")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class StreamExecution:
|
||||
iterator: AsyncIterator[dict[str, Any]]
|
||||
finalize: FinalizeHandler
|
||||
collect_chunks: bool = True
|
||||
|
||||
|
||||
StreamHandler = Callable[
|
||||
[str, dict[str, Any], object],
|
||||
AsyncIterator[dict[str, Any]]
|
||||
@@ -149,14 +141,18 @@ class _RegisteredPlugin:
|
||||
config: dict[str, Any]
|
||||
|
||||
|
||||
class CapabilityRouter:
|
||||
class CapabilityRouter(BuiltinCapabilityRouterMixin):
|
||||
def __init__(self) -> None:
|
||||
self._registrations: dict[str, _CapabilityRegistration] = {}
|
||||
self.db_store: dict[str, Any] = {}
|
||||
self.memory_store: dict[str, dict[str, Any]] = {}
|
||||
self.sent_messages: list[dict[str, Any]] = []
|
||||
self.event_actions: list[dict[str, Any]] = []
|
||||
self._event_streams: dict[str, dict[str, Any]] = {}
|
||||
self.http_api_store: list[dict[str, Any]] = []
|
||||
self._plugins: dict[str, _RegisteredPlugin] = {}
|
||||
self._system_data_root = Path.cwd() / ".astrbot_sdk_testing" / "plugin_data"
|
||||
self._session_waiters: dict[str, set[str]] = {}
|
||||
self._db_watch_subscriptions: dict[
|
||||
str, tuple[str | None, asyncio.Queue[dict[str, Any]]]
|
||||
] = {}
|
||||
@@ -212,9 +208,7 @@ class CapabilityRouter:
|
||||
queue.put_nowait(event)
|
||||
|
||||
def descriptors(self) -> list[CapabilityDescriptor]:
|
||||
return [
|
||||
entry.descriptor for entry in self._registrations.values() if entry.exposed
|
||||
]
|
||||
return [entry.descriptor for entry in self._registrations.values()]
|
||||
|
||||
def contains(self, name: str) -> bool:
|
||||
return name in self._registrations
|
||||
@@ -231,7 +225,13 @@ class CapabilityRouter:
|
||||
finalize: FinalizeHandler | None = None,
|
||||
exposed: bool = True,
|
||||
) -> None:
|
||||
if not CAPABILITY_NAME_PATTERN.fullmatch(descriptor.name):
|
||||
is_internal_reserved = not exposed and descriptor.name.startswith(
|
||||
RESERVED_CAPABILITY_PREFIXES
|
||||
)
|
||||
if (
|
||||
not CAPABILITY_NAME_PATTERN.fullmatch(descriptor.name)
|
||||
and not is_internal_reserved
|
||||
):
|
||||
raise ValueError(
|
||||
f"capability 名称必须匹配 {{namespace}}.{{method}}:{descriptor.name}"
|
||||
)
|
||||
@@ -320,599 +320,6 @@ class CapabilityRouter:
|
||||
collect_chunks=execution.collect_chunks,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Built-in capability registration
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _register_builtin_capabilities(self) -> None:
|
||||
"""注册全部内建 capability。"""
|
||||
self._register_llm_capabilities()
|
||||
self._register_memory_capabilities()
|
||||
self._register_db_capabilities()
|
||||
self._register_platform_capabilities()
|
||||
self._register_http_capabilities()
|
||||
self._register_metadata_capabilities()
|
||||
|
||||
def _builtin_descriptor(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
*,
|
||||
supports_stream: bool = False,
|
||||
cancelable: bool = False,
|
||||
) -> CapabilityDescriptor:
|
||||
"""构建内建 capability 描述符,schema 从注册表读取。"""
|
||||
schema = BUILTIN_CAPABILITY_SCHEMAS[name]
|
||||
return CapabilityDescriptor(
|
||||
name=name,
|
||||
description=description,
|
||||
input_schema=copy.deepcopy(schema["input"]),
|
||||
output_schema=copy.deepcopy(schema["output"]),
|
||||
supports_stream=supports_stream,
|
||||
cancelable=cancelable,
|
||||
)
|
||||
|
||||
def _resolve_target(
|
||||
self, payload: dict[str, Any]
|
||||
) -> tuple[str, dict[str, Any] | None]:
|
||||
"""从 payload 解析 session + target。"""
|
||||
target_payload = payload.get("target")
|
||||
if isinstance(target_payload, dict):
|
||||
target = SessionRef.model_validate(target_payload)
|
||||
return target.session, target.to_payload()
|
||||
return str(payload.get("session", "")), None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# LLM handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _llm_chat(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prompt = str(payload.get("prompt", ""))
|
||||
return {"text": f"Echo: {prompt}"}
|
||||
|
||||
async def _llm_chat_raw(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prompt = str(payload.get("prompt", ""))
|
||||
text = f"Echo: {prompt}"
|
||||
return {
|
||||
"text": text,
|
||||
"usage": {
|
||||
"input_tokens": len(prompt),
|
||||
"output_tokens": len(text),
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
"tool_calls": [],
|
||||
}
|
||||
|
||||
async def _llm_stream(
|
||||
self,
|
||||
_request_id: str,
|
||||
payload: dict[str, Any],
|
||||
token,
|
||||
) -> AsyncIterator[dict[str, Any]]: # type: ignore[override]
|
||||
text = f"Echo: {str(payload.get('prompt', ''))}"
|
||||
for char in text:
|
||||
token.raise_if_cancelled()
|
||||
await asyncio.sleep(0)
|
||||
yield {"text": char}
|
||||
|
||||
def _register_llm_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("llm.chat", "发送对话请求,返回文本"),
|
||||
call_handler=self._llm_chat,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("llm.chat_raw", "发送对话请求,返回完整响应"),
|
||||
call_handler=self._llm_chat_raw,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"llm.stream_chat",
|
||||
"流式对话",
|
||||
supports_stream=True,
|
||||
cancelable=True,
|
||||
),
|
||||
stream_handler=self._llm_stream,
|
||||
finalize=lambda chunks: {
|
||||
"text": "".join(item.get("text", "") for item in chunks)
|
||||
},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Memory handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _memory_search(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
query = str(payload.get("query", ""))
|
||||
items = [
|
||||
{"key": key, "value": value}
|
||||
for key, value in self.memory_store.items()
|
||||
if query in key or query in json.dumps(value, ensure_ascii=False)
|
||||
]
|
||||
return {"items": items}
|
||||
|
||||
async def _memory_save(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
if not isinstance(value, dict):
|
||||
raise AstrBotError.invalid_input("memory.save 的 value 必须是 object")
|
||||
self.memory_store[key] = value
|
||||
return {}
|
||||
|
||||
async def _memory_get(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
return {"value": self.memory_store.get(str(payload.get("key", "")))}
|
||||
|
||||
async def _memory_delete(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
self.memory_store.pop(str(payload.get("key", "")), None)
|
||||
return {}
|
||||
|
||||
async def _memory_save_with_ttl(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
"""保存带 TTL 的记忆项(测试实现,TTL 仅记录但不实际过期)。"""
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
ttl_seconds = payload.get("ttl_seconds", 0)
|
||||
if not isinstance(value, dict):
|
||||
raise AstrBotError.invalid_input(
|
||||
"memory.save_with_ttl 的 value 必须是 object"
|
||||
)
|
||||
# 在测试实现中,我们只存储值,TTL 由实际后端实现
|
||||
self.memory_store[key] = {"value": value, "ttl_seconds": ttl_seconds}
|
||||
return {}
|
||||
|
||||
async def _memory_get_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
"""批量获取多个记忆项。"""
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("memory.get_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
items = []
|
||||
for key in keys:
|
||||
stored = self.memory_store.get(key)
|
||||
# 如果存储的是带 TTL 的结构,提取实际值
|
||||
if (
|
||||
isinstance(stored, dict)
|
||||
and "value" in stored
|
||||
and "ttl_seconds" in stored
|
||||
):
|
||||
value = stored["value"]
|
||||
else:
|
||||
value = stored
|
||||
items.append({"key": key, "value": value})
|
||||
return {"items": items}
|
||||
|
||||
async def _memory_delete_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
"""批量删除多个记忆项。"""
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("memory.delete_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
deleted_count = 0
|
||||
for key in keys:
|
||||
if key in self.memory_store:
|
||||
del self.memory_store[key]
|
||||
deleted_count += 1
|
||||
return {"deleted_count": deleted_count}
|
||||
|
||||
async def _memory_stats(
|
||||
self, _request_id: str, _payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
"""获取记忆统计信息。"""
|
||||
total_items = len(self.memory_store)
|
||||
# 简单估算字节大小
|
||||
total_bytes = sum(
|
||||
len(str(key)) + len(str(value)) for key, value in self.memory_store.items()
|
||||
)
|
||||
return {"total_items": total_items, "total_bytes": total_bytes}
|
||||
|
||||
def _register_memory_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.search", "搜索记忆"),
|
||||
call_handler=self._memory_search,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.save", "保存记忆"),
|
||||
call_handler=self._memory_save,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.get", "读取单条记忆"),
|
||||
call_handler=self._memory_get,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.delete", "删除记忆"),
|
||||
call_handler=self._memory_delete,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.save_with_ttl", "保存带过期时间的记忆"),
|
||||
call_handler=self._memory_save_with_ttl,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.get_many", "批量获取记忆"),
|
||||
call_handler=self._memory_get_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.delete_many", "批量删除记忆"),
|
||||
call_handler=self._memory_delete_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("memory.stats", "获取记忆统计信息"),
|
||||
call_handler=self._memory_stats,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# DB handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _db_get(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
return {"value": self.db_store.get(str(payload.get("key", "")))}
|
||||
|
||||
async def _db_set(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
value = payload.get("value")
|
||||
self.db_store[key] = value
|
||||
self._emit_db_change(op="set", key=key, value=value)
|
||||
return {}
|
||||
|
||||
async def _db_delete(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
key = str(payload.get("key", ""))
|
||||
self.db_store.pop(key, None)
|
||||
self._emit_db_change(op="delete", key=key, value=None)
|
||||
return {}
|
||||
|
||||
async def _db_list(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
prefix = payload.get("prefix")
|
||||
keys = sorted(self.db_store.keys())
|
||||
if isinstance(prefix, str):
|
||||
keys = [item for item in keys if item.startswith(prefix)]
|
||||
return {"keys": keys}
|
||||
|
||||
async def _db_get_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
keys_payload = payload.get("keys")
|
||||
if not isinstance(keys_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("db.get_many 的 keys 必须是数组")
|
||||
keys = [str(item) for item in keys_payload]
|
||||
items = [{"key": key, "value": self.db_store.get(key)} for key in keys]
|
||||
return {"items": items}
|
||||
|
||||
async def _db_set_many(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
items_payload = payload.get("items")
|
||||
if not isinstance(items_payload, (list, tuple)):
|
||||
raise AstrBotError.invalid_input("db.set_many 的 items 必须是数组")
|
||||
for entry in items_payload:
|
||||
if not isinstance(entry, dict):
|
||||
raise AstrBotError.invalid_input(
|
||||
"db.set_many 的 items 必须是 object 数组"
|
||||
)
|
||||
key = str(entry.get("key", ""))
|
||||
value = entry.get("value")
|
||||
self.db_store[key] = value
|
||||
self._emit_db_change(op="set", key=key, value=value)
|
||||
return {}
|
||||
|
||||
async def _db_watch(
|
||||
self, request_id: str, payload: dict[str, Any], _token
|
||||
) -> StreamExecution:
|
||||
prefix = payload.get("prefix")
|
||||
prefix_value: str | None
|
||||
if isinstance(prefix, str):
|
||||
prefix_value = prefix
|
||||
elif prefix is None:
|
||||
prefix_value = None
|
||||
else:
|
||||
raise AstrBotError.invalid_input("db.watch 的 prefix 必须是 string 或 null")
|
||||
|
||||
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
|
||||
self._db_watch_subscriptions[request_id] = (prefix_value, queue)
|
||||
|
||||
async def iterator() -> AsyncIterator[dict[str, Any]]:
|
||||
try:
|
||||
while True:
|
||||
yield await queue.get()
|
||||
finally:
|
||||
self._db_watch_subscriptions.pop(request_id, None)
|
||||
|
||||
return StreamExecution(
|
||||
iterator=iterator(),
|
||||
finalize=lambda _chunks: {},
|
||||
collect_chunks=False,
|
||||
)
|
||||
|
||||
def _register_db_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("db.get", "读取 KV"),
|
||||
call_handler=self._db_get,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.set", "写入 KV"),
|
||||
call_handler=self._db_set,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.delete", "删除 KV"),
|
||||
call_handler=self._db_delete,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.list", "列出 KV"),
|
||||
call_handler=self._db_list,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.get_many", "批量读取 KV"),
|
||||
call_handler=self._db_get_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("db.set_many", "批量写入 KV"),
|
||||
call_handler=self._db_set_many,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"db.watch",
|
||||
"订阅 KV 变更",
|
||||
supports_stream=True,
|
||||
cancelable=True,
|
||||
),
|
||||
stream_handler=self._db_watch,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Platform handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _platform_send(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
text = str(payload.get("text", ""))
|
||||
message_id = f"msg_{len(self.sent_messages) + 1}"
|
||||
sent = {"message_id": message_id, "session": session, "text": text}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_send_image(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
image_url = str(payload.get("image_url", ""))
|
||||
message_id = f"img_{len(self.sent_messages) + 1}"
|
||||
sent = {
|
||||
"message_id": message_id,
|
||||
"session": session,
|
||||
"image_url": image_url,
|
||||
}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_send_chain(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, target = self._resolve_target(payload)
|
||||
chain = payload.get("chain")
|
||||
if not isinstance(chain, list) or not all(
|
||||
isinstance(item, dict) for item in chain
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"platform.send_chain 的 chain 必须是 object 数组"
|
||||
)
|
||||
message_id = f"chain_{len(self.sent_messages) + 1}"
|
||||
sent = {
|
||||
"message_id": message_id,
|
||||
"session": session,
|
||||
"chain": [dict(item) for item in chain],
|
||||
}
|
||||
if target is not None:
|
||||
sent["target"] = target
|
||||
self.sent_messages.append(sent)
|
||||
return {"message_id": message_id}
|
||||
|
||||
async def _platform_get_members(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
session, _target = self._resolve_target(payload)
|
||||
return {
|
||||
"members": [
|
||||
{"user_id": f"{session}:member-1", "nickname": "Member 1"},
|
||||
{"user_id": f"{session}:member-2", "nickname": "Member 2"},
|
||||
]
|
||||
}
|
||||
|
||||
def _register_platform_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send", "发送消息"),
|
||||
call_handler=self._platform_send,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send_image", "发送图片"),
|
||||
call_handler=self._platform_send_image,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.send_chain", "发送消息链"),
|
||||
call_handler=self._platform_send_chain,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("platform.get_members", "获取群成员"),
|
||||
call_handler=self._platform_get_members,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# HTTP handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _http_register_api(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
methods_payload = payload.get("methods")
|
||||
if not isinstance(methods_payload, list) or not all(
|
||||
isinstance(item, str) for item in methods_payload
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.register_api 的 methods 必须是 string 数组"
|
||||
)
|
||||
|
||||
route = str(payload.get("route", "")).strip()
|
||||
handler_capability = str(payload.get("handler_capability", "")).strip()
|
||||
if not route or not handler_capability:
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.register_api 需要 route 和 handler_capability"
|
||||
)
|
||||
|
||||
plugin_name = self._require_caller_plugin_id("http.register_api")
|
||||
methods = sorted({method.upper() for method in methods_payload if method})
|
||||
entry = {
|
||||
"route": route,
|
||||
"methods": methods,
|
||||
"handler_capability": handler_capability,
|
||||
"description": str(payload.get("description", "")),
|
||||
"plugin_id": plugin_name,
|
||||
}
|
||||
self.http_api_store = [
|
||||
item
|
||||
for item in self.http_api_store
|
||||
if not (
|
||||
item.get("route") == route
|
||||
and item.get("plugin_id") == entry["plugin_id"]
|
||||
and item.get("methods") == methods
|
||||
)
|
||||
]
|
||||
self.http_api_store.append(entry)
|
||||
return {}
|
||||
|
||||
async def _http_unregister_api(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
route = str(payload.get("route", "")).strip()
|
||||
methods_payload = payload.get("methods")
|
||||
if not isinstance(methods_payload, list) or not all(
|
||||
isinstance(item, str) for item in methods_payload
|
||||
):
|
||||
raise AstrBotError.invalid_input(
|
||||
"http.unregister_api 的 methods 必须是 string 数组"
|
||||
)
|
||||
|
||||
plugin_name = self._require_caller_plugin_id("http.unregister_api")
|
||||
methods = {method.upper() for method in methods_payload if method}
|
||||
updated: list[dict[str, Any]] = []
|
||||
for entry in self.http_api_store:
|
||||
if entry.get("route") != route:
|
||||
updated.append(entry)
|
||||
continue
|
||||
if entry.get("plugin_id") != plugin_name:
|
||||
updated.append(entry)
|
||||
continue
|
||||
if not methods:
|
||||
continue
|
||||
remaining_methods = [
|
||||
method for method in entry.get("methods", []) if method not in methods
|
||||
]
|
||||
if remaining_methods:
|
||||
updated.append({**entry, "methods": remaining_methods})
|
||||
self.http_api_store = updated
|
||||
return {}
|
||||
|
||||
async def _http_list_apis(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugin_name = self._require_caller_plugin_id("http.list_apis")
|
||||
apis = [
|
||||
dict(entry)
|
||||
for entry in self.http_api_store
|
||||
if entry.get("plugin_id") == plugin_name
|
||||
]
|
||||
return {"apis": apis}
|
||||
|
||||
def _register_http_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("http.register_api", "注册 HTTP 路由"),
|
||||
call_handler=self._http_register_api,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("http.unregister_api", "注销 HTTP 路由"),
|
||||
call_handler=self._http_unregister_api,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("http.list_apis", "列出 HTTP 路由"),
|
||||
call_handler=self._http_list_apis,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Metadata handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _metadata_get_plugin(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
name = str(payload.get("name", "")).strip()
|
||||
plugin = self._plugins.get(name)
|
||||
if plugin is None:
|
||||
return {"plugin": None}
|
||||
return {"plugin": dict(plugin.metadata)}
|
||||
|
||||
async def _metadata_list_plugins(
|
||||
self, _request_id: str, _payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
plugins = [
|
||||
dict(self._plugins[name].metadata) for name in sorted(self._plugins.keys())
|
||||
]
|
||||
return {"plugins": plugins}
|
||||
|
||||
async def _metadata_get_plugin_config(
|
||||
self, _request_id: str, payload: dict[str, Any], _token
|
||||
) -> dict[str, Any]:
|
||||
name = str(payload.get("name", "")).strip()
|
||||
caller_plugin_id = self._require_caller_plugin_id("metadata.get_plugin_config")
|
||||
if name != caller_plugin_id:
|
||||
return {"config": None}
|
||||
plugin = self._plugins.get(name)
|
||||
if plugin is None:
|
||||
return {"config": None}
|
||||
return {"config": dict(plugin.config)}
|
||||
|
||||
def _register_metadata_capabilities(self) -> None:
|
||||
self.register(
|
||||
self._builtin_descriptor("metadata.get_plugin", "获取单个插件元数据"),
|
||||
call_handler=self._metadata_get_plugin,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor("metadata.list_plugins", "列出插件元数据"),
|
||||
call_handler=self._metadata_list_plugins,
|
||||
)
|
||||
self.register(
|
||||
self._builtin_descriptor(
|
||||
"metadata.get_plugin_config",
|
||||
"获取插件配置",
|
||||
),
|
||||
call_handler=self._metadata_get_plugin_config,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Schema validation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -37,6 +37,10 @@ GROUP_STATE_FILE_NAME = ".group-venv-state.json"
|
||||
|
||||
_EXACT_PIN_PATTERN = re.compile(r"^([A-Za-z0-9_.-]+)==([^\s;]+)$")
|
||||
_NORMALIZE_PATTERN = re.compile(r"[-_.]+")
|
||||
_PYVENV_VERSION_PATTERN = re.compile(
|
||||
r"^(?:version|version_info)\s*=\s*(\d+\.\d+)(?:\.\d+)?\s*$",
|
||||
re.IGNORECASE | re.MULTILINE,
|
||||
)
|
||||
|
||||
|
||||
def _venv_python_path(venv_path: Path) -> Path:
|
||||
@@ -49,6 +53,19 @@ def _normalize_package_name(name: str) -> str:
|
||||
return _NORMALIZE_PATTERN.sub("-", name).lower()
|
||||
|
||||
|
||||
def _read_pyvenv_major_minor(pyvenv_cfg: Path) -> str | None:
|
||||
if not pyvenv_cfg.exists():
|
||||
return None
|
||||
try:
|
||||
content = pyvenv_cfg.read_text(encoding="utf-8")
|
||||
except OSError:
|
||||
return None
|
||||
match = _PYVENV_VERSION_PATTERN.search(content)
|
||||
if match is None:
|
||||
return None
|
||||
return match.group(1)
|
||||
|
||||
|
||||
def _requirement_lines(plugin: PluginSpec) -> list[str]:
|
||||
if not plugin.requirements_path.exists():
|
||||
return []
|
||||
@@ -612,15 +629,7 @@ class GroupEnvironmentManager:
|
||||
|
||||
@staticmethod
|
||||
def _matches_python_version(venv_path: Path, version: str) -> bool:
|
||||
pyvenv_cfg = venv_path / "pyvenv.cfg"
|
||||
if not pyvenv_cfg.exists():
|
||||
return False
|
||||
try:
|
||||
content = pyvenv_cfg.read_text(encoding="utf-8")
|
||||
except OSError:
|
||||
return False
|
||||
match = re.search(r"version\s*=\s*(\d+\.\d+)\.\d+", content, re.IGNORECASE)
|
||||
return match is not None and match.group(1) == version
|
||||
return _read_pyvenv_major_minor(venv_path / "pyvenv.cfg") == version
|
||||
|
||||
@staticmethod
|
||||
def _load_state(state_path: Path) -> dict[str, object]:
|
||||
|
||||
@@ -27,17 +27,25 @@ import inspect
|
||||
import re
|
||||
import shlex
|
||||
import typing
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, get_type_hints
|
||||
|
||||
from .._invocation_context import caller_plugin_scope
|
||||
from ..context import CancelToken, Context
|
||||
from ..errors import AstrBotError
|
||||
from ..events import MessageEvent
|
||||
from ..protocol.descriptors import CommandTrigger, MessageTrigger
|
||||
from ..filters import LocalFilterBinding
|
||||
from ..message_components import BaseMessageComponent
|
||||
from ..message_result import MessageChain, MessageEventResult, coerce_message_chain
|
||||
from ..protocol.descriptors import (
|
||||
CommandTrigger,
|
||||
MessageTrigger,
|
||||
ParamSpec,
|
||||
ScheduleTrigger,
|
||||
)
|
||||
from ..schedule import ScheduleContext
|
||||
from ..session_waiter import SessionWaiterManager
|
||||
from ..star import Star
|
||||
from .capability_router import StreamExecution
|
||||
from .loader import LoadedCapability, LoadedHandler
|
||||
from .capability_dispatcher import CapabilityDispatcher
|
||||
from .loader import LoadedHandler
|
||||
|
||||
|
||||
class HandlerDispatcher:
|
||||
@@ -46,9 +54,27 @@ class HandlerDispatcher:
|
||||
self._peer = peer
|
||||
self._handlers = {item.descriptor.id: item for item in handlers}
|
||||
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
|
||||
self._session_waiters = SessionWaiterManager(plugin_id=plugin_id, peer=peer)
|
||||
setattr(peer, "_session_waiter_manager", self._session_waiters)
|
||||
|
||||
async def invoke(self, message, cancel_token: CancelToken) -> dict[str, Any]:
|
||||
handler_id = str(message.input.get("handler_id", ""))
|
||||
if handler_id == "__sdk_session_waiter__":
|
||||
plugin_id = self._plugin_id
|
||||
ctx = Context(
|
||||
peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token
|
||||
)
|
||||
event = MessageEvent.from_payload(
|
||||
message.input.get("event", {}), context=ctx
|
||||
)
|
||||
event.bind_reply_handler(self._create_reply_handler(ctx, event))
|
||||
task = asyncio.create_task(self._session_waiters.dispatch(event))
|
||||
self._active[message.id] = (task, cancel_token)
|
||||
try:
|
||||
return await task
|
||||
finally:
|
||||
self._active.pop(message.id, None)
|
||||
|
||||
loaded = self._handlers.get(handler_id)
|
||||
if loaded is None:
|
||||
raise LookupError(f"handler not found: {handler_id}")
|
||||
@@ -57,6 +83,9 @@ class HandlerDispatcher:
|
||||
ctx = Context(peer=self._peer, plugin_id=plugin_id, cancel_token=cancel_token)
|
||||
event = MessageEvent.from_payload(message.input.get("event", {}), context=ctx)
|
||||
event.bind_reply_handler(self._create_reply_handler(ctx, event))
|
||||
schedule_context = self._build_schedule_context(
|
||||
loaded, message.input.get("event", {})
|
||||
)
|
||||
|
||||
# 提取 args 用于兼容 handler 签名
|
||||
raw_args = message.input.get("args") or {}
|
||||
@@ -65,7 +94,15 @@ class HandlerDispatcher:
|
||||
args = self._derive_args(loaded, event)
|
||||
|
||||
with caller_plugin_scope(plugin_id):
|
||||
task = asyncio.create_task(self._run_handler(loaded, event, ctx, args))
|
||||
task = asyncio.create_task(
|
||||
self._run_handler(
|
||||
loaded,
|
||||
event,
|
||||
ctx,
|
||||
args,
|
||||
schedule_context=schedule_context,
|
||||
)
|
||||
)
|
||||
self._active[message.id] = (task, cancel_token)
|
||||
try:
|
||||
return await task
|
||||
@@ -108,17 +145,31 @@ class HandlerDispatcher:
|
||||
event: MessageEvent,
|
||||
ctx: Context,
|
||||
args: dict[str, Any] | None = None,
|
||||
*,
|
||||
schedule_context: ScheduleContext | None = None,
|
||||
) -> dict[str, Any]:
|
||||
summary = {"sent_message": False, "stop": False, "call_llm": False}
|
||||
try:
|
||||
if not self._run_local_filters(
|
||||
loaded.local_filters,
|
||||
event=event,
|
||||
ctx=ctx,
|
||||
):
|
||||
return summary
|
||||
parsed_args = (
|
||||
self._parse_handler_args(loaded.descriptor.param_specs, args or {})
|
||||
if loaded.descriptor.param_specs
|
||||
else dict(args or {})
|
||||
)
|
||||
result = loaded.callable(
|
||||
*self._build_args(
|
||||
loaded.callable,
|
||||
event,
|
||||
ctx,
|
||||
args,
|
||||
parsed_args,
|
||||
plugin_id=self._resolve_plugin_id(loaded),
|
||||
handler_ref=loaded.descriptor.id,
|
||||
schedule_context=schedule_context,
|
||||
)
|
||||
)
|
||||
if inspect.isasyncgen(result):
|
||||
@@ -127,6 +178,7 @@ class HandlerDispatcher:
|
||||
summary,
|
||||
await self._handle_result_item(item, event, ctx),
|
||||
)
|
||||
summary["stop"] = bool(summary.get("stop")) or event.is_stopped()
|
||||
return summary
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
@@ -135,6 +187,7 @@ class HandlerDispatcher:
|
||||
summary,
|
||||
await self._handle_result_item(result, event, ctx),
|
||||
)
|
||||
summary["stop"] = bool(summary.get("stop")) or event.is_stopped()
|
||||
return summary
|
||||
except Exception as exc:
|
||||
await self._handle_error(
|
||||
@@ -154,16 +207,35 @@ class HandlerDispatcher:
|
||||
) -> dict[str, Any]:
|
||||
trigger = loaded.descriptor.trigger
|
||||
if isinstance(trigger, CommandTrigger):
|
||||
param_specs = loaded.descriptor.param_specs
|
||||
for command_name in [trigger.command, *trigger.aliases]:
|
||||
remainder = self._match_command_name(event.text, command_name)
|
||||
if remainder is not None:
|
||||
return self._build_command_args(loaded.callable, remainder)
|
||||
if param_specs:
|
||||
return self._build_command_args(param_specs, remainder)
|
||||
return self._build_command_args(
|
||||
[
|
||||
ParamSpec(name=name, type="str")
|
||||
for name in self._legacy_arg_parameter_names(
|
||||
loaded.callable
|
||||
)
|
||||
],
|
||||
remainder,
|
||||
)
|
||||
return {}
|
||||
if isinstance(trigger, MessageTrigger) and trigger.regex:
|
||||
match = re.search(trigger.regex, event.text)
|
||||
if match is None:
|
||||
return {}
|
||||
return self._build_regex_args(loaded.callable, match)
|
||||
if loaded.descriptor.param_specs:
|
||||
return self._build_regex_args(loaded.descriptor.param_specs, match)
|
||||
return self._build_regex_args(
|
||||
[
|
||||
ParamSpec(name=name, type="str")
|
||||
for name in self._legacy_arg_parameter_names(loaded.callable)
|
||||
],
|
||||
match,
|
||||
)
|
||||
return {}
|
||||
|
||||
def _build_args(
|
||||
@@ -175,6 +247,7 @@ class HandlerDispatcher:
|
||||
*,
|
||||
plugin_id: str | None = None,
|
||||
handler_ref: str | None = None,
|
||||
schedule_context: ScheduleContext | None = None,
|
||||
) -> list[Any]:
|
||||
"""构建 handler 参数列表。"""
|
||||
from loguru import logger
|
||||
@@ -201,7 +274,9 @@ class HandlerDispatcher:
|
||||
# 1. 优先按类型注解注入
|
||||
param_type = type_hints.get(parameter.name)
|
||||
if param_type is not None:
|
||||
injected = self._inject_by_type(param_type, event, ctx)
|
||||
injected = self._inject_by_type(
|
||||
param_type, event, ctx, schedule_context
|
||||
)
|
||||
|
||||
# 2. Fallback 按名字注入
|
||||
if injected is None:
|
||||
@@ -209,6 +284,8 @@ class HandlerDispatcher:
|
||||
injected = event
|
||||
elif parameter.name in {"ctx", "context"}:
|
||||
injected = ctx
|
||||
elif parameter.name in {"sched", "schedule"}:
|
||||
injected = schedule_context
|
||||
elif parameter.name in args:
|
||||
injected = args[parameter.name]
|
||||
|
||||
@@ -236,7 +313,11 @@ class HandlerDispatcher:
|
||||
return injected_args
|
||||
|
||||
def _inject_by_type(
|
||||
self, param_type: Any, event: MessageEvent, ctx: Context
|
||||
self,
|
||||
param_type: Any,
|
||||
event: MessageEvent,
|
||||
ctx: Context,
|
||||
schedule_context: ScheduleContext | None,
|
||||
) -> Any:
|
||||
"""根据类型注解注入参数。"""
|
||||
# 处理 Optional[Type] 情况
|
||||
@@ -263,6 +344,10 @@ class HandlerDispatcher:
|
||||
isinstance(param_type, type) and issubclass(param_type, Context)
|
||||
):
|
||||
return ctx
|
||||
if param_type is ScheduleContext or (
|
||||
isinstance(param_type, type) and issubclass(param_type, ScheduleContext)
|
||||
):
|
||||
return schedule_context
|
||||
|
||||
return None
|
||||
|
||||
@@ -341,6 +426,23 @@ class HandlerDispatcher:
|
||||
if isinstance(item, dict) and "text" in item:
|
||||
await event.reply(str(item["text"]))
|
||||
return True
|
||||
if isinstance(item, MessageEventResult):
|
||||
chain = item.chain
|
||||
if chain.components:
|
||||
await event.reply_chain(chain)
|
||||
return True
|
||||
return False
|
||||
chain = coerce_message_chain(item)
|
||||
if chain is not None:
|
||||
if chain.components:
|
||||
await event.reply_chain(chain)
|
||||
return True
|
||||
return False
|
||||
if isinstance(item, list) and all(
|
||||
isinstance(component, BaseMessageComponent) for component in item
|
||||
):
|
||||
await event.reply_chain(MessageChain(list(item)))
|
||||
return True
|
||||
# 支持带 text 属性的对象
|
||||
text = getattr(item, "text", None)
|
||||
if isinstance(text, str):
|
||||
@@ -358,27 +460,32 @@ class HandlerDispatcher:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _build_command_args(cls, handler, remainder: str) -> dict[str, Any]:
|
||||
names = cls._legacy_arg_parameter_names(handler)
|
||||
if not names or not remainder:
|
||||
def _build_command_args(
|
||||
cls, param_specs: list[ParamSpec], remainder: str
|
||||
) -> dict[str, Any]:
|
||||
if not param_specs or not remainder:
|
||||
return {}
|
||||
if len(names) == 1:
|
||||
return {names[0]: remainder}
|
||||
if len(param_specs) == 1:
|
||||
return {param_specs[0].name: remainder}
|
||||
parts = cls._split_command_remainder(remainder)
|
||||
return {
|
||||
name: parts[index] for index, name in enumerate(names) if index < len(parts)
|
||||
}
|
||||
values: dict[str, Any] = {}
|
||||
for index, spec in enumerate(param_specs):
|
||||
if index >= len(parts):
|
||||
break
|
||||
if spec.type == "greedy_str":
|
||||
values[spec.name] = " ".join(parts[index:])
|
||||
break
|
||||
values[spec.name] = parts[index]
|
||||
return values
|
||||
|
||||
@classmethod
|
||||
def _build_regex_args(cls, handler, match: re.Match[str]) -> dict[str, Any]:
|
||||
def _build_regex_args(
|
||||
cls, param_specs: list[ParamSpec], match: re.Match[str]
|
||||
) -> dict[str, Any]:
|
||||
named = {
|
||||
key: value for key, value in match.groupdict().items() if value is not None
|
||||
}
|
||||
names = [
|
||||
name
|
||||
for name in cls._legacy_arg_parameter_names(handler)
|
||||
if name not in named
|
||||
]
|
||||
names = [spec.name for spec in param_specs if spec.name not in named]
|
||||
positional = [value for value in match.groups() if value is not None]
|
||||
for index, value in enumerate(positional):
|
||||
if index >= len(names):
|
||||
@@ -386,6 +493,73 @@ class HandlerDispatcher:
|
||||
named[names[index]] = value
|
||||
return named
|
||||
|
||||
@staticmethod
|
||||
def _parse_handler_args(
|
||||
param_specs: list[ParamSpec],
|
||||
args: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
parsed: dict[str, Any] = {}
|
||||
for spec in param_specs:
|
||||
if spec.name not in args:
|
||||
if spec.type == "optional":
|
||||
parsed[spec.name] = None
|
||||
continue
|
||||
if spec.required:
|
||||
raise TypeError(f"缺少参数: {spec.name}")
|
||||
continue
|
||||
parsed[spec.name] = HandlerDispatcher._convert_param(spec, args[spec.name])
|
||||
return parsed
|
||||
|
||||
@staticmethod
|
||||
def _convert_param(spec: ParamSpec, value: Any) -> Any:
|
||||
if spec.type in {"str", "greedy_str"}:
|
||||
return str(value)
|
||||
if spec.type == "int":
|
||||
return int(str(value))
|
||||
if spec.type == "float":
|
||||
return float(str(value))
|
||||
if spec.type == "bool":
|
||||
normalized = str(value).strip().lower()
|
||||
if normalized in {"true", "1", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"false", "0", "no", "off"}:
|
||||
return False
|
||||
raise TypeError(f"无法解析布尔参数 {spec.name}: {value!r}")
|
||||
if spec.type == "optional":
|
||||
if value is None:
|
||||
return None
|
||||
inner = ParamSpec(
|
||||
name=spec.name,
|
||||
type=spec.inner_type or "str",
|
||||
required=False,
|
||||
)
|
||||
return HandlerDispatcher._convert_param(inner, value)
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _run_local_filters(
|
||||
bindings: list[LocalFilterBinding],
|
||||
*,
|
||||
event: MessageEvent,
|
||||
ctx: Context,
|
||||
) -> bool:
|
||||
for binding in bindings:
|
||||
if not binding.evaluate(event=event, ctx=ctx):
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _build_schedule_context(
|
||||
loaded: LoadedHandler,
|
||||
event_payload: dict[str, Any],
|
||||
) -> ScheduleContext | None:
|
||||
if not isinstance(loaded.descriptor.trigger, ScheduleTrigger):
|
||||
return None
|
||||
try:
|
||||
return ScheduleContext.from_payload(event_payload)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _split_command_remainder(remainder: str) -> list[str]:
|
||||
try:
|
||||
@@ -464,231 +638,4 @@ class HandlerDispatcher:
|
||||
await Star().on_error(exc, event, ctx)
|
||||
|
||||
|
||||
class CapabilityDispatcher:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plugin_id: str,
|
||||
peer,
|
||||
capabilities: list[LoadedCapability],
|
||||
) -> None:
|
||||
self._plugin_id = plugin_id
|
||||
self._peer = peer
|
||||
self._capabilities = {item.descriptor.name: item for item in capabilities}
|
||||
self._active: dict[str, tuple[asyncio.Task[Any], CancelToken]] = {}
|
||||
|
||||
async def invoke(
|
||||
self,
|
||||
message,
|
||||
cancel_token: CancelToken,
|
||||
) -> dict[str, Any] | StreamExecution:
|
||||
loaded = self._capabilities.get(message.capability)
|
||||
if loaded is None:
|
||||
raise LookupError(f"capability not found: {message.capability}")
|
||||
|
||||
plugin_id = self._resolve_plugin_id(loaded)
|
||||
ctx = Context(
|
||||
peer=self._peer,
|
||||
plugin_id=plugin_id,
|
||||
cancel_token=cancel_token,
|
||||
)
|
||||
|
||||
with caller_plugin_scope(plugin_id):
|
||||
task = asyncio.create_task(
|
||||
self._run_capability(
|
||||
loaded,
|
||||
payload=dict(message.input),
|
||||
ctx=ctx,
|
||||
cancel_token=cancel_token,
|
||||
stream=bool(message.stream),
|
||||
)
|
||||
)
|
||||
self._active[message.id] = (task, cancel_token)
|
||||
try:
|
||||
return await task
|
||||
finally:
|
||||
self._active.pop(message.id, None)
|
||||
|
||||
def _resolve_plugin_id(self, loaded: LoadedCapability) -> str:
|
||||
if loaded.plugin_id:
|
||||
return loaded.plugin_id
|
||||
return self._plugin_id
|
||||
|
||||
async def cancel(self, request_id: str) -> None:
|
||||
active = self._active.get(request_id)
|
||||
if active is None:
|
||||
return
|
||||
task, cancel_token = active
|
||||
cancel_token.cancel()
|
||||
task.cancel()
|
||||
|
||||
async def _run_capability(
|
||||
self,
|
||||
loaded: LoadedCapability,
|
||||
*,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
stream: bool,
|
||||
) -> dict[str, Any] | StreamExecution:
|
||||
result = loaded.callable(
|
||||
*self._build_args(
|
||||
loaded.callable,
|
||||
payload,
|
||||
ctx,
|
||||
cancel_token,
|
||||
plugin_id=self._resolve_plugin_id(loaded),
|
||||
capability_name=loaded.descriptor.name,
|
||||
)
|
||||
)
|
||||
if stream:
|
||||
if inspect.isasyncgen(result):
|
||||
return StreamExecution(
|
||||
iterator=self._iterate_generator(result),
|
||||
finalize=lambda chunks: {"items": chunks},
|
||||
)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
if isinstance(result, StreamExecution):
|
||||
return result
|
||||
raise AstrBotError.protocol_error(
|
||||
"stream=true 的插件 capability 必须返回 async generator 或 StreamExecution"
|
||||
)
|
||||
|
||||
if inspect.isasyncgen(result):
|
||||
raise AstrBotError.protocol_error(
|
||||
"stream=false 的插件 capability 不能返回 async generator"
|
||||
)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
return self._normalize_output(result)
|
||||
|
||||
def _build_args(
|
||||
self,
|
||||
handler,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
*,
|
||||
plugin_id: str | None = None,
|
||||
capability_name: str | None = None,
|
||||
) -> list[Any]:
|
||||
signature = inspect.signature(handler)
|
||||
args: list[Any] = []
|
||||
|
||||
type_hints: dict[str, Any] = {}
|
||||
try:
|
||||
type_hints = get_type_hints(handler)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for parameter in signature.parameters.values():
|
||||
if parameter.kind not in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
):
|
||||
continue
|
||||
|
||||
injected = None
|
||||
param_type = type_hints.get(parameter.name)
|
||||
if param_type is not None:
|
||||
injected = self._inject_by_type(param_type, payload, ctx, cancel_token)
|
||||
|
||||
if injected is None:
|
||||
if parameter.name in {"ctx", "context"}:
|
||||
injected = ctx
|
||||
elif parameter.name in {"payload", "input", "data"}:
|
||||
injected = payload
|
||||
elif parameter.name in {"cancel_token", "token"}:
|
||||
injected = cancel_token
|
||||
|
||||
if injected is None:
|
||||
if parameter.default is not parameter.empty:
|
||||
continue
|
||||
raise TypeError(
|
||||
self._format_capability_injection_error(
|
||||
handler=handler,
|
||||
parameter_name=parameter.name,
|
||||
plugin_id=plugin_id,
|
||||
capability_name=capability_name,
|
||||
payload=payload,
|
||||
)
|
||||
)
|
||||
args.append(injected)
|
||||
|
||||
return args
|
||||
|
||||
def _inject_by_type(
|
||||
self,
|
||||
param_type: Any,
|
||||
payload: dict[str, Any],
|
||||
ctx: Context,
|
||||
cancel_token: CancelToken,
|
||||
) -> Any:
|
||||
origin = typing.get_origin(param_type)
|
||||
if origin is typing.Union:
|
||||
type_args = typing.get_args(param_type)
|
||||
non_none_types = [item for item in type_args if item is not type(None)]
|
||||
if len(non_none_types) == 1:
|
||||
param_type = non_none_types[0]
|
||||
origin = typing.get_origin(param_type)
|
||||
|
||||
if param_type is Context or (
|
||||
isinstance(param_type, type) and issubclass(param_type, Context)
|
||||
):
|
||||
return ctx
|
||||
if param_type is CancelToken or (
|
||||
isinstance(param_type, type) and issubclass(param_type, CancelToken)
|
||||
):
|
||||
return cancel_token
|
||||
if param_type is dict or origin is dict:
|
||||
return payload
|
||||
return None
|
||||
|
||||
def _format_capability_injection_error(
|
||||
self,
|
||||
*,
|
||||
handler,
|
||||
parameter_name: str,
|
||||
plugin_id: str | None,
|
||||
capability_name: str | None,
|
||||
payload: dict[str, Any],
|
||||
) -> str:
|
||||
plugin_text = plugin_id or self._plugin_id
|
||||
target = capability_name or getattr(handler, "__name__", "<anonymous>")
|
||||
payload_keys = sorted(str(key) for key in payload.keys())
|
||||
payload_keys_text = ", ".join(payload_keys) if payload_keys else "<none>"
|
||||
return (
|
||||
f"插件 '{plugin_text}' 的 capability '{target}' 参数注入失败:"
|
||||
f"必填参数 '{parameter_name}' 无法注入。"
|
||||
f"签名: {getattr(handler, '__name__', '<anonymous>')}"
|
||||
f"{HandlerDispatcher._callable_signature(handler)}。"
|
||||
"当前支持按类型注入 Context / CancelToken / dict,"
|
||||
"按参数名注入 ctx / context / payload / input / data / cancel_token / token,"
|
||||
f"以及 payload 中现有键:{payload_keys_text}。"
|
||||
)
|
||||
|
||||
async def _iterate_generator(
|
||||
self,
|
||||
generator: AsyncIterator[Any],
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
async for item in generator:
|
||||
yield self._normalize_chunk(item)
|
||||
|
||||
def _normalize_chunk(self, item: Any) -> dict[str, Any]:
|
||||
output = self._normalize_output(item)
|
||||
if output:
|
||||
return output
|
||||
return {"ok": True}
|
||||
|
||||
def _normalize_output(self, result: Any) -> dict[str, Any]:
|
||||
if result is None:
|
||||
return {}
|
||||
if isinstance(result, dict):
|
||||
return result
|
||||
model_dump = getattr(result, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
dumped = model_dump()
|
||||
if isinstance(dumped, dict):
|
||||
return dumped
|
||||
raise AstrBotError.invalid_input("插件 capability 必须返回 dict 或可序列化对象")
|
||||
__all__ = ["CapabilityDispatcher", "HandlerDispatcher"]
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"]
|
||||
@@ -0,0 +1,239 @@
|
||||
"""Session-based conversational flow management.
|
||||
|
||||
本模块实现会话等待器 (session_waiter),用于构建多轮对话流程。
|
||||
|
||||
核心组件:
|
||||
- SessionController: 控制会话生命周期,支持超时管理、会话保持、历史记录
|
||||
- SessionWaiterManager: 管理活跃的会话等待器,处理事件分发和注册/注销
|
||||
- @session_waiter 装饰器: 将普通 handler 转换为会话式 handler
|
||||
|
||||
使用场景:
|
||||
当需要在用户首次触发后继续监听后续消息(如分步表单、问答游戏),
|
||||
可使用 @session_waiter 装饰器自动管理会话状态和超时。
|
||||
|
||||
注意事项:
|
||||
在当前桥接设计中,不应在普通 SDK handler 内直接 await session_waiter,
|
||||
这会导致首次 dispatch 保持打开直到下一条消息到达。
|
||||
如需非阻塞的会话等待,应从后台任务启动或添加显式的调度/恢复机制。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .events import MessageEvent
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SessionController:
|
||||
future: asyncio.Future[Any] = field(default_factory=asyncio.Future)
|
||||
current_event: asyncio.Event | None = None
|
||||
ts: float | None = None
|
||||
timeout: float | None = None
|
||||
history_chains: list[list[dict[str, Any]]] = field(default_factory=list)
|
||||
|
||||
def stop(self, error: Exception | None = None) -> None:
|
||||
if self.future.done():
|
||||
return
|
||||
if error is not None:
|
||||
self.future.set_exception(error)
|
||||
else:
|
||||
self.future.set_result(None)
|
||||
|
||||
def keep(self, timeout: float = 0, reset_timeout: bool = False) -> None:
|
||||
new_ts = time.time()
|
||||
if reset_timeout:
|
||||
if timeout <= 0:
|
||||
self.stop()
|
||||
return
|
||||
else:
|
||||
assert self.timeout is not None
|
||||
assert self.ts is not None
|
||||
left_timeout = self.timeout - (new_ts - self.ts)
|
||||
timeout = left_timeout + timeout
|
||||
if timeout <= 0:
|
||||
self.stop()
|
||||
return
|
||||
|
||||
if self.current_event and not self.current_event.is_set():
|
||||
self.current_event.set()
|
||||
|
||||
current_event = asyncio.Event()
|
||||
self.current_event = current_event
|
||||
self.ts = new_ts
|
||||
self.timeout = timeout
|
||||
asyncio.create_task(self._holding(current_event, timeout))
|
||||
|
||||
async def _holding(self, event: asyncio.Event, timeout: float) -> None:
|
||||
try:
|
||||
await asyncio.wait_for(event.wait(), timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
self.stop(exc)
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
def get_history_chains(self) -> list[list[dict[str, Any]]]:
|
||||
return list(self.history_chains)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _WaiterEntry:
|
||||
session_key: str
|
||||
handler: Callable[[SessionController, MessageEvent], Awaitable[Any]]
|
||||
controller: SessionController
|
||||
record_history_chains: bool
|
||||
|
||||
|
||||
class SessionWaiterManager:
|
||||
def __init__(self, *, plugin_id: str, peer) -> None:
|
||||
self._plugin_id = plugin_id
|
||||
self._peer = peer
|
||||
self._entries: dict[str, _WaiterEntry] = {}
|
||||
self._locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
async def register(
|
||||
self,
|
||||
*,
|
||||
event: MessageEvent,
|
||||
handler: Callable[[SessionController, MessageEvent], Awaitable[Any]],
|
||||
timeout: int,
|
||||
record_history_chains: bool,
|
||||
) -> Any:
|
||||
if event._context is None:
|
||||
raise RuntimeError("session_waiter requires runtime context")
|
||||
session_key = event.unified_msg_origin
|
||||
entry = _WaiterEntry(
|
||||
session_key=session_key,
|
||||
handler=handler,
|
||||
controller=SessionController(),
|
||||
record_history_chains=record_history_chains,
|
||||
)
|
||||
replaced = session_key in self._entries
|
||||
self._entries[session_key] = entry
|
||||
self._locks.setdefault(session_key, asyncio.Lock())
|
||||
if replaced:
|
||||
logger.warning(
|
||||
"Session waiter replaced: plugin_id=%s session_key=%s",
|
||||
self._plugin_id,
|
||||
session_key,
|
||||
)
|
||||
await self._peer.invoke(
|
||||
"system.session_waiter.register",
|
||||
{"session_key": session_key},
|
||||
)
|
||||
entry.controller.keep(timeout, reset_timeout=True)
|
||||
try:
|
||||
return await entry.controller.future
|
||||
finally:
|
||||
await self.unregister(session_key)
|
||||
|
||||
async def unregister(self, session_key: str) -> None:
|
||||
self._entries.pop(session_key, None)
|
||||
self._locks.pop(session_key, None)
|
||||
try:
|
||||
await self._peer.invoke(
|
||||
"system.session_waiter.unregister",
|
||||
{"session_key": session_key},
|
||||
)
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Failed to unregister session waiter: plugin_id=%s session_key=%s",
|
||||
self._plugin_id,
|
||||
session_key,
|
||||
)
|
||||
|
||||
def has_waiter(self, event: MessageEvent) -> bool:
|
||||
return event.unified_msg_origin in self._entries
|
||||
|
||||
async def dispatch(self, event: MessageEvent) -> dict[str, Any]:
|
||||
session_key = event.unified_msg_origin
|
||||
entry = self._entries.get(session_key)
|
||||
if entry is None:
|
||||
return {"sent_message": False, "stop": False, "call_llm": False}
|
||||
lock = self._locks.setdefault(session_key, asyncio.Lock())
|
||||
async with lock:
|
||||
if entry.record_history_chains:
|
||||
chain = []
|
||||
raw_chain = (
|
||||
event.raw.get("chain") if isinstance(event.raw, dict) else None
|
||||
)
|
||||
if isinstance(raw_chain, list):
|
||||
chain = [dict(item) for item in raw_chain if isinstance(item, dict)]
|
||||
entry.controller.history_chains.append(chain)
|
||||
await entry.handler(entry.controller, event)
|
||||
return {
|
||||
"sent_message": False,
|
||||
"stop": event.is_stopped(),
|
||||
"call_llm": False,
|
||||
}
|
||||
|
||||
|
||||
def session_waiter(
|
||||
timeout: int = 30,
|
||||
*,
|
||||
record_history_chains: bool = False,
|
||||
):
|
||||
def decorator(
|
||||
func: Callable[[SessionController, MessageEvent], Awaitable[Any]],
|
||||
):
|
||||
async def wrapper(*args, **kwargs):
|
||||
owner = None
|
||||
event: MessageEvent | None = None
|
||||
trailing_args = ()
|
||||
if args and isinstance(args[0], MessageEvent):
|
||||
event = args[0]
|
||||
trailing_args = args[1:]
|
||||
elif len(args) >= 2 and isinstance(args[1], MessageEvent):
|
||||
owner = args[0]
|
||||
event = args[1]
|
||||
trailing_args = args[2:]
|
||||
if event is None:
|
||||
raise RuntimeError("session_waiter requires a MessageEvent argument")
|
||||
if event._context is None:
|
||||
raise RuntimeError("session_waiter requires runtime context")
|
||||
manager = getattr(event._context.peer, "_session_waiter_manager", None)
|
||||
if manager is None:
|
||||
raise RuntimeError("session_waiter manager is unavailable")
|
||||
|
||||
if owner is None:
|
||||
|
||||
async def bound_handler(
|
||||
controller: SessionController,
|
||||
waiter_event: MessageEvent,
|
||||
) -> Any:
|
||||
return await func(
|
||||
controller,
|
||||
waiter_event,
|
||||
*trailing_args,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
|
||||
async def bound_handler(
|
||||
controller: SessionController,
|
||||
waiter_event: MessageEvent,
|
||||
) -> Any:
|
||||
return await func(
|
||||
owner,
|
||||
controller,
|
||||
waiter_event,
|
||||
*trailing_args,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
return await manager.register(
|
||||
event=event,
|
||||
handler=bound_handler,
|
||||
timeout=timeout,
|
||||
record_history_chains=record_history_chains,
|
||||
)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
+88
-495
@@ -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):
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user