Source code for codegen_database.pg_extension

"""Declarative PostgreSQL extension management for codegen_database.

:class:`PGExtension` describes a PostgreSQL extension that the
application requires.  :func:`register_pg_extension` stores it in
``metadata.info`` so the Alembic comparator can emit
``CREATE EXTENSION IF NOT EXISTS`` when the extension is missing from
the database.

Extensions are never dropped automatically.  Removing an extension
from metadata simply stops codegen_database from ensuring it is present; the
extension itself remains installed until a DBA drops it manually.

:data:`DEFAULT_PG_EXTENSIONS` (currently ``pg_trgm``) are registered
on every metadata by
:func:`codegen_database.alembic.register.configure_metadata` -- no
opt-in needed; a manual registration of the same name wins.

Usage::

    from codegen_database.pg_extension import PGExtension, register_pg_extension

    # Standalone registration
    register_pg_extension(metadata, PGExtension("btree_gist"))
    register_pg_extension(metadata, PGExtension("pg_cron", schema="cron"))

    # Alembic wiring (call once in env.py)
    from codegen_database.pg_extension import (
        register_pg_extension_alembic_events,
    )
    register_pg_extension_alembic_events()

Alembic autogenerate then emits::

    op.execute("CREATE EXTENSION IF NOT EXISTS btree_gist")
"""

from dataclasses import dataclass, field
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.errors import CodegenDatabaseValidationError

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

_alembic_registered = False


[docs] @dataclass(frozen=True) class PGExtension: """Describe a PostgreSQL extension to install. Args: name: Extension name, e.g. ``"btree_gist"``. schema: Schema in which to install the extension. When ``None`` the database default (usually ``public``) is used. cascade: When ``True``, adds ``CASCADE`` so dependent extensions are installed automatically. """ name: str schema: str | None = None cascade: bool = False
[docs] def to_sql_create(self) -> str: """Render ``CREATE EXTENSION IF NOT EXISTS`` DDL. Returns: A complete ``CREATE EXTENSION`` SQL string. """ 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: """Container for all extensions registered on a MetaData. Stored under ``metadata.info["pg_extensions"]`` by :func:`register_pg_extension`. """ extensions: list[PGExtension] = field(default_factory=list)
[docs] @classmethod def extract(cls, metadata: object) -> PGExtensions | None: """Return registered extensions or ``None`` if none exist. Args: metadata: SQLAlchemy :class:`~sqlalchemy.MetaData`, or ``None``. Returns: The :class:`PGExtensions` holder or ``None``. """ if metadata is None or not isinstance(metadata, MetaData): return None return metadata.info.get("pg_extensions")
[docs] def assert_pg_extension_declared( metadata: MetaData, name: str, ) -> None: """Raise if *name* has not been declared as a required extension. Call this inside a plugin's ``run()`` method to give a clear error when the user forgot to register the extension in metadata (e.g. via a :class:`~codegen_database.ext.pg_cron.PGCronExtension`). In the declarative model flow, extensions are configured via ``configure_metadata`` which runs after all models are imported. If a plugin calls this check at class-definition time and the extension is not yet in metadata but a :class:`~codegen_database.config.CodegenDatabaseConfig` is present on ``metadata.info``, the config's extension hooks are run eagerly so that the check can succeed. The full ``configure_metadata`` call in ``env.py`` will overwrite any interim state with the final correct values. Args: metadata: The :class:`~sqlalchemy.MetaData` to check. name: The PostgreSQL extension name to require. Raises: CodegenDatabaseValidationError: If *name* is not declared and cannot be resolved from the registered config. """ holder = PGExtensions.extract(metadata) if holder is not None and any(e.name == name for e in holder.extensions): return # Declarative flow: the config is present but configure_metadata hasn't # been called yet. Run extension hooks eagerly so the check can proceed. cfg = metadata.info.get("codegen_database_config") if cfg is not None: for ext in cfg._resolved_extensions(): ext.configure_metadata(metadata) holder = PGExtensions.extract(metadata) if holder is not None and any( e.name == name for e in holder.extensions ): return msg = ( f"Plugin requires PostgreSQL extension {name!r}, but it is not " f"declared in metadata. Register it via " f"register_pg_extension(metadata, PGExtension({name!r})) or " f"config.use(PGCronExtension()) for pg_cron." ) raise CodegenDatabaseValidationError(msg)
[docs] def register_pg_extension( metadata: MetaData, ext: PGExtension, ) -> None: """Store *ext* in *metadata* for Alembic autogenerate. Args: metadata: SQLAlchemy :class:`~sqlalchemy.MetaData` to register on. ext: The extension to register. """ holder: PGExtensions | None = metadata.info.get("pg_extensions") if holder is None: holder = PGExtensions() metadata.info["pg_extensions"] = holder holder.extensions.append(ext)
#: Extensions every codegen_database project gets without opting in. #: ``pg_trgm`` backs the ``%`` similarity operator that trigram text #: search (the default list-endpoint search strategy upstack) and #: ``trigram_indexes`` rely on; it ships with contrib and is harmless #: when unused, so registering it by default removes a foot-gun -- #: a project that declares searchable text columns no longer 500s at #: first search because nobody remembered ``CREATE EXTENSION``. DEFAULT_PG_EXTENSIONS: tuple[PGExtension, ...] = (PGExtension("pg_trgm"),)
[docs] def register_default_pg_extensions(metadata: MetaData) -> None: """Register :data:`DEFAULT_PG_EXTENSIONS` on *metadata*. Called by :func:`codegen_database.alembic.register.configure_metadata` so every project wired through the standard ``env.py`` hooks gets the defaults; safe to call repeatedly -- an extension already declared (by any path) is not added twice. Args: metadata: SQLAlchemy :class:`~sqlalchemy.MetaData` to register on. """ holder = PGExtensions.extract(metadata) declared = ( {e.name for e in holder.extensions} if holder is not None else set() ) for ext in DEFAULT_PG_EXTENSIONS: if ext.name not in declared: register_pg_extension(metadata, ext)
# --------------------------------------------------------------------------- # Alembic integration # ---------------------------------------------------------------------------
[docs] @dataclass class CreateExtensionOp(MigrateOperation): """Alembic operation: create a PostgreSQL extension. Inherits from :class:`alembic.operations.ops.MigrateOperation` so that codegen_database's Alembic rewriter passes it through during ``process_revision_directives`` traversal without error. """ extension: PGExtension
[docs] def to_sql(self) -> list[str]: """Return the DDL statement for this operation.""" return [self.extension.to_sql_create()]
[docs] def reverse(self) -> CreateExtensionOp: """Return a no-op downgrade: extensions are never dropped automatically. The downgrade renderer emits a SQL comment instead of a DROP statement so that autogenerate does not accidentally remove an extension that other objects may depend on. """ return self
def _fetch_current_extensions(connection: Connection) -> set[str]: """Return extension names currently installed in the database. Args: connection: Active database connection. Returns: Set of installed extension names. """ rows = connection.execute( text("SELECT extname FROM pg_extension") ).fetchall() return {row.extname for row in rows} #: Relations owned by an installed extension, keyed by ``(schema, name)``. #: ``schema`` is always a concrete name (never ``None``); ``public`` for #: the default schema. ExtensionOwnedRelations = set[tuple[str, str]]
[docs] def fetch_extension_owned_relations( connection: Connection, ) -> ExtensionOwnedRelations: """Return relations created and owned by installed extensions. A PostgreSQL extension may create tables, views, sequences, and other relations that it owns (recorded in ``pg_depend`` with ``deptype = 'e'``). The canonical example is PostGIS's ``spatial_ref_sys`` table. These relations are part of the extension, not the application schema, so Alembic autogenerate must not emit ``DROP TABLE`` / ``CREATE TABLE`` for them -- dropping one fails outright (``cannot drop table spatial_ref_sys because extension postgis requires it``). Args: connection: Active database connection. Returns: A set of ``(schema, name)`` pairs for every relation owned by an installed extension. Empty when no extensions own relations. """ 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}
[docs] def register_pg_extension_alembic_events() -> None: """Register the extension comparator and renderer with Alembic. Call once in ``env.py`` alongside :func:`~codegen_database.alembic.register.alembic_hook`:: from codegen_database.pg_extension import ( register_pg_extension_alembic_events, ) register_pg_extension_alembic_events() Safe to call multiple times; subsequent calls are no-ops. """ global _alembic_registered # noqa: PLW0603 if _alembic_registered: return _alembic_registered = True def _compare_extensions( autogen_context: AutogenContext, upgrade_ops: UpgradeOps, _schemas: object, ) -> None: holder = PGExtensions.extract(autogen_context.metadata) if not holder: return assert autogen_context.connection is not None # noqa: S101 current = _fetch_current_extensions(autogen_context.connection) seen: set[str] = set() for ext in holder.extensions: if ext.name in seen: continue seen.add(ext.name) if ext.name not in current: upgrade_ops.ops.append(CreateExtensionOp(extension=ext)) 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)