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:
fawney19
2026-03-09 13:26:54 +08:00
parent 0046123e22
commit d84c9d4b71
7 changed files with 291 additions and 55 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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