Improved log formatting support by @kanta.logfmt, which replaces old resolver and user_display arguments (breaking change).
This commit is contained in:
+50
-6
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
+16
-13
@@ -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,14 +79,18 @@ class PersistenceMixin:
|
||||
break
|
||||
except DatabaseError as e:
|
||||
self.background_error = e
|
||||
for callback in self.fatal_error_handlers:
|
||||
try:
|
||||
callback_result = callback(e)
|
||||
if inspect.isawaitable(callback_result):
|
||||
await callback_result
|
||||
except Exception as callback_error:
|
||||
|
||||
def _log_callback_error(callback_error, callback):
|
||||
_logger.exception(
|
||||
"Background error callback failed: %s", callback_error
|
||||
"Background error callback %r failed: %s",
|
||||
callback,
|
||||
callback_error,
|
||||
)
|
||||
|
||||
await self.callback_registry.invoke(
|
||||
"fatal_error",
|
||||
InjectionContext(error=e, kanta=self._kanta),
|
||||
on_error=_log_callback_error,
|
||||
)
|
||||
_logger.error("Background flush loop stopped: %s", e)
|
||||
break
|
||||
|
||||
+15
-5
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user