mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat: 配置导入导出支持 ProxyNode,号池 key 可用性检查延迟到排序后分页执行
配置导入导出: - 导出/导入新增 ProxyNode(代理节点)数据 - 导入时自动建立 old_id -> new_id 映射,重映射 Provider/Endpoint/Key 中的 node_id - 前端预览和结果展示新增代理节点统计 号池调度优化: - CandidateBuilder 不再逐 key 调用 _check_key_availability,直接收集全部 active key - 将可用性检查参数打包到 PoolCandidate._deferred_check_params - PoolManager.select_pool_keys 排序后分页调用 availability_checker,找到足够可用 key 即停止 - 减少大号池场景下不必要的可用性检查开销
This commit is contained in:
@@ -47,11 +47,28 @@ export interface ConfigExportData {
|
||||
exported_at: string
|
||||
global_models: GlobalModelExport[]
|
||||
providers: ProviderExport[]
|
||||
proxy_nodes?: ProxyNodeExport[]
|
||||
ldap_config?: LDAPConfigExport | null
|
||||
oauth_providers?: OAuthProviderExport[]
|
||||
system_configs?: SystemConfigExport[]
|
||||
}
|
||||
|
||||
export interface ProxyNodeExport {
|
||||
id: string
|
||||
name: string
|
||||
ip: string
|
||||
port: number
|
||||
region?: string | null
|
||||
is_manual: boolean
|
||||
proxy_url?: string | null
|
||||
proxy_username?: string | null
|
||||
proxy_password?: string | null
|
||||
tunnel_mode: boolean
|
||||
heartbeat_interval: number
|
||||
remote_config?: Record<string, unknown> | null
|
||||
config_version: number
|
||||
}
|
||||
|
||||
// 用户导出数据结构
|
||||
export interface UsersExportData {
|
||||
version: string
|
||||
@@ -308,6 +325,7 @@ export interface ConfigImportResponse {
|
||||
message: string
|
||||
stats: {
|
||||
global_models: { created: number; updated: number; skipped: number }
|
||||
proxy_nodes?: { created: number; updated: number; skipped: number }
|
||||
providers: { created: number; updated: number; skipped: number }
|
||||
endpoints: { created: number; updated: number; skipped: number }
|
||||
keys: { created: number; updated: number; skipped: number }
|
||||
|
||||
@@ -23,6 +23,9 @@
|
||||
<li>
|
||||
API Keys: {{ importPreview.providers?.reduce((sum: number, p: { api_keys?: unknown[] }) => sum + (p.api_keys?.length || 0), 0) }} 个
|
||||
</li>
|
||||
<li v-if="importPreview.proxy_nodes?.length">
|
||||
代理节点: {{ importPreview.proxy_nodes.length }} 个
|
||||
</li>
|
||||
<li v-if="importPreview.ldap_config">
|
||||
LDAP 配置: 1 个
|
||||
</li>
|
||||
@@ -169,6 +172,16 @@
|
||||
跳过: {{ importResult.stats.oauth.skipped }}
|
||||
</p>
|
||||
</div>
|
||||
<div v-if="importResult.stats.proxy_nodes">
|
||||
<p class="font-medium">
|
||||
代理节点
|
||||
</p>
|
||||
<p class="text-muted-foreground">
|
||||
创建: {{ importResult.stats.proxy_nodes.created }},
|
||||
更新: {{ importResult.stats.proxy_nodes.updated }},
|
||||
跳过: {{ importResult.stats.proxy_nodes.skipped }}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
|
||||
@@ -973,7 +973,13 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.models.database import GlobalModel, Model, ProviderAPIKey, ProviderEndpoint
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
ProxyNode,
|
||||
)
|
||||
|
||||
db = context.db
|
||||
|
||||
@@ -1138,11 +1144,34 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
}
|
||||
)
|
||||
|
||||
# 导出 ProxyNode(手动节点 + 隧道节点,不含运行时状态)
|
||||
proxy_nodes = db.query(ProxyNode).all()
|
||||
proxy_nodes_data = []
|
||||
for node in proxy_nodes:
|
||||
proxy_nodes_data.append(
|
||||
{
|
||||
"id": node.id,
|
||||
"name": node.name,
|
||||
"ip": node.ip,
|
||||
"port": node.port,
|
||||
"region": node.region,
|
||||
"is_manual": node.is_manual,
|
||||
"proxy_url": node.proxy_url,
|
||||
"proxy_username": node.proxy_username,
|
||||
"proxy_password": node.proxy_password,
|
||||
"tunnel_mode": node.tunnel_mode,
|
||||
"heartbeat_interval": node.heartbeat_interval,
|
||||
"remote_config": node.remote_config,
|
||||
"config_version": node.config_version,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"version": CONFIG_EXPORT_VERSION,
|
||||
"exported_at": datetime.now(timezone.utc).isoformat(),
|
||||
"global_models": global_models_data,
|
||||
"providers": providers_data,
|
||||
"proxy_nodes": proxy_nodes_data,
|
||||
"ldap_config": ldap_data,
|
||||
"oauth_providers": oauth_data,
|
||||
"system_configs": system_configs_data,
|
||||
@@ -1155,6 +1184,32 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
# Provider Ops 中需要加密的敏感字段
|
||||
SENSITIVE_CREDENTIALS = SENSITIVE_CREDENTIAL_FIELDS
|
||||
|
||||
@staticmethod
|
||||
def _remap_proxy_node_id(
|
||||
proxy: dict[str, Any] | None,
|
||||
node_id_map: dict[str, str],
|
||||
) -> dict[str, Any] | None:
|
||||
"""替换 proxy 配置中的 node_id 为新实例的 ID。
|
||||
|
||||
- node_id 在映射表中:替换为新 ID
|
||||
- node_id 不在映射表中(节点未导入):清除 proxy 配置
|
||||
- 无 node_id(手动 URL 模式):原样返回
|
||||
"""
|
||||
if not proxy or not isinstance(proxy, dict):
|
||||
return proxy
|
||||
|
||||
old_node_id = proxy.get("node_id")
|
||||
if not old_node_id or not isinstance(old_node_id, str):
|
||||
return proxy # 手动 URL 模式,无需映射
|
||||
|
||||
new_node_id = node_id_map.get(old_node_id)
|
||||
if new_node_id is None:
|
||||
return None # 节点未导入,清除代理配置
|
||||
|
||||
remapped = dict(proxy)
|
||||
remapped["node_id"] = new_node_id
|
||||
return remapped
|
||||
|
||||
@staticmethod
|
||||
def _extract_import_key_api_formats(
|
||||
key_data: dict[str, Any], endpoint_formats: set[str]
|
||||
@@ -1207,7 +1262,13 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.models.database import GlobalModel, Model, ProviderAPIKey, ProviderEndpoint
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
ProxyNode,
|
||||
)
|
||||
|
||||
# 检查请求体大小
|
||||
if context.raw_body and len(context.raw_body) > MAX_IMPORT_SIZE:
|
||||
@@ -1227,12 +1288,14 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
merge_mode = payload.get("merge_mode", "skip") # skip, overwrite, error
|
||||
global_models_data = payload.get("global_models", [])
|
||||
providers_data = payload.get("providers", [])
|
||||
proxy_nodes_data = payload.get("proxy_nodes", [])
|
||||
ldap_data = payload.get("ldap_config") # 2.1 新增
|
||||
oauth_data = payload.get("oauth_providers", []) # 2.1 新增
|
||||
system_configs_data = payload.get("system_configs", []) # 2.2 新增
|
||||
|
||||
stats = {
|
||||
"global_models": {"created": 0, "updated": 0, "skipped": 0},
|
||||
"proxy_nodes": {"created": 0, "updated": 0, "skipped": 0},
|
||||
"providers": {"created": 0, "updated": 0, "skipped": 0},
|
||||
"endpoints": {"created": 0, "updated": 0, "skipped": 0},
|
||||
"keys": {"created": 0, "updated": 0, "skipped": 0},
|
||||
@@ -1298,6 +1361,77 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
global_model_map[gm_data["name"]] = new_gm.id
|
||||
stats["global_models"]["created"] += 1
|
||||
|
||||
# 导入 ProxyNodes(在 Providers 之前,建立 old_id -> new_id 映射)
|
||||
proxy_node_id_map: dict[str, str] = {} # old_node_id -> new_node_id
|
||||
for node_data in proxy_nodes_data:
|
||||
old_id = node_data.get("id", "")
|
||||
ip = node_data.get("ip", "")
|
||||
port = node_data.get("port", 0)
|
||||
|
||||
if not ip or not port:
|
||||
stats["errors"].append(f"跳过无效的代理节点: {node_data.get('name', '?')}")
|
||||
continue
|
||||
|
||||
existing_node = (
|
||||
db.query(ProxyNode).filter(ProxyNode.ip == ip, ProxyNode.port == port).first()
|
||||
)
|
||||
|
||||
if existing_node:
|
||||
proxy_node_id_map[old_id] = existing_node.id
|
||||
if merge_mode == "skip":
|
||||
stats["proxy_nodes"]["skipped"] += 1
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException(
|
||||
f"代理节点 '{node_data.get('name')}' ({ip}:{port}) 已存在"
|
||||
)
|
||||
elif merge_mode == "overwrite":
|
||||
existing_node.name = node_data.get("name", existing_node.name)
|
||||
existing_node.region = node_data.get("region", existing_node.region)
|
||||
existing_node.is_manual = node_data.get(
|
||||
"is_manual", existing_node.is_manual
|
||||
)
|
||||
existing_node.proxy_url = node_data.get(
|
||||
"proxy_url", existing_node.proxy_url
|
||||
)
|
||||
existing_node.proxy_username = node_data.get(
|
||||
"proxy_username", existing_node.proxy_username
|
||||
)
|
||||
existing_node.proxy_password = node_data.get(
|
||||
"proxy_password", existing_node.proxy_password
|
||||
)
|
||||
existing_node.tunnel_mode = node_data.get(
|
||||
"tunnel_mode", existing_node.tunnel_mode
|
||||
)
|
||||
existing_node.remote_config = node_data.get(
|
||||
"remote_config", existing_node.remote_config
|
||||
)
|
||||
existing_node.updated_at = datetime.now(timezone.utc)
|
||||
stats["proxy_nodes"]["updated"] += 1
|
||||
else:
|
||||
from src.models.database import ProxyNodeStatus
|
||||
|
||||
is_manual = node_data.get("is_manual", False)
|
||||
new_node = ProxyNode(
|
||||
id=str(uuid.uuid4()),
|
||||
name=node_data.get("name", "Imported Node"),
|
||||
ip=ip,
|
||||
port=port,
|
||||
region=node_data.get("region"),
|
||||
is_manual=is_manual,
|
||||
proxy_url=node_data.get("proxy_url"),
|
||||
proxy_username=node_data.get("proxy_username"),
|
||||
proxy_password=node_data.get("proxy_password"),
|
||||
tunnel_mode=node_data.get("tunnel_mode", False),
|
||||
heartbeat_interval=node_data.get("heartbeat_interval", 0),
|
||||
remote_config=node_data.get("remote_config"),
|
||||
config_version=node_data.get("config_version", 0),
|
||||
status=ProxyNodeStatus.ONLINE if is_manual else ProxyNodeStatus.OFFLINE,
|
||||
)
|
||||
db.add(new_node)
|
||||
db.flush()
|
||||
proxy_node_id_map[old_id] = new_node.id
|
||||
stats["proxy_nodes"]["created"] += 1
|
||||
|
||||
# 导入 Providers
|
||||
for prov_data in providers_data:
|
||||
existing_provider = (
|
||||
@@ -1348,7 +1482,12 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
existing_provider.request_timeout = prov_data.get(
|
||||
"request_timeout", existing_provider.request_timeout
|
||||
)
|
||||
existing_provider.proxy = prov_data.get("proxy", existing_provider.proxy)
|
||||
if "proxy" in prov_data:
|
||||
existing_provider.proxy = self._remap_proxy_node_id(
|
||||
prov_data["proxy"],
|
||||
proxy_node_id_map,
|
||||
)
|
||||
# 未提供 proxy 字段时保留现有配置
|
||||
# 加密 provider_ops credentials 后再保存
|
||||
existing_provider.config = self._encrypt_provider_config(
|
||||
prov_data.get("config"), crypto_service
|
||||
@@ -1385,7 +1524,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
max_retries=prov_data.get("max_retries"),
|
||||
stream_first_byte_timeout=prov_data.get("stream_first_byte_timeout"),
|
||||
request_timeout=prov_data.get("request_timeout"),
|
||||
proxy=prov_data.get("proxy"),
|
||||
proxy=self._remap_proxy_node_id(prov_data.get("proxy"), proxy_node_id_map),
|
||||
config=encrypted_config,
|
||||
)
|
||||
db.add(new_provider)
|
||||
@@ -1428,7 +1567,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
existing_ep.format_acceptance_config = ep_data.get(
|
||||
"format_acceptance_config"
|
||||
)
|
||||
existing_ep.proxy = ep_data.get("proxy")
|
||||
existing_ep.proxy = self._remap_proxy_node_id(
|
||||
ep_data.get("proxy"), proxy_node_id_map
|
||||
)
|
||||
sig = parse_signature_key(ep_format)
|
||||
existing_ep.api_format = sig.key # 使用归一化后的格式
|
||||
existing_ep.api_family = sig.api_family.value
|
||||
@@ -1453,7 +1594,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
custom_path=ep_data.get("custom_path"),
|
||||
config=ep_data.get("config"),
|
||||
format_acceptance_config=ep_data.get("format_acceptance_config"),
|
||||
proxy=ep_data.get("proxy"),
|
||||
proxy=self._remap_proxy_node_id(
|
||||
ep_data.get("proxy"), proxy_node_id_map
|
||||
),
|
||||
)
|
||||
db.add(new_ep)
|
||||
db.flush()
|
||||
@@ -1521,8 +1664,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
if endpoint_formats and fmt_normalized not in endpoint_formats:
|
||||
missing_formats.append(fmt_normalized.upper())
|
||||
continue
|
||||
# 存储时使用大写格式以保持向后兼容
|
||||
normalized_formats.append(fmt_normalized.upper())
|
||||
normalized_formats.append(fmt_normalized)
|
||||
|
||||
if missing_formats:
|
||||
stats["errors"].append(
|
||||
@@ -1572,7 +1714,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
model_include_patterns=key_data.get("model_include_patterns"),
|
||||
model_exclude_patterns=key_data.get("model_exclude_patterns"),
|
||||
is_active=key_data.get("is_active", True),
|
||||
proxy=key_data.get("proxy"),
|
||||
proxy=self._remap_proxy_node_id(key_data.get("proxy"), proxy_node_id_map),
|
||||
fingerprint=generate_fingerprint(seed=new_key_id),
|
||||
health_by_format={},
|
||||
circuit_breaker_by_format={},
|
||||
@@ -1928,7 +2070,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
@staticmethod
|
||||
def _serialize_api_key(
|
||||
key: ApiKey, include_is_standalone: bool = False, db: Any = None,
|
||||
key: ApiKey,
|
||||
include_is_standalone: bool = False,
|
||||
db: Any = None,
|
||||
) -> dict[str, Any]:
|
||||
"""序列化用户 API Key 为导出格式。"""
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -2142,7 +2286,8 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
)
|
||||
wallet_limit_mode = (
|
||||
str(wallet_payload.get("limit_mode"))
|
||||
if wallet_payload and wallet_payload.get("limit_mode") in {"finite", "unlimited"}
|
||||
if wallet_payload
|
||||
and wallet_payload.get("limit_mode") in {"finite", "unlimited"}
|
||||
else ("unlimited" if user_data.get("unlimited") else "finite")
|
||||
)
|
||||
|
||||
@@ -2173,7 +2318,9 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
if wallet_payload:
|
||||
wallet.balance = wallet_payload.get("recharge_balance", 0) or 0
|
||||
wallet.gift_balance = wallet_payload.get("gift_balance", 0) or 0
|
||||
wallet.total_recharged = wallet_payload.get("total_recharged", 0) or 0
|
||||
wallet.total_recharged = (
|
||||
wallet_payload.get("total_recharged", 0) or 0
|
||||
)
|
||||
wallet.total_consumed = wallet_payload.get("total_consumed", 0) or 0
|
||||
wallet.total_refunded = wallet_payload.get("total_refunded", 0) or 0
|
||||
wallet.total_adjusted = wallet_payload.get("total_adjusted", 0) or 0
|
||||
@@ -2250,8 +2397,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
wallet.limit_mode = (
|
||||
str(wallet_payload.get("limit_mode"))
|
||||
if wallet_payload
|
||||
and wallet_payload.get("limit_mode")
|
||||
in {"finite", "unlimited"}
|
||||
and wallet_payload.get("limit_mode") in {"finite", "unlimited"}
|
||||
else ("unlimited" if key_data.get("unlimited") else "finite")
|
||||
)
|
||||
if wallet_payload:
|
||||
@@ -2269,7 +2415,9 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
wallet.total_adjusted = (
|
||||
wallet_payload.get("total_adjusted", 0) or 0
|
||||
)
|
||||
wallet.status = wallet_payload.get("status", "active") or "active"
|
||||
wallet.status = (
|
||||
wallet_payload.get("status", "active") or "active"
|
||||
)
|
||||
wallet.updated_at = datetime.now(timezone.utc)
|
||||
stats["standalone_keys"]["created"] += 1
|
||||
elif status == "skipped":
|
||||
|
||||
@@ -362,12 +362,23 @@ class PoolManager:
|
||||
self,
|
||||
session_uuid: str | None,
|
||||
keys: list[ProviderAPIKey],
|
||||
*,
|
||||
availability_checker: (
|
||||
Callable[[ProviderAPIKey], tuple[bool, str | None, str | None]] | None
|
||||
) = None,
|
||||
page_size: int = 50,
|
||||
) -> tuple[list[ProviderAPIKey], PoolSchedulingTrace]:
|
||||
"""Select and order pool keys with trace output.
|
||||
|
||||
Reuses :meth:`reorder_candidates` logic by adapting keys to lightweight
|
||||
candidate-like wrappers, then propagates skip/trace metadata back onto
|
||||
each key object for downstream execution/recording.
|
||||
|
||||
When *availability_checker* is provided, post-sort availability checks
|
||||
are performed lazily: only the top *page_size* non-skipped keys are
|
||||
checked at a time; if all fail, the next page is checked, and so on.
|
||||
Keys beyond the last checked page are marked as ``deferred`` (skipped
|
||||
without checking) to avoid unnecessary CPU work on large pools.
|
||||
"""
|
||||
if not keys:
|
||||
return (
|
||||
@@ -424,6 +435,30 @@ class PoolManager:
|
||||
)
|
||||
ordered_keys.append(key)
|
||||
|
||||
# -- 分页可用性检查 --
|
||||
# 排序后对非 skipped key 分页调用 availability_checker,
|
||||
# 找到 page_size 个可用 key 后停止检查,剩余标记 deferred。
|
||||
if availability_checker is not None:
|
||||
available_count = 0
|
||||
found_enough = False
|
||||
for key in ordered_keys:
|
||||
if getattr(key, "_pool_skipped", False):
|
||||
continue
|
||||
if found_enough:
|
||||
setattr(key, "_pool_skipped", True)
|
||||
setattr(key, "_pool_skip_reason", "deferred")
|
||||
continue
|
||||
is_available, skip_reason_check, mapping_model = availability_checker(key)
|
||||
if not is_available:
|
||||
setattr(key, "_pool_skipped", True)
|
||||
setattr(key, "_pool_skip_reason", skip_reason_check)
|
||||
else:
|
||||
if mapping_model:
|
||||
setattr(key, "_pool_mapping_matched_model", mapping_model)
|
||||
available_count += 1
|
||||
if available_count >= page_size:
|
||||
found_enough = True
|
||||
|
||||
return ordered_keys, trace
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -563,33 +563,9 @@ class CandidateBuilder:
|
||||
)
|
||||
|
||||
if pool_cfg is not None:
|
||||
# 号池 Provider 仅构建一个 PoolCandidate,内部 key 选择延迟到执行阶段。
|
||||
pool_keys: list[ProviderAPIKey] = []
|
||||
pool_miss_counts: list[int] = []
|
||||
pool_mapping: dict[str, str | None] = {}
|
||||
for key in keys_to_check:
|
||||
is_available, _key_skip_reason, mapping_matched_model = (
|
||||
self._check_key_availability(
|
||||
key,
|
||||
endpoint_format_str,
|
||||
model_name,
|
||||
capability_requirements,
|
||||
model_mappings=model_mappings,
|
||||
candidate_models=provider_model_names,
|
||||
provider_type=getattr(provider, "provider_type", None),
|
||||
)
|
||||
)
|
||||
if not is_available:
|
||||
continue
|
||||
|
||||
pool_keys.append(key)
|
||||
pool_miss_counts.append(
|
||||
compute_capability_score(
|
||||
key.capabilities or {},
|
||||
capability_requirements,
|
||||
)
|
||||
)
|
||||
pool_mapping[str(key.id)] = mapping_matched_model
|
||||
# 号池优化:跳过逐 key 的 _check_key_availability 检查,
|
||||
# 直接收集全部 active key,将检查推迟到 PoolManager 排序后分页执行。
|
||||
pool_keys = list(keys_to_check)
|
||||
|
||||
if not pool_keys:
|
||||
continue
|
||||
@@ -619,20 +595,21 @@ class CandidateBuilder:
|
||||
pool_keys=pool_keys,
|
||||
pool_config=pool_cfg,
|
||||
pool_priority=pool_priority,
|
||||
mapping_matched_model=pool_mapping.get(str(pool_keys[0].id)),
|
||||
needs_conversion=needs_conversion,
|
||||
provider_api_format=str(endpoint_format_str or ""),
|
||||
output_limit=output_limit,
|
||||
capability_miss_count=min(pool_miss_counts) if pool_miss_counts else 0,
|
||||
capability_miss_count=0,
|
||||
)
|
||||
|
||||
# 在 key 对象上附加映射结果,供 PoolCandidate 运行时切 key 后同步模型名。
|
||||
for pool_key in pool_keys:
|
||||
setattr(
|
||||
pool_key,
|
||||
"_pool_mapping_matched_model",
|
||||
pool_mapping.get(str(pool_key.id)),
|
||||
)
|
||||
# 打包延迟检查参数,供 PoolManager 排序后分页调用
|
||||
pool_candidate._deferred_check_params = {
|
||||
"endpoint_format": endpoint_format_str,
|
||||
"model_name": model_name,
|
||||
"capability_requirements": capability_requirements,
|
||||
"model_mappings": model_mappings,
|
||||
"candidate_models": provider_model_names,
|
||||
"provider_type": getattr(provider, "provider_type", None),
|
||||
}
|
||||
|
||||
if needs_conversion:
|
||||
convertible_candidates.append(pool_candidate)
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.models.database import (
|
||||
Provider,
|
||||
@@ -84,6 +84,8 @@ class PoolCandidate(ProviderCandidate):
|
||||
pool_config: PoolConfig | None = None
|
||||
pool_priority: int = 999999
|
||||
_pool_key_index: int = 0
|
||||
# 延迟可用性检查参数(号池优化:先排序再分页检查)
|
||||
_deferred_check_params: dict[str, Any] | None = field(default=None, repr=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -63,12 +63,25 @@ class TaskPoolOperationsService:
|
||||
if not candidate_keys and getattr(candidate, "key", None) is not None:
|
||||
candidate_keys = [candidate.key]
|
||||
|
||||
ordered_keys, trace = await manager.select_pool_keys(session_uuid, candidate_keys)
|
||||
candidate.pool_keys = ordered_keys
|
||||
# 构造延迟可用性检查回调(从 CandidateBuilder 打包的参数)
|
||||
checker = None
|
||||
deferred_params = candidate._deferred_check_params
|
||||
if deferred_params is not None:
|
||||
checker = self._build_availability_checker(deferred_params)
|
||||
|
||||
ordered_keys, trace = await manager.select_pool_keys(
|
||||
session_uuid,
|
||||
candidate_keys,
|
||||
availability_checker=checker,
|
||||
)
|
||||
# 移除 deferred key(未检查的),避免为其创建 DB 记录
|
||||
candidate.pool_keys = [
|
||||
k for k in ordered_keys if getattr(k, "_pool_skip_reason", None) != "deferred"
|
||||
]
|
||||
|
||||
selected_key_index = 0
|
||||
selected_key = None
|
||||
for idx, pool_key in enumerate(ordered_keys):
|
||||
for idx, pool_key in enumerate(candidate.pool_keys):
|
||||
if not bool(getattr(pool_key, "_pool_skipped", False)):
|
||||
selected_key = pool_key
|
||||
selected_key_index = idx
|
||||
@@ -220,3 +233,33 @@ class TaskPoolOperationsService:
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _build_availability_checker(
|
||||
params: dict[str, Any],
|
||||
) -> Any:
|
||||
"""Construct a key availability checker from deferred check params."""
|
||||
from src.services.scheduling.candidate_builder import CandidateBuilder
|
||||
|
||||
endpoint_format = params.get("endpoint_format")
|
||||
model_name = params.get("model_name", "")
|
||||
capability_requirements = params.get("capability_requirements")
|
||||
model_mappings = params.get("model_mappings")
|
||||
candidate_models = params.get("candidate_models")
|
||||
provider_type = params.get("provider_type")
|
||||
|
||||
# _check_key_availability 不依赖 _sorter,传 None 安全
|
||||
builder = CandidateBuilder(candidate_sorter=None) # type: ignore[arg-type]
|
||||
|
||||
def _checker(key: Any) -> tuple[bool, str | None, str | None]:
|
||||
return builder._check_key_availability(
|
||||
key,
|
||||
endpoint_format,
|
||||
model_name,
|
||||
capability_requirements,
|
||||
model_mappings=model_mappings,
|
||||
candidate_models=candidate_models,
|
||||
provider_type=provider_type,
|
||||
)
|
||||
|
||||
return _checker
|
||||
|
||||
Reference in New Issue
Block a user