Source code for codegen_database.alembic.dependency

from __future__ import annotations

import logging
import re
from dataclasses import dataclass, replace
from graphlib import TopologicalSorter
from itertools import pairwise
from typing import Any, Literal

import pglast
from alembic.operations import ops as alembic_ops
from pglast.error import Error as PglastError
from pglast.visitors import Visitor
from sqlalchemy import MetaData
from sqlalchemy import types as sa_types
from sqlalchemy_declarative_extensions.alembic.function import (
    CreateFunctionOp,
    DropFunctionOp,
    UpdateFunctionOp,
)
from sqlalchemy_declarative_extensions.alembic.procedure import (
    CreateProcedureOp,
    DropProcedureOp,
    UpdateProcedureOp,
)
from sqlalchemy_declarative_extensions.alembic.schema import (
    CreateSchemaOp,
    DropSchemaOp,
)
from sqlalchemy_declarative_extensions.alembic.trigger import (
    CreateTriggerOp,
    DropTriggerOp,
    UpdateTriggerOp,
)
from sqlalchemy_declarative_extensions.alembic.view import (
    CreateViewOp,
    DropViewOp,
    UpdateViewOp,
)
from sqlalchemy_declarative_extensions.dialects.postgresql.function import (
    type_map as _pg_type_map,
)
from sqlalchemy_declarative_extensions.dialects.postgresql.grant import (
    DefaultGrantStatement,
    GrantStatement,
)
from sqlalchemy_declarative_extensions.grant.compare import (
    GrantPrivilegesOp,
    RevokePrivilegesOp,
)
from sqlalchemy_declarative_extensions.op import MigrateOp
from sqlalchemy_declarative_extensions.role.compare import (
    CreateRoleOp,
    DropRoleOp,
)
from sqlalchemy_declarative_extensions.role.generic import Role

from codegen_database.ext.cron.compare import (
    ScheduleCronJobOp,
    UnscheduleCronJobOp,
)
from codegen_database.ext.rls.compare import (
    CreateRLSPolicyOp,
    DropRLSPolicyOp,
)
from codegen_database.pg_extension import CreateExtensionOp
from codegen_database.pg_type import CreateTypeOp, DropTypeOp, PGType

logger = logging.getLogger(__name__)

# Union of alembic's built-in ops and sqlalchemy-declarative-extensions ops,
# which don't share a common base class.
type AnyOp = alembic_ops.MigrateOperation | MigrateOp

Phase = Literal["drop", "create"]

_CREATE_OPS = (
    CreateRoleOp,
    CreateRLSPolicyOp,
    ScheduleCronJobOp,
    CreateSchemaOp,
    CreateTypeOp,
    alembic_ops.CreateTableOp,
    alembic_ops.CreateIndexOp,
    CreateViewOp,
    CreateFunctionOp,
    CreateProcedureOp,
    CreateTriggerOp,
    GrantPrivilegesOp,
)
_DROP_OPS = (
    DropRoleOp,
    DropRLSPolicyOp,
    UnscheduleCronJobOp,
    DropSchemaOp,
    alembic_ops.DropTableOp,
    alembic_ops.DropIndexOp,
    DropViewOp,
    DropFunctionOp,
    DropProcedureOp,
    DropTriggerOp,
    RevokePrivilegesOp,
    DropTypeOp,
)
_UPDATE_OPS = (
    UpdateViewOp,
    UpdateFunctionOp,
    UpdateProcedureOp,
    UpdateTriggerOp,
)


[docs] @dataclass(frozen=True) class EntityIdentifier: schema: str = "public" name: str | None = None phase: Phase | None = None kind: str | None = None signature: tuple[str, ...] | None = None target: tuple[str, str] | None = None
[docs] @dataclass(frozen=True) class OperationNode: index: int identity: EntityIdentifier
def _op_phase(op: AnyOp) -> Phase | None: if isinstance(op, _DROP_OPS): return "drop" if isinstance(op, _CREATE_OPS): return "create" return None def _entity_schema(op: AnyOp) -> str | None: if isinstance(op, (CreateSchemaOp, DropSchemaOp)): return op.schema.name if isinstance(op, _CREATE_OPS + _DROP_OPS): for attr in ("view", "function", "procedure"): entity = getattr(op, attr, None) if entity is not None: return entity.schema or "public" trigger = getattr(op, "trigger", None) if trigger is not None: # PostgreSQL triggers store the target as "schema.table" # in the `on` attribute; they have no `schema` field. on_schema, _ = _parse_qualified_name(trigger.on) return f"__triggers__{on_schema}" return None def _entity_name(op: AnyOp) -> str | None: for attr in ("view", "function", "procedure"): entity = getattr(op, attr, None) if entity is not None: return entity.name trigger = getattr(op, "trigger", None) if trigger is not None: return trigger.name return None def _entity_definition(op: AnyOp) -> str | None: view = getattr(op, "view", None) if view is not None and hasattr(view, "definition"): defn = view.definition if isinstance(defn, str): return defn return None def _plpgsql_queries(obj: object) -> list[str]: queries: list[str] = [] if isinstance(obj, dict): if "PLpgSQL_expr" in obj: query = obj["PLpgSQL_expr"].get("query", "") if query and query.upper() not in ("NEW", "OLD"): queries.append(query) else: for value in obj.values(): queries.extend(_plpgsql_queries(value)) elif isinstance(obj, list): for item in obj: queries.extend(_plpgsql_queries(item)) return queries def _fn_setof_ref(func: object) -> tuple[str, str] | None: returns = getattr(func, "returns", None) if not returns: return None # SDE normalises functions before creating ``CreateFunctionOp``, so # ``returns`` may be a ``FunctionReturn`` object rather than a plain string. if isinstance(returns, str): stripped = returns.strip() elif hasattr(returns, "value") and isinstance(returns.value, str): stripped = returns.value.strip() else: return None if not stripped.upper().startswith("SETOF "): return None ref = stripped[6:].strip() schema, sep, name = ref.partition(".") if sep: return (schema.lower(), name.lower()) return None def _function_return_refs(op: AnyOp) -> set[tuple[str, str]]: func = getattr(op, "function", None) if func is None: return set() ref = _fn_setof_ref(func) return {ref} if ref is not None else set() def _sql_function_body(op: AnyOp) -> str | None: func = getattr(op, "function", None) if func is None: return None defn = getattr(func, "definition", None) language = getattr(func, "language", "") if not defn or language.lower() != "sql": return None return defn def _plpgsql_function_queries(op: AnyOp) -> list[str]: func = getattr(op, "function", None) if func is None: return [] defn = getattr(func, "definition", None) language = getattr(func, "language", "") if not defn or language.lower() != "plpgsql": return [] # pglast.parse_plpgsql requires a full CREATE FUNCTION statement. schema_part = f"{func.schema}." if func.schema else "" wrapper = ( f"CREATE FUNCTION {schema_part}__cave_parse_helper()" f" RETURNS trigger LANGUAGE plpgsql AS $${defn}$$;" ) try: tree = pglast.parse_plpgsql(wrapper) except PglastError: logger.debug( "Could not parse PL/pgSQL body for %s", func.name, ) return [] return _plpgsql_queries(tree) def _sql_function_table_refs( op: AnyOp, ) -> set[tuple[str, str]]: defn = _sql_function_body(op) if defn is None: return set() return _view_table_refs(defn) def _plpgsql_table_refs( op: AnyOp, ) -> set[tuple[str, str]]: refs: set[tuple[str, str]] = set() for query in _plpgsql_function_queries(op): refs |= _view_table_refs(query) return refs def _role_name(member: Role | str) -> str: if isinstance(member, Role): return member.name return member def _callable_signature(entity: object) -> tuple[str, ...]: parameters: list[Any] = getattr(entity, "parameters", None) or [] return tuple( _canonical_type(getattr(parameter, "type", parameter)) for parameter in parameters ) def _id_for_declarative_op( op: AnyOp, phase: Phase | None, ) -> EntityIdentifier | None: trigger = getattr(op, "trigger", None) if trigger is not None: target = _parse_qualified_name(trigger.on) return EntityIdentifier( schema=f"__triggers__{target[0]}", name=trigger.name.lower(), phase=phase, kind="trigger", target=target, ) for attr, kind in ( ("view", "view"), ("function", "function"), ("procedure", "procedure"), ): entity = getattr(op, attr, None) if entity is not None: signature = ( _callable_signature(entity) if kind in {"function", "procedure"} else None ) return EntityIdentifier( schema=(entity.schema or "public").lower(), name=entity.name.lower(), phase=phase, kind=kind, signature=signature, ) return None def _entity_identifier( # noqa: C901, PLR0911 op: AnyOp, ) -> EntityIdentifier | None: phase = _op_phase(op) if isinstance(op, (CreateSchemaOp, DropSchemaOp)): return EntityIdentifier( schema=op.schema.name.lower(), phase=phase, ) if isinstance(op, (alembic_ops.CreateTableOp, alembic_ops.DropTableOp)): return EntityIdentifier( schema=(op.schema or "public").lower(), name=op.table_name.lower(), phase=phase, kind="relation", ) if isinstance(op, (CreateTypeOp, DropTypeOp)): return EntityIdentifier( schema=(op.type_.schema or "public").lower(), name=op.type_.name.lower(), phase=phase, kind="type", ) if isinstance(op, alembic_ops.ModifyTableOps): return EntityIdentifier( schema=(op.schema or "public").lower(), name=op.table_name.lower(), ) if isinstance(op, (alembic_ops.CreateIndexOp, alembic_ops.DropIndexOp)): return EntityIdentifier( schema=(op.schema or "public").lower(), name=f"__index__{(op.index_name or '').lower()}", phase=phase, ) if isinstance(op, (CreateRoleOp, DropRoleOp)): return EntityIdentifier( schema="__roles__", name=op.role.name.lower(), phase=phase, kind="role", ) if isinstance(op, (CreateRLSPolicyOp, DropRLSPolicyOp)): policy = op.policy return EntityIdentifier( schema=policy.schema.lower(), name=policy.name.lower(), phase=phase, kind="rls_policy", target=(policy.schema.lower(), policy.table.lower()), ) if isinstance(op, (ScheduleCronJobOp, UnscheduleCronJobOp)): name = op.job.name if isinstance(op, ScheduleCronJobOp) else op.name return EntityIdentifier( schema="cron", name=name.lower(), phase=phase, kind="cron_job", ) if isinstance(op, (GrantPrivilegesOp, RevokePrivilegesOp)): return EntityIdentifier( schema="__grants__", name=str(op.to_sql()).lower(), phase=phase, ) result = _id_for_declarative_op(op, phase) if result is not None: return result logger.warning( "Unhandled op type %s; ordering is unconstrained", type(op).__name__, ) return None def _refs_for_role( op: CreateRoleOp | DropRoleOp, phase: Phase | None, ) -> set[EntityIdentifier]: if not op.role.in_roles: return set() return { EntityIdentifier( schema="__roles__", name=_role_name(member).lower(), phase=phase, ) for member in op.role.in_roles } def _refs_for_grant( op: GrantPrivilegesOp | RevokePrivilegesOp, phase: Phase | None, ) -> set[EntityIdentifier]: grant_obj = op.grant refs: set[EntityIdentifier] = set() refs.add( EntityIdentifier( schema="__roles__", name=grant_obj.grant.target_role.lower(), phase=phase, ) ) if isinstance(grant_obj, GrantStatement): for target_name in grant_obj.targets: schema_part, sep, obj_name = target_name.partition(".") if sep: refs.add( EntityIdentifier( schema=schema_part.lower(), phase=phase, ) ) refs.add( EntityIdentifier( schema=schema_part.lower(), name=obj_name.lower(), phase=phase, ) ) else: refs.add( EntityIdentifier( schema=target_name.lower(), phase=phase, ) ) elif isinstance(grant_obj, DefaultGrantStatement): for schema_name in grant_obj.default_grant.in_schemas: refs.add( EntityIdentifier( schema=schema_name.lower(), phase=phase, ) ) return refs def _parse_qualified_name( qualified: str, ) -> tuple[str, str]: parts = qualified.split(".", 1) if len(parts) > 1: return parts[0].lower(), parts[1].lower() return "public", parts[0].lower() def _refs_for_trigger( op: AnyOp, phase: Phase | None, ) -> set[EntityIdentifier]: trigger = getattr(op, "trigger", None) if trigger is None: return set() refs: set[EntityIdentifier] = set() on_schema, on_name = _parse_qualified_name(trigger.on) refs.add( EntityIdentifier( schema=on_schema, name=on_name, phase=phase, ) ) refs.add(EntityIdentifier(schema=on_schema, phase=phase)) fn_schema, fn_name = _parse_qualified_name(trigger.execute) refs.add( EntityIdentifier( schema=fn_schema, name=fn_name, phase=phase, ) ) return refs def _add_named_refs( refs: set[EntityIdentifier], pairs: set[tuple[str, str]], phase: Phase | None, kind: str, ) -> None: for ref_schema, ref_name in pairs: refs.add(EntityIdentifier(ref_schema, ref_name, phase, kind=kind)) if phase is not None: refs.add(EntityIdentifier(ref_schema, ref_name, kind=kind)) def _view_table_refs(definition: str) -> set[tuple[str, str]]: refs: set[tuple[str, str]] = set() try: parsed = pglast.parse_sql(definition) class _TableFinder(Visitor): def visit_RangeVar( # noqa: N802 self, _ancestors: object, node: object, ) -> None: if name := getattr(node, "relname", None): schema = getattr(node, "schemaname", "public") or "public" refs.add((schema.lower(), name.lower())) _TableFinder()(parsed) except PglastError: logger.debug( "Could not parse view definition: %s", definition[:80], ) return refs def _sql_function_function_refs(op: AnyOp) -> set[tuple[str, str]]: """Extract ``(schema, name)`` function calls from a LANGUAGE sql function. Returns an empty set for non-sql-language ops or if parsing fails. """ defn = _sql_function_body(op) if defn is None: return set() return _view_function_refs(defn) def _plpgsql_function_refs(op: AnyOp) -> set[tuple[str, str]]: """Extract ``(schema, name)`` function calls from a PL/pgSQL function body. Returns an empty set for non-function ops or if parsing fails. """ queries = _plpgsql_function_queries(op) if not queries: return set() refs: set[tuple[str, str]] = set() for query in queries: refs |= _view_function_refs(query) return refs def _refs_from_definitions( op: AnyOp, phase: Phase | None, ) -> set[EntityIdentifier]: refs: set[EntityIdentifier] = set() definition = _entity_definition(op) if definition is not None: _add_named_refs(refs, _view_table_refs(definition), phase, "relation") _add_named_refs( refs, _view_function_refs(definition), phase, "function" ) _add_named_refs(refs, _plpgsql_table_refs(op), phase, "relation") _add_named_refs(refs, _plpgsql_function_refs(op), phase, "function") _add_named_refs(refs, _sql_function_table_refs(op), phase, "relation") _add_named_refs(refs, _sql_function_function_refs(op), phase, "function") _add_named_refs(refs, _function_return_refs(op), phase, "relation") return refs def _refs_for_declarative_op( op: AnyOp, phase: Phase | None, ) -> set[EntityIdentifier]: trigger_refs = _refs_for_trigger(op, phase) if trigger_refs: return trigger_refs name = _entity_name(op) if name is None: return set() schema = (_entity_schema(op) or "public").lower() self_id = EntityIdentifier( schema=schema, name=name.lower(), phase=phase, ) refs: set[EntityIdentifier] = set() # For split update ops: create must wait for drop. if phase == "create": refs.add( EntityIdentifier( schema=schema, name=name.lower(), phase="drop", ) ) # Depend on the containing schema: created after its schema on # create, and dropped before its schema on drop (the drop-phase # edge direction reverses this so DROP SCHEMA waits for the # function -- e.g. the ``codegen_database`` utility schema can't be # dropped while its ``codegen_database_date_bin`` chart polyfill # still lives in it). refs.add(EntityIdentifier(schema=schema, phase=phase)) refs |= _refs_from_definitions(op, phase) refs.discard(self_id) return refs def build_fk_graph_from_metadata( metadata: MetaData, ) -> dict[tuple[str, str], set[tuple[str, str]]]: graph: dict[tuple[str, str], set[tuple[str, str]]] = {} for table in metadata.tables.values(): key = ( (table.schema or "public").lower(), table.name.lower(), ) targets: set[tuple[str, str]] = set() for fk in table.foreign_keys: ref = fk.column.table targets.add( ( (ref.schema or "public").lower(), ref.name.lower(), ) ) targets.discard(key) if targets: graph[key] = targets return graph def _declared_type_specs( declared_types: tuple[PGType, ...], ) -> dict[str, tuple[str, str]]: specs: dict[str, tuple[str, str]] = {} for pg_type in declared_types: prefix = f"{pg_type.schema}." if pg_type.schema else "" key = ( (pg_type.schema or "public").lower(), pg_type.name.lower(), ) specs[f"{prefix}{pg_type.name.lower()}"] = key for companion in pg_type.companions: specs[f"{prefix}{companion.lower()}"] = key return specs def _column_type_specs(type_: sa_types.TypeEngine) -> set[str]: if isinstance(type_, sa_types.TypeDecorator): type_ = type_.impl_instance if isinstance(type_, sa_types.UserDefinedType): try: return {str(type_.get_col_spec()).lower()} except Exception: # noqa: BLE001 -- get_col_spec may be NotImplemented return set() return set() def build_type_graph_from_metadata( metadata: MetaData, declared_types: tuple[PGType, ...], ) -> dict[tuple[str, str], set[tuple[str, str]]]: specs = _declared_type_specs(declared_types) graph: dict[tuple[str, str], set[tuple[str, str]]] = {} for table in metadata.tables.values(): key = ( (table.schema or "public").lower(), table.name.lower(), ) targets: set[tuple[str, str]] = set() for column in table.columns: for spec in _column_type_specs(column.type): type_key = specs.get(spec) if type_key is not None: targets.add(type_key) if targets: graph[key] = targets return graph def _policy_references( op: CreateRLSPolicyOp | DropRLSPolicyOp, phase: Phase | None, ) -> set[EntityIdentifier]: policy = op.policy refs = { EntityIdentifier( schema=policy.schema.lower(), name=policy.table.lower(), phase=phase, kind="relation", ) } refs.update( EntityIdentifier( schema="__roles__", name=role.lower(), phase=phase, kind="role", ) for role in policy.roles if role.lower() != "public" ) return refs def _cron_references( op: ScheduleCronJobOp | UnscheduleCronJobOp, phase: Phase | None, ) -> set[EntityIdentifier]: refs = {EntityIdentifier(schema="cron", phase=phase)} if isinstance(op, ScheduleCronJobOp): for schema, name in _view_table_refs(op.job.command): refs.add(EntityIdentifier(schema, name, phase, kind="relation")) for schema, name in _view_function_refs(op.job.command): refs.add(EntityIdentifier(schema, name, phase, kind="function")) return refs def _entity_references( # noqa: PLR0911 op: AnyOp, ) -> set[EntityIdentifier]: phase = _op_phase(op) if isinstance(op, (CreateRLSPolicyOp, DropRLSPolicyOp)): return _policy_references(op, phase) if isinstance(op, (ScheduleCronJobOp, UnscheduleCronJobOp)): previous = getattr(op, "previous", None) if previous is not None and isinstance(op, UnscheduleCronJobOp): synthetic = ScheduleCronJobOp(previous) return _cron_references(synthetic, phase) return _cron_references(op, phase) # Tables depend on their containing schema. FK deps are # handled separately via fk_graph in sort_migration_ops. if isinstance(op, (alembic_ops.CreateTableOp, alembic_ops.DropTableOp)): return { EntityIdentifier( schema=(op.schema or "public").lower(), phase=phase, ) } if isinstance(op, (CreateTypeOp, DropTypeOp)): return { EntityIdentifier( schema=(op.type_.schema or "public").lower(), phase=phase, ) } if isinstance(op, (CreateRoleOp, DropRoleOp)): return _refs_for_role(op, phase) if isinstance(op, (GrantPrivilegesOp, RevokePrivilegesOp)): return _refs_for_grant(op, phase) return _refs_for_declarative_op(op, phase) def _op_label(op: AnyOp) -> str: identifier = _entity_identifier(op) if identifier is None: entity = "?" elif identifier.name is None: entity = identifier.schema else: entity = f"{identifier.schema}.{identifier.name}" return f"{type(op).__name__}({entity})" _QUALIFIED_FUNCNAME_MIN_PARTS = 2 _KIND_FUNCTION = "function" _KIND_VIEW = "view" def _view_function_refs(definition: str) -> set[tuple[str, str]]: refs: set[tuple[str, str]] = set() try: parsed = pglast.parse_sql(definition) class _FuncFinder(Visitor): def visit_FuncCall( # noqa: N802 self, _ancestors: object, node: object, ) -> None: funcname = getattr(node, "funcname", None) if funcname is None: return if len(funcname) < _QUALIFIED_FUNCNAME_MIN_PARTS: return schema = getattr(funcname[-2], "sval", None) name = getattr(funcname[-1], "sval", None) if schema and name: refs.add((schema.lower(), name.lower())) _FuncFinder()(parsed) except PglastError: logger.debug( "Could not parse view definition for function refs: %s", definition[:80], ) return refs type FunctionKey = tuple[str, str, tuple[str, ...]] def _function_key(function: Any) -> FunctionKey: # noqa: ANN401 return ( (function.schema or "public").lower(), function.name.lower(), _callable_signature(function), ) def _updated_function_keys( migration_ops: list[AnyOp], ) -> set[FunctionKey]: updated: set[FunctionKey] = set() for op in migration_ops: if isinstance(op, (UpdateFunctionOp, UpdateProcedureOp)): function = getattr(op, "function", None) or getattr( op, "procedure", None ) if function is not None: updated.add(_function_key(function)) return updated def _trigger_key(trigger: Any) -> tuple[str, str, str]: # noqa: ANN401 schema, table = _parse_qualified_name(trigger.on) return schema, table, trigger.name.lower() def _existing_trigger_keys( migration_ops: list[AnyOp], ) -> set[tuple[str, str, str]]: return { _trigger_key(trigger) for op in migration_ops if (trigger := getattr(op, "trigger", None)) is not None } def _existing_view_keys(migration_ops: list[AnyOp]) -> set[tuple[str, str]]: keys: set[tuple[str, str]] = set() for op in migration_ops: view = getattr(op, "view", None) if view is not None: keys.add(((view.schema or "public").lower(), view.name.lower())) return keys def _updated_view_names( migration_ops: list[AnyOp], ) -> set[tuple[str, str]]: updated: set[tuple[str, str]] = set() for op in migration_ops: if isinstance(op, UpdateViewOp): view = op.view updated.add(((view.schema or "public").lower(), view.name.lower())) return updated def _column_altered_table_keys( migration_ops: list[AnyOp], ) -> set[tuple[str, str]]: def has_type_change(op: alembic_ops.ModifyTableOps) -> bool: return any( isinstance(sub, alembic_ops.AlterColumnOp) and sub.modify_type is not None for sub in op.ops ) return { ((op.schema or "public").lower(), op.table_name.lower()) for op in migration_ops if isinstance(op, alembic_ops.ModifyTableOps) and has_type_change(op) } def _canonical_function_body(definition: str) -> str: return "\n".join(line.strip() for line in definition.splitlines()).strip() def _canonical_type(type_: Any) -> str: # noqa: ANN401 canon = re.sub(r"\s*\([^)]*\)", "", str(type_)).strip().lower() canon = re.sub(r"\bpublic\.", "", canon) return _pg_type_map.get(canon, canon) def _param_identities(parameters: Any) -> Any: # noqa: ANN401 if not parameters: return parameters return [ p if isinstance(p, str) else ( p.name.lower() if p.name is not None else None, _canonical_type(p.type), p.mode, ) for p in parameters ] _TABLE_RETURN_RE = re.compile(r"(?is)^\s*table\s*\((.*)\)\s*$") def _split_top_level(cols: str) -> list[str]: """Split a comma list, ignoring commas inside parens (``numeric(10,2)``).""" parts: list[str] = [] depth = 0 start = 0 for i, ch in enumerate(cols): if ch == "(": depth += 1 elif ch == ")": depth -= 1 elif ch == "," and depth == 0: parts.append(cols[start:i]) start = i + 1 parts.append(cols[start:]) return [p.strip() for p in parts if p.strip()] def _return_identity(returns: Any) -> Any: # noqa: ANN401 if returns is None: return returns table = getattr(returns, "table", None) value = getattr(returns, "value", None) if not table and isinstance(value, str): match = _TABLE_RETURN_RE.match(value) if match: table = [ col.split(maxsplit=1) for col in _split_top_level(match.group(1)) ] if table: return ("table", tuple(_col_identity(col) for col in table)) if isinstance(value, str): return ("scalar", _canonical_type(value)) return returns def _col_identity(col: Any) -> tuple[str, str]: # noqa: ANN401 if isinstance(col, str): name, _, type_ = col.strip().partition(" ") else: name, type_ = col[0], col[1] return str(name).lower(), _canonical_type(type_) def _qualify_default_schema(name: str | None) -> str | None: if name is not None and "." not in name: return f"public.{name}" return name def _function_update_is_spurious(op: UpdateFunctionOp) -> bool: def canon(fn: Any) -> Any: # noqa: ANN401 return replace( fn, schema=fn.schema or "public", definition=_canonical_function_body(fn.definition), parameters=_param_identities(fn.parameters), returns=_return_identity(fn.returns), ) return canon(op.from_function) == canon(op.function) def _view_update_is_spurious(op: UpdateViewOp) -> bool: def canon(view: Any) -> Any: # noqa: ANN401 return replace(view, schema=view.schema or "public") return canon(op.from_view) == canon(op.view) def _trigger_update_is_spurious(op: UpdateTriggerOp) -> bool: def canon(trigger: Any) -> Any: # noqa: ANN401 return replace( trigger, on=_qualify_default_schema(trigger.on), execute=_qualify_default_schema(trigger.execute), ) return canon(op.from_trigger) == canon(op.trigger) def drop_spurious_declarative_updates(ops: list[AnyOp]) -> list[AnyOp]: return [ op for op in ops if not ( ( isinstance(op, UpdateFunctionOp) and _function_update_is_spurious(op) ) or (isinstance(op, UpdateViewOp) and _view_update_is_spurious(op)) or ( isinstance(op, UpdateTriggerOp) and _trigger_update_is_spurious(op) ) ) ] def _inject_dependent_update_ops( # noqa: C901, PLR0915 migration_ops: list[AnyOp], metadata: MetaData | None, ) -> list[AnyOp]: if metadata is None: return list(migration_ops) seen_fns = _updated_function_keys(migration_ops) updated_views = _updated_view_names(migration_ops) altered_tables = _column_altered_table_keys(migration_ops) if not (seen_fns or updated_views or altered_tables): return list(migration_ops) all_views = list(metadata.info.get("views") or []) all_fns = list(metadata.info.get("functions") or []) views_by_key: dict[tuple[str, str], Any] = { ((v.schema or "public").lower(), v.name.lower()): v for v in all_views } fns_by_key: dict[FunctionKey, Any] = { _function_key(function): function for function in all_fns } all_triggers = list(metadata.info.get("triggers") or []) view_fn_refs = { k: _view_function_refs(getattr(v, "definition", "") or "") for k, v in views_by_key.items() } view_tbl_refs = { k: _view_table_refs(getattr(v, "definition", "") or "") for k, v in views_by_key.items() } fn_setof_refs = { key: ref for key, function in fns_by_key.items() if (ref := _fn_setof_ref(function)) is not None } seen_triggers = _existing_trigger_keys(migration_ops) seen_views = _existing_view_keys(migration_ops) extra: list[AnyOp] = [] queue: list[tuple[str, str, str]] = [ (schema, name, _KIND_FUNCTION) for schema, name, _signature in seen_fns ] # Recreated views drop their INSTEAD OF triggers; revisit to re-add them. queue += [(s, n, _KIND_VIEW) for s, n in updated_views] visited: set[tuple[str, str, str]] = set(queue) def _enqueue_view(vkey: tuple[str, str]) -> None: seen_views.add(vkey) view = views_by_key[vkey] extra.append(UpdateViewOp(view, view)) entry = (vkey[0], vkey[1], _KIND_VIEW) if entry not in visited: visited.add(entry) queue.append(entry) def _enqueue_trigger(trigger: Any) -> None: # noqa: ANN401 trigger_key = _trigger_key(trigger) if trigger_key not in seen_triggers: seen_triggers.add(trigger_key) extra.append(UpdateTriggerOp(trigger, trigger)) def _enqueue_fn(fn_key: FunctionKey) -> None: seen_fns.add(fn_key) function = fns_by_key[fn_key].normalize() extra.append(UpdateFunctionOp(function, function)) entry = (fn_key[0], fn_key[1], _KIND_FUNCTION) if entry not in visited: visited.add(entry) queue.append(entry) def _visit_fn(fn_key: tuple[str, str]) -> None: for trigger in all_triggers: if _parse_qualified_name(trigger.execute) == fn_key: _enqueue_trigger(trigger) for vkey, fn_refs in view_fn_refs.items(): if fn_key in fn_refs and vkey not in seen_views: _enqueue_view(vkey) def _visit_view(view_key: tuple[str, str]) -> None: # INSTEAD OF triggers drop with the view; recreate them. for trigger in all_triggers: if _parse_qualified_name(trigger.on) == view_key: _enqueue_trigger(trigger) for vkey, tbl_refs in view_tbl_refs.items(): if view_key in tbl_refs and vkey not in seen_views: _enqueue_view(vkey) for fn_key, return_ref in fn_setof_refs.items(): if return_ref == view_key and fn_key not in seen_fns: _enqueue_fn(fn_key) # A column TYPE change forces dependent views to drop/recreate around it. for tkey in altered_tables: for vkey, tbl_refs in view_tbl_refs.items(): if tkey in tbl_refs and vkey not in seen_views: _enqueue_view(vkey) while queue: schema, name, kind = queue.pop(0) if kind == _KIND_FUNCTION: _visit_fn((schema, name)) else: _visit_view((schema, name)) return list(migration_ops) + extra def expand_update_ops( migration_ops: list[AnyOp], ) -> list[AnyOp]: result: list[AnyOp] = [] for op in migration_ops: if isinstance(op, UpdateViewOp): result.append(DropViewOp(op.from_view)) result.append(CreateViewOp(op.view)) elif isinstance(op, UpdateFunctionOp): result.append(DropFunctionOp(op.from_function)) result.append(CreateFunctionOp(op.function)) elif isinstance(op, UpdateProcedureOp): result.append(DropProcedureOp(op.from_procedure)) result.append(CreateProcedureOp(op.procedure)) elif isinstance(op, UpdateTriggerOp): result.append(DropTriggerOp(op.from_trigger)) result.append(CreateTriggerOp(op.trigger)) else: result.append(op) return result def _partition_extension_ops( ops: list[AnyOp], ) -> tuple[list[AnyOp], list[AnyOp]]: extension_ops: list[AnyOp] = [] rest: list[AnyOp] = [] for op in ops: if isinstance(op, (CreateExtensionOp,)): extension_ops.append(op) else: rest.append(op) return extension_ops, rest def _index_and_modify_refs( op: AnyOp, phase: Phase | None, ) -> set[EntityIdentifier]: if isinstance( op, (alembic_ops.CreateIndexOp, alembic_ops.DropIndexOp), ): return { EntityIdentifier( schema=(op.schema or "public").lower(), name=(op.table_name or "").lower(), phase=phase, ) } if isinstance(op, alembic_ops.ModifyTableOps): return { EntityIdentifier( schema=(op.schema or "public").lower(), name=op.table_name.lower(), phase="create", ) } return set() def prune_redundant_index_drops[T: AnyOp](ops: list[T]) -> list[T]: dropped_tables = { ((op.schema or "public").lower(), op.table_name.lower()) for op in ops if isinstance(op, alembic_ops.DropTableOp) } if not dropped_tables: return ops def _index_table_dropped(op: alembic_ops.DropIndexOp) -> bool: return ( (op.schema or "public").lower(), (op.table_name or "").lower(), ) in dropped_tables result: list[T] = [] for op in ops: if isinstance(op, alembic_ops.DropIndexOp) and _index_table_dropped(op): continue if isinstance(op, alembic_ops.ModifyTableOps): op.ops[:] = [ sub for sub in op.ops if not ( isinstance(sub, alembic_ops.DropIndexOp) and _index_table_dropped(sub) ) ] if not op.ops: continue result.append(op) return result def sort_migration_ops( # noqa: C901, PLR0912 migration_ops: list[AnyOp], *, fk_graph: dict[tuple[str, str], set[tuple[str, str]]], type_graph: dict[tuple[str, str], set[tuple[str, str]]] | None = None, ) -> list[AnyOp]: logger.debug( "Sorting %d ops: %s", len(migration_ops), [_op_label(op) for op in migration_ops], ) type_graph = type_graph or {} extension_ops, migration_ops = _partition_extension_ops(migration_ops) op_by_node: dict[OperationNode, AnyOp] = {} nodes_by_identity: dict[EntityIdentifier, list[OperationNode]] = {} unkeyed_ops: list[AnyOp] = [] for index, op in enumerate(migration_ops): identity = _entity_identifier(op) if identity is None: unkeyed_ops.append(op) continue node = OperationNode(index=index, identity=identity) op_by_node[node] = op nodes_by_identity.setdefault(identity, []).append(node) sorter: TopologicalSorter[OperationNode] = TopologicalSorter() for current_node, current_op in op_by_node.items(): current_id = current_node.identity sorter.add(current_node) phase = _op_phase(current_op) refs = _entity_references(current_op) # Add FK-based refs for both create and drop table ops. # The fk_graph (built from metadata) is used because # neither op type carries FK info at this point. if isinstance( current_op, ( alembic_ops.CreateTableOp, alembic_ops.DropTableOp, ), ): table_key = ( (current_op.schema or "public").lower(), current_op.table_name.lower(), ) for ref_schema, ref_table in fk_graph.get(table_key, set()): refs.add( EntityIdentifier( schema=ref_schema, name=ref_table, phase=phase, ) ) if isinstance( current_op, ( alembic_ops.CreateTableOp, alembic_ops.DropTableOp, alembic_ops.ModifyTableOps, ), ): table_key = ( (current_op.schema or "public").lower(), current_op.table_name.lower(), ) type_phase = phase or "create" for type_schema, type_name in type_graph.get(table_key, set()): refs.add( EntityIdentifier( schema=type_schema, name=type_name, phase=type_phase, kind="type", ) ) refs |= _index_and_modify_refs(current_op, phase) # Replace-in-place across object kinds: when this migration both # drops and creates a relation under the same (schema, name) -- # e.g. a dimension flipping from a view-backed model to a # passthrough table -- the create must wait for the drop, or # Postgres rejects CREATE with "relation already exists". # ``_refs_for_declarative_op`` adds this drop->create edge only # for replaceable (view/function) create ops; generalise it to # any create op (notably ``CreateTableOp``) whose name a drop op # in this same migration also targets. if phase == "create" and current_id.name is not None: replacement_kind = ( "relation" if current_id.kind in {"relation", "view"} else current_id.kind ) refs.add( EntityIdentifier( schema=current_id.schema, name=current_id.name, phase="drop", kind=replacement_kind, signature=current_id.signature, target=current_id.target, ) ) for ref_id in refs: matching_nodes = [ candidate for candidate in op_by_node if candidate.identity.schema == ref_id.schema and candidate.identity.name == ref_id.name and candidate.identity.phase == ref_id.phase and ( ref_id.kind is None or candidate.identity.kind == ref_id.kind or ( ref_id.kind == "relation" and candidate.identity.kind in {"relation", "view"} ) ) and ( ref_id.signature is None or candidate.identity.signature == ref_id.signature ) and ( ref_id.target is None or candidate.identity.target == ref_id.target ) ] for matching_node in matching_nodes: if matching_node == current_node: continue if phase == "drop": node, prerequisite = matching_node, current_node else: node, prerequisite = current_node, matching_node logger.debug( "Edge: %s before %s", _op_label(op_by_node[prerequisite]), _op_label(op_by_node[node]), ) sorter.add(node, prerequisite) for matching_nodes in nodes_by_identity.values(): for before, after in pairwise(matching_nodes): sorter.add(after, before) sorted_ops = ( extension_ops + [op_by_node[node] for node in sorter.static_order()] + unkeyed_ops ) logger.debug( "Sorted order: %s", [_op_label(op) for op in sorted_ops], ) return sorted_ops