"""Import database/schema.sql into a disposable MySQL database and compare it to the ORM."""

from __future__ import annotations

import re
from pathlib import Path
from uuid import uuid4

from sqlalchemy import create_engine, inspect
from sqlalchemy.engine import make_url

from app.core.config import settings
from app.db.session import Base
import app.models.master_data  # noqa: F401
import app.models.procurement  # noqa: F401
import app.models.user  # noqa: F401


def _statements(sql: str) -> list[str]:
    without_comments = re.sub(r"(?m)^\s*--.*$", "", sql)
    return [statement.strip() for statement in without_comments.split(";") if statement.strip()]


def main() -> None:
    configured_url = make_url(settings.DATABASE_URL)
    if not configured_url.drivername.startswith("mysql"):
        raise RuntimeError("Schema validation requires a MySQL DATABASE_URL.")

    database_name = f"procurement_schema_validation_{uuid4().hex[:10]}"
    server_engine = create_engine(configured_url.set(database=None), pool_pre_ping=True)
    validation_engine = None
    created = False
    try:
        with server_engine.begin() as connection:
            existing = connection.exec_driver_sql("SHOW DATABASES LIKE %s", (database_name,)).first()
            if existing:
                raise RuntimeError(f"Refusing to reuse existing database {database_name}.")
            connection.exec_driver_sql(
                f"CREATE DATABASE `{database_name}` CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"
            )
            created = True

        validation_engine = create_engine(configured_url.set(database=database_name), pool_pre_ping=True)
        schema_path = Path(__file__).resolve().parents[1] / "database" / "schema.sql"
        with validation_engine.begin() as connection:
            for statement in _statements(schema_path.read_text(encoding="utf-8")):
                connection.exec_driver_sql(statement)

        inspector = inspect(validation_engine)
        actual_tables = set(inspector.get_table_names())
        expected_tables = set(Base.metadata.tables)
        errors: list[str] = []
        if expected_tables != actual_tables:
            errors.append(f"table mismatch: missing={sorted(expected_tables - actual_tables)}, extra={sorted(actual_tables - expected_tables)}")
        for table_name in sorted(expected_tables & actual_tables):
            expected_columns = set(Base.metadata.tables[table_name].columns.keys())
            actual_columns = {column["name"] for column in inspector.get_columns(table_name)}
            if expected_columns != actual_columns:
                errors.append(
                    f"{table_name}: missing={sorted(expected_columns - actual_columns)}, extra={sorted(actual_columns - expected_columns)}"
                )
        if errors:
            raise RuntimeError("Schema validation failed:\n" + "\n".join(errors))
        print(f"Schema validation passed: {len(actual_tables)} tables match the SQLAlchemy models.")
    finally:
        if validation_engine is not None:
            validation_engine.dispose()
        if created:
            with server_engine.begin() as connection:
                connection.exec_driver_sql(f"DROP DATABASE `{database_name}`")
        server_engine.dispose()


if __name__ == "__main__":
    main()
