from __future__ import annotations import re from dataclasses import dataclass from datetime import datetime, timezone from sqlalchemy import func, select from sqlalchemy.orm import Session from app.modules.services.models import ( ClientServiceSubscription, FirmServiceSelection, FirmServiceTaskTemplate, FirmTaskDocumentRequirement, FirmTaskDocumentTemplate, ServiceCatalogue, ServiceDefaultTaskTemplate, ServiceDueDateExtension, ServiceDueDateRule, ) @dataclass(frozen=True) class MergePreview: source: ServiceCatalogue target: ServiceCatalogue counts: dict[str, int] conflicts: dict[str, int] def _normalise(value: str | None) -> str: return re.sub(r"[^a-z0-9]+", "", (value or "").strip().lower()) def _count(db: Session, model, condition) -> int: return int(db.scalar(select(func.count(model.id)).where(condition)) or 0) def build_merge_preview(db: Session, *, source_id: int, target_id: int) -> MergePreview: if source_id == target_id: raise ValueError("Source and target services must be different.") source = db.get(ServiceCatalogue, source_id) target = db.get(ServiceCatalogue, target_id) if not source or not target: raise ValueError("The selected source or target service no longer exists.") counts = { "firm_selections": _count(db, FirmServiceSelection, FirmServiceSelection.service_catalogue_id == source_id), "system_default_tasks": _count(db, ServiceDefaultTaskTemplate, ServiceDefaultTaskTemplate.service_catalogue_id == source_id), "firm_task_templates": _count(db, FirmServiceTaskTemplate, FirmServiceTaskTemplate.service_catalogue_id == source_id), "document_requirements": _count(db, FirmTaskDocumentRequirement, FirmTaskDocumentRequirement.service_catalogue_id == source_id), "document_templates": _count(db, FirmTaskDocumentTemplate, FirmTaskDocumentTemplate.service_catalogue_id == source_id), "due_date_rules": _count(db, ServiceDueDateRule, ServiceDueDateRule.service_catalogue_id == source_id), "due_date_extensions": _count(db, ServiceDueDateExtension, ServiceDueDateExtension.service_catalogue_id == source_id), "client_subscriptions_preserved": _count(db, ClientServiceSubscription, ClientServiceSubscription.service_catalogue_id == source_id), } source_tenants = set(db.scalars(select(FirmServiceSelection.tenant_id).where(FirmServiceSelection.service_catalogue_id == source_id)).all()) target_tenants = set(db.scalars(select(FirmServiceSelection.tenant_id).where(FirmServiceSelection.service_catalogue_id == target_id)).all()) source_default_names = { _normalise(name) for name in db.scalars(select(ServiceDefaultTaskTemplate.task_name).where(ServiceDefaultTaskTemplate.service_catalogue_id == source_id)).all() } target_default_names = { _normalise(name) for name in db.scalars(select(ServiceDefaultTaskTemplate.task_name).where(ServiceDefaultTaskTemplate.service_catalogue_id == target_id)).all() } source_rules = { _normalise(name) for name in db.scalars(select(ServiceDueDateRule.rule_name).where(ServiceDueDateRule.service_catalogue_id == source_id)).all() } target_rules = { _normalise(name) for name in db.scalars(select(ServiceDueDateRule.rule_name).where(ServiceDueDateRule.service_catalogue_id == target_id)).all() } firm_task_conflicts = 0 source_firm_tasks = db.scalars(select(FirmServiceTaskTemplate).where(FirmServiceTaskTemplate.service_catalogue_id == source_id)).all() target_firm_tasks = db.scalars(select(FirmServiceTaskTemplate).where(FirmServiceTaskTemplate.service_catalogue_id == target_id)).all() target_task_keys = {(row.tenant_id, _normalise(row.task_name)) for row in target_firm_tasks} for row in source_firm_tasks: if (row.tenant_id, _normalise(row.task_name)) in target_task_keys: firm_task_conflicts += 1 conflicts = { "firm_selection_conflicts": len(source_tenants & target_tenants), "system_default_task_name_conflicts": len(source_default_names & target_default_names), "firm_task_name_conflicts": firm_task_conflicts, "due_rule_name_conflicts": len(source_rules & target_rules), } return MergePreview(source=source, target=target, counts=counts, conflicts=conflicts) def _merge_firm_selections(db: Session, *, source_id: int, target_id: int, actor_user_id: int) -> int: """Create/consolidate target selections while retaining source selections for history.""" processed = 0 source_rows = db.scalars( select(FirmServiceSelection) .where(FirmServiceSelection.service_catalogue_id == source_id) .order_by(FirmServiceSelection.id.asc()) ).all() for source_row in source_rows: target_row = db.scalar( select(FirmServiceSelection).where( FirmServiceSelection.tenant_id == source_row.tenant_id, FirmServiceSelection.service_catalogue_id == target_id, ) ) if target_row: target_row.is_enabled = bool(target_row.is_enabled or source_row.is_enabled) if target_row.default_branch_id is None and source_row.default_branch_id is not None: target_row.default_branch_id = source_row.default_branch_id target_row.updated_by_user_id = actor_user_id target_row.updated_at_utc = datetime.now(timezone.utc) else: db.add( FirmServiceSelection( tenant_id=source_row.tenant_id, service_catalogue_id=target_id, is_enabled=bool(source_row.is_enabled), default_branch_id=source_row.default_branch_id, activated_by_user_id=source_row.activated_by_user_id or actor_user_id, updated_by_user_id=actor_user_id, ) ) source_row.is_enabled = False source_row.updated_by_user_id = actor_user_id source_row.updated_at_utc = datetime.now(timezone.utc) processed += 1 db.flush() return processed def _merge_default_tasks(db: Session, *, source_id: int, target_id: int) -> tuple[int, int]: target_rows = db.scalars( select(ServiceDefaultTaskTemplate) .where(ServiceDefaultTaskTemplate.service_catalogue_id == target_id) .order_by(ServiceDefaultTaskTemplate.sequence_no.asc(), ServiceDefaultTaskTemplate.id.asc()) ).all() target_by_name = {_normalise(row.task_name): row for row in target_rows} next_sequence = max((row.sequence_no for row in target_rows), default=0) + 1 moved = 0 skipped = 0 source_rows = db.scalars( select(ServiceDefaultTaskTemplate) .where(ServiceDefaultTaskTemplate.service_catalogue_id == source_id) .order_by(ServiceDefaultTaskTemplate.sequence_no.asc(), ServiceDefaultTaskTemplate.id.asc()) ).all() for row in source_rows: key = _normalise(row.task_name) if key and key in target_by_name: skipped += 1 db.delete(row) continue row.service_catalogue_id = target_id row.sequence_no = next_sequence next_sequence += 1 target_by_name[key] = row moved += 1 db.flush() return moved, skipped def _merge_document_requirements(db: Session, *, source_template: FirmServiceTaskTemplate, target_template: FirmServiceTaskTemplate, target_id: int) -> None: target_names = { _normalise(row.document_name) for row in db.scalars( select(FirmTaskDocumentRequirement).where(FirmTaskDocumentRequirement.firm_task_template_id == target_template.id) ).all() } for row in db.scalars( select(FirmTaskDocumentRequirement).where(FirmTaskDocumentRequirement.firm_task_template_id == source_template.id) ).all(): key = _normalise(row.document_name) if key and key in target_names: db.delete(row) continue row.firm_task_template_id = target_template.id row.service_catalogue_id = target_id target_names.add(key) for row in db.scalars( select(FirmTaskDocumentTemplate).where(FirmTaskDocumentTemplate.firm_task_template_id == source_template.id) ).all(): row.firm_task_template_id = target_template.id row.service_catalogue_id = target_id def _merge_firm_tasks(db: Session, *, source_id: int, target_id: int) -> tuple[int, int]: target_rows = db.scalars( select(FirmServiceTaskTemplate) .where(FirmServiceTaskTemplate.service_catalogue_id == target_id) .order_by(FirmServiceTaskTemplate.tenant_id.asc(), FirmServiceTaskTemplate.sequence_no.asc()) ).all() target_by_key = {(row.tenant_id, _normalise(row.task_name)): row for row in target_rows} max_sequence_by_tenant: dict[int, int] = {} for row in target_rows: max_sequence_by_tenant[row.tenant_id] = max(max_sequence_by_tenant.get(row.tenant_id, 0), row.sequence_no) moved = 0 merged = 0 source_rows = db.scalars( select(FirmServiceTaskTemplate) .where(FirmServiceTaskTemplate.service_catalogue_id == source_id) .order_by(FirmServiceTaskTemplate.tenant_id.asc(), FirmServiceTaskTemplate.sequence_no.asc(), FirmServiceTaskTemplate.id.asc()) ).all() for row in source_rows: key = (row.tenant_id, _normalise(row.task_name)) target_row = target_by_key.get(key) if target_row: _merge_document_requirements(db, source_template=row, target_template=target_row, target_id=target_id) db.flush() db.delete(row) merged += 1 continue next_sequence = max_sequence_by_tenant.get(row.tenant_id, 0) + 1 max_sequence_by_tenant[row.tenant_id] = next_sequence row.service_catalogue_id = target_id row.sequence_no = next_sequence for requirement in row.document_requirements: requirement.service_catalogue_id = target_id for template in row.document_templates: template.service_catalogue_id = target_id target_by_key[key] = row moved += 1 db.flush() return moved, merged def _merge_due_rules(db: Session, *, source_id: int, target_id: int) -> tuple[int, int]: target_rows = db.scalars(select(ServiceDueDateRule).where(ServiceDueDateRule.service_catalogue_id == target_id)).all() target_by_name = {_normalise(row.rule_name): row for row in target_rows} moved = 0 merged = 0 source_rows = db.scalars(select(ServiceDueDateRule).where(ServiceDueDateRule.service_catalogue_id == source_id)).all() for row in source_rows: target_row = target_by_name.get(_normalise(row.rule_name)) if target_row: db.query(ServiceDueDateExtension).filter(ServiceDueDateExtension.due_date_rule_id == row.id).update( {ServiceDueDateExtension.due_date_rule_id: target_row.id}, synchronize_session=False ) db.query(ClientServiceSubscription).filter(ClientServiceSubscription.due_date_rule_id == row.id).update( {ClientServiceSubscription.due_date_rule_id: target_row.id}, synchronize_session=False ) db.delete(row) merged += 1 else: row.service_catalogue_id = target_id target_by_name[_normalise(row.rule_name)] = row moved += 1 db.query(ServiceDueDateExtension).filter(ServiceDueDateExtension.service_catalogue_id == source_id).update( {ServiceDueDateExtension.service_catalogue_id: target_id}, synchronize_session=False ) db.flush() return moved, merged def merge_service_catalogues( db: Session, *, source_id: int, target_id: int, actor_user_id: int, ) -> dict[str, int | str]: """Safely consolidate reusable configuration into the target service. Existing client subscriptions, execution tasks, invoices and consultant requests are intentionally retained against the source service to preserve historical records. The source service is disabled after configuration is moved, so all new work uses the target service. """ preview = build_merge_preview(db, source_id=source_id, target_id=target_id) source = preview.source target = preview.target selections_processed = _merge_firm_selections( db, source_id=source_id, target_id=target_id, actor_user_id=actor_user_id, ) default_moved, default_skipped = _merge_default_tasks(db, source_id=source_id, target_id=target_id) firm_moved, firm_merged = _merge_firm_tasks(db, source_id=source_id, target_id=target_id) rules_moved, rules_merged = _merge_due_rules(db, source_id=source_id, target_id=target_id) source.is_active = False source.is_client_requestable = False source.is_consultant_requestable = False source.updated_by_user_id = actor_user_id source.updated_at_utc = datetime.now(timezone.utc) target.is_active = True target.updated_by_user_id = actor_user_id target.updated_at_utc = datetime.now(timezone.utc) db.flush() return { "source_code": source.service_code, "target_code": target.service_code, "firm_selections_processed": selections_processed, "default_tasks_moved": default_moved, "default_tasks_deduplicated": default_skipped, "firm_tasks_moved": firm_moved, "firm_tasks_merged": firm_merged, "due_rules_moved": rules_moved, "due_rules_merged": rules_merged, "historical_subscriptions_preserved": preview.counts["client_subscriptions_preserved"], }