152 lines
4.4 KiB
Python
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
|