mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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,
|
User,
|
||||||
)
|
)
|
||||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
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):
|
class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||||
@@ -492,7 +492,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||||
request_body = self.prepare_provider_request_body(request_body)
|
request_body = self.prepare_provider_request_body(request_body)
|
||||||
|
|
||||||
# 构建请求
|
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||||
provider_payload, provider_headers = self._request_builder.build(
|
provider_payload, provider_headers = self._request_builder.build(
|
||||||
request_body,
|
request_body,
|
||||||
original_headers,
|
original_headers,
|
||||||
@@ -749,7 +749,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||||
request_body = self.prepare_provider_request_body(request_body)
|
request_body = self.prepare_provider_request_body(request_body)
|
||||||
|
|
||||||
# 构建请求
|
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||||
provider_payload, provider_hdrs = self._request_builder.build(
|
provider_payload, provider_hdrs = self._request_builder.build(
|
||||||
request_body,
|
request_body,
|
||||||
original_headers,
|
original_headers,
|
||||||
@@ -775,7 +775,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
f" [{self.request_id}] 发送非流式请求: Provider={provider.name}, "
|
f" [{self.request_id}] 发送非流式请求: Provider={provider.name}, "
|
||||||
f"模型={model} -> {mapped_model or '无映射'}"
|
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(
|
logger.debug(
|
||||||
f" [{self.request_id}] 请求体stream字段: {provider_payload.get('stream', 'N/A')}"
|
f" [{self.request_id}] 请求体stream字段: {provider_payload.get('stream', 'N/A')}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -474,6 +474,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 使用 RequestBuilder 构建请求体和请求头
|
# 使用 RequestBuilder 构建请求体和请求头
|
||||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||||
|
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||||
provider_payload, provider_headers = self._request_builder.build(
|
provider_payload, provider_headers = self._request_builder.build(
|
||||||
request_body,
|
request_body,
|
||||||
original_headers,
|
original_headers,
|
||||||
@@ -1638,6 +1639,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 使用 RequestBuilder 构建请求体和请求头
|
# 使用 RequestBuilder 构建请求体和请求头
|
||||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||||
|
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||||
provider_payload, provider_headers = self._request_builder.build(
|
provider_payload, provider_headers = self._request_builder.build(
|
||||||
request_body,
|
request_body,
|
||||||
original_headers,
|
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:
|
Returns:
|
||||||
Tuple[payload, headers]
|
Tuple[payload, headers]
|
||||||
"""
|
"""
|
||||||
@@ -124,6 +133,12 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
) -> Dict[str, str]:
|
) -> 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
|
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
|
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
|
||||||
|
# 认证头始终受保护,防止 header_rules 覆盖
|
||||||
protected_keys = {auth_header.lower(), "content-type"}
|
protected_keys = {auth_header.lower(), "content-type"}
|
||||||
|
|
||||||
builder = HeaderBuilder()
|
builder = HeaderBuilder()
|
||||||
@@ -147,7 +163,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
continue
|
continue
|
||||||
builder.add(name, value)
|
builder.add(name, value)
|
||||||
|
|
||||||
# 3. 应用 endpoint 的请求头规则
|
# 3. 应用 endpoint 的请求头规则(认证头受保护,无法通过 rules 设置)
|
||||||
header_rules = getattr(endpoint, "header_rules", None)
|
header_rules = getattr(endpoint, "header_rules", None)
|
||||||
if header_rules:
|
if header_rules:
|
||||||
builder.apply_rules(header_rules, protected_keys)
|
builder.apply_rules(header_rules, protected_keys)
|
||||||
@@ -156,7 +172,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
if extra_headers:
|
if extra_headers:
|
||||||
builder.add_many(extra_headers)
|
builder.add_many(extra_headers)
|
||||||
|
|
||||||
# 5. 设置认证头(最高优先级)
|
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
|
||||||
builder.add(auth_header, auth_value)
|
builder.add(auth_header, auth_value)
|
||||||
|
|
||||||
# 6. 确保有 Content-Type
|
# 6. 确保有 Content-Type
|
||||||
|
|||||||
@@ -7,13 +7,15 @@ Gemini Chat Adapter
|
|||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any, Dict, Optional, Tuple, Type
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
|
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
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.core.logger import logger
|
||||||
from src.models.gemini import GeminiRequest
|
from src.models.gemini import GeminiRequest
|
||||||
|
from src.services.provider.transport import redact_url_for_log
|
||||||
|
|
||||||
|
|
||||||
@register_adapter
|
@register_adapter
|
||||||
@@ -40,6 +42,20 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
super().__init__(allowed_api_formats or ["GEMINI"])
|
super().__init__(allowed_api_formats or ["GEMINI"])
|
||||||
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
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(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
@@ -171,7 +187,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
response = await client.get(models_url, headers=headers)
|
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:
|
if response.status_code == 200:
|
||||||
data = response.json()
|
data = response.json()
|
||||||
if "models" in data:
|
if "models" in data:
|
||||||
@@ -189,11 +205,13 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
else:
|
else:
|
||||||
error_body = response.text[:500] if response.text else "(empty)"
|
error_body = response.text[:500] if response.text else "(empty)"
|
||||||
error_msg = f"HTTP {response.status_code}: {error_body}"
|
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
|
return [], error_msg
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_msg = f"Request error: {str(e)}"
|
# 异常信息可能包含带 key 参数的 URL,需要脱敏
|
||||||
logger.warning(f"Failed to fetch Gemini models from {models_url}: {e}")
|
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
|
return [], error_msg
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -7,11 +7,13 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
|||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any, Dict, Optional, Tuple, Type
|
||||||
|
|
||||||
import httpx
|
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_adapter_base import CliAdapterBase, register_cli_adapter
|
||||||
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
|
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
|
||||||
from src.api.handlers.gemini.adapter import GeminiChatAdapter
|
from src.api.handlers.gemini.adapter import GeminiChatAdapter
|
||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
|
from src.core.api_format import extract_client_api_key_with_query
|
||||||
|
|
||||||
|
|
||||||
@register_cli_adapter
|
@register_cli_adapter
|
||||||
@@ -36,6 +38,20 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||||
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
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(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
|
|||||||
@@ -76,7 +76,8 @@ def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
|
|||||||
Returns:
|
Returns:
|
||||||
(api_format, api_key) 元组
|
(api_format, api_key) 元组
|
||||||
"""
|
"""
|
||||||
return detect_format_and_key_from_starlette(request)
|
format_name, api_key, _auth_method = detect_format_and_key_from_starlette(request)
|
||||||
|
return format_name, api_key
|
||||||
|
|
||||||
|
|
||||||
def _get_formats_for_api(api_format: str) -> list[str]:
|
def _get_formats_for_api(api_format: str) -> list[str]:
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ from src.core.api_format.headers import (
|
|||||||
build_upstream_headers,
|
build_upstream_headers,
|
||||||
detect_capabilities,
|
detect_capabilities,
|
||||||
extract_client_api_key,
|
extract_client_api_key,
|
||||||
|
extract_client_api_key_with_query,
|
||||||
extract_set_headers_from_rules,
|
extract_set_headers_from_rules,
|
||||||
filter_response_headers,
|
filter_response_headers,
|
||||||
get_adapter_protected_keys,
|
get_adapter_protected_keys,
|
||||||
@@ -113,6 +114,7 @@ __all__ = [
|
|||||||
"normalize_headers",
|
"normalize_headers",
|
||||||
"get_header_value",
|
"get_header_value",
|
||||||
"extract_client_api_key",
|
"extract_client_api_key",
|
||||||
|
"extract_client_api_key_with_query",
|
||||||
"detect_capabilities",
|
"detect_capabilities",
|
||||||
"HeaderBuilder",
|
"HeaderBuilder",
|
||||||
"build_upstream_headers",
|
"build_upstream_headers",
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ def _extract_api_key_by_definition(
|
|||||||
headers: Dict[str, str],
|
headers: Dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]],
|
query_params: Optional[Dict[str, str]],
|
||||||
definition: ApiFormatDefinition,
|
definition: ApiFormatDefinition,
|
||||||
) -> Optional[str]:
|
) -> Tuple[Optional[str], str]:
|
||||||
"""
|
"""
|
||||||
根据格式定义从请求中提取 API Key
|
根据格式定义从请求中提取 API Key
|
||||||
|
|
||||||
@@ -29,32 +29,44 @@ def _extract_api_key_by_definition(
|
|||||||
definition: API 格式定义
|
definition: API 格式定义
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
提取到的 API Key,或 None
|
(api_key, auth_method) 元组:
|
||||||
|
- api_key: 提取到的 API Key,或 None
|
||||||
|
- auth_method: 认证方式 ("header" 或 "query")
|
||||||
"""
|
"""
|
||||||
auth_header = definition.auth_header.lower()
|
auth_header = definition.auth_header.lower()
|
||||||
auth_type = definition.auth_type
|
auth_type = definition.auth_type
|
||||||
|
|
||||||
|
# Gemini 格式:query 参数优先(与 Google SDK 行为一致)
|
||||||
|
if definition.api_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||||
|
# 1. 优先检查 ?key= 参数
|
||||||
|
query_key = query_params.get("key") if query_params else None
|
||||||
|
if query_key:
|
||||||
|
return query_key, "query"
|
||||||
|
# 2. 再检查 x-goog-api-key 请求头
|
||||||
|
header_value = headers.get(auth_header)
|
||||||
|
if header_value:
|
||||||
|
return header_value, "header"
|
||||||
|
return None, "header"
|
||||||
|
|
||||||
|
# 其他格式:从 header 提取
|
||||||
header_value = headers.get(auth_header)
|
header_value = headers.get(auth_header)
|
||||||
if not header_value:
|
if not header_value:
|
||||||
# Gemini 还支持 ?key= 参数
|
return None, "header"
|
||||||
if definition.api_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
|
||||||
return query_params.get("key") if query_params else None
|
|
||||||
return None
|
|
||||||
|
|
||||||
if auth_type == "bearer":
|
if auth_type == "bearer":
|
||||||
# Bearer token: "Bearer xxx"
|
# Bearer token: "Bearer xxx"
|
||||||
if header_value.lower().startswith("bearer "):
|
if header_value.lower().startswith("bearer "):
|
||||||
return header_value[7:].strip()
|
return header_value[7:].strip(), "header"
|
||||||
return None
|
return None, "header"
|
||||||
else:
|
else:
|
||||||
# header 类型: 直接使用值
|
# header 类型: 直接使用值
|
||||||
return header_value
|
return header_value, "header"
|
||||||
|
|
||||||
|
|
||||||
def detect_format_from_request(
|
def detect_format_from_request(
|
||||||
headers: Dict[str, str],
|
headers: Dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: Optional[Dict[str, str]] = None,
|
||||||
) -> Tuple[APIFormat, Optional[str]]:
|
) -> Tuple[APIFormat, Optional[str], str]:
|
||||||
"""
|
"""
|
||||||
从请求头检测 API 格式和 API Key
|
从请求头检测 API 格式和 API Key
|
||||||
|
|
||||||
@@ -68,33 +80,35 @@ def detect_format_from_request(
|
|||||||
query_params: 查询参数字典(可选)
|
query_params: 查询参数字典(可选)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(APIFormat, api_key) 元组
|
(APIFormat, api_key, auth_method) 元组
|
||||||
|
- auth_method: 认证方式 ("header" 或 "query")
|
||||||
"""
|
"""
|
||||||
# Claude: x-api-key + anthropic-version (必须同时存在)
|
# Claude: x-api-key + anthropic-version (必须同时存在)
|
||||||
claude_def = API_FORMAT_DEFINITIONS[APIFormat.CLAUDE]
|
claude_def = API_FORMAT_DEFINITIONS[APIFormat.CLAUDE]
|
||||||
claude_key = _extract_api_key_by_definition(headers, query_params, claude_def)
|
claude_key, claude_auth_method = _extract_api_key_by_definition(headers, query_params, claude_def)
|
||||||
if claude_key and headers.get("anthropic-version"):
|
if claude_key and headers.get("anthropic-version"):
|
||||||
return APIFormat.CLAUDE, claude_key
|
return APIFormat.CLAUDE, claude_key, claude_auth_method
|
||||||
|
|
||||||
# Gemini: x-goog-api-key (header 类型) 或 ?key=
|
# Gemini: x-goog-api-key (header 类型) 或 ?key=
|
||||||
gemini_def = API_FORMAT_DEFINITIONS[APIFormat.GEMINI]
|
gemini_def = API_FORMAT_DEFINITIONS[APIFormat.GEMINI]
|
||||||
gemini_key = _extract_api_key_by_definition(headers, query_params, gemini_def)
|
gemini_key, gemini_auth_method = _extract_api_key_by_definition(headers, query_params, gemini_def)
|
||||||
if gemini_key:
|
if gemini_key:
|
||||||
return APIFormat.GEMINI, gemini_key
|
return APIFormat.GEMINI, gemini_key, gemini_auth_method
|
||||||
|
|
||||||
# OpenAI: Authorization: Bearer (默认)
|
# OpenAI: Authorization: Bearer (默认)
|
||||||
# 注意: 如果只有 x-api-key 但没有 anthropic-version,也走 OpenAI 格式
|
# 注意: 如果只有 x-api-key 但没有 anthropic-version,也走 OpenAI 格式
|
||||||
openai_def = API_FORMAT_DEFINITIONS[APIFormat.OPENAI]
|
openai_def = API_FORMAT_DEFINITIONS[APIFormat.OPENAI]
|
||||||
openai_key = _extract_api_key_by_definition(headers, query_params, openai_def)
|
openai_key, openai_auth_method = _extract_api_key_by_definition(headers, query_params, openai_def)
|
||||||
# 如果 OpenAI 格式没有 key,但有 x-api-key,也用它(兼容)
|
# 如果 OpenAI 格式没有 key,但有 x-api-key,也用它(兼容)
|
||||||
if not openai_key and claude_key:
|
if not openai_key and claude_key:
|
||||||
openai_key = claude_key
|
openai_key = claude_key
|
||||||
return APIFormat.OPENAI, openai_key
|
openai_auth_method = claude_auth_method
|
||||||
|
return APIFormat.OPENAI, openai_key, openai_auth_method
|
||||||
|
|
||||||
|
|
||||||
def detect_format_and_key_from_starlette(
|
def detect_format_and_key_from_starlette(
|
||||||
request: "Request",
|
request: "Request",
|
||||||
) -> Tuple[str, Optional[str]]:
|
) -> Tuple[str, Optional[str], str]:
|
||||||
"""
|
"""
|
||||||
从 Starlette Request 对象检测 API 格式和 API Key
|
从 Starlette Request 对象检测 API 格式和 API Key
|
||||||
|
|
||||||
@@ -104,17 +118,19 @@ def detect_format_and_key_from_starlette(
|
|||||||
request: Starlette Request 对象
|
request: Starlette Request 对象
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(format_name, api_key) 元组,format_name 为小写字符串
|
(format_name, api_key, auth_method) 元组
|
||||||
|
- format_name: 为小写字符串
|
||||||
|
- auth_method: 认证方式 ("header" 或 "query")
|
||||||
"""
|
"""
|
||||||
# 规范化 headers 为小写
|
# 规范化 headers 为小写
|
||||||
headers = {k.lower(): v for k, v in request.headers.items()}
|
headers = {k.lower(): v for k, v in request.headers.items()}
|
||||||
query_params = dict(request.query_params)
|
query_params = dict(request.query_params)
|
||||||
|
|
||||||
api_format, api_key = detect_format_from_request(headers, query_params)
|
api_format, api_key, auth_method = detect_format_from_request(headers, query_params)
|
||||||
|
|
||||||
# 返回小写格式名
|
# 返回小写格式名
|
||||||
format_name = api_format.value.lower()
|
format_name = api_format.value.lower()
|
||||||
return format_name, api_key
|
return format_name, api_key, auth_method
|
||||||
|
|
||||||
|
|
||||||
def detect_format_from_response(
|
def detect_format_from_response(
|
||||||
|
|||||||
@@ -146,6 +146,38 @@ def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Op
|
|||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def extract_client_api_key_with_query(
|
||||||
|
headers: Dict[str, str],
|
||||||
|
query_params: Optional[Dict[str, str]],
|
||||||
|
api_format: APIFormat,
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
从客户端请求头或 URL 参数提取 API Key
|
||||||
|
|
||||||
|
Gemini 格式优先级(与 Google SDK 行为一致):
|
||||||
|
1. URL 参数 ?key=
|
||||||
|
2. x-goog-api-key 请求头
|
||||||
|
|
||||||
|
其他格式仅从请求头提取。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
headers: 原始请求头(自动处理大小写)
|
||||||
|
query_params: URL 查询参数
|
||||||
|
api_format: API 格式
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
提取的 API Key,未找到返回 None
|
||||||
|
"""
|
||||||
|
# Gemini 格式:query 参数优先
|
||||||
|
if api_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||||
|
query_key = query_params.get("key") if query_params else None
|
||||||
|
if query_key:
|
||||||
|
return query_key
|
||||||
|
|
||||||
|
# 其他格式或 Gemini header 方式:使用现有逻辑
|
||||||
|
return extract_client_api_key(headers, api_format)
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# 能力需求检测
|
# 能力需求检测
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -342,7 +374,7 @@ def build_upstream_headers(
|
|||||||
if extra_headers:
|
if extra_headers:
|
||||||
builder.add_many(extra_headers)
|
builder.add_many(extra_headers)
|
||||||
|
|
||||||
# 4. 设置认证头(最高优先级)
|
# 4. 设置认证头(最高优先级,上游始终使用 header 认证)
|
||||||
builder.add(auth_header, auth_value)
|
builder.add(auth_header, auth_value)
|
||||||
|
|
||||||
# 5. 确保 Content-Type
|
# 5. 确保 Content-Type
|
||||||
|
|||||||
@@ -273,8 +273,8 @@ class CreateAPIKeyRequest(BaseModel):
|
|||||||
v = v.strip()
|
v = v.strip()
|
||||||
|
|
||||||
# 检查最小长度
|
# 检查最小长度
|
||||||
if len(v) < 10:
|
if len(v) < 3:
|
||||||
raise ValueError("API Key 长度不能少于 10 个字符")
|
raise ValueError("API Key 长度不能少于 3 个字符")
|
||||||
|
|
||||||
# 检查危险字符(不应包含 SQL 注入字符)
|
# 检查危险字符(不应包含 SQL 注入字符)
|
||||||
dangerous_chars = ["'", '"', ";", "--", "/*", "*/", "<", ">"]
|
dangerous_chars = ["'", '"', ";", "--", "/*", "*/", "<", ">"]
|
||||||
|
|||||||
@@ -382,8 +382,8 @@ class EndpointAPIKeyUpdate(BaseModel):
|
|||||||
return v
|
return v
|
||||||
|
|
||||||
v = v.strip()
|
v = v.strip()
|
||||||
if len(v) < 10:
|
if len(v) < 3:
|
||||||
raise ValueError("API Key 长度不能少于 10 个字符")
|
raise ValueError("API Key 长度不能少于 3 个字符")
|
||||||
|
|
||||||
dangerous_chars = ["'", '"', ";", "--", "/*", "*/", "<", ">"]
|
dangerous_chars = ["'", '"', ";", "--", "/*", "*/", "<", ">"]
|
||||||
for char in dangerous_chars:
|
for char in dangerous_chars:
|
||||||
|
|||||||
@@ -3,8 +3,10 @@
|
|||||||
|
|
||||||
负责:
|
负责:
|
||||||
- 根据 API 格式或端点配置生成请求 URL
|
- 根据 API 格式或端点配置生成请求 URL
|
||||||
|
- URL 脱敏(用于日志记录)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
@@ -14,6 +16,29 @@ from src.core.logger import logger
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.models.database import ProviderEndpoint
|
from src.models.database import ProviderEndpoint
|
||||||
|
|
||||||
|
|
||||||
|
# URL 中需要脱敏的查询参数(正则模式)
|
||||||
|
_SENSITIVE_QUERY_PARAMS_PATTERN = re.compile(
|
||||||
|
r"([?&])(key|api_key|apikey|token|secret|password|credential)=([^&]*)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def redact_url_for_log(url: str) -> str:
|
||||||
|
"""
|
||||||
|
对 URL 中的敏感查询参数进行脱敏,用于日志记录
|
||||||
|
|
||||||
|
将 ?key=xxx 替换为 ?key=***
|
||||||
|
|
||||||
|
Args:
|
||||||
|
url: 原始 URL
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
脱敏后的 URL
|
||||||
|
"""
|
||||||
|
return _SENSITIVE_QUERY_PARAMS_PATTERN.sub(r"\1\2=***", url)
|
||||||
|
|
||||||
|
|
||||||
def _normalize_base_url(base_url: str, path: str) -> str:
|
def _normalize_base_url(base_url: str, path: str) -> str:
|
||||||
"""
|
"""
|
||||||
规范化 base_url,去除末尾的斜杠和可能与 path 重复的版本前缀。
|
规范化 base_url,去除末尾的斜杠和可能与 path 重复的版本前缀。
|
||||||
@@ -96,9 +121,17 @@ def build_provider_url(
|
|||||||
base = _normalize_base_url(endpoint.base_url, path) # type: ignore[arg-type]
|
base = _normalize_base_url(endpoint.base_url, path) # type: ignore[arg-type]
|
||||||
url = f"{base}{path}"
|
url = f"{base}{path}"
|
||||||
|
|
||||||
|
# 合并查询参数
|
||||||
|
effective_query_params = dict(query_params) if query_params else {}
|
||||||
|
|
||||||
|
# Gemini 格式下清除可能存在的 key 参数(避免客户端传入的认证信息泄露到上游)
|
||||||
|
# 上游认证始终使用 header 方式,不使用 URL 参数
|
||||||
|
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||||
|
effective_query_params.pop("key", None)
|
||||||
|
|
||||||
# 添加查询参数
|
# 添加查询参数
|
||||||
if query_params:
|
if effective_query_params:
|
||||||
query_string = urlencode(query_params, doseq=True)
|
query_string = urlencode(effective_query_params, doseq=True)
|
||||||
if query_string:
|
if query_string:
|
||||||
url = f"{url}?{query_string}"
|
url = f"{url}?{query_string}"
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user