from __future__ import annotations

from collections.abc import Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Any


@dataclass(frozen=True)
class Migration:
    version: str
    description: str
    apply: Callable[[Any], None]


class MigrationRunner:
    def __init__(self, database: Any, migrations: list[Migration]) -> None:
        self.database = database
        self.migrations = list(sorted(migrations, key=lambda item: item.version))
        self.collection = self.database["schema_migrations"]

    def ensure_tracking(self) -> None:
        self.collection.create_index([("version", 1)], unique=True, name="schema_migrations_version")

    def applied_versions(self) -> set[str]:
        return {
            str(item.get("version"))
            for item in self.collection.find({}, {"version": 1})
            if item.get("version")
        }

    def run(self) -> list[str]:
        self.ensure_tracking()
        applied = self.applied_versions()
        executed: list[str] = []
        for migration in self.migrations:
            if migration.version in applied:
                continue
            migration.apply(self.database)
            self.collection.insert_one(
                {
                    "version": migration.version,
                    "description": migration.description,
                    "applied_at": datetime.now(UTC),
                }
            )
            executed.append(migration.version)
        return executed
