from __future__ import annotations import uuid from collections import defaultdict from typing import Iterable from fastapi import HTTPException, status from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.crud import document as document_crud from app.crud import etmf as etmf_crud from app.core.deps import get_operator_role_label from app.models.document import Document, DocumentStatus from app.models.etmf import EtmfNode from app.models.audit_log import AuditLog from app.schemas.document import DocumentCreate, DocumentSummary from app.schemas.etmf import EtmfNodeCreate, EtmfNodeRead, EtmfNodeStatus, EtmfTreeNode, EtmfNodeUpdate def calculate_node_status(node: EtmfNode, documents: Iterable[Document]) -> EtmfNodeStatus: docs = list(documents) if not node.is_active: return EtmfNodeStatus.INACTIVE if any(doc.current_effective_version_id for doc in docs): return EtmfNodeStatus.EFFECTIVE if docs: return EtmfNodeStatus.UPLOADED if node.required: return EtmfNodeStatus.MISSING return EtmfNodeStatus.NOT_REQUIRED def _tree_node(node: EtmfNode, documents: list[Document], children: list[EtmfTreeNode]) -> EtmfTreeNode: return EtmfTreeNode( id=node.id, study_id=node.study_id, parent_id=node.parent_id, code=node.code, name=node.name, description=node.description, scope_type=node.scope_type, required=node.required, expected_doc_type=node.expected_doc_type, sort_order=node.sort_order, is_active=node.is_active, created_at=node.created_at, updated_at=node.updated_at, status=calculate_node_status(node, documents), document_count=len(documents), effective_document_count=sum(1 for doc in documents if doc.current_effective_version_id), children=children, ) async def _documents_by_node(db: AsyncSession, study_id: uuid.UUID) -> dict[uuid.UUID, list[Document]]: result = await db.execute( select(Document).where( Document.trial_id == study_id, Document.etmf_node_id.is_not(None), Document.status != DocumentStatus.ARCHIVED, ) ) grouped: dict[uuid.UUID, list[Document]] = defaultdict(list) for document in result.scalars().all(): if document.etmf_node_id: grouped[document.etmf_node_id].append(document) return grouped async def list_etmf_tree(db: AsyncSession, *, study_id: uuid.UUID, current_user) -> list[EtmfTreeNode]: if current_user is not None: from app.services import document_service await document_service._ensure_study_access(db, study_id, current_user, action="view") nodes = list(await etmf_crud.list_nodes_by_study(db, study_id)) documents = await _documents_by_node(db, study_id) children_by_parent: dict[uuid.UUID | None, list[EtmfNode]] = defaultdict(list) for node in nodes: children_by_parent[node.parent_id].append(node) def build(node: EtmfNode) -> EtmfTreeNode: children = [build(child) for child in children_by_parent.get(node.id, [])] return _tree_node(node, documents.get(node.id, []), children) return [build(node) for node in children_by_parent.get(None, [])] async def create_node(db: AsyncSession, payload: EtmfNodeCreate, current_user) -> EtmfNode: from app.services import document_service await document_service._ensure_study_access(db, payload.study_id, current_user, action="create_document") if payload.parent_id: parent = await etmf_crud.get_node(db, payload.parent_id) if not parent or parent.study_id != payload.study_id: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="父级eTMF目录不存在或不属于当前项目") node = EtmfNode(**payload.model_dump()) db.add(node) db.add( AuditLog( study_id=payload.study_id, entity_type="ETMF_NODE", entity_id=node.id, action="ETMF_NODE_CREATED", detail=f'{{"code":"{node.code}","name":"{node.name}"}}', operator_id=current_user.id, operator_role=await get_operator_role_label(db, payload.study_id, current_user), ) ) await db.commit() await db.refresh(node) return node async def update_node(db: AsyncSession, node_id: uuid.UUID, payload: EtmfNodeUpdate, current_user) -> EtmfNode: from app.services import document_service node = await etmf_crud.get_node(db, node_id) if not node: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="eTMF目录不存在") await document_service._ensure_study_access(db, node.study_id, current_user, action="create_version") values = payload.model_dump(exclude_unset=True) if "parent_id" in values and values["parent_id"]: parent = await etmf_crud.get_node(db, values["parent_id"]) if not parent or parent.study_id != node.study_id or parent.id == node.id: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="父级eTMF目录不合法") await etmf_crud.update_node(db, node_id, values, commit=False) db.add( AuditLog( study_id=node.study_id, entity_type="ETMF_NODE", entity_id=node.id, action="ETMF_NODE_UPDATED", detail="{}", operator_id=current_user.id, operator_role=await get_operator_role_label(db, node.study_id, current_user), ) ) await db.commit() refreshed = await etmf_crud.get_node(db, node_id) if not refreshed: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="eTMF目录不存在") return refreshed async def list_node_documents( db: AsyncSession, *, node_id: uuid.UUID, site_id: uuid.UUID | None, current_user, ) -> list[DocumentSummary]: from app.services import document_service node = await etmf_crud.get_node(db, node_id) if not node: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="eTMF目录不存在") return await document_service.list_documents( db, trial_id=node.study_id, site_id=site_id, doc_type=None, status=None, scope_type=None, etmf_node_id=node_id, skip=0, limit=500, current_user=current_user, ) async def create_node_document( db: AsyncSession, *, node_id: uuid.UUID, payload: DocumentCreate, current_user, ) -> DocumentSummary: from app.services import document_service node = await etmf_crud.get_node(db, node_id) if not node: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="eTMF目录不存在") doc_payload = payload.model_copy(update={"trial_id": node.study_id, "etmf_node_id": node.id, "scope_type": node.scope_type}) document = await document_service.create_document(db, doc_payload, current_user) return DocumentSummary.model_validate(document)