from __future__ import annotations import re from difflib import SequenceMatcher from sqlalchemy import select from sqlalchemy.orm import Session from app.modules.services.models import ( FirmServiceTaskTemplate, FirmTaskDocumentRequirement, FirmTaskDocumentTemplate, ServiceDefaultTaskTemplate, ) _STOP_WORDS = {"the", "a", "an", "and", "of", "for", "to", "with", "as", "per", "verify", "verification", "check", "checking"} def _normalise(name: str | None) -> str: words = re.findall(r"[a-z0-9]+", (name or "").lower()) reduced = [w for w in words if w not in _STOP_WORDS] return " ".join(reduced or words) def _similarity(a: str | None, b: str | None) -> float: na, nb = _normalise(a), _normalise(b) if not na or not nb: return 0.0 if na == nb: return 1.0 ta, tb = set(na.split()), set(nb.split()) token_score = len(ta & tb) / max(len(ta | tb), 1) sequence_score = SequenceMatcher(None, na, nb).ratio() return max(sequence_score, token_score) def find_duplicate_pairs(tasks, *, threshold: float = 0.78) -> list[dict]: active = [t for t in tasks if bool(getattr(t, "is_active", True))] pairs: list[dict] = [] for i, left in enumerate(active): for right in active[i + 1:]: score = _similarity(left.task_name, right.task_name) if score >= threshold: pairs.append({"left": left, "right": right, "score": round(score * 100, 1)}) return sorted(pairs, key=lambda row: (-row["score"], row["left"].sequence_no, row["right"].sequence_no)) def merge_default_tasks(db: Session, *, catalogue_id: int, master_task_id: int, source_task_ids: list[int]) -> int: master = db.execute(select(ServiceDefaultTaskTemplate).where( ServiceDefaultTaskTemplate.id == master_task_id, ServiceDefaultTaskTemplate.service_catalogue_id == catalogue_id, )).scalar_one_or_none() if master is None: raise ValueError("Master task not found.") source_ids = {int(x) for x in source_task_ids if int(x) != master_task_id} if not source_ids: raise ValueError("Select at least one duplicate task to merge.") sources = list(db.execute(select(ServiceDefaultTaskTemplate).where( ServiceDefaultTaskTemplate.service_catalogue_id == catalogue_id, ServiceDefaultTaskTemplate.id.in_(source_ids), )).scalars().all()) if len(sources) != len(source_ids): raise ValueError("One or more duplicate tasks were not found.") for source in sources: source.is_active = False db.flush() return len(sources) def merge_firm_tasks(db: Session, *, tenant_id: int, catalogue_id: int, master_task_id: int, source_task_ids: list[int], user_id: int) -> int: master = db.execute(select(FirmServiceTaskTemplate).where( FirmServiceTaskTemplate.id == master_task_id, FirmServiceTaskTemplate.tenant_id == tenant_id, FirmServiceTaskTemplate.service_catalogue_id == catalogue_id, )).scalar_one_or_none() if master is None: raise ValueError("Master task not found.") source_ids = {int(x) for x in source_task_ids if int(x) != master_task_id} if not source_ids: raise ValueError("Select at least one duplicate task to merge.") sources = list(db.execute(select(FirmServiceTaskTemplate).where( FirmServiceTaskTemplate.tenant_id == tenant_id, FirmServiceTaskTemplate.service_catalogue_id == catalogue_id, FirmServiceTaskTemplate.id.in_(source_ids), )).scalars().all()) if len(sources) != len(source_ids): raise ValueError("One or more duplicate tasks were not found.") existing_req_names = { (r.document_name or "").strip().lower() for r in db.execute(select(FirmTaskDocumentRequirement).where( FirmTaskDocumentRequirement.tenant_id == tenant_id, FirmTaskDocumentRequirement.firm_task_template_id == master.id, )).scalars().all() } for source in sources: for req in db.execute(select(FirmTaskDocumentRequirement).where( FirmTaskDocumentRequirement.tenant_id == tenant_id, FirmTaskDocumentRequirement.firm_task_template_id == source.id, )).scalars().all(): key = (req.document_name or "").strip().lower() if key and key not in existing_req_names: req.firm_task_template_id = master.id req.updated_by_user_id = user_id existing_req_names.add(key) for template in db.execute(select(FirmTaskDocumentTemplate).where( FirmTaskDocumentTemplate.tenant_id == tenant_id, FirmTaskDocumentTemplate.firm_task_template_id == source.id, )).scalars().all(): template.firm_task_template_id = master.id source.is_active = False source.updated_by_user_id = user_id db.flush() return len(sources)