""" 项目域相关服务 """ from typing import Optional, Sequence from fastapi import HTTPException from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.project import Project, ProjectMember, ProjectMemberRole from app.models.user import User from app.schemas.project import ProjectResponse from app.services.storage import storage_service OWNER_ROLE = "owner" PUBLIC_ROLE = "public" async def get_project_or_404(db: AsyncSession, project_id: int) -> Project: """获取项目,不存在时抛出 404。""" result = await db.execute(select(Project).where(Project.id == project_id)) project = result.scalar_one_or_none() if not project: raise HTTPException(status_code=404, detail="项目不存在") return project async def get_project_member( db: AsyncSession, project_id: int, user_id: int, ) -> Optional[ProjectMember]: """查询项目成员记录。""" result = await db.execute( select(ProjectMember).where( ProjectMember.project_id == project_id, ProjectMember.user_id == user_id, ) ) return result.scalar_one_or_none() async def get_project_role( db: AsyncSession, project: Project, current_user: Optional[User], ) -> Optional[str]: """解析用户在项目中的角色。""" if not current_user: return None if project.owner_id == current_user.id: return OWNER_ROLE member = await get_project_member(db, project.id, current_user.id) return member.role if member else None async def require_project_read_access( db: AsyncSession, project_id: int, current_user: Optional[User], *, allow_public: bool = False, unauthenticated_detail: str = "请先登录", forbidden_detail: str = "无权访问该项目", ) -> tuple[Project, str]: """校验项目读取权限。""" project = await get_project_or_404(db, project_id) role = await get_project_role(db, project, current_user) if role: return project, role if allow_public and project.is_public == 1: return project, PUBLIC_ROLE if not current_user: raise HTTPException(status_code=401, detail=unauthenticated_detail) raise HTTPException(status_code=403, detail=forbidden_detail) async def require_project_roles( db: AsyncSession, project_id: int, current_user: Optional[User], *, allowed_roles: Sequence[str], allow_public: bool = False, unauthenticated_detail: str = "请先登录", forbidden_detail: str = "无权执行此操作", ) -> tuple[Project, str]: """校验项目角色权限。owner 始终视为通过。""" project, role = await require_project_read_access( db, project_id, current_user, allow_public=allow_public, unauthenticated_detail=unauthenticated_detail, forbidden_detail=forbidden_detail, ) if role in {OWNER_ROLE, *allowed_roles}: return project, role raise HTTPException(status_code=403, detail=forbidden_detail) async def require_project_write_access( db: AsyncSession, project_id: int, current_user: User, ) -> tuple[Project, str]: """校验项目写权限。""" return await require_project_roles( db, project_id, current_user, allowed_roles=[ ProjectMemberRole.ADMIN.value, ProjectMemberRole.EDITOR.value, ], forbidden_detail="无写入权限", ) def count_project_documents(storage_key: str) -> int: """统计项目中可见文档数量。""" try: project_path = storage_service.get_secure_path(storage_key) if not project_path.exists(): return 0 md_count = len(list(project_path.rglob("*.md"))) pdf_count = len(list(project_path.rglob("*.pdf"))) assets_dir = project_path / "_assets" assets_md = len(list(assets_dir.rglob("*.md"))) if assets_dir.exists() else 0 assets_pdf = len(list(assets_dir.rglob("*.pdf"))) if assets_dir.exists() else 0 return md_count + pdf_count - assets_md - assets_pdf except Exception: return 0 def serialize_project(project: Project, **extra_fields) -> dict: """序列化项目,并补充文档统计。""" project_data = ProjectResponse.from_orm(project).dict() project_data["doc_count"] = count_project_documents(project.storage_key) project_data.update(extra_fields) return project_data