from __future__ import annotations import re from datetime import date, datetime, timezone from typing import Iterable from sqlalchemy import func, select from sqlalchemy.orm import Session, selectinload from app.modules.services.models import ( ClientServiceSubscription, ServiceCatalogue, ServiceDueDateExtension, ServiceDueDateRule, ) DUE_PERIOD_TYPES = [ ("yearly", "Yearly / Annual"), ("monthly", "Monthly"), ("quarterly", "Quarterly"), ("one_time", "One Time"), ("event_based", "Event Based"), ("renewal_based", "Renewal Before Expiry"), ("custom", "Custom / Manual"), ] DUE_YEAR_BASIS_CHOICES = [ ("assessment_year_start", "Assessment year start year"), ("financial_year_start", "Financial year start year"), ("financial_year_end", "Financial year end year"), ("calendar_year", "Calendar year from period"), ] DUE_DATE_SOURCE_RULE = "rule" DUE_DATE_SOURCE_EXTENSION = "extension" DUE_DATE_SOURCE_MANUAL = "manual" def parse_optional_date(value: str | None) -> date | None: value = (value or "").strip() if not value: return None return date.fromisoformat(value) def _parse_year_pair(value: str | None) -> tuple[int, int] | None: value = (value or "").strip() match = re.match(r"^(\d{4})\s*-\s*(\d{2}|\d{4})$", value) if not match: return None start = int(match.group(1)) end_raw = match.group(2) end = int(end_raw) if len(end_raw) == 4 else int(str(start)[:2] + end_raw) return start, end def _month_add(year: int, month: int, offset: int) -> tuple[int, int]: index = (year * 12 + (month - 1)) + int(offset or 0) return index // 12, index % 12 + 1 def _safe_date(year: int, month: int, day: int) -> date | None: try: return date(int(year), int(month), int(day)) except Exception: return None def _period_month_from_label(financial_year: str | None, period_label: str | None) -> tuple[int, int] | None: label = (period_label or "").strip().lower() fy = _parse_year_pair(financial_year) iso_match = re.match(r"^(\d{4})[-/](\d{1,2})$", label) if iso_match: return int(iso_match.group(1)), int(iso_match.group(2)) month_names = { "apr": 4, "april": 4, "may": 5, "jun": 6, "june": 6, "jul": 7, "july": 7, "aug": 8, "august": 8, "sep": 9, "sept": 9, "september": 9, "oct": 10, "october": 10, "nov": 11, "november": 11, "dec": 12, "december": 12, "jan": 1, "january": 1, "feb": 2, "february": 2, "mar": 3, "march": 3, } if label in month_names and fy: month = month_names[label] year = fy[0] if month >= 4 else fy[1] return year, month return None def _quarter_end_from_label(financial_year: str | None, period_label: str | None) -> tuple[int, int] | None: label = (period_label or "").strip().lower().replace(" ", "") fy = _parse_year_pair(financial_year) if not fy: return None mapping = { "q1": (fy[0], 6), "quarter1": (fy[0], 6), "apr-jun": (fy[0], 6), "q2": (fy[0], 9), "quarter2": (fy[0], 9), "jul-sep": (fy[0], 9), "q3": (fy[0], 12), "quarter3": (fy[0], 12), "oct-dec": (fy[0], 12), "q4": (fy[1], 3), "quarter4": (fy[1], 3), "jan-mar": (fy[1], 3), } return mapping.get(label) def calculate_due_date( rule: ServiceDueDateRule | None, *, financial_year: str | None, assessment_year: str | None = None, period_label: str | None = None, expiry_date: date | None = None, ) -> date | None: """Calculate a statutory due date from a catalogue due-date rule. The function is intentionally conservative. If the rule needs a period label and the engagement does not yet carry one, it returns None instead of guessing. This preserves existing engagement creation behaviour. """ if not rule or not getattr(rule, "is_active", True): return None day = getattr(rule, "due_day", None) period_type = (getattr(rule, "period_type", None) or "yearly").strip().lower() due_month = getattr(rule, "due_month", None) month_offset = int(getattr(rule, "due_month_offset", None) or 0) if period_type in {"renewal_based", "before_expiry", "expiry_based"}: if not expiry_date: return None days_before = int(getattr(rule, "renewal_days_before_expiry", None) or 0) from datetime import timedelta return expiry_date - timedelta(days=days_before) if not day: return None if period_type in {"yearly", "one_time"}: if not due_month: return None basis = (getattr(rule, "due_year_basis", None) or "assessment_year_start").strip().lower() fy = _parse_year_pair(financial_year) ay = _parse_year_pair(assessment_year) if basis == "financial_year_start" and fy: year = fy[0] elif basis == "financial_year_end" and fy: year = fy[1] elif ay: year = ay[0] elif fy: year = fy[1] else: return None return _safe_date(year, int(due_month), int(day)) if period_type == "monthly": period = _period_month_from_label(financial_year, period_label) if not period: return None year, month = _month_add(period[0], period[1], month_offset) return _safe_date(year, month, int(day)) if period_type == "quarterly": period = _quarter_end_from_label(financial_year, period_label) if not period: return None year, month = _month_add(period[0], period[1], month_offset) return _safe_date(year, month, int(day)) return None def get_active_due_rule_for_catalogue(db: Session, *, catalogue_id: int) -> ServiceDueDateRule | None: return db.execute( select(ServiceDueDateRule) .where( ServiceDueDateRule.service_catalogue_id == catalogue_id, ServiceDueDateRule.is_active.is_(True), ) .order_by(ServiceDueDateRule.sort_order.asc(), ServiceDueDateRule.id.asc()) ).scalars().first() def list_due_rules(db: Session, *, catalogue_id: int, include_inactive: bool = True) -> list[ServiceDueDateRule]: query = select(ServiceDueDateRule).where(ServiceDueDateRule.service_catalogue_id == catalogue_id) if not include_inactive: query = query.where(ServiceDueDateRule.is_active.is_(True)) return db.execute(query.order_by(ServiceDueDateRule.sort_order.asc(), ServiceDueDateRule.id.asc())).scalars().all() def get_due_rule(db: Session, *, rule_id: int, catalogue_id: int | None = None) -> ServiceDueDateRule | None: query = select(ServiceDueDateRule).where(ServiceDueDateRule.id == rule_id) if catalogue_id: query = query.where(ServiceDueDateRule.service_catalogue_id == catalogue_id) return db.execute(query).scalar_one_or_none() def list_due_extensions(db: Session, *, catalogue_id: int, tenant_id: int | None = None, limit: int = 20) -> list[ServiceDueDateExtension]: query = ( select(ServiceDueDateExtension) .options(selectinload(ServiceDueDateExtension.due_rule)) .where(ServiceDueDateExtension.service_catalogue_id == catalogue_id) ) if tenant_id: query = query.where(ServiceDueDateExtension.tenant_id == tenant_id) return db.execute( query.order_by(ServiceDueDateExtension.id.desc()).limit(limit) ).scalars().all() def apply_due_date_rule_to_subscription( db: Session, subscription: ClientServiceSubscription, *, force: bool = False, ) -> date | None: if getattr(subscription, "is_locked", False): return getattr(subscription, "current_due_date", None) rule = get_active_due_rule_for_catalogue(db, catalogue_id=subscription.service_catalogue_id) rule_type = (getattr(rule, "period_type", None) or "").strip().lower() if rule else "" # For normal statutory rules, preserve an already calculated/extended due date. # For renewal-based rules, recalculate when expiry_date changes, unless the due date was manually overridden/extended. if getattr(subscription, "current_due_date", None) and not force: if rule_type not in {"renewal_based", "before_expiry", "expiry_based"}: return subscription.current_due_date if getattr(subscription, "due_date_source", None) in {DUE_DATE_SOURCE_EXTENSION, DUE_DATE_SOURCE_MANUAL}: return subscription.current_due_date calculated = calculate_due_date( rule, financial_year=subscription.financial_year, assessment_year=subscription.assessment_year, period_label=getattr(subscription, "period_label", None), expiry_date=getattr(subscription, "expiry_date", None), ) if calculated: subscription.due_date_rule_id = rule.id if rule else None subscription.original_due_date = calculated subscription.current_due_date = calculated subscription.due_date_source = DUE_DATE_SOURCE_RULE return calculated def _matching_extension_query( *, tenant_id: int, catalogue_id: int, rule_id: int | None, financial_year: str, assessment_year: str | None, period_label: str | None, ): query = select(ServiceDueDateExtension).where( ServiceDueDateExtension.tenant_id == tenant_id, ServiceDueDateExtension.service_catalogue_id == catalogue_id, ServiceDueDateExtension.financial_year == financial_year, ) if rule_id: query = query.where(ServiceDueDateExtension.due_date_rule_id == rule_id) if assessment_year: query = query.where(ServiceDueDateExtension.assessment_year == assessment_year) if period_label: query = query.where(ServiceDueDateExtension.period_label == period_label) else: query = query.where((ServiceDueDateExtension.period_label.is_(None)) | (ServiceDueDateExtension.period_label == "")) return query def create_due_date_extension( db: Session, *, tenant_id: int, catalogue_id: int, due_date_rule_id: int | None, financial_year: str, assessment_year: str | None, period_label: str | None, extended_due_date: date, notification_reference: str | None, notification_date: date | None, remarks: str | None, user_id: int, ) -> tuple[ServiceDueDateExtension, int, int]: rule = get_due_rule(db, rule_id=due_date_rule_id, catalogue_id=catalogue_id) if due_date_rule_id else get_active_due_rule_for_catalogue(db, catalogue_id=catalogue_id) latest = db.execute( _matching_extension_query( tenant_id=tenant_id, catalogue_id=catalogue_id, rule_id=rule.id if rule else None, financial_year=financial_year, assessment_year=assessment_year, period_label=period_label, ).order_by(ServiceDueDateExtension.extension_sequence.desc(), ServiceDueDateExtension.id.desc()) ).scalars().first() base_due = calculate_due_date(rule, financial_year=financial_year, assessment_year=assessment_year, period_label=period_label) previous_due = latest.extended_due_date if latest else base_due sequence = int((latest.extension_sequence if latest else 0) or 0) + 1 extension = ServiceDueDateExtension( tenant_id=tenant_id, service_catalogue_id=catalogue_id, due_date_rule_id=rule.id if rule else None, financial_year=financial_year, assessment_year=assessment_year, period_label=(period_label or "").strip() or None, previous_due_date=previous_due, extended_due_date=extended_due_date, extension_sequence=sequence, notification_reference=(notification_reference or "").strip() or None, notification_date=notification_date, remarks=(remarks or "").strip() or None, created_by_user_id=user_id, updated_by_user_id=user_id, ) db.add(extension) db.flush() updated = skipped_locked = 0 sub_query = select(ClientServiceSubscription).where( ClientServiceSubscription.tenant_id == tenant_id, ClientServiceSubscription.service_catalogue_id == catalogue_id, ClientServiceSubscription.financial_year == financial_year, ) if assessment_year: sub_query = sub_query.where(ClientServiceSubscription.assessment_year == assessment_year) subscriptions = db.execute(sub_query).scalars().all() for sub in subscriptions: if getattr(sub, "is_locked", False): skipped_locked += 1 continue if rule: sub.due_date_rule_id = rule.id if not getattr(sub, "original_due_date", None): sub.original_due_date = previous_due or base_due sub.current_due_date = extended_due_date sub.due_date_source = DUE_DATE_SOURCE_EXTENSION sub.updated_by_user_id = user_id updated += 1 return extension, updated, skipped_locked