ForcePilot/backend/package/yuxi/channel/plugins/registry_base.py
Kris bab30f2715
Some checks failed
Deploy VitePress site to Pages / build (push) Has been cancelled
Ruff Format Check / Ruff Format & Lint (push) Has been cancelled
Deploy VitePress site to Pages / Deploy (push) Has been cancelled
feat:0715
2026-07-15 12:30:58 +08:00

137 lines
5.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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