Source code for codegen_database.pg_type
from __future__ import annotations
from dataclasses import dataclass, field
from functools import cache
from typing import TYPE_CHECKING
from alembic.autogenerate.compare import comparators
from alembic.autogenerate.render import renderers
from alembic.operations.ops import MigrateOperation, UpgradeOps
from sqlalchemy import MetaData, text
from codegen_database.alembic.renderer import _prettify, _render_execute
if TYPE_CHECKING:
from alembic.autogenerate.api import AutogenContext
from sqlalchemy import Connection
[docs]
@dataclass(frozen=True, eq=True)
class PGType:
name: str
create_sql: str
schema: str | None = None
drop_sql: str | None = None
companions: tuple[str, ...] = ()
def _qualified_name(self) -> str:
return f"{self.schema}.{self.name}" if self.schema else self.name
def to_sql_create(self) -> str:
return self.create_sql
def to_sql_drop(self) -> str:
return self.drop_sql or f"DROP TYPE IF EXISTS {self._qualified_name()}"
[docs]
@dataclass
class PGTypes:
types: set[PGType] = field(default_factory=set)
@classmethod
def extract(cls, metadata: MetaData) -> PGTypes:
pg_types = metadata.info.get("pg_types") or PGTypes()
metadata.info["pg_types"] = pg_types
return pg_types
def register_pg_type(
metadata: MetaData,
*types: PGType,
) -> None:
holder = PGTypes.extract(metadata)
holder.types |= set(types)
[docs]
@dataclass
class CreateTypeOp(MigrateOperation):
type_: PGType
def to_sql(self) -> list[str]:
return [self.type_.to_sql_create()]
def reverse(self) -> DropTypeOp:
return DropTypeOp(type_=self.type_)
[docs]
@dataclass
class DropTypeOp(MigrateOperation):
type_: PGType
def to_sql(self) -> list[str]:
return [self.type_.to_sql_drop()]
def reverse(self) -> CreateTypeOp:
return CreateTypeOp(type_=self.type_)
def _fetch_current_types(connection: Connection) -> set[tuple[str, str]]:
rows = connection.execute(
text(
"SELECT n.nspname AS schema_name, t.typname AS type_name "
"FROM pg_type t "
"JOIN pg_namespace n ON n.oid = t.typnamespace "
"LEFT JOIN pg_class c ON c.oid = t.typrelid "
"WHERE t.typisdefined "
" AND t.typtype IN ('r','c','e','d','m') "
" AND (t.typrelid = 0 OR c.relkind = 'c')"
)
).fetchall()
return {(row.schema_name, row.type_name) for row in rows}
def _is_installed(type_: PGType, current: set[tuple[str, str]]) -> bool:
for schema_name, type_name in current:
if type_name != type_.name:
continue
if type_.schema is None or schema_name == type_.schema.casefold():
return True
return False
def _compare_types(
autogen_context: AutogenContext,
upgrade_ops: UpgradeOps,
_schemas: object,
) -> None:
holder = PGTypes.extract(autogen_context.metadata) # type: ignore[arg-type]
current = _fetch_current_types(autogen_context.connection) # type: ignore[arg-type]
upgrade_ops.ops.extend(
CreateTypeOp(type_=type_)
for type_ in holder.types
if not _is_installed(type_, current)
)
[docs]
@cache
def register_pg_type_alembic_events() -> None:
"""Wire the autogenerate comparator and renderer for custom types."""
def _render_type_op(
_autogen_context: AutogenContext,
op: CreateTypeOp,
) -> list[str]:
return [_render_execute(_prettify(cmd)) for cmd in op.to_sql()]
def _render_drop_type_op(
_autogen_context: AutogenContext,
op: DropTypeOp,
) -> list[str]:
return [_render_execute(_prettify(cmd)) for cmd in op.to_sql()]
comparators.dispatch_for("schema")(_compare_types)
renderers.dispatch_for(CreateTypeOp)(_render_type_op)
renderers.dispatch_for(DropTypeOp)(_render_drop_type_op)