mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
3
.gitignore
vendored
3
.gitignore
vendored
@@ -228,3 +228,6 @@ extracted_*.ts
|
|||||||
|
|
||||||
# Version file (auto-generated by hatch-vcs)
|
# Version file (auto-generated by hatch-vcs)
|
||||||
src/_version.py
|
src/_version.py
|
||||||
|
|
||||||
|
# Analysis folder (third-party code for reference)
|
||||||
|
analysis/
|
||||||
|
|||||||
@@ -141,7 +141,7 @@ RUN mkdir -p /var/log/supervisor /app/logs /app/data
|
|||||||
|
|
||||||
# 入口脚本(启动前执行迁移)
|
# 入口脚本(启动前执行迁移)
|
||||||
COPY entrypoint.sh /entrypoint.sh
|
COPY entrypoint.sh /entrypoint.sh
|
||||||
RUN chmod +x /entrypoint.sh
|
RUN sed -i 's/\r$//' /entrypoint.sh && chmod +x /entrypoint.sh
|
||||||
|
|
||||||
# 环境变量
|
# 环境变量
|
||||||
ENV PYTHONUNBUFFERED=1 \
|
ENV PYTHONUNBUFFERED=1 \
|
||||||
|
|||||||
@@ -156,7 +156,7 @@ async def get_usage_records(
|
|||||||
input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens, total_tokens,
|
input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens, total_tokens,
|
||||||
cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at, is_stream,
|
cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at, is_stream,
|
||||||
input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m, cache_read_price_per_1m,
|
input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m, cache_read_price_per_1m,
|
||||||
status_code, error_message, status, has_fallback, api_format, api_key_name, request_metadata
|
status_code, error_message, status, has_fallback, has_retry, has_rectified, api_format, api_key_name, request_metadata
|
||||||
- `total`: 符合条件的总记录数
|
- `total`: 符合条件的总记录数
|
||||||
- `limit`: 当前分页限制
|
- `limit`: 当前分页限制
|
||||||
- `offset`: 当前分页偏移量
|
- `offset`: 当前分页偏移量
|
||||||
@@ -719,6 +719,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
request_ids = [usage.request_id for usage, _, _, _, _ in records if usage.request_id]
|
request_ids = [usage.request_id for usage, _, _, _, _ in records if usage.request_id]
|
||||||
fallback_map = {}
|
fallback_map = {}
|
||||||
retry_map = {}
|
retry_map = {}
|
||||||
|
rectified_map = {}
|
||||||
if request_ids:
|
if request_ids:
|
||||||
# 查询每个请求的候选执行情况
|
# 查询每个请求的候选执行情况
|
||||||
# 只统计实际执行的候选(success 或 failed),不包括 skipped/pending/available
|
# 只统计实际执行的候选(success 或 failed),不包括 skipped/pending/available
|
||||||
@@ -727,6 +728,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
RequestCandidate.request_id,
|
RequestCandidate.request_id,
|
||||||
RequestCandidate.candidate_index,
|
RequestCandidate.candidate_index,
|
||||||
RequestCandidate.retry_index,
|
RequestCandidate.retry_index,
|
||||||
|
RequestCandidate.extra_data,
|
||||||
)
|
)
|
||||||
.filter(
|
.filter(
|
||||||
RequestCandidate.request_id.in_(request_ids),
|
RequestCandidate.request_id.in_(request_ids),
|
||||||
@@ -736,9 +738,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 按 request_id 分组分析
|
# 按 request_id 分组分析
|
||||||
request_candidates: dict[str, list[tuple[int, int]]] = defaultdict(list)
|
request_candidates: dict[str, list[tuple[int, int, dict]]] = defaultdict(list)
|
||||||
for req_id, candidate_idx, retry_idx in executed_candidates:
|
for req_id, candidate_idx, retry_idx, extra_data in executed_candidates:
|
||||||
request_candidates[req_id].append((candidate_idx, retry_idx))
|
request_candidates[req_id].append((candidate_idx, retry_idx, extra_data or {}))
|
||||||
|
|
||||||
for req_id, candidates in request_candidates.items():
|
for req_id, candidates in request_candidates.items():
|
||||||
# 提取所有不同的 candidate_index
|
# 提取所有不同的 candidate_index
|
||||||
@@ -755,6 +757,11 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
break
|
break
|
||||||
retry_map[req_id] = has_retry
|
retry_map[req_id] = has_retry
|
||||||
|
|
||||||
|
# 检查是否有整流:任意候选的 extra_data 中有 rectified=True
|
||||||
|
rectified_map[req_id] = any(
|
||||||
|
c[2].get("rectified", False) for c in candidates
|
||||||
|
)
|
||||||
|
|
||||||
context.add_audit_metadata(
|
context.add_audit_metadata(
|
||||||
action="usage_records",
|
action="usage_records",
|
||||||
start_date=self.start_date.isoformat() if self.start_date else None,
|
start_date=self.start_date.isoformat() if self.start_date else None,
|
||||||
@@ -834,6 +841,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
"status": usage.status, # 请求状态: pending, streaming, completed, failed
|
"status": usage.status, # 请求状态: pending, streaming, completed, failed
|
||||||
"has_fallback": fallback_map.get(usage.request_id, False),
|
"has_fallback": fallback_map.get(usage.request_id, False),
|
||||||
"has_retry": retry_map.get(usage.request_id, False),
|
"has_retry": retry_map.get(usage.request_id, False),
|
||||||
|
"has_rectified": rectified_map.get(usage.request_id, False),
|
||||||
"api_format": usage.api_format
|
"api_format": usage.api_format
|
||||||
or (endpoint.api_format if endpoint and endpoint.api_format else None),
|
or (endpoint.api_format if endpoint and endpoint.api_format else None),
|
||||||
"api_key_name": provider_api_key.name if provider_api_key else None,
|
"api_key_name": provider_api_key.name if provider_api_key else None,
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ from src.core.exceptions import (
|
|||||||
ProviderNotAvailableException,
|
ProviderNotAvailableException,
|
||||||
ProviderRateLimitException,
|
ProviderRateLimitException,
|
||||||
ProviderTimeoutException,
|
ProviderTimeoutException,
|
||||||
|
ThinkingSignatureException,
|
||||||
)
|
)
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
@@ -298,10 +299,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
||||||
api_format = self.allowed_api_formats[0]
|
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 = StreamContext(model=model, api_format=api_format)
|
||||||
ctx.request_id = self.request_id
|
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)
|
# 创建更新状态的回调闭包(可以访问 ctx)
|
||||||
def update_streaming_status() -> None:
|
def update_streaming_status() -> None:
|
||||||
@@ -327,7 +334,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
provider,
|
provider,
|
||||||
endpoint,
|
endpoint,
|
||||||
key,
|
key,
|
||||||
original_request_body,
|
request_body_ref["body"], # 使用容器中的请求体
|
||||||
original_headers,
|
original_headers,
|
||||||
query_params,
|
query_params,
|
||||||
candidate,
|
candidate,
|
||||||
@@ -356,6 +363,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
capability_requirements=capability_requirements or None,
|
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.provider_id = provider_id
|
||||||
ctx.endpoint_id = endpoint_id
|
ctx.endpoint_id = endpoint_id
|
||||||
ctx.key_id = key_id
|
ctx.key_id = key_id
|
||||||
|
# 同步整流状态(如果请求体被整流过)
|
||||||
|
ctx.rectified = request_body_ref.get("_rectified", False)
|
||||||
|
|
||||||
# 创建遥测记录器
|
# 创建遥测记录器
|
||||||
telemetry_recorder = StreamTelemetryRecorder(
|
telemetry_recorder = StreamTelemetryRecorder(
|
||||||
@@ -405,6 +415,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
background=background_tasks,
|
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:
|
except Exception as e:
|
||||||
self._log_request_error("流式请求失败", e)
|
self._log_request_error("流式请求失败", e)
|
||||||
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
|
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()
|
response_time_ms = self.elapsed_ms()
|
||||||
|
|
||||||
status_code = 503
|
status_code = 503
|
||||||
if isinstance(error, ProviderAuthException):
|
if isinstance(error, ThinkingSignatureException):
|
||||||
|
status_code = 400
|
||||||
|
elif isinstance(error, ProviderAuthException):
|
||||||
status_code = 503
|
status_code = 503
|
||||||
elif isinstance(error, ProviderRateLimitException):
|
elif isinstance(error, ProviderRateLimitException):
|
||||||
status_code = 429
|
status_code = 429
|
||||||
@@ -668,6 +687,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
||||||
api_format = self.allowed_api_formats[0]
|
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
|
provider_name: Optional[str] = None
|
||||||
response_json: Optional[Dict[str, Any]] = None
|
response_json: Optional[Dict[str, Any]] = None
|
||||||
@@ -709,9 +732,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
# 应用模型映射
|
# 应用模型映射
|
||||||
if mapped_model:
|
if mapped_model:
|
||||||
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
|
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:
|
else:
|
||||||
request_body = dict(original_request_body)
|
request_body = dict(request_body_ref["body"])
|
||||||
|
|
||||||
# 跨格式:先做请求体转换(严格模式,失败触发 failover)
|
# 跨格式:先做请求体转换(严格模式,失败触发 failover)
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
@@ -874,7 +897,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
response_json = converter_registry.convert_response_strict(
|
response_json = converter_registry.convert_response_strict(
|
||||||
response_json,
|
response_json,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
str(api_format),
|
client_api_format,
|
||||||
)
|
)
|
||||||
|
|
||||||
return response_json if isinstance(response_json, dict) else {}
|
return response_json if isinstance(response_json, dict) else {}
|
||||||
@@ -900,6 +923,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_func=sync_request_func,
|
request_func=sync_request_func,
|
||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
provider_name = actual_provider_name
|
provider_name = actual_provider_name
|
||||||
@@ -965,6 +989,23 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
headers=client_response_headers,
|
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:
|
except Exception as e:
|
||||||
response_time_ms = self.elapsed_ms()
|
response_time_ms = self.elapsed_ms()
|
||||||
|
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ from src.core.exceptions import (
|
|||||||
ProviderNotAvailableException,
|
ProviderNotAvailableException,
|
||||||
ProviderRateLimitException,
|
ProviderRateLimitException,
|
||||||
ProviderTimeoutException,
|
ProviderTimeoutException,
|
||||||
|
ThinkingSignatureException,
|
||||||
)
|
)
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
@@ -303,7 +304,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"""
|
"""
|
||||||
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
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 位置不同)
|
# 使用子类实现的方法提取 model(不同 API 格式的 model 位置不同)
|
||||||
|
# 注意:使用 original_request_body,因为整流只修改 messages,不影响 model 字段
|
||||||
model = self.extract_model_from_request(original_request_body, path_params)
|
model = self.extract_model_from_request(original_request_body, path_params)
|
||||||
|
|
||||||
# 创建流上下文
|
# 创建流上下文
|
||||||
@@ -327,7 +333,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
provider,
|
provider,
|
||||||
endpoint,
|
endpoint,
|
||||||
key,
|
key,
|
||||||
original_request_body,
|
request_body_ref["body"], # 使用容器中的请求体
|
||||||
original_headers,
|
original_headers,
|
||||||
query_params,
|
query_params,
|
||||||
candidate,
|
candidate,
|
||||||
@@ -356,6 +362,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
# 更新上下文(确保 provider 信息已设置,用于 streaming 状态更新)
|
# 更新上下文(确保 provider 信息已设置,用于 streaming 状态更新)
|
||||||
@@ -368,6 +375,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
ctx.endpoint_id = endpoint_id
|
ctx.endpoint_id = endpoint_id
|
||||||
if not ctx.key_id:
|
if not ctx.key_id:
|
||||||
ctx.key_id = key_id
|
ctx.key_id = key_id
|
||||||
|
# 同步整流状态(如果请求体被整流过)
|
||||||
|
ctx.rectified = request_body_ref.get("_rectified", False)
|
||||||
|
|
||||||
# 创建后台任务记录统计
|
# 创建后台任务记录统计
|
||||||
background_tasks = BackgroundTasks()
|
background_tasks = BackgroundTasks()
|
||||||
@@ -396,6 +405,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
background=background_tasks,
|
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:
|
except Exception as e:
|
||||||
self._log_request_error("流式请求失败", e)
|
self._log_request_error("流式请求失败", e)
|
||||||
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
|
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
|
||||||
@@ -856,7 +872,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"上游服务返回了非预期的响应格式",
|
"上游服务返回了非预期的响应格式",
|
||||||
provider_name=str(provider.name),
|
provider_name=str(provider.name),
|
||||||
upstream_status=200,
|
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(":"):
|
if not normalized_line or normalized_line.startswith(":"):
|
||||||
@@ -1350,7 +1368,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
# 记录失败的 Usage,但使用已收到的预估 token 信息(来自 message_start)
|
# 记录失败的 Usage,但使用已收到的预估 token 信息(来自 message_start)
|
||||||
# 这样即使请求中断,也能记录预估成本
|
# 这样即使请求中断,也能记录预估成本
|
||||||
# 失败时返回给客户端的是 JSON 错误响应,如果没有设置则使用默认值
|
# 失败时返回给客户端的是 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(
|
await bg_telemetry.record_failure(
|
||||||
provider=ctx.provider_name or "unknown",
|
provider=ctx.provider_name or "unknown",
|
||||||
model=ctx.model,
|
model=ctx.model,
|
||||||
@@ -1385,11 +1405,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
|
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
|
||||||
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
|
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
|
||||||
client_response_headers.update({
|
client_response_headers.update(
|
||||||
"Cache-Control": "no-cache, no-transform",
|
{
|
||||||
"X-Accel-Buffering": "no",
|
"Cache-Control": "no-cache, no-transform",
|
||||||
"content-type": "text/event-stream",
|
"X-Accel-Buffering": "no",
|
||||||
})
|
"content-type": "text/event-stream",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
total_cost = await bg_telemetry.record_success(
|
total_cost = await bg_telemetry.record_success(
|
||||||
provider=ctx.provider_name,
|
provider=ctx.provider_name,
|
||||||
@@ -1439,11 +1461,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
# 计算候选自身的 TTFB
|
# 计算候选自身的 TTFB
|
||||||
candidate_first_byte_time_ms: Optional[int] = None
|
candidate_first_byte_time_ms: Optional[int] = None
|
||||||
if ctx.first_byte_time_ms is not None:
|
if ctx.first_byte_time_ms is not None:
|
||||||
candidate_first_byte_time_ms = RequestCandidateService.calculate_candidate_ttfb(
|
candidate_first_byte_time_ms = (
|
||||||
db=bg_db,
|
RequestCandidateService.calculate_candidate_ttfb(
|
||||||
candidate_id=ctx.attempt_id,
|
db=bg_db,
|
||||||
request_start_time=self.start_time,
|
candidate_id=ctx.attempt_id,
|
||||||
global_first_byte_time_ms=ctx.first_byte_time_ms,
|
request_start_time=self.start_time,
|
||||||
|
global_first_byte_time_ms=ctx.first_byte_time_ms,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# 根据状态码决定是成功还是失败
|
# 根据状态码决定是成功还是失败
|
||||||
@@ -1451,7 +1475,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
# 503 = 服务不可用(如流中断),应标记为失败
|
# 503 = 服务不可用(如流中断),应标记为失败
|
||||||
if ctx.status_code and ctx.status_code >= 400:
|
if ctx.status_code and ctx.status_code >= 400:
|
||||||
# 请求链路追踪使用 upstream_response(原始响应),回退到 error_message(友好消息)
|
# 请求链路追踪使用 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 = {
|
extra_data = {
|
||||||
"stream_completed": False,
|
"stream_completed": False,
|
||||||
"chunk_count": ctx.chunk_count,
|
"chunk_count": ctx.chunk_count,
|
||||||
@@ -1476,6 +1502,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"chunk_count": ctx.chunk_count,
|
"chunk_count": ctx.chunk_count,
|
||||||
"data_count": ctx.data_count,
|
"data_count": ctx.data_count,
|
||||||
}
|
}
|
||||||
|
if ctx.rectified:
|
||||||
|
extra_data["rectified"] = True
|
||||||
if candidate_first_byte_time_ms is not None:
|
if candidate_first_byte_time_ms is not None:
|
||||||
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
|
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
|
||||||
RequestCandidateService.mark_candidate_success(
|
RequestCandidateService.mark_candidate_success(
|
||||||
@@ -1504,7 +1532,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||||
|
|
||||||
status_code = 503
|
status_code = 503
|
||||||
if isinstance(error, ProviderAuthException):
|
if isinstance(error, ThinkingSignatureException):
|
||||||
|
status_code = 400
|
||||||
|
elif isinstance(error, ProviderAuthException):
|
||||||
status_code = 503
|
status_code = 503
|
||||||
elif isinstance(error, ProviderRateLimitException):
|
elif isinstance(error, ProviderRateLimitException):
|
||||||
status_code = 429
|
status_code = 429
|
||||||
@@ -1574,6 +1604,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
||||||
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
|
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(
|
async def sync_request_func(
|
||||||
provider: Provider,
|
provider: Provider,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
@@ -1595,9 +1629,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
|
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
|
||||||
if mapped_model:
|
if mapped_model:
|
||||||
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
|
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:
|
else:
|
||||||
request_body = original_request_body
|
request_body = dict(request_body_ref["body"])
|
||||||
|
|
||||||
# 准备发送给 Provider 的请求体(子类可覆盖以移除不需要的字段)
|
# 准备发送给 Provider 的请求体(子类可覆盖以移除不需要的字段)
|
||||||
request_body = self.prepare_provider_request_body(request_body)
|
request_body = self.prepare_provider_request_body(request_body)
|
||||||
@@ -1749,6 +1783,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
request_func=sync_request_func,
|
request_func=sync_request_func,
|
||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
provider_name = actual_provider_name
|
provider_name = actual_provider_name
|
||||||
@@ -1830,6 +1865,24 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
headers=client_response_headers,
|
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:
|
except Exception as e:
|
||||||
response_time_ms = int((time.time() - sync_start_time) * 1000)
|
response_time_ms = int((time.time() - sync_start_time) * 1000)
|
||||||
|
|
||||||
|
|||||||
@@ -81,6 +81,9 @@ class StreamContext:
|
|||||||
# Provider 响应元数据(CLI handler 需要)
|
# Provider 响应元数据(CLI handler 需要)
|
||||||
response_metadata: Dict[str, Any] = field(default_factory=dict)
|
response_metadata: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
# 整流标记(Thinking Rectifier)
|
||||||
|
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
|
||||||
|
|
||||||
# 流式处理统计
|
# 流式处理统计
|
||||||
data_count: int = 0
|
data_count: int = 0
|
||||||
chunk_count: int = 0
|
chunk_count: int = 0
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from src.api.handlers.base.utils import (
|
|||||||
)
|
)
|
||||||
from src.config.constants import StreamDefaults
|
from src.config.constants import StreamDefaults
|
||||||
from src.config.settings import config
|
from src.config.settings import config
|
||||||
|
from src.core.api_format import FormatConversionError, converter_registry
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
EmbeddedErrorException,
|
EmbeddedErrorException,
|
||||||
ProviderNotAvailableException,
|
ProviderNotAvailableException,
|
||||||
@@ -296,6 +297,31 @@ class StreamProcessor:
|
|||||||
error_status=parsed.error_type,
|
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
|
should_stop = True
|
||||||
break
|
break
|
||||||
@@ -323,7 +349,12 @@ class StreamProcessor:
|
|||||||
base_url=endpoint.base_url,
|
base_url=endpoint.base_url,
|
||||||
)
|
)
|
||||||
|
|
||||||
except (EmbeddedErrorException, ProviderNotAvailableException, ProviderTimeoutException):
|
except (
|
||||||
|
EmbeddedErrorException,
|
||||||
|
ProviderNotAvailableException,
|
||||||
|
ProviderTimeoutException,
|
||||||
|
FormatConversionError,
|
||||||
|
):
|
||||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||||
raise
|
raise
|
||||||
except (OSError, IOError) as e:
|
except (OSError, IOError) as e:
|
||||||
|
|||||||
@@ -297,6 +297,8 @@ class StreamTelemetryRecorder:
|
|||||||
"stream_completed": ctx.is_success(),
|
"stream_completed": ctx.is_success(),
|
||||||
"data_count": ctx.data_count,
|
"data_count": ctx.data_count,
|
||||||
}
|
}
|
||||||
|
if ctx.rectified:
|
||||||
|
extra_data["rectified"] = True
|
||||||
if ctx.first_byte_time_ms is not None:
|
if ctx.first_byte_time_ms is not None:
|
||||||
# 计算候选自身的 TTFB
|
# 计算候选自身的 TTFB
|
||||||
first_byte_time_ms = RequestCandidateService.calculate_candidate_ttfb(
|
first_byte_time_ms = RequestCandidateService.calculate_candidate_ttfb(
|
||||||
|
|||||||
@@ -173,6 +173,14 @@ class Config:
|
|||||||
self.stream_stats_delay = float(os.getenv("STREAM_STATS_DELAY", "0.1"))
|
self.stream_stats_delay = float(os.getenv("STREAM_STATS_DELAY", "0.1"))
|
||||||
self.stream_first_byte_timeout = float(os.getenv("STREAM_FIRST_BYTE_TIMEOUT", "30.0"))
|
self.stream_first_byte_timeout = float(os.getenv("STREAM_FIRST_BYTE_TIMEOUT", "30.0"))
|
||||||
|
|
||||||
|
# Thinking 整流器配置
|
||||||
|
# THINKING_RECTIFIER_ENABLED: 是否启用 Thinking 整流器
|
||||||
|
# 当遇到跨 Provider 的 thinking 签名错误时,自动整流请求体后重试
|
||||||
|
# 默认启用,设为 false 可禁用此功能
|
||||||
|
self.thinking_rectifier_enabled = (
|
||||||
|
os.getenv("THINKING_RECTIFIER_ENABLED", "true").lower() == "true"
|
||||||
|
)
|
||||||
|
|
||||||
# 请求体读取超时(秒)
|
# 请求体读取超时(秒)
|
||||||
# REQUEST_BODY_TIMEOUT: 等待客户端发送完整请求体的超时时间
|
# REQUEST_BODY_TIMEOUT: 等待客户端发送完整请求体的超时时间
|
||||||
# 默认 60 秒,防止客户端发送不完整请求导致连接卡死
|
# 默认 60 秒,防止客户端发送不完整请求导致连接卡死
|
||||||
|
|||||||
@@ -17,9 +17,9 @@ import httpx
|
|||||||
from fastapi import HTTPException, status
|
from fastapi import HTTPException, status
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
from ..config import config
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
from ..config import config
|
||||||
|
|
||||||
# Pydantic 错误消息中英文翻译映射
|
# Pydantic 错误消息中英文翻译映射
|
||||||
PYDANTIC_ERROR_TRANSLATIONS = {
|
PYDANTIC_ERROR_TRANSLATIONS = {
|
||||||
@@ -139,7 +139,6 @@ def translate_pydantic_errors(errors: List[Dict[str, Any]]) -> str:
|
|||||||
return "; ".join(translated)
|
return "; ".join(translated)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# 延迟导入韧性管理器,避免循环导入
|
# 延迟导入韧性管理器,避免循环导入
|
||||||
def get_resilience_manager():
|
def get_resilience_manager():
|
||||||
try:
|
try:
|
||||||
@@ -530,6 +529,26 @@ class UpstreamClientException(ProxyException):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ThinkingSignatureException(UpstreamClientException):
|
||||||
|
"""Thinking 块签名验证失败异常"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
message: str,
|
||||||
|
provider_name: Optional[str] = None,
|
||||||
|
upstream_error: Optional[str] = None,
|
||||||
|
request_metadata: Any = None,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
message=message,
|
||||||
|
provider_name=provider_name,
|
||||||
|
status_code=400,
|
||||||
|
error_type="thinking_signature_error",
|
||||||
|
upstream_error=upstream_error,
|
||||||
|
request_metadata=request_metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ErrorResponse:
|
class ErrorResponse:
|
||||||
"""统一的错误响应格式化器"""
|
"""统一的错误响应格式化器"""
|
||||||
|
|
||||||
|
|||||||
@@ -1096,7 +1096,8 @@ class ProviderAPIKey(Base):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# API 格式支持列表(核心字段)
|
# API 格式支持列表(核心字段)
|
||||||
api_formats = Column(JSON, nullable=False, default=list) # ["CLAUDE", "CLAUDE_CLI"]
|
# None 表示支持所有格式(兼容历史数据),空列表 [] 表示不支持任何格式
|
||||||
|
api_formats = Column(JSON, nullable=True, default=list) # ["CLAUDE", "CLAUDE_CLI"]
|
||||||
|
|
||||||
# API密钥信息
|
# API密钥信息
|
||||||
api_key = Column(String(500), nullable=False) # API密钥(加密存储)
|
api_key = Column(String(500), nullable=False) # API密钥(加密存储)
|
||||||
|
|||||||
34
src/services/cache/aware_scheduler.py
vendored
34
src/services/cache/aware_scheduler.py
vendored
@@ -1084,10 +1084,12 @@ class CacheAwareScheduler:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
# Key 直属 Provider,通过 api_formats 按端点格式筛选
|
# Key 直属 Provider,通过 api_formats 按端点格式筛选
|
||||||
|
# api_formats=None 视为"全支持"(兼容历史数据)
|
||||||
active_keys = [
|
active_keys = [
|
||||||
key
|
key
|
||||||
for key in provider.api_keys
|
for key in provider.api_keys
|
||||||
if key.is_active and endpoint_format_str in (key.api_formats or [])
|
if key.is_active
|
||||||
|
and (key.api_formats is None or endpoint_format_str in key.api_formats)
|
||||||
]
|
]
|
||||||
if not active_keys:
|
if not active_keys:
|
||||||
continue
|
continue
|
||||||
@@ -1264,21 +1266,31 @@ class CacheAwareScheduler:
|
|||||||
"""
|
"""
|
||||||
根据优先级模式对候选列表排序(数字越小越优先)
|
根据优先级模式对候选列表排序(数字越小越优先)
|
||||||
|
|
||||||
- provider: 提供商优先模式,保持原有顺序(按 Provider.provider_priority -> Key.internal_priority 排序,已由查询保证)
|
排序规则:
|
||||||
Key.internal_priority 表示 Endpoint 内部优先级,同优先级内通过哈希分散负载均衡
|
1. exact 候选(needs_conversion=False)优先于 convertible 候选
|
||||||
- global_key: 全局 Key 优先模式,按 Key.global_priority_by_format 升序排序(数字小的优先)
|
2. 在同一类型内,按优先级模式排序:
|
||||||
有优先级的优先,NULL 的排后面
|
- provider: 提供商优先模式,按 Provider.provider_priority -> Key.internal_priority 排序
|
||||||
同优先级内通过哈希分散实现负载均衡
|
- global_key: 全局 Key 优先模式,按 Key.global_priority_by_format 排序
|
||||||
"""
|
"""
|
||||||
if not candidates:
|
if not candidates:
|
||||||
return candidates
|
return candidates
|
||||||
|
|
||||||
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
|
# 按 needs_conversion 分组:exact 优先
|
||||||
# 全局 Key 优先模式:按 global_priority 分组,同组内哈希分散负载均衡
|
exact_candidates = [c for c in candidates if not c.needs_conversion]
|
||||||
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
|
convertible_candidates = [c for c in candidates if c.needs_conversion]
|
||||||
|
|
||||||
# 提供商优先模式:保持原有顺序(provider_priority 排序已经由查询保证)
|
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
|
||||||
return candidates
|
# 全局 Key 优先模式:分别对两组排序后合并
|
||||||
|
sorted_exact = self._sort_by_global_priority_with_hash(
|
||||||
|
exact_candidates, affinity_key, api_format
|
||||||
|
)
|
||||||
|
sorted_convertible = self._sort_by_global_priority_with_hash(
|
||||||
|
convertible_candidates, affinity_key, api_format
|
||||||
|
)
|
||||||
|
return sorted_exact + sorted_convertible
|
||||||
|
|
||||||
|
# 提供商优先模式:exact 在前,convertible 在后(各组内部顺序已由构建时保证)
|
||||||
|
return exact_candidates + convertible_candidates
|
||||||
|
|
||||||
def _sort_by_global_priority_with_hash(
|
def _sort_by_global_priority_with_hash(
|
||||||
self,
|
self,
|
||||||
|
|||||||
7
src/services/message/__init__.py
Normal file
7
src/services/message/__init__.py
Normal file
@@ -0,0 +1,7 @@
|
|||||||
|
"""
|
||||||
|
消息处理服务模块
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .thinking_rectifier import ThinkingRectifier
|
||||||
|
|
||||||
|
__all__ = ["ThinkingRectifier"]
|
||||||
216
src/services/message/thinking_rectifier.py
Normal file
216
src/services/message/thinking_rectifier.py
Normal file
@@ -0,0 +1,216 @@
|
|||||||
|
"""
|
||||||
|
Thinking 整流器(Rectifier)
|
||||||
|
|
||||||
|
采用 cc-switch 的"错误触发"模式,在遇到 Thinking 签名/结构错误时触发整流。
|
||||||
|
|
||||||
|
核心功能:
|
||||||
|
1. 移除所有 thinking 和 redacted_thinking 块
|
||||||
|
2. 移除非 thinking 块上的 signature 字段
|
||||||
|
3. 条件删除顶层 thinking 参数
|
||||||
|
|
||||||
|
使用场景:
|
||||||
|
当遇到 ThinkingSignatureException 时,调用 rectify() 整流请求体后重试一次。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import copy
|
||||||
|
from typing import Any, Dict, List, Tuple
|
||||||
|
|
||||||
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
class ThinkingRectifier:
|
||||||
|
"""
|
||||||
|
Thinking 整流器
|
||||||
|
|
||||||
|
在遇到 Thinking 签名/结构错误时,整流请求体以便重试。
|
||||||
|
采用"彻底清洗 + 条件禁用 thinking"策略。
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def rectify(request_body: Dict[str, Any]) -> Tuple[Dict[str, Any], bool]:
|
||||||
|
"""
|
||||||
|
整流请求体
|
||||||
|
|
||||||
|
执行以下操作:
|
||||||
|
1. 移除所有 thinking 和 redacted_thinking 块
|
||||||
|
2. 移除非 thinking 块上的 signature 字段
|
||||||
|
3. 条件删除顶层 thinking 参数
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request_body: 原始请求体
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple[整流后的请求体, 是否有修改]
|
||||||
|
"""
|
||||||
|
if not request_body:
|
||||||
|
return request_body, False
|
||||||
|
|
||||||
|
# 深拷贝以避免修改原始数据
|
||||||
|
rectified_body = copy.deepcopy(request_body)
|
||||||
|
modified = False
|
||||||
|
|
||||||
|
# 1. 整流 messages
|
||||||
|
messages = rectified_body.get("messages", [])
|
||||||
|
if messages:
|
||||||
|
rectified_messages, messages_modified = ThinkingRectifier._rectify_messages(messages)
|
||||||
|
if messages_modified:
|
||||||
|
rectified_body["messages"] = rectified_messages
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
# 2. 条件删除顶层 thinking 参数(使用整流后的 messages 判断)
|
||||||
|
# 与 cc-switch 行为一致:在整流 messages 之后获取快照进行判断
|
||||||
|
if ThinkingRectifier._should_remove_top_level_thinking(rectified_body):
|
||||||
|
if "thinking" in rectified_body:
|
||||||
|
del rectified_body["thinking"]
|
||||||
|
modified = True
|
||||||
|
logger.info("ThinkingRectifier: 已移除顶层 thinking 参数")
|
||||||
|
|
||||||
|
return rectified_body, modified
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _rectify_messages(messages: List[Dict[str, Any]]) -> Tuple[List[Dict[str, Any]], bool]:
|
||||||
|
"""
|
||||||
|
整流消息列表
|
||||||
|
|
||||||
|
移除所有 thinking/redacted_thinking 块和 signature 字段
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: 原始消息列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple[整流后的消息列表, 是否有修改]
|
||||||
|
"""
|
||||||
|
if not messages:
|
||||||
|
return messages, False
|
||||||
|
|
||||||
|
modified = False
|
||||||
|
result_messages: List[Dict[str, Any]] = []
|
||||||
|
thinking_removed = 0
|
||||||
|
signature_removed = 0
|
||||||
|
|
||||||
|
for message in messages:
|
||||||
|
# 类型保护:跳过非 dict 消息
|
||||||
|
if not isinstance(message, dict):
|
||||||
|
result_messages.append(message)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 消息级浅拷贝:外层 rectify() 已深拷贝整个 request_body
|
||||||
|
# content 会被重建为新列表,不会影响原始数据
|
||||||
|
new_message = dict(message)
|
||||||
|
content = message.get("content")
|
||||||
|
|
||||||
|
if isinstance(content, list):
|
||||||
|
new_content = []
|
||||||
|
for block in content:
|
||||||
|
if isinstance(block, dict):
|
||||||
|
block_type = block.get("type")
|
||||||
|
|
||||||
|
# 移除 thinking 和 redacted_thinking 块
|
||||||
|
if block_type in ("thinking", "redacted_thinking"):
|
||||||
|
thinking_removed += 1
|
||||||
|
modified = True
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 移除非 thinking 块上的 signature 字段
|
||||||
|
if "signature" in block:
|
||||||
|
new_block = {k: v for k, v in block.items() if k != "signature"}
|
||||||
|
new_content.append(new_block)
|
||||||
|
signature_removed += 1
|
||||||
|
modified = True
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_content.append(block)
|
||||||
|
else:
|
||||||
|
new_content.append(block)
|
||||||
|
|
||||||
|
# 更新 content
|
||||||
|
new_message["content"] = new_content
|
||||||
|
|
||||||
|
# 如果整流后 assistant 消息的 content 为空,记录警告
|
||||||
|
# (空 content 本身不是"修改",只是检测到的状态,不设置 modified)
|
||||||
|
# 保留消息是必要的:跳过会破坏对话结构(后续 tool_result 消息需要前置 assistant 消息)
|
||||||
|
if new_message.get("role") == "assistant":
|
||||||
|
effective_content = new_message.get("content")
|
||||||
|
is_empty = not effective_content or (
|
||||||
|
isinstance(effective_content, list) and len(effective_content) == 0
|
||||||
|
)
|
||||||
|
if is_empty:
|
||||||
|
msg_idx = len(result_messages)
|
||||||
|
logger.warning(
|
||||||
|
f"ThinkingRectifier: assistant 消息整流后 content 为空 (message_index={msg_idx})"
|
||||||
|
)
|
||||||
|
|
||||||
|
result_messages.append(new_message)
|
||||||
|
|
||||||
|
if thinking_removed > 0 or signature_removed > 0:
|
||||||
|
logger.info(
|
||||||
|
f"ThinkingRectifier: 移除了 {thinking_removed} 个 thinking 块, "
|
||||||
|
f"{signature_removed} 个 signature 字段"
|
||||||
|
)
|
||||||
|
|
||||||
|
return result_messages, modified
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _should_remove_top_level_thinking(body: Dict[str, Any]) -> bool:
|
||||||
|
"""
|
||||||
|
判断是否应该删除顶层 thinking 参数
|
||||||
|
|
||||||
|
与 cc-switch 行为一致:只检查最后一条 assistant 消息
|
||||||
|
|
||||||
|
设计思路:
|
||||||
|
- body 中的 messages 是整流后的状态,thinking 块已被移除
|
||||||
|
- Claude API 只校验最后一条 assistant 消息的结构
|
||||||
|
- 如果最后一条有 tool_use 但首块不是 thinking,需要禁用 thinking 参数
|
||||||
|
|
||||||
|
Args:
|
||||||
|
body: 整流后的请求体
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
是否应该删除顶层 thinking 参数
|
||||||
|
"""
|
||||||
|
# 条件 1: thinking 参数存在且已启用
|
||||||
|
thinking_param = body.get("thinking")
|
||||||
|
if not isinstance(thinking_param, dict) or thinking_param.get("type") != "enabled":
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 从 body 中获取 messages
|
||||||
|
messages = body.get("messages", [])
|
||||||
|
|
||||||
|
# 类型保护:确保 messages 是 list
|
||||||
|
if not isinstance(messages, list) or not messages:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 条件 2: 找到最后一条 assistant 消息
|
||||||
|
last_assistant = None
|
||||||
|
for message in reversed(messages):
|
||||||
|
if isinstance(message, dict) and message.get("role") == "assistant":
|
||||||
|
last_assistant = message
|
||||||
|
break
|
||||||
|
|
||||||
|
if not last_assistant:
|
||||||
|
return False
|
||||||
|
|
||||||
|
content = last_assistant.get("content")
|
||||||
|
if not isinstance(content, list) or not content:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 注意:传入的 messages 是整流后的状态,thinking 块已被移除
|
||||||
|
# 因此只需检查是否有 tool_use,如果有则需要禁用 thinking 参数
|
||||||
|
# (因为整流后的 assistant 消息不再以 thinking 块开头)
|
||||||
|
|
||||||
|
# 检查是否有 tool_use
|
||||||
|
has_tool_use = any(
|
||||||
|
isinstance(block, dict) and block.get("type") == "tool_use" for block in content
|
||||||
|
)
|
||||||
|
|
||||||
|
# 整流后 assistant 消息不再以 thinking 块开头,如果有 tool_use 则需要禁用 thinking 参数
|
||||||
|
# (Claude API 要求:启用 thinking 时,有 tool_use 的 assistant 消息必须以 thinking 块开头)
|
||||||
|
if has_tool_use:
|
||||||
|
logger.info(
|
||||||
|
"ThinkingRectifier: 整流后 assistant 消息有 tool_use 但无 thinking 前缀,"
|
||||||
|
"禁用 thinking 参数以通过 API 校验"
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
logger.debug("ThinkingRectifier: 整流后 assistant 消息无 tool_use,保留 thinking 参数")
|
||||||
|
return False
|
||||||
@@ -19,6 +19,7 @@ from src.core.exceptions import (
|
|||||||
ProviderException,
|
ProviderException,
|
||||||
ProviderNotAvailableException,
|
ProviderNotAvailableException,
|
||||||
ProviderRateLimitException,
|
ProviderRateLimitException,
|
||||||
|
ThinkingSignatureException,
|
||||||
UpstreamClientException,
|
UpstreamClientException,
|
||||||
)
|
)
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
@@ -151,6 +152,30 @@ class ErrorClassifier:
|
|||||||
"not available for this model", # 此模型不可用
|
"not available for this model", # 此模型不可用
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Thinking 块相关错误模式 - 这类错误需要清洗 thinking 块或调整请求
|
||||||
|
# 场景:多供应商环境下,Provider A 生成的 thinking 块被发送到 Provider B 时签名验证失败
|
||||||
|
THINKING_ERROR_PATTERNS: Tuple[str, ...] = (
|
||||||
|
# 签名错误:跨 Provider 发送 thinking 块时,签名无法被目标 Provider 验证
|
||||||
|
# 例: "invalid `signature` in `thinking` block: signature is for a different request"
|
||||||
|
"invalid `signature` in `thinking` block",
|
||||||
|
"invalid signature in thinking block",
|
||||||
|
# 签名字段缺失或格式错误
|
||||||
|
# 例: "messages.0.content.0.thinking.signature: field required"
|
||||||
|
"thinking.signature: field required",
|
||||||
|
"thinking.signature:", # 匹配路径模式 messages.X.content.X.thinking.signature: xxx
|
||||||
|
"signature verification failed",
|
||||||
|
# 结构错误:启用 thinking 时,有 tool_use 的 assistant 消息必须以 thinking 块开头
|
||||||
|
# 例: "when `thinking` is enabled, the first content block ... must start with a `thinking` block"
|
||||||
|
"must start with a thinking block",
|
||||||
|
# 例: "expected thinking or redacted_thinking, found tool_use"
|
||||||
|
"expected thinking or redacted_thinking",
|
||||||
|
"expected `thinking`",
|
||||||
|
"expected thinking, found", # 统一匹配 "found tool_use/text" 等变体
|
||||||
|
"expected `thinking`, found", # 带反引号变体
|
||||||
|
"expected redacted_thinking, found",
|
||||||
|
"expected `redacted_thinking`, found",
|
||||||
|
)
|
||||||
|
|
||||||
def _parse_error_response(self, error_text: Optional[str]) -> Dict[str, Any]:
|
def _parse_error_response(self, error_text: Optional[str]) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
解析错误响应为结构化数据
|
解析错误响应为结构化数据
|
||||||
@@ -295,6 +320,25 @@ class ErrorClassifier:
|
|||||||
search_text = error_text.lower()
|
search_text = error_text.lower()
|
||||||
return any(pattern.lower() in search_text for pattern in self.COMPATIBILITY_ERROR_PATTERNS)
|
return any(pattern.lower() in search_text for pattern in self.COMPATIBILITY_ERROR_PATTERNS)
|
||||||
|
|
||||||
|
def _is_thinking_error(self, error_text: Optional[str]) -> bool:
|
||||||
|
"""
|
||||||
|
检测错误响应是否为 Thinking 块相关错误(签名错误或结构错误)
|
||||||
|
|
||||||
|
这类错误通常发生在:
|
||||||
|
1. 多供应商场景下,当一个供应商生成的 thinking 块被发送到另一个供应商时,签名验证会失败
|
||||||
|
2. 请求体中有 tool_use 但没有以 thinking 块开头时,Claude 会报结构错误
|
||||||
|
|
||||||
|
Args:
|
||||||
|
error_text: 错误响应文本
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
是否为 Thinking 相关错误
|
||||||
|
"""
|
||||||
|
if not error_text:
|
||||||
|
return False
|
||||||
|
search_text = error_text.lower()
|
||||||
|
return any(p.lower() in search_text for p in self.THINKING_ERROR_PATTERNS)
|
||||||
|
|
||||||
def _extract_error_message(self, error_text: Optional[str]) -> Optional[str]:
|
def _extract_error_message(self, error_text: Optional[str]) -> Optional[str]:
|
||||||
"""
|
"""
|
||||||
从错误响应中提取错误消息
|
从错误响应中提取错误消息
|
||||||
@@ -463,6 +507,15 @@ class ErrorClassifier:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 400 错误:检查是否为 Thinking 块签名错误
|
||||||
|
if status == 400 and self._is_thinking_error(error_response_text):
|
||||||
|
logger.info(f"检测到 Thinking 块错误: {extracted_message}")
|
||||||
|
return ThinkingSignatureException(
|
||||||
|
message=extracted_message or "Thinking block signature validation failed",
|
||||||
|
provider_name=provider_name,
|
||||||
|
upstream_error=error_response_text,
|
||||||
|
)
|
||||||
|
|
||||||
# 400 错误:先检查是否为 Provider 兼容性错误(应触发故障转移)
|
# 400 错误:先检查是否为 Provider 兼容性错误(应触发故障转移)
|
||||||
if status == 400 and self._is_compatibility_error(error_response_text):
|
if status == 400 and self._is_compatibility_error(error_response_text):
|
||||||
logger.info(f"检测到 Provider 兼容性错误,将触发故障转移: {extracted_message}")
|
logger.info(f"检测到 Provider 兼容性错误,将触发故障转移: {extracted_message}")
|
||||||
|
|||||||
@@ -29,12 +29,14 @@ import httpx
|
|||||||
from redis import Redis
|
from redis import Redis
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.config.settings import config
|
||||||
from src.core.api_format import APIFormat, FormatConversionError
|
from src.core.api_format import APIFormat, FormatConversionError
|
||||||
from src.core.error_utils import extract_error_message
|
from src.core.error_utils import extract_error_message
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
ConcurrencyLimitError,
|
ConcurrencyLimitError,
|
||||||
EmbeddedErrorException,
|
EmbeddedErrorException,
|
||||||
ProviderNotAvailableException,
|
ProviderNotAvailableException,
|
||||||
|
ThinkingSignatureException,
|
||||||
UpstreamClientException,
|
UpstreamClientException,
|
||||||
)
|
)
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
@@ -44,6 +46,7 @@ from src.services.cache.aware_scheduler import (
|
|||||||
ProviderCandidate,
|
ProviderCandidate,
|
||||||
get_cache_aware_scheduler,
|
get_cache_aware_scheduler,
|
||||||
)
|
)
|
||||||
|
from src.services.message.thinking_rectifier import ThinkingRectifier
|
||||||
from src.services.provider.format import normalize_api_format
|
from src.services.provider.format import normalize_api_format
|
||||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||||
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
||||||
@@ -298,6 +301,114 @@ class FallbackOrchestrator:
|
|||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _handle_thinking_signature_error(
|
||||||
|
self,
|
||||||
|
converted_error: ThinkingSignatureException,
|
||||||
|
request_id: Optional[str],
|
||||||
|
candidate_record_id: str,
|
||||||
|
elapsed_ms: int,
|
||||||
|
captured_key_concurrent: Optional[int],
|
||||||
|
serializable_extra_data: Dict[str, Any],
|
||||||
|
request_body_ref: Optional[Dict[str, Any]],
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
处理 ThinkingSignatureException 错误
|
||||||
|
|
||||||
|
尝试整流请求体后重试。如果无法整流或整流后仍失败,则抛出异常。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
converted_error: Thinking 签名异常
|
||||||
|
request_id: 请求 ID
|
||||||
|
candidate_record_id: 候选记录 ID
|
||||||
|
elapsed_ms: 耗时(毫秒)
|
||||||
|
captured_key_concurrent: 捕获的并发数
|
||||||
|
serializable_extra_data: 可序列化的额外数据
|
||||||
|
request_body_ref: 请求体引用容器
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
"continue" 表示整流成功应继续重试
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ThinkingSignatureException: 无法整流或整流后仍失败时
|
||||||
|
"""
|
||||||
|
# 检查整流器是否启用
|
||||||
|
if not config.thinking_rectifier_enabled:
|
||||||
|
logger.info(f" [{request_id}] Thinking 错误:整流器已禁用,终止重试")
|
||||||
|
self._mark_thinking_error_failed(
|
||||||
|
candidate_record_id, converted_error, elapsed_ms,
|
||||||
|
captured_key_concurrent, serializable_extra_data
|
||||||
|
)
|
||||||
|
raise converted_error
|
||||||
|
|
||||||
|
# 检查是否有请求体引用(由 Handler 层传入)
|
||||||
|
if request_body_ref is None:
|
||||||
|
logger.warning(f" [{request_id}] Thinking 错误:无法获取请求体引用,终止重试")
|
||||||
|
self._mark_thinking_error_failed(
|
||||||
|
candidate_record_id, converted_error, elapsed_ms,
|
||||||
|
captured_key_concurrent, serializable_extra_data
|
||||||
|
)
|
||||||
|
raise converted_error
|
||||||
|
|
||||||
|
# 检查是否已整流过(避免无限循环,单次重试)
|
||||||
|
if request_body_ref.get("_rectified", False):
|
||||||
|
logger.warning(f" [{request_id}] Thinking 错误:已整流仍失败,终止重试")
|
||||||
|
self._mark_thinking_error_failed(
|
||||||
|
candidate_record_id, converted_error, elapsed_ms,
|
||||||
|
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
|
||||||
|
)
|
||||||
|
raise converted_error
|
||||||
|
|
||||||
|
# 使用整流器
|
||||||
|
request_body = request_body_ref.get("body", {})
|
||||||
|
rectified_body, modified = ThinkingRectifier.rectify(request_body)
|
||||||
|
|
||||||
|
if modified:
|
||||||
|
# 更新容器中的请求体
|
||||||
|
request_body_ref["body"] = rectified_body
|
||||||
|
# _rectified: 全局标记,防止重复整流(整流只执行一次)
|
||||||
|
request_body_ref["_rectified"] = True
|
||||||
|
# _rectified_this_turn: 单轮标记,用于在当前 candidate 扩展重试次数
|
||||||
|
request_body_ref["_rectified_this_turn"] = True
|
||||||
|
|
||||||
|
logger.info(f" [{request_id}] 请求已整流,在当前候选上重试")
|
||||||
|
|
||||||
|
# 标记当前尝试为失败(整流前的状态)
|
||||||
|
# 注意:整流后重试会复用此记录 ID,成功时会覆盖为 success 状态
|
||||||
|
self._mark_thinking_error_failed(
|
||||||
|
candidate_record_id, converted_error, elapsed_ms,
|
||||||
|
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
|
||||||
|
)
|
||||||
|
|
||||||
|
# 返回 continue:在当前候选的重试循环中继续,使用整流后的请求体重试
|
||||||
|
return "continue"
|
||||||
|
else:
|
||||||
|
logger.warning(f" [{request_id}] Thinking 错误:无可整流内容")
|
||||||
|
self._mark_thinking_error_failed(
|
||||||
|
candidate_record_id, converted_error, elapsed_ms,
|
||||||
|
captured_key_concurrent, serializable_extra_data
|
||||||
|
)
|
||||||
|
raise converted_error
|
||||||
|
|
||||||
|
def _mark_thinking_error_failed(
|
||||||
|
self,
|
||||||
|
candidate_record_id: str,
|
||||||
|
error: ThinkingSignatureException,
|
||||||
|
elapsed_ms: int,
|
||||||
|
captured_key_concurrent: Optional[int],
|
||||||
|
extra_data: Dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
"""标记 Thinking 签名错误导致的候选失败"""
|
||||||
|
RequestCandidateService.mark_candidate_failed(
|
||||||
|
db=self.db,
|
||||||
|
candidate_id=candidate_record_id,
|
||||||
|
error_type="ThinkingSignatureException",
|
||||||
|
error_message=str(error),
|
||||||
|
status_code=400,
|
||||||
|
latency_ms=elapsed_ms,
|
||||||
|
concurrent_requests=captured_key_concurrent,
|
||||||
|
extra_data=extra_data,
|
||||||
|
)
|
||||||
|
|
||||||
async def _handle_candidate_error(
|
async def _handle_candidate_error(
|
||||||
self,
|
self,
|
||||||
exec_err: ExecutionError,
|
exec_err: ExecutionError,
|
||||||
@@ -311,6 +422,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: Optional[str],
|
request_id: Optional[str],
|
||||||
attempt: int,
|
attempt: int,
|
||||||
max_attempts: int,
|
max_attempts: int,
|
||||||
|
request_body_ref: Optional[Dict[str, Any]] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
处理候选执行错误
|
处理候选执行错误
|
||||||
@@ -327,6 +439,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: 请求 ID
|
request_id: 请求 ID
|
||||||
attempt: 当前尝试次数
|
attempt: 当前尝试次数
|
||||||
max_attempts: 最大尝试次数
|
max_attempts: 最大尝试次数
|
||||||
|
request_body_ref: 请求体引用容器(用于 Thinking 签名错误重试)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
action: "continue" (继续重试), "break" (跳到下一个候选), "raise" (抛出异常)
|
action: "continue" (继续重试), "break" (跳到下一个候选), "raise" (抛出异常)
|
||||||
@@ -431,6 +544,22 @@ class FallbackOrchestrator:
|
|||||||
k: v for k, v in extra_data.items() if k != "converted_error"
|
k: v for k, v in extra_data.items() if k != "converted_error"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 先检查 ThinkingSignatureException(它继承自 UpstreamClientException)
|
||||||
|
# 签名/结构错误需要特殊处理:整流请求体后重试一次
|
||||||
|
if isinstance(converted_error, ThinkingSignatureException):
|
||||||
|
action = self._handle_thinking_signature_error(
|
||||||
|
converted_error=converted_error,
|
||||||
|
request_id=request_id,
|
||||||
|
candidate_record_id=candidate_record_id,
|
||||||
|
elapsed_ms=elapsed_ms,
|
||||||
|
captured_key_concurrent=captured_key_concurrent,
|
||||||
|
serializable_extra_data=serializable_extra_data,
|
||||||
|
request_body_ref=request_body_ref,
|
||||||
|
)
|
||||||
|
if action == "continue":
|
||||||
|
return "continue"
|
||||||
|
# action == "raise" 时已在方法内部 raise
|
||||||
|
|
||||||
if isinstance(converted_error, UpstreamClientException):
|
if isinstance(converted_error, UpstreamClientException):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f" [{request_id}] 客户端请求错误,停止重试: {converted_error.message}"
|
f" [{request_id}] 客户端请求错误,停止重试: {converted_error.message}"
|
||||||
@@ -562,6 +691,7 @@ class FallbackOrchestrator:
|
|||||||
affinity_key: str,
|
affinity_key: str,
|
||||||
global_model_id: str,
|
global_model_id: str,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
|
request_body_ref: Optional[Dict[str, Any]] = None,
|
||||||
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
||||||
"""遍历所有候选执行请求,返回第一个成功的结果或抛出异常"""
|
"""遍历所有候选执行请求,返回第一个成功的结果或抛出异常"""
|
||||||
attempt_counter = 0
|
attempt_counter = 0
|
||||||
@@ -593,6 +723,7 @@ class FallbackOrchestrator:
|
|||||||
attempt_counter=attempt_counter,
|
attempt_counter=attempt_counter,
|
||||||
max_attempts=max_attempts,
|
max_attempts=max_attempts,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
|
request_body_ref=request_body_ref,
|
||||||
)
|
)
|
||||||
|
|
||||||
if result["success"]:
|
if result["success"]:
|
||||||
@@ -632,6 +763,7 @@ class FallbackOrchestrator:
|
|||||||
attempt_counter: int,
|
attempt_counter: int,
|
||||||
max_attempts: int,
|
max_attempts: int,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
|
request_body_ref: Optional[Dict[str, Any]] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""尝试单个候选(含重试逻辑),返回执行结果"""
|
"""尝试单个候选(含重试逻辑),返回执行结果"""
|
||||||
provider = candidate.provider
|
provider = candidate.provider
|
||||||
@@ -640,7 +772,8 @@ class FallbackOrchestrator:
|
|||||||
max_retries_for_candidate = int(provider.max_retries or 2) if candidate.is_cached else 1
|
max_retries_for_candidate = int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||||
last_error: Optional[Exception] = None
|
last_error: Optional[Exception] = None
|
||||||
|
|
||||||
for retry_index in range(max_retries_for_candidate):
|
retry_index = 0
|
||||||
|
while retry_index < max_retries_for_candidate:
|
||||||
attempt_counter += 1
|
attempt_counter += 1
|
||||||
max_attempts = max(max_attempts, attempt_counter)
|
max_attempts = max(max_attempts, attempt_counter)
|
||||||
|
|
||||||
@@ -655,7 +788,20 @@ class FallbackOrchestrator:
|
|||||||
f" [{request_id[:8] if request_id else 'N/A'}] -> {provider.name} (retry {retry_index})"
|
f" [{request_id[:8] if request_id else 'N/A'}] -> {provider.name} (retry {retry_index})"
|
||||||
)
|
)
|
||||||
|
|
||||||
candidate_record_id = candidate_record_map[(candidate_index, retry_index)]
|
# 获取候选记录 ID
|
||||||
|
# 正常情况下 record_key = (candidate_index, retry_index)
|
||||||
|
# 整流重试时 retry_index 可能超出预创建范围,复用最后一个有效记录
|
||||||
|
record_key = (candidate_index, retry_index)
|
||||||
|
if record_key not in candidate_record_map:
|
||||||
|
# 整流重试:复用该候选的最后一个有效记录(通常是 retry_index=0)
|
||||||
|
# 这样整流后的成功/失败会覆盖之前的记录状态
|
||||||
|
fallback_key = (candidate_index, 0)
|
||||||
|
candidate_record_id = candidate_record_map.get(fallback_key, "")
|
||||||
|
logger.debug(
|
||||||
|
f" [{request_id}] 整流重试:复用记录 {fallback_key} -> {candidate_record_id[:8] if candidate_record_id else 'N/A'}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
candidate_record_id = candidate_record_map[record_key]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = await self._try_single_candidate(
|
response = await self._try_single_candidate(
|
||||||
@@ -690,9 +836,20 @@ class FallbackOrchestrator:
|
|||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
attempt=attempt_counter,
|
attempt=attempt_counter,
|
||||||
max_attempts=max_attempts,
|
max_attempts=max_attempts,
|
||||||
|
request_body_ref=request_body_ref,
|
||||||
)
|
)
|
||||||
|
|
||||||
if action == "continue":
|
if action == "continue":
|
||||||
|
# 检查是否刚完成整流,需要额外重试一次
|
||||||
|
if request_body_ref and request_body_ref.get("_rectified_this_turn", False):
|
||||||
|
# 清除标记,扩展重试上限以允许整流后的请求发出
|
||||||
|
# 使用 max() 确保不会减少已有的重试次数
|
||||||
|
request_body_ref["_rectified_this_turn"] = False
|
||||||
|
max_retries_for_candidate = max(max_retries_for_candidate, retry_index + 2)
|
||||||
|
logger.debug(
|
||||||
|
f" [{request_id}] 整流后扩展重试次数至 {max_retries_for_candidate}"
|
||||||
|
)
|
||||||
|
retry_index += 1
|
||||||
continue
|
continue
|
||||||
elif action == "break":
|
elif action == "break":
|
||||||
break
|
break
|
||||||
@@ -704,6 +861,10 @@ class FallbackOrchestrator:
|
|||||||
"attempt_counter": attempt_counter,
|
"attempt_counter": attempt_counter,
|
||||||
"max_attempts": max_attempts,
|
"max_attempts": max_attempts,
|
||||||
}
|
}
|
||||||
|
else:
|
||||||
|
# 未知 action,安全起见跳出循环
|
||||||
|
logger.warning(f" [{request_id}] 未知 action: {action},跳出重试")
|
||||||
|
break
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -826,6 +987,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: Optional[str] = None,
|
request_id: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
|
request_body_ref: Optional[Dict[str, Any]] = None,
|
||||||
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
||||||
"""
|
"""
|
||||||
执行请求,并在失败时自动故障转移(缓存感知)
|
执行请求,并在失败时自动故障转移(缓存感知)
|
||||||
@@ -838,6 +1000,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: 请求 ID(用于日志)
|
request_id: 请求 ID(用于日志)
|
||||||
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
||||||
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
||||||
|
request_body_ref: 请求体引用容器(用于 Thinking 签名错误重试)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(请求响应, 实际Provider名称, RequestTraceAttempt ID, provider_id, endpoint_id, key_id)
|
(请求响应, 实际Provider名称, RequestTraceAttempt ID, provider_id, endpoint_id, key_id)
|
||||||
@@ -895,4 +1058,5 @@ class FallbackOrchestrator:
|
|||||||
affinity_key=affinity_key,
|
affinity_key=affinity_key,
|
||||||
global_model_id=global_model_id,
|
global_model_id=global_model_id,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
|
request_body_ref=request_body_ref,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -165,6 +165,9 @@ class RequestCandidateService:
|
|||||||
candidate.latency_ms = latency_ms
|
candidate.latency_ms = latency_ms
|
||||||
candidate.concurrent_requests = concurrent_requests
|
candidate.concurrent_requests = concurrent_requests
|
||||||
candidate.finished_at = datetime.now(timezone.utc)
|
candidate.finished_at = datetime.now(timezone.utc)
|
||||||
|
# 成功时清空错误字段(可能是整流重试后成功,之前记录过错误)
|
||||||
|
candidate.error_type = None
|
||||||
|
candidate.error_message = None
|
||||||
if extra_data:
|
if extra_data:
|
||||||
candidate.extra_data = {**(candidate.extra_data or {}), **extra_data}
|
candidate.extra_data = {**(candidate.extra_data or {}), **extra_data}
|
||||||
# 关键状态更新:立即提交,不使用批量提交
|
# 关键状态更新:立即提交,不使用批量提交
|
||||||
|
|||||||
164
tests/services/test_error_classifier_thinking.py
Normal file
164
tests/services/test_error_classifier_thinking.py
Normal file
@@ -0,0 +1,164 @@
|
|||||||
|
"""
|
||||||
|
ErrorClassifier Thinking 错误模式测试
|
||||||
|
|
||||||
|
测试 ErrorClassifier 对 Thinking 相关错误的识别:
|
||||||
|
- 签名验证失败
|
||||||
|
- 结构错误(缺少 thinking 块前缀)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||||
|
|
||||||
|
|
||||||
|
class TestThinkingErrorPatterns:
|
||||||
|
"""测试 Thinking 错误模式匹配"""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def classifier(self) -> ErrorClassifier:
|
||||||
|
mock_db = MagicMock()
|
||||||
|
return ErrorClassifier(db=mock_db)
|
||||||
|
|
||||||
|
# === 签名错误测试 ===
|
||||||
|
|
||||||
|
def test_detect_invalid_signature_with_backticks(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测带反引号的签名错误"""
|
||||||
|
error = '{"error": {"message": "invalid `signature` in `thinking` block"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_invalid_signature_without_backticks(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测不带反引号的签名错误"""
|
||||||
|
error = '{"error": {"message": "invalid signature in thinking block"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_signature_field_required(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测签名字段缺失错误"""
|
||||||
|
error = '{"error": {"message": "thinking.signature: field required"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_signature_path_pattern(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测签名路径模式(如 messages.0.content.0.thinking.signature)"""
|
||||||
|
error = '{"error": {"message": "messages.0.content.0.thinking.signature: invalid"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_signature_verification_failed(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测签名验证失败"""
|
||||||
|
error = '{"error": {"message": "signature verification failed"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
# === 结构错误测试 ===
|
||||||
|
|
||||||
|
def test_detect_must_start_with_thinking_block(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测必须以 thinking 块开头的错误"""
|
||||||
|
error = '{"error": {"message": "content must start with a thinking block"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_or_redacted(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测期望 thinking 或 redacted_thinking 的错误"""
|
||||||
|
error = '{"error": {"message": "expected thinking or redacted_thinking"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_with_backticks(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测带反引号的 expected thinking 错误"""
|
||||||
|
error = '{"error": {"message": "expected `thinking`"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_found_tool_use(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测 expected thinking, found xxx 错误(统一模式匹配)"""
|
||||||
|
error = '{"error": {"message": "expected thinking, found tool_use"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_found_tool_use_backticks(
|
||||||
|
self, classifier: ErrorClassifier
|
||||||
|
) -> None:
|
||||||
|
"""检测带反引号的 expected `thinking`, found xxx 错误"""
|
||||||
|
error = '{"error": {"message": "expected `thinking`, found `tool_use`"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_found_text(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测 expected thinking, found text 错误"""
|
||||||
|
error = '{"error": {"message": "expected thinking, found text"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_thinking_found_text_backticks(
|
||||||
|
self, classifier: ErrorClassifier
|
||||||
|
) -> None:
|
||||||
|
"""检测带反引号的 expected `thinking`, found text 错误"""
|
||||||
|
error = '{"error": {"message": "expected `thinking`, found `text`"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_redacted_thinking(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""检测 expected redacted_thinking 错误"""
|
||||||
|
error = '{"error": {"message": "expected redacted_thinking, found text"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_detect_expected_redacted_thinking_backticks(
|
||||||
|
self, classifier: ErrorClassifier
|
||||||
|
) -> None:
|
||||||
|
"""检测带反引号的 expected redacted_thinking 错误"""
|
||||||
|
error = '{"error": {"message": "expected `redacted_thinking`, found `text`"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
# === 真实错误响应测试 ===
|
||||||
|
|
||||||
|
def test_real_claude_signature_error(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""测试真实的 Claude 签名错误响应"""
|
||||||
|
error = """{
|
||||||
|
"type": "error",
|
||||||
|
"error": {
|
||||||
|
"type": "invalid_request_error",
|
||||||
|
"message": "messages.2.content.0: invalid `signature` in `thinking` block: signature is for a different request"
|
||||||
|
}
|
||||||
|
}"""
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
def test_real_claude_structure_error(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""测试真实的 Claude 结构错误响应"""
|
||||||
|
error = """{
|
||||||
|
"type": "error",
|
||||||
|
"error": {
|
||||||
|
"type": "invalid_request_error",
|
||||||
|
"message": "messages.1.content: when `thinking` is enabled, the first content block in an assistant turn containing `tool_use` blocks must start with a `thinking` or `redacted_thinking` block. expected thinking or redacted_thinking, found tool_use."
|
||||||
|
}
|
||||||
|
}"""
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
# === 负面测试 ===
|
||||||
|
|
||||||
|
def test_not_thinking_error_rate_limit(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""速率限制错误不应被识别为 thinking 错误"""
|
||||||
|
error = '{"error": {"message": "rate limit exceeded"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is False
|
||||||
|
|
||||||
|
def test_not_thinking_error_invalid_api_key(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""API Key 错误不应被识别为 thinking 错误"""
|
||||||
|
error = '{"error": {"message": "invalid api key"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is False
|
||||||
|
|
||||||
|
def test_not_thinking_error_model_not_found(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""模型不存在错误不应被识别为 thinking 错误"""
|
||||||
|
error = '{"error": {"message": "model not found"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is False
|
||||||
|
|
||||||
|
def test_not_thinking_error_empty(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""空错误文本不应被识别为 thinking 错误"""
|
||||||
|
assert classifier._is_thinking_error("") is False
|
||||||
|
assert classifier._is_thinking_error(None) is False
|
||||||
|
|
||||||
|
def test_not_thinking_error_generic_signature(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""通用 signature 字样(非 thinking 相关)不应被误判"""
|
||||||
|
# 这个应该不匹配,因为模式要求是 thinking 相关的 signature
|
||||||
|
error = '{"error": {"message": "request signature mismatch"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is False
|
||||||
|
|
||||||
|
# === 大小写不敏感测试 ===
|
||||||
|
|
||||||
|
def test_case_insensitive_matching(self, classifier: ErrorClassifier) -> None:
|
||||||
|
"""测试大小写不敏感匹配"""
|
||||||
|
error = '{"error": {"message": "INVALID `SIGNATURE` IN `THINKING` BLOCK"}}'
|
||||||
|
assert classifier._is_thinking_error(error) is True
|
||||||
|
|
||||||
|
error2 = '{"error": {"message": "Expected Thinking, Found Tool_Use"}}'
|
||||||
|
assert classifier._is_thinking_error(error2) is True
|
||||||
378
tests/services/test_thinking_rectifier.py
Normal file
378
tests/services/test_thinking_rectifier.py
Normal file
@@ -0,0 +1,378 @@
|
|||||||
|
"""
|
||||||
|
ThinkingRectifier 单元测试
|
||||||
|
|
||||||
|
测试 Thinking 整流器的核心功能:
|
||||||
|
- 移除 thinking 和 redacted_thinking 块
|
||||||
|
- 移除非 thinking 块上的 signature 字段
|
||||||
|
- 条件删除顶层 thinking 参数
|
||||||
|
"""
|
||||||
|
|
||||||
|
import copy
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.message.thinking_rectifier import ThinkingRectifier
|
||||||
|
|
||||||
|
|
||||||
|
class TestRectifyBasic:
|
||||||
|
"""测试基本整流功能"""
|
||||||
|
|
||||||
|
def test_empty_request_body(self) -> None:
|
||||||
|
"""空请求体应返回原值"""
|
||||||
|
result, modified = ThinkingRectifier.rectify({})
|
||||||
|
assert result == {}
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_none_request_body(self) -> None:
|
||||||
|
"""None 请求体应返回原值"""
|
||||||
|
result, modified = ThinkingRectifier.rectify(None) # type: ignore
|
||||||
|
assert result is None
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_no_messages(self) -> None:
|
||||||
|
"""无 messages 字段应返回原值"""
|
||||||
|
body = {"model": "claude-3-opus"}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
assert result == body
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_empty_messages(self) -> None:
|
||||||
|
"""空 messages 列表应返回原值"""
|
||||||
|
body = {"model": "claude-3-opus", "messages": []}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
assert result["messages"] == []
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestRemoveThinkingBlocks:
|
||||||
|
"""测试移除 thinking 块"""
|
||||||
|
|
||||||
|
def test_remove_thinking_block(self) -> None:
|
||||||
|
"""应移除 thinking 块"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "...", "signature": "abc"},
|
||||||
|
{"type": "text", "text": "Hello"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
content = result["messages"][0]["content"]
|
||||||
|
assert len(content) == 1
|
||||||
|
assert content[0]["type"] == "text"
|
||||||
|
|
||||||
|
def test_remove_redacted_thinking_block(self) -> None:
|
||||||
|
"""应移除 redacted_thinking 块"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "redacted_thinking", "data": "..."},
|
||||||
|
{"type": "text", "text": "Hello"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
content = result["messages"][0]["content"]
|
||||||
|
assert len(content) == 1
|
||||||
|
assert content[0]["type"] == "text"
|
||||||
|
|
||||||
|
def test_remove_multiple_thinking_blocks(self) -> None:
|
||||||
|
"""应移除多个 thinking 块"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "thought 1"},
|
||||||
|
{"type": "text", "text": "response 1"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "redacted_thinking", "data": "..."},
|
||||||
|
{"type": "text", "text": "response 2"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
assert len(result["messages"][0]["content"]) == 1
|
||||||
|
assert len(result["messages"][1]["content"]) == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestRemoveSignatureField:
|
||||||
|
"""测试移除 signature 字段"""
|
||||||
|
|
||||||
|
def test_remove_signature_from_text_block(self) -> None:
|
||||||
|
"""应从 text 块移除 signature 字段"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "Hello", "signature": "should_remove"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
block = result["messages"][0]["content"][0]
|
||||||
|
assert "signature" not in block
|
||||||
|
assert block["text"] == "Hello"
|
||||||
|
|
||||||
|
def test_remove_signature_from_tool_use_block(self) -> None:
|
||||||
|
"""应从 tool_use 块移除 signature 字段"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "tool_use",
|
||||||
|
"id": "tool_1",
|
||||||
|
"name": "search",
|
||||||
|
"input": {},
|
||||||
|
"signature": "should_remove",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
block = result["messages"][0]["content"][0]
|
||||||
|
assert "signature" not in block
|
||||||
|
assert block["id"] == "tool_1"
|
||||||
|
|
||||||
|
|
||||||
|
class TestTopLevelThinkingParam:
|
||||||
|
"""测试顶层 thinking 参数处理"""
|
||||||
|
|
||||||
|
def test_remove_thinking_param_when_tool_use_without_thinking_prefix(self) -> None:
|
||||||
|
"""整流后有 tool_use 但无 thinking 前缀时应移除 thinking 参数"""
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "enabled", "budget_tokens": 10000},
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "tool_use", "id": "t1", "name": "search", "input": {}},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
assert "thinking" not in result
|
||||||
|
# tool_use 应保留
|
||||||
|
assert result["messages"][0]["content"][0]["type"] == "tool_use"
|
||||||
|
|
||||||
|
def test_keep_thinking_param_when_no_tool_use(self) -> None:
|
||||||
|
"""整流后无 tool_use 时应保留 thinking 参数"""
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "enabled", "budget_tokens": 10000},
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "text", "text": "Hello"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
# thinking 参数应保留(只移除了 thinking 块)
|
||||||
|
assert "thinking" in result
|
||||||
|
assert result["thinking"]["type"] == "enabled"
|
||||||
|
|
||||||
|
def test_keep_thinking_param_when_disabled(self) -> None:
|
||||||
|
"""thinking 参数未启用时应保留"""
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "disabled"},
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "tool_use", "id": "t1", "name": "search", "input": {}},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
# 无 thinking 块需要移除,thinking 参数也不需要移除
|
||||||
|
assert modified is False
|
||||||
|
assert result["thinking"]["type"] == "disabled"
|
||||||
|
|
||||||
|
|
||||||
|
class TestEdgeCases:
|
||||||
|
"""测试边界情况"""
|
||||||
|
|
||||||
|
def test_non_dict_message_preserved(self) -> None:
|
||||||
|
"""非 dict 消息应保留"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
"string message", # 非 dict
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert result["messages"][0] == "string message"
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_string_content_preserved(self) -> None:
|
||||||
|
"""字符串 content 应保留"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
{"role": "assistant", "content": "Hi there"},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert result["messages"][1]["content"] == "Hi there"
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_non_dict_block_in_content_preserved(self) -> None:
|
||||||
|
"""content 中的非 dict 块应保留"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
"string block", # 非 dict
|
||||||
|
{"type": "text", "text": "Hello"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert result["messages"][0]["content"][0] == "string block"
|
||||||
|
assert modified is False
|
||||||
|
|
||||||
|
def test_empty_content_after_rectify_logs_warning(self) -> None:
|
||||||
|
"""整流后 assistant 消息 content 为空时应记录警告(不抛异常)"""
|
||||||
|
body = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "only thinking"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
# 应该不抛异常
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
# content 应为空列表
|
||||||
|
assert result["messages"][0]["content"] == []
|
||||||
|
|
||||||
|
def test_deep_copy_does_not_modify_original(self) -> None:
|
||||||
|
"""深拷贝应保护原始数据不被修改"""
|
||||||
|
original = {
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "text", "text": "Hello"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
original_copy = copy.deepcopy(original)
|
||||||
|
|
||||||
|
ThinkingRectifier.rectify(original)
|
||||||
|
|
||||||
|
# 原始数据不应被修改
|
||||||
|
assert original == original_copy
|
||||||
|
|
||||||
|
|
||||||
|
class TestLastAssistantCheck:
|
||||||
|
"""测试只检查最后一条 assistant 消息的逻辑"""
|
||||||
|
|
||||||
|
def test_only_last_assistant_matters_for_thinking_removal(self) -> None:
|
||||||
|
"""只有最后一条 assistant 消息决定是否移除 thinking 参数"""
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "enabled", "budget_tokens": 10000},
|
||||||
|
"messages": [
|
||||||
|
# 第一条 assistant 有 tool_use
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "tool_use", "id": "t1", "name": "search", "input": {}},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{"role": "user", "content": "result"},
|
||||||
|
# 最后一条 assistant 无 tool_use
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "text", "text": "Final answer"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
# thinking 参数应保留(最后一条 assistant 无 tool_use)
|
||||||
|
assert "thinking" in result
|
||||||
|
|
||||||
|
def test_last_assistant_with_tool_use_removes_thinking(self) -> None:
|
||||||
|
"""最后一条 assistant 有 tool_use 时应移除 thinking 参数"""
|
||||||
|
body = {
|
||||||
|
"thinking": {"type": "enabled", "budget_tokens": 10000},
|
||||||
|
"messages": [
|
||||||
|
# 第一条 assistant 无 tool_use
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "text", "text": "thinking..."},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{"role": "user", "content": "continue"},
|
||||||
|
# 最后一条 assistant 有 tool_use
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "..."},
|
||||||
|
{"type": "tool_use", "id": "t1", "name": "search", "input": {}},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result, modified = ThinkingRectifier.rectify(body)
|
||||||
|
|
||||||
|
assert modified is True
|
||||||
|
# thinking 参数应被移除
|
||||||
|
assert "thinking" not in result
|
||||||
Reference in New Issue
Block a user