mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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):
|
||||
"""
|
||||
|
||||
25
src/services/cache/aware_scheduler.py
vendored
25
src/services/cache/aware_scheduler.py
vendored
@@ -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(
|
||||
|
||||
@@ -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"),
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user