"""Add warehouse_id to inventory balances and ledger transactions.

Revision ID: 022_inventory_warehouse_balances
Revises: 021_stock_transfers
Create Date: 2026-08-09
"""

from typing import Sequence, Union

import sqlalchemy as sa
from alembic import op

revision: str = "022_inventory_warehouse_balances"
down_revision: Union[str, None] = "021_stock_transfers"
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)

    # ---- inventory_balances ----
    bal_cols = {col["name"] for col in inspector.get_columns("inventory_balances")}
    bal_indexes = {idx["name"] for idx in inspector.get_indexes("inventory_balances")}
    bal_uniques = {
        uq["name"] for uq in inspector.get_unique_constraints("inventory_balances")
    }
    bal_fks = {fk["name"] for fk in inspector.get_foreign_keys("inventory_balances")}

    if "warehouse_id" not in bal_cols:
        op.add_column(
            "inventory_balances",
            sa.Column("warehouse_id", sa.String(length=36), nullable=True),
        )

    # Backfill from item default warehouse
    op.execute(
        sa.text(
            """
            UPDATE inventory_balances b
            INNER JOIN items i ON i.id = b.item_id
            SET b.warehouse_id = i.warehouse_id
            WHERE b.warehouse_id IS NULL OR b.warehouse_id = ''
            """
        )
    )

    # Drop any orphan balances that still have no warehouse
    op.execute(
        sa.text(
            """
            DELETE FROM inventory_balances
            WHERE warehouse_id IS NULL OR warehouse_id = ''
            """
        )
    )

    if "uq_inventory_balances_company_item" in bal_uniques:
        op.drop_constraint(
            "uq_inventory_balances_company_item",
            "inventory_balances",
            type_="unique",
        )

    op.alter_column(
        "inventory_balances",
        "warehouse_id",
        existing_type=sa.String(length=36),
        nullable=False,
    )

    inspector = sa.inspect(conn)
    bal_indexes = {idx["name"] for idx in inspector.get_indexes("inventory_balances")}
    bal_uniques = {
        uq["name"] for uq in inspector.get_unique_constraints("inventory_balances")
    }
    bal_fks = {fk["name"] for fk in inspector.get_foreign_keys("inventory_balances")}

    if "ix_inventory_balances_warehouse_id" not in bal_indexes:
        op.create_index(
            "ix_inventory_balances_warehouse_id",
            "inventory_balances",
            ["warehouse_id"],
            unique=False,
        )
    if "uq_inventory_balances_company_item_warehouse" not in bal_uniques:
        op.create_unique_constraint(
            "uq_inventory_balances_company_item_warehouse",
            "inventory_balances",
            ["company_id", "item_id", "warehouse_id"],
        )
    if "fk_inventory_balances_warehouse_id" not in bal_fks:
        op.create_foreign_key(
            "fk_inventory_balances_warehouse_id",
            "inventory_balances",
            "warehouses",
            ["warehouse_id"],
            ["id"],
            ondelete="RESTRICT",
        )

    # ---- inventory_transactions ----
    inspector = sa.inspect(conn)
    txn_cols = {col["name"] for col in inspector.get_columns("inventory_transactions")}
    txn_indexes = {idx["name"] for idx in inspector.get_indexes("inventory_transactions")}
    txn_fks = {fk["name"] for fk in inspector.get_foreign_keys("inventory_transactions")}

    if "warehouse_id" not in txn_cols:
        op.add_column(
            "inventory_transactions",
            sa.Column("warehouse_id", sa.String(length=36), nullable=True),
        )

    op.execute(
        sa.text(
            """
            UPDATE inventory_transactions t
            INNER JOIN items i ON i.id = t.item_id
            SET t.warehouse_id = i.warehouse_id
            WHERE t.warehouse_id IS NULL OR t.warehouse_id = ''
            """
        )
    )

    if "ix_inventory_transactions_warehouse_id" not in txn_indexes:
        op.create_index(
            "ix_inventory_transactions_warehouse_id",
            "inventory_transactions",
            ["warehouse_id"],
            unique=False,
        )
    if "fk_inventory_transactions_warehouse_id" not in txn_fks:
        op.create_foreign_key(
            "fk_inventory_transactions_warehouse_id",
            "inventory_transactions",
            "warehouses",
            ["warehouse_id"],
            ["id"],
            ondelete="SET NULL",
        )


def downgrade() -> None:
    conn = op.get_bind()
    inspector = sa.inspect(conn)

    txn_indexes = {idx["name"] for idx in inspector.get_indexes("inventory_transactions")}
    txn_fks = {fk["name"] for fk in inspector.get_foreign_keys("inventory_transactions")}
    txn_cols = {col["name"] for col in inspector.get_columns("inventory_transactions")}

    if "fk_inventory_transactions_warehouse_id" in txn_fks:
        op.drop_constraint(
            "fk_inventory_transactions_warehouse_id",
            "inventory_transactions",
            type_="foreignkey",
        )
    if "ix_inventory_transactions_warehouse_id" in txn_indexes:
        op.drop_index(
            "ix_inventory_transactions_warehouse_id",
            table_name="inventory_transactions",
        )
    if "warehouse_id" in txn_cols:
        op.drop_column("inventory_transactions", "warehouse_id")

    inspector = sa.inspect(conn)
    bal_indexes = {idx["name"] for idx in inspector.get_indexes("inventory_balances")}
    bal_uniques = {
        uq["name"] for uq in inspector.get_unique_constraints("inventory_balances")
    }
    bal_fks = {fk["name"] for fk in inspector.get_foreign_keys("inventory_balances")}
    bal_cols = {col["name"] for col in inspector.get_columns("inventory_balances")}

    if "fk_inventory_balances_warehouse_id" in bal_fks:
        op.drop_constraint(
            "fk_inventory_balances_warehouse_id",
            "inventory_balances",
            type_="foreignkey",
        )
    if "uq_inventory_balances_company_item_warehouse" in bal_uniques:
        op.drop_constraint(
            "uq_inventory_balances_company_item_warehouse",
            "inventory_balances",
            type_="unique",
        )
    if "ix_inventory_balances_warehouse_id" in bal_indexes:
        op.drop_index(
            "ix_inventory_balances_warehouse_id",
            table_name="inventory_balances",
        )
    if "warehouse_id" in bal_cols:
        op.drop_column("inventory_balances", "warehouse_id")

    op.create_unique_constraint(
        "uq_inventory_balances_company_item",
        "inventory_balances",
        ["company_id", "item_id"],
    )
