Files
vps-manager/app/services/sync_service.py
T

434 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""云厂商同步服务
将适配器返回的标准化资产(NormalizedVPS / NormalizedDomain / AccountInfo
写入或更新到资产库。以 (provider_id, external_id) 作为去重键,
已存在则更新状态/详情,不存在则新建资产。
凭证层级:API 配置存在账号(Account.api_config_encrypted)上,
同步按账号维度进行(test_account/sync_account);同步产出的资产
自动挂到该账号名下(Asset.account)。平台不再持有凭证。
"""
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,
Account,
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 配置,仅为兼容历史数据保留"""
plain = crypto.decrypt(provider.api_config_encrypted)
if not plain:
return {}
try:
return json.loads(plain)
except (json.JSONDecodeError, TypeError):
return {}
def _load_account_config(account: Account) -> dict:
"""解密账号的 API 配置 JSON"""
plain = crypto.decrypt(account.api_config_encrypted)
if not plain:
return {}
try:
return json.loads(plain)
except (json.JSONDecodeError, TypeError):
return {}
def _find_provider_by_platform(session: Session, platform: str) -> Provider:
"""按账号的 platformslug 或名称)定位平台"""
provider = session.exec(select(Provider).where(Provider.slug == platform)).first()
if not provider:
provider = session.exec(select(Provider).where(Provider.name == platform)).first()
if not provider:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"账号所属平台不存在:{platform}(请先在账号编辑中选择有效平台)",
)
return provider
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, config: dict) -> 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, config)
def _get_account(session: Session, account_id: int) -> Account:
account = session.get(Account, account_id)
if not account:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="账号不存在")
return account
def _build_adapter_for_account(session: Session, account: Account):
"""按账号构建适配器:平台定 sdk_type,账号提供凭证"""
if not account.platform:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该账号未指定所属平台,无法确定 SDK 类型",
)
provider = _find_provider_by_platform(session, account.platform)
config = _load_account_config(account)
adapter = _build_adapter(provider, config)
return provider, adapter
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)
account = _first_account_with_config(session, provider)
if not account:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="凭证已改为在账号上配置:请先在该平台的账号管理中新建账号并填写 API 配置",
)
return test_account(session, account.id)
def _first_account_with_config(session: Session, provider: Provider):
"""找该平台下第一个配置了 API 凭证的账号(兼容旧入口)"""
return session.exec(
select(Account)
.where(Account.platform == provider.slug)
.where(Account.api_config_encrypted.is_not(None)) # type: ignore[union-attr]
.order_by(Account.id.asc())
).first()
def test_account(session: Session, account_id: int) -> dict:
"""测试账号凭证有效性(按账号 platform 定 SDK,凭证取自账号)"""
account = _get_account(session, account_id)
provider, adapter = _build_adapter_for_account(session, account)
base = {"capabilities": adapter.capabilities(), "sdk_type": provider.sdk_type, "account": account.name}
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, account_name: Optional[str] = None) -> 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
if account_name:
existing.account = account_name
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,
account=account_name,
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, account_name: Optional[str] = None) -> 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
if account_name:
existing.account = account_name
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,
account=account_name,
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)
account = _first_account_with_config(session, provider)
if not account:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="凭证已改为在账号上配置:请先在该平台的账号管理中新建账号并填写 API 配置",
)
return sync_account(session, account.id)
def sync_account(session: Session, account_id: int) -> dict:
"""同步账号资产到本地库(凭证取自账号,同步产出自动挂到该账号名下)"""
account = _get_account(session, account_id)
provider, adapter = _build_adapter_for_account(session, account)
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, "account": account.name}
if caps["list_vps"]:
try:
result["vps"] = _sync_vps(session, provider, adapter, account.name)
except Exception as e: # noqa: BLE001
result["vps_error"] = str(e)
if caps["list_domains"]:
try:
result["domains"] = _sync_domains(session, provider, adapter, account.name)
except Exception as e: # noqa: BLE001
result["domains_error"] = str(e)
if caps["get_account"]:
try:
result["account_info"] = adapter.get_account().to_dict()
except Exception as e: # noqa: BLE001
result["account_info_error"] = str(e)
account.last_synced_at = utcnow()
session.add(account)
session.commit()
result["last_synced_at"] = account.last_synced_at.isoformat()
logger.info("同步账号 account=%s provider=%s result=%s", account.name, 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(),
}