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

358 lines
13 KiB
Python

"""资产业务逻辑:CRUD + 统计聚合
统一处理 Asset 主表与其一对一详情表(VPSDetail/DomainDetail/AIAccount)的联动。
"""
from datetime import date, datetime
from typing import List, Optional
from fastapi import HTTPException, status
from sqlmodel import Session, select
from app.core import crypto
from app.models.asset import (
AIAccount,
Asset,
AssetStatus,
AssetType,
CloudflareDetail,
DomainDetail,
VPSDetail,
)
from app.models.provider import Provider
from app.schemas.asset import (
AIAccountRead,
AssetCreate,
AssetRead,
AssetUpdate,
CloudflareDetailRead,
DomainDetailRead,
VPSDetailRead,
)
# 资产类型 -> (AssetCreate/Update 中的字段名, 详情表模型)
DETAIL_MAP = {
AssetType.VPS: ("vps_detail", VPSDetail),
AssetType.DOMAIN: ("domain_detail", DomainDetail),
AssetType.AI_AGENT: ("ai_detail", AIAccount),
AssetType.CLOUDFLARE: ("cloudflare_detail", CloudflareDetail),
}
# 允许排序的字段白名单
SORTABLE_FIELDS = {"expiry_date", "name", "created_at", "cost", "updated_at"}
# --------------------------------------------------------------------------- #
# 辅助函数
# --------------------------------------------------------------------------- #
def _days_to_expiry(expiry_date: Optional[date]) -> Optional[int]:
"""计算距到期天数(已过期为负数),无到期日返回 None"""
if not expiry_date:
return None
return (expiry_date - date.today()).days
def _get_detail(session: Session, asset: Asset):
"""根据资产类型读取对应详情记录"""
item = DETAIL_MAP.get(asset.asset_type)
if not item:
return None
_, model = item
return session.exec(select(model).where(model.asset_id == asset.id)).first()
def _detail_to_read(detail):
"""详情模型转输出模型(VPS 计算凭证存在标记,不输出敏感明文)"""
if isinstance(detail, VPSDetail):
read = VPSDetailRead.model_validate(detail)
read.has_ssh_key = bool(detail.ssh_key_encrypted)
read.has_password = bool(detail.password_encrypted)
return read
if isinstance(detail, DomainDetail):
return DomainDetailRead.model_validate(detail)
if isinstance(detail, AIAccount):
read = AIAccountRead.model_validate(detail)
read.has_api_key = bool(detail.api_key_encrypted or detail.api_key)
return read
if isinstance(detail, CloudflareDetail):
return CloudflareDetailRead.model_validate(detail)
return None
def _provider_name(session: Session, asset: Asset) -> Optional[str]:
"""取关联平台显示名"""
if asset.provider_id:
provider = session.get(Provider, asset.provider_id)
if provider:
return provider.name
return None
def _to_read(session: Session, asset: Asset, detail) -> AssetRead:
"""组装 AssetRead 输出(主表 + 详情 + 计算字段)"""
read = AssetRead.model_validate(asset)
read.days_to_expiry = _days_to_expiry(asset.expiry_date)
read.provider_name = _provider_name(session, asset)
if isinstance(detail, VPSDetail):
read.vps_detail = _detail_to_read(detail)
elif isinstance(detail, DomainDetail):
read.domain_detail = _detail_to_read(detail)
elif isinstance(detail, AIAccount):
read.ai_detail = _detail_to_read(detail)
elif isinstance(detail, CloudflareDetail):
read.cloudflare_detail = _detail_to_read(detail)
return read
def _build_detail(asset_type: AssetType, asset_id: int, detail_in):
"""构建详情模型(VPS 加密凭证)"""
_, model = DETAIL_MAP[asset_type]
if asset_type == AssetType.VPS:
data = detail_in.model_dump(exclude={"ssh_key", "password"})
data["ssh_key_encrypted"] = crypto.encrypt(detail_in.ssh_key)
data["password_encrypted"] = crypto.encrypt(detail_in.password)
return model(asset_id=asset_id, **data)
if asset_type == AssetType.AI_AGENT:
data = detail_in.model_dump(exclude={"api_key"})
data["api_key_encrypted"] = crypto.encrypt(detail_in.api_key)
return model(asset_id=asset_id, **data)
return model(asset_id=asset_id, **detail_in.model_dump())
def _apply_detail_update(asset_type: AssetType, existing, detail_in) -> None:
"""更新详情字段(VPS 加密凭证;凭证为 None 时保留原值)"""
if asset_type == AssetType.VPS:
data = detail_in.model_dump(exclude={"ssh_key", "password"})
for key, value in data.items():
setattr(existing, key, value)
if detail_in.ssh_key is not None:
existing.ssh_key_encrypted = crypto.encrypt(detail_in.ssh_key)
if detail_in.password is not None:
existing.password_encrypted = crypto.encrypt(detail_in.password)
elif asset_type == AssetType.AI_AGENT:
data = detail_in.model_dump(exclude={"api_key"})
for key, value in data.items():
setattr(existing, key, value)
if detail_in.api_key is not None:
existing.api_key_encrypted = crypto.encrypt(detail_in.api_key)
else:
for key, value in detail_in.model_dump().items():
setattr(existing, key, value)
# --------------------------------------------------------------------------- #
# CRUD
# --------------------------------------------------------------------------- #
def create_asset(session: Session, data: AssetCreate) -> AssetRead:
"""创建资产及其详情"""
item = DETAIL_MAP.get(data.asset_type)
detail_in = None
model = None
if item:
field, model = item
detail_in = getattr(data, field)
if detail_in is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"asset_type={data.asset_type.value} 需要提供 {field}",
)
asset_data = data.model_dump(
exclude={"vps_detail", "domain_detail", "ai_detail", "cloudflare_detail"}
)
asset = Asset(**asset_data)
session.add(asset)
session.commit()
session.refresh(asset)
detail = None
if model is not None and detail_in is not None:
detail = _build_detail(data.asset_type, asset.id, detail_in)
session.add(detail)
session.commit()
session.refresh(detail)
return _to_read(session, asset, detail)
def get_asset(session: Session, asset_id: int) -> AssetRead:
"""读取单个资产(含详情)"""
asset = session.get(Asset, asset_id)
if not asset:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="资产不存在")
return _to_read(session, asset, _get_detail(session, asset))
def update_asset(session: Session, asset_id: int, data: AssetUpdate) -> AssetRead:
"""更新资产主表及详情(仅更新传入字段)"""
asset = session.get(Asset, asset_id)
if not asset:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="资产不存在")
main_fields = data.model_dump(
exclude_unset=True, exclude={"vps_detail", "domain_detail", "ai_detail", "cloudflare_detail"}
)
for key, value in main_fields.items():
setattr(asset, key, value)
session.add(asset)
# 详情表:以更新后的 asset_type 为准
detail = None
item = DETAIL_MAP.get(asset.asset_type)
if item:
field, model = item
detail_in = getattr(data, field)
existing = session.exec(
select(model).where(model.asset_id == asset.id)
).first()
if detail_in is not None:
if existing:
_apply_detail_update(asset.asset_type, existing, detail_in)
session.add(existing)
detail = existing
else:
detail = _build_detail(asset.asset_type, asset.id, detail_in)
session.add(detail)
else:
detail = existing
session.commit()
session.refresh(asset)
if detail is not None:
session.refresh(detail)
return _to_read(session, asset, detail)
def delete_asset(session: Session, asset_id: int) -> None:
"""删除资产及其详情"""
asset = session.get(Asset, asset_id)
if not asset:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="资产不存在")
detail = _get_detail(session, asset)
if detail is not None:
session.delete(detail)
session.delete(asset)
session.commit()
def export_assets(session: Session) -> dict:
"""导出所有资产(含 detail)为 JSON,用于数据迁移/备份"""
assets = session.exec(select(Asset)).all()
result = []
for asset in assets:
read = _to_read(session, asset, _get_detail(session, asset))
result.append(read.model_dump(mode="json"))
return {
"count": len(result),
"exported_at": datetime.utcnow().isoformat(),
"assets": result,
}
def import_assets(session: Session, assets_data: list) -> dict:
"""从导出数据导入资产(按 name+asset_type+provider 去重:存在则更新,不存在则创建)"""
created = 0
updated = 0
errors = []
skip_fields = {"id", "created_at", "updated_at", "days_to_expiry", "provider_name"}
for item in assets_data:
try:
payload = {k: v for k, v in item.items() if k not in skip_fields}
existing = session.exec(
select(Asset).where(
Asset.name == payload.get("name"),
Asset.asset_type == payload.get("asset_type"),
Asset.provider == payload.get("provider"),
)
).first()
if existing:
update_asset(session, existing.id, AssetUpdate(**payload))
updated += 1
else:
create_asset(session, AssetCreate(**payload))
created += 1
except Exception as e: # noqa: BLE001
errors.append(f"{item.get('name', '?')}: {e}")
return {"created": created, "updated": updated, "errors": errors}
def list_assets(
session: Session,
asset_type: Optional[AssetType] = None,
asset_status: Optional[AssetStatus] = None,
is_archived: Optional[bool] = None,
q: Optional[str] = None,
sort: str = "expiry_date",
order: str = "asc",
) -> List[AssetRead]:
"""资产列表(筛选 + 搜索 + 排序)"""
stmt = select(Asset)
if asset_type is not None:
stmt = stmt.where(Asset.asset_type == asset_type)
if asset_status is not None:
stmt = stmt.where(Asset.status == asset_status)
if is_archived is not None:
stmt = stmt.where(Asset.is_archived == is_archived)
if q:
pattern = f"%{q}%"
stmt = stmt.where(Asset.name.like(pattern) | Asset.provider.like(pattern))
sort_col = getattr(Asset, sort if sort in SORTABLE_FIELDS else "expiry_date")
stmt = stmt.order_by(sort_col.desc() if order == "desc" else sort_col.asc())
assets = session.exec(stmt).all()
return [_to_read(session, a, _get_detail(session, a)) for a in assets]
# --------------------------------------------------------------------------- #
# 统计
# --------------------------------------------------------------------------- #
def get_overview(session: Session) -> dict:
"""资产总览:数量分布、异常数、到期预警、支出合计"""
assets = session.exec(select(Asset)).all()
today = date.today()
by_type: dict = {}
by_status: dict = {}
abnormal = 0
expiring_30 = 0
year_cost = 0.0
month_cost = 0.0
for a in assets:
by_type[a.asset_type.value] = by_type.get(a.asset_type.value, 0) + 1
by_status[a.status.value] = by_status.get(a.status.value, 0) + 1
if a.status in (AssetStatus.STOPPED, AssetStatus.EXPIRED):
abnormal += 1
if a.expiry_date:
days = (a.expiry_date - today).days
if 0 <= days <= 30:
expiring_30 += 1
if a.expiry_date.year == today.year:
year_cost += a.cost
if a.expiry_date.month == today.month:
month_cost += a.cost
return {
"total": len(assets),
"by_type": by_type,
"by_status": by_status,
"abnormal_count": abnormal,
"expiring_30": expiring_30,
"year_cost": round(year_cost, 2),
"month_cost": round(month_cost, 2),
}
def get_expiring(session: Session, days: int = 30) -> List[AssetRead]:
"""N 天内到期资产列表(按剩余天数升序)"""
today = date.today()
assets = session.exec(select(Asset).where(Asset.expiry_date.is_not(None))).all()
result = []
for a in assets:
delta = (a.expiry_date - today).days
if 0 <= delta <= days:
result.append(_to_read(session, a, _get_detail(session, a)))
result.sort(key=lambda x: x.days_to_expiry if x.days_to_expiry is not None else 10**9)
return result