from __future__ import annotations

import logging
from decimal import Decimal

logger = logging.getLogger("dscons.postgres.project_crud.financial_mutations")


class FinancialMutationsMixin:
    """Financial transactions mutations (IPC, Retention, Advances)."""

    def create_ipc_financial_transactions(
        self,
        project_id: str,
        company_id: str,
        period_number: int,
        net_payable: Decimal,
        retention_amount: Decimal,
        advance_recovery_amount: Decimal,
    ) -> list[str]:
        """Tạo các giao dịch ghi nhận số thu IPC, tiền tạm ứng thu hồi và tiền giữ lại bảo hành."""

        transaction_ids = []
        sql = """
            INSERT INTO erp_financial_transactions (
                company_id, project_id, transaction_code, transaction_type, direction,
                amount, transaction_date, accounting_status, notes
            )
            VALUES (
                %(company_id)s, %(project_id)s, %(transaction_code)s, %(transaction_type)s, %(direction)s,
                %(amount)s, CURRENT_DATE, 'planned', %(notes)s
            )
            ON CONFLICT (transaction_code) DO NOTHING
            RETURNING id;
        """

        with self.get_connection() as conn, conn.cursor() as cur:
            # 1. Thu nhập chờ thu từ đợt IPC
            if net_payable > 0:
                code_net = f"IPC-NET-{str(project_id)[:8].upper()}-DOT{period_number}"
                cur.execute(
                    sql,
                    {
                        "company_id": company_id,
                        "project_id": project_id,
                        "transaction_code": code_net,
                        "transaction_type": "stage_payment",
                        "direction": "inflow",
                        "amount": net_payable,
                        "notes": f"Đề nghị thanh toán IPC đợt {period_number}",
                    },
                )
                row = cur.fetchone()
                if row:
                    transaction_ids.append(str(row["id"]))

            # 2. Tiền giữ lại bảo hành
            if retention_amount > 0:
                code_ret = f"IPC-RET-{str(project_id)[:8].upper()}-DOT{period_number}"
                cur.execute(
                    sql,
                    {
                        "company_id": company_id,
                        "project_id": project_id,
                        "transaction_code": code_ret,
                        "transaction_type": "retention_release",
                        "direction": "inflow",
                        "amount": retention_amount,
                        "notes": f"Trích lập giữ lại bảo hành IPC đợt {period_number}",
                    },
                )
                row = cur.fetchone()
                if row:
                    transaction_ids.append(str(row["id"]))

            # 3. Thu hồi tạm ứng
            if advance_recovery_amount > 0:
                code_adv = f"IPC-ADV-{str(project_id)[:8].upper()}-DOT{period_number}"
                cur.execute(
                    sql,
                    {
                        "company_id": company_id,
                        "project_id": project_id,
                        "transaction_code": code_adv,
                        "transaction_type": "advance_payment",
                        "direction": "inflow",
                        "amount": advance_recovery_amount,
                        "notes": f"Khấu trừ thu hồi tạm ứng đợt {period_number}",
                    },
                )
        return transaction_ids

    def list_project_ipcs(self, project_id: str) -> list[dict[str, Any]]:
        """Lấy danh sách các đợt tạm ứng, nghiệm thu thanh toán IPC."""
        sql = """
            SELECT i.*, inv.invoice_number, inv.invoice_series
            FROM erp_project_ipcs i
            LEFT JOIN erp_invoices inv ON i.matched_invoice_id = inv.id
            WHERE i.project_id = %s
            ORDER BY i.ipc_number ASC;
        """
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(sql, (project_id,))
            return cur.fetchall()

    def create_project_ipc(
        self,
        project_id: str,
        ipc_number: int,
        billing_period_start: Any,
        billing_period_end: Any,
        gross_claimed_amount_vnd: Decimal,
        advance_recovery_amount_vnd: Decimal,
        retention_withheld_amount_vnd: Decimal,
        net_certified_amount_vnd: Decimal,
        status: str,
        matched_invoice_id: Any | None = None,
    ) -> str:
        """Tạo đợt IPC thanh toán mới."""
        sql = """
            INSERT INTO erp_project_ipcs (
                project_id, ipc_number, billing_period_start, billing_period_end,
                gross_claimed_amount_vnd, advance_recovery_amount_vnd,
                retention_withheld_amount_vnd, net_certified_amount_vnd,
                status, matched_invoice_id
            ) VALUES (
                %s, %s, %s, %s, %s, %s, %s, %s, %s, %s
            ) RETURNING id;
        """
        with self.get_connection() as conn, conn.cursor() as cur:
            cur.execute(
                sql,
                (
                    project_id,
                    ipc_number,
                    billing_period_start,
                    billing_period_end,
                    gross_claimed_amount_vnd,
                    advance_recovery_amount_vnd,
                    retention_withheld_amount_vnd,
                    net_certified_amount_vnd,
                    status,
                    matched_invoice_id,
                ),
            )
            created_id = cur.fetchone()["id"]
            conn.commit()
            return str(created_id)

