feat: 模型列表 API 支持格式转换兼容性过滤

- 新增 get_compatible_provider_formats 函数,基于端点 format_acceptance_config 过滤兼容的 Provider
- 模型列表查询根据客户端格式和全局转换开关返回可用模型
- 缓存 key 增加 client_format 维度,避免不同格式的缓存混用
- GlobalModel 解析支持 provider_model_mappings 和 model_mappings 匹配
- 修复格式兼容性检查中 config 类型校验缺失问题
This commit is contained in:
fawney19
2026-01-23 14:33:42 +08:00
parent 08a3148ab3
commit 86b9bc3036
4 changed files with 425 additions and 54 deletions

View File

@@ -17,6 +17,7 @@ from src.api.base.models_service import (
AccessRestrictions,
ModelInfo,
find_model_by_id,
get_compatible_provider_formats,
get_available_provider_ids,
list_available_models,
)
@@ -26,10 +27,13 @@ from src.core.api_format import (
ApiFormatDefinition,
detect_format_and_key_from_starlette,
)
from src.core.api_format.conversion import converter_registry
from src.core.api_format.utils import is_cli_format
from src.core.logger import logger
from src.database import get_db
from src.models.database import ApiKey, User
from src.services.auth.service import AuthService
from src.services.system.config import SystemConfigService
router = APIRouter(tags=["System Catalog"])
@@ -39,6 +43,9 @@ _CLAUDE_FORMATS = [APIFormat.CLAUDE.value]
_OPENAI_FORMATS = [APIFormat.OPENAI.value]
_GEMINI_FORMATS = [APIFormat.GEMINI.value]
# 所有非 CLI 格式(用于格式转换时的查询)
_ALL_CHAT_FORMATS = [APIFormat.CLAUDE.value, APIFormat.OPENAI.value, APIFormat.GEMINI.value]
def _extract_api_key_from_request(
request: Request, definition: ApiFormatDefinition
@@ -90,6 +97,52 @@ def _get_formats_for_api(api_format: str) -> list[str]:
return _OPENAI_FORMATS
def _is_format_conversion_enabled(db: Session) -> bool:
"""检查全局格式转换开关"""
return bool(SystemConfigService.get_config(db, "format_conversion_enabled", False))
def _get_convertible_formats(client_format: str, global_conversion_enabled: bool) -> list[str]:
"""
获取客户端格式可转换到的所有目标格式列表
当启用格式转换时,返回所有可以转换的格式;
否则只返回客户端格式本身。
"""
if not global_conversion_enabled:
return _get_formats_for_api(client_format)
client_format_upper = client_format.upper()
# CLI 格式不支持转换
if is_cli_format(client_format_upper):
return _get_formats_for_api(client_format)
# 收集所有可转换的格式
convertible_formats = []
for target_format in _ALL_CHAT_FORMATS:
# 相同格式始终可用
if target_format == client_format_upper:
convertible_formats.append(target_format)
continue
# 检查是否有双向转换器
if converter_registry.can_convert_full(client_format_upper, target_format, require_stream=False):
convertible_formats.append(target_format)
return convertible_formats if convertible_formats else _get_formats_for_api(client_format)
def _flatten_provider_formats(provider_to_formats: dict[str, set[str]]) -> list[str]:
"""合并 Provider 格式映射为唯一格式列表"""
if not provider_to_formats:
return []
all_formats: set[str] = set()
for formats in provider_to_formats.values():
all_formats.update(formats)
return sorted(all_formats)
def _build_empty_list_response(api_format: str) -> dict:
"""根据 API 格式构建空列表响应"""
if api_format == "claude":
@@ -454,17 +507,32 @@ async def list_models(
# 构建访问限制
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 检查 API 格式限制
formats = _get_formats_for_api(api_format)
formats, empty_response = _filter_formats_by_restrictions(formats, restrictions, api_format)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)
if empty_response is not None:
return empty_response
available_provider_ids = get_available_provider_ids(db, formats)
provider_to_formats = get_compatible_provider_formats(
db, api_format, candidate_formats, global_conversion_enabled
)
formats = _flatten_provider_formats(provider_to_formats)
available_provider_ids = get_available_provider_ids(db, formats, provider_to_formats)
if not available_provider_ids:
return _build_empty_list_response(api_format)
models = await list_available_models(db, available_provider_ids, formats, restrictions)
models = await list_available_models(
db,
available_provider_ids,
formats,
restrictions,
provider_to_formats=provider_to_formats,
client_format=api_format,
)
logger.debug(f"[Models] 返回 {len(models)} 个模型")
if api_format == "claude":
@@ -543,14 +611,28 @@ async def retrieve_model(
# 构建访问限制
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 检查 API 格式限制
formats = _get_formats_for_api(api_format)
formats, _ = _filter_formats_by_restrictions(formats, restrictions, api_format)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)
provider_to_formats = get_compatible_provider_formats(
db, api_format, candidate_formats, global_conversion_enabled
)
formats = _flatten_provider_formats(provider_to_formats)
if not formats:
return _build_404_response(model_id, api_format)
available_provider_ids = get_available_provider_ids(db, formats)
model_info = find_model_by_id(db, model_id, available_provider_ids, formats, restrictions)
available_provider_ids = get_available_provider_ids(db, formats, provider_to_formats)
model_info = find_model_by_id(
db,
model_id,
available_provider_ids,
formats,
restrictions,
provider_to_formats=provider_to_formats,
)
if not model_info:
return _build_404_response(model_id, api_format)
@@ -614,18 +696,32 @@ async def list_models_gemini(
# 构建访问限制
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 检查 API 格式限制
formats, empty_response = _filter_formats_by_restrictions(
_GEMINI_FORMATS, restrictions, "gemini"
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats("gemini", global_conversion_enabled)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, "gemini"
)
if empty_response is not None:
return empty_response
available_provider_ids = get_available_provider_ids(db, formats)
provider_to_formats = get_compatible_provider_formats(
db, "gemini", candidate_formats, global_conversion_enabled
)
formats = _flatten_provider_formats(provider_to_formats)
available_provider_ids = get_available_provider_ids(db, formats, provider_to_formats)
if not available_provider_ids:
return {"models": []}
models = await list_available_models(db, available_provider_ids, formats, restrictions)
models = await list_available_models(
db,
available_provider_ids,
formats,
restrictions,
provider_to_formats=provider_to_formats,
client_format="gemini",
)
logger.debug(f"[Models] 返回 {len(models)} 个模型")
response = _build_gemini_list_response(models, page_size, page_token)
logger.debug(f"[Models] Gemini 响应: {response}")
@@ -681,14 +777,27 @@ async def get_model_gemini(
# 构建访问限制
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 检查 API 格式限制
formats, _ = _filter_formats_by_restrictions(_GEMINI_FORMATS, restrictions, "gemini")
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats("gemini", global_conversion_enabled)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, "gemini"
)
provider_to_formats = get_compatible_provider_formats(
db, "gemini", candidate_formats, global_conversion_enabled
)
formats = _flatten_provider_formats(provider_to_formats)
if not formats:
return _build_404_response(model_id, "gemini")
available_provider_ids = get_available_provider_ids(db, formats)
available_provider_ids = get_available_provider_ids(db, formats, provider_to_formats)
model_info = find_model_by_id(
db, model_id, available_provider_ids, formats, restrictions
db,
model_id,
available_provider_ids,
formats,
restrictions,
provider_to_formats=provider_to_formats,
)
if not model_info: