"""
PM Module — Domain Services

سرویس‌های حوزه‌ای — محاسبات خالص بدون وابستگی خارجی.
"""

from dataclasses import dataclass, field
from datetime import date, timedelta
from decimal import Decimal
from typing import List, Optional, Dict, Set, Tuple, FrozenSet
from uuid import UUID

from ..entities.task import Task, Dependency
from ..entities.cost import Budget, CostEntry
from ..value_objects.common import (
    DependencyType,
    TaskStatus,
    EVMMetrics,
)
from ..exceptions.pm_exceptions import CircularDependencyException


# ═══════════════════════════════════════════════════
# Calendar Helper
# ═══════════════════════════════════════════════════

@dataclass(frozen=True)
class CalendarData:
    """
    داده‌های تقویم کاری برای محاسبات CPM.

    working_weekdays: مجموعه شماره روزهای کاری هفته (0=دوشنبه ... 6=یکشنبه)
    holidays: مجموعه تاریخ‌های تعطیل
    hours_per_day: ساعت کاری روزانه
    """
    working_weekdays: FrozenSet[int] = frozenset({0, 1, 2, 3, 4, 5})  # شنبه تا پنجشنبه (default Iran)
    holidays: FrozenSet[date] = frozenset()
    hours_per_day: int = 8


# Default calendar: Saturday–Thursday (Iran standard), no holidays
DEFAULT_CALENDAR = CalendarData()


def _is_working_day(d: date, cal: CalendarData) -> bool:
    """آیا روز مشخص‌شده روز کاری است؟"""
    return d.weekday() in cal.working_weekdays and d not in cal.holidays


def _add_working_days(start: date, work_days: int, cal: CalendarData) -> date:
    """
    اضافه کردن تعداد مشخصی روز کاری به تاریخ.

    work_days=0 → همان تاریخ (اگر کاری باشد) یا اولین روز کاری بعدی
    work_days>0 → start + work_days روز کاری
    """
    if work_days <= 0:
        # Return start itself if it's a working day, otherwise next working day
        current = start
        while not _is_working_day(current, cal):
            current += timedelta(days=1)
        return current

    current = start
    # Ensure start is a working day
    while not _is_working_day(current, cal):
        current += timedelta(days=1)

    remaining = work_days - 1  # -1 because duration includes the start day
    while remaining > 0:
        current += timedelta(days=1)
        if _is_working_day(current, cal):
            remaining -= 1
        # Skip non-working days automatically
    return current


def _subtract_working_days(end: date, work_days: int, cal: CalendarData) -> date:
    """
    کم کردن تعداد مشخصی روز کاری از تاریخ.
    """
    if work_days <= 0:
        current = end
        while not _is_working_day(current, cal):
            current -= timedelta(days=1)
        return current

    current = end
    while not _is_working_day(current, cal):
        current -= timedelta(days=1)

    remaining = work_days - 1
    while remaining > 0:
        current -= timedelta(days=1)
        if _is_working_day(current, cal):
            remaining -= 1
    return current


def _working_days_between(start: date, end: date, cal: CalendarData) -> int:
    """تعداد روزهای کاری بین دو تاریخ (شامل هر دو سر)."""
    if end < start:
        return 0
    count = 0
    current = start
    while current <= end:
        if _is_working_day(current, cal):
            count += 1
        current += timedelta(days=1)
    return count


def _next_working_day(d: date, cal: CalendarData) -> date:
    """اولین روز کاری بعد از تاریخ مشخص‌شده."""
    result = d + timedelta(days=1)
    while not _is_working_day(result, cal):
        result += timedelta(days=1)
    return result


def _prev_working_day(d: date, cal: CalendarData) -> date:
    """آخرین روز کاری قبل از تاریخ مشخص‌شده."""
    result = d - timedelta(days=1)
    while not _is_working_day(result, cal):
        result -= timedelta(days=1)
    return result


class CriticalPathService:
    """
    محاسبه مسیر بحرانی (CPM — Critical Path Method).

    الگوریتم:
    1. Forward Pass: محاسبه ES/EF
    2. Backward Pass: محاسبه LS/LF
    3. Float = LS - ES (یا LF - EF)
    4. Critical Path = تسک‌هایی با Float = 0
    """

    def calculate(
        self,
        tasks: List[Task],
        dependencies: List[Dependency],
        calendar: Optional[CalendarData] = None,
    ) -> List[Task]:
        """
        محاسبه مسیر بحرانی و به‌روزرسانی ES/EF/LS/LF/Float.

        Args:
            tasks: لیست تسک‌ها
            dependencies: لیست وابستگی‌ها
            calendar: تقویم کاری (اختیاری — اگر None باشد از تقویم پیش‌فرض استفاده می‌شود)

        Returns: لیست تسک‌های به‌روزرسانی‌شده
        """
        if not tasks:
            return tasks

        cal = calendar or DEFAULT_CALENDAR
        task_map: Dict[UUID, Task] = {t.id: t for t in tasks}
        # فقط leaf tasks (بدون summary)
        leaf_ids = {t.id for t in tasks if t.task_type.value != 'summary'}

        # Build adjacency
        successors: Dict[UUID, List[Tuple[UUID, Dependency]]] = {tid: [] for tid in leaf_ids}
        predecessors: Dict[UUID, List[Tuple[UUID, Dependency]]] = {tid: [] for tid in leaf_ids}

        for dep in dependencies:
            if dep.predecessor_id in leaf_ids and dep.successor_id in leaf_ids:
                successors.setdefault(dep.predecessor_id, []).append((dep.successor_id, dep))
                predecessors.setdefault(dep.successor_id, []).append((dep.predecessor_id, dep))

        # Detect circular dependencies
        self._detect_cycles(leaf_ids, successors)

        # --- Forward Pass (calendar-aware) ---
        visited: Set[UUID] = set()

        def forward(tid: UUID) -> date:
            if tid in visited:
                t = task_map[tid]
                return t.early_finish if t.early_finish else t.planned_start or date.today()
            visited.add(tid)

            t = task_map[tid]
            preds = predecessors.get(tid, [])

            if not preds:
                raw_es = t.planned_start or date.today()
                es = _add_working_days(raw_es, 0, cal)  # snap to working day
            else:
                es_candidates = []
                for pred_id, dep in preds:
                    pred_ef = forward(pred_id)
                    lag = dep.lag_days
                    if dep.dependency_type == DependencyType.FS:
                        # ES = pred EF + 1 working day + lag working days
                        anchor = _next_working_day(pred_ef, cal)
                        if lag > 0:
                            anchor = _add_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _subtract_working_days(anchor, -lag, cal)
                        es_candidates.append(anchor)
                    elif dep.dependency_type == DependencyType.SS:
                        pred_t = task_map[pred_id]
                        anchor = pred_t.early_start or pred_t.planned_start or date.today()
                        if lag > 0:
                            anchor = _add_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _subtract_working_days(anchor, -lag, cal)
                        es_candidates.append(anchor)
                    elif dep.dependency_type == DependencyType.FF:
                        # EF = pred EF + lag, then ES = EF - duration + 1
                        anchor_ef = pred_ef
                        if lag > 0:
                            anchor_ef = _add_working_days(anchor_ef, lag, cal)
                        elif lag < 0:
                            anchor_ef = _subtract_working_days(anchor_ef, -lag, cal)
                        es_candidates.append(_subtract_working_days(anchor_ef, max(t.duration, 1), cal))
                    elif dep.dependency_type == DependencyType.SF:
                        pred_t = task_map[pred_id]
                        anchor = pred_t.early_start or pred_t.planned_start or date.today()
                        if lag > 0:
                            anchor = _add_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _subtract_working_days(anchor, -lag, cal)
                        es_candidates.append(_subtract_working_days(anchor, max(t.duration, 1), cal))
                es = max(es_candidates) if es_candidates else _add_working_days(t.planned_start or date.today(), 0, cal)

            t.early_start = es
            t.early_finish = _add_working_days(es, max(t.duration, 1), cal)
            return t.early_finish

        for tid in leaf_ids:
            forward(tid)

        # --- Backward Pass ---
        project_end = max(
            (task_map[tid].early_finish for tid in leaf_ids if task_map[tid].early_finish),
            default=date.today()
        )

        # --- Backward Pass (calendar-aware) ---
        project_end = max(
            (task_map[tid].early_finish for tid in leaf_ids if task_map[tid].early_finish),
            default=date.today()
        )

        visited_back: Set[UUID] = set()

        def backward(tid: UUID) -> date:
            if tid in visited_back:
                t = task_map[tid]
                return t.late_start if t.late_start else project_end
            visited_back.add(tid)

            t = task_map[tid]
            succs = successors.get(tid, [])

            if not succs:
                lf = project_end
            else:
                lf_candidates = []
                for succ_id, dep in succs:
                    succ_ls = backward(succ_id)
                    lag = dep.lag_days
                    if dep.dependency_type == DependencyType.FS:
                        # LF = succ LS - 1 working day - lag
                        anchor = _prev_working_day(succ_ls, cal)
                        if lag > 0:
                            anchor = _subtract_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _add_working_days(anchor, -lag, cal)
                        lf_candidates.append(anchor)
                    elif dep.dependency_type == DependencyType.SS:
                        # LS constraint → LF = succ LS - lag + duration - 1
                        anchor = succ_ls
                        if lag > 0:
                            anchor = _subtract_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _add_working_days(anchor, -lag, cal)
                        lf_candidates.append(_add_working_days(anchor, max(t.duration, 1), cal))
                    elif dep.dependency_type == DependencyType.FF:
                        succ_t = task_map[succ_id]
                        anchor = succ_t.late_finish or project_end
                        if lag > 0:
                            anchor = _subtract_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _add_working_days(anchor, -lag, cal)
                        lf_candidates.append(anchor)
                    elif dep.dependency_type == DependencyType.SF:
                        succ_t = task_map[succ_id]
                        anchor = succ_t.late_finish or project_end
                        if lag > 0:
                            anchor = _subtract_working_days(anchor, lag, cal)
                        elif lag < 0:
                            anchor = _add_working_days(anchor, -lag, cal)
                        lf_candidates.append(_add_working_days(anchor, max(t.duration, 1), cal))
                lf = min(lf_candidates) if lf_candidates else project_end

            t.late_finish = lf
            t.late_start = _subtract_working_days(lf, max(t.duration, 1), cal)

            # Float (in working days)
            if t.early_start and t.late_start:
                t.total_float = _working_days_between(t.early_start, t.late_start, cal) - 1
                if t.total_float < 0:
                    t.total_float = 0
            else:
                t.total_float = 0

            t.is_critical = (t.total_float == 0)
            return t.late_start

        for tid in leaf_ids:
            backward(tid)

        # Free Float (calendar-aware)
        for tid in leaf_ids:
            t = task_map[tid]
            succs = successors.get(tid, [])
            if not succs:
                t.free_float = t.total_float
            else:
                min_succ_es = min(
                    (task_map[s_id].early_start for s_id, _ in succs if task_map[s_id].early_start),
                    default=project_end,
                )
                if t.early_finish:
                    t.free_float = _working_days_between(
                        _next_working_day(t.early_finish, cal), min_succ_es, cal
                    )
                else:
                    t.free_float = 0

        return tasks

    def _detect_cycles(self, node_ids: Set[UUID], successors: Dict[UUID, List[Tuple[UUID, any]]]) -> None:
        """تشخیص وابستگی حلقوی با DFS."""
        WHITE, GRAY, BLACK = 0, 1, 2
        color = {nid: WHITE for nid in node_ids}

        def dfs(u: UUID):
            color[u] = GRAY
            for v, _ in successors.get(u, []):
                if v not in color:
                    continue
                if color[v] == GRAY:
                    raise CircularDependencyException()
                if color[v] == WHITE:
                    dfs(v)
            color[u] = BLACK

        for nid in node_ids:
            if color[nid] == WHITE:
                dfs(nid)


class EVMService:
    """
    محاسبه شاخص‌های مدیریت ارزش کسب‌شده (EVM).
    """

    def calculate_evm(
        self,
        tasks: List[Task],
        budget_at_completion: Decimal,
        report_date: Optional[date] = None,
    ) -> EVMMetrics:
        """
        محاسبه EVM بر اساس تسک‌ها و BAC.
        """
        today = report_date or date.today()

        if not tasks or budget_at_completion == 0:
            return EVMMetrics(
                planned_value=Decimal("0"),
                earned_value=Decimal("0"),
                actual_cost=Decimal("0"),
                budget_at_completion=budget_at_completion,
            )

        total_planned = sum(t.planned_cost for t in tasks)
        if total_planned == 0:
            total_planned = Decimal("1")  # avoid division by zero

        # PV — Planned Value: بودجه برنامه‌ریزی‌شده تسک‌هایی که باید تا امروز تمام شوند
        pv = Decimal("0")
        for t in tasks:
            if t.planned_end and t.planned_end <= today:
                pv += t.planned_cost
            elif t.planned_start and t.planned_end and t.planned_start <= today:
                # نسبت تکمیل برنامه‌ای
                total_days = (t.planned_end - t.planned_start).days or 1
                elapsed = (today - t.planned_start).days
                ratio = Decimal(min(elapsed / total_days, 1.0))
                pv += t.planned_cost * ratio

        # EV — Earned Value: ارزش کسب‌شده بر اساس پیشرفت واقعی
        ev = Decimal("0")
        for t in tasks:
            ev += t.planned_cost * Decimal(t.progress) / Decimal("100")

        # AC — Actual Cost
        ac = sum(t.actual_cost for t in tasks)

        # Scale to BAC
        scale = budget_at_completion / total_planned if total_planned else Decimal("1")

        return EVMMetrics(
            planned_value=pv * scale,
            earned_value=ev * scale,
            actual_cost=ac,
            budget_at_completion=budget_at_completion,
        )

    def calculate_s_curve_data(
        self,
        tasks: List[Task],
        budget_at_completion: Decimal,
        start_date: date,
        end_date: date,
    ) -> List[Dict]:
        """
        تولید داده‌های نمودار S-Curve.
        """
        result = []
        current = start_date
        while current <= end_date:
            evm = self.calculate_evm(tasks, budget_at_completion, current)
            result.append({
                'date': current.isoformat(),
                'pv': float(evm.planned_value),
                'ev': float(evm.earned_value),
                'ac': float(evm.actual_cost),
            })
            current += timedelta(days=7)  # هفتگی
        return result


class WBSService:
    """
    سرویس WBS — ساختار شکست کار.
    """

    def generate_wbs_codes(self, tasks: List[Task]) -> List[Task]:
        """
        تولید کدهای WBS بر اساس ساختار درختی و ترتیب.
        """
        task_map = {t.id: t for t in tasks}
        children_map: Dict[Optional[UUID], List[Task]] = {}

        for t in tasks:
            children_map.setdefault(t.parent_id, []).append(t)

        # مرتب‌سازی بر اساس sort_order
        for key in children_map:
            children_map[key].sort(key=lambda x: x.sort_order)

        def assign_codes(parent_id: Optional[UUID], prefix: str):
            children = children_map.get(parent_id, [])
            for idx, child in enumerate(children, 1):
                code = f"{prefix}.{idx}" if prefix else str(idx)
                child.code = code
                assign_codes(child.id, code)

        assign_codes(None, "")
        return tasks

    def generate_code(self, parent_code: str, sibling_index: int) -> str:
        """
        تولید یک کد WBS واحد.

        Args:
            parent_code: کد والد (خالی برای ریشه)
            sibling_index: شماره ترتیب فرزند

        Returns: کد WBS تولیدشده (مثال: "1.2.3")
        """
        if parent_code:
            return f"{parent_code}.{sibling_index}"
        return str(sibling_index)

    def calculate_summary_dates(self, tasks: List[Task]) -> List[Task]:
        """
        محاسبه تاریخ‌ها و پیشرفت Summary Tasks از فرزندان.
        """
        task_map = {t.id: t for t in tasks}
        children_map: Dict[Optional[UUID], List[Task]] = {}
        for t in tasks:
            children_map.setdefault(t.parent_id, []).append(t)

        def get_all_leaves(parent_id: UUID) -> List[Task]:
            children = children_map.get(parent_id, [])
            leaves = []
            for child in children:
                grandchildren = children_map.get(child.id, [])
                if grandchildren:
                    leaves.extend(get_all_leaves(child.id))
                else:
                    leaves.append(child)
            return leaves

        # Update summary tasks bottom-up
        def update_summary(parent_id: UUID):
            children = children_map.get(parent_id, [])
            if not children:
                return

            # First update children that are summaries
            for child in children:
                if children_map.get(child.id):
                    update_summary(child.id)

            parent = task_map.get(parent_id)
            if not parent:
                return

            leaves = get_all_leaves(parent_id)
            if not leaves:
                return

            start_dates = [l.planned_start for l in leaves if l.planned_start]
            end_dates = [l.planned_end for l in leaves if l.planned_end]

            if start_dates:
                parent.planned_start = min(start_dates)
            if end_dates:
                parent.planned_end = max(end_dates)

            if parent.planned_start and parent.planned_end:
                parent.duration = (parent.planned_end - parent.planned_start).days + 1

            # Average progress
            if leaves:
                parent.progress = sum(l.progress for l in leaves) // len(leaves)

            # Cost sum
            parent.planned_cost = sum(l.planned_cost for l in leaves)
            parent.actual_cost = sum(l.actual_cost for l in leaves)

        # Find root-level summaries and update
        for t in tasks:
            if t.parent_id is None and children_map.get(t.id):
                update_summary(t.id)

        return tasks
