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)