mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 增强 Gemini 认证支持与 URL 敏感信息脱敏
- 新增 extract_client_api_key_with_query 支持从 URL query 参数提取 API Key - Gemini 格式遵循 Google SDK 行为:query 参数优先于 header - 添加 redact_url_for_log 函数对日志中的敏感 URL 参数脱敏 - 上游请求始终使用 header 认证,清除 Gemini query 中的 key 参数 - 放宽 API Key 最小长度限制至 3 字符 - 完善认证方式相关代码注释
This commit is contained in:
@@ -56,7 +56,7 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.provider.transport import build_provider_url, redact_url_for_log
|
||||
|
||||
|
||||
class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
@@ -492,7 +492,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
|
||||
# 构建请求
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
@@ -749,7 +749,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
|
||||
# 构建请求
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_hdrs = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
@@ -775,7 +775,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
f" [{self.request_id}] 发送非流式请求: Provider={provider.name}, "
|
||||
f"模型={model} -> {mapped_model or '无映射'}"
|
||||
)
|
||||
logger.debug(f" [{self.request_id}] 请求URL: {url}")
|
||||
logger.debug(f" [{self.request_id}] 请求URL: {redact_url_for_log(url)}")
|
||||
logger.debug(
|
||||
f" [{self.request_id}] 请求体stream字段: {provider_payload.get('stream', 'N/A')}"
|
||||
)
|
||||
|
||||
@@ -474,6 +474,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
@@ -1638,6 +1639,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
|
||||
@@ -72,6 +72,15 @@ class RequestBuilder(ABC):
|
||||
"""
|
||||
构建完整的请求(请求体 + 请求头)
|
||||
|
||||
Args:
|
||||
original_body: 原始请求体
|
||||
original_headers: 原始请求头
|
||||
endpoint: 端点配置
|
||||
key: Provider API Key
|
||||
mapped_model: 映射后的模型名
|
||||
is_stream: 是否为流式请求
|
||||
extra_headers: 额外请求头
|
||||
|
||||
Returns:
|
||||
Tuple[payload, headers]
|
||||
"""
|
||||
@@ -124,6 +133,12 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||
|
||||
Args:
|
||||
original_headers: 原始请求头
|
||||
endpoint: 端点配置
|
||||
key: Provider API Key
|
||||
extra_headers: 额外请求头
|
||||
"""
|
||||
from src.core.api_format import get_auth_config, resolve_api_format
|
||||
|
||||
@@ -136,6 +151,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
)
|
||||
|
||||
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
|
||||
# 认证头始终受保护,防止 header_rules 覆盖
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
builder = HeaderBuilder()
|
||||
@@ -147,7 +163,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
continue
|
||||
builder.add(name, value)
|
||||
|
||||
# 3. 应用 endpoint 的请求头规则
|
||||
# 3. 应用 endpoint 的请求头规则(认证头受保护,无法通过 rules 设置)
|
||||
header_rules = getattr(endpoint, "header_rules", None)
|
||||
if header_rules:
|
||||
builder.apply_rules(header_rules, protected_keys)
|
||||
@@ -156,7 +172,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
if extra_headers:
|
||||
builder.add_many(extra_headers)
|
||||
|
||||
# 5. 设置认证头(最高优先级)
|
||||
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
|
||||
builder.add(auth_header, auth_value)
|
||||
|
||||
# 6. 确保有 Content-Type
|
||||
|
||||
@@ -7,13 +7,15 @@ Gemini Chat Adapter
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.core.api_format import extract_client_api_key_with_query
|
||||
from src.core.logger import logger
|
||||
from src.models.gemini import GeminiRequest
|
||||
from src.services.provider.transport import redact_url_for_log
|
||||
|
||||
|
||||
@register_adapter
|
||||
@@ -40,6 +42,20 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
super().__init__(allowed_api_formats or ["GEMINI"])
|
||||
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
||||
|
||||
优先级(与 Google SDK 行为一致):
|
||||
1. URL 参数 ?key=
|
||||
2. x-goog-api-key 请求头
|
||||
"""
|
||||
return extract_client_api_key_with_query(
|
||||
dict(request.headers),
|
||||
dict(request.query_params),
|
||||
self._get_api_format(),
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
@@ -171,7 +187,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
|
||||
try:
|
||||
response = await client.get(models_url, headers=headers)
|
||||
logger.debug(f"Gemini models request to {models_url}: status={response.status_code}")
|
||||
logger.debug(f"Gemini models request to {redact_url_for_log(models_url)}: status={response.status_code}")
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
if "models" in data:
|
||||
@@ -189,11 +205,13 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
else:
|
||||
error_body = response.text[:500] if response.text else "(empty)"
|
||||
error_msg = f"HTTP {response.status_code}: {error_body}"
|
||||
logger.warning(f"Gemini models request to {models_url} failed: {error_msg}")
|
||||
logger.warning(f"Gemini models request to {redact_url_for_log(models_url)} failed: {error_msg}")
|
||||
return [], error_msg
|
||||
except Exception as e:
|
||||
error_msg = f"Request error: {str(e)}"
|
||||
logger.warning(f"Failed to fetch Gemini models from {models_url}: {e}")
|
||||
# 异常信息可能包含带 key 参数的 URL,需要脱敏
|
||||
sanitized_error = redact_url_for_log(str(e))
|
||||
error_msg = f"Request error: {sanitized_error}"
|
||||
logger.warning(f"Failed to fetch Gemini models from {redact_url_for_log(models_url)}: {sanitized_error}")
|
||||
return [], error_msg
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -7,11 +7,13 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
|
||||
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
|
||||
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
|
||||
from src.api.handlers.gemini.adapter import GeminiChatAdapter
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import extract_client_api_key_with_query
|
||||
|
||||
|
||||
@register_cli_adapter
|
||||
@@ -36,6 +38,20 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
|
||||
|
||||
优先级(与 Google SDK 行为一致):
|
||||
1. URL 参数 ?key=
|
||||
2. x-goog-api-key 请求头
|
||||
"""
|
||||
return extract_client_api_key_with_query(
|
||||
dict(request.headers),
|
||||
dict(request.query_params),
|
||||
self._get_api_format(),
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
|
||||
Reference in New Issue
Block a user