"""allow immediate terminal compromise of staged signing keys

Revision ID: 4d8a2f6c9b31
Revises: f3a7c9d52b18
Create Date: 2026-08-16 23:45:00
"""

from datetime import timedelta
from typing import Sequence, Union

from alembic import op
import sqlalchemy as sa


revision: str = "4d8a2f6c9b31"
down_revision: Union[str, Sequence[str], None] = "f3a7c9d52b18"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None


def upgrade() -> None:
    with op.batch_alter_table("signing_key_metadata") as batch_op:
        batch_op.drop_constraint(
            "ck_signing_key_metadata_validity",
            type_="check",
        )
        batch_op.create_check_constraint(
            "ck_signing_key_metadata_validity",
            "status = 'compromised' OR expires_at IS NULL OR "
            "expires_at > not_before",
        )


def downgrade() -> None:
    signing_keys = sa.table(
        "signing_key_metadata",
        sa.column("id", sa.String()),
        sa.column("status", sa.String()),
        sa.column("not_before", sa.DateTime(timezone=True)),
        sa.column("expires_at", sa.DateTime(timezone=True)),
    )
    connection = op.get_bind()
    incompatible_rows = connection.execute(
        sa.select(
            signing_keys.c.id,
            signing_keys.c.expires_at,
        ).where(
            signing_keys.c.status == "compromised",
            signing_keys.c.expires_at <= signing_keys.c.not_before,
        )
    ).all()
    for row in incompatible_rows:
        connection.execute(
            signing_keys.update()
            .where(signing_keys.c.id == row.id)
            .values(not_before=row.expires_at - timedelta(microseconds=1))
        )
    with op.batch_alter_table("signing_key_metadata") as batch_op:
        batch_op.drop_constraint(
            "ck_signing_key_metadata_validity",
            type_="check",
        )
        batch_op.create_check_constraint(
            "ck_signing_key_metadata_validity",
            "expires_at IS NULL OR expires_at > not_before",
        )
