mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
✨ 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:
co-authored by
ATTomatoo
pre-commit-ci[bot]
parent
65b125dd07
commit
6da4f27b12
@@ -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(
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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] = []
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user