nex_docus/backend/app/services/project_service.py

152 lines
4.4 KiB
Python

"""
项目域相关服务
"""
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