From 6da4f27b12cf5814dccca48e575abe85e0fcea78 Mon Sep 17 00:00:00 2001 From: ManyManyTomato <93612024+ATTomatoo@users.noreply.github.com> Date: Thu, 26 Mar 2026 17:18:22 +0800 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20feat(auth):=20=E6=B7=BB=E5=8A=A0?= =?UTF-8?q?=E7=BE=A4=E7=BB=84=E5=92=8C=E6=9C=BA=E5=99=A8=E4=BA=BA=E5=94=A4?= =?UTF-8?q?=E9=86=92=E5=91=BD=E4=BB=A4=E6=94=AF=E6=8C=81=EF=BC=8C=E4=BC=98?= =?UTF-8?q?=E5=8C=96=E6=9D=83=E9=99=90=E6=A3=80=E6=9F=A5=E9=80=BB=E8=BE=91?= =?UTF-8?q?=20(#2113)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * ✨ feat(auth): 添加群组和机器人唤醒命令支持,优化权限检查逻辑 ✨ feat(llm): 增加额外请求头配置,改进API适配器请求头处理 * :rotating_light: auto fix by pre-commit hooks * ``` fix(auth): 优化bot权限验证逻辑并改进错误提示 - 将bot存在性检查与状态检查分离,提供更精确的错误信息 - 修复当bot为None时的状态访问问题 - 移除不必要的注释,保持代码简洁 - 优化日志记录的位置和条件判断 ``` --------- Co-authored-by: ATTomatoo <1126160939@qq.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- .../builtin_plugins/hooks/auth/auth_bot.py | 12 +- .../builtin_plugins/hooks/auth/auth_group.py | 22 +++- zhenxun/builtin_plugins/hooks/auth_checker.py | 109 +++++++++++++++++- zhenxun/builtin_plugins/init/init_plugin.py | 2 + zhenxun/services/cache/runtime_cache.py | 6 + zhenxun/services/llm/adapters/base.py | 25 +++- zhenxun/services/llm/adapters/gemini.py | 10 +- zhenxun/services/llm/adapters/openai.py | 4 +- zhenxun/services/llm/config/providers.py | 3 + zhenxun/services/llm/manager.py | 1 + zhenxun/services/llm/service.py | 20 +++- zhenxun/services/llm/tools.py | 14 ++- zhenxun/services/llm/types/models.py | 4 + 13 files changed, 214 insertions(+), 18 deletions(-) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py index e1e5ed86..20e64efb 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py @@ -15,6 +15,7 @@ async def auth_bot( bot_id: str, bot_data: BotConsole | BotSnapshot | None = None, skip_fetch: bool = False, + allow_sleep_bypass: bool = False, ): """bot层面的权限检查 @@ -33,16 +34,19 @@ async def auth_bot( if bot is None and not skip_fetch: bot = await BotMemoryCache.get(bot_id) - if not bot or not bot.status: - raise SkipPluginException("Bot不存在或休眠中阻断权限检测...") + if bot is None: + raise SkipPluginException("Bot不存在,阻断权限检测...") + + if not bot.status and not allow_sleep_bypass: + raise SkipPluginException("Bot休眠中阻断权限检测...") + if CommonUtils.format(plugin.module) in bot.block_plugins: raise SkipPluginException( f"Bot插件 {plugin.name}({plugin.module}) 权限检查结果为关闭..." ) finally: - # 记录执行时间 elapsed = time.time() - start_time - if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查 + if elapsed > WARNING_THRESHOLD: logger.warning( f"auth_bot 耗时: {elapsed:.3f}s, " f"bot_id={bot_id}, plugin={plugin.module}", diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_group.py b/zhenxun/builtin_plugins/hooks/auth/auth_group.py index a6ccf471..e59bc206 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_group.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_group.py @@ -1,3 +1,4 @@ +import re import time from zhenxun.models.group_console import GroupConsole @@ -8,6 +9,25 @@ from zhenxun.services.log import logger from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum from .exception import SkipPluginException +_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE) +_GROUP_WAKE_CANONICAL_PATTERN = re.compile(r"^group-status\s+wake$", re.IGNORECASE) + + +def _is_group_wake_command(plugin: PluginInfo, text: str) -> bool: + if "plugin_switch" not in (plugin.module or ""): + return False + normalized = re.sub(r"\s+", " ", (text or "").strip()) + if not normalized: + return False + if ( + _GROUP_WAKE_PATTERN.match(normalized) is not None + or _GROUP_WAKE_CANONICAL_PATTERN.match(normalized) is not None + ): + return True + # 兼容 to_me 前缀场景:如“真寻 醒来” + tokens = normalized.split(" ") + return len(tokens) == 2 and tokens[-1] == SwitchEnum.ENABLE + async def auth_group( plugin: PluginInfo, @@ -34,7 +54,7 @@ async def auth_group( raise SkipPluginException("群组信息不存在...") if group.level < 0: raise SkipPluginException("群组黑名单, 目标群组群权限权限-1...") - if text.strip() != SwitchEnum.ENABLE and not group.status: + if not _is_group_wake_command(plugin, text) and not group.status: raise SkipPluginException("群组休眠状态...") if plugin.level > group.level: raise SkipPluginException( diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index 03900194..fb18230a 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -136,6 +136,10 @@ _PREFILTER_STATS = { } _PREFILTER_LAST_LOG = 0.0 _CACHE_SWEEP_TASK: asyncio.Task | None = None +_BOT_WAKE_COMMAND_PATTERN = re.compile(r"^bot醒来(?:\s+\S+)?$", re.IGNORECASE) +_BOT_WAKE_CANONICAL_PATTERN = re.compile( + r"^bot_manage\s+bot_switch\s+enable(?:\s+\S+)?$", re.IGNORECASE +) class HookTraceRecorder: @@ -234,6 +238,20 @@ def _normalize_command(command: str) -> str: return text +def _is_bot_wake_command(module: str, text: str | None) -> bool: + if "bot_manage" not in (module or ""): + return False + if not text: + return False + normalized = re.sub(r"\s+", " ", text.strip()) + if not normalized: + return False + return ( + _BOT_WAKE_COMMAND_PATTERN.match(normalized) is not None + or _BOT_WAKE_CANONICAL_PATTERN.match(normalized) is not None + ) + + def _split_command_variants(command: str) -> tuple[str, ...]: text = command.strip() if not text: @@ -350,6 +368,79 @@ def _matcher_module_name(matcher_cls: type[Matcher]) -> str: return (getattr(plugin, "name", "") or "").strip() +def _collect_ai_route_modules(event: Event, state: dict | None = None) -> set[str]: + if state is not None: + cached = state.get("_zx_ai_route_modules") + if isinstance(cached, set): + return cached + + raw_value = getattr(event, "_ai_route_modules", None) + result: set[str] = set() + if isinstance(raw_value, str): + normalized = raw_value.strip() + if normalized: + result.add(normalized) + elif isinstance(raw_value, set | frozenset | list | tuple): + for item in raw_value: + if not isinstance(item, str): + continue + normalized = item.strip() + if normalized: + result.add(normalized) + + if state is not None and result: + state["_zx_ai_route_modules"] = result + return result + + +def _collect_ai_route_heads(event: Event, state: dict | None = None) -> set[str]: + if state is not None: + cached = state.get("_zx_ai_route_heads") + if isinstance(cached, set): + return cached + + raw_value = getattr(event, "_ai_route_heads", None) + result: set[str] = set() + if isinstance(raw_value, str): + normalized = raw_value.strip().casefold() + if normalized: + result.add(normalized) + elif isinstance(raw_value, set | frozenset | list | tuple): + for item in raw_value: + if not isinstance(item, str): + continue + normalized = item.strip().casefold() + if normalized: + result.add(normalized) + + if state is not None and result: + state["_zx_ai_route_heads"] = result + return result + + +def _matcher_matches_ai_route_heads( + matcher_cls: type[Matcher], + ai_route_heads: set[str], +) -> bool: + if not ai_route_heads: + return False + matcher_commands = _extract_matcher_command_literals(matcher_cls) + if not matcher_commands: + return False + for command in matcher_commands: + normalized_command = command.strip().casefold() + if not normalized_command: + continue + for head in ai_route_heads: + if not head: + continue + if _command_matches(head, normalized_command) or _command_matches( + normalized_command, head + ): + return True + return False + + def _is_command_matcher_class(matcher_cls: type[Matcher]) -> bool: if matcher_cls in _MATCHER_COMMAND_TYPE_CACHE: return _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] @@ -602,6 +693,12 @@ async def _check_matcher_prefilter( if not module: return False, None + ai_route_modules = _collect_ai_route_modules(event, state) + ai_route_heads = _collect_ai_route_heads(event, state) + if ai_route_modules and module not in ai_route_modules: + if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads): + return True, "route_miss" + if not _ROUTE_INDEX_READY: await _ensure_route_index() @@ -1425,12 +1522,21 @@ async def auth( # 并行执行所有 hook 检查,并记录执行时间 hooks_start = time.time() + allow_sleep_bypass = _is_bot_wake_command(module, text) # 创建所有 hook 任务 hook_tasks = [] if event_cache is None: hook_tasks.append( - time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_recorder) + time_hook( + auth_bot( + plugin, + bot.self_id, + allow_sleep_bypass=allow_sleep_bypass, + ), + "auth_bot", + hook_recorder, + ) ) else: if bot_timeout: @@ -1443,6 +1549,7 @@ async def auth( bot.self_id, bot_data=bot_data, skip_fetch=True, + allow_sleep_bypass=allow_sleep_bypass, ), "auth_bot", hook_recorder, diff --git a/zhenxun/builtin_plugins/init/init_plugin.py b/zhenxun/builtin_plugins/init/init_plugin.py index 96163d03..3020e318 100644 --- a/zhenxun/builtin_plugins/init/init_plugin.py +++ b/zhenxun/builtin_plugins/init/init_plugin.py @@ -131,6 +131,8 @@ async def _(): "admin_level", "plugin_type", "is_show", + "ignore_prompt", + "ignore_statistics", ] ) ) diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py index 6655e3f0..d656fc8b 100644 --- a/zhenxun/services/cache/runtime_cache.py +++ b/zhenxun/services/cache/runtime_cache.py @@ -586,6 +586,12 @@ class PluginInfoMemoryCache: await cls.ensure_loaded() return cls._by_module.get(module) + @classmethod + async def get_all(cls) -> dict[str, "PluginInfo"]: + if not cls._loaded: + await cls.ensure_loaded() + return dict(cls._by_module) + @classmethod def get_by_module_path(cls, module_path: str) -> "PluginInfo | None": return cls._by_module_path.get(module_path) diff --git a/zhenxun/services/llm/adapters/base.py b/zhenxun/services/llm/adapters/base.py index ca19aaeb..846f9b5d 100644 --- a/zhenxun/services/llm/adapters/base.py +++ b/zhenxun/services/llm/adapters/base.py @@ -202,7 +202,23 @@ class BaseAdapter(ABC): ) return f"{model.api_base.rstrip('/')}{endpoint}" - def get_base_headers(self, api_key: str) -> dict[str, str]: + def _get_provider_extra_headers(self, model: "LLMModel | None") -> dict[str, str]: + if not model: + return {} + raw_headers = getattr(model.provider_config, "extra_headers", None) + if not isinstance(raw_headers, dict): + return {} + headers: dict[str, str] = {} + for key, value in raw_headers.items(): + key_text = str(key).strip() + if not key_text or value is None: + continue + headers[key_text] = str(value) + return headers + + def get_base_headers( + self, api_key: str, model: "LLMModel | None" = None + ) -> dict[str, str]: """获取基础请求头""" from zhenxun.utils.user_agent import get_user_agent @@ -213,6 +229,7 @@ class BaseAdapter(ABC): "Authorization": f"Bearer {api_key}", } ) + headers.update(self._get_provider_extra_headers(model)) return headers def validate_response(self, response_json: dict[str, Any]) -> None: @@ -422,7 +439,7 @@ class OpenAICompatAdapter(BaseAdapter): ) -> RequestData: """准备简单文本生成请求 - OpenAI兼容API的通用实现""" url = self.get_api_url(model, self.get_chat_endpoint(model)) - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) messages = [] if history: @@ -449,7 +466,7 @@ class OpenAICompatAdapter(BaseAdapter): ) -> RequestData: """准备高级请求 - OpenAI兼容格式""" url = self.get_api_url(model, self.get_chat_endpoint(model)) - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) if model.api_type == "openrouter": headers.update( { @@ -523,7 +540,7 @@ class OpenAICompatAdapter(BaseAdapter): ) -> RequestData: """准备嵌入请求 - OpenAI兼容格式""" url = self.get_api_url(model, self.get_embedding_endpoint(model)) - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) body = { "model": model.model_name, diff --git a/zhenxun/services/llm/adapters/gemini.py b/zhenxun/services/llm/adapters/gemini.py index 720bdbe4..8247dbc4 100644 --- a/zhenxun/services/llm/adapters/gemini.py +++ b/zhenxun/services/llm/adapters/gemini.py @@ -39,14 +39,16 @@ class GeminiAdapter(BaseAdapter): def supported_api_types(self) -> list[str]: return ["gemini"] - def get_base_headers(self, api_key: str) -> dict[str, str]: + def get_base_headers( + self, api_key: str, model: "LLMModel | None" = None + ) -> dict[str, str]: """获取基础请求头""" from zhenxun.utils.user_agent import get_user_agent headers = get_user_agent() headers.update({"Content-Type": "application/json"}) headers["x-goog-api-key"] = api_key - + headers.update(self._get_provider_extra_headers(model)) return headers async def prepare_advanced_request( @@ -109,7 +111,7 @@ class GeminiAdapter(BaseAdapter): endpoint = self._get_gemini_endpoint(model, effective_config) url = self.get_api_url(model, endpoint) - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) converter = GeminiMessageConverter() system_instruction_parts: list[dict[str, Any]] | None = None @@ -252,7 +254,7 @@ class GeminiAdapter(BaseAdapter): base_url = model.api_base.rstrip("/") url = f"{base_url}/v1beta/{api_model_name}:batchEmbedContents" - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) requests_payload = [] for text_content in texts: diff --git a/zhenxun/services/llm/adapters/openai.py b/zhenxun/services/llm/adapters/openai.py index 16613524..42873ec5 100644 --- a/zhenxun/services/llm/adapters/openai.py +++ b/zhenxun/services/llm/adapters/openai.py @@ -240,7 +240,7 @@ class OpenAIAdapter(OpenAICompatAdapter): ) -> "RequestData": """根据不同协议策略构建高级请求""" url = self.get_api_url(model, self.get_chat_endpoint(model)) - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) if model.api_type == "openrouter": headers.update( { @@ -463,7 +463,7 @@ class OpenAIImageAdapter(BaseAdapter): ) -> RequestData: _ = tools, tool_choice effective_config = config if config is not None else model._generation_config - headers = self.get_base_headers(api_key) + headers = self.get_base_headers(api_key, model) prompt = "" images_bytes_list: list[bytes] = [] diff --git a/zhenxun/services/llm/config/providers.py b/zhenxun/services/llm/config/providers.py index b6cc0c16..e6869fbb 100644 --- a/zhenxun/services/llm/config/providers.py +++ b/zhenxun/services/llm/config/providers.py @@ -170,6 +170,7 @@ def get_default_providers() -> list[dict[str, Any]]: "api_key": "YOUR_ARK_API_KEY", "api_base": "https://api.deepseek.com", "api_type": "openai", + "extra_headers": {}, "models": [ { "model_name": "deepseek-chat", @@ -296,6 +297,8 @@ def register_llm_configs(): help=( "配置多个 AI 服务提供商及其模型信息。\n" "注意:可以在特定模型配置下添加 'api_type' 以覆盖提供商的全局设置。\n" + "可选:在 provider 下添加 'extra_headers' 传递额外请求头," + "用于 AI 网关鉴权(例如 cf-aig-authorization)。\n" "支持的 api_type 包括:\n" "- 'openai': 标准 OpenAI 格式 (DeepSeek, SiliconFlow, Moonshot 等)\n" "- 'gemini': Google Gemini API\n" diff --git a/zhenxun/services/llm/manager.py b/zhenxun/services/llm/manager.py index e69f9cec..e58d14fe 100644 --- a/zhenxun/services/llm/manager.py +++ b/zhenxun/services/llm/manager.py @@ -328,6 +328,7 @@ async def get_model_instance( openai_compat=provider_config_found.openai_compat, temperature=provider_config_found.temperature, max_tokens=provider_config_found.max_tokens, + extra_headers=provider_config_found.extra_headers, ) shared_http_client = await http_client_manager.get_client(config_for_http_client) diff --git a/zhenxun/services/llm/service.py b/zhenxun/services/llm/service.py index ecefd3a0..cec0c465 100644 --- a/zhenxun/services/llm/service.py +++ b/zhenxun/services/llm/service.py @@ -46,6 +46,24 @@ from .types.capabilities import ModelCapabilities, ModelModality T = TypeVar("T", bound=BaseModel) +def _sanitize_request_headers(headers: dict[str, Any]) -> dict[str, str]: + sanitized: dict[str, str] = {} + sensitive_parts = ("authorization", "token", "api-key", "api_key", "secret") + for key, value in headers.items(): + key_text = str(key) + value_text = str(value) + lowered = key_text.lower() + if any(part in lowered for part in sensitive_parts): + if " " in value_text: + prefix = value_text.split(" ", 1)[0] + sanitized[key_text] = f"{prefix} ***" + else: + sanitized[key_text] = "***" + else: + sanitized[key_text] = value_text + return sanitized + + class LLMContext(BaseModel): """LLM 执行上下文,用于在中间件管道中传递请求状态""" @@ -583,7 +601,7 @@ class NetworkRequestMiddleware(BaseLLMMiddleware): ) logger.debug(f"🔑 API密钥: {masked_key}") logger.debug(f"📡 请求URL: {request_data.url}") - logger.debug(f"📋 请求头: {dict(request_data.headers)}") + logger.debug(f"📋 请求头: {_sanitize_request_headers(request_data.headers)}") if self.model.api_type == "smart": effective_type = self.model._get_effective_api_type() diff --git a/zhenxun/services/llm/tools.py b/zhenxun/services/llm/tools.py index a0d79ee6..7b81377f 100644 --- a/zhenxun/services/llm/tools.py +++ b/zhenxun/services/llm/tools.py @@ -299,8 +299,20 @@ class FunctionExecutable(ToolExecutable): if self._params_model: try: _fields = model_fields(self._params_model) + if isinstance(_fields, dict): + field_names = set(_fields) + else: + field_names = { + name + for field in _fields + for name in ( + getattr(field, "name", None), + getattr(field, "alias", None), + ) + if name + } validation_input = { - key: value for key, value in kwargs.items() if key in _fields + key: value for key, value in kwargs.items() if key in field_names } validated_params = self._params_model(**validation_input) diff --git a/zhenxun/services/llm/types/models.py b/zhenxun/services/llm/types/models.py index e8c79078..d8fa5f66 100644 --- a/zhenxun/services/llm/types/models.py +++ b/zhenxun/services/llm/types/models.py @@ -609,6 +609,10 @@ class ProviderConfig(BaseModel): models: list[ModelDetail] = Field(..., description="支持的模型列表") timeout: int = Field(default=180, description="请求超时时间") proxy: str | None = Field(default=None, description="代理设置") + extra_headers: dict[str, str] | None = Field( + default=None, + description="额外请求头,用于网关鉴权等场景", + ) class LLMToolFunction(BaseModel):