Files
2026-07-17 23:20:08 +05:30

319 lines
14 KiB
Python

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.")
if not source.is_active or not target.is_active:
raise ValueError("Only active services can be selected for a merge.")
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"],
}