from __future__ import annotations from collections import defaultdict from dataclasses import dataclass import re from typing import Any from sqlalchemy import func, select from sqlalchemy.inspection import inspect as sa_inspect from sqlalchemy.orm import Session from app.modules.services.models import ( ClientServiceTaskInstance, FirmServiceTaskTemplate, ServiceDefaultTaskTemplate, ) @dataclass class FirmDefaultTaskSyncResult: created: int = 0 updated: int = 0 duplicates_disabled: int = 0 unchanged: int = 0 @property def active_total_change(self) -> int: return self.created - self.duplicates_disabled _SPACE_RE = re.compile(r"\s+") def _normalise_name(value: str | None) -> str: return _SPACE_RE.sub(" ", (value or "").strip()).casefold() def _column_names(model: type[Any]) -> set[str]: return {column.key for column in sa_inspect(model).mapper.column_attrs} def _copy_default_columns( source: ServiceDefaultTaskTemplate, target: FirmServiceTaskTemplate, ) -> bool: """Copy only columns that exist on both system and firm task models. Firm-only identity/audit fields are deliberately excluded. Related firm requirement/template-file records are untouched. """ source_cols = _column_names(ServiceDefaultTaskTemplate) target_cols = _column_names(FirmServiceTaskTemplate) excluded = { "id", "tenant_id", "service_catalogue_id", "created_at", "created_at_utc", "created_by_user_id", "updated_at", "updated_at_utc", "updated_by_user_id", } changed = False for name in sorted((source_cols & target_cols) - excluded): new_value = getattr(source, name, None) if getattr(target, name, None) != new_value: setattr(target, name, new_value) changed = True # System defaults being synced are expected to be usable in the firm. if hasattr(target, "is_active") and getattr(target, "is_active", None) is not True: target.is_active = True 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: """Prefer the row already used by the most engagement task instances. This avoids breaking historical links. If usage is equal, keep the oldest row (smallest id). """ return sorted( candidates, key=lambda row: ( -reference_counts.get(int(row.id), 0), int(row.id), ), )[0] def sync_firm_tasks_from_system_defaults( db: Session, *, tenant_id: int, service_catalogue_id: int, updated_by_user_id: int | None = None, ) -> FirmDefaultTaskSyncResult: """Idempotently synchronise system default tasks into one firm's template set. Behaviour: * existing firm task with the same normalised task name -> UPDATE, never INSERT; * duplicate firm rows with the same task name -> retain one canonical active row, mark the additional template rows inactive; * missing system default -> create one firm task; * unique firm-only/custom tasks are preserved; * existing ClientServiceTaskInstance rows are never deleted or reassigned. The function may safely be called repeatedly. """ 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() by_name: dict[str, list[FirmServiceTaskTemplate]] = defaultdict(list) for row in firm_rows: key = _normalise_name(getattr(row, "task_name", None)) if key: by_name[key].append(row) reference_counts = _reference_counts( db, firm_task_ids=[int(row.id) for row in firm_rows if getattr(row, "id", None)], ) # First clean pre-existing duplicate groups, regardless of whether the task # remains in the current system defaults. A unique firm-only task is untouched. for candidates in by_name.values(): if len(candidates) <= 1: continue canonical = _choose_canonical( candidates, reference_counts=reference_counts, ) for duplicate in candidates: if duplicate.id == canonical.id: continue if getattr(duplicate, "is_active", True): duplicate.is_active = False if ( updated_by_user_id is not None and hasattr(duplicate, "updated_by_user_id") ): duplicate.updated_by_user_id = updated_by_user_id result.duplicates_disabled += 1 # Upsert current system defaults. for default in defaults: key = _normalise_name(getattr(default, "task_name", None)) candidates = by_name.get(key, []) if key else [] if candidates: target = _choose_canonical( candidates, reference_counts=reference_counts, ) changed = _copy_default_columns(default, target) if ( updated_by_user_id is not None and hasattr(target, "updated_by_user_id") ): target.updated_by_user_id = updated_by_user_id if changed: result.updated += 1 else: result.unchanged += 1 continue # No name match. A unique sequence match is a conservative fallback for a # system-default rename while still avoiding arbitrary replacement. sequence_no = getattr(default, "sequence_no", None) sequence_candidates = [ row for row in firm_rows if getattr(row, "sequence_no", None) == sequence_no and getattr(row, "is_active", True) ] if len(sequence_candidates) == 1: target = sequence_candidates[0] changed = _copy_default_columns(default, target) if ( updated_by_user_id is not None and hasattr(target, "updated_by_user_id") ): target.updated_by_user_id = updated_by_user_id if changed: result.updated += 1 else: result.unchanged += 1 # Make subsequent defaults see the new name. by_name[_normalise_name(getattr(target, "task_name", None))].append(target) continue # Missing task: construct from the intersection of mapped columns. source_cols = _column_names(ServiceDefaultTaskTemplate) target_cols = _column_names(FirmServiceTaskTemplate) excluded = { "id", "tenant_id", "service_catalogue_id", "created_at", "created_at_utc", "created_by_user_id", "updated_at", "updated_at_utc", "updated_by_user_id", } payload = { name: getattr(default, name, None) for name in sorted((source_cols & target_cols) - excluded) } payload["tenant_id"] = tenant_id payload["service_catalogue_id"] = service_catalogue_id if "is_active" in target_cols: payload["is_active"] = True if ( updated_by_user_id is not None and "created_by_user_id" in target_cols ): payload["created_by_user_id"] = updated_by_user_id if ( updated_by_user_id is not None and "updated_by_user_id" in target_cols ): payload["updated_by_user_id"] = updated_by_user_id new_row = FirmServiceTaskTemplate(**payload) db.add(new_row) db.flush() firm_rows.append(new_row) if key: by_name[key].append(new_row) result.created += 1 db.flush() return result