"""资产业务逻辑:CRUD + 统计聚合 统一处理 Asset 主表与其一对一详情表(VPSDetail/DomainDetail/AIAccount)的联动。 """ from datetime import date, datetime, timedelta, timezone from typing import List, Optional from fastapi import HTTPException, status from sqlalchemy import or_ from sqlmodel import Session, select from app.core import crypto from app.models.asset import ( Account, 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, ) import logging logger = logging.getLogger("vps-manager.assets") # 资产类型 -> (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 _account_name(session: Session, asset: Asset) -> Optional[str]: """取关联账号标识(展示用)""" if asset.account_id: account = session.get(Account, asset.account_id) if account: return account.name return None def _validate_account(session: Session, account_id: Optional[int]) -> None: """account_id 非空时确认账号存在,防外键悬空""" if account_id is not None and not session.get(Account, account_id): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=f"账号不存在:id={account_id}" ) 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) read.account_name = _account_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) def _get_details_batch(session: Session, assets: list) -> dict: """批量预加载所有资产的详情和 Provider 名,避免 N+1 查询""" if not assets: return {} asset_ids = [a.id for a in assets] # 批量查 Provider 名 provider_ids = {a.provider_id for a in assets if a.provider_id} provider_map = {} if provider_ids: for p in session.exec(select(Provider).where(Provider.id.in_(provider_ids))).all(): provider_map[p.id] = p.name # 批量查 Account 名 account_ids = {a.account_id for a in assets if a.account_id} account_map = {} if account_ids: for acc in session.exec(select(Account).where(Account.id.in_(account_ids))).all(): account_map[acc.id] = acc.name # 批量查各类型详情 detail_map = {} for asset_type, (_, model) in DETAIL_MAP.items(): typed_ids = [a.id for a in assets if a.asset_type == asset_type] if typed_ids: for d in session.exec(select(model).where(model.asset_id.in_(typed_ids))).all(): detail_map[d.asset_id] = d return {"providers": provider_map, "accounts": account_map, "details": detail_map} def _to_read_batch(session: Session, asset: Asset, batch: dict) -> AssetRead: """用批量预加载的数据组装 AssetRead(避免逐个查询)""" read = AssetRead.model_validate(asset) read.days_to_expiry = _days_to_expiry(asset.expiry_date) read.provider_name = batch["providers"].get(asset.provider_id) read.account_name = batch.get("accounts", {}).get(asset.account_id) detail = batch["details"].get(asset.id) 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 # --------------------------------------------------------------------------- # # CRUD # --------------------------------------------------------------------------- # def _create_asset_no_commit(session: Session, data: AssetCreate) -> Asset: """创建资产及详情(不 commit,由调用方统一提交)""" 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"} ) _validate_account(session, asset_data.get("account_id")) asset = Asset(**asset_data) session.add(asset) session.flush() # 获取 asset.id,但不提交 if model is not None and detail_in is not None: detail = _build_detail(data.asset_type, asset.id, detail_in) session.add(detail) return asset def create_asset(session: Session, data: AssetCreate) -> AssetRead: """创建资产及其详情(单次事务提交)""" asset = _create_asset_no_commit(session, data) session.commit() session.refresh(asset) logger.info("创建资产 id=%s name=%s type=%s", asset.id, asset.name, asset.asset_type) return _to_read(session, asset, _get_detail(session, asset)) 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_no_commit(session: Session, asset_id: int, data: AssetUpdate) -> Asset: """更新资产主表及详情(不 commit,由调用方统一提交)""" 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"} ) if "account_id" in main_fields: _validate_account(session, main_fields["account_id"]) for key, value in main_fields.items(): setattr(asset, key, value) session.add(asset) # 详情表:以更新后的 asset_type 为准 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) else: detail = _build_detail(asset.asset_type, asset.id, detail_in) session.add(detail) return asset def update_asset(session: Session, asset_id: int, data: AssetUpdate) -> AssetRead: """更新资产主表及详情(单次事务提交)""" asset = _update_asset_no_commit(session, asset_id, data) session.commit() session.refresh(asset) return _to_read(session, asset, _get_detail(session, asset)) def delete_asset(session: Session, asset_id: int) -> None: """删除资产及其详情,并清理 metrics.db 中的关联监控数据(避免孤儿数据)""" 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() logger.info("删除资产 id=%s name=%s", asset_id, asset.name) _cleanup_metrics_for_asset(asset_id) def _cleanup_metrics_for_asset(asset_id: int) -> None: """清理 metrics.db 中该资产的 MetricPoint/ServerInfo/SecurityCheck/EventLog 使用批量 DELETE(而非逐行加载后删除):监控时序数据可能上万行, 逐行删除会全部载入内存且产生数万次 ORM 操作。 """ from sqlmodel import delete from app.database import metrics_engine from app.models.monitor import EventLog, MetricPoint, SecurityCheck, ServerInfo with Session(metrics_engine) as ms: for model in (MetricPoint, ServerInfo, SecurityCheck, EventLog): ms.exec(delete(model).where(model.asset_id == asset_id)) ms.commit() def export_assets(session: Session) -> dict: """导出所有资产(含 detail)为 JSON,用于数据迁移/备份""" assets = session.exec(select(Asset)).all() batch = _get_details_batch(session, assets) result = [] for asset in assets: read = _to_read_batch(session, asset, batch) result.append(read.model_dump(mode="json")) return { "count": len(result), "exported_at": datetime.now(timezone.utc).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"} try: 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_no_commit(session, existing.id, AssetUpdate(**payload)) updated += 1 else: _create_asset_no_commit(session, AssetCreate(**payload)) created += 1 except Exception as e: # noqa: BLE001 errors.append(f"{item.get('name', '?')}: {e}") if errors: session.rollback() return {"created": 0, "updated": 0, "errors": errors} session.commit() except Exception: # noqa: BLE001 session.rollback() raise 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, provider: 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 provider: # 兼容两种关联方式:provider slug 文本或 provider_id 外键 conds = [Asset.provider == provider] if provider.isdigit(): conds.append(Asset.provider_id == int(provider)) stmt = stmt.where(or_(*conds)) 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() batch = _get_details_batch(session, assets) return [_to_read_batch(session, a, batch) for a in assets] # --------------------------------------------------------------------------- # # 统计 # --------------------------------------------------------------------------- # def get_overview(session: Session) -> dict: """资产总览:数量分布、异常数、到期预警、支出合计 性能优化:仅查询统计所需的列,避免加载完整 Asset 实体。 """ stmt = select(Asset.asset_type, Asset.status, Asset.expiry_date, Asset.cost) rows = session.exec(stmt).all() today = date.today() by_type: dict = {} by_status: dict = {} abnormal = 0 expiring_30 = 0 year_cost = 0.0 month_cost = 0.0 for asset_type, status_val, expiry_date, cost in rows: tval = asset_type.value if hasattr(asset_type, "value") else asset_type sval = status_val.value if hasattr(status_val, "value") else status_val by_type[tval] = by_type.get(tval, 0) + 1 by_status[sval] = by_status.get(sval, 0) + 1 if status_val in (AssetStatus.STOPPED, AssetStatus.EXPIRED): abnormal += 1 if expiry_date: days = (expiry_date - today).days if 0 <= days <= 30: expiring_30 += 1 if expiry_date.year == today.year: year_cost += cost if expiry_date.month == today.month: month_cost += cost return { "total": len(rows), "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 天内到期资产列表(按剩余天数升序) 性能优化:过滤条件下推到 SQL 层,仅查询 [today, today+days] 区间内的资产。 """ today = date.today() deadline = today + timedelta(days=days) assets = session.exec( select(Asset).where( Asset.expiry_date.is_not(None), Asset.expiry_date >= today, Asset.expiry_date <= deadline, ) ).all() batch = _get_details_batch(session, assets) result = [_to_read_batch(session, a, batch) for a in assets] result.sort(key=lambda x: x.days_to_expiry if x.days_to_expiry is not None else 10**9) return result