from uuid import UUID
from typing import Optional
from app.modules.invoices.domain.entities import Invoice
from decimal import Decimal

class InvoiceRepository:
    """
    Kết nối Domain Entity Invoice xuống bảng PostgreSQL.
    Sử dụng schema 'invoices'.
    """
    def __init__(self, conn):
        self._conn = conn

    def ensure_schema(self):
        with self._conn.cursor() as cur:
            cur.execute("""
                CREATE SCHEMA IF NOT EXISTS invoices;
                CREATE TABLE IF NOT EXISTS invoices.invoices (
                    id UUID PRIMARY KEY,
                    invoice_no TEXT NOT NULL,
                    supplier_tax_code TEXT NOT NULL,
                    total_amount NUMERIC(18, 4) NOT NULL,
                    tax_amount NUMERIC(18, 4) NOT NULL,
                    issue_date DATE NOT NULL,
                    status TEXT NOT NULL,
                    project_id UUID,
                    created_at TIMESTAMP WITH TIME ZONE
                );
                CREATE INDEX IF NOT EXISTS idx_invoices_tax_code ON invoices.invoices(supplier_tax_code);
            """)

    def save(self, invoice: Invoice) -> None:
        self.ensure_schema()
        with self._conn.cursor() as cur:
            cur.execute("""
                INSERT INTO invoices.invoices 
                (id, invoice_no, supplier_tax_code, total_amount, tax_amount, issue_date, status, project_id, created_at)
                VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)
                ON CONFLICT (id) DO UPDATE SET
                    invoice_no = EXCLUDED.invoice_no,
                    supplier_tax_code = EXCLUDED.supplier_tax_code,
                    total_amount = EXCLUDED.total_amount,
                    tax_amount = EXCLUDED.tax_amount,
                    issue_date = EXCLUDED.issue_date,
                    status = EXCLUDED.status,
                    project_id = EXCLUDED.project_id;
            """, (
                str(invoice.id),
                invoice.invoice_no,
                invoice.supplier_tax_code,
                invoice.total_amount,
                invoice.tax_amount,
                invoice.issue_date,
                invoice.status,
                str(invoice.project_id) if invoice.project_id else None,
                invoice.created_at
            ))

    def get_by_id(self, invoice_id: UUID) -> Optional[Invoice]:
        with self._conn.cursor() as cur:
            cur.execute("SELECT * FROM invoices.invoices WHERE id = %s", (str(invoice_id),))
            row = cur.fetchone()
            if row:
                return Invoice(**row)
        return None
