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