137 lines
5.0 KiB
Python
137 lines
5.0 KiB
Python
|
|
"""可命名组件的泛型注册表基类。
|
|||
|
|
|
|||
|
|
提供注册、解析、排序、缓存与生命周期管理的通用实现,
|
|||
|
|
供 InboundMiddlewareRegistry、OutboundMiddlewareRegistry、SecurityCheckerRegistry 复用。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
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
|