Source code for sqlspec.migrations.context

"""Migration context for passing runtime information to migrations."""

import asyncio
import inspect
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any

from sqlglot.dialects.dialect import Dialect

from sqlspec.protocols import HasStatementConfigProtocol
from sqlspec.utils.logging import get_logger
from sqlspec.utils.type_guards import has_statement_config_factory

if TYPE_CHECKING:
    from sqlspec.driver import AsyncDriverAdapterBase, SyncDriverAdapterBase

__all__ = ("MigrationContext",)

logger = get_logger("sqlspec.migrations.context")


@dataclass(slots=True)
class MigrationContext:
    """Context object passed to migration functions.

    Provides runtime information about the database environment
    to migration functions, allowing them to generate dialect-specific SQL.
    """

    config: "Any | None" = None
    """Database configuration object."""
    dialect: "str | None" = None
    """Database dialect."""
    metadata: "dict[str, Any] | None" = None
    """Additional metadata for the migration."""
    extension_config: "dict[str, Any] | None" = None
    """Extension-specific configuration options."""

    driver: "SyncDriverAdapterBase | AsyncDriverAdapterBase | None" = None
    """Database driver instance (available during execution)."""

    _execution_metadata: "dict[str, Any]" = field(default_factory=dict)
    """Internal execution metadata for tracking async operations."""

[docs]
def __post_init__(self) -> None: """Initialize metadata and extension config if not provided.""" if not self.metadata: self.metadata = {} if not self.extension_config: self.extension_config = {}
@classmethod def from_config(cls, config: Any) -> "MigrationContext": """Create context from database configuration. Args: config: Database configuration object. Returns: Migration context with dialect information. """ dialect: Any | None = None try: if isinstance(config, HasStatementConfigProtocol) and config.statement_config: dialect = config.statement_config.dialect elif has_statement_config_factory(config): stmt_config = config._create_statement_config() # pyright: ignore[reportPrivateUsage] dialect = stmt_config.dialect except Exception: logger.debug("Unable to extract dialect from config") return cls(dialect=_normalize_dialect_name(dialect), config=config)