"""production authentication lifecycle

Revision ID: d7e5a9206c13
Revises: b4b19c645cff
"""
from collections.abc import Sequence

import sqlalchemy as sa
from alembic import op

revision: str = "d7e5a9206c13"
down_revision: str | Sequence[str] | None = "b4b19c645cff"
branch_labels = None
depends_on = None


def upgrade() -> None:
    op.add_column(
        "users", sa.Column("auth_version", sa.Integer(), server_default="1", nullable=False)
    )
    op.add_column("users", sa.Column("password_changed_at", sa.DateTime(timezone=True)))
    op.add_column("users", sa.Column("mfa_secret_ciphertext", sa.LargeBinary()))
    op.add_column("users", sa.Column("mfa_key_version", sa.Integer()))
    op.add_column("users", sa.Column("mfa_enabled_at", sa.DateTime(timezone=True)))
    op.add_column("users", sa.Column("mfa_pending_ciphertext", sa.LargeBinary()))
    op.add_column("users", sa.Column("mfa_pending_key_version", sa.Integer()))
    op.add_column("users", sa.Column("mfa_pending_expires_at", sa.DateTime(timezone=True)))
    op.add_column("users", sa.Column("mfa_last_used_step", sa.BigInteger()))
    with op.batch_alter_table("users") as batch:
        batch.create_check_constraint("ck_user_auth_version", "auth_version > 0")

    op.add_column("device_sessions", sa.Column("idle_expires_at", sa.DateTime(timezone=True)))
    op.add_column("device_sessions", sa.Column("last_rotated_at", sa.DateTime(timezone=True)))
    op.add_column(
        "device_sessions",
        sa.Column("session_version", sa.Integer(), server_default="1", nullable=False),
    )
    op.execute(
        sa.text(
            "UPDATE device_sessions SET idle_expires_at = expires_at, "
            "last_rotated_at = created_at"
        )
    )
    with op.batch_alter_table("device_sessions") as batch:
        batch.alter_column("idle_expires_at", existing_type=sa.DateTime(timezone=True), nullable=False)
        batch.alter_column("last_rotated_at", existing_type=sa.DateTime(timezone=True), nullable=False)

    op.create_table(
        "mfa_recovery_codes",
        sa.Column("id", sa.Uuid(), nullable=False),
        sa.Column("user_id", sa.Uuid(), nullable=False),
        sa.Column("code_hash", sa.String(length=64), nullable=False),
        sa.Column(
            "created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
        ),
        sa.Column("used_at", sa.DateTime(timezone=True)),
        sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
        sa.PrimaryKeyConstraint("id"),
        sa.UniqueConstraint("code_hash"),
    )
    op.create_index(
        "ix_mfa_recovery_codes_user_id", "mfa_recovery_codes", ["user_id"], unique=False
    )
    op.create_index(
        "ix_mfa_recovery_user_used",
        "mfa_recovery_codes",
        ["user_id", "used_at"],
        unique=False,
    )


def downgrade() -> None:
    op.drop_index("ix_mfa_recovery_user_used", table_name="mfa_recovery_codes")
    op.drop_index("ix_mfa_recovery_codes_user_id", table_name="mfa_recovery_codes")
    op.drop_table("mfa_recovery_codes")

    with op.batch_alter_table("device_sessions") as batch:
        batch.drop_column("session_version")
        batch.drop_column("last_rotated_at")
        batch.drop_column("idle_expires_at")

    with op.batch_alter_table("users") as batch:
        batch.drop_constraint("ck_user_auth_version", type_="check")
        batch.drop_column("mfa_last_used_step")
        batch.drop_column("mfa_pending_expires_at")
        batch.drop_column("mfa_pending_key_version")
        batch.drop_column("mfa_pending_ciphertext")
        batch.drop_column("mfa_enabled_at")
        batch.drop_column("mfa_key_version")
        batch.drop_column("mfa_secret_ciphertext")
        batch.drop_column("password_changed_at")
        batch.drop_column("auth_version")
