from __future__ import annotations

import logging
from typing import Any

logger = logging.getLogger("dscons.postgres.takeoff.rules")


class TakeoffRulesMixin:
    """Learned rules memory persistence and application metrics."""

    def save_takeoff_learned_rule(
        self,
        drawing_type: str,
        category: str,
        error_pattern: str,
        correction_rule: str,
        learned_from_user: str = "Kỹ Sư Trưởng QS",
        takeoff_id: str | None = None,
    ) -> dict[str, Any]:
        """Save or reinforce a learned rule with automatic deduplication."""
        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:
                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,
                            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,
                        (
                            learned_from_user,
                            learned_from_user,
                            learned_from_user,
                            rule_id,
                        ),
                    )
                    row = cur.fetchone()
                else:
                    insert_sql = """
                        INSERT INTO erp_ai_takeoff_learned_rules (
                            drawing_type, category, error_pattern, correction_rule,
                            learned_from_user, takeoff_id, apply_count, created_at, updated_at
                        )
                        VALUES (%s, %s, %s, %s, %s, %s, 1, NOW(), NOW())
                        RETURNING *;
                    """
                    cur.execute(
                        insert_sql,
                        (
                            drawing_type,
                            category,
                            error_pattern,
                            correction_rule,
                            learned_from_user,
                            takeoff_id,
                        ),
                    )
                    row = cur.fetchone()

                conn.commit()
                res = dict(row)
                res["rule_title"] = (
                    f"Quy tắc {category.upper()}: {correction_rule[:60]}..."
                )
                res["rule_category"] = res.get("category")
                res["trigger_pattern"] = res.get("error_pattern")
                res["correction_directive"] = res.get("correction_rule")
                res["created_by_role"] = res.get("learned_from_user")
                return res

    def list_takeoff_learned_rules(
        self,
        drawing_type: str | None = None,
        is_active: bool = True,
    ) -> list[dict[str, Any]]:
        """List active learned rules."""
        where_clauses = ["is_active = %s"]
        params: list[Any] = [is_active]

        if drawing_type and drawing_type != "all":
            where_clauses.append("(drawing_type = %s OR drawing_type = 'all')")
            params.append(drawing_type)

        sql = f"""
            SELECT * FROM erp_ai_takeoff_learned_rules
            WHERE {" AND ".join(where_clauses)}
            ORDER BY apply_count DESC, created_at DESC;
        """
        with self.get_connection() as conn:
            with conn.cursor() as cur:
                cur.execute(sql, tuple(params))
                rows = cur.fetchall()
                result = []
                for r in rows:
                    item = dict(r)
                    cat = item.get("category", "general")
                    rule_text = item.get("correction_rule", "")
                    title = (
                        rule_text.split("\n")[0][:80]
                        if rule_text
                        else f"Quy tắc {cat.upper()}"
                    )
                    item["rule_title"] = title
                    item["rule_category"] = cat
                    item["trigger_pattern"] = item.get("error_pattern")
                    item["correction_directive"] = rule_text
                    item["created_by_role"] = item.get("learned_from_user")
                    result.append(item)
                return result

    def delete_takeoff_learned_rule(self, rule_id: str) -> bool:
        """Delete a learned rule by ID."""
        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,))
            deleted = cur.fetchone()
            conn.commit()
            return bool(deleted)

    def record_learned_rule_application(self, rule_id: str) -> None:
        """Increment application counter for a learned rule."""
        sql = """
            UPDATE erp_ai_takeoff_learned_rules
            SET apply_count = apply_count + 1, updated_at = NOW()
            WHERE id = %s;
        """
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, (rule_id,))
            conn.commit()
