import hashlib
import json
import logging
from datetime import date, datetime
from decimal import ROUND_HALF_UP, Decimal
from typing import Any

from .common import to_decimal

logger = logging.getLogger(__name__)


class InvoiceStorageGdtRawMixin:
    """Mixin for persisting GDT raw invoice records."""

    def _persist_gdt_raw_invoice(
        self, raw: dict[str, Any], fallback_direction: str, my_tax_code: str
    ) -> str:
        """Parse và lưu trữ bản ghi hóa đơn từ API Cổng Thuế vào CSDL PostgreSQL."""
        seller_tax_code = str(
            raw.get("nbmst") or raw.get("seller_tax_code") or ""
        ).strip()
        seller_name = str(
            raw.get("nbten") or raw.get("seller_name") or "Đơn vị bán"
        ).strip()
        buyer_tax_code = str(
            raw.get("nmmst") or raw.get("buyer_tax_code") or ""
        ).strip()
        buyer_name = str(
            raw.get("nmten") or raw.get("buyer_name") or "Đơn vị mua"
        ).strip()

        invoice_series = str(
            raw.get("khhdon") or raw.get("invoice_series") or "C24TAA"
        ).strip()
        inv_num_raw = str(raw.get("shdon") or raw.get("invoice_number") or "1").strip()
        invoice_number = inv_num_raw.zfill(8) if inv_num_raw.isdigit() else inv_num_raw
        template_code = str(
            raw.get("khmshdon") or raw.get("template_code") or "1"
        ).strip()

        tdlap = str(raw.get("tdlap") or raw.get("issue_date") or "").strip()
        try:
            if "T" in tdlap:
                issue_date = datetime.fromisoformat(tdlap.split("+")[0]).date()
            elif "/" in tdlap:
                issue_date = datetime.strptime(tdlap, "%d/%m/%Y").date()
            else:
                issue_date = datetime.strptime(tdlap[:10], "%Y-%m-%d").date()
        except Exception:
            issue_date = date.today()

        total_val = (
            raw.get("tgtttbso") or raw.get("tgtttgthue") or raw.get("total_amount_vnd")
        )
        vat_val = raw.get("tgtthue") or raw.get("vat_amount_vnd") or 0

        total_amount = to_decimal(total_val)
        vat_amount = to_decimal(vat_val)

        thttltsuat = raw.get("thttltsuat") or []
        if thttltsuat and isinstance(thttltsuat, list) and len(thttltsuat) > 0:
            subtotal = sum(to_decimal(b.get("thtien", 0)) for b in thttltsuat)
            calc_vat = sum(to_decimal(b.get("tthue", 0)) for b in thttltsuat)
            if calc_vat > 0:
                vat_amount = calc_vat
        elif raw.get("tgtttbthue"):
            subtotal = to_decimal(raw.get("tgtttbthue"))
        else:
            subtotal = (
                (total_amount - vat_amount)
                if total_amount >= vat_amount
                else Decimal("0.0000")
            )

        if total_amount == Decimal("0.0000") and (subtotal > 0 or vat_amount > 0):
            total_amount = subtotal + vat_amount

        vat_rate = Decimal("8.00")
        if subtotal > 0 and vat_amount > 0:
            vat_rate = ((vat_amount / subtotal) * Decimal(100)).quantize(
                Decimal("0.01"), rounding=ROUND_HALF_UP
            )
        elif subtotal > 0 and vat_amount == 0:
            vat_rate = Decimal("0.00")

        if buyer_tax_code and my_tax_code and buyer_tax_code == my_tax_code:
            direction = "input"
        elif seller_tax_code and my_tax_code and seller_tax_code == my_tax_code:
            direction = "output"
        else:
            direction = fallback_direction

        hash_seed = f"{seller_tax_code}_{invoice_series}_{invoice_number}_{total_amount}_{template_code}"
        xml_hash = hashlib.sha256(hash_seed.encode("utf-8")).hexdigest()

        tax_authority_code = str(
            raw.get("hsgcma") or raw.get("tax_authority_code") or ""
        )

        status_map = {1: "valid", 2: "replaced", 3: "adjusted", 4: "cancelled"}
        tthai = raw.get("tthai", 1)
        inv_status = status_map.get(tthai, "valid")

        # Extract true signing timestamp
        signed_at_val = None
        if raw.get("nky"):
            try:
                signed_at_val = datetime.fromisoformat(
                    str(raw["nky"]).replace("Z", "+00:00")
                )
            except Exception:
                pass
        if (
            not signed_at_val
            and raw.get("nbcks")
            and "SigningTime" in str(raw["nbcks"])
        ):
            try:
                nbcks_dict = (
                    json.loads(raw["nbcks"])
                    if isinstance(raw["nbcks"], str)
                    else raw["nbcks"]
                )
                if nbcks_dict.get("SigningTime"):
                    signed_at_val = datetime.fromisoformat(
                        str(nbcks_dict["SigningTime"]).replace("Z", "+00:00")
                    )
            except Exception:
                pass
        if not signed_at_val:
            signed_at_val = datetime(issue_date.year, issue_date.month, issue_date.day)

        with self.get_connection() as conn:
            with conn.cursor() as cur:
                cur.execute(
                    """
                    SELECT id FROM erp_invoices 
                    WHERE xml_hash_sha256 = %s 
                       OR (seller_tax_code = %s AND invoice_series = %s AND invoice_number = %s)
                    LIMIT 1;
                    """,
                    (xml_hash, seller_tax_code, invoice_series, invoice_number),
                )
                if cur.fetchone():
                    return "dup"

                try:
                    cur.execute(
                        """
                        INSERT INTO erp_invoices (
                            direction, invoice_type, invoice_number, invoice_series, template_code,
                            issue_date, seller_tax_code, seller_name, seller_address,
                            buyer_tax_code, buyer_name, buyer_address,
                            subtotal_amount_vnd, vat_rate_percent, vat_amount_vnd, total_amount_vnd,
                            amount_in_words, status, tax_authority_code, xml_raw_content, xml_hash_sha256,
                            signature_valid, signed_by, signed_at, source_channel, reconciliation_status
                        ) VALUES (
                            %s, 'vat', %s, %s, %s,
                            %s, %s, %s, %s,
                            %s, %s, %s,
                            %s, %s, %s, %s,
                            %s, %s, %s, %s, %s,
                            TRUE, %s, %s, 'gdt_sync', 'unreconciled'
                        ) ON CONFLICT (seller_tax_code, invoice_series, invoice_number) DO NOTHING
                        RETURNING id;
                    """,
                        (
                            direction,
                            invoice_number,
                            invoice_series,
                            template_code,
                            issue_date,
                            seller_tax_code,
                            seller_name,
                            raw.get("nbdchi", ""),
                            buyer_tax_code,
                            buyer_name,
                            raw.get("nmdchi", ""),
                            subtotal,
                            vat_rate,
                            vat_amount,
                            total_amount,
                            raw.get("amount_in_words", ""),
                            inv_status,
                            tax_authority_code,
                            json.dumps(raw, ensure_ascii=False),
                            xml_hash,
                            seller_name,
                            signed_at_val,
                        ),
                    )
                    row = cur.fetchone()
                    if not row:
                        return "dup"
                    new_id = (
                        row["id"]
                        if (row and isinstance(row, dict))
                        else (row[0] if row else None)
                    )
                except Exception as exc:
                    if "duplicate key" in str(exc).lower() or "unique constraint" in str(exc).lower():
                        return "dup"
                    raise

                # Insert genuine line items or synthesize summary lines from GDT tax breakdown (Single Source of Truth)
                raw_items = raw.get("hdhhdvu")
                if not raw_items or not isinstance(raw_items, list) or len(raw_items) == 0:
                    thttltsuat = raw.get("thttltsuat") or []
                    if thttltsuat and isinstance(thttltsuat, list) and len(thttltsuat) > 0:
                        raw_items = []
                        for idx, vb in enumerate(thttltsuat, start=1):
                            t_rate_str = str(vb.get("tsuat") or "8%").replace("%", "").strip()
                            t_amt_before = to_decimal(vb.get("thtien"), default=Decimal("0.0000"))
                            t_vat = to_decimal(vb.get("tthue"), default=Decimal("0.0000"))
                            raw_items.append({
                                "stt": idx,
                                "ten": f"Cung cấp hàng hóa, dịch vụ theo HĐ {invoice_number} ({invoice_series})",
                                "mhhdvu": f"GDT-SUM-{idx:02d}",
                                "dvtinh": "Gói",
                                "sluong": Decimal("1.0000"),
                                "dgia": t_amt_before,
                                "thtien": t_amt_before,
                                "ltsuat": f"{t_rate_str}%",
                                "tthue": t_vat,
                            })
                    elif subtotal > 0 or total_amount > 0:
                        raw_items = [{
                            "stt": 1,
                            "ten": f"Cung cấp hàng hóa, dịch vụ theo HĐ {invoice_number} ({invoice_series})",
                            "mhhdvu": "GDT-SUM-01",
                            "dvtinh": "Gói",
                            "sluong": Decimal("1.0000"),
                            "dgia": subtotal if subtotal > 0 else total_amount,
                            "thtien": subtotal if subtotal > 0 else total_amount,
                            "ltsuat": f"{vat_rate}%",
                            "tthue": vat_amount,
                        }]

                if raw_items and isinstance(raw_items, list) and len(raw_items) > 0:
                    for itm in raw_items:
                        stt = int(itm.get("stt") or 1)
                        name = str(itm.get("ten") or "Hàng hóa / Dịch vụ").strip()
                        code = str(itm.get("mhhdvu") or "").strip()
                        unit = str(itm.get("dvtinh") or "Gói").strip()
                        qty = to_decimal(itm.get("sluong"), default=Decimal("0.0000"))
                        unit_price = to_decimal(
                            itm.get("dgia"), default=Decimal("0.0000")
                        )
                        amount_before = to_decimal(
                            itm.get("thtien"), default=Decimal("0.0000")
                        )
                        vat_str = (
                            str(itm.get("ltsuat") or "8%").replace("%", "").strip()
                        )
                        try:
                            v_rate = to_decimal(vat_str, default=Decimal("8.0000"))
                        except Exception:
                            v_rate = Decimal("8.0000")
                        vat_amt = (
                            to_decimal(itm.get("tthue"))
                            if itm.get("tthue") is not None
                            else ((amount_before * v_rate) / Decimal(100)).quantize(
                                Decimal("0.0001"), rounding=ROUND_HALF_UP
                            )
                        )
                        tot_amt = amount_before + vat_amt
                        cost_cat = self.categorize_item_cost(name, code)

                        cur.execute(
                            """
                            INSERT INTO erp_invoice_items (
                                invoice_id, item_order, item_name, item_code, unit, quantity,
                                unit_price_vnd, amount_before_vat_vnd, vat_rate_percent, vat_amount_vnd, total_item_amount_vnd, cost_category
                            ) VALUES (
                                %s, %s, %s, %s, %s, %s,
                                %s, %s, %s, %s, %s, %s
                            );
                        """,
                            (
                                new_id,
                                stt,
                                name,
                                code,
                                unit,
                                qty,
                                unit_price,
                                amount_before,
                                v_rate,
                                vat_amt,
                                tot_amt,
                                cost_cat,
                            ),
                        )

                cur.execute(
                    """
                    INSERT INTO erp_invoice_audit_logs (invoice_id, action_type, performed_by, details_json)
                    VALUES (%s, 'GDT_SYNC_IMPORT', 'SystemBot', %s);
                """,
                    (
                        new_id,
                        json.dumps(
                            {
                                "source": "hoadondientu.gdt.gov.vn",
                                "raw_id": raw.get("id", ""),
                            },
                            ensure_ascii=False,
                        ),
                    ),
                )

                conn.commit()
                return "new"
