Source code for codegen_database.extension
"""Extension base class and discovery for codegen_database."""
from __future__ import annotations
from dataclasses import dataclass, field
from importlib.metadata import entry_points
from typing import TYPE_CHECKING, ClassVar
from codegen_database.errors import CodegenDatabaseValidationError
if TYPE_CHECKING:
from sqlalchemy import MetaData
from codegen_database.plugin import Plugin
[docs]
@dataclass
class CodegenDatabaseExtension:
"""Base class for codegen_database extensions.
An extension bundles plugins, metadata hooks, Alembic hooks,
and CLI commands into a single installable unit.
Subclasses override hook methods to participate in the codegen_database
lifecycle. Extensions declare inter-extension dependencies via
the ``depends_on`` class variable.
Example::
@dataclass
class MyExtension(CodegenDatabaseExtension):
name: str = "my-ext"
def plugins(self) -> list[Plugin]:
return [MyGlobalPlugin()]
def configure_metadata(self, metadata: MetaData) -> None:
# register roles, grants, schemas, etc.
...
"""
name: str
depends_on: ClassVar[list[str]] = field(default=[], init=False, repr=False)
[docs]
def plugins(self) -> list[Plugin]:
"""Global plugins prepended to every factory.
Returns:
List of plugin instances. Empty by default.
"""
return []
[docs]
def register_cli(self, app: object) -> None:
"""Add subcommands during CLI setup.
Args:
app: The ``typer.Typer`` application instance.
"""
[docs]
def validate(self, registered_names: frozenset[str]) -> None:
"""Validate after all extensions are loaded.
Override to check that required peer extensions are
present or that configuration is consistent.
Args:
registered_names: Names of all loaded extensions.
"""
[docs]
def discover_extensions() -> dict[str, type[CodegenDatabaseExtension]]:
"""Discover extensions via the ``codegen_database.ext`` entry point group.
Returns:
Mapping of extension name to extension class.
"""
eps = entry_points(group="codegen_database.ext")
return {ep.name: ep.load() for ep in eps}
[docs]
def validate_extension_deps(
extensions: list[CodegenDatabaseExtension],
) -> None:
"""Check that every extension's ``depends_on`` is satisfied.
Args:
extensions: The resolved list of extension instances.
Raises:
CodegenDatabaseValidationError: If a dependency is missing.
"""
names = frozenset(ext.name for ext in extensions)
for ext in extensions:
deps: list[str] = getattr(type(ext), "depends_on", [])
missing = [d for d in deps if d not in names]
if missing:
msg = (
f"Extension {ext.name!r} depends on "
f"{missing!r}, which are not registered."
)
raise CodegenDatabaseValidationError(msg)
for ext in extensions:
ext.validate(names)