from __future__ import annotations from dataclasses import dataclass, field from datetime import datetime, timezone import hashlib import json import re from typing import Any from sqlalchemy import func, select from sqlalchemy.orm import Session from app.modules.services.models import ( ClientServiceTaskInstance, FirmServiceSelection, FirmServiceTaskTemplate, ServiceDefaultTaskTemplate, ServiceTaskCategory, ) # Fields that define the centrally maintained system-default task snapshot. # task_category_id is deliberately excluded because system and firm category # masters use different scoped rows; task_category text is the portable value. SYSTEM_TASK_FIELDS: tuple[str, ...] = ( "task_name", "description", "sequence_no", "default_role_name", "eligible_role_names", "is_mandatory", "requires_review", "normal_review_role", "task_category", "response_required", "response_type", "evidence_required", "remarks_required_if_no", "task_tool_code", "is_aqmm_task", "aqmm_mandatory", "aqmm_evidence_required", "aqmm_manager_review_required", "aqmm_partner_review_required", "aqmm_review_partner_required", "aqmm_blocks_final_release", "aqmm_reference", "is_active", ) _SPACE_RE = re.compile(r"\s+") @dataclass class FirmDefaultTaskSyncResult: created: int = 0 updated: int = 0 duplicates_disabled: int = 0 unchanged: int = 0 custom_updates_available: int = 0 linked_legacy: int = 0 engagement_created: int = 0 engagement_updated_pending: int = 0 engagement_deactivated_pending: int = 0 engagement_preserved_history: int = 0 @property def active_total_change(self) -> int: return self.created - self.duplicates_disabled @dataclass class SystemDefaultRolloutResult: firms_processed: int = 0 firms_changed: int = 0 created: int = 0 updated: int = 0 unchanged: int = 0 custom_updates_available: int = 0 duplicates_disabled: int = 0 engagement_created: int = 0 engagement_updated_pending: int = 0 engagement_deactivated_pending: int = 0 engagement_preserved_history: int = 0 tenant_results: dict[int, FirmDefaultTaskSyncResult] = field(default_factory=dict) def _normalise_name(value: str | None) -> str: return _SPACE_RE.sub(" ", (value or "").strip()).casefold() def _portable_value(value: Any) -> Any: if isinstance(value, (str, int, float, bool)) or value is None: return value return str(value) def _snapshot_dict(row: Any) -> dict[str, Any]: return {name: _portable_value(getattr(row, name, None)) for name in SYSTEM_TASK_FIELDS} def system_task_hash(source: ServiceDefaultTaskTemplate) -> str: payload = json.dumps(_snapshot_dict(source), sort_keys=True, separators=(",", ":"), ensure_ascii=False) return hashlib.sha256(payload.encode("utf-8")).hexdigest() def firm_task_hash(target: FirmServiceTaskTemplate) -> str: payload = json.dumps(_snapshot_dict(target), sort_keys=True, separators=(",", ":"), ensure_ascii=False) return hashlib.sha256(payload.encode("utf-8")).hexdigest() def _legacy_retirement_equivalent( source: ServiceDefaultTaskTemplate, target: FirmServiceTaskTemplate, ) -> bool: """Return True when a legacy firm row differs only because the system row was retired. Full-sync retirement moves omitted system rows into a high temporary sequence range before setting is_active=False. Legacy firm rows created before provenance tracking still carry their old active flag and old sequence. Those two differences must not make an otherwise inherited task look like a deliberate firm customization. Any substantive field difference still protects the firm row as customized. """ source_snapshot = _snapshot_dict(source) target_snapshot = _snapshot_dict(target) source_snapshot.pop("is_active", None) target_snapshot.pop("is_active", None) if int(getattr(source, "sequence_no", 0) or 0) >= 100000: source_snapshot.pop("sequence_no", None) target_snapshot.pop("sequence_no", None) return source_snapshot == target_snapshot def system_task_diff(source: ServiceDefaultTaskTemplate, target: FirmServiceTaskTemplate) -> list[dict[str, Any]]: labels = { "task_name": "Task name", "description": "Description", "sequence_no": "Sequence", "default_role_name": "Default role", "eligible_role_names": "Eligible roles", "is_mandatory": "Mandatory", "requires_review": "Review required", "normal_review_role": "Normal reviewer", "task_category": "Task category", "response_required": "Response required", "response_type": "Response type", "evidence_required": "Evidence required", "remarks_required_if_no": "Remarks required if No", "task_tool_code": "Task tool", "is_aqmm_task": "AQMM task", "aqmm_mandatory": "AQMM mandatory", "aqmm_evidence_required": "AQMM evidence", "aqmm_manager_review_required": "AQMM manager review", "aqmm_partner_review_required": "AQMM partner review", "aqmm_review_partner_required": "AQMM review partner", "aqmm_blocks_final_release": "Blocks final release", "aqmm_reference": "AQMM reference", "is_active": "Active", } changes: list[dict[str, Any]] = [] for name in SYSTEM_TASK_FIELDS: old = getattr(target, name, None) new = getattr(source, name, None) if old != new: changes.append({"field": name, "label": labels.get(name, name), "firm": old, "system": new}) return changes def mark_firm_task_customized(task: FirmServiceTaskTemplate) -> None: """Mark an explicitly edited firm task as protected from automatic overwrite.""" task.is_customized = True # Do not clear an already pending system update. If there is no pending update, # the next system revision/hash change will create one automatically. def _ensure_firm_category( db: Session, *, source: ServiceDefaultTaskTemplate, tenant_id: int, user_id: int | None, ) -> ServiceTaskCategory | None: name = (getattr(source, "task_category", None) or "").strip() if not name: return None system_category = None source_category_id = getattr(source, "task_category_id", None) if source_category_id: system_category = db.get(ServiceTaskCategory, int(source_category_id)) code = ((getattr(system_category, "code", None) or "").strip().upper() if system_category else "") if code: existing = db.execute( select(ServiceTaskCategory).where( ServiceTaskCategory.tenant_id == tenant_id, ServiceTaskCategory.service_catalogue_id == source.service_catalogue_id, ServiceTaskCategory.code == code, ) ).scalar_one_or_none() else: existing = db.execute( select(ServiceTaskCategory).where( ServiceTaskCategory.tenant_id == tenant_id, ServiceTaskCategory.service_catalogue_id == source.service_catalogue_id, func.lower(ServiceTaskCategory.name) == name.lower(), ) ).scalar_one_or_none() if existing: # Central category rename/status/order should remain aligned for inherited use. # Firm category rows are shared by firm templates, so update only metadata that # does not destroy task history. if system_category: existing.name = system_category.name existing.sort_order = system_category.sort_order existing.is_active = system_category.is_active if user_id is not None: existing.updated_by_user_id = user_id return existing if not code: # Match the application's existing category-code convention sufficiently for # automatic inherited categories without importing services.py (avoids cycles). code = re.sub(r"[^A-Z0-9]+", "_", name.upper()).strip("_")[:50] or "CATEGORY" base = code suffix = 2 while db.execute( select(ServiceTaskCategory.id).where( ServiceTaskCategory.tenant_id == tenant_id, ServiceTaskCategory.service_catalogue_id == source.service_catalogue_id, ServiceTaskCategory.code == code, ) ).first(): code = f"{base[:45]}_{suffix}" suffix += 1 row = ServiceTaskCategory( tenant_id=tenant_id, service_catalogue_id=source.service_catalogue_id, code=code, name=(system_category.name if system_category else name), sort_order=(system_category.sort_order if system_category else 100), is_active=(system_category.is_active if system_category else True), created_by_user_id=user_id, updated_by_user_id=user_id, ) db.add(row) db.flush() return row def _copy_default_columns( db: Session, source: ServiceDefaultTaskTemplate, target: FirmServiceTaskTemplate, *, tenant_id: int, user_id: int | None, allow_sequence_change: bool = True, ) -> bool: changed = False for name in SYSTEM_TASK_FIELDS: if name == "sequence_no" and not allow_sequence_change: continue new_value = getattr(source, name, None) if getattr(target, name, None) != new_value: setattr(target, name, new_value) changed = True category = _ensure_firm_category(db, source=source, tenant_id=tenant_id, user_id=user_id) category_id = category.id if category else None category_name = category.name if category else None if getattr(target, "task_category_id", None) != category_id: target.task_category_id = category_id changed = True if getattr(target, "task_category", None) != category_name: target.task_category = category_name changed = True return changed def _reference_counts(db: Session, *, firm_task_ids: list[int]) -> dict[int, int]: if not firm_task_ids: return {} rows = db.execute( select(ClientServiceTaskInstance.firm_task_template_id, func.count(ClientServiceTaskInstance.id)) .where(ClientServiceTaskInstance.firm_task_template_id.in_(firm_task_ids)) .group_by(ClientServiceTaskInstance.firm_task_template_id) ).all() return {int(template_id): int(count) for template_id, count in rows if template_id is not None} def _choose_canonical(candidates: list[FirmServiceTaskTemplate], *, reference_counts: dict[int, int]) -> FirmServiceTaskTemplate: return sorted(candidates, key=lambda row: (-reference_counts.get(int(row.id), 0), int(row.id)))[0] def _next_free_sequence(used: set[int], preferred: int) -> int: if preferred > 0 and preferred not in used: return preferred candidate = max(used or {0}) + 1 while candidate in used: candidate += 1 return candidate def sync_firm_tasks_from_system_defaults( db: Session, *, tenant_id: int, service_catalogue_id: int, updated_by_user_id: int | None = None, sync_open_engagements: bool = False, ) -> FirmDefaultTaskSyncResult: """Synchronise one firm's task templates from centrally maintained defaults. Inherited tasks are updated automatically. Explicitly customized linked tasks are never overwritten; a system_update_available flag is raised for Firm Admin review. Firm-only tasks remain untouched. Legacy rows are linked by exact normalized name. """ defaults = db.execute( select(ServiceDefaultTaskTemplate) .where(ServiceDefaultTaskTemplate.service_catalogue_id == service_catalogue_id) .order_by(ServiceDefaultTaskTemplate.sequence_no.asc(), ServiceDefaultTaskTemplate.id.asc()) ).scalars().all() firm_rows = db.execute( select(FirmServiceTaskTemplate) .where( FirmServiceTaskTemplate.tenant_id == tenant_id, FirmServiceTaskTemplate.service_catalogue_id == service_catalogue_id, ) .order_by(FirmServiceTaskTemplate.sequence_no.asc(), FirmServiceTaskTemplate.id.asc()) ).scalars().all() result = FirmDefaultTaskSyncResult() reference_counts = _reference_counts(db, firm_task_ids=[int(r.id) for r in firm_rows if r.id]) by_source: dict[int, list[FirmServiceTaskTemplate]] = {} legacy_by_name: dict[str, list[FirmServiceTaskTemplate]] = {} for row in firm_rows: if row.source_system_task_id: by_source.setdefault(int(row.source_system_task_id), []).append(row) else: key = _normalise_name(row.task_name) if key: legacy_by_name.setdefault(key, []).append(row) used_sequences = {int(r.sequence_no) for r in firm_rows if r.sequence_no is not None} for source in defaults: latest_hash = system_task_hash(source) candidates = by_source.get(int(source.id), []) linked_from_legacy = False if not candidates: legacy_candidates = legacy_by_name.get(_normalise_name(source.task_name), []) if not legacy_candidates: # One-time migration fallback for a firm that renamed a previously # copied default before provenance columns existed. Sequence is used # only when there is exactly one unlinked candidate at that position. legacy_candidates = [ row for row in firm_rows if row.source_system_task_id is None and int(row.sequence_no or 0) == int(source.sequence_no or 0) ] if len(legacy_candidates) != 1: legacy_candidates = [] if legacy_candidates: target = _choose_canonical(legacy_candidates, reference_counts=reference_counts) candidates = [target] target.source_system_task_id = source.id linked_from_legacy = True result.linked_legacy += 1 by_source.setdefault(int(source.id), []).append(target) # Remove from future legacy matching. key = _normalise_name(target.task_name) if key in legacy_by_name: legacy_by_name[key] = [r for r in legacy_by_name[key] if r.id != target.id] if candidates: target = _choose_canonical(candidates, reference_counts=reference_counts) for duplicate in candidates: if duplicate.id == target.id: continue if duplicate.is_active: duplicate.is_active = False if updated_by_user_id is not None: duplicate.updated_by_user_id = updated_by_user_id result.duplicates_disabled += 1 # For legacy rows with no inheritance history, infer whether they were # already firm-customized by comparing the full portable snapshot. if linked_from_legacy and target.last_synced_system_hash is None: if not source.is_active and _legacy_retirement_equivalent(source, target): # This was an inherited legacy row whose system source has just been # retired by Full Synchronization. Do not misclassify the active/old # sequence difference as a firm customization; allow normal retirement # propagation below. target.is_customized = False else: target.is_customized = firm_task_hash(target) != latest_hash if not target.is_customized: target.last_synced_system_hash = latest_hash target.last_reviewed_system_hash = latest_hash # Repair rows linked by an earlier rollout where retirement-only differences # were incorrectly interpreted as customization. This is deliberately limited # to rows that have never had a successful inherited snapshot recorded. if ( not source.is_active and target.is_customized and target.last_synced_system_hash is None and _legacy_retirement_equivalent(source, target) ): target.is_customized = False target.system_update_available = False target.system_update_detected_at_utc = None # Detect out-of-band edits to a linked inherited row as customization. if ( not target.is_customized and target.last_synced_system_hash and firm_task_hash(target) != target.last_synced_system_hash ): target.is_customized = True if target.is_customized: pending = target.last_reviewed_system_hash != latest_hash target.system_update_available = pending target.system_update_detected_at_utc = datetime.now(timezone.utc) if pending else None if pending: result.custom_updates_available += 1 else: result.unchanged += 1 continue # Inherited task: copy system values. Sequence changes are applied when # the desired number is free or currently owned by this same row. If a # firm-only/custom row owns it, preserve that row and place this inherited # task at the next free sequence instead of overwriting customization. desired_seq = int(source.sequence_no or 0) current_seq = int(target.sequence_no or 0) allow_seq = desired_seq == current_seq or desired_seq not in (used_sequences - {current_seq}) if allow_seq: used_sequences.discard(current_seq) used_sequences.add(desired_seq) changed = _copy_default_columns( db, source, target, tenant_id=tenant_id, user_id=updated_by_user_id, allow_sequence_change=allow_seq, ) target.source_system_task_id = source.id target.last_synced_system_hash = latest_hash target.last_reviewed_system_hash = latest_hash target.system_update_available = False target.system_update_detected_at_utc = None if updated_by_user_id is not None: target.updated_by_user_id = updated_by_user_id result.updated += 1 if changed else 0 result.unchanged += 0 if changed else 1 continue # Missing default: create a new inherited row. Avoid colliding with a # firm-only sequence; identity is source_system_task_id, not sequence number. preferred = int(source.sequence_no or 0) seq = _next_free_sequence(used_sequences, preferred) used_sequences.add(seq) row = FirmServiceTaskTemplate( tenant_id=tenant_id, service_catalogue_id=service_catalogue_id, task_name=source.task_name, sequence_no=seq, source_system_task_id=source.id, is_customized=False, last_synced_system_hash=latest_hash, last_reviewed_system_hash=latest_hash, system_update_available=False, created_by_user_id=updated_by_user_id, updated_by_user_id=updated_by_user_id, ) db.add(row) db.flush() _copy_default_columns( db, source, row, tenant_id=tenant_id, user_id=updated_by_user_id, allow_sequence_change=(seq == preferred), ) result.created += 1 firm_rows.append(row) by_source.setdefault(int(source.id), []).append(row) db.flush() if sync_open_engagements: from app.modules.services.execution import sync_open_engagement_tasks_for_service engagement = sync_open_engagement_tasks_for_service( db, tenant_id=tenant_id, catalogue_id=service_catalogue_id, user_id=updated_by_user_id or 0, include_started_open_tasks=False, safe_system_rollout=True, ) result.engagement_created = engagement.get("created", 0) result.engagement_updated_pending = engagement.get("updated_pending", 0) result.engagement_deactivated_pending = engagement.get("deactivated_pending", 0) result.engagement_preserved_history = engagement.get("preserved_history", 0) return result def sync_system_defaults_to_all_firms( db: Session, *, service_catalogue_id: int, updated_by_user_id: int | None = None, sync_open_engagements: bool = True, ) -> SystemDefaultRolloutResult: """Roll a system-default service checklist to every firm that enabled it.""" selections = db.execute( select(FirmServiceSelection).where( FirmServiceSelection.service_catalogue_id == service_catalogue_id, FirmServiceSelection.is_enabled.is_(True), ) ).scalars().all() aggregate = SystemDefaultRolloutResult() for selection in selections: result = sync_firm_tasks_from_system_defaults( db, tenant_id=int(selection.tenant_id), service_catalogue_id=service_catalogue_id, updated_by_user_id=updated_by_user_id, sync_open_engagements=sync_open_engagements, ) aggregate.firms_processed += 1 if result.created or result.updated or result.duplicates_disabled or result.custom_updates_available: aggregate.firms_changed += 1 aggregate.created += result.created aggregate.updated += result.updated aggregate.unchanged += result.unchanged aggregate.custom_updates_available += result.custom_updates_available aggregate.duplicates_disabled += result.duplicates_disabled aggregate.engagement_created += result.engagement_created aggregate.engagement_updated_pending += result.engagement_updated_pending aggregate.engagement_deactivated_pending += result.engagement_deactivated_pending aggregate.engagement_preserved_history += result.engagement_preserved_history aggregate.tenant_results[int(selection.tenant_id)] = result return aggregate def accept_system_update_for_firm_task( db: Session, *, task: FirmServiceTaskTemplate, user_id: int, sync_open_engagements: bool = True, ) -> FirmDefaultTaskSyncResult: if not task.source_system_task_id: raise ValueError("This firm task is not linked to a system default.") source = db.get(ServiceDefaultTaskTemplate, int(task.source_system_task_id)) if source is None: raise ValueError("The linked system default no longer exists.") latest_hash = system_task_hash(source) used = { int(v) for v in db.scalars( select(FirmServiceTaskTemplate.sequence_no).where( FirmServiceTaskTemplate.tenant_id == task.tenant_id, FirmServiceTaskTemplate.service_catalogue_id == task.service_catalogue_id, FirmServiceTaskTemplate.id != task.id, ) ).all() if v is not None } desired = int(source.sequence_no or 0) allow_seq = desired not in used changed = _copy_default_columns( db, source, task, tenant_id=int(task.tenant_id), user_id=user_id, allow_sequence_change=allow_seq, ) task.is_customized = False task.last_synced_system_hash = latest_hash task.last_reviewed_system_hash = latest_hash task.system_update_available = False task.system_update_detected_at_utc = None task.updated_by_user_id = user_id db.flush() result = FirmDefaultTaskSyncResult(updated=1 if changed else 0, unchanged=0 if changed else 1) if sync_open_engagements: from app.modules.services.execution import sync_open_engagement_tasks_for_service engagement = sync_open_engagement_tasks_for_service( db, tenant_id=int(task.tenant_id), catalogue_id=int(task.service_catalogue_id), user_id=user_id, include_started_open_tasks=False, safe_system_rollout=True, ) result.engagement_created = engagement.get("created", 0) result.engagement_updated_pending = engagement.get("updated_pending", 0) result.engagement_deactivated_pending = engagement.get("deactivated_pending", 0) result.engagement_preserved_history = engagement.get("preserved_history", 0) return result def keep_firm_customization_for_system_revision( db: Session, *, task: FirmServiceTaskTemplate, user_id: int, ) -> None: if not task.source_system_task_id: raise ValueError("This firm task is not linked to a system default.") source = db.get(ServiceDefaultTaskTemplate, int(task.source_system_task_id)) if source is None: raise ValueError("The linked system default no longer exists.") task.is_customized = True task.last_reviewed_system_hash = system_task_hash(source) task.system_update_available = False task.system_update_detected_at_utc = None task.updated_by_user_id = user_id db.flush()