from __future__ import annotations

"""XML Parser for Decree 123 General Department of Taxation format."""


import logging
import re
import xml.etree.ElementTree as ET
from datetime import date, datetime
from decimal import ROUND_HALF_UP, Decimal
from typing import Any

from .common import DSCONS_DEFAULT_TAX_CODE, to_decimal

logger = logging.getLogger(__name__)


class InvoiceGdtXmlMixin:
    """Mixin for parsing GDT ND123 XML invoices."""

    def _parse_gdt_nd123_xml(
        self, root: ET.Element, raw_str: str, sha256_hash: str, source_channel: str
    ) -> dict[str, Any]:
        """Bóc tách cấu trúc XML chuẩn Nghị định 123 / Thông tư 78 của Tổng Cục Thuế."""

        def get_text(parent: ET.Element | None, tag: str, default: str = "") -> str:
            if parent is None:
                return default
            el = parent.find(f".//{tag}")
            return el.text.strip() if el is not None and el.text else default

        tt_chung = root.find(".//TTChung")
        nd_hdon = root.find(".//NDHDon")
        n_ban = nd_hdon.find(".//NBan") if nd_hdon is not None else root.find(".//NBan")
        n_mua = nd_hdon.find(".//NMua") if nd_hdon is not None else root.find(".//NMua")
        t_toan = root.find(".//TToan")

        invoice_number = get_text(tt_chung, "SHDon") or get_text(
            root, "SHDon", "00000000"
        )
        if invoice_number.isdigit() and len(invoice_number) < 8:
            invoice_number = invoice_number.zfill(8)

        invoice_series = get_text(tt_chung, "KHHDon") or get_text(
            root, "KHHDon", "1C26TDS"
        )
        template_code = get_text(tt_chung, "KHMSHDon") or get_text(
            root, "KHMSHDon", "1"
        )
        issue_date_str = get_text(tt_chung, "NLap") or get_text(
            root, "NLap", str(date.today())
        )

        try:
            if "T" in issue_date_str:
                issue_date = datetime.fromisoformat(
                    issue_date_str.replace("Z", "+00:00")
                ).date()
            elif "/" in issue_date_str:
                issue_date = datetime.strptime(issue_date_str, "%d/%m/%Y").date()
            else:
                issue_date = datetime.strptime(issue_date_str[:10], "%Y-%m-%d").date()
        except Exception:
            issue_date = date.today()

        seller_tax_code = get_text(n_ban, "MST", "")
        seller_name = get_text(n_ban, "Ten", "Người Bán Chưa Xác Định")
        seller_address = get_text(n_ban, "DChi", "")

        buyer_tax_code = get_text(n_mua, "MST", "")
        buyer_name = get_text(n_mua, "Ten", "Người Mua Chưa Xác Định")
        buyer_address = get_text(n_mua, "DChi", "")

        currency_code = get_text(tt_chung, "DVTTe", "VND")
        exchange_rate = to_decimal(get_text(tt_chung, "TGia", "1.0000"))
        subtotal_amount = to_decimal(
            get_text(t_toan, "TgTCThue") or get_text(root, "TgTCThue", "0")
        )
        vat_amount = to_decimal(
            get_text(t_toan, "TgTThue") or get_text(root, "TgTThue", "0")
        )
        total_amount = to_decimal(
            get_text(t_toan, "TgTTTBSo") or get_text(root, "TgTTTBSo", "0")
        )
        amount_in_words = get_text(t_toan, "TgTTTBChu") or get_text(
            root, "TgTTTBChu", ""
        )
        tax_authority_code = get_text(root, "MCCQT") or get_text(tt_chung, "MCCQT", "")

        if (
            seller_tax_code == DSCONS_DEFAULT_TAX_CODE
            or "ĐỊNH SƠN" in seller_name.upper()
        ):
            direction = "output"
        else:
            direction = "input"

        is_mtt = self.is_cash_register_invoice(invoice_series, template_code)
        sig_element = root.find(".//Signature")
        signature_valid = sig_element is not None or is_mtt

        if is_mtt and not sig_element:
            signed_by = f"Mã CQT xác thực máy tính tiền (Điều 11 NĐ123) - {seller_name}"
            signed_at = datetime(issue_date.year, issue_date.month, issue_date.day)
        else:
            signed_by = get_text(sig_element, "Subject") or (
                seller_name if direction == "output" else seller_name
            )
            signed_at_str = get_text(sig_element, "SigningTime")
            try:
                signed_at = (
                    datetime.fromisoformat(signed_at_str.replace("Z", "+00:00"))
                    if signed_at_str
                    else datetime(issue_date.year, issue_date.month, issue_date.day)
                )
            except Exception:
                signed_at = datetime(issue_date.year, issue_date.month, issue_date.day)

        items: list[dict[str, Any]] = []
        item_nodes = (
            root.findall(".//HHDVu")
            or root.findall(".//HHDVuMTT")
            or root.findall(".//Item")
        )
        item_idx = 1
        calc_subtotal = Decimal("0.0000")
        calc_vat = Decimal("0.0000")

        for node in item_nodes:
            item_name = get_text(node, "THHVu", f"Hàng hóa / Dịch vụ #{item_idx}")
            item_code = get_text(node, "MHHDVu", "")
            unit = get_text(node, "DVTinh", "Lô")
            quantity = to_decimal(get_text(node, "SLuong", "1.0000"))
            unit_price = to_decimal(get_text(node, "DGia", "0.0000"))
            line_amount = to_decimal(get_text(node, "ThTien", "0.0000"))

            if line_amount == Decimal("0.0000") and unit_price > Decimal("0.0000"):
                line_amount = (quantity * unit_price).quantize(
                    Decimal("0.0001"), rounding=ROUND_HALF_UP
                )

            vat_rate_str = get_text(node, "TSuat", "8%")
            vat_rate_num = to_decimal(
                re.sub(r"[^\d.]", "", vat_rate_str), Decimal("8.00")
            )
            line_vat = to_decimal(get_text(node, "TThue", ""))
            if line_vat == Decimal("0.0000"):
                line_vat = (line_amount * vat_rate_num / Decimal("100.00")).quantize(
                    Decimal("0.0001"), rounding=ROUND_HALF_UP
                )

            line_total = line_amount + line_vat
            calc_subtotal += line_amount
            calc_vat += line_vat
            cost_cat = self.categorize_item_cost(item_name, item_code)

            items.append(
                {
                    "item_order": item_idx,
                    "item_name": item_name,
                    "item_code": item_code,
                    "unit": unit,
                    "quantity": quantity,
                    "unit_price_vnd": unit_price,
                    "amount_before_vat_vnd": line_amount,
                    "vat_rate_percent": vat_rate_num,
                    "vat_amount_vnd": line_vat,
                    "total_item_amount_vnd": line_total,
                    "cost_category": cost_cat,
                }
            )
            item_idx += 1

        if subtotal_amount == Decimal("0.0000") and calc_subtotal > Decimal("0.0000"):
            subtotal_amount = calc_subtotal
        if vat_amount == Decimal("0.0000") and calc_vat > Decimal("0.0000"):
            vat_amount = calc_vat
        if total_amount == Decimal("0.0000"):
            total_amount = subtotal_amount + vat_amount

        avg_vat_rate = Decimal("8.00")
        if subtotal_amount > Decimal("0.0000"):
            avg_vat_rate = (
                (vat_amount / subtotal_amount) * Decimal("100.00")
            ).quantize(Decimal("0.01"), rounding=ROUND_HALF_UP)

        return {
            "direction": direction,
            "invoice_type": "vat",
            "invoice_number": invoice_number,
            "invoice_series": invoice_series,
            "template_code": template_code,
            "issue_date": str(issue_date),
            "seller_tax_code": seller_tax_code,
            "seller_name": seller_name,
            "seller_address": seller_address,
            "buyer_tax_code": buyer_tax_code,
            "buyer_name": buyer_name,
            "buyer_address": buyer_address,
            "currency_code": currency_code,
            "exchange_rate": exchange_rate,
            "subtotal_amount_vnd": subtotal_amount,
            "vat_rate_percent": avg_vat_rate,
            "vat_amount_vnd": vat_amount,
            "total_amount_vnd": total_amount,
            "amount_in_words": amount_in_words,
            "status": "valid",
            "tax_authority_code": tax_authority_code,
            "xml_raw_content": raw_str,
            "xml_hash_sha256": sha256_hash,
            "signature_valid": signature_valid,
            "signed_by": signed_by,
            "signed_at": signed_at.isoformat(),
            "source_channel": source_channel,
            "reconciliation_status": "unreconciled",
            "items": items,
        }
