feat: 跨格式转换时按 ApiFamily 优先级排序端点

- 为 ApiFamily 枚举添加 priority 属性 (OpenAI=1, Claude=2, Gemini=3)
- 新增 _sort_endpoints_by_family_priority 函数按优先级排序端点
- 在调度器中对各分组内的端点应用优先级排序
- 修正测试文件的 type ignore 注解 (attr-defined -> method-assign)
- 新增 7 个端点排序相关的单元测试
This commit is contained in:
fawney19
2026-02-05 02:01:06 +08:00
parent 9a8f25d1a9
commit 34272af711
3 changed files with 162 additions and 11 deletions

View File

@@ -17,6 +17,15 @@ class ApiFamily(str, Enum):
CLAUDE = "claude" # claude-compatible
GEMINI = "gemini" # gemini-compatible
@property
def priority(self) -> int:
"""基础优先级(数字越小越优先)"""
return {
ApiFamily.OPENAI: 1,
ApiFamily.CLAUDE: 2,
ApiFamily.GEMINI: 3,
}.get(self, 99)
class EndpointKind(str, Enum):
"""

View File

@@ -34,12 +34,13 @@ import hashlib
import random
import re
import time
from collections.abc import Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from sqlalchemy.orm import Session, selectinload
from src.core.api_format.enums import EndpointKind
from src.core.api_format.enums import ApiFamily, EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
from src.core.logger import logger
@@ -103,6 +104,21 @@ class ConcurrencySnapshot:
)
def _sort_endpoints_by_family_priority(
eps: Sequence[ProviderEndpoint],
) -> list[ProviderEndpoint]:
"""按 ApiFamily 优先级对端点排序(同分组内使用)。"""
def sort_key(ep: ProviderEndpoint) -> int:
family_str = str(getattr(ep, "api_family", "") or "").strip().lower()
try:
return ApiFamily(family_str).priority
except ValueError:
return 99
return sorted(eps, key=sort_key)
class CacheAwareScheduler:
"""
缓存感知调度器
@@ -1149,7 +1165,12 @@ class CacheAwareScheduler:
else:
fallback_other_family.append(ep)
endpoints = preferred + preferred_other_family + fallback + fallback_other_family
endpoints = (
_sort_endpoints_by_family_priority(preferred)
+ _sort_endpoints_by_family_priority(preferred_other_family)
+ _sort_endpoints_by_family_priority(fallback)
+ _sort_endpoints_by_family_priority(fallback_other_family)
)
for endpoint in endpoints:
logger.debug(

View File

@@ -1,9 +1,13 @@
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
from src.core.api_format.conversion import register_default_normalizers
from src.services.cache.aware_scheduler import CacheAwareScheduler
from src.services.cache.aware_scheduler import (
CacheAwareScheduler,
_sort_endpoints_by_family_priority,
)
def _mock_key(key_id: str, api_formats: list[str]) -> MagicMock:
@@ -34,8 +38,8 @@ async def test_build_candidates_allows_cross_format_when_endpoint_accepts_and_ov
register_default_normalizers()
scheduler = CacheAwareScheduler()
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[attr-defined]
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[method-assign]
provider = MagicMock()
provider.name = "p1"
@@ -67,8 +71,8 @@ async def test_build_candidates_blocks_cross_format_when_master_switch_off() ->
register_default_normalizers()
scheduler = CacheAwareScheduler()
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[attr-defined]
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[method-assign]
provider = MagicMock()
provider.name = "p1"
@@ -98,8 +102,8 @@ async def test_build_candidates_includes_cross_format_when_enabled() -> None:
register_default_normalizers()
scheduler = CacheAwareScheduler()
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[attr-defined]
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[method-assign]
provider = MagicMock()
provider.name = "p1"
@@ -129,8 +133,8 @@ async def test_exact_matches_rank_before_convertible() -> None:
register_default_normalizers()
scheduler = CacheAwareScheduler()
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[attr-defined]
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[method-assign]
scheduler._check_key_availability = MagicMock(return_value=(True, None, None)) # type: ignore[method-assign]
provider = MagicMock()
provider.name = "p1"
@@ -162,3 +166,120 @@ async def test_exact_matches_rank_before_convertible() -> None:
assert candidates[0].provider_api_format == "claude:chat"
assert candidates[1].needs_conversion is True
assert candidates[1].provider_api_format == "openai:chat"
def test_sort_endpoints_by_family_priority_orders_openai_claude_gemini() -> None:
eps = [
_mock_endpoint("gemini:chat"),
_mock_endpoint("claude:chat"),
_mock_endpoint("openai:chat"),
]
result = _sort_endpoints_by_family_priority(eps)
assert [e.api_family for e in result] == ["openai", "claude", "gemini"]
def test_sort_endpoints_by_family_priority_unknown_family_sorted_last() -> None:
eps = [
_mock_endpoint("unknown:chat"),
_mock_endpoint("openai:chat"),
]
result = _sort_endpoints_by_family_priority(eps)
assert [e.api_family for e in result] == ["openai", "unknown"]
def test_sort_endpoints_by_family_priority_stable_for_same_family() -> None:
ep1 = _mock_endpoint("openai:chat")
ep1.base_url = "url_1"
ep2 = _mock_endpoint("openai:chat")
ep2.base_url = "url_2"
result = _sort_endpoints_by_family_priority([ep1, ep2])
assert result[0].base_url == "url_1"
assert result[1].base_url == "url_2"
def test_sort_endpoints_by_family_priority_empty_list() -> None:
assert _sort_endpoints_by_family_priority([]) == []
def _group_and_sort_endpoints(
client_family: str, client_kind: str, endpoints: list[MagicMock]
) -> list[Any]:
preferred, preferred_other, fallback, fallback_other = [], [], [], []
for ep in endpoints:
same_family = ep.api_family == client_family
same_kind = ep.endpoint_kind == client_kind
if same_family and same_kind:
preferred.append(ep)
elif same_kind:
preferred_other.append(ep)
elif same_family:
fallback.append(ep)
else:
fallback_other.append(ep)
return (
_sort_endpoints_by_family_priority(preferred)
+ _sort_endpoints_by_family_priority(preferred_other)
+ _sort_endpoints_by_family_priority(fallback)
+ _sort_endpoints_by_family_priority(fallback_other)
)
def test_group_and_sort_endpoints_client_openai_chat() -> None:
endpoints = [
_mock_endpoint("gemini:cli"),
_mock_endpoint("claude:chat"),
_mock_endpoint("openai:cli"),
_mock_endpoint("gemini:chat"),
_mock_endpoint("claude:cli"),
_mock_endpoint("openai:chat"),
]
result = _group_and_sort_endpoints("openai", "chat", endpoints)
assert [(e.api_family, e.endpoint_kind) for e in result] == [
("openai", "chat"),
("claude", "chat"),
("gemini", "chat"),
("openai", "cli"),
("claude", "cli"),
("gemini", "cli"),
]
def test_group_and_sort_endpoints_client_claude_chat() -> None:
endpoints = [
_mock_endpoint("gemini:cli"),
_mock_endpoint("claude:chat"),
_mock_endpoint("openai:cli"),
_mock_endpoint("gemini:chat"),
_mock_endpoint("claude:cli"),
_mock_endpoint("openai:chat"),
]
result = _group_and_sort_endpoints("claude", "chat", endpoints)
assert [(e.api_family, e.endpoint_kind) for e in result] == [
("claude", "chat"),
("openai", "chat"),
("gemini", "chat"),
("claude", "cli"),
("openai", "cli"),
("gemini", "cli"),
]
def test_group_and_sort_endpoints_client_openai_cli() -> None:
endpoints = [
_mock_endpoint("gemini:cli"),
_mock_endpoint("claude:chat"),
_mock_endpoint("openai:cli"),
_mock_endpoint("gemini:chat"),
_mock_endpoint("claude:cli"),
_mock_endpoint("openai:chat"),
]
result = _group_and_sort_endpoints("openai", "cli", endpoints)
assert [(e.api_family, e.endpoint_kind) for e in result] == [
("openai", "cli"),
("claude", "cli"),
("gemini", "cli"),
("openai", "chat"),
("claude", "chat"),
("gemini", "chat"),
]