Source code for codegen_database.pg_extension

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

if TYPE_CHECKING:
    from alembic.autogenerate.api import AutogenContext
    from sqlalchemy import Connection


[docs] @dataclass(frozen=True, eq=True) class PGExtension: name: str schema: str | None = None cascade: bool = False def to_sql_create(self) -> str: parts = [f"CREATE EXTENSION IF NOT EXISTS {self.name}"] if self.schema is not None: parts.append(f"SCHEMA {self.schema}") if self.cascade: parts.append("CASCADE") return " ".join(parts)
[docs] @dataclass class PGExtensions: extensions: set[PGExtension] = field(default_factory=set) @classmethod def extract(cls, metadata: MetaData) -> PGExtensions: pg_extensions = metadata.info.get("pg_extensions") or PGExtensions() metadata.info["pg_extensions"] = pg_extensions return pg_extensions
def register_pg_extension( metadata: MetaData, *ext: PGExtension, ) -> None: holder = PGExtensions.extract(metadata) holder.extensions |= set(ext) def register_default_pg_extensions(metadata: MetaData) -> None: register_pg_extension(metadata, PGExtension("pg_trgm"))
[docs] @dataclass class CreateExtensionOp(MigrateOperation): extension: PGExtension def to_sql(self) -> list[str]: return [self.extension.to_sql_create()] def reverse(self) -> CreateExtensionOp: return self
def _fetch_current_extensions(connection: Connection) -> set[str]: rows = connection.execute( text("SELECT extname FROM pg_extension") ).fetchall() return {row.extname for row in rows} def fetch_extension_owned_relations( connection: Connection, ) -> set[tuple[str, str]]: rows = connection.execute( text( "SELECT n.nspname AS schema_name, c.relname AS rel_name " "FROM pg_depend d " "JOIN pg_class c " " ON c.oid = d.objid AND d.classid = 'pg_class'::regclass " "JOIN pg_namespace n ON n.oid = c.relnamespace " "WHERE d.deptype = 'e' " " AND d.refclassid = 'pg_extension'::regclass" ) ).fetchall() return {(row.schema_name, row.rel_name) for row in rows} def _compare_extensions( autogen_context: AutogenContext, upgrade_ops: UpgradeOps, _schemas: object, ) -> None: holder = PGExtensions.extract(autogen_context.metadata) # type: ignore[arg-type] current = _fetch_current_extensions(autogen_context.connection) # type: ignore[arg-type] upgrade_ops.ops.extend( CreateExtensionOp(extension=extension) for extension in holder.extensions if extension.name not in current ) @cache def register_pg_extension_alembic_events() -> None: def _render_extension_op( _autogen_context: AutogenContext, op: CreateExtensionOp, ) -> list[str]: return [f'op.execute("{cmd}")' for cmd in op.to_sql()] comparators.dispatch_for("schema")(_compare_extensions) renderers.dispatch_for(CreateExtensionOp)(_render_extension_op)