mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 格式转换支持三层开关控制(全局/Provider/端点)
- _get_convertible_formats 始终返回所有可转换格式,由下游精确过滤 - get_compatible_provider_formats 增加 Provider 级别 enable_format_conversion 判断 - Provider 设置更新后自动失效相关缓存(模型列表、解析缓存、Provider 缓存) - 前端 toggleFormatConversion 使用接口返回的完整对象更新本地状态
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user