use std::{error::Error, fmt};
use lenso_kernel::{InvocationContext, InvocationContextError};
use serde::{Serialize, de::DeserializeOwned};
pub trait TypedExtension: Serialize + DeserializeOwned {
const KEY: &'static str;
}
#[derive(Debug)]
pub enum TypedExtensionError {
Encode(serde_json::Error),
Decode(serde_json::Error),
Context(InvocationContextError),
}
impl fmt::Display for TypedExtensionError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Encode(error) => write!(formatter, "typed extension encoding failed: {error}"),
Self::Decode(error) => write!(formatter, "typed extension decoding failed: {error}"),
Self::Context(error) => write!(formatter, "typed extension attachment failed: {error}"),
}
}
}
impl Error for TypedExtensionError {}
pub trait CtxExt: Sized {
fn typed_extension<T: TypedExtension>(&self) -> Result<Option<T>, TypedExtensionError>;
fn with_typed_extension<T: TypedExtension>(
self,
value: &T,
) -> Result<Self, TypedExtensionError>;
}
impl CtxExt for InvocationContext {
fn typed_extension<T: TypedExtension>(&self) -> Result<Option<T>, TypedExtensionError> {
self.extension(T::KEY)
.map(|bytes| serde_json::from_slice(bytes).map_err(TypedExtensionError::Decode))
.transpose()
}
fn with_typed_extension<T: TypedExtension>(
self,
value: &T,
) -> Result<Self, TypedExtensionError> {
let bytes = serde_json::to_vec(value).map_err(TypedExtensionError::Encode)?;
self.with_extension(T::KEY, bytes)
.map_err(TypedExtensionError::Context)
}
}