"""Allow refresh tokens for company_users as well as user_registrations.

Revision ID: 032_refresh_tokens_company_users
Revises: 031_company_users_dashboard
Create Date: 2026-08-15
"""

from typing import Sequence, Union

import sqlalchemy as sa
from alembic import op

revision: str = "032_refresh_tokens_company_users"
down_revision: Union[str, None] = "031_company_users_dashboard"
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)
    tables = set(inspector.get_table_names())
    if "refresh_tokens" not in tables:
        return

    for fk in inspector.get_foreign_keys("refresh_tokens"):
        if "user_id" in (fk.get("constrained_columns") or []):
            op.drop_constraint(fk["name"], "refresh_tokens", type_="foreignkey")


def downgrade() -> None:
    conn = op.get_bind()
    inspector = sa.inspect(conn)
    tables = set(inspector.get_table_names())
    if "refresh_tokens" not in tables or "user_registrations" not in tables:
        return
    op.create_foreign_key(
        "refresh_tokens_ibfk_1",
        "refresh_tokens",
        "user_registrations",
        ["user_id"],
        ["id"],
        ondelete="CASCADE",
    )
