Fix duplicate firm task creation when syncing system defaults
This commit is contained in:
@@ -0,0 +1,290 @@
|
||||
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
|
||||
Reference in New Issue
Block a user