feat: 新增 Thinking 整流器处理跨 Provider 签名错误 (#115)

当 Provider A 生成的 thinking 块被发送到 Provider B 时,签名验证会失败。
本次更新实现了自动整流机制,在遇到签名错误时自动清洗 thinking 块后重试。

主要更改:
- 新增 ThinkingRectifier 整流器,移除 thinking 块和 signature 字段
- 新增 ThinkingSignatureException 异常类型
- ErrorClassifier 新增 Thinking 错误模式检测
- FallbackOrchestrator 支持整流后在当前候选重试
- Handler 层传递 request_body_ref 容器支持请求体动态修改
- Usage API 新增 has_rectified 字段标识整流过的请求
- 新增 THINKING_RECTIFIER_ENABLED 配置项控制功能开关

其他改进:
- CacheAwareScheduler 支持 exact/convertible 候选分组排序
- StreamProcessor 预读阶段新增格式转换试验
- ProviderAPIKey.api_formats 改为可空(None 表示支持所有格式)
- Dockerfile 修复 entrypoint.sh 换行符问题

Closes #115
Co-Authored-By: FredericMN <FredericMN@users.noreply.github.com>
This commit is contained in:
fawney19
2026-01-22 14:17:59 +08:00
parent cc5db20c58
commit af1828dd32
19 changed files with 1211 additions and 45 deletions

View File

@@ -45,6 +45,7 @@ from src.core.exceptions import (
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
from src.models.database import (
@@ -298,10 +299,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = self.allowed_api_formats[0]
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
# 创建类型安全的流式上下文
ctx = StreamContext(model=model, api_format=api_format)
ctx.request_id = self.request_id
ctx.client_api_format = api_format.value if hasattr(api_format, "value") else str(api_format)
ctx.client_api_format = (
api_format.value if hasattr(api_format, "value") else str(api_format)
)
# 创建更新状态的回调闭包(可以访问 ctx
def update_streaming_status() -> None:
@@ -327,7 +334,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider,
endpoint,
key,
original_request_body,
request_body_ref["body"], # 使用容器中的请求体
original_headers,
query_params,
candidate,
@@ -356,6 +363,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_id=self.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
request_body_ref=request_body_ref, # 传递容器引用
)
# 更新上下文
@@ -364,6 +372,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.provider_id = provider_id
ctx.endpoint_id = endpoint_id
ctx.key_id = key_id
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
# 创建遥测记录器
telemetry_recorder = StreamTelemetryRecorder(
@@ -405,6 +415,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
background=background_tasks,
)
except ThinkingSignatureException as e:
# Thinking 签名错误orchestrator 层已处理整流重试但仍失败
# 记录 original_request_body客户端原始请求便于排查问题根因
self._log_request_error("流式请求失败(签名错误)", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
raise
except Exception as e:
self._log_request_error("流式请求失败", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
@@ -622,7 +639,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
response_time_ms = self.elapsed_ms()
status_code = 503
if isinstance(error, ProviderAuthException):
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
@@ -668,6 +687,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = self.allowed_api_formats[0]
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
# 用于跟踪的变量
provider_name: Optional[str] = None
response_json: Optional[Dict[str, Any]] = None
@@ -709,9 +732,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 应用模型映射
if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(original_request_body, mapped_model)
request_body = self.apply_mapped_model(request_body_ref["body"], mapped_model)
else:
request_body = dict(original_request_body)
request_body = dict(request_body_ref["body"])
# 跨格式:先做请求体转换(严格模式,失败触发 failover
if needs_conversion:
@@ -874,7 +897,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
response_json = converter_registry.convert_response_strict(
response_json,
provider_api_format,
str(api_format),
client_api_format,
)
return response_json if isinstance(response_json, dict) else {}
@@ -900,6 +923,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_func=sync_request_func,
request_id=self.request_id,
capability_requirements=capability_requirements or None,
request_body_ref=request_body_ref, # 传递容器引用
)
provider_name = actual_provider_name
@@ -965,6 +989,23 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
headers=client_response_headers,
)
except ThinkingSignatureException as e:
# Thinking 签名错误orchestrator 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = self.elapsed_ms()
actual_request_body = provider_request_body or original_request_body
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=e.status_code or 400,
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
)
raise
except Exception as e:
response_time_ms = self.elapsed_ms()

View File

@@ -56,6 +56,7 @@ from src.core.exceptions import (
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
from src.database import get_db
@@ -303,7 +304,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
"""
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
# 注意:使用 original_request_body因为整流只修改 messages不影响 model 字段
model = self.extract_model_from_request(original_request_body, path_params)
# 创建流上下文
@@ -327,7 +333,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider,
endpoint,
key,
original_request_body,
request_body_ref["body"], # 使用容器中的请求体
original_headers,
query_params,
candidate,
@@ -356,6 +362,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
request_id=self.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
request_body_ref=request_body_ref, # 传递容器引用
)
# 更新上下文(确保 provider 信息已设置,用于 streaming 状态更新)
@@ -368,6 +375,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx.endpoint_id = endpoint_id
if not ctx.key_id:
ctx.key_id = key_id
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
# 创建后台任务记录统计
background_tasks = BackgroundTasks()
@@ -396,6 +405,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
background=background_tasks,
)
except ThinkingSignatureException as e:
# Thinking 签名错误orchestrator 层已处理整流重试但仍失败
# 记录 original_request_body客户端原始请求便于排查问题根因
self._log_request_error("流式请求失败(签名错误)", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
raise
except Exception as e:
self._log_request_error("流式请求失败", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
@@ -856,7 +872,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
"上游服务返回了非预期的响应格式",
provider_name=str(provider.name),
upstream_status=200,
upstream_response=normalized_line[:500] if normalized_line else "(empty)",
upstream_response=(
normalized_line[:500] if normalized_line else "(empty)"
),
)
if not normalized_line or normalized_line.startswith(":"):
@@ -1350,7 +1368,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 记录失败的 Usage但使用已收到的预估 token 信息(来自 message_start
# 这样即使请求中断,也能记录预估成本
# 失败时返回给客户端的是 JSON 错误响应,如果没有设置则使用默认值
client_response_headers = ctx.client_response_headers or {"content-type": "application/json"}
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
await bg_telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
@@ -1385,11 +1405,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers.update({
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
})
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name,
@@ -1439,11 +1461,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 计算候选自身的 TTFB
candidate_first_byte_time_ms: Optional[int] = None
if ctx.first_byte_time_ms is not None:
candidate_first_byte_time_ms = RequestCandidateService.calculate_candidate_ttfb(
db=bg_db,
candidate_id=ctx.attempt_id,
request_start_time=self.start_time,
global_first_byte_time_ms=ctx.first_byte_time_ms,
candidate_first_byte_time_ms = (
RequestCandidateService.calculate_candidate_ttfb(
db=bg_db,
candidate_id=ctx.attempt_id,
request_start_time=self.start_time,
global_first_byte_time_ms=ctx.first_byte_time_ms,
)
)
# 根据状态码决定是成功还是失败
@@ -1451,7 +1475,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 503 = 服务不可用(如流中断),应标记为失败
if ctx.status_code and ctx.status_code >= 400:
# 请求链路追踪使用 upstream_response原始响应回退到 error_message友好消息
trace_error_message = ctx.upstream_response or ctx.error_message or f"HTTP {ctx.status_code}"
trace_error_message = (
ctx.upstream_response or ctx.error_message or f"HTTP {ctx.status_code}"
)
extra_data = {
"stream_completed": False,
"chunk_count": ctx.chunk_count,
@@ -1476,6 +1502,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
"chunk_count": ctx.chunk_count,
"data_count": ctx.data_count,
}
if ctx.rectified:
extra_data["rectified"] = True
if candidate_first_byte_time_ms is not None:
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
RequestCandidateService.mark_candidate_success(
@@ -1504,7 +1532,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
response_time_ms = int((time.time() - self.start_time) * 1000)
status_code = 503
if isinstance(error, ProviderAuthException):
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
@@ -1574,6 +1604,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
@@ -1595,9 +1629,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(original_request_body, mapped_model)
request_body = self.apply_mapped_model(request_body_ref["body"], mapped_model)
else:
request_body = original_request_body
request_body = dict(request_body_ref["body"])
# 准备发送给 Provider 的请求体(子类可覆盖以移除不需要的字段)
request_body = self.prepare_provider_request_body(request_body)
@@ -1749,6 +1783,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
request_func=sync_request_func,
request_id=self.request_id,
capability_requirements=capability_requirements or None,
request_body_ref=request_body_ref, # 传递容器引用
)
provider_name = actual_provider_name
@@ -1830,6 +1865,24 @@ class CliMessageHandlerBase(BaseMessageHandler):
headers=client_response_headers,
)
except ThinkingSignatureException as e:
# Thinking 签名错误orchestrator 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = int((time.time() - sync_start_time) * 1000)
actual_request_body = provider_request_body or original_request_body
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=e.status_code or 400,
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
)
raise
except Exception as e:
response_time_ms = int((time.time() - sync_start_time) * 1000)

View File

@@ -81,6 +81,9 @@ class StreamContext:
# Provider 响应元数据CLI handler 需要)
response_metadata: Dict[str, Any] = field(default_factory=dict)
# 整流标记Thinking Rectifier
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
# 流式处理统计
data_count: int = 0
chunk_count: int = 0

View File

@@ -31,6 +31,7 @@ from src.api.handlers.base.utils import (
)
from src.config.constants import StreamDefaults
from src.config.settings import config
from src.core.api_format import FormatConversionError, converter_registry
from src.core.exceptions import (
EmbeddedErrorException,
ProviderNotAvailableException,
@@ -296,6 +297,31 @@ class StreamProcessor:
error_status=parsed.error_type,
)
# 预读阶段格式转换试验:首字节前可 failover
# 如果需要跨格式转换,对首个有效数据块做试转换
if ctx.needs_conversion and isinstance(data, dict):
client_format = (ctx.client_api_format or "").upper()
provider_format = (ctx.provider_api_format or "").upper()
if client_format and provider_format:
try:
# 试转换:传 state=None不保留状态
# 如果失败触发 failover下一个候选会使用干净的 state
converter_registry.convert_stream_chunk_strict(
data,
provider_format,
client_format,
state=None,
)
except FormatConversionError as conv_err:
# 格式转换失败:抛出异常触发 failover
logger.debug(
f" [{self.request_id}] 预读阶段格式转换试验失败: "
f"Provider={provider.name}, "
f"{provider_format} -> {client_format}, "
f"error={conv_err}"
)
raise
# 预读到有效数据,没有错误,停止预读
should_stop = True
break
@@ -323,7 +349,12 @@ class StreamProcessor:
base_url=endpoint.base_url,
)
except (EmbeddedErrorException, ProviderNotAvailableException, ProviderTimeoutException):
except (
EmbeddedErrorException,
ProviderNotAvailableException,
ProviderTimeoutException,
FormatConversionError,
):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except (OSError, IOError) as e:

View File

@@ -297,6 +297,8 @@ class StreamTelemetryRecorder:
"stream_completed": ctx.is_success(),
"data_count": ctx.data_count,
}
if ctx.rectified:
extra_data["rectified"] = True
if ctx.first_byte_time_ms is not None:
# 计算候选自身的 TTFB
first_byte_time_ms = RequestCandidateService.calculate_candidate_ttfb(