9 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
LeoVasanko bec4635460 Add kanta.ctx SimpleNamespace for user variables. This is passed to migration functions if they take a second argument. Remove the old MigrationCtx system. 2026-06-14 19:54:23 +00:00
LeoVasanko c04a245366 Add read only mode that doesn't use locking. Useful for inspection while the database is in use or when no changes are intended (except in RAM). 2026-06-14 19:35:06 +00:00
LeoVasanko b501fcca86 Only Kanta exported at root, others available from submodules. 2026-06-13 21:26:35 +00:00
LeoVasanko 7e553bd868 Improved log formatting support by @kanta.logfmt, which replaces old resolver and user_display arguments (breaking change). 2026-06-13 21:22:05 +00:00
26 changed files with 2293 additions and 376 deletions
+21 -5
View File
@@ -53,10 +53,13 @@ asyncio.run(main())
## Bootstrap and Open Modes
Kanta supports open-time bootstrap callbacks for initializing a brand-new
database before `open()` returns.
When `open()` creates a brand-new database, it always writes a single bootstrap
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
kanta = Kanta("data.kantadb", Data())
@@ -76,9 +79,10 @@ async def bootstrap_async(data) -> None:
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,
- exactly one bootstrap change record is queued,
- bootstrap metadata (`action`, `user`, `mtime`) is taken from the last
registration.
@@ -94,6 +98,18 @@ await kanta.open(create=False)
With `create=False`, open fails if the database file does not exist or is
empty.
Read-only mode opens an existing database without locking it or starting the
background flush task. This is useful for readers that must not block the
writer or modify the file:
```python
await kanta.open(readonly=True)
```
In read-only mode, records are replayed and migrations are applied in memory,
but transactions and explicit flushes are rejected and the file is never
created if missing.
## Fatal Error Handlers
Fatal background write errors can be observed with a decorator:
+63 -7
View File
@@ -118,28 +118,84 @@ reloads, while system operations such as migrations leave it unchanged.
- `await kanta.open()` (default) creates the database file if missing.
- `await kanta.open(create=False)` fails when the file is missing or empty.
- `await kanta.open(readonly=True)` opens an existing database read-only.
- The file is opened without acquiring a lock and without a background flush
task.
- Existing records are replayed and migrations are still applied in memory.
- Transactions and explicit flushes are rejected.
- The file is never created if missing.
### Bootstrap Callbacks
### Callbacks
- Bootstrap callbacks run during `open()` when the database is empty.
All callbacks are registered via decorators and receive arguments by their
annotation types. Parameters without a supported annotation are only allowed
when they have a default value.
#### Bootstrap Callbacks
- 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:
- `@kanta.bootstrap`
- `@kanta.bootstrap(action=..., user=..., mtime=...)`
- Bootstrap callbacks may be sync or async and receive the live root data
object.
- Bootstrap callbacks may be sync or async. The live root data object is
injected by annotating a parameter with the struct type passed to `Kanta`,
and the `Kanta` instance itself can be injected by annotating a parameter
with `Kanta`.
- Multiple bootstrap callbacks are supported:
- callbacks execute in registration order,
- exactly one bootstrap `ChangeRecord` is queued,
- bootstrap metadata (`action`, `user`, `mtime`) is taken from the last
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,
then re-raises the exception.
### Fatal Error Handlers
#### Fatal Error Handlers
- Fatal background persistence errors can be handled with `@kanta.fatal_error`.
- Handlers may be sync or async.
- Multiple handlers are supported and invoked in registration order.
- Handlers may be sync or async. The `DatabaseError` is injected by annotating
a parameter with `DatabaseError`; `Kanta` may also be injected.
- Multiple handlers are supported and invoked in registration order. A failing
handler is logged and does not prevent subsequent handlers from running.
#### Transaction Log Formatting
- Logfmt callbacks prettify identifiers in the change log and are registered with
`@kanta.logfmt`.
- A logfmt callback is called for every value Kanta renders: diff values, path
components, and the transaction `user`. It receives the value as its first
parameter and optionally a `path: str` parameter with the dot-notation path
to the value. The special path `"$user"` is used when rendering the
transaction actor, replacing the old `user_display` parameter.
- The callback returns `str | None`: a string replaces the default rendering,
while `None` means "fall through to the next formatter".
- State dicts can be injected via `DictPre` (`Annotated[dict, "pre"]`)
and `DictPost` (`Annotated[dict, "post"]`); the `Kanta` instance can also be
injected.
- Alternatively, a logfmt callback can be a class inheriting from `LogFmt`; the
framework instantiates it with the state dicts and calls its
`resolve(value, path) -> str | None` method.
- Multiple logfmt callbacks are stacked in registration order; the first
callback to return a non-`None` result wins. If none handle a value, Kanta
falls back to its default formatting.
The decorator accepts an optional ``path`` so the callback only runs for
values at that exact path:
```python
@kanta.logfmt(path="$user")
def resolve_user(value: str, current: DictPost) -> str | None:
return current.get("users", {}).get(value, {}).get("name")
@kanta.logfmt(path="users.uuid-1")
def resolve_user_key(value: str) -> str | None:
return names_by_id.get(value)
```
## Migrations
-21
View File
@@ -1,26 +1,5 @@
from .diff import compute_diff
from .diff import replay_jsonl as replay
from .exceptions import DatabaseError, DataIntegrityError, FileLockError, ReplayError
from .filelock import LockedFile
from .kanta import Kanta
from .logging import configure_logging, format_diff, log_change
from .serialization import JsonSerializer, MsgPackSerializer
from .structs import ChangeRecord, Snapshot
__all__ = [
"ChangeRecord",
"compute_diff",
"configure_logging",
"DataIntegrityError",
"DatabaseError",
"FileLockError",
"format_diff",
"JsonSerializer",
"Kanta",
"LockedFile",
"log_change",
"MsgPackSerializer",
"ReplayError",
"replay",
"Snapshot",
]
+539
View File
@@ -0,0 +1,539 @@
"""Unified decorator-based callback registry for Kanta.
Callbacks are registered once and invoked with arguments filled by their
annotation types. Unknown arguments are only permitted when they have a
default value.
Log formatters are a special case: they are called per value being rendered
and receive the value plus an optional ``path`` string. They return
``str | None``; ``None`` means "fall through to the next formatter".
"""
from __future__ import annotations
import inspect
import types
from collections.abc import Callable
from dataclasses import dataclass
from typing import Annotated, Any, Union, get_args, get_origin
from kanta.exceptions import DatabaseError
from kanta.migrations import MigrationResult
DictPre = Annotated[dict, "pre"]
DictPost = Annotated[dict, "post"]
class LogFmt:
"""Base class for stateful logfmt callbacks.
Subclasses only need to override :meth:`resolve`. The framework injects
``previous_state`` and ``current_state`` through ``__init__``.
"""
def __init__(
self,
previous: DictPre | None = None,
current: DictPost | None = None,
) -> None:
self.previous_state = previous
self.current_state = current
def __call__(self, value: Any, path: str) -> str | None:
return self.resolve(value, path)
def resolve(self, value: Any, path: str) -> str | None:
"""Resolve *value* into a display string.
The default implementation returns ``None`` so other formatters are
tried.
"""
return None
@dataclass
class InjectionContext:
"""Runtime values available for injection into callbacks."""
kanta: Any | None = None
data: Any | None = None
error: DatabaseError | None = None
previous_state: dict | None = None
current_state: dict | None = None
migration_result: MigrationResult | None = None
@dataclass
class _CallbackRegistration:
callback: Callable[..., Any]
params: list[tuple[str, type]]
is_async: bool = False
@dataclass
class _LogFmtFunctionSpec:
callback: Callable[..., Any]
value_type: type | Any
has_path: bool
inject_params: list[tuple[str, type]]
path: str | None = None
@dataclass
class _LogFmtClassSpec:
cls: type[LogFmt]
inject_params: list[tuple[str, type]]
path: str | None = None
class CallbackRegistry:
"""Stores and invokes callbacks, resolving arguments by annotation."""
def __init__(
self,
*,
kanta_class: type | None = None,
data_type: type | None = None,
) -> None:
self._kanta_class = kanta_class
self._data_type = data_type
self._callbacks: dict[str, list[_CallbackRegistration]] = {
"bootstrap": [],
"fatal_error": [],
"logmigr": [],
}
self._logfmt_callbacks: list[_LogFmtFunctionSpec | _LogFmtClassSpec] = []
def register(
self,
kind: str,
callback: Callable[..., Any],
*,
path: str | None = None,
) -> Callable[..., Any]:
"""Register *callback* for *kind* after validating its signature."""
if kind == "logfmt":
if inspect.isclass(callback):
self._logfmt_callbacks.append(
self._validate_logfmt_class(callback, path=path)
)
else:
self._logfmt_callbacks.append(
self._validate_logfmt_function(callback, path=path)
)
return callback
if kind not in self._callbacks:
raise ValueError(f"unknown callback kind: {kind}")
if inspect.isclass(callback):
raise TypeError(f"{kind} callbacks must be functions, not classes")
if not callable(callback):
raise TypeError(f"{kind} callback must be callable")
params = self._validate_function(callback, kind)
is_async = inspect.iscoroutinefunction(callback)
self._callbacks[kind].append(
_CallbackRegistration(
callback=callback,
params=params,
is_async=is_async,
)
)
return callback
async def invoke(
self,
kind: str,
ctx: InjectionContext,
*,
on_error: Callable[[Exception, Callable[..., Any]], bool | None] | None = None,
) -> list[Any]:
"""Invoke all callbacks of *kind* with arguments from *ctx*.
If *on_error* is provided it is called for each exception and may return
``False`` to stop invoking further callbacks. When *on_error* is not
provided the first exception is raised immediately.
"""
results: list[Any] = []
for reg in self._callbacks[kind]:
try:
kwargs = self._build_kwargs(reg.params, ctx)
result = reg.callback(**kwargs)
if inspect.isawaitable(result):
result = await result
results.append(result)
except Exception as exc:
if on_error is None:
raise
if on_error(exc, reg.callback) is False:
break
return results
def has(self, kind: str) -> bool:
"""Return True if any callback of *kind* is registered."""
if kind == "logfmt":
return bool(self._logfmt_callbacks)
return bool(self._callbacks[kind])
def build_logfmt(self, ctx: InjectionContext) -> Callable[[Any, str], str | None]:
"""Build a chained formatter from registered logfmt callbacks."""
formatters: list[tuple[Callable[[Any, str], str | None], str | None]] = []
for spec in self._logfmt_callbacks:
if isinstance(spec, _LogFmtClassSpec):
kwargs = self._build_kwargs(spec.inject_params, ctx)
instance: Callable[[Any, str], str | None] = spec.cls(**kwargs)
formatters.append((instance, spec.path))
else:
kwargs = self._build_kwargs(spec.inject_params, ctx)
def make_formatter(
callback: Callable[..., Any] = spec.callback,
value_type: type | Any = spec.value_type,
has_path: bool = spec.has_path,
state_kwargs: dict[str, Any] = kwargs,
) -> Callable[[Any, str], str | None]:
def formatter(value: Any, path: str) -> str | None:
if value_type is str and not isinstance(value, str):
return None
call_kwargs = dict(state_kwargs)
if has_path:
call_kwargs["path"] = path
return callback(value, **call_kwargs)
return formatter
formatters.append((make_formatter(), spec.path))
def format_value(value: Any, path: str) -> str | None:
for fn, pattern in formatters:
if pattern is not None and path != pattern:
continue
resolved = fn(value, path)
if resolved is not None:
return resolved
return None
return format_value
def _validate_function(
self,
callback: Callable[..., Any],
kind: str,
) -> list[tuple[str, type]]:
sig = inspect.signature(callback)
params: list[tuple[str, type]] = []
for name, param in sig.parameters.items():
if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD):
raise TypeError(
f"{kind} callback {callback.__name__} must not use "
f"*args or **kwargs"
)
if param.annotation is inspect.Parameter.empty:
if param.default is inspect.Parameter.empty:
raise TypeError(
f"{kind} callback {callback.__name__} has parameter "
f"'{name}' without an annotation or default value"
)
continue
ann = self._resolve_raw_annotation(param.annotation, callback)
if not self._is_allowed(kind, ann):
if param.default is inspect.Parameter.empty:
raise TypeError(
f"{kind} callback {callback.__name__} has parameter "
f"'{name}' with unsupported annotation {ann!r}. "
f"Allowed: {self._allowed_message(kind)}"
)
continue
params.append((name, ann))
return params
def _validate_logfmt_function(
self,
callback: Callable[..., Any],
*,
path: str | None = None,
) -> _LogFmtFunctionSpec:
sig = inspect.signature(callback)
if inspect.iscoroutinefunction(callback):
raise TypeError("logfmt callbacks must not be async")
params = list(sig.parameters.items())
if not params:
raise TypeError(
f"logfmt callback {callback.__name__} must accept a value parameter"
)
value_name, value_param = params[0]
if value_param.kind in (value_param.VAR_POSITIONAL, value_param.VAR_KEYWORD):
raise TypeError(
f"logfmt callback {callback.__name__} must not use *args or **kwargs"
)
if value_param.annotation is inspect.Parameter.empty:
raise TypeError(
f"logfmt callback {callback.__name__} value parameter "
f"'{value_name}' must be annotated as str or Any"
)
value_ann = self._resolve_raw_annotation(value_param.annotation, callback)
value_bare = self._unwrap_optional(value_ann)
if value_bare is str:
value_type = str
elif value_bare is Any:
value_type = Any
else:
raise TypeError(
f"logfmt callback {callback.__name__} value parameter "
f"'{value_name}' must be annotated as str or Any, got {value_ann!r}"
)
has_path = False
inject_params: list[tuple[str, type]] = []
for name, param in params[1:]:
if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD):
raise TypeError(
f"logfmt callback {callback.__name__} must not use "
f"*args or **kwargs"
)
if param.annotation is inspect.Parameter.empty:
if param.default is inspect.Parameter.empty:
raise TypeError(
f"logfmt callback {callback.__name__} has parameter "
f"'{name}' without an annotation or default value"
)
continue
ann = self._resolve_raw_annotation(param.annotation, callback)
if name == "path" and self._unwrap_optional(ann) is str:
has_path = True
continue
if self._is_allowed("logfmt", ann):
inject_params.append((name, ann))
continue
if param.default is inspect.Parameter.empty:
raise TypeError(
f"logfmt callback {callback.__name__} has parameter "
f"'{name}' with unsupported annotation {ann!r}. "
f"Allowed: str path, {self._allowed_message('logfmt')}"
)
if sig.return_annotation is not inspect.Signature.empty:
return_ann = self._resolve_raw_annotation(sig.return_annotation, callback)
if not self._is_optional_str(return_ann):
raise TypeError(
f"logfmt callback {callback.__name__} must return str | None, "
f"got {return_ann!r}"
)
return _LogFmtFunctionSpec(
callback=callback,
value_type=value_type,
has_path=has_path,
inject_params=inject_params,
path=path,
)
def _validate_logfmt_class(
self,
cls: type[LogFmt],
*,
path: str | None = None,
) -> _LogFmtClassSpec:
if not issubclass(cls, LogFmt):
raise TypeError("logfmt classes must inherit from kanta.callbacks.LogFmt")
if inspect.iscoroutinefunction(cls.__init__):
raise TypeError("logfmt class __init__ must not be async")
sig = inspect.signature(cls.__init__)
inject_params: list[tuple[str, type]] = []
first = True
for name, param in sig.parameters.items():
if first and name == "self":
first = False
continue
first = False
if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD):
raise TypeError(
f"logfmt class {cls.__name__}.__init__ must not use "
f"*args or **kwargs"
)
if param.annotation is inspect.Parameter.empty:
if param.default is inspect.Parameter.empty:
raise TypeError(
f"logfmt class {cls.__name__}.__init__ has parameter "
f"'{name}' without an annotation or default value"
)
continue
ann = self._resolve_raw_annotation(param.annotation, cls.__init__)
if self._is_allowed("logfmt", ann):
inject_params.append((name, ann))
continue
if param.default is inspect.Parameter.empty:
raise TypeError(
f"logfmt class {cls.__name__}.__init__ has parameter "
f"'{name}' with unsupported annotation {ann!r}. "
f"Allowed: {self._allowed_message('logfmt')}"
)
resolve = getattr(cls, "resolve", None)
if resolve is None:
raise TypeError(f"logfmt class {cls.__name__} must define a resolve method")
resolve_sig = inspect.signature(resolve)
resolve_params = list(resolve_sig.parameters.items())
if not resolve_params or resolve_params[0][0] != "self":
raise TypeError(
f"logfmt class {cls.__name__}.resolve must have 'self' as first parameter"
)
if len(resolve_params) < 2:
raise TypeError(
f"logfmt class {cls.__name__}.resolve must accept a value parameter"
)
value_name, value_param = resolve_params[1]
value_ann = self._resolve_raw_annotation(value_param.annotation, resolve)
value_bare = self._unwrap_optional(value_ann)
if value_bare not in (inspect.Parameter.empty, str, Any):
raise TypeError(
f"logfmt class {cls.__name__}.resolve value parameter "
f"'{value_name}' must be annotated as str or Any, got {value_ann!r}"
)
path_found = False
for name, param in resolve_params[2:]:
path_ann = self._resolve_raw_annotation(param.annotation, resolve)
path_bare = self._unwrap_optional(path_ann)
if name == "path" and path_bare in (inspect.Parameter.empty, str):
path_found = True
break
if not path_found:
raise TypeError(
f"logfmt class {cls.__name__}.resolve must accept a 'path: str' parameter"
)
if resolve_sig.return_annotation is not inspect.Signature.empty:
return_ann = self._resolve_raw_annotation(
resolve_sig.return_annotation, resolve
)
if not self._is_optional_str(return_ann):
raise TypeError(
f"logfmt class {cls.__name__}.resolve must return str | None, "
f"got {return_ann!r}"
)
return _LogFmtClassSpec(cls=cls, inject_params=inject_params, path=path)
def _build_kwargs(
self,
params: list[tuple[str, type]],
ctx: InjectionContext,
) -> dict[str, Any]:
kwargs: dict[str, Any] = {}
for name, ann in params:
value = self._resolve_annotation(ann, ctx)
if value is _UNRESOLVED:
raise RuntimeError(f"no value available for annotation {ann!r}")
kwargs[name] = value
return kwargs
def _is_allowed(self, kind: str, ann: Any) -> bool:
bare = self._unwrap_optional(ann)
if self._matches_state_annotation(bare, "pre"):
return kind == "logfmt"
if self._matches_state_annotation(bare, "post"):
return kind == "logfmt"
if bare is DatabaseError:
return kind == "fatal_error"
if bare is MigrationResult:
return kind == "logmigr"
if self._data_type is not None and bare is self._data_type:
return kind == "bootstrap"
if self._kanta_class is not None and bare is self._kanta_class:
return kind in {"bootstrap", "fatal_error", "logfmt", "logmigr"}
return False
def _allowed_message(self, kind: str) -> str:
parts: list[str] = []
if kind == "bootstrap":
if self._data_type is not None:
parts.append(self._data_type.__name__)
if kind in {"bootstrap", "fatal_error", "logfmt"}:
if self._kanta_class is not None:
parts.append(self._kanta_class.__name__)
if kind == "fatal_error":
parts.append("DatabaseError")
if kind == "logmigr":
parts.append("MigrationResult")
if kind == "logfmt":
parts.append("Annotated[dict, 'pre']")
parts.append("Annotated[dict, 'post']")
return ", ".join(parts) if parts else "none"
def _resolve_annotation(self, ann: Any, ctx: InjectionContext) -> Any:
bare = self._unwrap_optional(ann)
if self._matches_state_annotation(bare, "pre"):
return ctx.previous_state
if self._matches_state_annotation(bare, "post"):
return ctx.current_state
if bare is DatabaseError:
return ctx.error
if bare is MigrationResult:
return ctx.migration_result
if self._data_type is not None and bare is self._data_type:
return ctx.data
if self._kanta_class is not None and bare is self._kanta_class:
return ctx.kanta
return _UNRESOLVED
def _resolve_raw_annotation(
self,
raw_ann: Any,
callback: Callable[..., Any],
) -> Any:
if isinstance(raw_ann, str):
try:
return eval(raw_ann, callback.__globals__)
except Exception as exc:
raise TypeError(
f"could not resolve annotation {raw_ann!r} for "
f"{callback.__name__}: {exc}"
) from exc
return raw_ann
@staticmethod
def _matches_state_annotation(ann: Any, marker: str) -> bool:
origin = get_origin(ann)
if origin is not Annotated:
return False
args = get_args(ann)
if not args:
return False
return args[0] is dict and marker in args[1:]
@staticmethod
def _unwrap_optional(ann: Any) -> Any:
origin = get_origin(ann)
if origin not in (Union, types.UnionType):
return ann
args = [arg for arg in get_args(ann) if arg is not type(None)]
return args[0] if len(args) == 1 else ann
@staticmethod
def _is_optional_str(ann: Any) -> bool:
origin = get_origin(ann)
if origin not in (Union, types.UnionType):
return ann is str
args = get_args(ann)
return type(None) in args and any(arg is str for arg in args)
class _Unresolved:
pass
_UNRESOLVED = _Unresolved()
+28 -12
View File
@@ -34,6 +34,7 @@ if sys.platform == "win32":
_GENERIC_READ = 0x80000000
_GENERIC_WRITE = 0x40000000
_FILE_SHARE_READ = 0x00000001
_FILE_SHARE_WRITE = 0x00000002
_OPEN_EXISTING = 3
_OPEN_ALWAYS = 4
_FILE_ATTRIBUTE_NORMAL = 0x80
@@ -91,12 +92,13 @@ else:
class LockedFile:
"""A file opened with an exclusive write lock.
"""A file opened for read+write with an optional exclusive lock.
Usage::
f = LockedFile()
f.open(path) # open + lock (read+write)
f.open(path, readonly=True) # open read-only without locking
content = f.read() # read entire content
f.write(data) # append data (seeks to end first)
f.close() # release lock + close fd
@@ -108,12 +110,13 @@ class LockedFile:
def __init__(self) -> None:
self._fd: int | None = None # Unix fd or Windows HANDLE
def open(self, path: Path, *, create: bool = False) -> None:
"""Open *path* for read+write with an exclusive lock.
def open(self, path: Path, *, create: bool = False, readonly: bool = False) -> None:
"""Open *path* and optionally acquire an exclusive lock.
Args:
path: File to open and lock.
create: If True, create the file if it doesn't exist (bootstrap).
readonly: If True, open read-only without acquiring a lock.
Raises:
FileLockError: If the file is locked by another process or not found.
@@ -122,16 +125,18 @@ class LockedFile:
return # Already open (idempotent)
if sys.platform == "win32":
self._open_win32(path, create)
self._open_win32(path, create, readonly)
else:
self._open_unix(path, create)
self._open_unix(path, create, readonly)
def open_and_read(self, path: Path, create: bool = False) -> bytes:
"""Open *path* with exclusive lock and read all content.
def open_and_read(
self, path: Path, create: bool = False, readonly: bool = False
) -> bytes:
"""Open *path* and read all content.
Combined operation for efficient use with asyncio.to_thread().
"""
self.open(path, create=create)
self.open(path, create=create, readonly=readonly)
return self.read()
def read(self) -> bytes:
@@ -188,12 +193,16 @@ class LockedFile:
# -- Unix ----------------------------------------------------------------
def _open_unix(self, path: Path, create: bool) -> None:
def _open_unix(self, path: Path, create: bool, readonly: bool) -> None:
if readonly:
flags = os.O_RDONLY
else:
flags = os.O_RDWR | (os.O_CREAT if create else 0)
try:
fd = os.open(path, flags, 0o666)
except FileNotFoundError:
_fatal(f"Database file not found: {path.resolve()}", db_path=path)
if not readonly:
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except OSError:
@@ -220,12 +229,19 @@ class LockedFile:
# -- Windows -------------------------------------------------------------
def _open_win32(self, path: Path, create: bool) -> None:
def _open_win32(self, path: Path, create: bool, readonly: bool) -> None:
if readonly:
disposition = _OPEN_EXISTING
access = _GENERIC_READ
share = _FILE_SHARE_READ | _FILE_SHARE_WRITE
else:
disposition = _OPEN_ALWAYS if create else _OPEN_EXISTING
access = _GENERIC_READ | _GENERIC_WRITE
share = _FILE_SHARE_READ
handle = _kernel32.CreateFileW(
str(path),
_GENERIC_READ | _GENERIC_WRITE,
_FILE_SHARE_READ,
access,
share,
None,
disposition,
_FILE_ATTRIBUTE_NORMAL,
+78 -20
View File
@@ -1,12 +1,12 @@
"""JSONL persistence layer with background flush task."""
"""Kanta DB main public API"""
from __future__ import annotations
import logging
from datetime import datetime
from pathlib import Path
from types import ModuleType
from typing import Any, Generic, TypeVar
from types import ModuleType, SimpleNamespace
from typing import Generic, TypeVar
from kanta.exceptions import DatabaseError
from kanta.kantaimpl import KantaImpl
from kanta.serialization import JsonSerializer, Serializer
from kanta.transaction import transaction as _transaction
@@ -51,7 +51,6 @@ class Kanta(Generic[T]):
*,
type: type[T] | None = None,
migrations: ModuleType | str | None = None,
migration_ctx: Any | None = None,
serializer: Serializer | None = None,
flush_interval: float = 0.1,
):
@@ -62,7 +61,6 @@ class Kanta(Generic[T]):
data: Caller-owned root msgspec.Struct state instance.
type: Optional explicit root type. Defaults to ``type(data)``.
migrations: Optional migrations module object or import path.
migration_ctx: Optional context object passed to migration functions.
flush_interval: Background flush interval in seconds.
serializer: Optional serializer implementation.
@@ -79,8 +77,8 @@ class Kanta(Generic[T]):
data=data,
type=data_type,
migrations=migrations,
migration_ctx=migration_ctx,
flush_interval=flush_interval,
kanta=self,
)
@property
@@ -127,6 +125,15 @@ class Kanta(Generic[T]):
"""
return self._impl.filename
@property
def ctx(self) -> SimpleNamespace:
"""User-writable context namespace.
Migration functions receive the ``Kanta`` instance and can read or
mutate ``kanta.ctx`` during migrations.
"""
return self._impl.ctx
@property
def mtime(self) -> datetime | None:
"""Last modification time carried forward from change records.
@@ -138,7 +145,13 @@ class Kanta(Generic[T]):
"""
return self._impl.mtime
async def open(self, *, create: bool = True) -> 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.
This loads existing records, applies configured migrations, and starts
@@ -147,6 +160,14 @@ class Kanta(Generic[T]):
Args:
create: Whether to create the database file when missing.
If False, opening fails when the file does not exist or is empty.
readonly: If True, open the database read-only. No lock is acquired,
no background flush task is started, and transactions are
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.
@@ -154,7 +175,7 @@ class Kanta(Generic[T]):
kanta.exceptions.DatabaseError: If replay or decoding fails.
kanta.exceptions.DataIntegrityError: If the instance is already open.
"""
await self._impl.open(create=create)
await self._impl.open(create=create, readonly=readonly, log=log)
async def __aenter__(self) -> Kanta[T]:
"""Enter async context manager and open the database.
@@ -202,8 +223,6 @@ class Kanta(Generic[T]):
"""
def _register(callback):
if not callable(callback):
raise TypeError("bootstrap callback must be callable")
self._impl.add_bootstrap(
callback=callback,
action=action,
@@ -225,8 +244,6 @@ class Kanta(Generic[T]):
"""
def _register(callback):
if not callable(callback):
raise TypeError("fatal error callback must be callable")
self._impl.add_fatal_error(callback)
return callback
@@ -234,28 +251,70 @@ class Kanta(Generic[T]):
return _register
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):
"""Register a transaction logfmt callback.
Can be used as ``@kanta.logfmt`` or ``@kanta.logfmt(path=...)``.
The callback is called for each value being rendered and receives the
value plus an optional ``path: str`` parameter. It must return
``str | None`` (or inherit from :class:`kanta.callbacks.LogFmt`).
When ``path`` is given, the callback is only invoked for values whose
dot-notation path matches the pattern (full match, shell-style wildcards
such as ``*`` are supported).
"""
def _register(callback):
self._impl.add_logfmt(callback, path=path)
return callback
if fn is None:
return _register
return _register(fn)
def transaction(
self,
action: str,
*,
user: str | None = None,
user_display: str | None = None,
resolver: Any = None,
mtime: bool | datetime = True,
log: bool | logging.Logger = True,
):
"""Create a transactional mutation context manager.
Args:
action: Action label stored in the change record.
user: Optional user identifier stored in metadata.
user_display: Optional display name used for logging/resolution.
resolver: Optional callable for resolving identifiers in logs.
user: Optional user identifier stored in metadata and rendered in
the log header. Register a ``@kanta.logfmt`` callback to format
the user value; the path ``"$user"`` is passed for this case.
mtime: Controls the modification time ``m``. ``True`` (default)
sets ``m`` to the current UTC time. ``False`` omits ``m`` so the
previous modification time remains in effect; this is used for
system operations that are not considered modifications. A
:class:`~datetime.datetime` value sets ``m`` to that explicit
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:
A context manager yielding the live state object for mutation.
@@ -269,7 +328,6 @@ class Kanta(Generic[T]):
self._impl,
action,
user=user,
user_display=user_display,
resolver=resolver,
mtime=mtime,
log=log,
)
+159 -34
View File
@@ -5,13 +5,15 @@ from __future__ import annotations
import asyncio
import copy
import importlib
import inspect
import logging
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import Any, Generic, TypeVar
from kanta.callbacks import CallbackRegistry, InjectionContext
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.serialization import restore_data_in_place, struct_to_dict
from kanta.serialization.base import replay
@@ -27,31 +29,35 @@ class KantaImpl(PersistenceMixin, Generic[T]):
def __init__(self, **kwargs: Any):
self.data_type = kwargs.pop("type")
self.data: T = kwargs.pop("data")
self.migrations = kwargs.pop("migrations", None)
self.migration_ctx = kwargs.pop("migration_ctx", None)
self._kanta = kwargs.pop("kanta", None)
migrations = kwargs.pop("migrations", None)
self.ctx = SimpleNamespace()
super().__init__(**kwargs)
self.migration_registry: MigrationRegistry | None = None
if self.migrations is not None:
self.migrations: Migrations | None = None
if migrations is not None:
module = (
importlib.import_module(self.migrations)
if isinstance(self.migrations, str)
else self.migrations
importlib.import_module(migrations)
if isinstance(migrations, str)
else migrations
)
self.migration_registry = MigrationRegistry.from_module(module)
self.migrations = Migrations.from_module(module)
self.in_transaction = False
self.transaction_snapshot: dict[str, Any] | None = None
self.opened = False
self.bootstrap_callbacks: list[Any] = []
self.readonly = False
self.bootstrap_action = "bootstrap"
self.bootstrap_user: str | None = None
self.bootstrap_mtime: bool | datetime = True
self.statedict = struct_to_dict(self.data, serializer=self.serializer)
self.version = (
self.migration_registry.dbver if self.migration_registry is not None else 0
self.callback_registry = CallbackRegistry(
kanta_class=type(self._kanta) if self._kanta is not None else None,
data_type=self.data_type,
)
self.statedict = struct_to_dict(self.data, serializer=self.serializer)
self.version = self.migrations.dbver if self.migrations is not None else 0
def add_bootstrap(
self,
*,
@@ -61,12 +67,73 @@ class KantaImpl(PersistenceMixin, Generic[T]):
mtime: bool | datetime,
) -> None:
"""Add bootstrap callback and update bootstrap metadata."""
self.bootstrap_callbacks.append(callback)
self.callback_registry.register("bootstrap", callback)
self.bootstrap_action = action
self.bootstrap_user = user
self.bootstrap_mtime = mtime
async def open(self, *, create: bool = True) -> None:
def add_logfmt(self, callback, *, path: str | None = None) -> None:
"""Register one transaction logfmt callback."""
self.callback_registry.register("logfmt", callback, path=path)
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."""
if self.opened:
raise DataIntegrityError(
@@ -75,12 +142,17 @@ class KantaImpl(PersistenceMixin, Generic[T]):
action="open",
)
self.readonly = readonly
existed_before_open = self.filename.exists()
# Read-only mode never creates the file.
open_create = create and not readonly
content = await asyncio.to_thread(
self.file.open_and_read,
self.filename,
create=create,
create=open_create,
readonly=readonly,
)
if not create and (not existed_before_open or not content):
@@ -96,6 +168,9 @@ class KantaImpl(PersistenceMixin, Generic[T]):
action="open",
)
# From this point the file is open and must be closed via close().
self.opened = True
if content:
try:
rr = replay(
@@ -125,12 +200,33 @@ class KantaImpl(PersistenceMixin, Generic[T]):
cause_type=type(e).__name__,
) from e
if self.migration_registry is not None:
rr.version = self.migration_registry.apply(
rr.state, rr.version, self.migration_ctx
migration_result = None
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.version = migration_result.version
migrations_ran = rr.version != previous_version
migration_state_changed = (
state_before_migrations is not None
and state_before_migrations != rr.state
)
self.statedict = copy.deepcopy(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,
rr.state,
@@ -140,27 +236,56 @@ class KantaImpl(PersistenceMixin, Generic[T]):
self.version = rr.version
self.mtime = rr.m
normalized = struct_to_dict(self.data, serializer=self.serializer)
self.queue_change("migrate:msgspec", normalized, mtime=False)
self.snapshot.ts = (
datetime.fromtimestamp(rr.last_snapshot_mtime, UTC)
if rr.last_snapshot_mtime is not None
else None
if self.readonly:
self.statedict = copy.deepcopy(normalized)
else:
if migrations_ran and migration_state_changed:
self.queue_change(
f"migrate:v{self.version}",
rr.state,
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
)
elif self.bootstrap_callbacks:
try:
for callback in self.bootstrap_callbacks:
callback_result = callback(self.data)
if inspect.isawaitable(callback_result):
await callback_result
if migrations_ran and migration_result is not None:
await self._handle_migration_log(
migration_result, previous_version, log
)
elif self.readonly:
self.opened = False
self.file.close()
raise DataIntegrityError(
"Cannot open empty database in read-only mode",
db_path=self.filename,
action="open",
)
else:
try:
if self.callback_registry.has("bootstrap"):
await self.callback_registry.invoke(
"bootstrap",
InjectionContext(data=self.data, kanta=self._kanta),
)
self.statedict = {}
current = struct_to_dict(self.data, serializer=self.serializer)
self.queue_change(
self.bootstrap_action,
current,
user=self.bootstrap_user,
mtime=self.bootstrap_mtime,
force=True,
)
except Exception:
self.opened = False
self.file.close()
try:
await asyncio.to_thread(self.filename.unlink, missing_ok=True)
@@ -168,8 +293,7 @@ class KantaImpl(PersistenceMixin, Generic[T]):
pass
raise
self.opened = True
if not self.readonly:
self.background_task = asyncio.create_task(self._background_loop())
async def close(self) -> None:
@@ -187,6 +311,7 @@ class KantaImpl(PersistenceMixin, Generic[T]):
# Always run a final flush in case the background task never reached
# its cancellation handler.
if not self.readonly:
await self.flush()
self.file.close()
+100 -74
View File
@@ -10,7 +10,8 @@ import sys
from collections.abc import Callable
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
_UNSAFE_CHARS = re.compile(
@@ -31,11 +32,30 @@ _ADD = "\033[0;32m" # Green for additions
_ACTION = "\033[1;34m" # Bold blue for action name
_USER = "\033[0;34m" # Blue for user display
# Metadata path used when formatting the transaction actor.
_USER_PATH = "$user"
def _join_path(path: str, key: str) -> str:
"""Append *key* to a dot-notation *path*."""
if not path:
return key
return f"{path}.{key}"
def _format_value(
value: Any, max_len: int = 60, resolver: Callable[[str], str] | None = None
value: Any,
path: str,
*,
max_len: int = 60,
logfmt: Callable[[Any, str], str | None] | None = None,
) -> str:
"""Format a value for display, truncating if needed."""
if logfmt is not None:
resolved = logfmt(value, path)
if resolved is not None:
return resolved
if value is None:
return "null"
if isinstance(value, bool):
@@ -44,10 +64,6 @@ def _format_value(
return str(value)
if isinstance(value, str):
value = _UNSAFE_CHARS.sub("", value)
if resolver is not None:
resolved = resolver(value)
if resolved != value:
return resolved
if len(value) > max_len:
return value[: max_len - 3] + "..."
return value
@@ -57,17 +73,21 @@ def _format_value(
all_true = all(v is True for v in value.values())
parts = []
for k, v in value.items():
key_display = resolver(k) if resolver is not None else k
key_path = _join_path(path, str(k))
key_display = _format_value(k, key_path, max_len=30, logfmt=logfmt)
if all_true:
parts.append(key_display)
else:
val_display = _format_value(v, max_len=30, resolver=resolver)
val_display = _format_value(v, key_path, max_len=30, logfmt=logfmt)
parts.append(f"{key_display}: {val_display}")
return "{" + ", ".join(parts) + "}"
if isinstance(value, list):
if not value:
return "[]"
parts = [_format_value(v, max_len=30, resolver=resolver) for v in value]
parts = []
for i, v in enumerate(value):
item_path = _join_path(path, str(i))
parts.append(_format_value(v, item_path, max_len=30, logfmt=logfmt))
return "[" + ", ".join(parts) + "]"
text = str(value)
if len(text) > max_len:
@@ -75,16 +95,35 @@ def _format_value(
return text
def _format_path(path: list[str], resolver: Callable[[str], str] | None = None) -> str:
"""Format a path as dot notation with prefix in dark grey, final in default."""
def _format_path_components(
path: list[str], logfmt: Callable[[Any, str], str | None] | None
) -> list[str]:
"""Return path components after applying formatters."""
if not path:
return []
result = []
for i, component in enumerate(path):
prefix_path = ".".join(path[: i + 1])
display = component
if logfmt is not None:
resolved = logfmt(component, prefix_path)
if resolved is not None:
display = resolved
result.append(display)
return result
def _format_path(
path: list[str], logfmt: Callable[[Any, str], str | None] | None
) -> str:
"""Format a path as dot notation with prefix in dark grey, final in default."""
components = _format_path_components(path, logfmt)
if not components:
return ""
if resolver is not None:
path = [resolver(p) for p in path]
if len(path) == 1:
return f"{_PATH_FINAL}{path[0]}{_RESET}"
prefix = ".".join(path[:-1])
final = path[-1]
if len(components) == 1:
return f"{_PATH_FINAL}{components[0]}{_RESET}"
prefix = ".".join(components[:-1])
final = components[-1]
return f"{_PATH_PREFIX}{prefix}.{_RESET}{_PATH_FINAL}{final}{_RESET}"
@@ -158,74 +197,56 @@ def _format_change_lines(
change_type: str,
path: list[str],
value: Any,
resolver: Callable[[str], str] | None = None,
logfmt: Callable[[Any, str], str | None] | None = None,
) -> list[str]:
"""Format a single change as one or more lines."""
def fmt_value(v: Any, child_path: list[str]) -> str:
return _format_value(v, resolver=resolver)
formatted_path = list(path)
if resolver is not None:
formatted_path = [resolver(p) for p in formatted_path]
path_str = _format_path(path, logfmt=logfmt)
if change_type == "delete":
if len(formatted_path) == 1:
return [f" {_DELETE}{formatted_path[0]}{_RESET}"]
prefix = ".".join(formatted_path[:-1])
final = formatted_path[-1]
components = _format_path_components(path, logfmt)
if len(components) == 1:
return [f" {_DELETE}{components[0]}{_RESET}"]
prefix = ".".join(components[:-1])
final = components[-1]
return [f" {_PATH_PREFIX}{prefix}.{_RESET}{_DELETE}{final}{_RESET}"]
if change_type == "add":
if isinstance(value, dict) and value:
lines = []
if len(formatted_path) == 1:
lines.append(f" {_ADD}{formatted_path[0]}{_RESET} {_SEP}={_RESET}")
else:
prefix = ".".join(formatted_path[:-1])
final = formatted_path[-1]
lines.append(
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_SEP}={_RESET}"
)
lines = [f" {path_str} {_SEP}={_RESET}"]
formatted_items = []
base_path = ".".join(path)
for k, v in value.items():
k_display = resolver(k) if resolver is not None else k
v_str = fmt_value(v, path + [k])
formatted_items.append((k_display, v_str))
key_path = _join_path(base_path, str(k))
key_display = _format_value(k, key_path, max_len=30, logfmt=logfmt)
v_str = _format_value(v, key_path, max_len=30, logfmt=logfmt)
formatted_items.append((key_display, v_str))
max_key_len = max(len(k) for k, _ in formatted_items)
field_width = max(max_key_len, 12)
for k_display, v_str in formatted_items:
padding = " " * (field_width - len(k_display))
lines.append(f" {k_display}{_SEP}:{_RESET}{padding} {v_str}")
return lines
else:
value_str = fmt_value(value, path)
if len(formatted_path) == 1:
return [
f" {_ADD}{formatted_path[0]}{_RESET} {_SEP}={_RESET} {value_str}"
]
prefix = ".".join(formatted_path[:-1])
final = formatted_path[-1]
return [
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_SEP}={_RESET} {value_str}"
]
value_str = _format_value(value, ".".join(path), logfmt=logfmt)
return [f" {path_str} {_SEP}={_RESET} {value_str}"]
value_str = fmt_value(value, path)
path_str = _format_path(path, resolver=resolver)
value_str = _format_value(value, ".".join(path), logfmt=logfmt)
return [f" {path_str} {_SEP}={_RESET} {value_str}"]
def format_diff(
diff: dict,
previous: dict | None = None,
resolver: Callable[[str], str] | None = None,
logfmt: Callable[[Any, str], str | None] | None = None,
) -> list[str]:
"""Format a JSON diff as human-readable lines.
Args:
diff: The JSON diff dict.
previous: The previous state dict (for determining add vs update).
resolver: Optional callable to resolve path components (e.g. UUID→name).
logfmt: Optional formatter callable ``(value, path) -> str | None``.
``path`` is a dot-notation string; ``"$user"`` is used for the
transaction actor. If the callable returns ``None``, default
formatting is used.
Returns a list of formatted lines (without newlines).
"""
@@ -235,15 +256,15 @@ def format_diff(
return []
lines = []
for change_type, path, value in changes:
lines.extend(_format_change_lines(change_type, path, value, resolver))
lines.extend(_format_change_lines(change_type, path, value, logfmt))
return lines
def format_action_header(action: str, user_display: str | None = None) -> str:
def format_action_header(action: str, user: str | None = None) -> str:
"""Format the action header line."""
action_str = f"{_ACTION}{action}{_RESET}"
if user_display:
user_str = f"{_USER}{user_display}{_RESET}"
if user:
user_str = f"{_USER}{user}{_RESET}"
return f"{action_str} by {user_str}"
return action_str
@@ -251,39 +272,44 @@ def format_action_header(action: str, user_display: str | None = None) -> str:
def log_change(
action: str,
diff: dict,
user_display: str | None = None,
user: str | None = None,
previous: dict | None = None,
resolver: Callable[[str], str] | None = None,
logfmt: Callable[[Any, str], str | None] | None = None,
*,
logger: logging.Logger = changes_logger,
level: int = logging.INFO,
) -> None:
"""Log a database change with pretty-printed diff.
Args:
action: The action name (e.g., "login", "admin:delete_user").
diff: The JSON diff dict.
user_display: Optional display name of the user who performed the action.
user: Optional already-formatted user name to show in the header.
previous: The previous state dict (for determining add vs update).
resolver: Optional callable to resolve path components (e.g. UUID→name).
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_display)
diff_lines = format_diff(diff, previous, resolver)
header = format_action_header(action, user)
diff_lines = format_diff(diff, previous, logfmt)
if not diff_lines:
logger.info(header)
logger.log(level, header)
return
if len(diff_lines) == 1:
logger.info(f"{header}{diff_lines[0]}")
logger.log(level, f"{header}{diff_lines[0]}")
else:
logger.info(header)
logger.log(level, header)
for line in diff_lines:
logger.info(line)
logger.log(level, line)
def configure_logging() -> None:
"""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.setFormatter(logging.Formatter("%(message)s"))
logger.addHandler(handler)
logger.setLevel(logging.INFO)
logger.propagate = False
changes_logger.addHandler(handler)
changes_logger.setLevel(logging.INFO)
changes_logger.propagate = False
-117
View File
@@ -1,117 +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 logging
from types import ModuleType
from typing import Any
import msgspec
_logger = logging.getLogger(__name__)
class MigrationCtx(msgspec.Struct, omit_defaults=True):
"""Context passed to each migration function.
Subclass or replace this with your own context type.
"""
pass
class MigrationRegistry:
"""Registry of schema migration functions.
Usage::
registry = MigrationRegistry()
@registry.register
def migrate_v1(d: dict, ctx: MigrationCtx) -> None:
d.setdefault("version", 1)
new_version = registry.apply(state, current_version=0)
Or load from a module::
registry = MigrationRegistry.from_module("myapp.migrations")
new_version = registry.apply(state, current_version=0)
"""
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)
def apply(
self,
data_dict: dict[str, Any],
current_version: int,
ctx: MigrationCtx | None = None,
*,
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})"
)
fn(data_dict, ctx or MigrationCtx())
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)
+39 -14
View File
@@ -4,14 +4,13 @@ from __future__ import annotations
import asyncio
import copy
import inspect
import logging
from collections import deque
from collections.abc import Callable
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
from kanta.callbacks import CallbackRegistry, InjectionContext
from kanta.diff import compute_diff
from kanta.exceptions import DatabaseError, DataIntegrityError
from kanta.filelock import LockedFile
@@ -35,11 +34,12 @@ class PersistenceMixin:
serializer: Serializer
framer: Framer
background_task: asyncio.Task | None
fatal_error_handlers: list[Callable[[DatabaseError], Any]]
callback_registry: CallbackRegistry
background_error: DatabaseError | None
flush_interval: float
version: int
opened: bool
readonly: bool
mtime: datetime | None
def __init__(self, **kwargs: Any) -> None:
@@ -57,18 +57,20 @@ class PersistenceMixin:
self.framer = self.serializer.framer_cls()
self.snapshot = SnapshotState(serializer=self.serializer, framer=self.framer)
self.background_task = None
self.fatal_error_handlers = []
self.callback_registry = CallbackRegistry()
self.background_error = None
self.flush_interval = flush_interval
self.version = 0
self.mtime: datetime | None = None
def add_fatal_error(self, callback: Callable[[DatabaseError], Any]) -> None:
def add_fatal_error(self, callback) -> None:
"""Register one fatal error callback in call order."""
self.fatal_error_handlers.append(callback)
self.callback_registry.register("fatal_error", callback)
async def _background_loop(self) -> None:
"""Background task that periodically flushes changes to disk."""
if self.readonly:
return
while True:
try:
await asyncio.sleep(self.flush_interval)
@@ -80,14 +82,18 @@ class PersistenceMixin:
break
except DatabaseError as e:
self.background_error = e
for callback in self.fatal_error_handlers:
try:
callback_result = callback(e)
if inspect.isawaitable(callback_result):
await callback_result
except Exception as callback_error:
def _log_callback_error(callback_error, callback):
_logger.exception(
"Background error callback failed: %s", callback_error
"Background error callback %r failed: %s",
callback,
callback_error,
)
await self.callback_registry.invoke(
"fatal_error",
InjectionContext(error=e, kanta=self._kanta),
on_error=_log_callback_error,
)
_logger.error("Background flush loop stopped: %s", e)
break
@@ -103,6 +109,7 @@ class PersistenceMixin:
*,
user: str | None = None,
mtime: bool | datetime = True,
force: bool = False,
) -> ChangeRecord | None:
"""Queue a change record internally (thread-safe).
@@ -115,9 +122,11 @@ class PersistenceMixin:
previous modification time remains in effect; this is used for
system operations that are not considered modifications. A
:class:`~datetime.datetime` value sets ``m`` to that explicit time.
force: If ``True``, queue the record even when the diff is empty.
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)
@@ -132,7 +141,9 @@ class PersistenceMixin:
diff = compute_diff(self.statedict, current)
if not diff:
if not force:
return None
diff = {}
record = ChangeRecord(
ts=now,
@@ -157,6 +168,13 @@ class PersistenceMixin:
action="flush_sync",
)
if self.readonly:
raise DataIntegrityError(
"Cannot flush in read-only mode",
db_path=self.filename,
action="flush_sync",
)
if self.flush_failed:
return
@@ -204,6 +222,13 @@ class PersistenceMixin:
action="flush",
)
if self.readonly:
raise DataIntegrityError(
"Cannot flush in read-only mode",
db_path=self.filename,
action="flush",
)
if self.flush_failed:
return
+5 -4
View File
@@ -41,14 +41,15 @@ class SnapshotState:
self, file, version: int, state: dict, m: datetime | None = None
) -> None:
"""Write snapshot when thresholds/time policy allows it."""
if self.changes < self._min_diffs:
return
force = self._force_pending
now = datetime.now(UTC)
if not force and now.weekday() != 6: # 6 = Sunday
if not force:
if self.changes < self._min_diffs:
return
if now.weekday() != 6: # 6 = Sunday
return
sunday_midnight = now.replace(hour=0, minute=0, second=0, microsecond=0)
if not force and self.ts is not None and self.ts >= sunday_midnight:
if self.ts is not None and self.ts >= sunday_midnight:
return
if not file.is_open:
return
+1 -1
View File
@@ -23,7 +23,7 @@ class ChangeRecord(msgspec.Struct, omit_defaults=True, kw_only=True):
v: int = 0
u: str | None = None
m: datetime | None = None
diff: dict
diff: dict = {}
class Snapshot(msgspec.Struct, omit_defaults=True):
+32 -5
View File
@@ -5,11 +5,11 @@ from __future__ import annotations
import logging
from contextlib import contextmanager
from datetime import datetime
from typing import Any
from kanta.diff import compute_diff
from kanta.exceptions import DataIntegrityError
from kanta.logging import log_change
from kanta.callbacks import InjectionContext
from kanta.logging import _USER_PATH, changes_logger, log_change
from kanta.serialization import restore_data_in_place, struct_to_dict
_logger = logging.getLogger(__name__)
@@ -21,11 +21,17 @@ def transaction(
action: str,
*,
user: str | None = None,
user_display: str | None = None,
resolver: Any = None,
mtime: bool | datetime = True,
log: bool | logging.Logger = True,
):
"""Wrap writes in a transaction and yield the live db object."""
if impl.readonly:
raise DataIntegrityError(
"Cannot start transaction in read-only mode",
db_path=impl.filename,
action=action,
)
if impl.in_transaction:
raise RuntimeError(
"Nested or simultaneous transactions are not supported "
@@ -63,7 +69,28 @@ def transaction(
previous = impl.statedict
record = impl.queue_change(action, new_dict, user=user, mtime=mtime)
if record is not None:
log_change(action, record.diff, user_display, previous, resolver)
logfmt = impl.callback_registry.build_logfmt(
InjectionContext(
previous_state=previous,
current_state=new_dict,
kanta=impl._kanta,
)
)
formatted_user = user
if user is not None and logfmt is not None:
resolved = logfmt(user, _USER_PATH)
if resolved is not None:
formatted_user = resolved
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:
_logger.warning("Transaction '%s' failed, rolling back changes", action)
if impl.transaction_snapshot is not None:
+1 -1
View File
@@ -1,6 +1,6 @@
import pytest
from kanta import JsonSerializer, MsgPackSerializer
from kanta.serialization import JsonSerializer, MsgPackSerializer
@pytest.fixture(
+25 -1
View File
@@ -6,7 +6,8 @@ from uuid import UUID
import msgspec
from kanta import ChangeRecord, Kanta
from kanta.kanta import Kanta
from kanta.structs import ChangeRecord, Snapshot
class User(msgspec.Struct):
@@ -69,6 +70,18 @@ def change_actions(path: Path, format_config) -> list[str]:
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):
mod = ModuleType(name)
mod.__dict__[fn_name] = fn
@@ -76,6 +89,17 @@ def make_migrations_module(name: str, fn_name: str, fn):
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:
return ChangeRecord(
ts=datetime(2026, 1, 1, tzinfo=UTC), a=action, v=version, diff=diff
+349
View File
@@ -0,0 +1,349 @@
from typing import Any, Optional, Union
import pytest
from kanta import Kanta
from kanta.callbacks import DictPost, DictPre, LogFmt
from kanta.exceptions import DatabaseError
from .support import Data, User, make_kanta
def test_bootstrap_rejects_unannotated_param(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
with pytest.raises(TypeError, match="without an annotation or default"):
@kanta.bootstrap
def seed(data):
data.counter = 1
def test_bootstrap_accepts_unknown_with_default(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
@kanta.bootstrap
def seed(data: Data, extra: int = 0) -> None:
data.counter = extra + 1
# Should register without error.
def test_bootstrap_rejects_unknown_annotation(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
with pytest.raises(TypeError, match="unsupported annotation"):
@kanta.bootstrap
def seed(data: int):
pass
def test_logfmt_requires_value_annotation(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
with pytest.raises(TypeError, match="value parameter.*must be annotated"):
@kanta.logfmt
def resolve_names(previous: DictPre, current: DictPost) -> str | None:
return None
def test_logfmt_allows_missing_return_annotation(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
@kanta.logfmt
def resolve_names(value: str, current: DictPost):
return None
def test_logfmt_class_allows_missing_return_annotation(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
@kanta.logfmt
class UserLogFmt(LogFmt):
def resolve(self, value: str, path: str):
return None
# fmt: off
def test_logfmt_accepts_optional_return_typing_forms(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
@kanta.logfmt
def resolve_optional(value: str) -> Optional[str]: # noqa: UP007
return value
@kanta.logfmt
def resolve_union(value: str) -> Union[str, None]: # noqa: UP007
return value
@kanta.logfmt
def resolve_pipe(value: "str") -> "str | None":
return value
def test_logfmt_class_accepts_optional_return_typing_forms(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
@kanta.logfmt
class OptionalStyle(LogFmt):
def resolve(self, value: str, path: str) -> Optional[str]: # noqa: UP007
return value
@kanta.logfmt
class UnionStyle(LogFmt):
def resolve(self, value: str, path: str) -> Union[str, None]: # noqa: UP007
return value
@kanta.logfmt
class StringStyle(LogFmt):
def resolve(self, value: "str", path: "str") -> "str | None":
return value
# fmt: on
def test_logfmt_rejects_async_callback(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
with pytest.raises(TypeError, match="must not be async"):
@kanta.logfmt
async def resolve_names(value: str, current: DictPost) -> str | None:
return None
@pytest.mark.asyncio
async def test_bootstrap_injects_data_by_type(tmp_path, format_config):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap
def seed(data: Data) -> None:
data.counter = 7
await kanta.open()
assert kanta.data.counter == 7
await kanta.close()
@pytest.mark.asyncio
async def test_bootstrap_injects_kanta(tmp_path, format_config):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
seen: list[Kanta] = []
@kanta.bootstrap
def seed(data: Data, kanta_ref: Kanta) -> None:
seen.append(kanta_ref)
data.counter = 8
await kanta.open()
assert seen == [kanta]
assert kanta.data.counter == 8
await kanta.close()
@pytest.mark.asyncio
async def test_logfmt_injects_states(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt
def resolve_users(value: str, current: DictPost) -> str | None:
return current.get("users", {}).get(value, {}).get("name")
await kanta.open()
with kanta.transaction(action="create_user") as data:
data.users["uuid-1"] = User(name="Alice")
await kanta.close()
assert "Alice" in caplog.text
@pytest.mark.asyncio
async def test_logfmt_class_injection(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt
class UserLogFmt(LogFmt):
def resolve(self, value: str, path: str) -> str | None:
return self.current_state.get("users", {}).get(value, {}).get("name")
await kanta.open()
with kanta.transaction(action="create_user") as data:
data.users["uuid-2"] = User(name="Bob")
await kanta.close()
assert "Bob" in caplog.text
@pytest.mark.asyncio
async def test_multiple_logfmt_chain(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt
def resolve_a(value: str) -> str | None:
return "A" if value == "a" else None
@kanta.logfmt
def resolve_b(value: str) -> str | None:
return "B" if value == "b" else None
await kanta.open()
with kanta.transaction(action="create_user") as data:
data.users["a"] = User(name="first")
data.users["b"] = User(name="second")
await kanta.close()
assert "A" in caplog.text
assert "B" in caplog.text
@pytest.mark.asyncio
async def test_logfmt_path_context(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt(path="users.uuid-1")
def resolve_user_key(value: str) -> str | None:
if value == "uuid-1":
return "user-alice"
return None
await kanta.open()
with kanta.transaction(action="create_user") as data:
data.users["uuid-1"] = User(name="Alice")
await kanta.close()
assert "user-alice" in caplog.text
@pytest.mark.asyncio
async def test_logfmt_decorator_path_filters_calls(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt(path="counter")
def fmt_counter(value: Any) -> str | None:
if value == 1:
return "one"
return None
await kanta.open()
with kanta.transaction(action="create_user") as data:
data.users["uuid-1"] = User(name="Alice")
data.counter = 1
await kanta.close()
assert "one" in caplog.text
assert "uuid-1" in caplog.text
@pytest.mark.asyncio
async def test_logfmt_user_path_replaces_user_display(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt(path="$user")
def resolve_user(value: str, current: DictPost) -> str | None:
return current.get("users", {}).get(value, {}).get("name")
await kanta.open()
with kanta.transaction(action="create_user", user="uuid-1") as data:
data.users["uuid-1"] = User(name="Alice")
await kanta.close()
assert "by Alice" in caplog.text
@pytest.mark.asyncio
async def test_logfmt_non_string_value(tmp_path, format_config, caplog):
import logging
caplog.set_level(logging.INFO, logger="kanta.changes")
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
@kanta.logfmt
def fmt_count(value: Any, path: str) -> str | None:
if path == "counter" and value == 1:
return "one"
return None
await kanta.open()
with kanta.transaction(action="inc") as data:
data.counter = 1
await kanta.close()
assert "one" in caplog.text
@pytest.mark.asyncio
async def test_fatal_error_injects_kanta_and_error(
tmp_path, format_config, monkeypatch
):
import asyncio
path = tmp_path / "test.db"
errors: list[DatabaseError] = []
kantas: list[Kanta] = []
signaled = asyncio.Event()
kanta = make_kanta(path, Data, format_config, flush_interval=0.01)
@kanta.fatal_error
def on_fatal(error: DatabaseError, kanta_ref: Kanta) -> None:
errors.append(error)
kantas.append(kanta_ref)
signaled.set()
await kanta.open()
with kanta.transaction(action="inc") as data:
data.counter = 1
def fail_write(_data: bytes) -> None:
raise OSError("simulated background write failure")
monkeypatch.setattr(kanta._impl.file, "write", fail_write)
await asyncio.wait_for(signaled.wait(), timeout=1.0)
assert errors
assert kantas == [kanta]
await kanta.close()
+1 -1
View File
@@ -1,4 +1,4 @@
from kanta import compute_diff
from kanta.diff import compute_diff
def test_no_diff():
+27 -3
View File
@@ -1,4 +1,4 @@
from kanta import format_diff
from kanta.logging import format_diff
def test_add():
@@ -16,10 +16,34 @@ def test_delete():
assert any("old_key" in line for line in lines)
def test_resolver():
def test_logfmt():
lines = format_diff(
{"users": {"uuid-1": {"name": "Alice"}}},
previous={},
resolver=lambda x: "Alice" if x == "uuid-1" else x,
logfmt=lambda value, path: "Alice" if value == "uuid-1" else None,
)
assert any("Alice" in line for line in lines)
def test_logfmt_uses_path_context():
lines = format_diff(
{
"users": {"uuid-1": {"name": "Alice"}},
"groups": {"uuid-1": {"name": "Admins"}},
},
previous={},
logfmt=lambda value, path: (
"User Alice" if path.startswith("users.") and value == "uuid-1" else None
),
)
assert any("User Alice" in line for line in lines)
assert any("uuid-1" in line for line in lines)
def test_logfmt_formats_non_string_value():
lines = format_diff(
{"count": 42},
previous={},
logfmt=lambda value, path: "forty-two" if value == 42 else None,
)
assert any("forty-two" in line for line in lines)
+268 -9
View File
@@ -1,4 +1,5 @@
import asyncio
import logging
import sys
from datetime import UTC, datetime
from uuid import uuid4
@@ -6,6 +7,7 @@ from uuid import uuid4
import pytest
from kanta.exceptions import DatabaseError, DataIntegrityError, FileLockError
from kanta.migrations import MigrationResult
from kanta.serialization import struct_to_dict
from .support import (
@@ -17,6 +19,9 @@ from .support import (
change_actions,
fixed_change,
make_kanta,
make_migrations_module,
read_changes,
read_last_snapshot,
seed_single_change,
)
@@ -30,6 +35,63 @@ async def test_load_empty(tmp_path, format_config):
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
async def test_open_overwrites_caller_owned_root_data(tmp_path, format_config):
path = tmp_path / "test.db"
@@ -106,7 +168,7 @@ async def test_bootstrap_decorator_with_args(tmp_path, format_config):
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap(action="seed_init", user="system")
def seed(data):
def seed(data: Data):
data.counter = 3
await kanta.open()
@@ -121,7 +183,7 @@ async def test_bootstrap_decorator_without_args(tmp_path, format_config):
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap
def seed(data):
def seed(data: Data):
data.counter = 4
await kanta.open()
@@ -136,7 +198,7 @@ async def test_bootstrap_decorator_async(tmp_path, format_config):
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap(action="async_seed")
async def seed(data):
async def seed(data: Data):
await asyncio.sleep(0)
data.counter = 5
@@ -152,11 +214,11 @@ async def test_bootstrap_decorator_multiple_handlers_in_order(tmp_path, format_c
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap(action="boot_1")
def seed_one(data):
def seed_one(data: Data):
data.counter = 1
@kanta.bootstrap(action="boot_2")
async def seed_two(data):
async def seed_two(data: Data):
await asyncio.sleep(0)
data.counter = 2
@@ -172,7 +234,7 @@ async def test_bootstrap_failure_removes_database_file(tmp_path, format_config):
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap(action="boot_fail")
def seed_fail(data):
def seed_fail(data: Data):
data.counter = 10
raise RuntimeError("bootstrap failed")
@@ -188,7 +250,7 @@ async def test_bootstrap_async_failure_removes_database_file(tmp_path, format_co
kanta = make_kanta(path, Data, format_config)
@kanta.bootstrap(action="boot_fail_async")
async def seed_fail(data):
async def seed_fail(data: Data):
await asyncio.sleep(0)
data.counter = 10
raise RuntimeError("bootstrap async failed")
@@ -404,7 +466,7 @@ async def test_migrations_from_module(tmp_path, format_config):
mod = type(sys)("test_migrations")
def migrate_v1(d, ctx):
def migrate_v1(d, kanta):
d["version"] = 1
mod.__dict__["migrate_v1"] = migrate_v1
@@ -434,6 +496,203 @@ 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_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
async def test_open_locked_file_raises_filelock_error(tmp_path, format_config):
path = tmp_path / "test.db"
@@ -514,7 +773,7 @@ async def test_migrations_from_module_path(tmp_path, format_config):
module_name = "test_migrations_path"
mod = type(sys)(module_name)
def migrate_v1(d, ctx):
def migrate_v1(d, kanta):
d["counter"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
+3 -4
View File
@@ -1,16 +1,15 @@
import logging
from kanta import configure_logging, log_change
from kanta.logging import logger
from kanta.logging import changes_logger, configure_logging, log_change
def test_configure_logging():
configure_logging()
assert logger.level == logging.INFO
assert changes_logger.level == logging.INFO
def test_log_change_no_diff(capsys):
logger.handlers.clear()
changes_logger.handlers.clear()
configure_logging()
log_change("test", {})
captured = capsys.readouterr()
+171 -16
View File
@@ -1,53 +1,208 @@
from types import ModuleType
from types import ModuleType, SimpleNamespace
from kanta.migrate import MigrationRegistry
import pytest
from kanta.exceptions import DatabaseError
from kanta.migrations import Migrations
class _DummyKanta:
def __init__(self):
self.ctx = SimpleNamespace()
def test_register_and_apply():
reg = MigrationRegistry()
reg = Migrations()
kanta = _DummyKanta()
@reg.register
def migrate_v1(d, ctx):
def migrate_v1(d, kanta):
d["version"] = 1
@reg.register
def migrate_v2(d, ctx):
def migrate_v2(d, kanta):
d["version"] = 2
state = {}
new_ver = reg.apply(state, current_version=0, silent=True)
assert new_ver == 2
result = reg.apply(state, current_version=0, kanta=kanta)
assert result.version == 2
assert state["version"] == 2
def test_no_migrations_needed():
reg = MigrationRegistry()
reg = Migrations()
kanta = _DummyKanta()
@reg.register
def migrate_v1(d, ctx):
def migrate_v1(d, kanta):
d["x"] = 1
state = {"x": 1}
new_ver = reg.apply(state, current_version=1, silent=True)
assert new_ver == 1
result = reg.apply(state, current_version=1, kanta=kanta)
assert result.version == 1
def test_from_module():
mod = ModuleType("fake_migrations")
kanta = _DummyKanta()
def migrate_v1(d, ctx):
def migrate_v1(d, kanta):
d["v"] = 1
def migrate_v2(d, ctx):
def migrate_v2(d, kanta):
d["v"] = 2
mod.__dict__["migrate_v1"] = migrate_v1
mod.__dict__["migrate_v2"] = migrate_v2
reg = MigrationRegistry.from_module(mod)
reg = Migrations.from_module(mod)
assert reg.dbver == 2
state = {}
new_ver = reg.apply(state, current_version=0, silent=True)
assert new_ver == 2
result = reg.apply(state, current_version=0, kanta=kanta)
assert result.version == 2
assert state["v"] == 2
def test_migrations_can_use_kanta_ctx():
reg = Migrations()
kanta = _DummyKanta()
@reg.register
def migrate_v1(d, kanta):
kanta.ctx.source = "migration"
d["source"] = kanta.ctx.source
state = {}
result = reg.apply(state, current_version=0, kanta=kanta)
assert result.version == 1
assert state["source"] == "migration"
assert kanta.ctx.source == "migration"
def test_migration_can_omit_kanta_argument():
reg = Migrations()
kanta = _DummyKanta()
@reg.register
def migrate_v1(d):
d["x"] = 1
state = {}
result = reg.apply(state, current_version=0, kanta=kanta)
assert result.version == 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"
+4 -3
View File
@@ -4,7 +4,7 @@ from datetime import UTC, datetime
import pytest
from kanta import ChangeRecord
from kanta.structs import ChangeRecord
from .support import Data, make_kanta, seed_single_change
@@ -82,8 +82,9 @@ async def test_transaction_mtime_false_preserves_mtime(tmp_path, format_config):
continue
records.append(serializer.decode(payload, type=ChangeRecord))
assert records[0].m == first_m
assert records[1].m is None
assert records[0].a == "bootstrap"
assert records[1].m == first_m
assert records[2].m is None
assert kanta.mtime == first_m
+153
View File
@@ -0,0 +1,153 @@
"""Tests for Kanta read-only mode."""
import pytest
from kanta.exceptions import DataIntegrityError, FileLockError
from kanta.serialization import struct_to_dict
from .support import (
Data,
EvolvableDataV2,
fixed_change,
make_kanta,
make_migrations_module,
seed_single_change,
)
@pytest.mark.asyncio
async def test_readonly_opens_existing_database(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("seed", {"counter": 7}), format_config)
kanta = make_kanta(path, Data, format_config)
await kanta.open(readonly=True)
assert isinstance(kanta.data, Data)
assert kanta.data.counter == 7
assert kanta._impl.readonly is True
assert kanta._impl.background_task is None
await kanta.close()
@pytest.mark.asyncio
async def test_readonly_missing_file_fails(tmp_path, format_config):
path = tmp_path / "missing.db"
kanta = make_kanta(path, Data, format_config)
with pytest.raises(FileLockError):
await kanta.open(readonly=True)
assert not path.exists()
@pytest.mark.asyncio
async def test_readonly_empty_file_fails(tmp_path, format_config):
path = tmp_path / "empty.db"
path.touch()
kanta = make_kanta(path, Data, format_config)
with pytest.raises(DataIntegrityError, match="empty"):
await kanta.open(readonly=True)
@pytest.mark.asyncio
async def test_readonly_transaction_fails(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("seed", {"counter": 1}), format_config)
kanta = make_kanta(path, Data, format_config)
await kanta.open(readonly=True)
with pytest.raises(DataIntegrityError, match="read-only"):
with kanta.transaction(action="inc") as data:
data.counter = 2
# In-memory state must remain unchanged.
assert kanta.data.counter == 1
await kanta.close()
@pytest.mark.asyncio
async def test_readonly_flush_fails(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("seed", {"counter": 1}), format_config)
kanta = make_kanta(path, Data, format_config)
await kanta.open(readonly=True)
with pytest.raises(DataIntegrityError, match="read-only"):
await kanta.flush()
await kanta.close()
@pytest.mark.asyncio
async def test_readonly_create_true_does_not_create_file(tmp_path, format_config):
path = tmp_path / "test.db"
kanta = make_kanta(path, Data, format_config)
with pytest.raises(FileLockError):
await kanta.open(create=True, readonly=True)
assert not path.exists()
@pytest.mark.asyncio
async def test_readonly_does_not_persist_changes(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("seed", {"counter": 1}), format_config)
original_content = path.read_bytes()
kanta = make_kanta(path, Data, format_config)
await kanta.open(readonly=True)
await kanta.close()
assert path.read_bytes() == original_content
@pytest.mark.asyncio
async def test_readonly_runs_migrations(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(
path,
fixed_change("seed", {"counter": 1}, version=0),
format_config,
)
def migrate_v1(data, kanta):
data.setdefault("enabled", True)
migrations = make_migrations_module("readonly_migrations", "migrate_v1", migrate_v1)
kanta = make_kanta(path, EvolvableDataV2, format_config, migrations=migrations)
await kanta.open(readonly=True)
assert kanta.data.counter == 1
# Migration ran in memory even though no change was persisted.
assert struct_to_dict(kanta.data, serializer=kanta._impl.serializer) == {
"counter": 1,
"enabled": True,
}
assert not kanta._impl.pending_changes
await kanta.close()
@pytest.mark.asyncio
async def test_readwrite_and_readonly_can_open_together(tmp_path, format_config):
path = tmp_path / "test.db"
seed_single_change(path, fixed_change("seed", {"counter": 1}), format_config)
rw = make_kanta(path, Data, format_config)
await rw.open()
ro = make_kanta(path, Data, format_config)
await ro.open(readonly=True)
assert rw.data.counter == 1
assert ro.data.counter == 1
await ro.close()
await rw.close()
+2 -1
View File
@@ -1,6 +1,7 @@
from datetime import UTC, datetime
from kanta import ChangeRecord, Snapshot, replay
from kanta.diff import replay_jsonl as replay
from kanta.structs import ChangeRecord, Snapshot
from kanta.serialization.framing import LineFramer
+17
View File
@@ -32,3 +32,20 @@ def test_force_writes():
f = FakeFile()
ss.maybe_write(f, 1, {"x": 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