ForcePilot/backend/package/yuxi/channel/plugins/registry_base.py

137 lines
5.0 KiB
Python
Raw Normal View History

2026-07-15 12:30:58 +08:00
"""可命名组件的泛型注册表基类。
提供注册解析排序缓存与生命周期管理的通用实现
InboundMiddlewareRegistryOutboundMiddlewareRegistrySecurityCheckerRegistry 复用
"""
from __future__ import annotations
import hashlib
import json
from typing import TYPE_CHECKING, Generic, TypeVar
from yuxi.utils.logging_config import logger
if TYPE_CHECKING:
pass
T = TypeVar("T")
class GenericRegistry(Generic[T]):
"""泛型注册表基类:负责注册、解析、排序、缓存与生命周期管理。
子类需定义
- ``_CONFIG_KEY``: 配置中对应列表的键名
- ``_ORDER_FIELD``: 配置条目中的排序字段名 "order" "priority"
- ``_DEFAULT_ORDER_ATTR``: 注册项上默认排序值的属性名 "default_order" "default_priority"
- ``_ITEM_LABEL``: 日志中使用的条目标签 "inbound_middlewares" "security_checker"
"""
_CONFIG_KEY: str
_ORDER_FIELD: str
_DEFAULT_ORDER_ATTR: str
_ITEM_LABEL: str
def __init__(self) -> None:
self._items: dict[str, T] = {}
self._cache: dict[tuple[str, str, str], list[T]] = {}
def register(self, item: T) -> None:
"""注册一个组件实例。"""
name = getattr(item, "name", "")
if not name:
raise ValueError(f"{self._ITEM_LABEL} must have a non-empty name")
self._items[name] = item
def unregister(self, name: str) -> T | None:
"""注销指定名称的组件,返回被移除的实例(若存在)。"""
return self._items.pop(name, None)
def get(self, name: str) -> T | None:
"""按名称获取已注册的组件实例。"""
return self._items.get(name)
def get_item_config(self, config: dict, name: str) -> dict:
"""从账户配置中提取指定组件的专属配置。"""
for entry in config.get(self._CONFIG_KEY, []):
if isinstance(entry, dict) and entry.get("name") == name:
return entry.get("config") or {}
return {}
def resolve_chain(self, config: dict) -> list[T]:
"""根据账户配置解析启用的组件链并按排序字段升序排序。"""
channel_type = config.get("channel_type") or ""
account_id = config.get("account_id") or ""
config_hash = self._make_config_hash(config)
cache_key = (channel_type, account_id, config_hash)
cached = self._cache.get(cache_key)
if cached is not None:
return cached
configured = config.get(self._CONFIG_KEY)
if not configured:
chain = sorted(
self._items.values(),
key=lambda m: getattr(m, self._DEFAULT_ORDER_ATTR),
)
else:
chain = self._resolve_configured_chain(configured)
self._validate_chain_order([getattr(m, "name", "") for m in chain])
self._cache[cache_key] = chain
return chain
def _resolve_configured_chain(self, configured: list[dict]) -> list[T]:
chain: list[tuple[int, T]] = []
for entry in configured:
if not isinstance(entry, dict):
logger.warning("Ignoring invalid %s entry: %s", self._ITEM_LABEL, entry)
continue
name = entry.get("name")
if not isinstance(name, str) or not name:
logger.warning("Ignoring %s entry with missing name: %s", self._ITEM_LABEL, entry)
continue
if not entry.get("enabled", True):
continue
item = self._items.get(name)
if item is None:
logger.warning("Unknown %s %r, ignoring", self._ITEM_LABEL, name)
continue
order = entry.get(self._ORDER_FIELD)
if not isinstance(order, int):
order = getattr(item, self._DEFAULT_ORDER_ATTR)
chain.append((order, item))
chain.sort(key=lambda x: x[0])
return [item for _, item in chain]
def _validate_chain_order(self, names: list[str]) -> None:
"""子类可覆写以校验链的顺序约束。"""
def invalidate(self, channel_type: str, account_id: str) -> None:
"""清除指定渠道账户的缓存。"""
keys_to_remove = [key for key in self._cache if key[0] == channel_type and key[1] == account_id]
for key in keys_to_remove:
self._cache.pop(key, None)
async def start_all(self) -> None:
"""启动所有已注册组件。"""
for item in self._items.values():
start = getattr(item, "start", None)
if start is not None:
await start()
async def stop_all(self) -> None:
"""停止所有已注册组件。"""
for item in self._items.values():
stop = getattr(item, "stop", None)
if stop is not None:
await stop()
@staticmethod
def _make_config_hash(config: dict) -> str:
"""子类需覆写以指定从 config 中提取哪个键计算哈希。"""
raise NotImplementedError