refactor: 统一请求头处理逻辑到 headers.py

- 新增 src/core/headers.py 集中管理头部处理函数
- 扩展 api_format_metadata.py 添加 extra_headers 和 protected_keys 配置
- 将各 adapter 中重复的头部方法提取到基类,统一调用 headers.py
- 移除 transport.py 中未使用的 build_provider_headers 函数
- 添加 headers 模块的单元测试
This commit is contained in:
fawney19
2026-01-15 22:25:56 +08:00
parent 80887dd2cc
commit ea45918561
17 changed files with 811 additions and 353 deletions

View File

@@ -22,12 +22,13 @@ from abc import abstractmethod
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.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.enums import APIFormat
from src.core.exceptions import (
InvalidRequestException,
ModelNotSupportedException,
@@ -39,6 +40,12 @@ from src.core.exceptions import (
QuotaExceededException,
UpstreamClientException,
)
from src.core.headers import (
build_adapter_base_headers,
build_adapter_headers,
extract_client_api_key,
get_adapter_protected_keys,
)
from src.core.logger import logger
from src.services.billing import calculate_request_cost as _calculate_request_cost
from src.services.request.result import RequestResult
@@ -67,6 +74,14 @@ class ChatAdapterBase(ApiAdapter):
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini"
BILLING_TEMPLATE: str = "claude"
@classmethod
def _get_api_format(cls) -> APIFormat:
"""获取 API 格式枚举,用于调用 headers.py 的统一函数"""
try:
return APIFormat[cls.FORMAT_ID]
except KeyError:
return APIFormat.OPENAI # 默认回退
# 子类可以配置的特殊方法用于check_endpoint
@classmethod
def build_endpoint_url(cls, base_url: str) -> str:
@@ -76,18 +91,20 @@ class ChatAdapterBase(ApiAdapter):
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
"""构建基础请求头,子类可以覆盖以自定义认证头"""
# 默认实现Bearer token认证
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
"""构建基础请求头,使用统一的 headers.py 实现"""
return build_adapter_base_headers(cls._get_api_format(), api_key)
@classmethod
def get_protected_header_keys(cls) -> tuple:
"""返回不应被extra_headers覆盖的头部key子类可以覆盖"""
# 默认保护认证相关头部
return ("authorization", "content-type")
"""返回不应被extra_headers覆盖的头部key使用统一的 headers.py 实现"""
return get_adapter_protected_keys(cls._get_api_format())
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
"""构建完整请求头(包含 extra_headers使用统一的 headers.py 实现"""
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
@@ -95,6 +112,10 @@ class ChatAdapterBase(ApiAdapter):
# 默认实现:直接使用请求数据
return request_data.copy()
def extract_api_key(self, request: Request) -> Optional[str]:
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@@ -626,13 +647,11 @@ class ChatAdapterBase(ApiAdapter):
Returns:
测试响应数据
"""
from src.api.handlers.base.endpoint_checker import build_safe_headers, run_endpoint_check
from src.api.handlers.base.endpoint_checker import run_endpoint_check
# 使用子类配置方法构建请求组件
url = cls.build_endpoint_url(base_url)
base_headers = cls.build_base_headers(api_key)
protected_keys = cls.get_protected_header_keys()
headers = build_safe_headers(base_headers, extra_headers, protected_keys)
headers = cls.build_headers_with_extra(api_key, extra_headers)
body = cls.build_request_body(request_data)
# 使用通用的endpoint checker执行请求

View File

@@ -20,12 +20,13 @@ import traceback
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.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.core.enums import APIFormat
from src.core.exceptions import (
InvalidRequestException,
ModelNotSupportedException,
@@ -37,6 +38,12 @@ from src.core.exceptions import (
QuotaExceededException,
UpstreamClientException,
)
from src.core.headers import (
build_adapter_base_headers,
build_adapter_headers,
extract_client_api_key,
get_adapter_protected_keys,
)
from src.core.logger import logger
from src.services.billing import calculate_request_cost as _calculate_request_cost
from src.services.request.result import RequestResult
@@ -68,6 +75,55 @@ class CliAdapterBase(ApiAdapter):
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
# =========================================================================
# API 格式与头部处理 - 使用统一的 headers.py 函数
# =========================================================================
@classmethod
def _get_api_format(cls) -> APIFormat:
"""将 FORMAT_ID 转换为 APIFormat 枚举"""
try:
return APIFormat[cls.FORMAT_ID]
except KeyError:
return APIFormat.OPENAI
def extract_api_key(self, request: Request) -> Optional[str]:
"""
从请求中提取 API 密钥
使用统一的头部处理函数,根据 API 格式自动识别认证头。
"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
"""
构建 CLI API 认证头
使用统一的头部处理函数。
"""
return build_adapter_base_headers(cls._get_api_format(), api_key)
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
"""
构建带额外头部的完整请求头
使用统一的头部处理函数,自动保护关键头部不被覆盖。
"""
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
@classmethod
def get_protected_header_keys(cls) -> tuple[str, ...]:
"""
返回 CLI API 的保护头部 key
使用统一的头部处理函数。
"""
return get_adapter_protected_keys(cls._get_api_format())
async def handle(self, context: ApiRequestContext):
"""处理 CLI API 请求"""
http_request = context.request
@@ -575,20 +631,19 @@ class CliAdapterBase(ApiAdapter):
Returns:
测试响应数据
"""
from src.api.handlers.base.endpoint_checker import build_safe_headers, run_endpoint_check
from src.api.handlers.base.endpoint_checker import run_endpoint_check
# 构建请求组件
url = cls.build_endpoint_url(base_url, request_data, model_name)
base_headers = cls.build_base_headers(api_key)
protected_keys = cls.get_protected_header_keys()
# 添加CLI User-Agent
# 添加 CLI User-Agent 到 extra_headers
cli_user_agent = cls.get_cli_user_agent()
merged_extra = dict(extra_headers) if extra_headers else {}
if cli_user_agent:
base_headers["User-Agent"] = cli_user_agent
protected_keys = tuple(list(protected_keys) + ["user-agent"])
merged_extra["User-Agent"] = cli_user_agent
headers = build_safe_headers(base_headers, extra_headers, protected_keys)
# 使用统一的头部构建函数
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
body = cls.build_request_body(request_data)
# 获取有效的模型名称
@@ -610,7 +665,7 @@ class CliAdapterBase(ApiAdapter):
)
# =========================================================================
# CLI Adapter 配置方法 - 子类应覆盖这些方法而不是整个 check_endpoint
# CLI Adapter 配置方法 - 子类应覆盖这些方法
# =========================================================================
@classmethod
@@ -628,29 +683,6 @@ class CliAdapterBase(ApiAdapter):
"""
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
"""
构建CLI API认证头 - 子类应覆盖
Args:
api_key: API密钥
Returns:
基础认证头部字典
"""
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_base_headers")
@classmethod
def get_protected_header_keys(cls) -> tuple:
"""
返回CLI API的保护头部key - 子类应覆盖
Returns:
保护头部key的元组
"""
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement get_protected_header_keys")
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""

View File

@@ -27,23 +27,11 @@ from collections import defaultdict
import httpx
from src.core.logger import logger
_SENSITIVE_HEADER_KEYS = {
"authorization",
"x-api-key",
"x-goog-api-key",
}
from src.core.headers import CORE_REDACT_HEADERS, merge_headers_with_protection, redact_headers_for_log
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
redacted: Dict[str, str] = {}
for key, value in headers.items():
if key.lower() in _SENSITIVE_HEADER_KEYS:
redacted[key] = "***"
else:
redacted[key] = value
return redacted
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
def _truncate_repr(value: Any, limit: int = 1200) -> str:
@@ -64,14 +52,7 @@ def build_safe_headers(
"""
合并 extra_headers但防止覆盖 protected_keys大小写不敏感
"""
headers = dict(base_headers)
if not extra_headers:
return headers
protected = {k.lower() for k in protected_keys}
safe_headers = {k: v for k, v in extra_headers.items() if k.lower() not in protected}
headers.update(safe_headers)
return headers
return merge_headers_with_protection(base_headers, extra_headers, set(protected_keys))
async def run_endpoint_check(

View File

@@ -17,27 +17,14 @@ from abc import ABC, abstractmethod
from typing import Any, Dict, FrozenSet, Optional, Tuple
from src.core.crypto import crypto_service
from src.core.headers import HeaderBuilder, UPSTREAM_DROP_HEADERS
# ==============================================================================
# 统一的头部配置常量
# ==============================================================================
# 敏感头部 - 透传时需要清理(黑名单)
# 这些头部要么包含认证信息,要么由代理层重新生成
SENSITIVE_HEADERS: FrozenSet[str] = frozenset(
{
"authorization",
"x-api-key",
"x-goog-api-key", # Gemini API 认证头
"host",
"content-length",
"transfer-encoding",
"connection",
# 不透传 accept-encoding让 httpx 自己协商压缩格式
# 避免客户端请求 brotli/zstd 但 httpx 不支持解压的问题
"accept-encoding",
}
)
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
# ==============================================================================
@@ -140,8 +127,6 @@ class PassthroughRequestBuilder(RequestBuilder):
"""
from src.core.api_format_metadata import get_auth_config, resolve_api_format
headers: Dict[str, str] = {}
# 1. 根据 API 格式自动设置认证头
decrypted_key = crypto_service.decrypt(key.api_key)
api_format = getattr(endpoint, "api_format", None)
@@ -150,32 +135,32 @@ class PassthroughRequestBuilder(RequestBuilder):
get_auth_config(resolved_format) if resolved_format else ("Authorization", "bearer")
)
if auth_type == "bearer":
headers[auth_header] = f"Bearer {decrypted_key}"
else:
headers[auth_header] = decrypted_key
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
protected_keys = {auth_header.lower(), "content-type"}
# 2. 添加 endpoint 配置的额外头部
if endpoint.headers:
headers.update(endpoint.headers)
builder = HeaderBuilder()
# 3. 透传原始头部(排除敏感头部 - 黑名单模式)
# 2. 透传原始头部(排除敏感头部 - 黑名单模式)
if original_headers:
for name, value in original_headers.items():
lower_name = name.lower()
# 跳过敏感头部
if lower_name in SENSITIVE_HEADERS:
if name.lower() in SENSITIVE_HEADERS:
continue
builder.add(name, value)
headers[name] = value
# 3. 添加 endpoint 配置的额外头部(不能覆盖认证头/Content-Type
if endpoint.headers:
builder.add_protected(endpoint.headers, protected_keys)
# 4. 添加额外头部
if extra_headers:
headers.update(extra_headers)
builder.add_many(extra_headers)
# 5. 确保有 Content-Type
if "Content-Type" not in headers and "content-type" not in headers:
# 5. 设置认证头(最高优先级)
builder.add(auth_header, auth_value)
# 6. 确保有 Content-Type
headers = builder.build()
if not any(k.lower() == "content-type" for k in headers):
headers["Content-Type"] = "application/json"
return headers

View File

@@ -6,6 +6,7 @@ import json
from typing import Any, Dict, Optional
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
from src.core.headers import filter_response_headers
from src.core.logger import logger
@@ -94,25 +95,6 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
return headers
_PROXY_RESPONSE_HEADER_BLOCKLIST = frozenset(
{
# Body-dependent headers: 我们会重编码响应体JSONResponse / SSE不能透传上游值
"content-length",
"content-encoding",
"transfer-encoding",
"content-type",
# Hop-by-hop headers (RFC 7230)
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"upgrade",
}
)
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
"""
过滤上游响应头中不应透传给客户端的字段。
@@ -123,9 +105,7 @@ def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str
如果透传上游的 `content-length/content-encoding/...`,会导致客户端解码失败或等待更多字节。
"""
if not headers:
return {}
return {k: v for k, v in headers.items() if k.lower() not in _PROXY_RESPONSE_HEADER_BLOCKLIST}
return filter_response_headers(headers)
def check_html_response(line: str) -> bool: