from __future__ import annotations import os import shutil from pathlib import Path from uuid import uuid4 from sqlalchemy import select from sqlalchemy.orm import Session, joinedload from app.modules.documents.models import EngagementDocument from app.modules.documents.services import DEFAULT_STORAGE_ROOT, sanitize_segment, save_uploaded_revision from app.modules.services.models import ( ClientServiceSubscription, ClientServiceTaskInstance, FirmServiceTaskTemplate, FirmTaskDocumentRequirement, FirmTaskDocumentTemplate, ) TEMPLATE_UPLOAD_ROOT = Path(os.getenv("DOCUMENT_TEMPLATE_STORAGE_ROOT", str(DEFAULT_STORAGE_ROOT.parent / "document_templates"))).resolve() def list_task_document_requirements(db: Session, *, tenant_id: int, firm_task_template_id: int) -> list[FirmTaskDocumentRequirement]: return db.execute( select(FirmTaskDocumentRequirement) .where( FirmTaskDocumentRequirement.tenant_id == int(tenant_id), FirmTaskDocumentRequirement.firm_task_template_id == int(firm_task_template_id), ) .order_by(FirmTaskDocumentRequirement.sort_order.asc(), FirmTaskDocumentRequirement.id.asc()) ).scalars().all() def get_task_document_requirement(db: Session, *, requirement_id: int, tenant_id: int | None = None) -> FirmTaskDocumentRequirement | None: stmt = select(FirmTaskDocumentRequirement).where(FirmTaskDocumentRequirement.id == int(requirement_id)) if tenant_id is not None: stmt = stmt.where(FirmTaskDocumentRequirement.tenant_id == int(tenant_id)) return db.execute(stmt).scalar_one_or_none() def create_task_document_requirement( db: Session, *, task_template: FirmServiceTaskTemplate, document_name: str, document_type: str, is_mandatory: bool, allowed_file_types: str | None, instructions: str | None, sort_order: int, user, ) -> FirmTaskDocumentRequirement: row = FirmTaskDocumentRequirement( tenant_id=task_template.tenant_id, service_catalogue_id=task_template.service_catalogue_id, firm_task_template_id=task_template.id, document_name=document_name.strip()[:200], document_type=(document_type or "GENERAL").strip().upper()[:80] or "GENERAL", is_mandatory=bool(is_mandatory), allowed_file_types=(allowed_file_types or "").strip()[:255] or None, instructions=(instructions or "").strip() or None, sort_order=int(sort_order or 100), is_active=True, created_by_user_id=getattr(user, "id", None), updated_by_user_id=getattr(user, "id", None), ) db.add(row) db.flush() return row def update_task_document_requirement( db: Session, *, requirement: FirmTaskDocumentRequirement, document_name: str, document_type: str, is_mandatory: bool, allowed_file_types: str | None, instructions: str | None, sort_order: int, is_active: bool, user, ) -> FirmTaskDocumentRequirement: requirement.document_name = document_name.strip()[:200] requirement.document_type = (document_type or "GENERAL").strip().upper()[:80] or "GENERAL" requirement.is_mandatory = bool(is_mandatory) requirement.allowed_file_types = (allowed_file_types or "").strip()[:255] or None requirement.instructions = (instructions or "").strip() or None requirement.sort_order = int(sort_order or 100) requirement.is_active = bool(is_active) requirement.updated_by_user_id = getattr(user, "id", None) db.flush() return requirement def list_task_document_templates(db: Session, *, tenant_id: int, firm_task_template_id: int) -> list[FirmTaskDocumentTemplate]: return db.execute( select(FirmTaskDocumentTemplate) .where( FirmTaskDocumentTemplate.tenant_id == int(tenant_id), FirmTaskDocumentTemplate.firm_task_template_id == int(firm_task_template_id), FirmTaskDocumentTemplate.is_active.is_(True), ) .order_by(FirmTaskDocumentTemplate.uploaded_at_utc.desc(), FirmTaskDocumentTemplate.id.desc()) ).scalars().all() def _template_relative_path(task_template: FirmServiceTaskTemplate, original_filename: str, template_id: int) -> Path: suffix = Path(original_filename or "template.bin").suffix or ".bin" safe_name = sanitize_segment(Path(original_filename or "template.bin").stem, "template")[:80] return ( Path(f"tenant_{task_template.tenant_id}") / f"service_{task_template.service_catalogue_id}" / f"task_{task_template.id}" / f"TPL{template_id:06d}_{safe_name}_{uuid4().hex[:8]}{suffix}" ) def save_task_document_template( db: Session, *, task_template: FirmServiceTaskTemplate, template_name: str, template_category: str | None, description: str | None, upload_file, user, ) -> FirmTaskDocumentTemplate: original_filename = Path(upload_file.filename or "template.bin").name row = FirmTaskDocumentTemplate( tenant_id=task_template.tenant_id, service_catalogue_id=task_template.service_catalogue_id, firm_task_template_id=task_template.id, template_name=(template_name or original_filename).strip()[:200], template_category=(template_category or "").strip()[:100] or None, description=(description or "").strip() or None, original_filename=original_filename, stored_filename="PENDING", content_type=getattr(upload_file, "content_type", None), file_size_bytes=0, local_relative_path="PENDING", uploaded_by_user_id=getattr(user, "id", None), ) db.add(row) db.flush() rel_path = _template_relative_path(task_template, original_filename, row.id) abs_path = TEMPLATE_UPLOAD_ROOT / rel_path abs_path.parent.mkdir(parents=True, exist_ok=True) total = 0 with abs_path.open("wb") as out: while True: chunk = upload_file.file.read(1024 * 1024) if not chunk: break total += len(chunk) out.write(chunk) row.stored_filename = abs_path.name row.file_size_bytes = total row.local_relative_path = str(rel_path).replace("\\", "/") db.flush() return row def template_absolute_path(template: FirmTaskDocumentTemplate) -> Path: return TEMPLATE_UPLOAD_ROOT / (template.local_relative_path or "") def get_task_with_subscription(db: Session, task_id: int) -> ClientServiceTaskInstance | None: return db.execute( select(ClientServiceTaskInstance) .options( joinedload(ClientServiceTaskInstance.subscription).joinedload(ClientServiceSubscription.client), joinedload(ClientServiceTaskInstance.subscription).joinedload(ClientServiceSubscription.catalogue), joinedload(ClientServiceTaskInstance.template), ) .where(ClientServiceTaskInstance.id == int(task_id)) ).unique().scalar_one_or_none() def list_documents_for_task(db: Session, task_id: int) -> list[EngagementDocument]: return db.execute( select(EngagementDocument) .options(joinedload(EngagementDocument.versions), joinedload(EngagementDocument.document_requirement)) .where(EngagementDocument.task_instance_id == int(task_id), EngagementDocument.is_deleted.is_(False)) .order_by(EngagementDocument.updated_at_utc.desc(), EngagementDocument.id.desc()) ).unique().scalars().all() def requirement_upload_status(requirements: list[FirmTaskDocumentRequirement], documents: list[EngagementDocument]) -> list[dict]: by_req: dict[int, list[EngagementDocument]] = {} for doc in documents: if doc.document_requirement_id: by_req.setdefault(int(doc.document_requirement_id), []).append(doc) payload = [] for req in requirements: docs = by_req.get(int(req.id), []) payload.append({ "requirement": req, "documents": docs, "is_uploaded": bool(docs), "is_pending_mandatory": bool(req.is_mandatory and not docs), }) return payload def save_uploaded_task_document( db: Session, *, task: ClientServiceTaskInstance, requirement: FirmTaskDocumentRequirement | None, upload_file, title: str, document_type: str, description: str | None, remarks: str | None, user, existing_document_id: int | None = None, ) -> EngagementDocument: engagement = task.subscription or db.get(ClientServiceSubscription, task.subscription_id) if engagement is None: raise ValueError("Task is not linked to a valid engagement.") if requirement: title = title or requirement.document_name document_type = requirement.document_type or document_type description = description or requirement.instructions doc = save_uploaded_revision( db, engagement=engagement, upload_file=upload_file, title=title, document_type=document_type, description=description, remarks=remarks, user=user, existing_document_id=existing_document_id, ) doc.task_instance_id = task.id doc.document_requirement_id = requirement.id if requirement else None db.flush() return doc