from __future__ import annotations

import json
import logging
from typing import Any

logger = logging.getLogger("dscons.postgres.takeoff.chat")


class TakeoffChatMixin:
    """Takeoff agent chat message persistence and inline suggestion applications."""

    def save_takeoff_chat_message(
        self,
        takeoff_id: str,
        sender_role: str,
        sender_name: str,
        message_text: str,
        suggested_corrections: list[dict[str, Any]] | None = None,
        extracted_rules: list[dict[str, Any]] | None = None,
        applied: bool = False,
    ) -> dict[str, Any]:
        """Save a chat message in the AI Quỳnh QS Interactive Studio."""
        sql = """
            INSERT INTO erp_ai_takeoff_chat_messages (
                takeoff_id, sender_role, sender_name, message_text,
                suggested_corrections, extracted_rules, applied
            )
            VALUES (%s, %s, %s, %s, %s, %s, %s)
            RETURNING *;
        """
        corrections_json = (
            json.dumps(suggested_corrections, ensure_ascii=False)
            if suggested_corrections
            else None
        )
        rules_json = (
            json.dumps(extracted_rules, ensure_ascii=False) if extracted_rules else None
        )

        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(
                sql,
                (
                    takeoff_id,
                    sender_role,
                    sender_name,
                    message_text,
                    corrections_json,
                    rules_json,
                    applied,
                ),
            )
            row = cur.fetchone()
            conn.commit()
            res = dict(row)
            res["message"] = res.get("message_text")
            res["ai_response"] = res.get("message_text")
            res["reply_message"] = res.get("message_text")
            res["sender"] = res.get("sender_role")
            res["suggested_boq_items"] = res.get("suggested_corrections")
            return res

    def list_takeoff_chat_messages(self, takeoff_id: str) -> list[dict[str, Any]]:
        """List all chat messages for a specific drawing takeoff studio."""
        sql = """
            SELECT * FROM erp_ai_takeoff_chat_messages
            WHERE takeoff_id = %s
            ORDER BY created_at ASC;
        """
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, (takeoff_id,))
            rows = cur.fetchall()
            result = []
            for r in rows:
                item = dict(r)
                item["message"] = item.get("message_text")
                item["ai_response"] = item.get("message_text")
                item["reply_message"] = item.get("message_text")
                item["sender"] = item.get("sender_role")
                item["suggested_boq_items"] = item.get("suggested_corrections")
                result.append(item)
            return result

    def clear_takeoff_chat_history(self, takeoff_id: str) -> int:
        """Clear all chat messages for a takeoff session."""
        sql = "DELETE FROM erp_ai_takeoff_chat_messages WHERE takeoff_id = %s RETURNING id;"
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, (takeoff_id,))
            rows = cur.fetchall()
            conn.commit()
            return len(rows)

    def apply_chat_suggested_corrections(
        self,
        takeoff_id: str,
        suggested_items: list[dict[str, Any]],
        message_id: str | None = None,
    ) -> dict[str, Any]:
        """Apply suggested corrected items into the BoQ items table."""
        if not suggested_items:
            return {"status": "noop", "updated_count": 0}

        self.delete_drawing_takeoff_items(takeoff_id)
        saved = self.save_drawing_takeoff_items(takeoff_id, suggested_items)

        tot_cost = sum(float(i.get("total_amount_vnd") or 0) for i in suggested_items)
        tot_concrete = sum(
            float(i.get("quantity") or 0)
            for i in suggested_items
            if i.get("category") == "concrete"
        )
        tot_rebar = sum(
            float(i.get("quantity") or 0)
            for i in suggested_items
            if i.get("category") == "rebar"
        )
        tot_formwork = sum(
            float(i.get("quantity") or 0)
            for i in suggested_items
            if i.get("category") == "formwork"
        )
        tot_earthwork = sum(
            float(i.get("quantity") or 0)
            for i in suggested_items
            if i.get("category") == "earthwork"
        )

        self.update_drawing_takeoff(
            takeoff_id,
            {
                "total_estimated_cost_vnd": tot_cost,
                "total_concrete_volume_m3": tot_concrete,
                "total_rebar_weight_tons": tot_rebar,
                "total_formwork_area_m2": tot_formwork,
                "total_earthwork_volume_m3": tot_earthwork,
            },
        )

        if message_id:
            with self.get_connection() as conn:
                with conn.cursor() as cur:
                    cur.execute(
                        "UPDATE erp_ai_takeoff_chat_messages SET applied = TRUE WHERE id = %s;",
                        (message_id,),
                    )
                    conn.commit()

        return {
            "status": "success",
            "takeoff_id": takeoff_id,
            "updated_count": len(saved),
            "total_estimated_cost_vnd": tot_cost,
        }
