✨ feat(auth): 添加群组和机器人唤醒命令支持,优化权限检查逻辑 (#2113)

* ✨ feat(auth): 添加群组和机器人唤醒命令支持,优化权限检查逻辑
✨ feat(llm): 增加额外请求头配置,改进API适配器请求头处理

* 🚨 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>
This commit is contained in:
ManyManyTomato
2026-03-26 17:18:22 +08:00
committed by GitHub
co-authored by ATTomatoo pre-commit-ci[bot]
parent 65b125dd07
commit 6da4f27b12
13 changed files with 214 additions and 18 deletions
@@ -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}",
@@ -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(
+108 -1
View File
@@ -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,
@@ -131,6 +131,8 @@ async def _():
"admin_level",
"plugin_type",
"is_show",
"ignore_prompt",
"ignore_statistics",
]
)
)
+6
View File
@@ -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)
+21 -4
View File
@@ -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,
+6 -4
View File
@@ -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:
+2 -2
View File
@@ -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] = []
+3
View File
@@ -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"
+1
View File
@@ -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)
+19 -1
View File
@@ -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()
+13 -1
View File
@@ -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)
+4
View File
@@ -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):