Improved log formatting support by @kanta.logfmt, which replaces old resolver and user_display arguments (breaking change).

This commit is contained in:
2026-06-13 21:22:05 +00:00
parent 4db046d627
commit 7e553bd868
11 changed files with 1091 additions and 122 deletions
+50 -6
View File
@@ -119,14 +119,22 @@ 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.
### Bootstrap Callbacks
### Callbacks
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
- Bootstrap callbacks run during `open()` when the database is empty.
- 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,
@@ -135,11 +143,47 @@ reloads, while system operations such as migrations leave it unchanged.
- 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
+5 -1
View File
@@ -1,3 +1,4 @@
from .callbacks import DictPost, DictPre, LogFmt
from .diff import compute_diff
from .diff import replay_jsonl as replay
from .exceptions import DatabaseError, DataIntegrityError, FileLockError, ReplayError
@@ -20,7 +21,10 @@ __all__ = [
"LockedFile",
"log_change",
"MsgPackSerializer",
"DictPost",
"DictPre",
"ReplayError",
"replay",
"LogFmt",
"Snapshot",
"replay",
]
+537
View File
@@ -0,0 +1,537 @@
"""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
from collections.abc import Callable
from dataclasses import dataclass
from typing import Annotated, Any, Union, get_args, get_origin
from kanta.exceptions import DatabaseError
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
@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": [],
}
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 inspect.Signature.empty:
raise TypeError(
f"logfmt callback {callback.__name__} must annotate its "
f"return type as str | None"
)
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 inspect.Signature.empty:
raise TypeError(
f"logfmt class {cls.__name__}.resolve must annotate its "
f"return type as str | None"
)
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 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"}
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 == "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 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 is not Union:
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 is not Union:
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()
+26 -13
View File
@@ -1,4 +1,4 @@
"""JSONL persistence layer with background flush task."""
"""Kanta DB main public API"""
from __future__ import annotations
from datetime import datetime
@@ -6,7 +6,6 @@ from pathlib import Path
from types import ModuleType
from typing import Any, 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
@@ -81,6 +80,7 @@ class Kanta(Generic[T]):
migrations=migrations,
migration_ctx=migration_ctx,
flush_interval=flush_interval,
kanta=self,
)
@property
@@ -202,8 +202,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 +223,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,22 +230,41 @@ class Kanta(Generic[T]):
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,
):
"""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
@@ -269,7 +284,5 @@ class Kanta(Generic[T]):
self._impl,
action,
user=user,
user_display=user_display,
resolver=resolver,
mtime=mtime,
)
+17 -8
View File
@@ -5,11 +5,11 @@ from __future__ import annotations
import asyncio
import copy
import importlib
import inspect
import logging
from datetime import UTC, datetime
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.persistence import PersistenceMixin
@@ -29,6 +29,7 @@ class KantaImpl(PersistenceMixin, Generic[T]):
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)
super().__init__(**kwargs)
self.migration_registry: MigrationRegistry | None = None
if self.migrations is not None:
@@ -42,11 +43,15 @@ class KantaImpl(PersistenceMixin, Generic[T]):
self.in_transaction = False
self.transaction_snapshot: dict[str, Any] | None = None
self.opened = False
self.bootstrap_callbacks: list[Any] = []
self.bootstrap_action = "bootstrap"
self.bootstrap_user: str | None = None
self.bootstrap_mtime: bool | datetime = True
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.migration_registry.dbver if self.migration_registry is not None else 0
@@ -61,11 +66,15 @@ 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
def add_logfmt(self, callback, *, path: str | None = None) -> None:
"""Register one transaction logfmt callback."""
self.callback_registry.register("logfmt", callback, path=path)
async def open(self, *, create: bool = True) -> None:
"""Open the database: load from disk, apply migrations, start background task."""
if self.opened:
@@ -146,12 +155,12 @@ class KantaImpl(PersistenceMixin, Generic[T]):
if rr.last_snapshot_mtime is not None
else None
)
elif self.bootstrap_callbacks:
elif self.callback_registry.has("bootstrap"):
try:
for callback in self.bootstrap_callbacks:
callback_result = callback(self.data)
if inspect.isawaitable(callback_result):
await callback_result
await self.callback_registry.invoke(
"bootstrap",
InjectionContext(data=self.data, kanta=self._kanta),
)
current = struct_to_dict(self.data, serializer=self.serializer)
self.queue_change(
+85 -65
View File
@@ -31,11 +31,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 +63,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 +72,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 +94,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 +196,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 +255,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,21 +271,21 @@ 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,
) -> 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``.
"""
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)
+18 -15
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,7 +34,7 @@ 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
@@ -57,15 +56,15 @@ 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."""
@@ -80,15 +79,19 @@ 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:
_logger.exception(
"Background error callback failed: %s", callback_error
)
def _log_callback_error(callback_error, callback):
_logger.exception(
"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
+15 -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, log_change
from kanta.serialization import restore_data_in_place, struct_to_dict
_logger = logging.getLogger(__name__)
@@ -21,8 +21,6 @@ def transaction(
action: str,
*,
user: str | None = None,
user_display: str | None = None,
resolver: Any = None,
mtime: bool | datetime = True,
):
"""Wrap writes in a transaction and yield the live db object."""
@@ -63,7 +61,19 @@ 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
log_change(action, record.diff, formatted_user, previous, logfmt)
except Exception:
_logger.warning("Transaction '%s' failed, rolling back changes", action)
if impl.transaction_snapshot is not None:
+305
View File
@@ -0,0 +1,305 @@
from typing import Any
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_requires_return_annotation(tmp_path, format_config):
kanta = make_kanta(tmp_path / "test.db", Data, format_config)
with pytest.raises(TypeError, match="must annotate its return"):
@kanta.logfmt
def resolve_names(value: str, current: DictPost):
return None
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()
+26 -2
View File
@@ -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)
+7 -7
View File
@@ -106,7 +106,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 +121,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 +136,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 +152,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 +172,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 +188,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")