from __future__ import annotations

"""Storage handler for saving and parsing invoice records into DB."""


import json
import logging
from decimal import Decimal
from typing import Any

logger = logging.getLogger(__name__)


class InvoiceStorageSaveMixin:
    """Mixin for invoice insert and update operations."""

    def save_invoice(
        self,
        parsed: dict[str, Any],
        user_id: str | None = None,
        auto_reconcile: bool = True,
    ) -> dict[str, Any]:
        """Lưu hóa đơn và danh mục hàng hóa vào PostgreSQL trong một Atomic Transaction."""
        with self.get_connection() as conn:
            with conn.cursor() as cur:
                # 1. Check if hash already exists (Deduplication)
                cur.execute(
                    "SELECT id, invoice_number, invoice_series FROM erp_invoices WHERE xml_hash_sha256 = %s",
                    (parsed["xml_hash_sha256"],),
                )
                existing = cur.fetchone()
                if existing:
                    return {
                        "status": "duplicate",
                        "message": f"Hóa đơn số {existing['invoice_number']} ký hiệu {existing['invoice_series']} đã tồn tại trong hệ thống (Hash: {parsed['xml_hash_sha256'][:8]}...).",
                        "invoice_id": str(existing["id"]),
                    }

                # 2. Auto-match Project ID if possible
                matched_proj_id: str | None = None
                matched_wbs_id: str | None = None

                if auto_reconcile:
                    cur.execute(
                        "SELECT id, project_code, project_name, contract_number FROM projects WHERE status != 'archived' ORDER BY created_at DESC LIMIT 20"
                    )
                    projects = cur.fetchall()

                    search_corpus = f"{parsed['buyer_name']} {parsed['seller_name']} {parsed.get('notes', '')} "
                    for itm in parsed.get("items", []):
                        search_corpus += f"{itm['item_name']} "

                    for p in projects:
                        p_code = p["project_code"].upper()
                        p_name = p["project_name"].upper()
                        if p_code in search_corpus.upper() or any(
                            w in search_corpus.upper()
                            for w in p_name.split()
                            if len(w) > 4
                        ):
                            matched_proj_id = str(p["id"])
                            cur.execute(
                                "SELECT id FROM erp_project_wbs WHERE project_id = %s LIMIT 1",
                                (matched_proj_id,),
                            )
                            wbs_row = cur.fetchone()
                            if wbs_row:
                                matched_wbs_id = str(wbs_row["id"])
                            break

                reconciliation_status = (
                    "matched_project" if matched_proj_id else "unreconciled"
                )

                # 3. Insert Invoice
                insert_inv_sql = """
                    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, currency_code, exchange_rate,
                        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,
                        matched_project_id, matched_wbs_id, notes
                    ) VALUES (
                        %s, %s, %s, %s, %s,
                        %s, %s, %s, %s,
                        %s, %s, %s, %s, %s,
                        %s, %s, %s, %s,
                        %s, %s, %s, %s, %s,
                        %s, %s, %s, %s, %s,
                        %s, %s, %s
                    ) RETURNING id;
                """
                cur.execute(
                    insert_inv_sql,
                    (
                        parsed["direction"],
                        parsed["invoice_type"],
                        parsed["invoice_number"],
                        parsed["invoice_series"],
                        parsed["template_code"],
                        parsed["issue_date"],
                        parsed["seller_tax_code"],
                        parsed["seller_name"],
                        parsed.get("seller_address", ""),
                        parsed["buyer_tax_code"],
                        parsed["buyer_name"],
                        parsed.get("buyer_address", ""),
                        parsed.get("currency_code", "VND"),
                        parsed.get("exchange_rate", Decimal("1.0000")),
                        parsed["subtotal_amount_vnd"],
                        parsed["vat_rate_percent"],
                        parsed["vat_amount_vnd"],
                        parsed["total_amount_vnd"],
                        parsed.get("amount_in_words", ""),
                        parsed.get("status", "valid"),
                        parsed.get("tax_authority_code", ""),
                        parsed.get("xml_raw_content", ""),
                        parsed["xml_hash_sha256"],
                        parsed.get("signature_valid", True),
                        parsed.get("signed_by", ""),
                        parsed.get("signed_at"),
                        parsed.get("source_channel", "manual_xml"),
                        reconciliation_status,
                        matched_proj_id,
                        matched_wbs_id,
                        parsed.get("notes", ""),
                    ),
                )
                invoice_id = str(cur.fetchone()["id"])

                # 4. Insert Invoice Line Items
                insert_item_sql = """
                    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, matched_material_code, matched_wbs_code, cost_category
                    ) VALUES (
                        %s, %s, %s, %s, %s, %s,
                        %s, %s, %s, %s,
                        %s, %s, %s, %s
                    );
                """
                for item in parsed.get("items", []):
                    cost_cat = item.get("cost_category") or self.categorize_item_cost(
                        item.get("item_name", ""), item.get("item_code", "")
                    )
                    cur.execute(
                        insert_item_sql,
                        (
                            invoice_id,
                            item.get("item_order", 1),
                            item["item_name"],
                            item.get("item_code", ""),
                            item.get("unit", "Lô"),
                            item["quantity"],
                            item["unit_price_vnd"],
                            item["amount_before_vat_vnd"],
                            item["vat_rate_percent"],
                            item["vat_amount_vnd"],
                            item["total_item_amount_vnd"],
                            item.get("matched_material_code", ""),
                            item.get("matched_wbs_code", ""),
                            cost_cat,
                        ),
                    )

                # 5. Insert Audit Log
                insert_audit_sql = """
                    INSERT INTO erp_invoice_audit_logs (
                        invoice_id, action_type, performed_by, details_json
                    ) VALUES (%s, %s, %s, %s);
                """
                details = {
                    "source_channel": parsed.get("source_channel", "manual_xml"),
                    "total_amount_vnd": str(parsed["total_amount_vnd"]),
                    "matched_project_id": matched_proj_id,
                    "items_count": len(parsed.get("items", [])),
                }
                cur.execute(
                    insert_audit_sql,
                    (
                        invoice_id,
                        "IMPORT_XML",
                        user_id or "System / AI Processor",
                        json.dumps(details, ensure_ascii=False),
                    ),
                )

                conn.commit()
                return {
                    "status": "success",
                    "message": f"Đã nạp thành công hóa đơn số {parsed['invoice_number']} ({'Đầu Ra' if parsed['direction'] == 'output' else 'Đầu Vào'}).",
                    "invoice_id": invoice_id,
                    "matched_project_id": matched_proj_id,
                    "reconciliation_status": reconciliation_status,
                }
