from __future__ import annotations

import logging
from typing import Any

logger = logging.getLogger(__name__)


class ErpTakeoffRulesCrudMixin:
    def save_takeoff_learned_rule(
        self,
        drawing_type: str | dict[str, Any] = "general",
        category: str = "all",
        error_pattern: str = "",
        correction_rule: str = "",
        learned_from_user: str = "Kỹ Sư Trưởng QS",
        takeoff_id: str | None = None,
        confidence_weight: float = 1.0,
    ) -> dict[str, Any]:
        """Lưu trữ hoặc củng cố một quy tắc rút kinh nghiệm của AI Quỳnh, tự động chống trùng lặp."""
        if isinstance(drawing_type, dict):
            payload = drawing_type
            drawing_type = payload.get("drawing_type") or "general"
            category = payload.get("category") or payload.get("rule_category") or "all"
            error_pattern = (
                payload.get("error_pattern")
                or payload.get("trigger_pattern")
                or payload.get("rule_title")
                or ""
            )
            correction_rule = (
                payload.get("correction_rule")
                or payload.get("correction_directive")
                or ""
            )
            learned_from_user = (
                payload.get("learned_from_user")
                or payload.get("created_by_role")
                or "Kỹ Sư Trưởng QS"
            )
            takeoff_id = payload.get("takeoff_id")
            confidence_weight = float(payload.get("confidence_weight") or 1.0)

        error_pattern = str(error_pattern or "").strip()
        correction_rule = str(correction_rule or "").strip()
        drawing_type = str(drawing_type or "general").strip()
        category = str(category or "all").strip()

        with self.get_connection() as conn:
            with conn.cursor() as cur:
                # 1. Kiểm tra xem quy tắc với nội dung tương tự đã tồn tại trong CSDL chưa (chống trùng lặp)
                find_sql = """
                    SELECT id, apply_count, confidence_weight, learned_from_user
                    FROM erp_ai_takeoff_learned_rules
                    WHERE (
                        (LOWER(TRIM(correction_rule)) = LOWER(TRIM(%s)) AND LOWER(TRIM(error_pattern)) = LOWER(TRIM(%s)))
                        OR (LENGTH(TRIM(%s)) >= 25 AND LOWER(TRIM(correction_rule)) = LOWER(TRIM(%s)))
                        OR (LENGTH(TRIM(%s)) >= 25 AND LOWER(TRIM(error_pattern)) = LOWER(TRIM(%s)))
                    )
                    AND is_active = TRUE
                    ORDER BY apply_count DESC, created_at ASC
                    LIMIT 1;
                """
                cur.execute(
                    find_sql,
                    (
                        correction_rule,
                        error_pattern,
                        correction_rule,
                        correction_rule,
                        error_pattern,
                        error_pattern,
                    ),
                )
                existing_match = cur.fetchone()

                if existing_match:
                    rule_id = existing_match["id"]
                    update_sql = """
                        UPDATE erp_ai_takeoff_learned_rules
                        SET apply_count = apply_count + 1,
                            confidence_weight = GREATEST(confidence_weight, %s),
                            learned_from_user = CASE 
                                WHEN %s <> '' AND %s NOT ILIKE '%%test%%' THEN %s 
                                ELSE learned_from_user 
                            END,
                            updated_at = NOW()
                        WHERE id = %s
                        RETURNING *;
                    """
                    cur.execute(
                        update_sql,
                        (
                            confidence_weight,
                            learned_from_user,
                            learned_from_user,
                            learned_from_user,
                            rule_id,
                        ),
                    )
                    row = cur.fetchone()
                else:
                    insert_sql = """
                        INSERT INTO erp_ai_takeoff_learned_rules (
                            takeoff_id, drawing_type, category, error_pattern,
                            correction_rule, learned_from_user, is_active, confidence_weight,
                            apply_count, created_at, updated_at
                        ) VALUES (
                            %s, %s, %s, %s, %s, %s, TRUE, %s, 1, NOW(), NOW()
                        ) RETURNING *;
                    """
                    cur.execute(
                        insert_sql,
                        (
                            takeoff_id,
                            drawing_type,
                            category,
                            error_pattern,
                            correction_rule,
                            learned_from_user,
                            confidence_weight,
                        ),
                    )
                    row = cur.fetchone()

                conn.commit()
                if row:
                    row["id"] = str(row["id"])
                    if row.get("takeoff_id") is not None:
                        row["takeoff_id"] = str(row["takeoff_id"])
                    err_pat = row.get("error_pattern") or ""
                    corr_rule = row.get("correction_rule") or ""
                    clean_title = (
                        corr_rule.split("\n")[0]
                        .replace("Bắt buộc rà soát và hiệu chỉnh:", "")
                        .strip()
                    )
                    if not clean_title:
                        clean_title = (
                            err_pat.split("\n")[0]
                            .replace("Phản hồi từ kỹ sư:", "")
                            .strip()
                        )
                    if len(clean_title) > 70:
                        clean_title = clean_title[:70] + "..."
                    row["rule_title"] = (
                        clean_title or f"Quy tắc {row.get('drawing_type', 'xây dựng')}"
                    )
                    row["rule_category"] = row.get("category") or "TẤT CẢ"
                    row["trigger_pattern"] = err_pat
                    row["correction_directive"] = corr_rule
                    row["created_by_role"] = (
                        row.get("learned_from_user") or "Kỹ Sư Trưởng QS"
                    )
                return row

    def list_takeoff_learned_rules(
        self,
        drawing_type: str | None = None,
        category: str | None = None,
        is_active: bool = True,
        limit: int = 100,
    ) -> list[dict[str, Any]]:
        """Lấy danh sách các quy tắc rút kinh nghiệm của AI Quỳnh."""
        conditions = ["1=1"]
        params: list[Any] = []

        if is_active is not None:
            conditions.append("is_active = %s")
            params.append(is_active)

        if drawing_type and drawing_type != "all":
            conditions.append(
                "(drawing_type ILIKE %s OR drawing_type ILIKE 'general' OR drawing_type ILIKE 'all')"
            )
            params.append(drawing_type)

        if category and category != "all":
            conditions.append("(category ILIKE %s OR category ILIKE 'all')")
            params.append(category)

        params.append(limit)
        sql = f"""
            SELECT * FROM erp_ai_takeoff_learned_rules
            WHERE {" AND ".join(conditions)}
            ORDER BY apply_count DESC, created_at DESC
            LIMIT %s;
        """
        with self.get_connection() as conn:
            with conn.cursor() as cur:
                cur.execute(sql, tuple(params))
                rows = cur.fetchall()
                # Format UUID to str & enrich aliases
                for r in rows:
                    if "id" in r and r["id"] is not None:
                        r["id"] = str(r["id"])
                    if "takeoff_id" in r and r["takeoff_id"] is not None:
                        r["takeoff_id"] = str(r["takeoff_id"])

                    err_pat = r.get("error_pattern") or ""
                    corr_rule = r.get("correction_rule") or ""
                    clean_title = (
                        corr_rule.split("\n")[0]
                        .replace("Bắt buộc rà soát và hiệu chỉnh:", "")
                        .strip()
                    )
                    if not clean_title:
                        clean_title = (
                            err_pat.split("\n")[0]
                            .replace("Phản hồi từ kỹ sư:", "")
                            .strip()
                        )
                    if len(clean_title) > 70:
                        clean_title = clean_title[:70] + "..."
                    r["rule_title"] = (
                        clean_title or f"Quy tắc {r.get('drawing_type', 'xây dựng')}"
                    )
                    r["rule_category"] = r.get("category") or "TẤT CẢ"
                    r["trigger_pattern"] = err_pat
                    r["correction_directive"] = corr_rule
                    r["created_by_role"] = (
                        r.get("learned_from_user") or "Kỹ Sư Trưởng QS"
                    )
                return rows

    def update_takeoff_learned_rule(
        self, rule_id: str, payload: dict[str, Any]
    ) -> dict[str, Any] | None:
        """Cập nhật hoặc bật/tắt một quy tắc rút kinh nghiệm."""
        if not payload:
            return None
        set_clauses = []
        params = []
        for k, v in payload.items():
            set_clauses.append(f"{k} = %s")
            params.append(v)
        params.append(rule_id)
        sql = f"UPDATE erp_ai_takeoff_learned_rules SET {', '.join(set_clauses)}, updated_at = NOW() WHERE id = %s RETURNING *;"
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, params)
            row = cur.fetchone()
            conn.commit()
            return row

    def delete_takeoff_learned_rule(self, rule_id: str) -> bool:
        """Xóa vĩnh viễn một quy tắc rút kinh nghiệm."""
        sql = "DELETE FROM erp_ai_takeoff_learned_rules WHERE id = %s RETURNING id;"
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, (rule_id,))
            row = cur.fetchone()
            conn.commit()
            return bool(row)

    def deduplicate_takeoff_learned_rules(self) -> dict[str, Any]:
        """Gộp các quy tắc trùng lặp trong kho tri thức thành 1 quy tắc duy nhất, cộng dồn apply_count."""
        with self.get_connection() as conn:
            with conn.cursor() as cur:
                # 1. Tìm tất cả các nhóm trùng lặp theo LOWER(TRIM(correction_rule)) và LOWER(TRIM(error_pattern))
                cur.execute("""
                    SELECT 
                        LOWER(TRIM(correction_rule)) as norm_corr,
                        LOWER(TRIM(error_pattern)) as norm_err,
                        COUNT(*) as group_cnt,
                        SUM(COALESCE(apply_count, 1)) as total_applied,
                        ARRAY_AGG(id ORDER BY apply_count DESC, created_at ASC) as rule_ids,
                        MAX(updated_at) as latest_updated
                    FROM erp_ai_takeoff_learned_rules
                    GROUP BY LOWER(TRIM(correction_rule)), LOWER(TRIM(error_pattern))
                    HAVING COUNT(*) > 1;
                """)
                duplicate_groups = cur.fetchall()
                merged_groups_count = len(duplicate_groups)
                deleted_count = 0

                for group in duplicate_groups:
                    rule_ids = group["rule_ids"]
                    primary_id = rule_ids[0]
                    dup_ids = rule_ids[1:]
                    total_applied = int(group["total_applied"] or len(rule_ids))
                    latest_updated = group["latest_updated"]

                    # Cập nhật bản ghi chính với tổng apply_count và thời gian mới nhất
                    cur.execute(
                        """
                        UPDATE erp_ai_takeoff_learned_rules
                        SET apply_count = %s,
                            updated_at = %s
                        WHERE id = %s;
                    """,
                        (total_applied, latest_updated, primary_id),
                    )

                    # Xóa các bản ghi trùng lặp
                    cur.execute(
                        """
                        DELETE FROM erp_ai_takeoff_learned_rules
                        WHERE id = ANY(%s);
                    """,
                        (dup_ids,),
                    )
                    deleted_count += len(dup_ids)

                conn.commit()
                return {
                    "status": "success",
                    "duplicate_groups_merged": merged_groups_count,
                    "deleted_duplicate_rules": deleted_count,
                }
