from __future__ import annotations

import logging
from typing import Any

logger = logging.getLogger("dscons.postgres.project_crud.wbs")


class WbsOperationsMixin:
    """WBS node hierarchy management and automatic progress recalculation."""

    def _build_wbs_tree(self, wbs_rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
        """Chuyển đổi danh sách phẳng WBS thành cây phân cấp."""
        id_map = {r["id"]: {**r, "subtasks": []} for r in wbs_rows}
        tree = []

        for r in wbs_rows:
            node = id_map[r["id"]]
            parent_id = r.get("parent_wbs_id")
            if parent_id and parent_id in id_map:
                id_map[parent_id]["subtasks"].append(node)
            else:
                tree.append(node)

        return tree

    def _insert_wbs_items_batch(
        self, cur: Any, project_id: Any, items: list[dict[str, Any]]
    ) -> None:
        """Helper to insert batch WBS hierarchy."""
        sql_insert = """
            INSERT INTO erp_project_wbs (
                project_id, parent_wbs_id, wbs_code, task_name, task_type,
                unit, quantity, unit_price_vnd, total_amount_vnd, start_date,
                end_date, duration_days, progress_percent, status, assigned_role,
                sort_order, notes
            ) VALUES (
                %(project_id)s, %(parent_wbs_id)s, %(wbs_code)s, %(task_name)s, %(task_type)s,
                %(unit)s, %(quantity)s, %(unit_price_vnd)s, %(total_amount_vnd)s, %(start_date)s,
                %(end_date)s, %(duration_days)s, %(progress_percent)s, %(status)s, %(assigned_role)s,
                %(sort_order)s, %(notes)s
            ) RETURNING id, wbs_code;
        """
        code_to_id = {}
        for item in items:
            parent_code = item.get("parent_wbs_code")
            parent_id = item.get("parent_wbs_id") or code_to_id.get(parent_code)

            cur.execute(
                sql_insert,
                {
                    "project_id": project_id,
                    "parent_wbs_id": parent_id,
                    "wbs_code": item.get("wbs_code", "1.0"),
                    "task_name": item.get("task_name", "Hạng mục công việc"),
                    "task_type": item.get("task_type", "task"),
                    "unit": item.get("unit"),
                    "quantity": item.get("quantity", 0),
                    "unit_price_vnd": item.get("unit_price_vnd", 0),
                    "total_amount_vnd": item.get("total_amount_vnd", 0),
                    "start_date": item.get("start_date"),
                    "end_date": item.get("end_date"),
                    "duration_days": item.get("duration_days", 1),
                    "progress_percent": item.get("progress_percent", 0),
                    "status": item.get("status", "todo"),
                    "assigned_role": item.get("assigned_role"),
                    "sort_order": item.get("sort_order", 0),
                    "notes": item.get("notes", ""),
                },
            )
            res = cur.fetchone()
            if res:
                code_to_id[res["wbs_code"]] = res["id"]

    def create_wbs_item(
        self, project_id: str, payload: dict[str, Any]
    ) -> dict[str, Any]:
        """Tạo 1 công việc WBS mới."""
        sql = """
            INSERT INTO erp_project_wbs (
                project_id, parent_wbs_id, wbs_code, task_name, task_type,
                unit, quantity, unit_price_vnd, total_amount_vnd, start_date,
                end_date, duration_days, progress_percent, status, assigned_role,
                sort_order, notes
            ) VALUES (
                %(project_id)s, %(parent_wbs_id)s, %(wbs_code)s, %(task_name)s, %(task_type)s,
                %(unit)s, %(quantity)s, %(unit_price_vnd)s, %(total_amount_vnd)s, %(start_date)s,
                %(end_date)s, %(duration_days)s, %(progress_percent)s, %(status)s, %(assigned_role)s,
                %(sort_order)s, %(notes)s
            ) RETURNING *;
        """
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(
                sql,
                {
                    "project_id": project_id,
                    "parent_wbs_id": payload.get("parent_wbs_id"),
                    "wbs_code": payload.get("wbs_code", "1.1"),
                    "task_name": payload.get("task_name", "Công việc mới"),
                    "task_type": payload.get("task_type", "task"),
                    "unit": payload.get("unit"),
                    "quantity": payload.get("quantity", 0),
                    "unit_price_vnd": payload.get("unit_price_vnd", 0),
                    "total_amount_vnd": payload.get("total_amount_vnd", 0),
                    "start_date": payload.get("start_date"),
                    "end_date": payload.get("end_date"),
                    "duration_days": payload.get("duration_days", 1),
                    "progress_percent": payload.get("progress_percent", 0),
                    "status": payload.get("status", "todo"),
                    "assigned_role": payload.get("assigned_role"),
                    "sort_order": payload.get("sort_order", 0),
                    "notes": payload.get("notes", ""),
                },
            )
            item = cur.fetchone()
            self._recalculate_project_progress(cur, project_id)
            conn.commit()
            return item

    def update_wbs_item(self, wbs_id: str, updates: dict[str, Any]) -> dict[str, Any]:
        """Cập nhật tiến độ hoặc trạng thái công việc WBS."""
        allowed_fields = [
            "task_name",
            "task_type",
            "unit",
            "quantity",
            "unit_price_vnd",
            "total_amount_vnd",
            "start_date",
            "end_date",
            "duration_days",
            "progress_percent",
            "status",
            "assigned_role",
            "sort_order",
            "notes",
        ]

        set_clauses = []
        params: dict[str, Any] = {"id": wbs_id}

        for k, v in updates.items():
            if k in allowed_fields:
                set_clauses.append(f"{k} = %({k})s")
                params[k] = v

        if not set_clauses:
            return {}

        set_clauses.append("updated_at = NOW()")
        sql = f"""
            UPDATE erp_project_wbs
            SET {", ".join(set_clauses)}
            WHERE id::text = %(id)s
            RETURNING *;
        """

        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, params)
            item = cur.fetchone()
            if item:
                self._recalculate_project_progress(cur, str(item["project_id"]))
            conn.commit()
            return item

    def delete_wbs_item(self, wbs_id: str) -> bool:
        """Xóa công việc WBS."""
        sql_find = "SELECT project_id FROM erp_project_wbs WHERE id::text = %s;"
        sql_del = "DELETE FROM erp_project_wbs WHERE id::text = %s;"
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql_find, (wbs_id,))
            row = cur.fetchone()
            if not row:
                return False
            project_id = str(row["project_id"])

            cur.execute(sql_del, (wbs_id,))
            self._recalculate_project_progress(cur, project_id)
            conn.commit()
            return True

    def _recalculate_project_progress(self, cur: Any, project_id: str) -> None:
        """Tự động tính toán lại overall_progress_percent cho dự án từ các tasks WBS."""
        sql_calc = """
            SELECT 
                COALESCE(AVG(progress_percent), 0) AS avg_progress
            FROM erp_project_wbs
            WHERE project_id::text = %s AND task_type = 'task';
        """
        cur.execute(sql_calc, (project_id,))
        res = cur.fetchone()
        if res:
            avg_progress = int(round(float(res["avg_progress"])))
            sql_up = "UPDATE projects SET progress_percent = %s, updated_at = NOW() WHERE id::text = %s;"
            cur.execute(sql_up, (avg_progress, project_id))
