mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 09:50:21 +08:00
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:
@@ -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_format(Key 直属 Provider,通过 api_formats 过滤)
|
- Provider 下有活跃的 Key 且支持该 api_format(Key 直属 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
|
||||||
|
|||||||
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
264
src/services/model/availability.py
Normal file
264
src/services/model/availability.py
Normal 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_available(is_available=NULL 视为可用,兼容历史数据)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def base_active_models(db: Session, eager_load: bool = False) -> Query:
|
||||||
|
"""
|
||||||
|
返回基础的可用模型查询(实体为 Model)
|
||||||
|
|
||||||
|
已包含条件:
|
||||||
|
- Model.is_active = True
|
||||||
|
- Model.is_available = True 或 NULL(NULL 视为可用)
|
||||||
|
- 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
|
||||||
|
|
||||||
@@ -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()
|
||||||
|
|||||||
728
tests/services/test_model_availability.py
Normal file
728
tests/services/test_model_availability.py
Normal file
@@ -0,0 +1,728 @@
|
|||||||
|
# mypy: disable-error-code="arg-type"
|
||||||
|
"""
|
||||||
|
ModelAvailabilityQuery 单元测试
|
||||||
|
|
||||||
|
说明:
|
||||||
|
- 本测试使用静态源码检查 + FakeSession/FakeQuery 风格,避免引入真实 DB 依赖
|
||||||
|
- 项目日志使用 loguru,pytest 的 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
|
||||||
|
|
||||||
|
|
||||||
Reference in New Issue
Block a user