mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 修复候选分页逻辑并拆分配额检查模块
- list_all_candidates 返回 provider_batch_count,区分"无候选"与"无 Provider" 避免分页在有 Provider 但无候选时提前终止 - 将配额检查逻辑从 aware_scheduler.py 拆分到 quota_skipper.py - 提取 reorder_candidates 方法,支持跨页汇总后全局重排序 - CandidateResolver 分页循环改用 provider_batch_count 判断终止 - candidate extra_data 中新增 mapping_matched_model 字段 - 新增契约测试和分页行为测试
This commit is contained in:
40
tests/contracts/test_public_api_routes_contract.py
Normal file
40
tests/contracts/test_public_api_routes_contract.py
Normal file
@@ -0,0 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
|
||||
def _build_contract_app() -> FastAPI:
|
||||
from src.api.public.claude import router as claude_router
|
||||
from src.api.public.gemini import router as gemini_router
|
||||
from src.api.public.openai import router as openai_router
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(claude_router)
|
||||
app.include_router(openai_router)
|
||||
app.include_router(gemini_router)
|
||||
return app
|
||||
|
||||
|
||||
def test_public_api_routes_contract_paths_and_tags() -> None:
|
||||
app = _build_contract_app()
|
||||
schema = app.openapi()
|
||||
|
||||
paths = schema.get("paths") or {}
|
||||
|
||||
expected = [
|
||||
("/v1/messages", "post", "Claude API"),
|
||||
("/v1/messages/count_tokens", "post", "Claude API"),
|
||||
("/v1/chat/completions", "post", "OpenAI API"),
|
||||
("/v1/responses", "post", "OpenAI API"),
|
||||
("/v1beta/models/{model}:generateContent", "post", "Gemini API"),
|
||||
("/v1beta/models/{model}:streamGenerateContent", "post", "Gemini API"),
|
||||
("/v1/models/{model}:generateContent", "post", "Gemini API"),
|
||||
("/v1/models/{model}:streamGenerateContent", "post", "Gemini API"),
|
||||
]
|
||||
|
||||
for path, method, expected_tag in expected:
|
||||
assert path in paths, f"missing path {path}"
|
||||
operations = paths.get(path) or {}
|
||||
assert method in operations, f"missing {method.upper()} {path}"
|
||||
tags = operations.get(method, {}).get("tags") or []
|
||||
assert expected_tag in tags, f"{method.upper()} {path} missing tag {expected_tag}"
|
||||
132
tests/contracts/test_scheduler_list_all_candidates_contract.py
Normal file
132
tests/contracts/test_scheduler_list_all_candidates_contract.py
Normal file
@@ -0,0 +1,132 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
def _make_db() -> MagicMock:
|
||||
db = MagicMock()
|
||||
db.new = []
|
||||
db.dirty = []
|
||||
db.deleted = []
|
||||
db.in_transaction.return_value = False
|
||||
return db
|
||||
|
||||
|
||||
def _make_global_model(*, gid: str, name: str) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id=gid,
|
||||
name=name,
|
||||
is_active=True,
|
||||
config={},
|
||||
supported_capabilities=[],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_all_candidates_returns_provider_batch_count_even_when_candidates_empty() -> (
|
||||
None
|
||||
):
|
||||
"""契约:候选为空不代表 Provider 页为空(用于分页继续拉取下一页)。"""
|
||||
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
|
||||
|
||||
db = _make_db()
|
||||
|
||||
providers = [
|
||||
SimpleNamespace(
|
||||
id="p1",
|
||||
name="p1",
|
||||
is_active=True,
|
||||
endpoints=[],
|
||||
models=[],
|
||||
provider_priority=1,
|
||||
),
|
||||
SimpleNamespace(
|
||||
id="p2",
|
||||
name="p2",
|
||||
is_active=True,
|
||||
endpoints=[],
|
||||
models=[],
|
||||
provider_priority=2,
|
||||
),
|
||||
]
|
||||
|
||||
# allowed_providers 会把本页的 provider 全过滤掉,导致 candidates 为空;但 provider_batch_count 应保留过滤前数量。
|
||||
user_api_key = SimpleNamespace(
|
||||
id="ak1",
|
||||
allowed_providers=["not-matching"],
|
||||
allowed_models=None,
|
||||
allowed_api_formats=None,
|
||||
user=None,
|
||||
)
|
||||
|
||||
global_model = _make_global_model(gid="gm1", name="gpt-4o")
|
||||
|
||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||
with patch.object(scheduler, "_query_providers", return_value=providers):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
new=AsyncMock(return_value=global_model),
|
||||
):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
return_value=True,
|
||||
):
|
||||
candidates, global_model_id, provider_batch_count = (
|
||||
await scheduler.list_all_candidates(
|
||||
db=db,
|
||||
api_format="openai:chat",
|
||||
model_name="gpt-4o",
|
||||
affinity_key=None,
|
||||
user_api_key=user_api_key, # type: ignore[arg-type]
|
||||
provider_offset=0,
|
||||
provider_limit=20,
|
||||
)
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert global_model_id == "gm1"
|
||||
assert provider_batch_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_all_candidates_returns_zero_provider_batch_count_when_provider_page_empty() -> (
|
||||
None
|
||||
):
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
|
||||
|
||||
db = _make_db()
|
||||
global_model = _make_global_model(gid="gm1", name="gpt-4o")
|
||||
|
||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||
with patch.object(scheduler, "_query_providers", return_value=[]):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
new=AsyncMock(return_value=global_model),
|
||||
):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
return_value=True,
|
||||
):
|
||||
candidates, global_model_id, provider_batch_count = (
|
||||
await scheduler.list_all_candidates(
|
||||
db=db,
|
||||
api_format="openai:chat",
|
||||
model_name="gpt-4o",
|
||||
affinity_key=None,
|
||||
user_api_key=None,
|
||||
provider_offset=0,
|
||||
provider_limit=20,
|
||||
)
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert global_model_id == "gm1"
|
||||
assert provider_batch_count == 0
|
||||
97
tests/services/test_candidate_resolver_pagination.py
Normal file
97
tests/services/test_candidate_resolver_pagination.py
Normal file
@@ -0,0 +1,97 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||
|
||||
|
||||
class _FakeScheduler:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[int] = []
|
||||
self.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
|
||||
self.priority_mode = CacheAwareScheduler.PRIORITY_MODE_PROVIDER
|
||||
|
||||
async def list_all_candidates(
|
||||
self,
|
||||
*,
|
||||
db: Any,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str | None = None,
|
||||
user_api_key: Any | None = None,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
max_candidates: int | None = None,
|
||||
is_stream: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
) -> tuple[list[Any], str, int]:
|
||||
_ = (
|
||||
db,
|
||||
api_format,
|
||||
model_name,
|
||||
affinity_key,
|
||||
user_api_key,
|
||||
max_candidates,
|
||||
is_stream,
|
||||
capability_requirements,
|
||||
)
|
||||
assert provider_limit is not None
|
||||
|
||||
self.calls.append(int(provider_offset))
|
||||
|
||||
if provider_offset == 0:
|
||||
# Simulate a provider page that has providers but no eligible candidates.
|
||||
return [], "gm1", int(provider_limit)
|
||||
|
||||
if provider_offset == int(provider_limit):
|
||||
# Next page yields one eligible candidate and is also the last provider page.
|
||||
cand = SimpleNamespace(
|
||||
provider=SimpleNamespace(id="p1", name="prov"),
|
||||
endpoint=SimpleNamespace(id="e1"),
|
||||
key=SimpleNamespace(id="k1"),
|
||||
is_skipped=False,
|
||||
skip_reason=None,
|
||||
is_cached=False,
|
||||
needs_conversion=False,
|
||||
provider_api_format=str(api_format),
|
||||
mapping_matched_model=None,
|
||||
)
|
||||
return [cand], "gm1", 5
|
||||
|
||||
return [], "gm1", 0
|
||||
|
||||
async def reorder_candidates(
|
||||
self,
|
||||
candidates: list[Any],
|
||||
db: Any = None,
|
||||
affinity_key: str | None = None,
|
||||
api_format: str | None = None,
|
||||
global_model_id: str | None = None,
|
||||
) -> list[Any]:
|
||||
return candidates
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_candidate_resolver_pagination_continues_on_empty_candidate_batch() -> None:
|
||||
db = MagicMock()
|
||||
scheduler = _FakeScheduler()
|
||||
resolver = CandidateResolver(db=db, cache_scheduler=scheduler) # type: ignore[arg-type]
|
||||
|
||||
candidates, global_model_id = await resolver.fetch_candidates(
|
||||
api_format="openai:chat",
|
||||
model_name="gpt-4o",
|
||||
affinity_key="a1",
|
||||
user_api_key=None,
|
||||
request_id="r1",
|
||||
is_stream=False,
|
||||
capability_requirements=None,
|
||||
)
|
||||
|
||||
assert global_model_id == "gm1"
|
||||
assert len(candidates) == 1
|
||||
assert scheduler.calls == [0, 20]
|
||||
Reference in New Issue
Block a user