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