"""Extend purchase orders for Create Purchase Order form.

Revision ID: 015_purchase_order_form
Revises: 014_vendor_form_fields
Create Date: 2026-08-06

"""

from typing import Sequence, Union

import sqlalchemy as sa
from alembic import op

revision: str = "015_purchase_order_form"
down_revision: Union[str, None] = "014_vendor_form_fields"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None


def upgrade() -> None:
    conn = op.get_bind()
    inspector = sa.inspect(conn)
    po_cols = {col["name"] for col in inspector.get_columns("purchase_orders")}
    line_cols = {col["name"] for col in inspector.get_columns("purchase_order_lines")}

    po_adds = [
        ("delivery_date", sa.Column("delivery_date", sa.Date(), nullable=True)),
        ("payment_terms", sa.Column("payment_terms", sa.String(length=30), nullable=True)),
        ("warehouse_id", sa.Column("warehouse_id", sa.String(length=36), nullable=True)),
        ("ship_to", sa.Column("ship_to", sa.String(length=255), nullable=True)),
        ("purchase_type", sa.Column("purchase_type", sa.String(length=30), nullable=True)),
        ("reference_number", sa.Column("reference_number", sa.String(length=100), nullable=True)),
        ("department", sa.Column("department", sa.String(length=100), nullable=True)),
        ("terms_and_conditions", sa.Column("terms_and_conditions", sa.Text(), nullable=True)),
        ("attachments", sa.Column("attachments", sa.Text(), nullable=True)),
        (
            "discount_amount",
            sa.Column(
                "discount_amount",
                sa.Numeric(18, 2),
                nullable=False,
                server_default="0.00",
            ),
        ),
    ]
    for name, column in po_adds:
        if name not in po_cols:
            op.add_column("purchase_orders", column)

    line_adds = [
        ("base_unit_id", sa.Column("base_unit_id", sa.String(length=36), nullable=True)),
        (
            "discount_type",
            sa.Column(
                "discount_type",
                sa.String(length=20),
                nullable=False,
                server_default="percent",
            ),
        ),
        (
            "discount_value",
            sa.Column(
                "discount_value",
                sa.Numeric(18, 4),
                nullable=False,
                server_default="0.0000",
            ),
        ),
        (
            "discount_amount",
            sa.Column(
                "discount_amount",
                sa.Numeric(18, 2),
                nullable=False,
                server_default="0.00",
            ),
        ),
        (
            "tax_type",
            sa.Column(
                "tax_type",
                sa.String(length=20),
                nullable=False,
                server_default="percent",
            ),
        ),
        (
            "tax_amount",
            sa.Column(
                "tax_amount",
                sa.Numeric(18, 2),
                nullable=False,
                server_default="0.00",
            ),
        ),
    ]
    for name, column in line_adds:
        if name not in line_cols:
            op.add_column("purchase_order_lines", column)

    # Backfill tax_amount from existing tax_rate percent on base amount
    op.execute(
        sa.text(
            """
            UPDATE purchase_order_lines
            SET tax_amount = ROUND((quantity * unit_price) * tax_rate / 100, 2)
            WHERE tax_amount = 0 AND tax_rate <> 0
            """
        )
    )

    op.alter_column("purchase_orders", "discount_amount", server_default=None)
    op.alter_column("purchase_order_lines", "discount_type", server_default=None)
    op.alter_column("purchase_order_lines", "discount_value", server_default=None)
    op.alter_column("purchase_order_lines", "discount_amount", server_default=None)
    op.alter_column("purchase_order_lines", "tax_type", server_default=None)
    op.alter_column("purchase_order_lines", "tax_amount", server_default=None)

    indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("purchase_orders")}
    if "ix_po_warehouse_id" not in indexes:
        op.create_index("ix_po_warehouse_id", "purchase_orders", ["warehouse_id"], unique=False)

    line_indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("purchase_order_lines")}
    if "ix_po_lines_base_unit_id" not in line_indexes:
        op.create_index(
            "ix_po_lines_base_unit_id",
            "purchase_order_lines",
            ["base_unit_id"],
            unique=False,
        )

    fks = {fk["name"] for fk in sa.inspect(conn).get_foreign_keys("purchase_orders")}
    if "fk_po_warehouse_id" not in fks:
        op.create_foreign_key(
            "fk_po_warehouse_id",
            "purchase_orders",
            "warehouses",
            ["warehouse_id"],
            ["id"],
            ondelete="SET NULL",
        )

    line_fks = {fk["name"] for fk in sa.inspect(conn).get_foreign_keys("purchase_order_lines")}
    if "fk_po_lines_base_unit_id" not in line_fks:
        op.create_foreign_key(
            "fk_po_lines_base_unit_id",
            "purchase_order_lines",
            "base_units",
            ["base_unit_id"],
            ["id"],
            ondelete="SET NULL",
        )


def downgrade() -> None:
    conn = op.get_bind()
    line_fks = {fk["name"] for fk in sa.inspect(conn).get_foreign_keys("purchase_order_lines")}
    if "fk_po_lines_base_unit_id" in line_fks:
        op.drop_constraint("fk_po_lines_base_unit_id", "purchase_order_lines", type_="foreignkey")
    fks = {fk["name"] for fk in sa.inspect(conn).get_foreign_keys("purchase_orders")}
    if "fk_po_warehouse_id" in fks:
        op.drop_constraint("fk_po_warehouse_id", "purchase_orders", type_="foreignkey")

    line_indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("purchase_order_lines")}
    if "ix_po_lines_base_unit_id" in line_indexes:
        op.drop_index("ix_po_lines_base_unit_id", table_name="purchase_order_lines")
    indexes = {idx["name"] for idx in sa.inspect(conn).get_indexes("purchase_orders")}
    if "ix_po_warehouse_id" in indexes:
        op.drop_index("ix_po_warehouse_id", table_name="purchase_orders")

    for col in (
        "tax_amount",
        "tax_type",
        "discount_amount",
        "discount_value",
        "discount_type",
        "base_unit_id",
    ):
        cols = {c["name"] for c in sa.inspect(conn).get_columns("purchase_order_lines")}
        if col in cols:
            op.drop_column("purchase_order_lines", col)

    for col in (
        "discount_amount",
        "attachments",
        "terms_and_conditions",
        "department",
        "reference_number",
        "purchase_type",
        "ship_to",
        "warehouse_id",
        "payment_terms",
        "delivery_date",
    ):
        cols = {c["name"] for c in sa.inspect(conn).get_columns("purchase_orders")}
        if col in cols:
            op.drop_column("purchase_orders", col)
