✨ 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
+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):