mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user