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)