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