mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -590,7 +590,7 @@ async def test_model(
|
||||
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})")
|
||||
|
||||
# 准备测试请求数据
|
||||
# 准备测试请求数据(优先使用流式)
|
||||
check_request = {
|
||||
"model": request.model_name,
|
||||
"messages": [
|
||||
@@ -598,6 +598,7 @@ async def test_model(
|
||||
],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
# 获取端点规则(不在此处应用,传递给 check_endpoint 在格式转换后应用)
|
||||
@@ -619,27 +620,57 @@ async def test_model(
|
||||
# Provider 上下文:auth_type 用于 OAuth 认证头处理,provider_type 用于特殊路由
|
||||
p_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
|
||||
response = await adapter_class.check_endpoint(
|
||||
client,
|
||||
endpoint_config["base_url"],
|
||||
endpoint_config["api_key"],
|
||||
check_request,
|
||||
extra_headers if extra_headers else None,
|
||||
# 端点规则(在 check_endpoint 内部格式转换后应用)
|
||||
body_rules=body_rules,
|
||||
header_rules=header_rules,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db=db,
|
||||
user=current_user,
|
||||
provider_name=provider.name,
|
||||
provider_id=provider.id,
|
||||
api_key_id=endpoint_config.get("api_key_id"),
|
||||
model_name=request.model_name,
|
||||
# Provider 上下文
|
||||
auth_type=auth_type,
|
||||
provider_type=p_type if p_type else None,
|
||||
decrypted_auth_config=oauth_meta if oauth_meta else None,
|
||||
)
|
||||
async def _do_check(req: dict) -> dict:
|
||||
return await adapter_class.check_endpoint(
|
||||
client,
|
||||
endpoint_config["base_url"],
|
||||
endpoint_config["api_key"],
|
||||
req,
|
||||
extra_headers if extra_headers else None,
|
||||
body_rules=body_rules,
|
||||
header_rules=header_rules,
|
||||
db=db,
|
||||
user=current_user,
|
||||
provider_name=provider.name,
|
||||
provider_id=provider.id,
|
||||
api_key_id=endpoint_config.get("api_key_id"),
|
||||
model_name=request.model_name,
|
||||
auth_type=auth_type,
|
||||
provider_type=p_type if p_type 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] 端点测试结果:")
|
||||
@@ -745,7 +776,7 @@ async def test_model(
|
||||
return {
|
||||
"success": is_success,
|
||||
"data": {
|
||||
"stream": False,
|
||||
"stream": used_stream,
|
||||
"response": response,
|
||||
},
|
||||
"provider": {
|
||||
|
||||
@@ -19,7 +19,6 @@ Chat Adapter 通用基类
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import traceback
|
||||
from abc import abstractmethod
|
||||
from typing import Any, ClassVar
|
||||
|
||||
@@ -438,18 +437,10 @@ class ChatAdapterBase(ApiAdapter):
|
||||
) -> JSONResponse:
|
||||
"""处理未预期的异常"""
|
||||
if isinstance(e, ProxyException):
|
||||
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}")
|
||||
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}: {e}")
|
||||
else:
|
||||
logger.error(
|
||||
f"{self.FORMAT_ID} 请求处理意外异常",
|
||||
exception=e,
|
||||
extra_data={
|
||||
"exception_class": e.__class__.__name__,
|
||||
"processing_stage": "request_processing",
|
||||
"model": model,
|
||||
"stream": stream,
|
||||
"traceback_preview": str(traceback.format_exc())[:500],
|
||||
},
|
||||
logger.opt(exception=e).error(
|
||||
f"{self.FORMAT_ID} 请求处理意外异常: {type(e).__name__}: {e}"
|
||||
)
|
||||
|
||||
response_time = int((time.time() - start_time) * 1000)
|
||||
|
||||
@@ -18,7 +18,6 @@ CLI Adapter 通用基类
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import traceback
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import httpx
|
||||
@@ -417,18 +416,10 @@ class CliAdapterBase(ApiAdapter):
|
||||
) -> JSONResponse:
|
||||
"""处理未预期的异常"""
|
||||
if isinstance(e, ProxyException):
|
||||
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}")
|
||||
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}: {e}")
|
||||
else:
|
||||
logger.error(
|
||||
f"{self.FORMAT_ID} 请求处理意外异常",
|
||||
exception=e,
|
||||
extra_data={
|
||||
"exception_class": e.__class__.__name__,
|
||||
"processing_stage": "request_processing",
|
||||
"model": model,
|
||||
"stream": stream,
|
||||
"traceback_preview": str(traceback.format_exc())[:500],
|
||||
},
|
||||
logger.opt(exception=e).error(
|
||||
f"{self.FORMAT_ID} 请求处理意外异常: {type(e).__name__}: {e}"
|
||||
)
|
||||
|
||||
response_time = int((time.time() - start_time) * 1000)
|
||||
|
||||
@@ -1353,6 +1353,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except GeneratorExit:
|
||||
raise
|
||||
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:
|
||||
# 流已开始,发送错误事件而不是抛出异常
|
||||
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()
|
||||
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:
|
||||
error_event = {
|
||||
"type": "error",
|
||||
@@ -1390,6 +1398,106 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except Exception:
|
||||
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(
|
||||
self,
|
||||
byte_iterator: Any,
|
||||
@@ -1787,6 +1895,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except GeneratorExit:
|
||||
raise
|
||||
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:
|
||||
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()
|
||||
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:
|
||||
error_event = {
|
||||
"type": "error",
|
||||
@@ -2357,6 +2473,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
|
||||
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 必需头
|
||||
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
|
||||
client_response_headers.update(
|
||||
|
||||
41
src/services/cache/aware_scheduler.py
vendored
41
src/services/cache/aware_scheduler.py
vendored
@@ -83,6 +83,43 @@ class ProviderCandidate:
|
||||
needs_conversion: bool = False # 是否需要格式转换
|
||||
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
|
||||
class ConcurrencySnapshot:
|
||||
@@ -1585,7 +1622,7 @@ class CacheAwareScheduler:
|
||||
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)
|
||||
else:
|
||||
# 单个候选或没有 affinity_key,按次要排序条件排序
|
||||
@@ -1706,7 +1743,7 @@ class CacheAwareScheduler:
|
||||
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)
|
||||
else:
|
||||
# 没有 affinity_key 时按 ID 排序保持稳定性
|
||||
|
||||
Reference in New Issue
Block a user