From 7e553bd86810491df5a88852bb3d3f4415014e99 Mon Sep 17 00:00:00 2001 From: Leo Vasanko Date: Sat, 13 Jun 2026 21:22:05 +0000 Subject: [PATCH] Improved log formatting support by @kanta.logfmt, which replaces old resolver and user_display arguments (breaking change). --- docs/database.md | 56 +++- kanta/__init__.py | 6 +- kanta/callbacks.py | 537 ++++++++++++++++++++++++++++++++ kanta/kanta.py | 39 ++- kanta/kantaimpl.py | 25 +- kanta/logging.py | 150 +++++---- kanta/persistence.py | 33 +- kanta/transaction.py | 20 +- tests/test_callbacks.py | 305 ++++++++++++++++++ tests/test_format_diff.py | 28 +- tests/test_kanta_integration.py | 14 +- 11 files changed, 1091 insertions(+), 122 deletions(-) create mode 100644 kanta/callbacks.py create mode 100644 tests/test_callbacks.py diff --git a/docs/database.md b/docs/database.md index 394b3ef..141c6db 100644 --- a/docs/database.md +++ b/docs/database.md @@ -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 diff --git a/kanta/__init__.py b/kanta/__init__.py index 4655912..a741010 100644 --- a/kanta/__init__.py +++ b/kanta/__init__.py @@ -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", ] diff --git a/kanta/callbacks.py b/kanta/callbacks.py new file mode 100644 index 0000000..7dfaa7a --- /dev/null +++ b/kanta/callbacks.py @@ -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() diff --git a/kanta/kanta.py b/kanta/kanta.py index a029d5b..3060359 100644 --- a/kanta/kanta.py +++ b/kanta/kanta.py @@ -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, ) diff --git a/kanta/kantaimpl.py b/kanta/kantaimpl.py index c298fac..a171fb7 100644 --- a/kanta/kantaimpl.py +++ b/kanta/kantaimpl.py @@ -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( diff --git a/kanta/logging.py b/kanta/logging.py index 9c2f650..48d1ec6 100644 --- a/kanta/logging.py +++ b/kanta/logging.py @@ -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) diff --git a/kanta/persistence.py b/kanta/persistence.py index 64f1214..9b7d000 100644 --- a/kanta/persistence.py +++ b/kanta/persistence.py @@ -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 diff --git a/kanta/transaction.py b/kanta/transaction.py index 8e37afb..d5a9559 100644 --- a/kanta/transaction.py +++ b/kanta/transaction.py @@ -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: diff --git a/tests/test_callbacks.py b/tests/test_callbacks.py new file mode 100644 index 0000000..9fe462d --- /dev/null +++ b/tests/test_callbacks.py @@ -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() diff --git a/tests/test_format_diff.py b/tests/test_format_diff.py index 098f49a..0c91d97 100644 --- a/tests/test_format_diff.py +++ b/tests/test_format_diff.py @@ -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) diff --git a/tests/test_kanta_integration.py b/tests/test_kanta_integration.py index 19ff268..eff6179 100644 --- a/tests/test_kanta_integration.py +++ b/tests/test_kanta_integration.py @@ -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")