mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +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] 端点 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": {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
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 # 是否需要格式转换
|
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 排序保持稳定性
|
||||||
|
|||||||
81
tests/unit/test_provider_candidate_ordering.py
Normal file
81
tests/unit/test_provider_candidate_ordering.py
Normal 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]
|
||||||
Reference in New Issue
Block a user