5 Commits
Author SHA1 Message Date
LeoVasanko 753b7eba86 Add separate migration logging, logmigr callback, and log= override
- Split transaction logger (kanta.changes) and migration logger (kanta.migrations).
- Migrations.apply() now returns MigrationResult instead of logging.
- Kanta.open() emits one info summary per DB and debug transaction per migration.
- Add open(log=...) to suppress/redirect default migration logging.
- Add @kanta.logmigr callback for custom migration logging/summaries.
- Add transaction(log=...) to suppress/redirect transaction logging.
- Update tests for the new MigrationResult API and logging behaviour.
2026-06-15 23:13:03 +00:00
LeoVasanko 42789e6619 Revised migration context: don't store changerecord unless something was changed, and always snapshot if and after changes done. 2026-06-15 22:14:03 +00:00
LeoVasanko c4726e6728 Stricter versioning: always store migrate version record, file must be within versions included in migrations. Log only for migrations that made changes. Allow deleting older migration functions when no longer required. 2026-06-15 03:09:45 +00:00
LeoVasanko 66e92739ab Refactor Migrations internals, add caching to avoid reloading per each Kanta instance. Naming changed from MigrationRegistry to Migration, module from migrate to kanta.migrations. 2026-06-15 01:12:42 +00:00
LeoVasanko 4dc2f0648e Fix creation of new database, ensuring that a bootstrap record of the initial state is always written. 2026-06-15 00:27:49 +00:00
19 changed files with 859 additions and 205 deletions
+9 -5
View File
@@ -53,10 +53,13 @@ asyncio.run(main())
## Bootstrap and Open Modes ## Bootstrap and Open Modes
Kanta supports open-time bootstrap callbacks for initializing a brand-new When `open()` creates a brand-new database, it always writes a single bootstrap
database before `open()` returns. change record from the initial data object you passed to `Kanta(...)`. The
simplest bootstrap is therefore the object itself — no extra code is required.
Register bootstrap handlers with a decorator: Bootstrap handlers are optional. Use them only when you need to modify the
initial state at creation time, for example to seed defaults or perform
expensive/external setup that should happen exactly once:
```python ```python
kanta = Kanta("data.kantadb", Data()) kanta = Kanta("data.kantadb", Data())
@@ -76,9 +79,10 @@ async def bootstrap_async(data) -> None:
data.counter = 1 data.counter = 1
``` ```
When multiple bootstrap handlers are registered: Whether or not handlers are registered, exactly one bootstrap change record is
written when a new database is created. The record contains the initial object,
or the state after all bootstrap handlers have run. When handlers are present:
- they run in registration order, - they run in registration order,
- exactly one bootstrap change record is queued,
- bootstrap metadata (`action`, `user`, `mtime`) is taken from the last - bootstrap metadata (`action`, `user`, `mtime`) is taken from the last
registration. registration.
+7 -1
View File
@@ -133,7 +133,11 @@ when they have a default value.
#### Bootstrap Callbacks #### Bootstrap Callbacks
- Bootstrap callbacks run during `open()` when the database is empty. - When `open()` creates a new database, it always writes a single bootstrap
`ChangeRecord`.
- The simplest bootstrap is the initial data object passed to `Kanta(...)`;
bootstrap callbacks are optional and only needed when you want to modify or
enrich that object at creation time.
- Register callbacks via: - Register callbacks via:
- `@kanta.bootstrap` - `@kanta.bootstrap`
- `@kanta.bootstrap(action=..., user=..., mtime=...)` - `@kanta.bootstrap(action=..., user=..., mtime=...)`
@@ -146,6 +150,8 @@ when they have a default value.
- exactly one bootstrap `ChangeRecord` is queued, - exactly one bootstrap `ChangeRecord` is queued,
- bootstrap metadata (`action`, `user`, `mtime`) is taken from the last - bootstrap metadata (`action`, `user`, `mtime`) is taken from the last
callback registration. callback registration.
- If no bootstrap callbacks are registered, the bootstrap record still uses
`action="bootstrap"` and contains the initial data object.
- If any bootstrap callback raises, Kanta closes and removes the database file, - If any bootstrap callback raises, Kanta closes and removes the database file,
then re-raises the exception. then re-raises the exception.
+10 -1
View File
@@ -18,6 +18,7 @@ from dataclasses import dataclass
from typing import Annotated, Any, Union, get_args, get_origin from typing import Annotated, Any, Union, get_args, get_origin
from kanta.exceptions import DatabaseError from kanta.exceptions import DatabaseError
from kanta.migrations import MigrationResult
DictPre = Annotated[dict, "pre"] DictPre = Annotated[dict, "pre"]
DictPost = Annotated[dict, "post"] DictPost = Annotated[dict, "post"]
@@ -59,6 +60,7 @@ class InjectionContext:
error: DatabaseError | None = None error: DatabaseError | None = None
previous_state: dict | None = None previous_state: dict | None = None
current_state: dict | None = None current_state: dict | None = None
migration_result: MigrationResult | None = None
@dataclass @dataclass
@@ -98,6 +100,7 @@ class CallbackRegistry:
self._callbacks: dict[str, list[_CallbackRegistration]] = { self._callbacks: dict[str, list[_CallbackRegistration]] = {
"bootstrap": [], "bootstrap": [],
"fatal_error": [], "fatal_error": [],
"logmigr": [],
} }
self._logfmt_callbacks: list[_LogFmtFunctionSpec | _LogFmtClassSpec] = [] self._logfmt_callbacks: list[_LogFmtFunctionSpec | _LogFmtClassSpec] = []
@@ -446,10 +449,12 @@ class CallbackRegistry:
return kind == "logfmt" return kind == "logfmt"
if bare is DatabaseError: if bare is DatabaseError:
return kind == "fatal_error" return kind == "fatal_error"
if bare is MigrationResult:
return kind == "logmigr"
if self._data_type is not None and bare is self._data_type: if self._data_type is not None and bare is self._data_type:
return kind == "bootstrap" return kind == "bootstrap"
if self._kanta_class is not None and bare is self._kanta_class: if self._kanta_class is not None and bare is self._kanta_class:
return kind in {"bootstrap", "fatal_error", "logfmt"} return kind in {"bootstrap", "fatal_error", "logfmt", "logmigr"}
return False return False
def _allowed_message(self, kind: str) -> str: def _allowed_message(self, kind: str) -> str:
@@ -462,6 +467,8 @@ class CallbackRegistry:
parts.append(self._kanta_class.__name__) parts.append(self._kanta_class.__name__)
if kind == "fatal_error": if kind == "fatal_error":
parts.append("DatabaseError") parts.append("DatabaseError")
if kind == "logmigr":
parts.append("MigrationResult")
if kind == "logfmt": if kind == "logfmt":
parts.append("Annotated[dict, 'pre']") parts.append("Annotated[dict, 'pre']")
parts.append("Annotated[dict, 'post']") parts.append("Annotated[dict, 'post']")
@@ -475,6 +482,8 @@ class CallbackRegistry:
return ctx.current_state return ctx.current_state
if bare is DatabaseError: if bare is DatabaseError:
return ctx.error return ctx.error
if bare is MigrationResult:
return ctx.migration_result
if self._data_type is not None and bare is self._data_type: if self._data_type is not None and bare is self._data_type:
return ctx.data return ctx.data
if self._kanta_class is not None and bare is self._kanta_class: if self._kanta_class is not None and bare is self._kanta_class:
+3 -1
View File
@@ -129,7 +129,9 @@ class LockedFile:
else: else:
self._open_unix(path, create, readonly) self._open_unix(path, create, readonly)
def open_and_read(self, path: Path, create: bool = False, readonly: bool = False) -> bytes: def open_and_read(
self, path: Path, create: bool = False, readonly: bool = False
) -> bytes:
"""Open *path* and read all content. """Open *path* and read all content.
Combined operation for efficient use with asyncio.to_thread(). Combined operation for efficient use with asyncio.to_thread().
+39 -3
View File
@@ -1,10 +1,11 @@
"""Kanta DB main public API""" """Kanta DB main public API"""
from __future__ import annotations from __future__ import annotations
import logging
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from types import ModuleType, SimpleNamespace from types import ModuleType, SimpleNamespace
from typing import Any, Generic, TypeVar from typing import Generic, TypeVar
from kanta.kantaimpl import KantaImpl from kanta.kantaimpl import KantaImpl
from kanta.serialization import JsonSerializer, Serializer from kanta.serialization import JsonSerializer, Serializer
@@ -144,7 +145,13 @@ class Kanta(Generic[T]):
""" """
return self._impl.mtime return self._impl.mtime
async def open(self, *, create: bool = True, readonly: bool = False) -> None: async def open(
self,
*,
create: bool = True,
readonly: bool = False,
log: bool | logging.Logger = True,
) -> None:
"""Open the database file and start background persistence. """Open the database file and start background persistence.
This loads existing records, applies configured migrations, and starts This loads existing records, applies configured migrations, and starts
@@ -156,6 +163,11 @@ class Kanta(Generic[T]):
readonly: If True, open the database read-only. No lock is acquired, readonly: If True, open the database read-only. No lock is acquired,
no background flush task is started, and transactions are no background flush task is started, and transactions are
rejected. The file is not created if missing. rejected. The file is not created if missing.
log: Controls migration logging. ``True`` (default) uses the
``kanta.migrations`` logger. ``False`` suppresses the default
migration log. A :class:`~logging.Logger` instance writes
default migration output to that logger instead. Custom
``@kanta.logmigr`` callbacks run regardless of this setting.
Calling ``open`` more than once on the same instance is not allowed. Calling ``open`` more than once on the same instance is not allowed.
@@ -163,7 +175,7 @@ class Kanta(Generic[T]):
kanta.exceptions.DatabaseError: If replay or decoding fails. kanta.exceptions.DatabaseError: If replay or decoding fails.
kanta.exceptions.DataIntegrityError: If the instance is already open. kanta.exceptions.DataIntegrityError: If the instance is already open.
""" """
await self._impl.open(create=create, readonly=readonly) await self._impl.open(create=create, readonly=readonly, log=log)
async def __aenter__(self) -> Kanta[T]: async def __aenter__(self) -> Kanta[T]:
"""Enter async context manager and open the database. """Enter async context manager and open the database.
@@ -239,6 +251,24 @@ class Kanta(Generic[T]):
return _register return _register
return _register(fn) return _register(fn)
def logmigr(self, fn=None):
"""Register a migration logging callback.
Can be used as ``@kanta.logmigr``.
The callback receives a :class:`kanta.migrations.MigrationResult` and
may be sync or async. If registered, it replaces the default migration
logger output; the application is responsible for emitting any log
messages.
"""
def _register(callback):
self._impl.add_logmigr(callback)
return callback
if fn is None:
return _register
return _register(fn)
def logfmt(self, fn=None, *, path: str | None = None): def logfmt(self, fn=None, *, path: str | None = None):
"""Register a transaction logfmt callback. """Register a transaction logfmt callback.
@@ -266,6 +296,7 @@ class Kanta(Generic[T]):
*, *,
user: str | None = None, user: str | None = None,
mtime: bool | datetime = True, mtime: bool | datetime = True,
log: bool | logging.Logger = True,
): ):
"""Create a transactional mutation context manager. """Create a transactional mutation context manager.
@@ -280,6 +311,10 @@ class Kanta(Generic[T]):
system operations that are not considered modifications. A system operations that are not considered modifications. A
:class:`~datetime.datetime` value sets ``m`` to that explicit :class:`~datetime.datetime` value sets ``m`` to that explicit
time. time.
log: Controls transaction logging. ``True`` (default) uses the
``kanta.changes`` logger. ``False`` suppresses the transaction
log. A :class:`~logging.Logger` instance writes output to that
logger instead.
Returns: Returns:
A context manager yielding the live state object for mutation. A context manager yielding the live state object for mutation.
@@ -294,4 +329,5 @@ class Kanta(Generic[T]):
action, action,
user=user, user=user,
mtime=mtime, mtime=mtime,
log=log,
) )
+125 -28
View File
@@ -12,7 +12,8 @@ from typing import Any, Generic, TypeVar
from kanta.callbacks import CallbackRegistry, InjectionContext from kanta.callbacks import CallbackRegistry, InjectionContext
from kanta.exceptions import DatabaseError, DataIntegrityError, ReplayError from kanta.exceptions import DatabaseError, DataIntegrityError, ReplayError
from kanta.migrate import MigrationRegistry from kanta.logging import log_change, migration_logger
from kanta.migrations import MigrationResult, Migrations
from kanta.persistence import PersistenceMixin from kanta.persistence import PersistenceMixin
from kanta.serialization import restore_data_in_place, struct_to_dict from kanta.serialization import restore_data_in_place, struct_to_dict
from kanta.serialization.base import replay from kanta.serialization.base import replay
@@ -28,18 +29,18 @@ class KantaImpl(PersistenceMixin, Generic[T]):
def __init__(self, **kwargs: Any): def __init__(self, **kwargs: Any):
self.data_type = kwargs.pop("type") self.data_type = kwargs.pop("type")
self.data: T = kwargs.pop("data") self.data: T = kwargs.pop("data")
self.migrations = kwargs.pop("migrations", None)
self._kanta = kwargs.pop("kanta", None) self._kanta = kwargs.pop("kanta", None)
migrations = kwargs.pop("migrations", None)
self.ctx = SimpleNamespace() self.ctx = SimpleNamespace()
super().__init__(**kwargs) super().__init__(**kwargs)
self.migration_registry: MigrationRegistry | None = None self.migrations: Migrations | None = None
if self.migrations is not None: if migrations is not None:
module = ( module = (
importlib.import_module(self.migrations) importlib.import_module(migrations)
if isinstance(self.migrations, str) if isinstance(migrations, str)
else self.migrations else migrations
) )
self.migration_registry = MigrationRegistry.from_module(module) self.migrations = Migrations.from_module(module)
self.in_transaction = False self.in_transaction = False
self.transaction_snapshot: dict[str, Any] | None = None self.transaction_snapshot: dict[str, Any] | None = None
@@ -55,9 +56,7 @@ class KantaImpl(PersistenceMixin, Generic[T]):
) )
self.statedict = struct_to_dict(self.data, serializer=self.serializer) self.statedict = struct_to_dict(self.data, serializer=self.serializer)
self.version = ( self.version = self.migrations.dbver if self.migrations is not None else 0
self.migration_registry.dbver if self.migration_registry is not None else 0
)
def add_bootstrap( def add_bootstrap(
self, self,
@@ -77,7 +76,64 @@ class KantaImpl(PersistenceMixin, Generic[T]):
"""Register one transaction logfmt callback.""" """Register one transaction logfmt callback."""
self.callback_registry.register("logfmt", callback, path=path) self.callback_registry.register("logfmt", callback, path=path)
async def open(self, *, create: bool = True, readonly: bool = False) -> None: def add_logmigr(self, callback) -> None:
"""Register one migration logging callback."""
self.callback_registry.register("logmigr", callback)
async def _handle_migration_log(
self,
migration_result: MigrationResult,
previous_version: int,
log: bool | logging.Logger,
) -> None:
"""Route migration logging to callback or default logger."""
assert isinstance(migration_result, MigrationResult)
if self.callback_registry.has("logmigr"):
await self.callback_registry.invoke(
"logmigr",
InjectionContext(
kanta=self._kanta,
migration_result=migration_result,
),
)
return
if log is False:
return
migration_log = log if isinstance(log, logging.Logger) else migration_logger
changed = [m for m in migration_result.migrations if m.changed]
if not changed:
return
for info in changed:
if info.diff:
log_change(
info.name,
info.diff,
previous=info.before,
logger=migration_log,
level=logging.DEBUG,
)
descriptions = [f"{m.name} ({m.description})" for m in changed]
migration_log.info(
"Migrated %s v%s -> v%s: %s",
self.filename,
previous_version,
migration_result.version,
", ".join(descriptions),
)
async def open(
self,
*,
create: bool = True,
readonly: bool = False,
log: bool | logging.Logger = True,
) -> None:
"""Open the database: load from disk, apply migrations, start background task.""" """Open the database: load from disk, apply migrations, start background task."""
if self.opened: if self.opened:
raise DataIntegrityError( raise DataIntegrityError(
@@ -112,6 +168,9 @@ class KantaImpl(PersistenceMixin, Generic[T]):
action="open", action="open",
) )
# From this point the file is open and must be closed via close().
self.opened = True
if content: if content:
try: try:
rr = replay( rr = replay(
@@ -141,12 +200,33 @@ class KantaImpl(PersistenceMixin, Generic[T]):
cause_type=type(e).__name__, cause_type=type(e).__name__,
) from e ) from e
if self.migration_registry is not None: migration_result = None
rr.version = self.migration_registry.apply( state_before_migrations = None
previous_version = rr.version
if self.migrations is not None:
state_before_migrations = copy.deepcopy(rr.state)
migration_result = self.migrations.apply(
rr.state, rr.version, self._kanta rr.state, rr.version, self._kanta
) )
rr.version = migration_result.version
self.statedict = copy.deepcopy(rr.state) migrations_ran = rr.version != previous_version
migration_state_changed = (
state_before_migrations is not None
and state_before_migrations != rr.state
)
self.snapshot.ts = (
datetime.fromtimestamp(rr.last_snapshot_mtime, UTC)
if rr.last_snapshot_mtime is not None
else None
)
self.statedict = copy.deepcopy(
state_before_migrations
if state_before_migrations is not None
else rr.state
)
self.data = restore_data_in_place( self.data = restore_data_in_place(
self.data, self.data,
rr.state, rr.state,
@@ -159,34 +239,53 @@ class KantaImpl(PersistenceMixin, Generic[T]):
if self.readonly: if self.readonly:
self.statedict = copy.deepcopy(normalized) self.statedict = copy.deepcopy(normalized)
else: else:
self.queue_change("migrate:msgspec", normalized, mtime=False) if migrations_ran and migration_state_changed:
self.snapshot.ts = ( self.queue_change(
datetime.fromtimestamp(rr.last_snapshot_mtime, UTC) f"migrate:v{self.version}",
if rr.last_snapshot_mtime is not None rr.state,
else None mtime=False,
) )
msgspec_record = self.queue_change(
"migrate:msgspec", normalized, mtime=False
)
if migrations_ran or msgspec_record is not None:
self.snapshot.request_force()
await self.flush()
self.snapshot.maybe_write(
self.file, self.version, self.statedict, m=self.mtime
)
if migrations_ran and migration_result is not None:
await self._handle_migration_log(
migration_result, previous_version, log
)
elif self.readonly: elif self.readonly:
self.opened = False
self.file.close() self.file.close()
raise DataIntegrityError( raise DataIntegrityError(
"Cannot open empty database in read-only mode", "Cannot open empty database in read-only mode",
db_path=self.filename, db_path=self.filename,
action="open", action="open",
) )
elif self.callback_registry.has("bootstrap"): else:
try: try:
await self.callback_registry.invoke( if self.callback_registry.has("bootstrap"):
"bootstrap", await self.callback_registry.invoke(
InjectionContext(data=self.data, kanta=self._kanta), "bootstrap",
) InjectionContext(data=self.data, kanta=self._kanta),
)
self.statedict = {}
current = struct_to_dict(self.data, serializer=self.serializer) current = struct_to_dict(self.data, serializer=self.serializer)
self.queue_change( self.queue_change(
self.bootstrap_action, self.bootstrap_action,
current, current,
user=self.bootstrap_user, user=self.bootstrap_user,
mtime=self.bootstrap_mtime, mtime=self.bootstrap_mtime,
force=True,
) )
except Exception: except Exception:
self.opened = False
self.file.close() self.file.close()
try: try:
await asyncio.to_thread(self.filename.unlink, missing_ok=True) await asyncio.to_thread(self.filename.unlink, missing_ok=True)
@@ -194,8 +293,6 @@ class KantaImpl(PersistenceMixin, Generic[T]):
pass pass
raise raise
self.opened = True
if not self.readonly: if not self.readonly:
self.background_task = asyncio.create_task(self._background_loop()) self.background_task = asyncio.create_task(self._background_loop())
+15 -9
View File
@@ -10,7 +10,8 @@ import sys
from collections.abc import Callable from collections.abc import Callable
from typing import Any from typing import Any
logger = logging.getLogger("kanta.changes") changes_logger = logging.getLogger("kanta.changes")
migration_logger = logging.getLogger("kanta.migrations")
# Pattern to match control characters and bidirectional overrides # Pattern to match control characters and bidirectional overrides
_UNSAFE_CHARS = re.compile( _UNSAFE_CHARS = re.compile(
@@ -274,6 +275,9 @@ def log_change(
user: str | None = None, user: str | None = None,
previous: dict | None = None, previous: dict | None = None,
logfmt: Callable[[Any, str], str | None] | None = None, logfmt: Callable[[Any, str], str | None] | None = None,
*,
logger: logging.Logger = changes_logger,
level: int = logging.INFO,
) -> None: ) -> None:
"""Log a database change with pretty-printed diff. """Log a database change with pretty-printed diff.
@@ -283,27 +287,29 @@ def log_change(
user: Optional already-formatted user name to show in the header. user: Optional already-formatted user name to show in the header.
previous: The previous state dict (for determining add vs update). previous: The previous state dict (for determining add vs update).
logfmt: Optional formatter callable ``(value, path) -> str | None``. logfmt: Optional formatter callable ``(value, path) -> str | None``.
logger: Logger to write to. Defaults to the ``kanta.changes`` logger.
level: Log level to use. Defaults to ``logging.INFO``.
""" """
header = format_action_header(action, user) header = format_action_header(action, user)
diff_lines = format_diff(diff, previous, logfmt) diff_lines = format_diff(diff, previous, logfmt)
if not diff_lines: if not diff_lines:
logger.info(header) logger.log(level, header)
return return
if len(diff_lines) == 1: if len(diff_lines) == 1:
logger.info(f"{header}{diff_lines[0]}") logger.log(level, f"{header}{diff_lines[0]}")
else: else:
logger.info(header) logger.log(level, header)
for line in diff_lines: for line in diff_lines:
logger.info(line) logger.log(level, line)
def configure_logging() -> None: def configure_logging() -> None:
"""Configure the database logger to output to stderr without prefix.""" """Configure the database logger to output to stderr without prefix."""
if not logger.handlers: if not changes_logger.handlers:
handler = logging.StreamHandler(sys.stderr) handler = logging.StreamHandler(sys.stderr)
handler.setFormatter(logging.Formatter("%(message)s")) handler.setFormatter(logging.Formatter("%(message)s"))
logger.addHandler(handler) changes_logger.addHandler(handler)
logger.setLevel(logging.INFO) changes_logger.setLevel(logging.INFO)
logger.propagate = False changes_logger.propagate = False
-122
View File
@@ -1,122 +0,0 @@
"""Database schema migration framework.
Migrations are numbered functions discovered automatically via a decorator
or by prefix. Each runs exactly once based on the current version.
"""
from __future__ import annotations
import importlib
import inspect
import logging
from types import ModuleType
from typing import Any
_logger = logging.getLogger(__name__)
class MigrationRegistry:
"""Registry of schema migration functions.
Usage::
registry = MigrationRegistry()
@registry.register
def migrate_v1(d: dict, kanta) -> None:
d.setdefault("version", 1)
kanta.ctx.note = "migrated"
@registry.register
def migrate_v2(d: dict) -> None:
d.setdefault("version", 2)
new_version = registry.apply(state, current_version=0, kanta=kanta)
Or load from a module::
registry = MigrationRegistry.from_module("myapp.migrations")
new_version = registry.apply(state, current_version=0, kanta=kanta)
"""
def __init__(self) -> None:
self._migrations: dict[int, Any] = {}
@staticmethod
def _migration_version(fn: Any) -> int:
name = getattr(fn, "__name__", "")
if not name.startswith("migrate_v"):
raise ValueError(f"Invalid migration function name: {name!r}")
suffix = name.removeprefix("migrate_v")
if not suffix.isdigit() or int(suffix) <= 0:
raise ValueError(f"Invalid migration version in function name: {name!r}")
return int(suffix)
def register(self, fn):
"""Decorator to register a migration function."""
version = self._migration_version(fn)
self._migrations[version] = fn
return fn
@classmethod
def from_module(cls, module: str | ModuleType) -> MigrationRegistry:
"""Create a registry by scanning a module for ``migrate_vN`` functions.
Args:
module: A module name (string) or an imported module object.
"""
reg = cls()
if isinstance(module, str):
mod = importlib.import_module(module)
else:
mod = module
for name in dir(mod):
if name.startswith("migrate_v"):
fn = getattr(mod, name)
if callable(fn):
version = reg._migration_version(fn)
reg._migrations[version] = fn
return reg
@property
def dbver(self) -> int:
"""Current schema version (= highest discovered migration, or 0)."""
return max(self._migrations.keys(), default=0)
@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."""
try:
inspect.signature(fn).bind(data_dict, kanta)
except TypeError:
fn(data_dict)
else:
fn(data_dict, kanta)
def apply(
self,
data_dict: dict[str, Any],
current_version: int,
kanta: Any,
*,
silent: bool = False,
) -> int:
"""Apply pending migrations to *data_dict* in place.
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})"
)
self._call_migration(fn, data_dict, kanta)
current_version = next_version
if not silent:
desc = (fn.__doc__ or fn.__name__).split("\n")[0].rstrip(".")
_logger.info("Applied migration %s: %s", fn.__name__, desc)
return current_version
+184
View File
@@ -0,0 +1,184 @@
"""Database schema migration framework.
Migrations are numbered functions discovered automatically via a decorator
or by prefix. Each runs exactly once based on the current version.
"""
from __future__ import annotations
import copy
import importlib
import inspect
from dataclasses import dataclass
from types import ModuleType
from typing import Any
from kanta.diff import compute_diff
from kanta.exceptions import DatabaseError
# Cache registries by imported module object so that many Kanta instances using
# the same migrations module do not re-scan it each time.
_module_registry_cache: dict[ModuleType, Migrations] = {}
@dataclass
class MigrationInfo:
"""Information about a single migration that ran."""
name: str
description: str
version: int
changed: bool
diff: dict | None = None
before: dict | None = None
@dataclass
class MigrationResult:
"""Result of applying migrations."""
version: int
migrations: list[MigrationInfo]
class Migrations:
"""Registry of schema migration functions.
Usage::
migrations = Migrations()
@migrations.register
def migrate_v1(d: dict, kanta) -> None:
d.setdefault("version", 1)
kanta.ctx.note = "migrated"
@migrations.register
def migrate_v2(d: dict) -> None:
d.setdefault("version", 2)
result = migrations.apply(state, current_version=0, kanta=kanta)
new_version = result.version
Or load from a module::
migrations = Migrations.from_module("myapp.migrations")
result = migrations.apply(state, current_version=0, kanta=kanta)
"""
def __init__(self) -> None:
self._migrations: dict[int, Any] = {}
@staticmethod
def _migration_version(fn: Any) -> int:
name = getattr(fn, "__name__", "")
if not name.startswith("migrate_v"):
raise ValueError(f"Invalid migration function name: {name!r}")
suffix = name.removeprefix("migrate_v")
if not suffix.isdigit() or int(suffix) <= 0:
raise ValueError(f"Invalid migration version in function name: {name!r}")
return int(suffix)
def register(self, fn):
"""Decorator to register a migration function."""
version = self._migration_version(fn)
self._migrations[version] = fn
return fn
@classmethod
def from_module(cls, module: str | ModuleType) -> Migrations:
"""Create or retrieve a cached registry by scanning a module.
Args:
module: A module name (string) or an imported module object.
"""
if isinstance(module, str):
mod = importlib.import_module(module)
else:
mod = module
try:
return _module_registry_cache[mod]
except KeyError:
pass
reg = cls()
for name in dir(mod):
if name.startswith("migrate_v"):
fn = getattr(mod, name)
if callable(fn):
version = reg._migration_version(fn)
reg._migrations[version] = fn
_module_registry_cache[mod] = reg
return reg
@property
def dbver(self) -> int:
"""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."""
try:
inspect.signature(fn).bind(data_dict, kanta)
except TypeError:
fn(data_dict)
else:
fn(data_dict, kanta)
def apply(
self,
data_dict: dict[str, Any],
current_version: int,
kanta: Any,
) -> MigrationResult:
"""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 a :class:`MigrationResult` describing the new version and every
migration that ran.
"""
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}"
)
migrations: list[MigrationInfo] = []
for version in sorted(self._migrations.keys()):
if version <= current_version:
continue
fn = self._migrations[version]
before = copy.deepcopy(data_dict)
self._call_migration(fn, data_dict, kanta)
current_version = version
changed = before != data_dict
diff = compute_diff(before, data_dict) if changed else None
desc = (fn.__doc__ or f"v{version}").split("\n")[0].rstrip(".")
migrations.append(
MigrationInfo(
name=fn.__name__,
description=desc,
version=version,
changed=changed,
diff=diff,
before=before,
)
)
return MigrationResult(version=current_version, migrations=migrations)
+7 -2
View File
@@ -109,6 +109,7 @@ class PersistenceMixin:
*, *,
user: str | None = None, user: str | None = None,
mtime: bool | datetime = True, mtime: bool | datetime = True,
force: bool = False,
) -> ChangeRecord | None: ) -> ChangeRecord | None:
"""Queue a change record internally (thread-safe). """Queue a change record internally (thread-safe).
@@ -121,9 +122,11 @@ class PersistenceMixin:
previous modification time remains in effect; this is used for previous modification time remains in effect; this is used for
system operations that are not considered modifications. A system operations that are not considered modifications. A
:class:`~datetime.datetime` value sets ``m`` to that explicit time. :class:`~datetime.datetime` value sets ``m`` to that explicit time.
force: If ``True``, queue the record even when the diff is empty.
Returns: Returns:
The queued :class:`ChangeRecord`, or ``None`` if the diff was empty. The queued :class:`ChangeRecord`, or ``None`` if the diff was empty
and *force* is ``False``.
""" """
now = datetime.now(UTC) now = datetime.now(UTC)
@@ -138,7 +141,9 @@ class PersistenceMixin:
diff = compute_diff(self.statedict, current) diff = compute_diff(self.statedict, current)
if not diff: if not diff:
return None if not force:
return None
diff = {}
record = ChangeRecord( record = ChangeRecord(
ts=now, ts=now,
+8 -7
View File
@@ -41,15 +41,16 @@ class SnapshotState:
self, file, version: int, state: dict, m: datetime | None = None self, file, version: int, state: dict, m: datetime | None = None
) -> None: ) -> None:
"""Write snapshot when thresholds/time policy allows it.""" """Write snapshot when thresholds/time policy allows it."""
if self.changes < self._min_diffs:
return
force = self._force_pending force = self._force_pending
now = datetime.now(UTC) now = datetime.now(UTC)
if not force and now.weekday() != 6: # 6 = Sunday if not force:
return if self.changes < self._min_diffs:
sunday_midnight = now.replace(hour=0, minute=0, second=0, microsecond=0) return
if not force and self.ts is not None and self.ts >= sunday_midnight: if now.weekday() != 6: # 6 = Sunday
return return
sunday_midnight = now.replace(hour=0, minute=0, second=0, microsecond=0)
if self.ts is not None and self.ts >= sunday_midnight:
return
if not file.is_open: if not file.is_open:
return return
try: try:
+1 -1
View File
@@ -23,7 +23,7 @@ class ChangeRecord(msgspec.Struct, omit_defaults=True, kw_only=True):
v: int = 0 v: int = 0
u: str | None = None u: str | None = None
m: datetime | None = None m: datetime | None = None
diff: dict diff: dict = {}
class Snapshot(msgspec.Struct, omit_defaults=True): class Snapshot(msgspec.Struct, omit_defaults=True):
+12 -2
View File
@@ -9,7 +9,7 @@ from datetime import datetime
from kanta.diff import compute_diff from kanta.diff import compute_diff
from kanta.exceptions import DataIntegrityError from kanta.exceptions import DataIntegrityError
from kanta.callbacks import InjectionContext from kanta.callbacks import InjectionContext
from kanta.logging import _USER_PATH, log_change from kanta.logging import _USER_PATH, changes_logger, log_change
from kanta.serialization import restore_data_in_place, struct_to_dict from kanta.serialization import restore_data_in_place, struct_to_dict
_logger = logging.getLogger(__name__) _logger = logging.getLogger(__name__)
@@ -22,6 +22,7 @@ def transaction(
*, *,
user: str | None = None, user: str | None = None,
mtime: bool | datetime = True, mtime: bool | datetime = True,
log: bool | logging.Logger = True,
): ):
"""Wrap writes in a transaction and yield the live db object.""" """Wrap writes in a transaction and yield the live db object."""
if impl.readonly: if impl.readonly:
@@ -80,7 +81,16 @@ def transaction(
resolved = logfmt(user, _USER_PATH) resolved = logfmt(user, _USER_PATH)
if resolved is not None: if resolved is not None:
formatted_user = resolved formatted_user = resolved
log_change(action, record.diff, formatted_user, previous, logfmt) if log is not False:
logger = log if isinstance(log, logging.Logger) else changes_logger
log_change(
action,
record.diff,
formatted_user,
previous,
logfmt,
logger=logger,
)
except Exception: except Exception:
_logger.warning("Transaction '%s' failed, rolling back changes", action) _logger.warning("Transaction '%s' failed, rolling back changes", action)
if impl.transaction_snapshot is not None: if impl.transaction_snapshot is not None:
+24 -1
View File
@@ -7,7 +7,7 @@ from uuid import UUID
import msgspec import msgspec
from kanta.kanta import Kanta from kanta.kanta import Kanta
from kanta.structs import ChangeRecord from kanta.structs import ChangeRecord, Snapshot
class User(msgspec.Struct): class User(msgspec.Struct):
@@ -70,6 +70,18 @@ def change_actions(path: Path, format_config) -> list[str]:
return actions return actions
def read_changes(path: Path, format_config) -> list[ChangeRecord]:
_, serializer_cls = format_config
serializer = serializer_cls()
framer = serializer.framer_cls()
records: list[ChangeRecord] = []
for is_snapshot, payload, _, _ in framer.iter_records(path.read_bytes(), 0):
if is_snapshot:
continue
records.append(serializer.decode(payload, type=ChangeRecord))
return records
def make_migrations_module(name: str, fn_name: str, fn): def make_migrations_module(name: str, fn_name: str, fn):
mod = ModuleType(name) mod = ModuleType(name)
mod.__dict__[fn_name] = fn mod.__dict__[fn_name] = fn
@@ -77,6 +89,17 @@ def make_migrations_module(name: str, fn_name: str, fn):
return mod return mod
def read_last_snapshot(path: Path, format_config) -> Snapshot | None:
_, serializer_cls = format_config
serializer = serializer_cls()
framer = serializer.framer_cls()
data = path.read_bytes()
payload, _, _ = framer.scan_last_snapshot(data)
if payload is None:
return None
return serializer.decode(payload, type=Snapshot)
def fixed_change(action: str, diff: dict, *, version: int = 0) -> ChangeRecord: def fixed_change(action: str, diff: dict, *, version: int = 0) -> ChangeRecord:
return ChangeRecord( return ChangeRecord(
ts=datetime(2026, 1, 1, tzinfo=UTC), a=action, v=version, diff=diff ts=datetime(2026, 1, 1, tzinfo=UTC), a=action, v=version, diff=diff
+259
View File
@@ -1,4 +1,5 @@
import asyncio import asyncio
import logging
import sys import sys
from datetime import UTC, datetime from datetime import UTC, datetime
from uuid import uuid4 from uuid import uuid4
@@ -6,6 +7,7 @@ from uuid import uuid4
import pytest import pytest
from kanta.exceptions import DatabaseError, DataIntegrityError, FileLockError from kanta.exceptions import DatabaseError, DataIntegrityError, FileLockError
from kanta.migrations import MigrationResult
from kanta.serialization import struct_to_dict from kanta.serialization import struct_to_dict
from .support import ( from .support import (
@@ -17,6 +19,9 @@ from .support import (
change_actions, change_actions,
fixed_change, fixed_change,
make_kanta, make_kanta,
make_migrations_module,
read_changes,
read_last_snapshot,
seed_single_change, seed_single_change,
) )
@@ -30,6 +35,63 @@ async def test_load_empty(tmp_path, format_config):
await kanta.close() await kanta.close()
@pytest.mark.asyncio
async def test_new_file_writes_bootstrap_record_without_handlers(
tmp_path, format_config
):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
await kanta.open()
await kanta.close()
records = read_changes(path, format_config)
assert len(records) == 1
assert records[0].a == "bootstrap"
assert records[0].diff == {"$replace": {"users": {}, "counter": 0}}
@pytest.mark.asyncio
async def test_new_file_persists_initial_state_for_roundtrip(tmp_path, format_config):
path = tmp_path / "test.db"
kanta = make_kanta(
path, Data(counter=5, users={"alice": User(name="Alice")}), format_config
)
await kanta.open()
await kanta.close()
records = read_changes(path, format_config)
assert len(records) == 1
assert records[0].a == "bootstrap"
assert records[0].diff == {
"$replace": {"users": {"alice": {"name": "Alice", "age": 0}}, "counter": 5}
}
kanta2 = make_kanta(path, Data, format_config)
await kanta2.open()
assert kanta2.data.counter == 5
assert kanta2.data.users["alice"].name == "Alice"
await kanta2.close()
@pytest.mark.asyncio
async def test_reopen_without_changes_does_not_force_snapshot(tmp_path, format_config):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data(counter=5), format_config)
await kanta.open()
await kanta.close()
# No snapshot should exist after the initial bootstrap and close.
assert read_last_snapshot(path, format_config) is None
kanta2 = make_kanta(path, Data, format_config)
await kanta2.open()
assert kanta2.data.counter == 5
await kanta2.close()
# Re-opening without migrations or normalization changes must not force one.
assert read_last_snapshot(path, format_config) is None
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_open_overwrites_caller_owned_root_data(tmp_path, format_config): async def test_open_overwrites_caller_owned_root_data(tmp_path, format_config):
path = tmp_path / "test.db" path = tmp_path / "test.db"
@@ -434,6 +496,203 @@ async def test_msgspec_normalization_logs_migration(tmp_path, format_config):
assert "migrate:msgspec" in change_actions(path, format_config) assert "migrate:msgspec" in change_actions(path, format_config)
@pytest.mark.asyncio
async def test_empty_migration_writes_snapshot_and_is_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.close()
# Empty migrations must not produce empty change records.
records = read_changes(path, format_config)
migration_records = [r for r in records if r.a.startswith("migrate")]
assert not migration_records
# The version bump is persisted via a snapshot instead.
snap = read_last_snapshot(path, format_config)
assert snap is not None
assert snap.v == 1
assert snap.state == {"counter": 0, "users": {}}
kanta2 = make_kanta(path, Data, format_config, migrations=mod)
await kanta2.open()
assert kanta2.version == 1
await kanta2.close()
# Re-opening must not create additional migration records or snapshots.
records2 = read_changes(path, format_config)
assert not [r for r in records2 if r.a.startswith("migrate")]
finally:
sys.modules.pop("empty_migration_mod", None)
@pytest.mark.asyncio
async def test_migration_with_changes_records_diff_and_snapshot(
tmp_path, format_config
):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("init", {"counter": 0}), format_config)
mod = type(sys)("test_migrations_changes")
def migrate_v1(d, kanta):
d["counter"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
kanta = make_kanta(path, Data, format_config, migrations=mod)
await kanta.open()
assert kanta.version == 1
assert kanta.data.counter == 2
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) == 2
assert migration_records[0].a == "migrate:v1"
assert migration_records[0].v == 1
assert migration_records[0].diff == {"counter": 2}
assert migration_records[1].a == "migrate:msgspec"
assert migration_records[1].v == 1
assert migration_records[1].diff == {"users": {}}
snap = read_last_snapshot(path, format_config)
assert snap is not None
assert snap.v == 1
assert snap.state == {"counter": 2, "users": {}}
@pytest.mark.asyncio
async def test_migration_summary_log_includes_filename(tmp_path, format_config, caplog):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("init", {"counter": 0}), format_config)
mod = type(sys)("test_migrations_log")
def migrate_v1(d, kanta):
"""Bump counter."""
d["counter"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
with caplog.at_level(logging.INFO, logger="kanta.migrations"):
kanta = make_kanta(path, Data, format_config, migrations=mod)
await kanta.open()
assert kanta.version == 1
await kanta.close()
info_messages = [r.message for r in caplog.records if r.levelno == logging.INFO]
assert len(info_messages) == 1
assert str(path) in info_messages[0]
assert "v0 -> v1" in info_messages[0]
assert "migrate_v1 (Bump counter)" in info_messages[0]
@pytest.mark.asyncio
async def test_open_log_false_suppresses_migration_log(tmp_path, format_config, caplog):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("init", {"counter": 0}), format_config)
mod = type(sys)("test_migrations_silent")
def migrate_v1(d, kanta):
d["counter"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
with caplog.at_level(logging.INFO, logger="kanta.migrations"):
kanta = make_kanta(path, Data, format_config, migrations=mod)
await kanta.open(log=False)
await kanta.close()
info_messages = [r for r in caplog.records if r.levelno == logging.INFO]
assert not info_messages
@pytest.mark.asyncio
async def test_logmigr_callback_replaces_default_logging(
tmp_path, format_config, caplog
):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("init", {"counter": 0}), format_config)
mod = type(sys)("test_migrations_callback")
def migrate_v1(d, kanta):
"""Bump counter."""
d["counter"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
summaries = []
kanta = make_kanta(path, Data, format_config, migrations=mod)
@kanta.logmigr
def collect(summary: MigrationResult):
summaries.append(summary)
with caplog.at_level(logging.INFO, logger="kanta.migrations"):
await kanta.open()
await kanta.close()
assert len(summaries) == 1
assert summaries[0].version == 1
assert summaries[0].migrations[0].name == "migrate_v1"
info_messages = [r for r in caplog.records if r.levelno == logging.INFO]
assert not info_messages
@pytest.mark.asyncio
async def test_transaction_log_false_suppresses_log(tmp_path, format_config, caplog):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
await kanta.open()
with caplog.at_level(logging.INFO, logger="kanta.changes"):
with kanta.transaction(action="inc", log=False) as data:
data.counter = 1
await kanta.close()
info_messages = [r for r in caplog.records if r.levelno == logging.INFO]
assert not info_messages
@pytest.mark.asyncio
async def test_transaction_log_custom_logger(tmp_path, format_config, caplog):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
await kanta.open()
custom_logger = logging.getLogger("custom.transaction")
custom_logger.setLevel(logging.INFO)
with caplog.at_level(logging.INFO, logger="custom.transaction"):
with kanta.transaction(action="inc", log=custom_logger) as data:
data.counter = 1
await kanta.close()
info_messages = [r for r in caplog.records if r.levelno == logging.INFO]
assert len(info_messages) >= 1
assert "inc" in info_messages[0].message
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_open_locked_file_raises_filelock_error(tmp_path, format_config): async def test_open_locked_file_raises_filelock_error(tmp_path, format_config):
path = tmp_path / "test.db" path = tmp_path / "test.db"
+3 -4
View File
@@ -1,16 +1,15 @@
import logging import logging
from kanta.logging import configure_logging, log_change from kanta.logging import changes_logger, configure_logging, log_change
from kanta.logging import logger
def test_configure_logging(): def test_configure_logging():
configure_logging() configure_logging()
assert logger.level == logging.INFO assert changes_logger.level == logging.INFO
def test_log_change_no_diff(capsys): def test_log_change_no_diff(capsys):
logger.handlers.clear() changes_logger.handlers.clear()
configure_logging() configure_logging()
log_change("test", {}) log_change("test", {})
captured = capsys.readouterr() captured = capsys.readouterr()
+133 -16
View File
@@ -1,6 +1,9 @@
from types import ModuleType, SimpleNamespace from types import ModuleType, SimpleNamespace
from kanta.migrate import MigrationRegistry import pytest
from kanta.exceptions import DatabaseError
from kanta.migrations import Migrations
class _DummyKanta: class _DummyKanta:
@@ -9,7 +12,7 @@ class _DummyKanta:
def test_register_and_apply(): def test_register_and_apply():
reg = MigrationRegistry() reg = Migrations()
kanta = _DummyKanta() kanta = _DummyKanta()
@reg.register @reg.register
@@ -21,13 +24,13 @@ def test_register_and_apply():
d["version"] = 2 d["version"] = 2
state = {} state = {}
new_ver = reg.apply(state, current_version=0, kanta=kanta, silent=True) result = reg.apply(state, current_version=0, kanta=kanta)
assert new_ver == 2 assert result.version == 2
assert state["version"] == 2 assert state["version"] == 2
def test_no_migrations_needed(): def test_no_migrations_needed():
reg = MigrationRegistry() reg = Migrations()
kanta = _DummyKanta() kanta = _DummyKanta()
@reg.register @reg.register
@@ -35,8 +38,8 @@ def test_no_migrations_needed():
d["x"] = 1 d["x"] = 1
state = {"x": 1} state = {"x": 1}
new_ver = reg.apply(state, current_version=1, kanta=kanta, silent=True) result = reg.apply(state, current_version=1, kanta=kanta)
assert new_ver == 1 assert result.version == 1
def test_from_module(): def test_from_module():
@@ -52,17 +55,17 @@ def test_from_module():
mod.__dict__["migrate_v1"] = migrate_v1 mod.__dict__["migrate_v1"] = migrate_v1
mod.__dict__["migrate_v2"] = migrate_v2 mod.__dict__["migrate_v2"] = migrate_v2
reg = MigrationRegistry.from_module(mod) reg = Migrations.from_module(mod)
assert reg.dbver == 2 assert reg.dbver == 2
state = {} state = {}
new_ver = reg.apply(state, current_version=0, kanta=kanta, silent=True) result = reg.apply(state, current_version=0, kanta=kanta)
assert new_ver == 2 assert result.version == 2
assert state["v"] == 2 assert state["v"] == 2
def test_migrations_can_use_kanta_ctx(): def test_migrations_can_use_kanta_ctx():
reg = MigrationRegistry() reg = Migrations()
kanta = _DummyKanta() kanta = _DummyKanta()
@reg.register @reg.register
@@ -71,14 +74,14 @@ def test_migrations_can_use_kanta_ctx():
d["source"] = kanta.ctx.source d["source"] = kanta.ctx.source
state = {} state = {}
new_ver = reg.apply(state, current_version=0, kanta=kanta, silent=True) result = reg.apply(state, current_version=0, kanta=kanta)
assert new_ver == 1 assert result.version == 1
assert state["source"] == "migration" assert state["source"] == "migration"
assert kanta.ctx.source == "migration" assert kanta.ctx.source == "migration"
def test_migration_can_omit_kanta_argument(): def test_migration_can_omit_kanta_argument():
reg = MigrationRegistry() reg = Migrations()
kanta = _DummyKanta() kanta = _DummyKanta()
@reg.register @reg.register
@@ -86,6 +89,120 @@ def test_migration_can_omit_kanta_argument():
d["x"] = 1 d["x"] = 1
state = {} state = {}
new_ver = reg.apply(state, current_version=0, kanta=kanta, silent=True) result = reg.apply(state, current_version=0, kanta=kanta)
assert new_ver == 1 assert result.version == 1
assert state["x"] == 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)
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)
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}
result = reg.apply(state, current_version=1, kanta=kanta)
assert result.version == 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}
result = reg.apply(state, current_version=2, kanta=kanta)
assert result.version == 3
assert state["x"] == 3
def test_apply_returns_change_information():
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
result = reg.apply({}, current_version=0, kanta=kanta)
assert result.version == 3
assert len(result.migrations) == 3
assert result.migrations[0].name == "migrate_v1"
assert result.migrations[0].description == "Set x"
assert result.migrations[0].changed is True
assert result.migrations[0].diff == {"$replace": {"x": 1}}
assert result.migrations[1].name == "migrate_v2"
assert result.migrations[1].description == "No-op"
assert result.migrations[1].changed is False
assert result.migrations[1].diff is None
assert result.migrations[2].name == "migrate_v3"
assert result.migrations[2].description == "Set y"
assert result.migrations[2].changed is True
assert result.migrations[2].diff == {"y": 3}
def test_description_defaults_to_version_when_no_docstring():
reg = Migrations()
kanta = _DummyKanta()
@reg.register
def migrate_v1(d):
d["x"] = 1
result = reg.apply({}, current_version=0, kanta=kanta)
assert result.migrations[0].description == "v1"
+3 -2
View File
@@ -82,8 +82,9 @@ async def test_transaction_mtime_false_preserves_mtime(tmp_path, format_config):
continue continue
records.append(serializer.decode(payload, type=ChangeRecord)) records.append(serializer.decode(payload, type=ChangeRecord))
assert records[0].m == first_m assert records[0].a == "bootstrap"
assert records[1].m is None assert records[1].m == first_m
assert records[2].m is None
assert kanta.mtime == first_m assert kanta.mtime == first_m
+17
View File
@@ -32,3 +32,20 @@ def test_force_writes():
f = FakeFile() f = FakeFile()
ss.maybe_write(f, 1, {"x": 1}) ss.maybe_write(f, 1, {"x": 1})
assert len(f.written) == 1 assert len(f.written) == 1
def test_force_bypasses_min_diffs():
class FakeFile:
def __init__(self):
self.written = []
self.is_open = True
def write(self, data: bytes):
self.written.append(data)
ss = SnapshotState(min_diffs=100)
ss.record_changes(5)
ss.request_force()
f = FakeFile()
ss.maybe_write(f, 1, {"x": 1})
assert len(f.written) == 1