Co-authored-by: Charles-Edouard Cady <charles.edouard.cady@mistral.ai> Co-authored-by: Guillaume LE GOFF <guillaume.lgf@gmail.com> Co-authored-by: Pierre Rossinès <pierre.rossines@mistral.ai> Co-authored-by: Sylvain Afchain <safchain@gmail.com> Co-authored-by: Mistral Vibe <vibe@mistral.ai>
395 lines
13 KiB
Python
395 lines
13 KiB
Python
from __future__ import annotations
|
|
|
|
from abc import ABC, abstractmethod
|
|
import asyncio
|
|
from dataclasses import dataclass, replace
|
|
from typing import Any
|
|
|
|
from jsonpatch import JsonPatchException, apply_patch
|
|
from jsonpointer import JsonPointerException
|
|
from pydantic import BaseModel, ConfigDict, ValidationError
|
|
|
|
from vibe.core.config.patch import ConfigPatch
|
|
from vibe.core.config.types import (
|
|
ConcurrencyConflictError,
|
|
ConflictStrategy,
|
|
LayerConfigSnapshot,
|
|
)
|
|
|
|
|
|
class RawConfig(BaseModel):
|
|
"""Permissive default schema that preserves all fields as extras."""
|
|
|
|
model_config = ConfigDict(extra="allow")
|
|
|
|
|
|
class ConfigLayerError(Exception):
|
|
"""Base error for all ConfigLayer errors."""
|
|
|
|
def __init__(self, layer_name: str, message: str) -> None:
|
|
super().__init__(message)
|
|
self.layer_name = layer_name
|
|
|
|
|
|
class UntrustedLayerError(ConfigLayerError):
|
|
"""Raised when attempting to load data from an untrusted layer."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(layer_name, f"Layer '{layer_name}' is not trusted")
|
|
|
|
|
|
class EmptyLayerError(ConfigLayerError):
|
|
"""Raised when a trusted layer has no data."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(layer_name, f"Layer '{layer_name}' has no data after load")
|
|
|
|
|
|
class TrustNotResolvedError(ConfigLayerError):
|
|
"""Raised when grant_trust/revoke_trust is called before trust has been resolved."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(
|
|
layer_name, f"Layer '{layer_name}': trust has not been resolved yet"
|
|
)
|
|
|
|
|
|
class LayerNotLoadedError(ConfigLayerError):
|
|
"""Raised when a layer operation requires cached data and fingerprint."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(
|
|
layer_name, f"Layer '{layer_name}' must be loaded before applying patches"
|
|
)
|
|
|
|
|
|
class ConfigPatchApplicationError(ConfigLayerError):
|
|
"""Raised when a patch cannot be applied to cached layer data."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(layer_name, f"Layer '{layer_name}': failed to apply patch")
|
|
|
|
|
|
class TrustResolutionError(ConfigLayerError):
|
|
"""Raised when trust status is not resolvable."""
|
|
|
|
def __init__(self, layer_name: str) -> None:
|
|
super().__init__(
|
|
layer_name,
|
|
f"Layer '{layer_name}': _check_trust() must return bool, got None",
|
|
)
|
|
|
|
|
|
class LayerImplementationError(ConfigLayerError):
|
|
"""Raised when a subclass-provided method fails."""
|
|
|
|
def __init__(self, layer_name: str, method_name: str) -> None:
|
|
super().__init__(layer_name, f"Layer '{layer_name}': {method_name}() failed")
|
|
self.method_name = method_name
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _LayerState[S: BaseModel]:
|
|
is_trusted: bool | None = None
|
|
data: S | None = None
|
|
fingerprint: str | None = None
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _GrantTrust:
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _RevokeTrust:
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _ResolveTrust:
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _Load:
|
|
force: bool = False
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _InvalidateCache:
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _ApplyPatch:
|
|
patch: ConfigPatch
|
|
on_conflict: ConflictStrategy
|
|
|
|
|
|
class ConfigLayer[S: BaseModel](ABC):
|
|
"""Each layer represents a named source that produces a sparse
|
|
dictionary of configuration values.
|
|
"""
|
|
|
|
def __init__(self, *, name: str, output_schema: type[S] = RawConfig) -> None:
|
|
self.name = name
|
|
self.output_schema = output_schema
|
|
|
|
self._state: _LayerState[S] = _LayerState()
|
|
self._lock = asyncio.Lock()
|
|
|
|
# --- Overridable ---
|
|
|
|
async def _check_trust(self) -> bool:
|
|
"""Resolve whether this layer should be trusted.
|
|
|
|
Override in subclasses to implement custom trust logic.
|
|
E.g. a user-local layer returns ``True``,
|
|
a project layer checks persisted permissions.
|
|
|
|
Returns ``False`` by default (untrusted).
|
|
"""
|
|
return False
|
|
|
|
@abstractmethod
|
|
async def _build_config_snapshot(self) -> LayerConfigSnapshot:
|
|
"""Read and return sparse config with its backing-store fingerprint.
|
|
|
|
Subclasses only need to implement the raw read logic; caching
|
|
is handled by the base.
|
|
"""
|
|
...
|
|
|
|
async def _on_trust_changed(self, old: bool | None, new: bool | None) -> None:
|
|
"""Called when the trust status changes.
|
|
|
|
Override to persist trust status or react to trust transitions.
|
|
Default is a no-op.
|
|
"""
|
|
return
|
|
|
|
@abstractmethod
|
|
async def _save_to_store(self, _next_config: S) -> str:
|
|
"""Persist full layer data and return the store's new fingerprint.
|
|
|
|
The base class applies patches and validates the result before calling this.
|
|
"""
|
|
...
|
|
|
|
# --- Internal ---
|
|
|
|
async def _notify_trust_change(self, old: bool | None, new: bool | None) -> None:
|
|
"""Call ``_on_trust_changed`` and wrap any error."""
|
|
if old == new:
|
|
return
|
|
try:
|
|
await self._on_trust_changed(old, new)
|
|
except Exception as e:
|
|
raise LayerImplementationError(self.name, "_on_trust_changed") from e
|
|
|
|
async def _resolve_check_trust(self) -> bool:
|
|
"""Call ``_check_trust`` and wrap any error."""
|
|
try:
|
|
return await self._check_trust()
|
|
except Exception as e:
|
|
raise LayerImplementationError(self.name, "_check_trust") from e
|
|
|
|
async def _dispatch(self, action: Any) -> _LayerState[S]:
|
|
"""Serialize all state mutations through a single lock."""
|
|
async with self._lock:
|
|
state = _LayerState(
|
|
is_trusted=self._state.is_trusted,
|
|
data=self._state.data,
|
|
fingerprint=self._state.fingerprint,
|
|
)
|
|
|
|
match action:
|
|
case _GrantTrust():
|
|
new_state = await self._handle_grant_trust(state)
|
|
case _RevokeTrust():
|
|
new_state = await self._handle_revoke_trust(state)
|
|
case _ResolveTrust():
|
|
new_state = await self._handle_resolve_trust(state)
|
|
case _Load(force=force):
|
|
new_state = await self._handle_load(state, force)
|
|
case _InvalidateCache():
|
|
new_state = await self._handle_invalidate_cache(state)
|
|
case _ApplyPatch(patch=patch, on_conflict=on_conflict):
|
|
new_state = await self._handle_apply_patch(
|
|
state, patch, on_conflict
|
|
)
|
|
case _:
|
|
raise NotImplementedError(f"Unknown action: {action!r}")
|
|
|
|
self._state = new_state
|
|
|
|
return new_state
|
|
|
|
async def _handle_grant_trust(self, state: _LayerState[S]) -> _LayerState[S]:
|
|
if state.is_trusted is None:
|
|
raise TrustNotResolvedError(self.name)
|
|
|
|
if state.is_trusted is True:
|
|
return state
|
|
|
|
await self._notify_trust_change(state.is_trusted, True)
|
|
|
|
return _LayerState(
|
|
is_trusted=True, data=state.data, fingerprint=state.fingerprint
|
|
)
|
|
|
|
async def _handle_revoke_trust(self, state: _LayerState[S]) -> _LayerState[S]:
|
|
if state.is_trusted is None:
|
|
raise TrustNotResolvedError(self.name)
|
|
|
|
if state.is_trusted is False:
|
|
return state
|
|
|
|
await self._notify_trust_change(state.is_trusted, False)
|
|
|
|
return _LayerState(is_trusted=False, data=None, fingerprint=None)
|
|
|
|
async def _handle_resolve_trust(self, state: _LayerState[S]) -> _LayerState[S]:
|
|
is_trusted = await self._resolve_check_trust()
|
|
|
|
await self._notify_trust_change(state.is_trusted, is_trusted)
|
|
|
|
return _LayerState(
|
|
is_trusted=is_trusted,
|
|
data=state.data if is_trusted else None,
|
|
fingerprint=state.fingerprint if is_trusted else None,
|
|
)
|
|
|
|
async def _handle_load(self, state: _LayerState[S], force: bool) -> _LayerState[S]:
|
|
if state.is_trusted is not None:
|
|
is_trusted = state.is_trusted
|
|
else:
|
|
is_trusted = await self._resolve_check_trust()
|
|
|
|
await self._notify_trust_change(state.is_trusted, is_trusted)
|
|
|
|
if not is_trusted:
|
|
return _LayerState(is_trusted=is_trusted, data=None, fingerprint=None)
|
|
|
|
next_state = _LayerState(
|
|
is_trusted=is_trusted, data=state.data, fingerprint=state.fingerprint
|
|
)
|
|
|
|
if next_state.data is None or force:
|
|
try:
|
|
snapshot = await self._build_config_snapshot()
|
|
next_state = _LayerState(
|
|
is_trusted=next_state.is_trusted,
|
|
data=self.validate_output(snapshot.data),
|
|
fingerprint=snapshot.fingerprint,
|
|
)
|
|
except ConcurrencyConflictError:
|
|
raise
|
|
except Exception as e:
|
|
raise LayerImplementationError(
|
|
self.name, "_build_config_snapshot"
|
|
) from e
|
|
|
|
return next_state
|
|
|
|
async def _handle_invalidate_cache(self, state: _LayerState[S]) -> _LayerState[S]:
|
|
return _LayerState(is_trusted=state.is_trusted, data=None, fingerprint=None)
|
|
|
|
async def _handle_apply_patch(
|
|
self, state: _LayerState[S], patch: ConfigPatch, on_conflict: ConflictStrategy
|
|
) -> _LayerState[S]:
|
|
|
|
if state.data is None or state.fingerprint is None:
|
|
raise LayerNotLoadedError(self.name)
|
|
|
|
match on_conflict:
|
|
case ConflictStrategy.CANCEL:
|
|
if patch.fingerprint != state.fingerprint:
|
|
raise ConcurrencyConflictError(
|
|
expected_fp=patch.fingerprint, actual_fp=state.fingerprint
|
|
)
|
|
case ConflictStrategy.REPLACE:
|
|
pass
|
|
case _:
|
|
raise ValueError(f"Unsupported conflict strategy: {on_conflict!r}")
|
|
|
|
try:
|
|
new_data = apply_patch(state.data.model_dump(), patch.to_json_patch())
|
|
validated_new_data = self.validate_output(new_data)
|
|
except (JsonPatchException, JsonPointerException, ValidationError) as e:
|
|
raise ConfigPatchApplicationError(self.name) from e
|
|
|
|
try:
|
|
new_fingerprint = await self._save_to_store(validated_new_data)
|
|
except NotImplementedError:
|
|
raise
|
|
except Exception as e:
|
|
raise LayerImplementationError(self.name, "_save_to_store") from e
|
|
return replace(state, data=validated_new_data, fingerprint=new_fingerprint)
|
|
|
|
# --- Public ---
|
|
|
|
@property
|
|
def is_trusted(self) -> bool | None:
|
|
"""Current trust status. ``None`` if unresolved."""
|
|
return self._state.is_trusted
|
|
|
|
async def resolve_trust(self) -> bool:
|
|
"""Resolve and cache trust status."""
|
|
state = await self._dispatch(_ResolveTrust())
|
|
|
|
if state.is_trusted is None:
|
|
raise TrustResolutionError(self.name)
|
|
|
|
return state.is_trusted
|
|
|
|
async def grant_trust(self) -> None:
|
|
"""Explicitly mark this layer as trusted."""
|
|
await self._dispatch(_GrantTrust())
|
|
|
|
async def revoke_trust(self) -> None:
|
|
"""Explicitly mark this layer as untrusted and clear cached data."""
|
|
await self._dispatch(_RevokeTrust())
|
|
|
|
async def invalidate_cache(self) -> None:
|
|
"""Clear the in-memory cache so the next ``load()`` re-reads."""
|
|
await self._dispatch(_InvalidateCache())
|
|
|
|
async def load(self, *, force: bool = False) -> S:
|
|
"""Load data from this layer, enforcing trust and caching the result.
|
|
|
|
Use ``force=True`` to bypass caching.
|
|
"""
|
|
state = await self._dispatch(_Load(force=force))
|
|
|
|
if state.is_trusted is None:
|
|
raise TrustResolutionError(self.name)
|
|
|
|
if not state.is_trusted:
|
|
raise UntrustedLayerError(self.name)
|
|
|
|
if state.data is None:
|
|
raise EmptyLayerError(self.name)
|
|
|
|
return state.data.model_copy(deep=True)
|
|
|
|
@property
|
|
def fingerprint(self) -> str | None:
|
|
"""Cached opaque fingerprint token for this layer. ``None`` if unresolved."""
|
|
return self._state.fingerprint
|
|
|
|
async def apply(
|
|
self,
|
|
patch: ConfigPatch,
|
|
*,
|
|
on_conflict: ConflictStrategy = ConflictStrategy.CANCEL,
|
|
) -> None:
|
|
"""Persist a patch to this layer's backing store."""
|
|
await self._dispatch(_ApplyPatch(patch=patch, on_conflict=on_conflict))
|
|
|
|
def validate_output(self, data: dict[str, Any]) -> S:
|
|
"""Validate *data* against ``output_schema``."""
|
|
return self.output_schema.model_validate(data)
|
|
|
|
def __repr__(self) -> str:
|
|
return f"{type(self).__name__}(name={self.name!r})"
|