"""Extend unit_types to match Unit Type form.

Revision ID: 008_unit_type_fields
Revises: 007_item_group_item_type
Create Date: 2026-08-05

"""

from typing import Sequence, Union

import sqlalchemy as sa
from alembic import op

revision: str = "008_unit_type_fields"
down_revision: Union[str, None] = "007_item_group_item_type"
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)
    columns = {col["name"]: col for col in inspector.get_columns("unit_types")}

    if "pack_size" in columns and "conversion_rate" not in columns:
        op.alter_column(
            "unit_types",
            "pack_size",
            new_column_name="conversion_rate",
            existing_type=sa.Numeric(18, 4),
            existing_nullable=False,
        )
    elif "conversion_rate" not in columns:
        op.add_column(
            "unit_types",
            sa.Column(
                "conversion_rate",
                sa.Numeric(18, 4),
                nullable=False,
                server_default="1.0000",
            ),
        )
        op.alter_column("unit_types", "conversion_rate", server_default=None)

    if "decimal_places" not in columns:
        op.add_column(
            "unit_types",
            sa.Column("decimal_places", sa.Integer(), nullable=False, server_default="2"),
        )
        op.alter_column("unit_types", "decimal_places", server_default=None)

    if "sort_order" not in columns:
        op.add_column("unit_types", sa.Column("sort_order", sa.Integer(), nullable=True))

    if "description" not in columns:
        op.add_column(
            "unit_types", sa.Column("description", sa.String(length=255), nullable=True)
        )

    if "remarks" not in columns:
        op.add_column(
            "unit_types", sa.Column("remarks", sa.String(length=255), nullable=True)
        )

    indexes = {idx["name"] for idx in inspector.get_indexes("unit_types")}
    if "ix_unit_types_base_unit_id" not in indexes:
        op.create_index(
            "ix_unit_types_base_unit_id", "unit_types", ["base_unit_id"], unique=False
        )

    fks = {fk["name"] for fk in inspector.get_foreign_keys("unit_types")}
    if "fk_unit_types_base_unit_id" not in fks:
        op.create_foreign_key(
            "fk_unit_types_base_unit_id",
            "unit_types",
            "unit_types",
            ["base_unit_id"],
            ["id"],
            ondelete="SET NULL",
        )


def downgrade() -> None:
    op.drop_constraint("fk_unit_types_base_unit_id", "unit_types", type_="foreignkey")
    op.drop_index("ix_unit_types_base_unit_id", table_name="unit_types")
    op.drop_column("unit_types", "remarks")
    op.drop_column("unit_types", "description")
    op.drop_column("unit_types", "sort_order")
    op.drop_column("unit_types", "decimal_places")
    op.alter_column(
        "unit_types",
        "conversion_rate",
        new_column_name="pack_size",
        existing_type=sa.Numeric(18, 4),
        existing_nullable=False,
    )
