feat: 格式转换支持三层开关控制(全局/Provider/端点)

- _get_convertible_formats 始终返回所有可转换格式,由下游精确过滤
- get_compatible_provider_formats 增加 Provider 级别 enable_format_conversion 判断
- Provider 设置更新后自动失效相关缓存(模型列表、解析缓存、Provider 缓存)
- 前端 toggleFormatConversion 使用接口返回的完整对象更新本地状态
This commit is contained in:
fawney19
2026-02-06 23:01:56 +08:00
parent a1ea060cd9
commit 18d0af3dbf
5 changed files with 43 additions and 15 deletions

View File

@@ -12,6 +12,7 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.models_service import invalidate_models_list_cache
from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import ProviderBillingType
from src.core.exceptions import NotFoundException
@@ -31,6 +32,8 @@ from src.models.endpoint_models import (
ProviderUpdateRequest,
ProviderWithEndpointsSummary,
)
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
router = APIRouter(tags=["Provider Summary"])
pipeline = ApiRequestPipeline()
@@ -499,10 +502,21 @@ class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
for key, value in update_dict.items():
setattr(provider, key, value)
provider.updated_at = datetime.now(timezone.utc)
db.commit()
db.refresh(provider)
admin_name = context.user.username if context.user else "admin"
logger.info(f"Provider {provider.name} updated by {admin_name}: {update_dict}")
# 缓存失效
affects_model_visibility = {"is_active", "enable_format_conversion"} & update_dict.keys()
if affects_model_visibility:
await invalidate_models_list_cache()
if "is_active" in update_dict:
await ModelCacheService.invalidate_all_resolve_cache()
if "billing_type" in update_dict:
await ProviderCacheService.invalidate_provider_cache(provider.id)
return _build_provider_summary(db, provider)

View File

@@ -289,6 +289,7 @@ def get_compatible_provider_formats(
ProviderEndpoint.api_family,
ProviderEndpoint.endpoint_kind,
ProviderEndpoint.format_acceptance_config,
Provider.enable_format_conversion,
)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
@@ -302,16 +303,24 @@ def get_compatible_provider_formats(
)
provider_to_formats: dict[str, set[str]] = {}
for provider_id, api_family, endpoint_kind, format_acceptance_config in endpoint_rows:
for (
provider_id,
api_family,
endpoint_kind,
format_acceptance_config,
provider_conversion_enabled,
) in endpoint_rows:
if not provider_id or not api_family or not endpoint_kind:
continue
endpoint_format = normalize_endpoint_signature(f"{api_family}:{endpoint_kind}")
skip_endpoint_check = global_conversion_enabled or bool(provider_conversion_enabled)
is_compatible, _needs_conversion, _reason = is_format_compatible(
client_format_norm,
endpoint_format,
format_acceptance_config,
is_stream=False,
effective_conversion_enabled=global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
if not is_compatible:
continue

View File

@@ -80,19 +80,15 @@ def _is_format_conversion_enabled(db: Session) -> bool:
return SystemConfigService.is_format_conversion_enabled(db)
def _get_convertible_formats(client_format: str, global_conversion_enabled: bool) -> list[str]:
def _get_convertible_formats(client_format: str) -> list[str]:
"""
获取客户端格式可转换到的所有目标格式列表
当启用格式转换时,返回所有可以转换的格式
否则只返回客户端格式本身(不包括同族的其他格式)
始终返回所有转换的格式(包括客户端格式本身),
由下游 get_compatible_provider_formats 按三层开关(全局/Provider/端点)精确过滤
"""
client_format_norm = normalize_endpoint_signature(client_format)
# 格式转换关闭时,只返回客户端格式本身
if not global_conversion_enabled:
return [client_format_norm]
# 收集所有可转换的格式
register_default_normalizers()
convertible_formats: list[str] = []
@@ -501,7 +497,7 @@ async def list_models(
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats = _get_convertible_formats(api_format)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)
@@ -605,7 +601,7 @@ async def retrieve_model(
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats = _get_convertible_formats(api_format)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)
@@ -688,7 +684,7 @@ async def list_models_gemini(
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats = _get_convertible_formats(api_format)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)
@@ -767,7 +763,7 @@ async def get_model_gemini(
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats = _get_convertible_formats(api_format)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
)

View File

@@ -1261,6 +1261,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
ProviderEndpoint.api_family,
ProviderEndpoint.endpoint_kind,
ProviderEndpoint.format_acceptance_config,
Provider.enable_format_conversion,
)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
@@ -1282,11 +1283,18 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
# 只要端点能被任意一种客户端格式访问,就将其 Provider 加入结果
provider_to_formats: dict[str, set[str]] = {}
for provider_id, api_family, endpoint_kind, format_acceptance_config in endpoint_rows:
for (
provider_id,
api_family,
endpoint_kind,
format_acceptance_config,
provider_conversion_enabled,
) in endpoint_rows:
if not provider_id or not api_family or not endpoint_kind:
continue
endpoint_format = make_signature_key(str(api_family), str(endpoint_kind))
skip_endpoint_check = global_conversion_enabled or bool(provider_conversion_enabled)
# 检查该端点是否能被任意客户端格式访问
for client_format in all_formats:
@@ -1296,6 +1304,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
format_acceptance_config,
is_stream=False,
effective_conversion_enabled=global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
if is_compatible:
provider_to_formats.setdefault(provider_id, set()).add(endpoint_format)