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
|