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:
fawney19
2026-02-09 12:47:47 +08:00
parent 8702786fa2
commit 8684548072
7 changed files with 71 additions and 55 deletions

View File

@@ -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

View File

@@ -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 {

View File

@@ -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 {

View File

@@ -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 的模型列表"""

View File

@@ -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:

View File

@@ -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}")

View File

@@ -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