fix: 流式响应健壮性增强与候选排序稳定性修复

- test-model 接口优先使用流式请求,失败自动回退非流式
- CLI 流式处理在 StreamClosed/RemoteProtocolError 时 flush 残余 SSE 数据,捕获尾部 usage
- 流未正常完成时兜底估算 tokens,避免 usage 记录为 0
- ProviderCandidate 添加 __lt__ 解决 tuple 排序 TypeError
- sorted() 添加显式 key 参数避免隐式比较候选对象
- 异常日志改用 logger.opt(exception=e) 替代手动 traceback
This commit is contained in:
fawney19
2026-02-06 18:06:38 +08:00
parent b88fb6273b
commit 62dae22a2c
6 changed files with 306 additions and 49 deletions

View File

@@ -590,7 +590,7 @@ async def test_model(
logger.debug(f"[test-model] 端点 API Format: {endpoint.api_format}") logger.debug(f"[test-model] 端点 API Format: {endpoint.api_format}")
logger.debug(f"[test-model] 使用 Key: {api_key.name or api_key.id} (auth_type={auth_type})") logger.debug(f"[test-model] 使用 Key: {api_key.name or api_key.id} (auth_type={auth_type})")
# 准备测试请求数据 # 准备测试请求数据(优先使用流式)
check_request = { check_request = {
"model": request.model_name, "model": request.model_name,
"messages": [ "messages": [
@@ -598,6 +598,7 @@ async def test_model(
], ],
"max_tokens": 30, "max_tokens": 30,
"temperature": 0.7, "temperature": 0.7,
"stream": True,
} }
# 获取端点规则(不在此处应用,传递给 check_endpoint 在格式转换后应用) # 获取端点规则(不在此处应用,传递给 check_endpoint 在格式转换后应用)
@@ -619,28 +620,58 @@ async def test_model(
# Provider 上下文auth_type 用于 OAuth 认证头处理provider_type 用于特殊路由 # Provider 上下文auth_type 用于 OAuth 认证头处理provider_type 用于特殊路由
p_type = str(getattr(provider, "provider_type", "") or "").lower() p_type = str(getattr(provider, "provider_type", "") or "").lower()
response = await adapter_class.check_endpoint( async def _do_check(req: dict) -> dict:
return await adapter_class.check_endpoint(
client, client,
endpoint_config["base_url"], endpoint_config["base_url"],
endpoint_config["api_key"], endpoint_config["api_key"],
check_request, req,
extra_headers if extra_headers else None, extra_headers if extra_headers else None,
# 端点规则(在 check_endpoint 内部格式转换后应用)
body_rules=body_rules, body_rules=body_rules,
header_rules=header_rules, header_rules=header_rules,
# 用量计算参数(现在强制记录)
db=db, db=db,
user=current_user, user=current_user,
provider_name=provider.name, provider_name=provider.name,
provider_id=provider.id, provider_id=provider.id,
api_key_id=endpoint_config.get("api_key_id"), api_key_id=endpoint_config.get("api_key_id"),
model_name=request.model_name, model_name=request.model_name,
# Provider 上下文
auth_type=auth_type, auth_type=auth_type,
provider_type=p_type if p_type else None, provider_type=p_type if p_type else None,
decrypted_auth_config=oauth_meta if oauth_meta else None, decrypted_auth_config=oauth_meta if oauth_meta else None,
) )
def _response_has_error(resp: dict) -> bool:
"""快速判断响应是否包含错误"""
if "error" in resp:
return True
if resp.get("status_code", 0) != 200:
return True
resp_data = resp.get("response", {})
resp_body = resp_data.get("response_body", {})
parsed = resp_body
if isinstance(resp_body, str):
try:
parsed = json.loads(resp_body)
except (json.JSONDecodeError, ValueError):
pass
if isinstance(parsed, dict) and "error" in parsed:
return True
return False
# 策略:优先流式,若失败回退到非流式
used_stream = True
logger.debug("[test-model] 尝试流式请求...")
response = await _do_check(check_request)
if _response_has_error(response):
logger.info(
"[test-model] 流式请求失败 (status={}),回退到非流式请求",
response.get("status_code", "?"),
)
check_request["stream"] = False
used_stream = False
response = await _do_check(check_request)
# 记录提供商返回信息 # 记录提供商返回信息
logger.debug("[test-model] 端点测试结果:") logger.debug("[test-model] 端点测试结果:")
logger.debug(f"[test-model] Status Code: {response.get('status_code')}") logger.debug(f"[test-model] Status Code: {response.get('status_code')}")
@@ -745,7 +776,7 @@ async def test_model(
return { return {
"success": is_success, "success": is_success,
"data": { "data": {
"stream": False, "stream": used_stream,
"response": response, "response": response,
}, },
"provider": { "provider": {

View File

@@ -19,7 +19,6 @@ Chat Adapter 通用基类
from __future__ import annotations from __future__ import annotations
import time import time
import traceback
from abc import abstractmethod from abc import abstractmethod
from typing import Any, ClassVar from typing import Any, ClassVar
@@ -438,18 +437,10 @@ class ChatAdapterBase(ApiAdapter):
) -> JSONResponse: ) -> JSONResponse:
"""处理未预期的异常""" """处理未预期的异常"""
if isinstance(e, ProxyException): if isinstance(e, ProxyException):
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}") logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}: {e}")
else: else:
logger.error( logger.opt(exception=e).error(
f"{self.FORMAT_ID} 请求处理意外异常", f"{self.FORMAT_ID} 请求处理意外异常: {type(e).__name__}: {e}"
exception=e,
extra_data={
"exception_class": e.__class__.__name__,
"processing_stage": "request_processing",
"model": model,
"stream": stream,
"traceback_preview": str(traceback.format_exc())[:500],
},
) )
response_time = int((time.time() - start_time) * 1000) response_time = int((time.time() - start_time) * 1000)

View File

@@ -18,7 +18,6 @@ CLI Adapter 通用基类
from __future__ import annotations from __future__ import annotations
import time import time
import traceback
from typing import Any, ClassVar from typing import Any, ClassVar
import httpx import httpx
@@ -417,18 +416,10 @@ class CliAdapterBase(ApiAdapter):
) -> JSONResponse: ) -> JSONResponse:
"""处理未预期的异常""" """处理未预期的异常"""
if isinstance(e, ProxyException): if isinstance(e, ProxyException):
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}") logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}: {e}")
else: else:
logger.error( logger.opt(exception=e).error(
f"{self.FORMAT_ID} 请求处理意外异常", f"{self.FORMAT_ID} 请求处理意外异常: {type(e).__name__}: {e}"
exception=e,
extra_data={
"exception_class": e.__class__.__name__,
"processing_stage": "request_processing",
"model": model,
"stream": stream,
"traceback_preview": str(traceback.format_exc())[:500],
},
) )
response_time = int((time.time() - start_time) * 1000) response_time = int((time.time() - start_time) * 1000)

View File

@@ -1353,6 +1353,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
except GeneratorExit: except GeneratorExit:
raise raise
except httpx.StreamClosed: except httpx.StreamClosed:
# 连接关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count == 0: if ctx.data_count == 0:
# 流已开始,发送错误事件而不是抛出异常 # 流已开始,发送错误事件而不是抛出异常
logger.warning(f"Provider '{ctx.provider_name}' 流连接关闭且无数据") logger.warning(f"Provider '{ctx.provider_name}' 流连接关闭且无数据")
@@ -1369,6 +1373,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
} }
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode() yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError: except httpx.RemoteProtocolError:
# 连接异常关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0: if ctx.data_count > 0:
error_event = { error_event = {
"type": "error", "type": "error",
@@ -1390,6 +1398,106 @@ class CliMessageHandlerBase(BaseMessageHandler):
except Exception: except Exception:
pass pass
def _flush_remaining_sse_data(
self,
ctx: StreamContext,
buffer: bytes,
decoder: codecs.IncrementalDecoder,
sse_parser: SSEEventParser,
*,
record_chunk: bool = True,
) -> None:
"""
异常发生时 flush 残留的字节 buffer 和 SSE parser 内部缓冲区。
用于 StreamClosed / RemoteProtocolError 等场景:
连接断开可能恰好发生在最后一个 SSE 事件(如 response.completed
的 data 行已收到、但终止空行尚未到达之时。此方法确保这些事件仍能被处理,
从而正确捕获 usage 等关键信息。
"""
try:
# 1) flush 字节 buffer 中的残余行
if buffer:
remaining = decoder.decode(buffer, True)
for line in remaining.split("\n"):
stripped = line.rstrip("\r")
events = sse_parser.feed_line(stripped)
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
# 2) flush SSE parser 内部累积的未完成事件
for event in sse_parser.flush():
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
except Exception:
# best-effort: 不应因 flush 失败影响后续流程
pass
def _estimate_tokens_for_incomplete_stream(
self,
ctx: StreamContext,
request_body: dict[str, Any],
) -> None:
"""
流未正常完成(无 response.completed且 token 均为 0 时的兜底估算。
从已收集的输出文本和请求体粗略估算 token 数,确保 usage 记录不为 0。
估算采用 ~4 字符/token 的保守比例。
"""
# 输出 tokens从已收集的文本估算
collected = ctx.collected_text
if collected:
ctx.output_tokens = max(1, len(collected) // 4)
# 输入 tokens从请求体文本内容估算
try:
total_input_len = 0
instructions = request_body.get("instructions")
if isinstance(instructions, str):
total_input_len += len(instructions)
# OpenAI Responses API 使用 input 字段Claude 使用 messages
input_items = request_body.get("input") or request_body.get("messages") or []
if isinstance(input_items, list):
for item in input_items:
if isinstance(item, str):
total_input_len += len(item)
elif isinstance(item, dict):
content = item.get("content", "")
if isinstance(content, str):
total_input_len += len(content)
elif isinstance(content, list):
for block in content:
if isinstance(block, dict):
text = block.get("text", "")
if isinstance(text, str):
total_input_len += len(text)
if total_input_len > 0:
ctx.input_tokens = max(1, total_input_len // 4)
else:
# fallback: 整个请求体 JSON 大小
body_str = json.dumps(request_body, ensure_ascii=False)
ctx.input_tokens = max(1, len(body_str) // 4)
except Exception:
pass
if ctx.input_tokens > 0 or ctx.output_tokens > 0:
logger.warning(
"[{}] 流未正常完成 (has_completion=False, data_count={}), "
"使用估算 tokens: in={}, out={}",
ctx.request_id,
ctx.data_count,
ctx.input_tokens,
ctx.output_tokens,
)
async def _prefetch_and_check_embedded_error( async def _prefetch_and_check_embedded_error(
self, self,
byte_iterator: Any, byte_iterator: Any,
@@ -1787,6 +1895,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
except GeneratorExit: except GeneratorExit:
raise raise
except httpx.StreamClosed: except httpx.StreamClosed:
# 连接关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count == 0: if ctx.data_count == 0:
logger.warning(f"Provider '{ctx.provider_name}' 流连接关闭且无数据") logger.warning(f"Provider '{ctx.provider_name}' 流连接关闭且无数据")
# 设置错误状态用于后续记录 # 设置错误状态用于后续记录
@@ -1802,6 +1914,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
} }
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode() yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError: except httpx.RemoteProtocolError:
# 连接异常关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0: if ctx.data_count > 0:
error_event = { error_event = {
"type": "error", "type": "error",
@@ -2357,6 +2473,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据 # 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
self._finalize_stream_metadata(ctx) self._finalize_stream_metadata(ctx)
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
# 从已收集的文本和请求体估算 tokens避免 usage 记录为 0
if (
not ctx.has_completion
and ctx.data_count > 0
and ctx.input_tokens == 0
and ctx.output_tokens == 0
):
self._estimate_tokens_for_incomplete_stream(ctx, actual_request_body)
# 流式成功时,返回给客户端的是提供商响应头 + 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(

View File

@@ -83,6 +83,43 @@ class ProviderCandidate:
needs_conversion: bool = False # 是否需要格式转换 needs_conversion: bool = False # 是否需要格式转换
provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket
def _stable_order_key(self) -> tuple[int, int, str, str, str]:
"""
为排序/优先队列提供稳定的比较键。
说明:
- 运行时偶发会出现对 ProviderCandidate 做 tuple 排序/heap 排序的场景;
当主键相同需要比较候选本身时,若候选不可比较会触发:
TypeError: '<' not supported between instances of 'ProviderCandidate' and 'ProviderCandidate'
- 这里提供一个与调度逻辑无关、但足够稳定且可比的兜底顺序。
"""
provider_priority_raw = getattr(self.provider, "provider_priority", None)
internal_priority_raw = getattr(self.key, "internal_priority", None)
try:
provider_priority = (
int(provider_priority_raw) if provider_priority_raw is not None else 999999
)
except Exception:
provider_priority = 999999
try:
internal_priority = (
int(internal_priority_raw) if internal_priority_raw is not None else 999999
)
except Exception:
internal_priority = 999999
provider_id = str(getattr(self.provider, "id", "") or "")
endpoint_id = str(getattr(self.endpoint, "id", "") or "")
key_id = str(getattr(self.key, "id", "") or "")
return (provider_priority, internal_priority, provider_id, endpoint_id, key_id)
def __lt__(self, other: object) -> bool:
if not isinstance(other, ProviderCandidate):
return NotImplemented
return self._stable_order_key() < other._stable_order_key()
@dataclass @dataclass
class ConcurrencySnapshot: class ConcurrencySnapshot:
@@ -1585,7 +1622,7 @@ class CacheAwareScheduler:
scored_candidates.append((hash_value, candidate)) scored_candidates.append((hash_value, candidate))
# 按哈希值排序 # 按哈希值排序
sorted_group = [c for _, c in sorted(scored_candidates)] sorted_group = [c for _, c in sorted(scored_candidates, key=lambda x: x[0])]
result.extend(sorted_group) result.extend(sorted_group)
else: else:
# 单个候选或没有 affinity_key按次要排序条件排序 # 单个候选或没有 affinity_key按次要排序条件排序
@@ -1706,7 +1743,7 @@ class CacheAwareScheduler:
key_scores.append((hash_value, key)) key_scores.append((hash_value, key))
# 按哈希值排序 # 按哈希值排序
sorted_group = [key for _, key in sorted(key_scores)] sorted_group = [key for _, key in sorted(key_scores, key=lambda x: x[0])]
result.extend(sorted_group) result.extend(sorted_group)
else: else:
# 没有 affinity_key 时按 ID 排序保持稳定性 # 没有 affinity_key 时按 ID 排序保持稳定性

View File

@@ -0,0 +1,81 @@
from __future__ import annotations
from types import SimpleNamespace
from src.services.cache.aware_scheduler import ProviderCandidate
def _make_candidate(
*,
provider_priority: int | None,
internal_priority: int | None,
provider_id: str,
endpoint_id: str,
key_id: str,
) -> ProviderCandidate:
provider = SimpleNamespace(id=provider_id, name="p", provider_priority=provider_priority)
endpoint = SimpleNamespace(id=endpoint_id)
key = SimpleNamespace(id=key_id, internal_priority=internal_priority)
return ProviderCandidate(provider=provider, endpoint=endpoint, key=key) # type: ignore[arg-type]
def test_provider_candidate_is_orderable() -> None:
c1 = _make_candidate(
provider_priority=1,
internal_priority=1,
provider_id="p1",
endpoint_id="e1",
key_id="k1",
)
c2 = _make_candidate(
provider_priority=1,
internal_priority=2,
provider_id="p1",
endpoint_id="e1",
key_id="k2",
)
assert c1 < c2
assert sorted([c2, c1]) == [c1, c2]
def test_provider_candidate_can_break_ties_in_tuple_sort() -> None:
c1 = _make_candidate(
provider_priority=1,
internal_priority=1,
provider_id="p1",
endpoint_id="e1",
key_id="k1",
)
c2 = _make_candidate(
provider_priority=1,
internal_priority=2,
provider_id="p1",
endpoint_id="e1",
key_id="k2",
)
# When the first tuple element ties, Python will compare the second element.
# This should not raise:
# TypeError: '<' not supported between instances of 'ProviderCandidate' and 'ProviderCandidate'
pairs = [(0, c2), (0, c1)]
assert [c for _score, c in sorted(pairs)] == [c1, c2]
def test_provider_candidate_none_priorities_sort_last() -> None:
c1 = _make_candidate(
provider_priority=1,
internal_priority=None,
provider_id="p1",
endpoint_id="e1",
key_id="k1",
)
c2 = _make_candidate(
provider_priority=None,
internal_priority=None,
provider_id="p2",
endpoint_id="e1",
key_id="k2",
)
assert sorted([c2, c1]) == [c1, c2]