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
@@ -1025,8 +1025,8 @@ async function toggleFormatConversion() {
if (!provider.value) return if (!provider.value) return
const newValue = !provider.value.enable_format_conversion const newValue = !provider.value.enable_format_conversion
try { try {
await updateProvider(provider.value.id, { enable_format_conversion: newValue }) const updated = await updateProvider(provider.value.id, { enable_format_conversion: newValue })
provider.value.enable_format_conversion = newValue provider.value = updated
showSuccess(newValue ? '已启用格式转换' : '已禁用格式转换') showSuccess(newValue ? '已启用格式转换' : '已禁用格式转换')
emit('refresh') emit('refresh')
} catch { } catch {
+14
View File
@@ -12,6 +12,7 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext 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.api.base.pipeline import ApiRequestPipeline
from src.core.enums import ProviderBillingType from src.core.enums import ProviderBillingType
from src.core.exceptions import NotFoundException from src.core.exceptions import NotFoundException
@@ -31,6 +32,8 @@ from src.models.endpoint_models import (
ProviderUpdateRequest, ProviderUpdateRequest,
ProviderWithEndpointsSummary, ProviderWithEndpointsSummary,
) )
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
router = APIRouter(tags=["Provider Summary"]) router = APIRouter(tags=["Provider Summary"])
pipeline = ApiRequestPipeline() pipeline = ApiRequestPipeline()
@@ -499,10 +502,21 @@ class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
for key, value in update_dict.items(): for key, value in update_dict.items():
setattr(provider, key, value) setattr(provider, key, value)
provider.updated_at = datetime.now(timezone.utc)
db.commit() db.commit()
db.refresh(provider) db.refresh(provider)
admin_name = context.user.username if context.user else "admin" admin_name = context.user.username if context.user else "admin"
logger.info(f"Provider {provider.name} updated by {admin_name}: {update_dict}") 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) return _build_provider_summary(db, provider)
+10 -1
View File
@@ -289,6 +289,7 @@ def get_compatible_provider_formats(
ProviderEndpoint.api_family, ProviderEndpoint.api_family,
ProviderEndpoint.endpoint_kind, ProviderEndpoint.endpoint_kind,
ProviderEndpoint.format_acceptance_config, ProviderEndpoint.format_acceptance_config,
Provider.enable_format_conversion,
) )
.join(Provider, ProviderEndpoint.provider_id == Provider.id) .join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter( .filter(
@@ -302,16 +303,24 @@ def get_compatible_provider_formats(
) )
provider_to_formats: dict[str, set[str]] = {} 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: if not provider_id or not api_family or not endpoint_kind:
continue continue
endpoint_format = normalize_endpoint_signature(f"{api_family}:{endpoint_kind}") 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( is_compatible, _needs_conversion, _reason = is_format_compatible(
client_format_norm, client_format_norm,
endpoint_format, endpoint_format,
format_acceptance_config, format_acceptance_config,
is_stream=False, is_stream=False,
effective_conversion_enabled=global_conversion_enabled, effective_conversion_enabled=global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
) )
if not is_compatible: if not is_compatible:
continue continue
+7 -11
View File
@@ -80,19 +80,15 @@ def _is_format_conversion_enabled(db: Session) -> bool:
return SystemConfigService.is_format_conversion_enabled(db) 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) client_format_norm = normalize_endpoint_signature(client_format)
# 格式转换关闭时,只返回客户端格式本身
if not global_conversion_enabled:
return [client_format_norm]
# 收集所有可转换的格式 # 收集所有可转换的格式
register_default_normalizers() register_default_normalizers()
convertible_formats: list[str] = [] convertible_formats: list[str] = []
@@ -501,7 +497,7 @@ async def list_models(
# 获取可用格式(包括可转换的格式) # 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db) 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, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format candidate_formats, restrictions, api_format
) )
@@ -605,7 +601,7 @@ async def retrieve_model(
# 获取可用格式(包括可转换的格式) # 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db) 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, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format candidate_formats, restrictions, api_format
) )
@@ -688,7 +684,7 @@ async def list_models_gemini(
# 获取可用格式(包括可转换的格式) # 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db) 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, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format candidate_formats, restrictions, api_format
) )
@@ -767,7 +763,7 @@ async def get_model_gemini(
# 获取可用格式(包括可转换的格式) # 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled(db) 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, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format candidate_formats, restrictions, api_format
) )
+10 -1
View File
@@ -1261,6 +1261,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
ProviderEndpoint.api_family, ProviderEndpoint.api_family,
ProviderEndpoint.endpoint_kind, ProviderEndpoint.endpoint_kind,
ProviderEndpoint.format_acceptance_config, ProviderEndpoint.format_acceptance_config,
Provider.enable_format_conversion,
) )
.join(Provider, ProviderEndpoint.provider_id == Provider.id) .join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter( .filter(
@@ -1282,11 +1283,18 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
# 只要端点能被任意一种客户端格式访问,就将其 Provider 加入结果 # 只要端点能被任意一种客户端格式访问,就将其 Provider 加入结果
provider_to_formats: dict[str, set[str]] = {} 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: if not provider_id or not api_family or not endpoint_kind:
continue continue
endpoint_format = make_signature_key(str(api_family), str(endpoint_kind)) 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: for client_format in all_formats:
@@ -1296,6 +1304,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
format_acceptance_config, format_acceptance_config,
is_stream=False, is_stream=False,
effective_conversion_enabled=global_conversion_enabled, effective_conversion_enabled=global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
) )
if is_compatible: if is_compatible:
provider_to_formats.setdefault(provider_id, set()).add(endpoint_format) provider_to_formats.setdefault(provider_id, set()).add(endpoint_format)