mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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
|
exported_at: string
|
||||||
global_models: GlobalModelExport[]
|
global_models: GlobalModelExport[]
|
||||||
providers: ProviderExport[]
|
providers: ProviderExport[]
|
||||||
|
proxy_nodes?: ProxyNodeExport[]
|
||||||
ldap_config?: LDAPConfigExport | null
|
ldap_config?: LDAPConfigExport | null
|
||||||
oauth_providers?: OAuthProviderExport[]
|
oauth_providers?: OAuthProviderExport[]
|
||||||
system_configs?: SystemConfigExport[]
|
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 {
|
export interface UsersExportData {
|
||||||
version: string
|
version: string
|
||||||
@@ -308,6 +325,7 @@ export interface ConfigImportResponse {
|
|||||||
message: string
|
message: string
|
||||||
stats: {
|
stats: {
|
||||||
global_models: { created: number; updated: number; skipped: number }
|
global_models: { created: number; updated: number; skipped: number }
|
||||||
|
proxy_nodes?: { created: number; updated: number; skipped: number }
|
||||||
providers: { created: number; updated: number; skipped: number }
|
providers: { created: number; updated: number; skipped: number }
|
||||||
endpoints: { created: number; updated: number; skipped: number }
|
endpoints: { created: number; updated: number; skipped: number }
|
||||||
keys: { created: number; updated: number; skipped: number }
|
keys: { created: number; updated: number; skipped: number }
|
||||||
|
|||||||
@@ -23,6 +23,9 @@
|
|||||||
<li>
|
<li>
|
||||||
API Keys: {{ importPreview.providers?.reduce((sum: number, p: { api_keys?: unknown[] }) => sum + (p.api_keys?.length || 0), 0) }} 个
|
API Keys: {{ importPreview.providers?.reduce((sum: number, p: { api_keys?: unknown[] }) => sum + (p.api_keys?.length || 0), 0) }} 个
|
||||||
</li>
|
</li>
|
||||||
|
<li v-if="importPreview.proxy_nodes?.length">
|
||||||
|
代理节点: {{ importPreview.proxy_nodes.length }} 个
|
||||||
|
</li>
|
||||||
<li v-if="importPreview.ldap_config">
|
<li v-if="importPreview.ldap_config">
|
||||||
LDAP 配置: 1 个
|
LDAP 配置: 1 个
|
||||||
</li>
|
</li>
|
||||||
@@ -169,6 +172,16 @@
|
|||||||
跳过: {{ importResult.stats.oauth.skipped }}
|
跳过: {{ importResult.stats.oauth.skipped }}
|
||||||
</p>
|
</p>
|
||||||
</div>
|
</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>
|
||||||
|
|
||||||
<div
|
<div
|
||||||
|
|||||||
@@ -973,7 +973,13 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
|||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from src.core.crypto import crypto_service
|
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
|
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 {
|
return {
|
||||||
"version": CONFIG_EXPORT_VERSION,
|
"version": CONFIG_EXPORT_VERSION,
|
||||||
"exported_at": datetime.now(timezone.utc).isoformat(),
|
"exported_at": datetime.now(timezone.utc).isoformat(),
|
||||||
"global_models": global_models_data,
|
"global_models": global_models_data,
|
||||||
"providers": providers_data,
|
"providers": providers_data,
|
||||||
|
"proxy_nodes": proxy_nodes_data,
|
||||||
"ldap_config": ldap_data,
|
"ldap_config": ldap_data,
|
||||||
"oauth_providers": oauth_data,
|
"oauth_providers": oauth_data,
|
||||||
"system_configs": system_configs_data,
|
"system_configs": system_configs_data,
|
||||||
@@ -1155,6 +1184,32 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
# Provider Ops 中需要加密的敏感字段
|
# Provider Ops 中需要加密的敏感字段
|
||||||
SENSITIVE_CREDENTIALS = SENSITIVE_CREDENTIAL_FIELDS
|
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
|
@staticmethod
|
||||||
def _extract_import_key_api_formats(
|
def _extract_import_key_api_formats(
|
||||||
key_data: dict[str, Any], endpoint_formats: set[str]
|
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.crypto import crypto_service
|
||||||
from src.core.enums import ProviderBillingType
|
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:
|
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
|
merge_mode = payload.get("merge_mode", "skip") # skip, overwrite, error
|
||||||
global_models_data = payload.get("global_models", [])
|
global_models_data = payload.get("global_models", [])
|
||||||
providers_data = payload.get("providers", [])
|
providers_data = payload.get("providers", [])
|
||||||
|
proxy_nodes_data = payload.get("proxy_nodes", [])
|
||||||
ldap_data = payload.get("ldap_config") # 2.1 新增
|
ldap_data = payload.get("ldap_config") # 2.1 新增
|
||||||
oauth_data = payload.get("oauth_providers", []) # 2.1 新增
|
oauth_data = payload.get("oauth_providers", []) # 2.1 新增
|
||||||
system_configs_data = payload.get("system_configs", []) # 2.2 新增
|
system_configs_data = payload.get("system_configs", []) # 2.2 新增
|
||||||
|
|
||||||
stats = {
|
stats = {
|
||||||
"global_models": {"created": 0, "updated": 0, "skipped": 0},
|
"global_models": {"created": 0, "updated": 0, "skipped": 0},
|
||||||
|
"proxy_nodes": {"created": 0, "updated": 0, "skipped": 0},
|
||||||
"providers": {"created": 0, "updated": 0, "skipped": 0},
|
"providers": {"created": 0, "updated": 0, "skipped": 0},
|
||||||
"endpoints": {"created": 0, "updated": 0, "skipped": 0},
|
"endpoints": {"created": 0, "updated": 0, "skipped": 0},
|
||||||
"keys": {"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
|
global_model_map[gm_data["name"]] = new_gm.id
|
||||||
stats["global_models"]["created"] += 1
|
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
|
# 导入 Providers
|
||||||
for prov_data in providers_data:
|
for prov_data in providers_data:
|
||||||
existing_provider = (
|
existing_provider = (
|
||||||
@@ -1348,7 +1482,12 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
existing_provider.request_timeout = prov_data.get(
|
existing_provider.request_timeout = prov_data.get(
|
||||||
"request_timeout", existing_provider.request_timeout
|
"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 后再保存
|
# 加密 provider_ops credentials 后再保存
|
||||||
existing_provider.config = self._encrypt_provider_config(
|
existing_provider.config = self._encrypt_provider_config(
|
||||||
prov_data.get("config"), crypto_service
|
prov_data.get("config"), crypto_service
|
||||||
@@ -1385,7 +1524,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
max_retries=prov_data.get("max_retries"),
|
max_retries=prov_data.get("max_retries"),
|
||||||
stream_first_byte_timeout=prov_data.get("stream_first_byte_timeout"),
|
stream_first_byte_timeout=prov_data.get("stream_first_byte_timeout"),
|
||||||
request_timeout=prov_data.get("request_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,
|
config=encrypted_config,
|
||||||
)
|
)
|
||||||
db.add(new_provider)
|
db.add(new_provider)
|
||||||
@@ -1428,7 +1567,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
existing_ep.format_acceptance_config = ep_data.get(
|
existing_ep.format_acceptance_config = ep_data.get(
|
||||||
"format_acceptance_config"
|
"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)
|
sig = parse_signature_key(ep_format)
|
||||||
existing_ep.api_format = sig.key # 使用归一化后的格式
|
existing_ep.api_format = sig.key # 使用归一化后的格式
|
||||||
existing_ep.api_family = sig.api_family.value
|
existing_ep.api_family = sig.api_family.value
|
||||||
@@ -1453,7 +1594,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
custom_path=ep_data.get("custom_path"),
|
custom_path=ep_data.get("custom_path"),
|
||||||
config=ep_data.get("config"),
|
config=ep_data.get("config"),
|
||||||
format_acceptance_config=ep_data.get("format_acceptance_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.add(new_ep)
|
||||||
db.flush()
|
db.flush()
|
||||||
@@ -1521,8 +1664,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
if endpoint_formats and fmt_normalized not in endpoint_formats:
|
if endpoint_formats and fmt_normalized not in endpoint_formats:
|
||||||
missing_formats.append(fmt_normalized.upper())
|
missing_formats.append(fmt_normalized.upper())
|
||||||
continue
|
continue
|
||||||
# 存储时使用大写格式以保持向后兼容
|
normalized_formats.append(fmt_normalized)
|
||||||
normalized_formats.append(fmt_normalized.upper())
|
|
||||||
|
|
||||||
if missing_formats:
|
if missing_formats:
|
||||||
stats["errors"].append(
|
stats["errors"].append(
|
||||||
@@ -1572,7 +1714,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
model_include_patterns=key_data.get("model_include_patterns"),
|
model_include_patterns=key_data.get("model_include_patterns"),
|
||||||
model_exclude_patterns=key_data.get("model_exclude_patterns"),
|
model_exclude_patterns=key_data.get("model_exclude_patterns"),
|
||||||
is_active=key_data.get("is_active", True),
|
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),
|
fingerprint=generate_fingerprint(seed=new_key_id),
|
||||||
health_by_format={},
|
health_by_format={},
|
||||||
circuit_breaker_by_format={},
|
circuit_breaker_by_format={},
|
||||||
@@ -1928,7 +2070,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
|||||||
class AdminExportUsersAdapter(AdminApiAdapter):
|
class AdminExportUsersAdapter(AdminApiAdapter):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _serialize_api_key(
|
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]:
|
) -> dict[str, Any]:
|
||||||
"""序列化用户 API Key 为导出格式。"""
|
"""序列化用户 API Key 为导出格式。"""
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
@@ -2142,7 +2286,8 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
wallet_limit_mode = (
|
wallet_limit_mode = (
|
||||||
str(wallet_payload.get("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")
|
else ("unlimited" if user_data.get("unlimited") else "finite")
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2173,7 +2318,9 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
|||||||
if wallet_payload:
|
if wallet_payload:
|
||||||
wallet.balance = wallet_payload.get("recharge_balance", 0) or 0
|
wallet.balance = wallet_payload.get("recharge_balance", 0) or 0
|
||||||
wallet.gift_balance = wallet_payload.get("gift_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_consumed = wallet_payload.get("total_consumed", 0) or 0
|
||||||
wallet.total_refunded = wallet_payload.get("total_refunded", 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
|
wallet.total_adjusted = wallet_payload.get("total_adjusted", 0) or 0
|
||||||
@@ -2250,8 +2397,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
|||||||
wallet.limit_mode = (
|
wallet.limit_mode = (
|
||||||
str(wallet_payload.get("limit_mode"))
|
str(wallet_payload.get("limit_mode"))
|
||||||
if wallet_payload
|
if wallet_payload
|
||||||
and wallet_payload.get("limit_mode")
|
and wallet_payload.get("limit_mode") in {"finite", "unlimited"}
|
||||||
in {"finite", "unlimited"}
|
|
||||||
else ("unlimited" if key_data.get("unlimited") else "finite")
|
else ("unlimited" if key_data.get("unlimited") else "finite")
|
||||||
)
|
)
|
||||||
if wallet_payload:
|
if wallet_payload:
|
||||||
@@ -2269,7 +2415,9 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
|||||||
wallet.total_adjusted = (
|
wallet.total_adjusted = (
|
||||||
wallet_payload.get("total_adjusted", 0) or 0
|
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)
|
wallet.updated_at = datetime.now(timezone.utc)
|
||||||
stats["standalone_keys"]["created"] += 1
|
stats["standalone_keys"]["created"] += 1
|
||||||
elif status == "skipped":
|
elif status == "skipped":
|
||||||
|
|||||||
@@ -362,12 +362,23 @@ class PoolManager:
|
|||||||
self,
|
self,
|
||||||
session_uuid: str | None,
|
session_uuid: str | None,
|
||||||
keys: list[ProviderAPIKey],
|
keys: list[ProviderAPIKey],
|
||||||
|
*,
|
||||||
|
availability_checker: (
|
||||||
|
Callable[[ProviderAPIKey], tuple[bool, str | None, str | None]] | None
|
||||||
|
) = None,
|
||||||
|
page_size: int = 50,
|
||||||
) -> tuple[list[ProviderAPIKey], PoolSchedulingTrace]:
|
) -> tuple[list[ProviderAPIKey], PoolSchedulingTrace]:
|
||||||
"""Select and order pool keys with trace output.
|
"""Select and order pool keys with trace output.
|
||||||
|
|
||||||
Reuses :meth:`reorder_candidates` logic by adapting keys to lightweight
|
Reuses :meth:`reorder_candidates` logic by adapting keys to lightweight
|
||||||
candidate-like wrappers, then propagates skip/trace metadata back onto
|
candidate-like wrappers, then propagates skip/trace metadata back onto
|
||||||
each key object for downstream execution/recording.
|
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:
|
if not keys:
|
||||||
return (
|
return (
|
||||||
@@ -424,6 +435,30 @@ class PoolManager:
|
|||||||
)
|
)
|
||||||
ordered_keys.append(key)
|
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
|
return ordered_keys, trace
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|||||||
@@ -563,33 +563,9 @@ class CandidateBuilder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if pool_cfg is not None:
|
if pool_cfg is not None:
|
||||||
# 号池 Provider 仅构建一个 PoolCandidate,内部 key 选择延迟到执行阶段。
|
# 号池优化:跳过逐 key 的 _check_key_availability 检查,
|
||||||
pool_keys: list[ProviderAPIKey] = []
|
# 直接收集全部 active key,将检查推迟到 PoolManager 排序后分页执行。
|
||||||
pool_miss_counts: list[int] = []
|
pool_keys = list(keys_to_check)
|
||||||
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
|
|
||||||
|
|
||||||
if not pool_keys:
|
if not pool_keys:
|
||||||
continue
|
continue
|
||||||
@@ -619,20 +595,21 @@ class CandidateBuilder:
|
|||||||
pool_keys=pool_keys,
|
pool_keys=pool_keys,
|
||||||
pool_config=pool_cfg,
|
pool_config=pool_cfg,
|
||||||
pool_priority=pool_priority,
|
pool_priority=pool_priority,
|
||||||
mapping_matched_model=pool_mapping.get(str(pool_keys[0].id)),
|
|
||||||
needs_conversion=needs_conversion,
|
needs_conversion=needs_conversion,
|
||||||
provider_api_format=str(endpoint_format_str or ""),
|
provider_api_format=str(endpoint_format_str or ""),
|
||||||
output_limit=output_limit,
|
output_limit=output_limit,
|
||||||
capability_miss_count=min(pool_miss_counts) if pool_miss_counts else 0,
|
capability_miss_count=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 在 key 对象上附加映射结果,供 PoolCandidate 运行时切 key 后同步模型名。
|
# 打包延迟检查参数,供 PoolManager 排序后分页调用
|
||||||
for pool_key in pool_keys:
|
pool_candidate._deferred_check_params = {
|
||||||
setattr(
|
"endpoint_format": endpoint_format_str,
|
||||||
pool_key,
|
"model_name": model_name,
|
||||||
"_pool_mapping_matched_model",
|
"capability_requirements": capability_requirements,
|
||||||
pool_mapping.get(str(pool_key.id)),
|
"model_mappings": model_mappings,
|
||||||
)
|
"candidate_models": provider_model_names,
|
||||||
|
"provider_type": getattr(provider, "provider_type", None),
|
||||||
|
}
|
||||||
|
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
convertible_candidates.append(pool_candidate)
|
convertible_candidates.append(pool_candidate)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
Provider,
|
Provider,
|
||||||
@@ -84,6 +84,8 @@ class PoolCandidate(ProviderCandidate):
|
|||||||
pool_config: PoolConfig | None = None
|
pool_config: PoolConfig | None = None
|
||||||
pool_priority: int = 999999
|
pool_priority: int = 999999
|
||||||
_pool_key_index: int = 0
|
_pool_key_index: int = 0
|
||||||
|
# 延迟可用性检查参数(号池优化:先排序再分页检查)
|
||||||
|
_deferred_check_params: dict[str, Any] | None = field(default=None, repr=False)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -63,12 +63,25 @@ class TaskPoolOperationsService:
|
|||||||
if not candidate_keys and getattr(candidate, "key", None) is not None:
|
if not candidate_keys and getattr(candidate, "key", None) is not None:
|
||||||
candidate_keys = [candidate.key]
|
candidate_keys = [candidate.key]
|
||||||
|
|
||||||
ordered_keys, trace = await manager.select_pool_keys(session_uuid, candidate_keys)
|
# 构造延迟可用性检查回调(从 CandidateBuilder 打包的参数)
|
||||||
candidate.pool_keys = ordered_keys
|
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_index = 0
|
||||||
selected_key = None
|
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)):
|
if not bool(getattr(pool_key, "_pool_skipped", False)):
|
||||||
selected_key = pool_key
|
selected_key = pool_key
|
||||||
selected_key_index = idx
|
selected_key_index = idx
|
||||||
@@ -220,3 +233,33 @@ class TaskPoolOperationsService:
|
|||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
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