345 lines
12 KiB
Python
345 lines
12 KiB
Python
"""云厂商同步服务
|
||
|
||
将适配器返回的标准化资产(NormalizedVPS / NormalizedDomain / AccountInfo)
|
||
写入或更新到资产库。以 (provider_id, external_id) 作为去重键,
|
||
已存在则更新状态/详情,不存在则新建资产。
|
||
"""
|
||
|
||
import json
|
||
import logging
|
||
from typing import Optional
|
||
|
||
from fastapi import HTTPException, status
|
||
from sqlmodel import Session, select
|
||
|
||
logger = logging.getLogger("vps-manager.sync")
|
||
|
||
from app.adapters import registry
|
||
from app.adapters.base import BaseAdapter
|
||
from app.core import crypto
|
||
from app.core.timeutils import utcnow
|
||
from app.models.asset import (
|
||
AIAccount,
|
||
Asset,
|
||
AssetStatus,
|
||
AssetType,
|
||
DomainDetail,
|
||
VPSDetail,
|
||
)
|
||
from app.models.provider import Provider
|
||
|
||
_VALID_STATUS = {s.value for s in AssetStatus}
|
||
|
||
|
||
def _load_config(provider: Provider) -> dict:
|
||
"""解密平台的 API 配置 JSON"""
|
||
plain = crypto.decrypt(provider.api_config_encrypted)
|
||
if not plain:
|
||
return {}
|
||
try:
|
||
return json.loads(plain)
|
||
except (json.JSONDecodeError, TypeError):
|
||
return {}
|
||
|
||
|
||
def _get_provider(session: Session, provider_id: int) -> Provider:
|
||
provider = session.get(Provider, provider_id)
|
||
if not provider:
|
||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="平台不存在")
|
||
return provider
|
||
|
||
|
||
def _build_adapter(provider: Provider) -> BaseAdapter:
|
||
if not provider.sdk_type:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST, detail="该平台未配置 SDK 类型(sdk_type)"
|
||
)
|
||
if not registry.is_supported(provider.sdk_type):
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST,
|
||
detail=f"暂不支持的 SDK 类型:{provider.sdk_type}",
|
||
)
|
||
return registry.get_adapter(provider.sdk_type, _load_config(provider))
|
||
|
||
|
||
def _norm_status(raw: Optional[str]) -> AssetStatus:
|
||
return AssetStatus(raw) if raw in _VALID_STATUS else AssetStatus.UNKNOWN
|
||
|
||
|
||
def _missing_config(adapter: BaseAdapter) -> list:
|
||
"""检查适配器所需凭证是否已配置"""
|
||
return [f for f in adapter.required_config if not adapter.config.get(f)]
|
||
|
||
|
||
def test_provider(session: Session, provider_id: int) -> dict:
|
||
"""测试平台连接 / 凭证有效性"""
|
||
provider = _get_provider(session, provider_id)
|
||
adapter = _build_adapter(provider)
|
||
base = {"capabilities": adapter.capabilities(), "sdk_type": provider.sdk_type}
|
||
missing = _missing_config(adapter)
|
||
if missing:
|
||
return {
|
||
"ok": False,
|
||
"message": f"缺少凭证配置:{', '.join(missing)}(请在平台编辑里填写 API 配置)",
|
||
**base,
|
||
}
|
||
result = adapter.test_connection()
|
||
result.update(base)
|
||
return result
|
||
|
||
|
||
def _find_asset(session: Session, provider_id: int, external_id: str, asset_type: AssetType):
|
||
return session.exec(
|
||
select(Asset).where(
|
||
Asset.provider_id == provider_id,
|
||
Asset.external_id == external_id,
|
||
Asset.asset_type == asset_type,
|
||
)
|
||
).first()
|
||
|
||
|
||
def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter) -> dict:
|
||
created = updated = 0
|
||
for vps in adapter.list_vps():
|
||
existing = _find_asset(session, provider.id, vps.external_id, AssetType.VPS)
|
||
if existing:
|
||
existing.name = vps.name or existing.name
|
||
existing.status = _norm_status(vps.status)
|
||
if vps.monthly_cost is not None:
|
||
existing.cost = vps.monthly_cost
|
||
existing.currency = vps.currency
|
||
session.add(existing)
|
||
detail = session.exec(
|
||
select(VPSDetail).where(VPSDetail.asset_id == existing.id)
|
||
).first()
|
||
if detail:
|
||
detail.ip_address = vps.ip_address or detail.ip_address
|
||
detail.region = vps.region or detail.region
|
||
detail.os = vps.os or detail.os
|
||
detail.cpu_cores = vps.cpu_cores or detail.cpu_cores
|
||
detail.memory_gb = vps.memory_gb or detail.memory_gb
|
||
detail.disk_gb = vps.disk_gb or detail.disk_gb
|
||
session.add(detail)
|
||
updated += 1
|
||
else:
|
||
asset = Asset(
|
||
name=vps.name or vps.external_id,
|
||
asset_type=AssetType.VPS,
|
||
provider=provider.slug,
|
||
provider_id=provider.id,
|
||
external_id=vps.external_id,
|
||
status=_norm_status(vps.status),
|
||
cost=vps.monthly_cost or 0,
|
||
currency=vps.currency,
|
||
)
|
||
session.add(asset)
|
||
session.flush() # 获取 asset.id,统一在循环外提交
|
||
session.add(
|
||
VPSDetail(
|
||
asset_id=asset.id,
|
||
ip_address=vps.ip_address or "",
|
||
region=vps.region,
|
||
os=vps.os,
|
||
cpu_cores=vps.cpu_cores or 1,
|
||
memory_gb=vps.memory_gb or 1,
|
||
disk_gb=vps.disk_gb or 20,
|
||
)
|
||
)
|
||
created += 1
|
||
session.commit()
|
||
return {"created": created, "updated": updated}
|
||
|
||
|
||
def _sync_domains(session: Session, provider: Provider, adapter: BaseAdapter) -> dict:
|
||
created = updated = 0
|
||
for dom in adapter.list_domains():
|
||
existing = _find_asset(session, provider.id, dom.external_id, AssetType.DOMAIN)
|
||
if existing:
|
||
existing.status = _norm_status(dom.status)
|
||
if dom.expiry_date:
|
||
existing.expiry_date = dom.expiry_date
|
||
session.add(existing)
|
||
detail = session.exec(
|
||
select(DomainDetail).where(DomainDetail.asset_id == existing.id)
|
||
).first()
|
||
if detail:
|
||
detail.domain_name = dom.domain_name or detail.domain_name
|
||
detail.registrar = dom.registrar or detail.registrar
|
||
session.add(detail)
|
||
updated += 1
|
||
else:
|
||
asset = Asset(
|
||
name=dom.domain_name,
|
||
asset_type=AssetType.DOMAIN,
|
||
provider=provider.slug,
|
||
provider_id=provider.id,
|
||
external_id=dom.external_id,
|
||
status=_norm_status(dom.status),
|
||
expiry_date=dom.expiry_date,
|
||
)
|
||
session.add(asset)
|
||
session.flush() # 获取 asset.id,统一在循环外提交
|
||
session.add(
|
||
DomainDetail(
|
||
asset_id=asset.id,
|
||
domain_name=dom.domain_name,
|
||
registrar=dom.registrar or provider.slug,
|
||
)
|
||
)
|
||
created += 1
|
||
session.commit()
|
||
return {"created": created, "updated": updated}
|
||
|
||
|
||
def sync_provider(session: Session, provider_id: int) -> dict:
|
||
"""同步平台资产到本地库"""
|
||
provider = _get_provider(session, provider_id)
|
||
adapter = _build_adapter(provider)
|
||
missing = _missing_config(adapter)
|
||
if missing:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST,
|
||
detail=f"缺少凭证配置:{', '.join(missing)}(请先在平台编辑里填写 API 配置)",
|
||
)
|
||
caps = adapter.capabilities()
|
||
result = {"provider": provider.slug, "sdk_type": provider.sdk_type}
|
||
|
||
if caps["list_vps"]:
|
||
try:
|
||
result["vps"] = _sync_vps(session, provider, adapter)
|
||
except Exception as e: # noqa: BLE001
|
||
result["vps_error"] = str(e)
|
||
if caps["list_domains"]:
|
||
try:
|
||
result["domains"] = _sync_domains(session, provider, adapter)
|
||
except Exception as e: # noqa: BLE001
|
||
result["domains_error"] = str(e)
|
||
if caps["get_account"]:
|
||
try:
|
||
result["account"] = adapter.get_account().to_dict()
|
||
except Exception as e: # noqa: BLE001
|
||
result["account_error"] = str(e)
|
||
|
||
provider.last_synced_at = utcnow()
|
||
session.add(provider)
|
||
session.commit()
|
||
result["last_synced_at"] = provider.last_synced_at.isoformat()
|
||
logger.info("同步平台 provider=%s result=%s", provider.slug, result)
|
||
return result
|
||
|
||
|
||
# AI 服务商 slug -> sdk_type 映射(Provider 未配 sdk_type 时的回退推断)
|
||
_AI_SDK_MAP = {
|
||
"openai": "openai-api",
|
||
"deepseek": "deepseek-api",
|
||
"kimi": "moonshot-api",
|
||
"moonshot": "moonshot-api",
|
||
"minimax": "minimax-api",
|
||
}
|
||
|
||
|
||
def refresh_ai_balance(session: Session, asset_id: int) -> dict:
|
||
"""刷新 AI 账号余额:解密 api_key → 适配器 get_account → 更新余额"""
|
||
asset = session.get(Asset, asset_id)
|
||
if not asset or asset.asset_type != AssetType.AI_AGENT:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST, detail="非 AI 账号资产"
|
||
)
|
||
ai = session.exec(select(AIAccount).where(AIAccount.asset_id == asset_id)).first()
|
||
if not ai:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_404_NOT_FOUND, detail="AI 账号详情不存在"
|
||
)
|
||
|
||
sdk_type = None
|
||
if asset.provider_id:
|
||
provider = session.get(Provider, asset.provider_id)
|
||
if provider:
|
||
sdk_type = provider.sdk_type
|
||
if not sdk_type:
|
||
sdk_type = _AI_SDK_MAP.get((ai.provider or "").lower())
|
||
if not sdk_type or not registry.is_supported(sdk_type):
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST,
|
||
detail=f"无法确定 AI 适配器({ai.provider})",
|
||
)
|
||
|
||
api_key = crypto.decrypt(ai.api_key_encrypted)
|
||
if not api_key:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST, detail="未配置 API Key"
|
||
)
|
||
|
||
adapter = registry.get_adapter(sdk_type, {"api_key": api_key})
|
||
try:
|
||
acc = adapter.get_account()
|
||
except NotImplementedError:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_400_BAD_REQUEST,
|
||
detail=f"{sdk_type} 暂不支持余额查询",
|
||
)
|
||
except HTTPException:
|
||
raise
|
||
except Exception as e: # noqa: BLE001
|
||
raise HTTPException(
|
||
status_code=status.HTTP_502_BAD_GATEWAY, detail=f"查询余额失败:{e}"
|
||
)
|
||
|
||
if acc.balance is not None:
|
||
ai.balance = acc.balance
|
||
ai.currency = acc.currency
|
||
ai.last_synced_at = utcnow()
|
||
session.add(ai)
|
||
session.commit()
|
||
session.refresh(ai)
|
||
return {
|
||
"balance": ai.balance,
|
||
"currency": ai.currency,
|
||
"last_synced_at": ai.last_synced_at.isoformat(),
|
||
}
|
||
|
||
|
||
def sync_all_ai_balances(session: Session, max_workers: int = 5) -> dict:
|
||
"""并发刷新所有 AI 账号余额(单个失败不中断整体)
|
||
|
||
每个账号需调用外部 API(单次最长 30s),串行时总耗时随账号数线性增长;
|
||
改为线程池并发后显著提速。注意:SQLite Session 不能跨线程共享,
|
||
每个 worker 使用独立 Session(WAL 模式下多连接读写安全)。
|
||
"""
|
||
from concurrent.futures import ThreadPoolExecutor
|
||
|
||
from app.database import assets_engine
|
||
|
||
ai_assets = session.exec(
|
||
select(Asset).where(Asset.asset_type == AssetType.AI_AGENT)
|
||
).all()
|
||
asset_ids = [(a.id, a.name) for a in ai_assets]
|
||
|
||
def _refresh_one(asset_id: int):
|
||
with Session(assets_engine) as worker_session:
|
||
refresh_ai_balance(worker_session, asset_id)
|
||
|
||
success = 0
|
||
failed = 0
|
||
errors = []
|
||
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
||
futures = {pool.submit(_refresh_one, aid): name for aid, name in asset_ids}
|
||
for future in futures:
|
||
name = futures[future]
|
||
try:
|
||
future.result()
|
||
success += 1
|
||
except HTTPException as e:
|
||
failed += 1
|
||
errors.append(f"{name}: {e.detail}")
|
||
except Exception as e: # noqa: BLE001
|
||
failed += 1
|
||
errors.append(f"{name}: {e}")
|
||
return {
|
||
"total": len(asset_ids),
|
||
"success": success,
|
||
"failed": failed,
|
||
"errors": errors,
|
||
"synced_at": utcnow().isoformat(),
|
||
}
|