from __future__ import annotations import re from datetime import datetime, timezone from fastapi import HTTPException from sqlalchemy import case, func, select, update from sqlalchemy.ext.asyncio import AsyncSession from app.enums.private_portrait import ( PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME, PrivatePortraitAssetGroupStatus, PrivatePortraitAssetStatus, PrivatePortraitEventSource, PrivatePortraitEventStatus, PrivatePortraitEventType, PrivatePortraitProjectStatus, PrivatePortraitRemoteDeleteStatus, ) from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject, PrivatePortraitValidateSession from app.schemas.private_portrait import PrivatePortraitProjectCreate, PrivatePortraitProjectOut, PrivatePortraitProjectUpdate from app.services.operation_log_service import log_operation_event from app.utils.id_gen import generate_id DOMAIN = "private_portrait" def _safe_slug(value: str, *, max_length: int = 80) -> str: value = (value or "").strip().lower() # 先保留常见英文数字连字符;中文等字符统一转 _,仅用于本地项目 slug。 value = re.sub(r"[^a-z0-9_-]+", "_", value) value = re.sub(r"_+", "_", value).strip("_-") return (value[:max_length] or "project") def project_to_out(project: PrivatePortraitProject, *, include_user: bool = False) -> PrivatePortraitProjectOut: return PrivatePortraitProjectOut( id=project.id, user_id=project.user_id if include_user else None, name=project.name, name_slug=project.name_slug, remote_project_name=project.remote_project_name, description=project.description, status=project.status, asset_group_count=project.asset_group_count or 0, asset_count=project.asset_count or 0, active_asset_count=project.active_asset_count or 0, last_used_at=project.last_used_at, created_at=project.created_at, updated_at=project.updated_at, ) async def get_user_project(db: AsyncSession, *, user_id: str, project_id: str) -> PrivatePortraitProject: result = await db.execute( select(PrivatePortraitProject).where( PrivatePortraitProject.id == project_id, PrivatePortraitProject.user_id == user_id, PrivatePortraitProject.deleted_at.is_(None), ).limit(1) ) project = result.scalar_one_or_none() if not project: raise HTTPException(status_code=404, detail="真人素材项目不存在") return project async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePortraitProjectCreate) -> PrivatePortraitProject: slug = _safe_slug(payload.name) project = PrivatePortraitProject( id=generate_id(), user_id=user_id, name=payload.name.strip(), name_slug=slug, remote_project_name=PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME, description=payload.description, status=PrivatePortraitProjectStatus.ACTIVE.value, ) db.add(project) await db.flush() log_operation_event( domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_CREATE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, message="创建真人素材项目", detail={"name": project.name, "remote_project_name": project.remote_project_name}, ) return project async def update_project(db: AsyncSession, *, user_id: str, project_id: str, payload: PrivatePortraitProjectUpdate) -> PrivatePortraitProject: project = await get_user_project(db, user_id=user_id, project_id=project_id) before = { "name": project.name, "name_slug": project.name_slug, "remote_project_name": project.remote_project_name, "description": project.description, "status": project.status, } if payload.name is not None: new_name = payload.name.strip() if new_name and new_name != project.name: project.name = new_name project.name_slug = _safe_slug(new_name) project.remote_project_name = PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME if payload.description is not None: project.description = payload.description if payload.status is not None: if payload.status not in {PrivatePortraitProjectStatus.ACTIVE.value}: raise HTTPException(status_code=400, detail="项目状态不支持") project.status = payload.status await db.flush() after = { "name": project.name, "name_slug": project.name_slug, "remote_project_name": project.remote_project_name, "description": project.description, "status": project.status, } log_operation_event( domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_UPDATE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, message="更新真人素材项目", detail={"before": before, "after": after}, ) return project async def list_projects(db: AsyncSession, *, user_id: str | None, page: int = 1, page_size: int = 20, keyword: str | None = None, status: str | None = None) -> tuple[list[PrivatePortraitProject], int]: page = max(1, page) page_size = min(max(1, page_size), 100) filters = [PrivatePortraitProject.deleted_at.is_(None)] if user_id: filters.append(PrivatePortraitProject.user_id == user_id) if keyword: filters.append(PrivatePortraitProject.name.ilike(f"%{keyword.strip()}%")) if status: filters.append(PrivatePortraitProject.status == status) total = (await db.execute(select(func.count(PrivatePortraitProject.id)).where(*filters))).scalar_one() result = await db.execute(select(PrivatePortraitProject).where(*filters).order_by(PrivatePortraitProject.created_at.desc()).offset((page - 1) * page_size).limit(page_size)) return list(result.scalars().all()), int(total or 0) async def refresh_project_counters(db: AsyncSession, project_ids: list[str]) -> None: project_ids = list({pid for pid in project_ids if pid}) if not project_ids: return group_rows = await db.execute( select(PrivatePortraitAssetGroup.project_id, func.count(PrivatePortraitAssetGroup.id)) .where(PrivatePortraitAssetGroup.project_id.in_(project_ids), PrivatePortraitAssetGroup.deleted_at.is_(None)) .group_by(PrivatePortraitAssetGroup.project_id) ) asset_rows = await db.execute( select( PrivatePortraitAsset.project_id, func.count(PrivatePortraitAsset.id), func.sum(case((PrivatePortraitAsset.status == PrivatePortraitAssetStatus.ACTIVE.value, 1), else_=0)), ) .where(PrivatePortraitAsset.project_id.in_(project_ids), PrivatePortraitAsset.deleted_at.is_(None)) .group_by(PrivatePortraitAsset.project_id) ) group_count_map = {pid: int(count or 0) for pid, count in group_rows.all()} asset_count_map: dict[str, tuple[int, int]] = {} for pid, total, active_total in asset_rows.all(): asset_count_map[pid] = (int(total or 0), int(active_total or 0)) for pid in project_ids: total, active_total = asset_count_map.get(pid, (0, 0)) await db.execute(update(PrivatePortraitProject).where(PrivatePortraitProject.id == pid).values(asset_group_count=group_count_map.get(pid, 0), asset_count=total, active_asset_count=active_total)) async def soft_delete_project(db: AsyncSession, *, user_id: str, project_id: str) -> PrivatePortraitProject: project = await get_user_project(db, user_id=user_id, project_id=project_id) now = datetime.now(timezone.utc) project.deleted_at = now project.status = PrivatePortraitProjectStatus.DELETED.value await db.execute(update(PrivatePortraitAsset).where(PrivatePortraitAsset.project_id == project_id, PrivatePortraitAsset.deleted_at.is_(None)).values(deleted_at=now, status=PrivatePortraitAssetStatus.LOCAL_DELETED.value, remote_delete_status=PrivatePortraitRemoteDeleteStatus.PENDING.value)) await db.execute(update(PrivatePortraitAssetGroup).where(PrivatePortraitAssetGroup.project_id == project_id, PrivatePortraitAssetGroup.deleted_at.is_(None)).values(deleted_at=now, status=PrivatePortraitAssetGroupStatus.LOCAL_DELETED.value, remote_delete_status=PrivatePortraitRemoteDeleteStatus.PENDING.value)) await db.flush() log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, message="本地软删真人素材项目", detail={"remote_project_name": project.remote_project_name}) return project