mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-11 06:39:59 +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
+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