Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,7 @@ output-format = "concise"
[tool.ruff.lint]
select = ["ALL"]
ignore = [
"CPY",
"D",
"FIX",
"TD003",
Expand Down
3 changes: 2 additions & 1 deletion src/pogo_core/migration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 5 additions & 3 deletions src/pogo_core/util/migrate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down
21 changes: 21 additions & 0 deletions src/pogo_core/util/plugins.py
Original file line number Diff line number Diff line change
@@ -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)
21 changes: 19 additions & 2 deletions src/pogo_core/util/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 = """
Expand Down
67 changes: 66 additions & 1 deletion tests/util/test_migrate.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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")
Expand Down
55 changes: 55 additions & 0 deletions tests/util/test_plugins.py
Original file line number Diff line number Diff line change
@@ -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")),
]
Loading