mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 用 EndpointFetchConfig 纯数据类替代 ORM 对象传递,统一上游模型缓存管理
- 引入 EndpointFetchConfig dataclass 替代直接传递 ProviderEndpoint ORM 对象, 避免 DB session 关闭后 DetachedInstanceError - 新增 build_format_to_config() 统一构建 api_format -> EndpointFetchConfig 映射 - KeyAllowedModelsDialog 改用 useUpstreamModelsCache composable 管理上游模型获取 - useUpstreamModelsCache 增加 error 字段透传部分格式获取失败的 warning - 删除废弃的 queryProviderUpstreamModels API 函数 - ProviderCandidate 添加 __lt__ 方法支持排序比较
This commit is contained in:
@@ -122,29 +122,6 @@ export async function batchAssignModelsToProvider(
|
|||||||
return response.data
|
return response.data
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
|
||||||
* 查询提供商的上游模型列表
|
|
||||||
*/
|
|
||||||
export async function queryProviderUpstreamModels(
|
|
||||||
providerId: string
|
|
||||||
): Promise<{
|
|
||||||
success: boolean
|
|
||||||
data: {
|
|
||||||
models: UpstreamModel[]
|
|
||||||
error: string | null
|
|
||||||
}
|
|
||||||
provider: {
|
|
||||||
id: string
|
|
||||||
name: string
|
|
||||||
display_name: string
|
|
||||||
}
|
|
||||||
}> {
|
|
||||||
const response = await client.post('/api/admin/provider-query/models', {
|
|
||||||
provider_id: providerId,
|
|
||||||
})
|
|
||||||
return response.data
|
|
||||||
}
|
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 从上游提供商导入模型
|
* 从上游提供商导入模型
|
||||||
* @param providerId 提供商 ID
|
* @param providerId 提供商 ID
|
||||||
|
|||||||
@@ -192,7 +192,6 @@ import Badge from '@/components/ui/badge.vue'
|
|||||||
import Checkbox from '@/components/ui/checkbox.vue'
|
import Checkbox from '@/components/ui/checkbox.vue'
|
||||||
import { useToast } from '@/composables/useToast'
|
import { useToast } from '@/composables/useToast'
|
||||||
import { parseUpstreamModelError } from '@/utils/errorParser'
|
import { parseUpstreamModelError } from '@/utils/errorParser'
|
||||||
import { adminApi } from '@/api/admin'
|
|
||||||
import {
|
import {
|
||||||
importModelsFromUpstream,
|
importModelsFromUpstream,
|
||||||
getProviderModels,
|
getProviderModels,
|
||||||
@@ -200,6 +199,7 @@ import {
|
|||||||
type UpstreamModel,
|
type UpstreamModel,
|
||||||
API_FORMAT_LABELS,
|
API_FORMAT_LABELS,
|
||||||
} from '@/api/endpoints'
|
} from '@/api/endpoints'
|
||||||
|
import { useUpstreamModelsCache } from '../composables/useUpstreamModelsCache'
|
||||||
|
|
||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
open: boolean
|
open: boolean
|
||||||
@@ -213,6 +213,7 @@ const emit = defineEmits<{
|
|||||||
}>()
|
}>()
|
||||||
|
|
||||||
const { success, error: showError } = useToast()
|
const { success, error: showError } = useToast()
|
||||||
|
const { fetchModels: fetchCachedModels } = useUpstreamModelsCache()
|
||||||
|
|
||||||
const isOpen = computed(() => props.open)
|
const isOpen = computed(() => props.open)
|
||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
@@ -275,7 +276,7 @@ async function loadExistingModels() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 获取上游模型(获取所有 Key 的聚合结果)
|
// 获取上游模型(获取所有 Key 的聚合结果,通过 useUpstreamModelsCache 统一管理)
|
||||||
async function fetchUpstreamModels() {
|
async function fetchUpstreamModels() {
|
||||||
if (!props.providerId) return
|
if (!props.providerId) return
|
||||||
|
|
||||||
@@ -285,27 +286,26 @@ async function fetchUpstreamModels() {
|
|||||||
try {
|
try {
|
||||||
// 不传 apiKeyId,后端会遍历所有 Key 并聚合结果。
|
// 不传 apiKeyId,后端会遍历所有 Key 并聚合结果。
|
||||||
// 已查询过再点“刷新”时,强制跳过后端缓存,避免长期 TTL 导致模型列表不更新。
|
// 已查询过再点“刷新”时,强制跳过后端缓存,避免长期 TTL 导致模型列表不更新。
|
||||||
const response = await adminApi.queryProviderModels(props.providerId, undefined, hasQueried.value)
|
const result = await fetchCachedModels(props.providerId, undefined, hasQueried.value)
|
||||||
|
|
||||||
if (response.success && response.data?.models) {
|
if (result.models.length > 0) {
|
||||||
upstreamModels.value = response.data.models
|
upstreamModels.value = result.models
|
||||||
// 默认选中所有新模型
|
// 默认选中所有新模型
|
||||||
selectedModels.value = response.data.models
|
selectedModels.value = result.models
|
||||||
.filter((m: UpstreamModel) => !existingModelIds.value.has(m.id))
|
.filter((m: UpstreamModel) => !existingModelIds.value.has(m.id))
|
||||||
.map((m: UpstreamModel) => m.id)
|
.map((m: UpstreamModel) => m.id)
|
||||||
hasQueried.value = true
|
hasQueried.value = true
|
||||||
// 如果有部分失败,显示警告提示
|
// 如果有部分失败,显示警告提示
|
||||||
if (response.data.error) {
|
if (result.error) {
|
||||||
// 使用友好的错误解析
|
showError(`部分格式获取失败: ${result.error}`, '警告')
|
||||||
showError(`部分格式获取失败: ${parseUpstreamModelError(response.data.error)}`, '警告')
|
|
||||||
}
|
}
|
||||||
|
} else if (result.error) {
|
||||||
|
errorMessage.value = result.error
|
||||||
} else {
|
} else {
|
||||||
// 使用友好的错误解析
|
// 上游返回空列表但无错误
|
||||||
const rawError = response.data?.error || '获取上游模型失败'
|
hasQueried.value = true
|
||||||
errorMessage.value = parseUpstreamModelError(rawError)
|
|
||||||
}
|
}
|
||||||
} catch (err: any) {
|
} catch (err: any) {
|
||||||
// 使用友好的错误解析
|
|
||||||
const rawError = err.response?.data?.detail || err.message || '获取上游模型失败'
|
const rawError = err.response?.data?.detail || err.message || '获取上游模型失败'
|
||||||
errorMessage.value = parseUpstreamModelError(rawError)
|
errorMessage.value = parseUpstreamModelError(rawError)
|
||||||
} finally {
|
} finally {
|
||||||
|
|||||||
@@ -54,6 +54,8 @@ export function useUpstreamModelsCache() {
|
|||||||
if (response.success && response.data?.models) {
|
if (response.success && response.data?.models) {
|
||||||
return {
|
return {
|
||||||
models: response.data.models,
|
models: response.data.models,
|
||||||
|
// 传递部分格式获取失败的 warning(后端 success=true 但仍可能附带 error)
|
||||||
|
error: response.data.error ? parseUpstreamModelError(response.data.error) : undefined,
|
||||||
fromCache: response.data.from_cache
|
fromCache: response.data.from_cache
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -27,7 +27,9 @@ from src.services.model.fetch_scheduler import (
|
|||||||
set_upstream_models_to_cache,
|
set_upstream_models_to_cache,
|
||||||
)
|
)
|
||||||
from src.services.model.upstream_fetcher import (
|
from src.services.model.upstream_fetcher import (
|
||||||
|
EndpointFetchConfig,
|
||||||
UpstreamModelsFetchContext,
|
UpstreamModelsFetchContext,
|
||||||
|
build_format_to_config,
|
||||||
fetch_models_for_key,
|
fetch_models_for_key,
|
||||||
get_adapter_for_format,
|
get_adapter_for_format,
|
||||||
)
|
)
|
||||||
@@ -171,11 +173,8 @@ async def query_available_models(
|
|||||||
if not provider:
|
if not provider:
|
||||||
raise HTTPException(status_code=404, detail="Provider not found")
|
raise HTTPException(status_code=404, detail="Provider not found")
|
||||||
|
|
||||||
# 构建 api_format -> endpoint 映射
|
# 构建 api_format -> EndpointFetchConfig 映射(纯数据,不依赖 ORM session)
|
||||||
format_to_endpoint: dict[str, ProviderEndpoint] = {}
|
format_to_endpoint = build_format_to_config(provider.endpoints)
|
||||||
for endpoint in provider.endpoints:
|
|
||||||
if endpoint.is_active:
|
|
||||||
format_to_endpoint[endpoint.api_format] = endpoint
|
|
||||||
|
|
||||||
if not format_to_endpoint:
|
if not format_to_endpoint:
|
||||||
raise HTTPException(status_code=400, detail="No active endpoints found for this provider")
|
raise HTTPException(status_code=400, detail="No active endpoints found for this provider")
|
||||||
@@ -327,7 +326,7 @@ def _aggregate_models_by_id(models: list[dict]) -> list[dict]:
|
|||||||
async def _fetch_models_for_single_key(
|
async def _fetch_models_for_single_key(
|
||||||
provider: Provider,
|
provider: Provider,
|
||||||
api_key_id: str,
|
api_key_id: str,
|
||||||
format_to_endpoint: dict[str, ProviderEndpoint],
|
format_to_endpoint: dict[str, EndpointFetchConfig],
|
||||||
force_refresh: bool,
|
force_refresh: bool,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""获取单个 Key 的模型列表"""
|
"""获取单个 Key 的模型列表"""
|
||||||
|
|||||||
@@ -28,6 +28,16 @@ class ProviderCandidate:
|
|||||||
if self.metadata is None:
|
if self.metadata is None:
|
||||||
self.metadata = {}
|
self.metadata = {}
|
||||||
|
|
||||||
|
def __lt__(self, other: object) -> bool:
|
||||||
|
if not isinstance(other, ProviderCandidate):
|
||||||
|
return NotImplemented
|
||||||
|
# 优先级数字越大越优先,权重越大越优先
|
||||||
|
return (-self.priority, -self.weight, str(getattr(self.provider, "id", ""))) < (
|
||||||
|
-other.priority,
|
||||||
|
-other.weight,
|
||||||
|
str(getattr(other.provider, "id", "")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class SelectionResult:
|
class SelectionResult:
|
||||||
|
|||||||
@@ -28,7 +28,9 @@ from src.core.provider_types import ProviderType
|
|||||||
from src.database import create_session
|
from src.database import create_session
|
||||||
from src.models.database import Provider, ProviderAPIKey
|
from src.models.database import Provider, ProviderAPIKey
|
||||||
from src.services.model.upstream_fetcher import (
|
from src.services.model.upstream_fetcher import (
|
||||||
|
EndpointFetchConfig,
|
||||||
UpstreamModelsFetchContext,
|
UpstreamModelsFetchContext,
|
||||||
|
build_format_to_config,
|
||||||
fetch_models_for_key,
|
fetch_models_for_key,
|
||||||
merge_upstream_metadata,
|
merge_upstream_metadata,
|
||||||
)
|
)
|
||||||
@@ -62,7 +64,7 @@ class PreparedModelsFetchContext:
|
|||||||
auth_type: str
|
auth_type: str
|
||||||
encrypted_api_key: str
|
encrypted_api_key: str
|
||||||
encrypted_auth_config: str | None
|
encrypted_auth_config: str | None
|
||||||
format_to_endpoint: dict[str, Any]
|
format_to_endpoint: dict[str, EndpointFetchConfig]
|
||||||
proxy_config: dict[str, Any] | None
|
proxy_config: dict[str, Any] | None
|
||||||
|
|
||||||
|
|
||||||
@@ -422,11 +424,8 @@ class ModelFetchScheduler:
|
|||||||
db.commit()
|
db.commit()
|
||||||
return "error"
|
return "error"
|
||||||
|
|
||||||
# 构建 api_format -> endpoint 映射
|
# 构建 api_format -> EndpointFetchConfig 映射(纯数据,session 无关)
|
||||||
format_to_endpoint: dict[str, Any] = {}
|
format_to_endpoint = build_format_to_config(provider.endpoints) # type: ignore[attr-defined]
|
||||||
for endpoint in provider.endpoints: # type: ignore[attr-defined]
|
|
||||||
if endpoint.is_active:
|
|
||||||
format_to_endpoint[endpoint.api_format] = endpoint
|
|
||||||
|
|
||||||
if not format_to_endpoint:
|
if not format_to_endpoint:
|
||||||
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
||||||
|
|||||||
@@ -7,6 +7,7 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from collections.abc import Iterable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Awaitable, Callable
|
from typing import Any, Awaitable, Callable
|
||||||
|
|
||||||
@@ -14,7 +15,6 @@ import httpx
|
|||||||
|
|
||||||
from src.core.api_format import get_extra_headers_from_endpoint
|
from src.core.api_format import get_extra_headers_from_endpoint
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ProviderEndpoint
|
|
||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context
|
||||||
|
|
||||||
# 并发请求限制
|
# 并发请求限制
|
||||||
@@ -35,13 +35,42 @@ _ModelsFetcher = Callable[
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class EndpointFetchConfig:
|
||||||
|
"""端点获取配置(纯数据,不依赖 DB session)。
|
||||||
|
|
||||||
|
从 ProviderEndpoint ORM 对象提取必要字段,确保在 DB session 关闭后
|
||||||
|
仍可安全使用(避免 DetachedInstanceError)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
base_url: str
|
||||||
|
extra_headers: dict[str, str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def build_format_to_config(endpoints: Iterable[Any]) -> dict[str, EndpointFetchConfig]:
|
||||||
|
"""将活跃的 ProviderEndpoint 转换为 api_format -> EndpointFetchConfig 映射。
|
||||||
|
|
||||||
|
应在 DB session 活跃时调用,提取 ORM 对象上的 base_url 和 header_rules,
|
||||||
|
转换为 session 无关的纯数据结构。
|
||||||
|
"""
|
||||||
|
result: dict[str, EndpointFetchConfig] = {}
|
||||||
|
for ep in endpoints:
|
||||||
|
if not getattr(ep, "is_active", False):
|
||||||
|
continue
|
||||||
|
result[ep.api_format] = EndpointFetchConfig(
|
||||||
|
base_url=ep.base_url,
|
||||||
|
extra_headers=get_extra_headers_from_endpoint(ep),
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class UpstreamModelsFetchContext:
|
class UpstreamModelsFetchContext:
|
||||||
"""上游模型获取上下文(Key 级别)。"""
|
"""上游模型获取上下文(Key 级别)。"""
|
||||||
|
|
||||||
provider_type: str
|
provider_type: str
|
||||||
api_key_value: str
|
api_key_value: str
|
||||||
format_to_endpoint: dict[str, Any]
|
format_to_endpoint: dict[str, EndpointFetchConfig]
|
||||||
proxy_config: dict[str, Any] | None = None
|
proxy_config: dict[str, Any] | None = None
|
||||||
auth_config: dict[str, Any] | None = None
|
auth_config: dict[str, Any] | None = None
|
||||||
|
|
||||||
@@ -149,7 +178,7 @@ def get_adapter_for_format(api_format: str) -> type | None:
|
|||||||
|
|
||||||
def build_all_format_configs(
|
def build_all_format_configs(
|
||||||
api_key_value: str,
|
api_key_value: str,
|
||||||
format_to_endpoint: dict[str, ProviderEndpoint],
|
format_to_endpoint: dict[str, EndpointFetchConfig],
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
"""
|
"""
|
||||||
构建所有 API 格式的端点配置
|
构建所有 API 格式的端点配置
|
||||||
@@ -159,7 +188,7 @@ def build_all_format_configs(
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
api_key_value: 解密后的 API Key
|
api_key_value: 解密后的 API Key
|
||||||
format_to_endpoint: API 格式到端点的映射
|
format_to_endpoint: API 格式到 EndpointFetchConfig 的映射
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
端点配置列表,每个配置包含 api_key, base_url, api_format, extra_headers
|
端点配置列表,每个配置包含 api_key, base_url, api_format, extra_headers
|
||||||
@@ -172,13 +201,13 @@ def build_all_format_configs(
|
|||||||
for candidates in MODEL_FETCH_FORMAT_PRIORITY:
|
for candidates in MODEL_FETCH_FORMAT_PRIORITY:
|
||||||
fmt = next((f for f in candidates if f in format_to_endpoint), None)
|
fmt = next((f for f in candidates if f in format_to_endpoint), None)
|
||||||
if fmt is not None:
|
if fmt is not None:
|
||||||
ep = format_to_endpoint[fmt]
|
cfg = format_to_endpoint[fmt]
|
||||||
configs.append(
|
configs.append(
|
||||||
{
|
{
|
||||||
"api_key": api_key_value,
|
"api_key": api_key_value,
|
||||||
"base_url": ep.base_url,
|
"base_url": cfg.base_url,
|
||||||
"api_format": fmt,
|
"api_format": fmt,
|
||||||
"extra_headers": get_extra_headers_from_endpoint(ep),
|
"extra_headers": cfg.extra_headers,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return configs
|
return configs
|
||||||
|
|||||||
Reference in New Issue
Block a user