From dd446c96ed6bef4bfeae0c6748dac1f5a6421aee Mon Sep 17 00:00:00 2001 From: Daniel Edgecombe Date: Mon, 21 Sep 2026 10:42:54 +0100 Subject: [PATCH] feat: Support packages defining pogo plugins to provide additional schema management. --- .pre-commit-config.yaml | 6 ++-- pyproject.toml | 1 + src/pogo_core/migration.py | 3 +- src/pogo_core/util/migrate.py | 8 +++-- src/pogo_core/util/plugins.py | 21 +++++++++++ src/pogo_core/util/sql.py | 21 +++++++++-- tests/util/test_migrate.py | 67 ++++++++++++++++++++++++++++++++++- tests/util/test_plugins.py | 55 ++++++++++++++++++++++++++++ 8 files changed, 172 insertions(+), 10 deletions(-) create mode 100644 src/pogo_core/util/plugins.py create mode 100644 tests/util/test_plugins.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index d430b43..c623ba4 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -11,17 +11,17 @@ repos: - id: end-of-file-fixer - repo: https://github.com/crate-ci/typos - rev: v1 + rev: v1.50.2 hooks: - id: typos args: [--force-exclude] - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.5 + rev: v0.16.8 hooks: - id: ruff-format args: [--preview, -s] - - id: ruff + - id: ruff-check args: [--fix] - repo: local diff --git a/pyproject.toml b/pyproject.toml index 5a31d0f..ff36b83 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -124,6 +124,7 @@ output-format = "concise" [tool.ruff.lint] select = ["ALL"] ignore = [ + "CPY", "D", "FIX", "TD003", diff --git a/src/pogo_core/migration.py b/src/pogo_core/migration.py index f8fcd0d..1d2a4ec 100644 --- a/src/pogo_core/migration.py +++ b/src/pogo_core/migration.py @@ -90,11 +90,12 @@ async def rollback(db: asyncpg.Connection) -> None: class Migration: __migrations: t.ClassVar[dict[str, Migration]] = {} - def __init__(self, mig_id: str, path: Path, applied_migrations: set[str] | None) -> None: + def __init__(self, mig_id: str, path: Path, applied_migrations: set[str] | None, plugin: str = "") -> None: applied_migrations = applied_migrations or set() self.id = mig_id self.path = path self.hash: str = hashlib.sha256(mig_id.encode("utf-8")).hexdigest() + self.plugin = plugin self._use_transaction: bool = True self._doc: str | None = None self._depends: set[Migration] | None = None diff --git a/src/pogo_core/util/migrate.py b/src/pogo_core/util/migrate.py index 40c8350..9992e88 100644 --- a/src/pogo_core/util/migrate.py +++ b/src/pogo_core/util/migrate.py @@ -42,10 +42,11 @@ async def apply( *, schema_name: str, logger: Logger | logging.Logger | None = None, + include_plugins: bool = False, ) -> None: logger = logger or logger_ await sql.ensure_pogo_sync(db) - migrations = await sql.read_migrations(migrations_dir, db, schema_name=schema_name) + migrations = await sql.read_migrations(migrations_dir, db, schema_name=schema_name, include_plugins=include_plugins) migrations = topological_sort([m.load() for m in migrations]) for migration in migrations: @@ -61,17 +62,18 @@ async def apply( raise error.BadMigrationError(msg) from e -async def rollback( +async def rollback( # noqa: PLR0913 db: asyncpg.Connection, migrations_dir: Path, *, schema_name: str, count: int | None = None, logger: Logger | logging.Logger | None = None, + include_plugins: bool = False, ) -> None: logger = logger or logger_ await sql.ensure_pogo_sync(db) - migrations = await sql.read_migrations(migrations_dir, db, schema_name=schema_name) + migrations = await sql.read_migrations(migrations_dir, db, schema_name=schema_name, include_plugins=include_plugins) migrations = reversed(list(topological_sort([m.load() for m in migrations]))) i = 0 diff --git a/src/pogo_core/util/plugins.py b/src/pogo_core/util/plugins.py new file mode 100644 index 0000000..daa557d --- /dev/null +++ b/src/pogo_core/util/plugins.py @@ -0,0 +1,21 @@ +import dataclasses +from importlib.metadata import entry_points +from pathlib import Path + + +@dataclasses.dataclass +class Plugin: + migrations: Path + schema: str | None = None # Inherit parent schema + + +def discover_migrations() -> list[tuple[str, Plugin]]: + plugins = entry_points(group="pogo") + + ret = [] + for plugin in plugins: + p = plugin.load() + if hasattr(p, "plugin") and isinstance(p.plugin, Plugin): + ret.append((plugin.name, p.plugin)) + + return sorted(ret) diff --git a/src/pogo_core/util/sql.py b/src/pogo_core/util/sql.py index dff214d..cbda026 100644 --- a/src/pogo_core/util/sql.py +++ b/src/pogo_core/util/sql.py @@ -5,6 +5,7 @@ import asyncpg from pogo_core.migration import Migration +from pogo_core.util import plugins if t.TYPE_CHECKING: from pathlib import Path @@ -32,14 +33,30 @@ async def read_migrations( db: asyncpg.Connection | None, *, schema_name: str, + plugin: str = "", + include_plugins: bool = False, ) -> list[Migration]: applied_migrations = await get_applied_migrations(db, schema_name=schema_name) if db else set() - return [ - Migration(path.stem, path, applied_migrations) + + migrations = [ + Migration(path.stem, path, applied_migrations, plugin=plugin) for path in migrations_location.iterdir() # noqa: ASYNC240 if path.suffix in {".py", ".sql"} ] + if include_plugins: + plugins_ = plugins.discover_migrations() + for plugin_name, plugin_ in plugins_: + migrations_ = await read_migrations( + plugin_.migrations, + db, + schema_name=plugin_.schema or schema_name, + plugin=plugin_name, + ) + migrations.extend(migrations_) + + return migrations + async def get_applied_migrations(db: asyncpg.Connection, *, schema_name: str) -> set[str]: stmt = """ diff --git a/tests/util/test_migrate.py b/tests/util/test_migrate.py index d9a983b..958fd84 100644 --- a/tests/util/test_migrate.py +++ b/tests/util/test_migrate.py @@ -1,9 +1,54 @@ +import sys +from importlib.metadata import EntryPoint from unittest import mock import pytest from pogo_core import error -from pogo_core.util import migrate, sql +from pogo_core.util import migrate, plugins, sql + + +@pytest.fixture +def package(cwd): + p = cwd / "package" + p.mkdir() + + with (p / "pogo.py").open("w") as f: + f.write(""" +from pathlib import Path + +from pogo_core.util.plugins import Plugin + +plugin = Plugin(Path(__file__).parent / "migrations") +""") + + sys.path.insert(0, str(cwd)) + return p + + +@pytest.fixture +def package_migrations(package): + p = package / "migrations" + p.mkdir() + + return p + + +@pytest.fixture +def _package_migration_one(package_migrations): + p = package_migrations / "20250317_01_abcde-initial-migration.sql" + + with p.open("w") as f: + f.write(""" +-- initial migration +-- depends: + +-- migrate: apply +CREATE TABLE package_table_one(); + +-- migrate: rollback +DROP TABLE package_table_one; +""") @pytest.fixture @@ -154,6 +199,26 @@ async def test_broken_migration_not_applied(self, migrations, db_session): ) assert str(e.value) == "Failed to apply 20240318_01_12345-broken-apply" + @pytest.mark.usefixtures("_migration_two", "_package_migration_one") + async def test_package_migrations_applied(self, migrations, db_session, monkeypatch): + monkeypatch.setattr( + plugins, + "entry_points", + mock.Mock(return_value=[EntryPoint(name="name", group=None, value="package.pogo")]), + ) + await migrate.apply(db_session, migrations, schema_name="public", include_plugins=True) + + await self.assert_tables( + db_session, + [ + "public._pogo_migration", + "public._pogo_version", + "public.package_table_one", + "public.table_one", + "public.table_two", + ], + ) + class TestRollback(Base): @pytest.mark.usefixtures("migrations") diff --git a/tests/util/test_plugins.py b/tests/util/test_plugins.py new file mode 100644 index 0000000..e0a81a8 --- /dev/null +++ b/tests/util/test_plugins.py @@ -0,0 +1,55 @@ +from unittest import mock + +from pogo_core.util import plugins + + +class EntryPoint: + def __init__(self, name: str) -> None: + self.name = name + + def load(self) -> object: + return self + + +class PluginEntryPoint: + def __init__(self, name: str, plugin: plugins.Plugin) -> None: + self.name = name + self.plugin = plugin + + def load(self) -> object: + return mock.Mock( + plugin=self.plugin, + ) + + +def test_invalid_plugins_ignored(monkeypatch): + monkeypatch.setattr( + plugins, + "entry_points", + mock.Mock( + return_value=[ + EntryPoint("fail"), + PluginEntryPoint("invalid_type", mock.Mock()), + ], + ), + ) + + assert plugins.discover_migrations() == [] + + +def test_plugins_discovered(monkeypatch): + monkeypatch.setattr( + plugins, + "entry_points", + mock.Mock( + return_value=[ + PluginEntryPoint("no_schema", plugins.Plugin("./")), + PluginEntryPoint("schema", plugins.Plugin("./", schema="schema")), + ], + ), + ) + + assert plugins.discover_migrations() == [ + ("no_schema", plugins.Plugin("./", schema=None)), + ("schema", plugins.Plugin("./", schema="schema")), + ]