from __future__ import annotations import uuid from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession from app.models.api_endpoint_permission import ApiEndpointPermission from app.core.api_permissions import API_ENDPOINT_PERMISSIONS, OPERATION_PREREQUISITES from app.core.permission_cache import get_permission_cache async def role_has_api_permission( db: AsyncSession, study_id: uuid.UUID, role: str | None, endpoint_key: str, check_prerequisites: bool = True, ) -> bool: """检查角色是否有权访问特定接口""" if role == "ADMIN": return True result = await db.execute( select(ApiEndpointPermission).where( ApiEndpointPermission.study_id == study_id, ApiEndpointPermission.role == role, ApiEndpointPermission.endpoint_key == endpoint_key, ) ) perm = result.scalar_one_or_none() if perm is None: return False if not perm.allowed: return False if check_prerequisites: prerequisites = OPERATION_PREREQUISITES.get(endpoint_key, []) for prereq_endpoint in prerequisites: has_prereq = await role_has_api_permission( db, study_id, role, prereq_endpoint, check_prerequisites=False ) if not has_prereq: return False return True async def get_missing_prerequisites( db: AsyncSession, study_id: uuid.UUID, role: str | None, endpoint_key: str, ) -> list[str]: """获取缺失的前置权限列表""" if role == "ADMIN": return [] missing = [] prerequisites = OPERATION_PREREQUISITES.get(endpoint_key, []) for prereq_endpoint in prerequisites: has_prereq = await role_has_api_permission( db, study_id, role, prereq_endpoint, check_prerequisites=False ) if not has_prereq: missing.append(prereq_endpoint) return missing async def get_api_endpoint_permissions( db: AsyncSession, study_id: uuid.UUID, ) -> dict[str, dict[str, dict[str, bool]]]: """获取项目的接口级权限矩阵 返回格式: {role: {endpoint_key: {allowed: bool}}} """ result = await db.execute( select(ApiEndpointPermission).where( ApiEndpointPermission.study_id == study_id, ) ) rows = result.scalars().all() roles = ["PM", "CRA", "PV", "MEDICAL_REVIEW", "IMP", "QA"] matrix: dict[str, dict[str, dict[str, bool]]] = {} for role in roles: matrix[role] = {} for endpoint_key, config in API_ENDPOINT_PERMISSIONS.items(): default_allowed = role in config.get("default_roles", []) matrix[role][endpoint_key] = {"allowed": default_allowed} for row in rows: if row.role not in matrix: matrix[row.role] = {} matrix[row.role][row.endpoint_key] = {"allowed": row.allowed} return matrix async def replace_api_endpoint_permissions( db: AsyncSession, study_id: uuid.UUID, payload: dict[str, dict[str, bool]], ) -> dict[str, dict[str, dict[str, bool]]]: """替换项目的接口级权限矩阵""" await db.execute( delete(ApiEndpointPermission).where( ApiEndpointPermission.study_id == study_id, ) ) for role, endpoints in payload.items(): if role == "ADMIN": continue for endpoint_key, allowed in endpoints.items(): if endpoint_key not in API_ENDPOINT_PERMISSIONS: continue db.add( ApiEndpointPermission( study_id=study_id, role=role, endpoint_key=endpoint_key, allowed=allowed, ) ) await db.commit() cache = get_permission_cache() cache.invalidate_project_permissions(study_id) cache.invalidate_all_member_roles(study_id) return await get_api_endpoint_permissions(db, study_id)