refactor: 提取模型可用性查询到独立模块 ModelAvailabilityQuery

Close #101

- 新增 src/services/model/availability.py,集中管理系统级模型可用性查询条件
- 重构 models_service.py,复用 ModelAvailabilityQuery 的查询方法
- 删除 find_model_by_id 的 provider_model_name 回退逻辑
- 支持 model_mappings 正则匹配(使用 check_model_allowed_with_mappings)
- 修复 test_auth.py 中缺少的 is_locked mock 属性
- 新增 ModelAvailabilityQuery 单元测试
This commit is contained in:
fawney19
2026-01-18 01:45:25 +08:00
parent b98fc8497f
commit 0c5c6b3288
5 changed files with 1052 additions and 185 deletions

View File

@@ -13,18 +13,15 @@
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
from typing import Any, Optional from typing import Any, Optional
from sqlalchemy.orm import Session, joinedload from sqlalchemy.orm import Session
from src.config.constants import CacheTTL from src.config.constants import CacheTTL
from src.core.cache_service import CacheService from src.core.cache_service import CacheService
from src.core.logger import logger from src.core.logger import logger
from src.services.model.availability import ModelAvailabilityQuery
from src.models.database import ( from src.models.database import (
ApiKey, ApiKey,
GlobalModel,
Model, Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
User, User,
) )
@@ -201,58 +198,17 @@ def get_available_provider_ids(db: Session, api_formats: list[str]) -> set[str]:
- 端点是活跃的 - 端点是活跃的
- Provider 下有活跃的 Key 且支持该 api_formatKey 直属 Provider通过 api_formats 过滤) - Provider 下有活跃的 Key 且支持该 api_formatKey 直属 Provider通过 api_formats 过滤)
""" """
target_formats = {f.upper() for f in api_formats} provider_to_formats = ModelAvailabilityQuery.get_providers_with_active_endpoints(db, api_formats)
if not provider_to_formats:
# 1) 先找出有活跃端点的 Provider记录每个 Provider 支持的格式集合)
endpoint_rows = (
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
.filter(
ProviderEndpoint.api_format.in_(list(target_formats)),
ProviderEndpoint.is_active.is_(True),
)
.all()
)
if not endpoint_rows:
return set() return set()
provider_to_formats: dict[str, set[str]] = {} return ModelAvailabilityQuery.get_providers_with_active_keys(
for provider_id, fmt in endpoint_rows: db,
if not provider_id or not fmt: set(provider_to_formats.keys()),
continue api_formats,
provider_to_formats.setdefault(provider_id, set()).add(str(fmt).upper()) provider_to_formats,
provider_ids_with_endpoints = set(provider_to_formats.keys())
if not provider_ids_with_endpoints:
return set()
# 2) 再检查这些 Provider 是否至少有一个活跃 Key 支持对应格式
key_rows = (
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
.filter(
ProviderAPIKey.provider_id.in_(provider_ids_with_endpoints),
ProviderAPIKey.is_active.is_(True),
)
.all()
) )
available_provider_ids: set[str] = set()
for provider_id, key_formats in key_rows:
if not provider_id:
continue
endpoint_formats = provider_to_formats.get(provider_id)
if not endpoint_formats:
continue
formats_list = key_formats if isinstance(key_formats, list) else []
key_formats_upper = {str(f).upper() for f in formats_list}
# 只有同时满足:请求格式 ∩ Provider 端点格式 ∩ Key 支持格式 非空,才算可用
if key_formats_upper & endpoint_formats & target_formats:
available_provider_ids.add(provider_id)
return available_provider_ids
def _get_available_model_ids_for_format(db: Session, api_formats: list[str]) -> set[str]: def _get_available_model_ids_for_format(db: Session, api_formats: list[str]) -> set[str]:
""" """
@@ -264,74 +220,24 @@ def _get_available_model_ids_for_format(db: Session, api_formats: list[str]) ->
3. **该端点的 Provider 关联了该模型** 3. **该端点的 Provider 关联了该模型**
4. Key 的 allowed_models 允许该模型null = 允许该 Provider 关联的所有模型) 4. Key 的 allowed_models 允许该模型null = 允许该 Provider 关联的所有模型)
""" """
target_formats = {f.upper() for f in api_formats} provider_to_formats = ModelAvailabilityQuery.get_providers_with_active_endpoints(db, api_formats)
if not provider_to_formats:
# 1) 找出有活跃端点的 Provider记录每个 Provider 支持的格式集合)
endpoint_rows = (
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
.filter(
ProviderEndpoint.api_format.in_(list(target_formats)),
ProviderEndpoint.is_active.is_(True),
)
.all()
)
if not endpoint_rows:
return set() return set()
provider_to_formats: dict[str, set[str]] = {} provider_key_rules = ModelAvailabilityQuery.get_provider_key_rules(
for provider_id, fmt in endpoint_rows: db,
if not provider_id or not fmt: provider_ids=set(provider_to_formats.keys()),
continue api_formats=api_formats,
provider_to_formats.setdefault(provider_id, set()).add(str(fmt).upper()) provider_to_endpoint_formats=provider_to_formats,
provider_ids_with_endpoints = set(provider_to_formats.keys())
if not provider_ids_with_endpoints:
return set()
# 2) 收集每个 Provider 下「支持对应格式」的活跃 Key 的 allowed_models
# Key 直属 Provider通过 key.api_formats 与 Provider 端点格式交集筛选
key_rows = (
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.allowed_models, ProviderAPIKey.api_formats)
.filter(
ProviderAPIKey.provider_id.in_(provider_ids_with_endpoints),
ProviderAPIKey.is_active.is_(True),
)
.all()
) )
# provider_id -> list[(allowed_models, usable_formats)]
provider_key_rules: dict[str, list[tuple[object, set[str]]]] = {}
for provider_id, allowed_models, key_formats in key_rows:
if not provider_id:
continue
endpoint_formats = provider_to_formats.get(provider_id)
if not endpoint_formats:
continue
formats_list = key_formats if isinstance(key_formats, list) else []
key_formats_upper = {str(f).upper() for f in formats_list}
usable_formats = key_formats_upper & endpoint_formats & target_formats
if not usable_formats:
continue
provider_key_rules.setdefault(provider_id, []).append((allowed_models, usable_formats))
provider_ids_with_format = set(provider_key_rules.keys()) provider_ids_with_format = set(provider_key_rules.keys())
if not provider_ids_with_format: if not provider_ids_with_format:
return set() return set()
# 只查询那些有匹配格式端点的 Provider 下的模型
models = ( models = (
db.query(Model) ModelAvailabilityQuery.base_active_models(db, eager_load=True)
.options(joinedload(Model.global_model)) .filter(Model.provider_id.in_(provider_ids_with_format))
.join(Provider)
.filter(
Model.provider_id.in_(provider_ids_with_format),
Model.is_active.is_(True),
Provider.is_active.is_(True),
)
.all() .all()
) )
@@ -340,9 +246,7 @@ def _get_available_model_ids_for_format(db: Session, api_formats: list[str]) ->
for model in models: for model in models:
model_provider_id = model.provider_id model_provider_id = model.provider_id
global_model = model.global_model global_model = model.global_model
model_id = global_model.name if global_model else model.provider_model_name # type: ignore if not model_provider_id or not global_model or not global_model.name:
if not model_provider_id or not model_id:
continue continue
# 该模型的 Provider 必须有匹配格式的端点 # 该模型的 Provider 必须有匹配格式的端点
@@ -350,32 +254,46 @@ def _get_available_model_ids_for_format(db: Session, api_formats: list[str]) ->
continue continue
# 检查该 provider 下是否有 Key 允许这个模型 # 检查该 provider 下是否有 Key 允许这个模型
from src.core.model_permissions import check_model_allowed from src.core.model_permissions import check_model_allowed_with_mappings
model_id = global_model.name
model_mappings = (global_model.config or {}).get("model_mappings")
rules = provider_key_rules.get(model_provider_id, []) rules = provider_key_rules.get(model_provider_id, [])
for allowed_models, usable_formats in rules: for allowed_models, _usable_formats in rules:
# None = 不限制 # None = 不限制
if allowed_models is None: if allowed_models is None:
available_model_ids.add(model_id) available_model_ids.add(model_id)
break break
# 检查是否允许该模型 # 检查是否允许该模型(支持 model_mappings 正则匹配)
if check_model_allowed( is_allowed, _ = check_model_allowed_with_mappings(
model_name=model_id, model_name=model_id,
allowed_models=allowed_models, # type: ignore[arg-type] allowed_models=allowed_models,
resolved_model_name=(model.provider_model_name if global_model else None), resolved_model_name=model.provider_model_name,
): model_mappings=model_mappings,
)
if is_allowed:
available_model_ids.add(model_id) available_model_ids.add(model_id)
break break
return available_model_ids return available_model_ids
def _extract_model_info(model: Any) -> ModelInfo: def _extract_model_info(model: Any) -> Optional[ModelInfo]:
"""从 Model 对象提取 ModelInfo""" """
从 Model 对象提取 ModelInfo
前置条件model 必须关联 GlobalModel由 base_active_models 内连接保证)
如果 global_model 为 None不应发生返回 None 并记录日志。
"""
global_model = model.global_model global_model = model.global_model
model_id: str = global_model.name if global_model else model.provider_model_name if global_model is None:
display_name: str = global_model.display_name if global_model else model.provider_model_name logger.warning(f"[ModelService] Model {getattr(model, 'id', 'unknown')} 缺少 global_model跳过")
return None
model_id: str = global_model.name
display_name: str = global_model.display_name
created_at: Optional[str] = ( created_at: Optional[str] = (
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
) )
@@ -384,11 +302,8 @@ def _extract_model_info(model: Any) -> ModelInfo:
provider_id: str = model.provider_id or "" provider_id: str = model.provider_id or ""
# 从 GlobalModel.config 提取配置信息 # 从 GlobalModel.config 提取配置信息
config: dict = {} config: dict = global_model.config or {}
description: Optional[str] = None description: Optional[str] = config.get("description")
if global_model:
config = global_model.config or {}
description = config.get("description")
return ModelInfo( return ModelInfo(
id=model_id, id=model_id,
@@ -458,24 +373,20 @@ async def list_available_models(
if not available_model_ids: if not available_model_ids:
return [] return []
query = ( all_models = (
db.query(Model) ModelAvailabilityQuery.base_active_models(db, eager_load=True)
.options(joinedload(Model.global_model), joinedload(Model.provider)) .filter(Model.provider_id.in_(available_provider_ids))
.join(Provider)
.filter(
Model.is_active.is_(True),
Provider.is_active.is_(True),
Model.provider_id.in_(available_provider_ids),
)
.order_by(Model.created_at.desc()) .order_by(Model.created_at.desc())
.all()
) )
all_models = query.all()
result: list[ModelInfo] = [] result: list[ModelInfo] = []
seen_model_ids: set[str] = set() seen_model_ids: set[str] = set()
for model in all_models: for model in all_models:
info = _extract_model_info(model) info = _extract_model_info(model)
if info is None:
continue
# 如果有 available_model_ids 限制,检查是否在其中 # 如果有 available_model_ids 限制,检查是否在其中
if available_model_ids is not None and info.id not in available_model_ids: if available_model_ids is not None and info.id not in available_model_ids:
@@ -507,12 +418,7 @@ def find_model_by_id(
restrictions: Optional[AccessRestrictions] = None, restrictions: Optional[AccessRestrictions] = None,
) -> Optional[ModelInfo]: ) -> Optional[ModelInfo]:
""" """
按 ID 查找模型 按 ID 查找模型(仅支持 GlobalModel.name
查找顺序:
1. 先按 GlobalModel.name 查找
2. 如果没找到任何候选,再按 provider_model_name 查找
3. 如果有候选但都不可用,返回 None不回退
Args: Args:
db: 数据库会话 db: 数据库会话
@@ -540,17 +446,8 @@ def find_model_by_id(
if model_id not in restrictions.allowed_models: if model_id not in restrictions.allowed_models:
return None return None
# 先按 GlobalModel.name 查找
models_by_global = ( models_by_global = (
db.query(Model) ModelAvailabilityQuery.find_by_global_model_name(db, model_id, eager_load=True)
.options(joinedload(Model.global_model), joinedload(Model.provider))
.join(Provider)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.filter(
GlobalModel.name == model_id,
Model.is_active.is_(True),
Provider.is_active.is_(True),
)
.order_by(Model.created_at.desc()) .order_by(Model.created_at.desc())
.all() .all()
) )
@@ -566,34 +463,7 @@ def find_model_by_id(
return False return False
return True return True
model = next( model = next((m for m in models_by_global if is_model_accessible(m)), None)
(m for m in models_by_global if is_model_accessible(m)),
None,
)
# 如果有候选但都不可用,直接返回 None不回退 provider_model_name
if not model and models_by_global:
return None
# 如果找不到任何候选,按 provider_model_name 查找
if not model:
models_by_provider_name = (
db.query(Model)
.options(joinedload(Model.global_model), joinedload(Model.provider))
.join(Provider)
.filter(
Model.provider_model_name == model_id,
Model.is_active.is_(True),
Provider.is_active.is_(True),
)
.order_by(Model.created_at.desc())
.all()
)
model = next(
(m for m in models_by_provider_name if is_model_accessible(m)),
None,
)
if not model: if not model:
return None return None

View File

@@ -7,12 +7,14 @@
from src.services.model.cost import ModelCostService from src.services.model.cost import ModelCostService
from src.services.model.fetch_scheduler import ModelFetchScheduler, get_model_fetch_scheduler from src.services.model.fetch_scheduler import ModelFetchScheduler, get_model_fetch_scheduler
from src.services.model.global_model import GlobalModelService from src.services.model.global_model import GlobalModelService
from src.services.model.availability import ModelAvailabilityQuery
from src.services.model.service import ModelService from src.services.model.service import ModelService
__all__ = [ __all__ = [
"ModelService", "ModelService",
"GlobalModelService", "GlobalModelService",
"ModelCostService", "ModelCostService",
"ModelAvailabilityQuery",
"ModelFetchScheduler", "ModelFetchScheduler",
"get_model_fetch_scheduler", "get_model_fetch_scheduler",
] ]

View File

@@ -0,0 +1,264 @@
"""
模型可用性查询模块
将所有系统级「可用性」条件集中管理,作为模型查询的单一来源。
职责边界:
- 本模块只负责系统级可用性(对所有请求一致)
- API Key/User 的请求级访问限制由 models_service.AccessRestrictions 处理
"""
from __future__ import annotations
from typing import Optional
from sqlalchemy import or_
from sqlalchemy.orm import Query, Session, contains_eager
from src.core.logger import logger
from src.models.database import (
GlobalModel,
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
class ModelAvailabilityQuery:
"""
模型可用性查询构建器
设计原则:
1. 单一来源:所有可用性条件定义在此类中
2. 内连接 GlobalModel未关联的 Model 不参与路由global_model_id=NULL 不返回)
3. 完整过滤:包含 is_active 与 is_availableis_available=NULL 视为可用,兼容历史数据)
"""
@staticmethod
def base_active_models(db: Session, eager_load: bool = False) -> Query:
"""
返回基础的可用模型查询(实体为 Model
已包含条件:
- Model.is_active = True
- Model.is_available = True 或 NULLNULL 视为可用)
- Provider.is_active = True
- GlobalModel.is_active = True
- Model 必须关联 GlobalModel内连接排除 global_model_id=NULL
Args:
db: 数据库会话
eager_load: 是否预加载 Provider 与 GlobalModel复用 join避免重复 JOIN
"""
# 使用关系路径 join与 contains_eager 兼容
query = (
db.query(Model)
.join(Model.provider)
.join(Model.global_model)
.filter(
Model.is_active.is_(True),
or_(Model.is_available.is_(True), Model.is_available.is_(None)),
Provider.is_active.is_(True),
GlobalModel.is_active.is_(True),
)
)
if eager_load:
query = query.options(
contains_eager(Model.provider),
contains_eager(Model.global_model),
)
return query
@staticmethod
def get_providers_with_active_endpoints(
db: Session,
api_formats: list[str],
) -> dict[str, set[str]]:
"""
获取有活跃端点的 Provider 及其支持的格式集合
条件:
- Provider.is_active = True提前过滤减少无效候选
- ProviderEndpoint.is_active = True
- ProviderEndpoint.api_format 匹配
Returns:
{provider_id: {format1, format2, ...}}
"""
target_formats = {f.upper() for f in api_formats}
endpoint_rows = (
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
Provider.is_active.is_(True),
ProviderEndpoint.api_format.in_(list(target_formats)),
ProviderEndpoint.is_active.is_(True),
)
.all()
)
provider_to_formats: dict[str, set[str]] = {}
for provider_id, fmt in endpoint_rows:
if provider_id and fmt:
provider_to_formats.setdefault(provider_id, set()).add(str(fmt).upper())
return provider_to_formats
@staticmethod
def get_providers_with_active_keys(
db: Session,
provider_ids: set[str],
api_formats: list[str],
provider_to_endpoint_formats: dict[str, set[str]],
) -> set[str]:
"""
过滤出有活跃 Key 支持指定格式的 Provider
条件:
- ProviderAPIKey.is_active = True
- Key.api_formats 与 Endpoint 格式与请求格式有交集
"""
if not provider_ids:
return set()
target_formats = {f.upper() for f in api_formats}
key_rows = (
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
.filter(
ProviderAPIKey.provider_id.in_(provider_ids),
ProviderAPIKey.is_active.is_(True),
)
.all()
)
available_provider_ids: set[str] = set()
for provider_id, key_formats in key_rows:
if not provider_id:
continue
endpoint_formats = provider_to_endpoint_formats.get(provider_id)
if not endpoint_formats:
continue
# 类型兜底key_formats 是 JSON 字段
if not isinstance(key_formats, list):
if key_formats is not None:
logger.warning(
f"[ModelAvailability] Key api_formats 类型异常, "
f"provider_id={provider_id}, type={type(key_formats).__name__}"
)
continue
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
if key_formats_upper & endpoint_formats & target_formats:
available_provider_ids.add(provider_id)
return available_provider_ids
@staticmethod
def get_provider_key_rules(
db: Session,
provider_ids: set[str],
api_formats: list[str],
provider_to_endpoint_formats: dict[str, set[str]],
) -> dict[str, list[tuple[Optional[list[str]], set[str]]]]:
"""
获取每个 Provider 的 Key 权限规则
Returns:
{provider_id: [(allowed_models, usable_formats), ...]}
注意:
- allowed_models 是 JSON 字段,此方法会进行类型兜底处理
- 非预期类型会跳过该 Key 并打日志(安全优先,不放大权限)
"""
if not provider_ids:
return {}
target_formats = {f.upper() for f in api_formats}
key_rows = (
db.query(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.allowed_models,
ProviderAPIKey.api_formats,
)
.filter(
ProviderAPIKey.provider_id.in_(provider_ids),
ProviderAPIKey.is_active.is_(True),
)
.all()
)
provider_key_rules: dict[str, list[tuple[Optional[list[str]], set[str]]]] = {}
for key_id, provider_id, allowed_models_raw, key_formats in key_rows:
if not provider_id:
continue
endpoint_formats = provider_to_endpoint_formats.get(provider_id)
if not endpoint_formats:
continue
# 类型兜底key_formats
if not isinstance(key_formats, list):
if key_formats is not None:
logger.warning(
f"[ModelAvailability] Key api_formats 类型异常, "
f"key_id={key_id}, type={type(key_formats).__name__}"
)
continue
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
usable_formats = key_formats_upper & endpoint_formats & target_formats
if not usable_formats:
continue
# 类型兜底allowed_models安全优先
allowed_models: Optional[list[str]]
if allowed_models_raw is None:
# None = 不限制
allowed_models = None
elif isinstance(allowed_models_raw, list):
allowed_models = [m for m in allowed_models_raw if isinstance(m, str)]
else:
logger.warning(
f"[ModelAvailability] Key allowed_models 类型异常, "
f"key_id={key_id}, type={type(allowed_models_raw).__name__}, 跳过该 Key"
)
continue
provider_key_rules.setdefault(provider_id, []).append((allowed_models, usable_formats))
return provider_key_rules
@staticmethod
def find_by_global_model_name(
db: Session,
model_name: str,
provider_ids: Optional[set[str]] = None,
eager_load: bool = False,
) -> Query:
"""
按 GlobalModel.name 查找模型
条件:
- 基础可用性条件base_active_models
- GlobalModel.name = model_name
- 可选:限制到指定 Provider
"""
query = ModelAvailabilityQuery.base_active_models(db, eager_load=eager_load).filter(
GlobalModel.name == model_name
)
if provider_ids is not None:
query = query.filter(Model.provider_id.in_(provider_ids))
return query

View File

@@ -235,6 +235,7 @@ class TestAPIKeyAuthentication:
mock_api_key = MagicMock() mock_api_key = MagicMock()
mock_api_key.is_active = True mock_api_key.is_active = True
mock_api_key.is_locked = False
mock_api_key.expires_at = None mock_api_key.expires_at = None
mock_api_key.user = mock_user mock_api_key.user = mock_user
mock_api_key.balance_used_usd = 0.0 mock_api_key.balance_used_usd = 0.0
@@ -271,6 +272,7 @@ class TestAPIKeyAuthentication:
"""测试 API Key 已禁用""" """测试 API Key 已禁用"""
mock_api_key = MagicMock() mock_api_key = MagicMock()
mock_api_key.is_active = False mock_api_key.is_active = False
mock_api_key.is_locked = False
mock_db = MagicMock() mock_db = MagicMock()
mock_db.query.return_value.options.return_value.filter.return_value.first.return_value = ( mock_db.query.return_value.options.return_value.filter.return_value.first.return_value = (
@@ -286,6 +288,7 @@ class TestAPIKeyAuthentication:
"""测试 API Key 已过期""" """测试 API Key 已过期"""
mock_api_key = MagicMock() mock_api_key = MagicMock()
mock_api_key.is_active = True mock_api_key.is_active = True
mock_api_key.is_locked = False
mock_api_key.expires_at = datetime.now(timezone.utc) - timedelta(days=1) mock_api_key.expires_at = datetime.now(timezone.utc) - timedelta(days=1)
mock_db = MagicMock() mock_db = MagicMock()

View File

@@ -0,0 +1,728 @@
# mypy: disable-error-code="arg-type"
"""
ModelAvailabilityQuery 单元测试
说明:
- 本测试使用静态源码检查 + FakeSession/FakeQuery 风格,避免引入真实 DB 依赖
- 项目日志使用 logurupytest 的 caplog 默认无法直接捕获,因此用 loguru sink 验证 warning
"""
import inspect
from io import StringIO
from typing import Any
from loguru import logger
from src.services.model.availability import ModelAvailabilityQuery
# ============================================================================
# 测试辅助类
# ============================================================================
class FakeQuery:
"""通用 Fake Query支持链式调用和自定义返回数据"""
def __init__(self, data: list[Any]) -> None:
self._data = data
def filter(self, *_args: Any, **_kwargs: Any) -> "FakeQuery":
return self
def join(self, *_args: Any, **_kwargs: Any) -> "FakeQuery":
return self
def all(self) -> list[Any]:
return self._data
class FakeSession:
"""
通用 Fake Session
FakeSession 仅实现了 Session 的 query 方法,用于测试场景。
类型检查会报错,但运行时正常工作。
"""
def __init__(self, query_data: list[Any]) -> None:
self._query_data = query_data
def query(self, *_entities: Any) -> FakeQuery:
return FakeQuery(self._query_data)
class TestBaseActiveModelsSourceCode:
"""
静态验证 base_active_models 的核心实现
通过检查源码确保关键条件不会被遗漏。
"""
def test_uses_inner_join_for_global_model(self) -> None:
"""应使用内连接 GlobalModel排除 global_model_id=NULL"""
source = inspect.getsource(ModelAvailabilityQuery.base_active_models)
assert "join(Model.global_model)" in source or ".join(Model.global_model)" in source, (
"base_active_models 应使用 join(Model.global_model) 内连接"
)
assert "outerjoin" not in source.lower(), (
"base_active_models 不应使用 outerjoin会返回 global_model_id=NULL 的记录)"
)
def test_filters_is_available_true_or_null(self) -> None:
"""应过滤 is_available = True 或 NULL"""
source = inspect.getsource(ModelAvailabilityQuery.base_active_models)
assert "or_(" in source, "base_active_models 应使用 or_() 处理 is_available"
assert "is_available" in source, "base_active_models 应包含 is_available 条件"
assert "is_(True)" in source and "is_(None)" in source, (
"base_active_models 应同时检查 is_available = True 和 NULL"
)
def test_filters_all_is_active_fields(self) -> None:
"""应过滤 Model/Provider/GlobalModel 的 is_active"""
source = inspect.getsource(ModelAvailabilityQuery.base_active_models)
assert "Model.is_active" in source, "base_active_models 应检查 Model.is_active"
assert "Provider.is_active" in source, "base_active_models 应检查 Provider.is_active"
assert "GlobalModel.is_active" in source, "base_active_models 应检查 GlobalModel.is_active"
class TestGetProviderKeyRules:
"""测试 get_provider_key_rules"""
def test_empty_provider_ids_returns_empty_dict(self) -> None:
"""空 provider_ids 应返回空字典"""
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession([]),
provider_ids=set(),
api_formats=["OPENAI"],
provider_to_endpoint_formats={},
)
assert result == {}
def test_skips_key_with_invalid_allowed_models_type(self) -> None:
"""allowed_models 类型异常时应跳过该 Key 并打日志"""
log_output = StringIO()
handler_id = logger.add(
log_output,
format="{message}",
level="WARNING",
filter=lambda record: "[ModelAvailability]" in record["message"],
)
try:
# (key_id, provider_id, allowed_models, api_formats)
data = [("key-1", "provider-1", "invalid-string-type", ["OPENAI"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
log_content = log_output.getvalue()
assert "allowed_models 类型异常" in log_content
assert "key-1" in log_content
finally:
logger.remove(handler_id)
def test_skips_key_with_invalid_api_formats_type(self) -> None:
"""api_formats 类型异常时应跳过该 Key 并打日志"""
log_output = StringIO()
handler_id = logger.add(
log_output,
format="{message}",
level="WARNING",
filter=lambda record: "[ModelAvailability]" in record["message"],
)
try:
# api_formats 是字符串而非列表
data = [("key-1", "provider-1", None, "OPENAI")]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
log_content = log_output.getvalue()
assert "api_formats 类型异常" in log_content
finally:
logger.remove(handler_id)
def test_skips_key_with_none_api_formats(self) -> None:
"""api_formats 为 None 时应跳过该 Key不打日志"""
data = [("key-1", "provider-1", None, None)]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
def test_allowed_models_none_means_no_restriction(self) -> None:
"""allowed_models = None 表示不限制模型"""
data = [("key-1", "provider-1", None, ["OPENAI"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" in result
rules = result["provider-1"]
assert len(rules) == 1
allowed_models, usable_formats = rules[0]
assert allowed_models is None # None = 不限制
assert "OPENAI" in usable_formats
def test_allowed_models_list_is_preserved(self) -> None:
"""allowed_models 为有效列表时应正常返回"""
data = [("key-1", "provider-1", ["claude-3-opus", "claude-3-sonnet"], ["OPENAI"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" in result
rules = result["provider-1"]
assert len(rules) == 1
allowed_models, _ = rules[0]
assert allowed_models == ["claude-3-opus", "claude-3-sonnet"]
def test_skips_key_when_format_intersection_empty(self) -> None:
"""格式交集为空时不包含该 Key"""
# Key 支持 CLAUDE但请求的是 OPENAI
data = [("key-1", "provider-1", None, ["CLAUDE"])]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result or len(result.get("provider-1", [])) == 0
def test_multiple_keys_same_provider(self) -> None:
"""同一 Provider 下多个 Key 应合并规则"""
data = [
("key-1", "provider-1", None, ["OPENAI"]),
("key-2", "provider-1", ["claude-3-opus"], ["OPENAI"]),
]
result = ModelAvailabilityQuery.get_provider_key_rules(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" in result
rules = result["provider-1"]
assert len(rules) == 2
class TestGetProvidersWithActiveKeys:
"""测试 get_providers_with_active_keys"""
def test_empty_provider_ids_returns_empty_set(self) -> None:
"""空 provider_ids 应返回空集合"""
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession([]),
provider_ids=set(),
api_formats=["OPENAI"],
provider_to_endpoint_formats={},
)
assert result == set()
def test_skips_key_with_invalid_api_formats_type(self) -> None:
"""api_formats 类型异常时应跳过该 Key 并打日志"""
log_output = StringIO()
handler_id = logger.add(
log_output,
format="{message}",
level="WARNING",
filter=lambda record: "[ModelAvailability]" in record["message"],
)
try:
# (provider_id, api_formats) - api_formats 是字符串
data = [("provider-1", "OPENAI")]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result
log_content = log_output.getvalue()
assert "api_formats 类型异常" in log_content
finally:
logger.remove(handler_id)
def test_skips_key_with_none_api_formats(self) -> None:
"""api_formats 为 None 时应跳过该 Key不打日志"""
data = [("provider-1", None)]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result
def test_returns_provider_when_format_matches(self) -> None:
"""格式匹配时应返回该 Provider"""
data = [("provider-1", ["OPENAI", "CLAUDE"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" in result
def test_skips_provider_when_format_intersection_empty(self) -> None:
"""格式交集为空时不返回该 Provider"""
# Key 支持 GEMINI但请求的是 OPENAI且端点支持 OPENAI
data = [("provider-1", ["GEMINI"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" not in result
def test_skips_provider_not_in_endpoint_formats(self) -> None:
"""Provider 不在 endpoint_formats 中时应跳过"""
data = [("provider-1", ["OPENAI"])]
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"],
provider_to_endpoint_formats={}, # 空,无匹配端点
)
assert "provider-1" not in result
def test_case_insensitive_format_matching(self) -> None:
"""格式匹配应忽略大小写"""
data = [("provider-1", ["openai"])] # 小写
result = ModelAvailabilityQuery.get_providers_with_active_keys(
FakeSession(data),
provider_ids={"provider-1"},
api_formats=["OPENAI"], # 大写
provider_to_endpoint_formats={"provider-1": {"OPENAI"}},
)
assert "provider-1" in result
class TestGetProvidersWithActiveEndpointsSourceCode:
"""静态验证 get_providers_with_active_endpoints 的核心实现"""
def test_checks_provider_is_active(self) -> None:
"""应检查 Provider.is_active"""
source = inspect.getsource(ModelAvailabilityQuery.get_providers_with_active_endpoints)
assert "Provider.is_active" in source
def test_checks_endpoint_is_active(self) -> None:
"""应检查 ProviderEndpoint.is_active"""
source = inspect.getsource(ModelAvailabilityQuery.get_providers_with_active_endpoints)
assert "ProviderEndpoint.is_active" in source
def test_joins_provider_table(self) -> None:
"""应 JOIN Provider 表"""
source = inspect.getsource(ModelAvailabilityQuery.get_providers_with_active_endpoints)
assert ".join(" in source.lower() or "join(" in source.lower()
class TestFindByGlobalModelNameSourceCode:
"""静态验证 find_by_global_model_name 的核心实现"""
def test_uses_base_active_models(self) -> None:
"""应复用 base_active_models"""
source = inspect.getsource(ModelAvailabilityQuery.find_by_global_model_name)
assert "base_active_models" in source
def test_filters_by_global_model_name(self) -> None:
"""应按 GlobalModel.name 过滤"""
source = inspect.getsource(ModelAvailabilityQuery.find_by_global_model_name)
assert "GlobalModel.name" in source
class TestFindModelByIdNoFallback:
"""测试 find_model_by_id 不再回退到 provider_model_name"""
def test_source_code_no_provider_model_name_in_function(self) -> None:
"""
静态验证find_model_by_id 函数中不应出现 provider_model_name
直接检查字符串是否出现,覆盖所有可能的写法。
"""
from src.api.base import models_service
source = inspect.getsource(models_service.find_model_by_id)
assert "provider_model_name" not in source, (
"find_model_by_id 中不应出现 provider_model_name已删除回退逻辑"
)
class TestGetAvailableModelIdsWithMappings:
"""测试 _get_available_model_ids_for_format 支持 model_mappings"""
def test_source_code_uses_check_model_allowed_with_mappings(self) -> None:
"""静态验证:应使用 check_model_allowed_with_mappings 而非 check_model_allowed"""
from src.api.base import models_service
source = inspect.getsource(models_service._get_available_model_ids_for_format)
assert "check_model_allowed_with_mappings" in source, (
"_get_available_model_ids_for_format 应使用 check_model_allowed_with_mappings"
)
def test_source_code_extracts_model_mappings_from_config(self) -> None:
"""静态验证:应从 global_model.config 提取 model_mappings"""
from src.api.base import models_service
source = inspect.getsource(models_service._get_available_model_ids_for_format)
assert "model_mappings" in source, (
"_get_available_model_ids_for_format 应提取 model_mappings"
)
assert ".config" in source or "config" in source, (
"_get_available_model_ids_for_format 应从 config 获取 model_mappings"
)
class TestModelMappingsIntegration:
"""
测试 model_mappings 正则映射的集成场景
场景:
- GlobalModel.name = "claude-haiku-4-5-20251001"
- Model.provider_model_name = "claude-haiku-4-5-20251001"
- GlobalModel.config.model_mappings = ["claude-3-5-haiku-.*"]
- Key.allowed_models = ["claude-3-5-haiku-20251001"]
预期:通过 model_mappings 正则匹配,模型应出现在可用列表中
"""
def test_mapping_pattern_matches_allowed_model(self) -> None:
"""model_mappings 正则应能匹配 Key.allowed_models 中的模型"""
from src.core.model_permissions import check_model_allowed_with_mappings
# 模拟你描述的场景
is_allowed, matched = check_model_allowed_with_mappings(
model_name="claude-haiku-4-5-20251001", # GlobalModel.name
allowed_models=["claude-3-5-haiku-20251001"], # Key.allowed_models
resolved_model_name="claude-haiku-4-5-20251001", # Model.provider_model_name
model_mappings=["claude-3-5-haiku-.*"], # GlobalModel.config.model_mappings
)
assert is_allowed is True
assert matched == "claude-3-5-haiku-20251001"
def test_exact_match_takes_priority_over_mapping(self) -> None:
"""精确匹配优先于映射匹配"""
from src.core.model_permissions import check_model_allowed_with_mappings
is_allowed, matched = check_model_allowed_with_mappings(
model_name="claude-haiku-4-5-20251001",
allowed_models=["claude-haiku-4-5-20251001"], # 精确匹配
resolved_model_name="claude-haiku-4-5-20251001",
model_mappings=["claude-3-5-haiku-.*"],
)
assert is_allowed is True
assert matched is None # 精确匹配时 matched 为 None
def test_no_mapping_match_returns_false(self) -> None:
"""映射不匹配时返回 False"""
from src.core.model_permissions import check_model_allowed_with_mappings
is_allowed, matched = check_model_allowed_with_mappings(
model_name="claude-haiku-4-5-20251001",
allowed_models=["gpt-4o"], # 与映射模式不匹配
resolved_model_name="claude-haiku-4-5-20251001",
model_mappings=["claude-3-5-haiku-.*"],
)
assert is_allowed is False
assert matched is None
def test_multiple_mappings_first_match_wins(self) -> None:
"""多个映射时,按 allowed_models 排序后第一个匹配的生效"""
from src.core.model_permissions import check_model_allowed_with_mappings
is_allowed, matched = check_model_allowed_with_mappings(
model_name="target-model",
allowed_models=["b-model-1", "a-model-1"],
resolved_model_name="target-model",
model_mappings=[".*-model-1"], # 匹配两个
)
assert is_allowed is True
# 按字母排序a-model-1 先匹配
assert matched == "a-model-1"
# ============================================================================
# AccessRestrictions 测试
# ============================================================================
class TestAccessRestrictionsFromApiKeyAndUser:
"""测试 AccessRestrictions.from_api_key_and_user 合并逻辑"""
def test_both_none_returns_no_restrictions(self) -> None:
"""API Key 和 User 都为 None 时返回无限制"""
from src.api.base.models_service import AccessRestrictions
result = AccessRestrictions.from_api_key_and_user(None, None)
assert result.allowed_providers is None
assert result.allowed_models is None
assert result.allowed_api_formats is None
def test_api_key_restrictions_take_priority(self) -> None:
"""API Key 的限制优先于 User 的限制"""
from unittest.mock import MagicMock
from src.api.base.models_service import AccessRestrictions
api_key = MagicMock()
api_key.allowed_providers = ["provider-a"]
api_key.allowed_models = ["model-a"]
api_key.allowed_api_formats = ["OPENAI"]
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-a"]
assert result.allowed_models == ["model-a"]
assert result.allowed_api_formats == ["OPENAI"]
def test_user_restrictions_used_when_api_key_has_none(self) -> None:
"""API Key 无限制时使用 User 的限制"""
from unittest.mock import MagicMock
from src.api.base.models_service import AccessRestrictions
api_key = MagicMock()
api_key.allowed_providers = None
api_key.allowed_models = None
api_key.allowed_api_formats = None
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-b"]
assert result.allowed_models == ["model-b"]
assert result.allowed_api_formats == ["CLAUDE"]
def test_partial_api_key_restrictions(self) -> None:
"""API Key 部分限制时,其余字段从 User 获取"""
from unittest.mock import MagicMock
from src.api.base.models_service import AccessRestrictions
api_key = MagicMock()
api_key.allowed_providers = ["provider-a"]
api_key.allowed_models = None # 无限制
api_key.allowed_api_formats = None # 无限制
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
result = AccessRestrictions.from_api_key_and_user(api_key, user)
assert result.allowed_providers == ["provider-a"] # 来自 API Key
assert result.allowed_models == ["model-b"] # 来自 User
assert result.allowed_api_formats == ["CLAUDE"] # 来自 User
def test_only_api_key_provided(self) -> None:
"""只提供 API Key 时使用其限制"""
from unittest.mock import MagicMock
from src.api.base.models_service import AccessRestrictions
api_key = MagicMock()
api_key.allowed_providers = ["provider-a"]
api_key.allowed_models = ["model-a"]
api_key.allowed_api_formats = ["OPENAI"]
result = AccessRestrictions.from_api_key_and_user(api_key, None)
assert result.allowed_providers == ["provider-a"]
assert result.allowed_models == ["model-a"]
assert result.allowed_api_formats == ["OPENAI"]
def test_only_user_provided(self) -> None:
"""只提供 User 时使用其限制"""
from unittest.mock import MagicMock
from src.api.base.models_service import AccessRestrictions
user = MagicMock()
user.allowed_providers = ["provider-b"]
user.allowed_models = ["model-b"]
user.allowed_api_formats = ["CLAUDE"]
result = AccessRestrictions.from_api_key_and_user(None, user)
assert result.allowed_providers == ["provider-b"]
assert result.allowed_models == ["model-b"]
assert result.allowed_api_formats == ["CLAUDE"]
class TestAccessRestrictionsIsApiFormatAllowed:
"""测试 AccessRestrictions.is_api_format_allowed"""
def test_none_allowed_formats_means_all_allowed(self) -> None:
"""allowed_api_formats = None 表示所有格式都允许"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=None)
assert restrictions.is_api_format_allowed("OPENAI") is True
assert restrictions.is_api_format_allowed("CLAUDE") is True
assert restrictions.is_api_format_allowed("GEMINI") is True
def test_format_in_allowed_list(self) -> None:
"""格式在允许列表中时返回 True"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=["OPENAI", "CLAUDE"])
assert restrictions.is_api_format_allowed("OPENAI") is True
assert restrictions.is_api_format_allowed("CLAUDE") is True
def test_format_not_in_allowed_list(self) -> None:
"""格式不在允许列表中时返回 False"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=["OPENAI"])
assert restrictions.is_api_format_allowed("CLAUDE") is False
assert restrictions.is_api_format_allowed("GEMINI") is False
def test_empty_allowed_list_blocks_all(self) -> None:
"""空允许列表阻止所有格式"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_api_formats=[])
assert restrictions.is_api_format_allowed("OPENAI") is False
assert restrictions.is_api_format_allowed("CLAUDE") is False
class TestAccessRestrictionsIsModelAllowed:
"""测试 AccessRestrictions.is_model_allowed"""
def test_no_restrictions_allows_all(self) -> None:
"""无限制时允许所有模型"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_providers=None, allowed_models=None)
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is True
assert restrictions.is_model_allowed("gpt-4", "provider-b") is True
def test_provider_restriction_blocks_unallowed_provider(self) -> None:
"""Provider 限制阻止不在列表中的 Provider"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_providers=["provider-a"])
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is True
assert restrictions.is_model_allowed("claude-3-opus", "provider-b") is False
def test_model_restriction_blocks_unallowed_model(self) -> None:
"""模型限制阻止不在列表中的模型"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_models=["claude-3-opus"])
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is True
assert restrictions.is_model_allowed("gpt-4", "provider-a") is False
def test_both_restrictions_must_pass(self) -> None:
"""Provider 和模型限制都必须通过"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(
allowed_providers=["provider-a"],
allowed_models=["claude-3-opus"],
)
# 两者都满足
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is True
# Provider 不满足
assert restrictions.is_model_allowed("claude-3-opus", "provider-b") is False
# 模型不满足
assert restrictions.is_model_allowed("gpt-4", "provider-a") is False
# 两者都不满足
assert restrictions.is_model_allowed("gpt-4", "provider-b") is False
def test_empty_provider_list_blocks_all(self) -> None:
"""空 Provider 列表阻止所有"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_providers=[])
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is False
def test_empty_model_list_blocks_all(self) -> None:
"""空模型列表阻止所有"""
from src.api.base.models_service import AccessRestrictions
restrictions = AccessRestrictions(allowed_models=[])
assert restrictions.is_model_allowed("claude-3-opus", "provider-a") is False