291 lines
9.0 KiB
Python
291 lines
9.0 KiB
Python
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
|