diff --git a/kanta/kantaimpl.py b/kanta/kantaimpl.py index 186f667..5f83a3c 100644 --- a/kanta/kantaimpl.py +++ b/kanta/kantaimpl.py @@ -139,8 +139,11 @@ class KantaImpl(PersistenceMixin, Generic[T]): cause_type=type(e).__name__, ) from e + migrations_ran = False if self.migrations is not None: + previous_version = rr.version rr.version = self.migrations.apply(rr.state, rr.version, self._kanta) + migrations_ran = rr.version != previous_version self.statedict = copy.deepcopy(rr.state) self.data = restore_data_in_place( @@ -155,6 +158,13 @@ class KantaImpl(PersistenceMixin, Generic[T]): if self.readonly: self.statedict = copy.deepcopy(normalized) else: + if migrations_ran: + self.queue_change( + f"migrate:v{self.version}", + self.statedict, + mtime=False, + force=True, + ) self.queue_change("migrate:msgspec", normalized, mtime=False) self.snapshot.ts = ( datetime.fromtimestamp(rr.last_snapshot_mtime, UTC) diff --git a/kanta/migrations.py b/kanta/migrations.py index 6e89882..b561f3d 100644 --- a/kanta/migrations.py +++ b/kanta/migrations.py @@ -6,12 +6,15 @@ or by prefix. Each runs exactly once based on the current version. from __future__ import annotations +import copy import importlib import inspect import logging from types import ModuleType from typing import Any +from kanta.exceptions import DatabaseError + _logger = logging.getLogger(__name__) # Cache registries by imported module object so that many Kanta instances using @@ -94,6 +97,11 @@ class Migrations: """Current schema version (= highest discovered migration, or 0).""" return max(self._migrations.keys(), default=0) + @property + def minver(self) -> int: + """Minimum supported current version (first migration minus 1, or 0).""" + return min(self._migrations.keys(), default=1) - 1 + @staticmethod def _call_migration(fn: Any, data_dict: dict[str, Any], kanta: Any) -> None: """Call *fn* with the data dict and, if accepted, the Kanta instance.""" @@ -114,19 +122,33 @@ class Migrations: ) -> int: """Apply pending migrations to *data_dict* in place. + Missing intermediate migration steps are silently skipped. + + Raises: + DatabaseError: If the database version is newer than the highest + supported version or older than the minimum supported version. + Returns the new version after all migrations. """ - while current_version < self.dbver: - next_version = current_version + 1 - fn = self._migrations.get(next_version) - if fn is None: - raise ValueError( - f"Missing migration step migrate_v{next_version} " - f"(highest discovered is v{self.dbver})" - ) + if current_version > self.dbver: + raise DatabaseError( + f"Database version v{current_version} is newer than the " + f"highest supported version v{self.dbver}" + ) + if current_version < self.minver: + raise DatabaseError( + f"Database version v{current_version} is older than the " + f"minimum supported version v{self.minver}" + ) + + for version in sorted(self._migrations.keys()): + if version <= current_version: + continue + fn = self._migrations[version] + before = copy.deepcopy(data_dict) if not silent else None self._call_migration(fn, data_dict, kanta) - current_version = next_version - if not silent: + current_version = version + if not silent and before != data_dict: desc = (fn.__doc__ or fn.__name__).split("\n")[0].rstrip(".") _logger.info("Applied migration %s: %s", fn.__name__, desc) return current_version diff --git a/kanta/persistence.py b/kanta/persistence.py index 10413e9..5113256 100644 --- a/kanta/persistence.py +++ b/kanta/persistence.py @@ -140,8 +140,10 @@ class PersistenceMixin: raise TypeError("mtime must be True, False, or a datetime") diff = compute_diff(self.statedict, current) - if not diff and not force: - return None + if not diff: + if not force: + return None + diff = {} record = ChangeRecord( ts=now, diff --git a/tests/test_kanta_integration.py b/tests/test_kanta_integration.py index bcd7a16..fa6e444 100644 --- a/tests/test_kanta_integration.py +++ b/tests/test_kanta_integration.py @@ -17,6 +17,7 @@ from .support import ( change_actions, fixed_change, make_kanta, + make_migrations_module, read_changes, seed_single_change, ) @@ -473,6 +474,43 @@ async def test_msgspec_normalization_logs_migration(tmp_path, format_config): assert "migrate:msgspec" in change_actions(path, format_config) +@pytest.mark.asyncio +async def test_empty_migration_is_recorded_and_not_reapplied(tmp_path, format_config): + path = tmp_path / "test.db" + seed_single_change( + path, fixed_change("init", {"counter": 0, "users": {}}), format_config + ) + + def migrate_v1(d, kanta): + """No-op migration that only bumps the schema version.""" + pass + + mod = make_migrations_module("empty_migration_mod", "migrate_v1", migrate_v1) + + try: + kanta = make_kanta(path, Data, format_config, migrations=mod) + await kanta.open() + assert kanta.version == 1 + await kanta.flush() + await kanta.close() + + records = read_changes(path, format_config) + migration_records = [r for r in records if r.a.startswith("migrate")] + assert len(migration_records) == 1 + assert migration_records[0].v == 1 + assert migration_records[0].diff == {} + + kanta2 = make_kanta(path, Data, format_config, migrations=mod) + await kanta2.open() + assert kanta2.version == 1 + await kanta2.close() + + records2 = read_changes(path, format_config) + assert len([r for r in records2 if r.a.startswith("migrate")]) == 1 + finally: + sys.modules.pop("empty_migration_mod", None) + + @pytest.mark.asyncio async def test_open_locked_file_raises_filelock_error(tmp_path, format_config): path = tmp_path / "test.db" diff --git a/tests/test_migrations.py b/tests/test_migrations.py index a2b789d..179ffec 100644 --- a/tests/test_migrations.py +++ b/tests/test_migrations.py @@ -1,5 +1,9 @@ +import logging from types import ModuleType, SimpleNamespace +import pytest + +from kanta.exceptions import DatabaseError from kanta.migrations import Migrations @@ -89,3 +93,112 @@ def test_migration_can_omit_kanta_argument(): new_ver = reg.apply(state, current_version=0, kanta=kanta, silent=True) assert new_ver == 1 assert state["x"] == 1 + + +def test_version_too_new(): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v1(d): + d["x"] = 1 + + with pytest.raises( + DatabaseError, + match="Database version v2 is newer than the highest supported version v1", + ): + reg.apply({}, current_version=2, kanta=kanta, silent=True) + + +def test_version_too_old(): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v3(d): + d["x"] = 3 + + with pytest.raises( + DatabaseError, + match="Database version v1 is older than the minimum supported version v2", + ): + reg.apply({}, current_version=1, kanta=kanta, silent=True) + + +def test_missing_middle_migration_is_skipped(): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v1(d): + d["x"] = 1 + + @reg.register + def migrate_v3(d): + d["y"] = 3 + + state = {"x": 1} + new_ver = reg.apply(state, current_version=1, kanta=kanta, silent=True) + assert new_ver == 3 + assert state["x"] == 1 + assert state["y"] == 3 + + +def test_old_migrations_deleted_current_supported(): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v3(d): + d["x"] = 3 + + state = {"x": 2} + new_ver = reg.apply(state, current_version=2, kanta=kanta, silent=True) + assert new_ver == 3 + assert state["x"] == 3 + + +def test_migration_log_only_when_changed(caplog): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v1(d): + """Set x.""" + d["x"] = 1 + + @reg.register + def migrate_v2(d): + """No-op.""" + pass + + @reg.register + def migrate_v3(d): + """Set y.""" + d["y"] = 3 + + with caplog.at_level(logging.INFO, logger="kanta.migrations"): + reg.apply({}, current_version=0, kanta=kanta) + + messages = [r.message for r in caplog.records if r.levelno == logging.INFO] + assert len(messages) == 2 + assert "migrate_v1" in messages[0] + assert "Set x" in messages[0] + assert "migrate_v3" in messages[1] + assert "Set y" in messages[1] + + +def test_no_op_migration_produces_no_log(caplog): + reg = Migrations() + kanta = _DummyKanta() + + @reg.register + def migrate_v1(d): + """No-op.""" + pass + + with caplog.at_level(logging.INFO, logger="kanta.migrations"): + reg.apply({}, current_version=0, kanta=kanta) + + info_messages = [r for r in caplog.records if r.levelno == logging.INFO] + assert not info_messages