use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Visitor};
use std::{
cmp::Ordering,
fmt,
hash::{Hash, Hasher},
marker::PhantomData,
ops::{Deref, DerefMut},
};
pub trait Recover<T> {
fn recover<E: serde::de::Error>(err: E) -> Result<T, E>;
}
pub struct RecoverDefault;
impl<T: Default> Recover<T> for RecoverDefault {
fn recover<E: serde::de::Error>(_err: E) -> Result<T, E> {
Ok(T::default())
}
}
pub struct Recoverable<T, P = RecoverDefault> {
value: T,
recovered: bool,
_policy: PhantomData<fn() -> P>,
}
impl Recoverable<(), ()> {
pub(crate) const NEWTYPE_NAME: &str = "$postbag::recoverable::Recoverable";
}
impl<T, P> Recoverable<T, P> {
pub fn new(value: T) -> Self {
Self { value, recovered: false, _policy: PhantomData }
}
pub fn into_inner(this: Self) -> T {
this.value
}
pub fn is_recovered(this: &Self) -> bool {
this.recovered
}
fn recovered(value: T) -> Self {
Self { value, recovered: true, _policy: PhantomData }
}
}
impl<T, P> From<T> for Recoverable<T, P> {
fn from(value: T) -> Self {
Self::new(value)
}
}
impl<T, P> Deref for Recoverable<T, P> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.value
}
}
impl<T, P> DerefMut for Recoverable<T, P> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.value
}
}
impl<T: fmt::Debug, P> fmt::Debug for Recoverable<T, P> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Recoverable").field("value", &self.value).field("recovered", &self.recovered).finish()
}
}
impl<T: Clone, P> Clone for Recoverable<T, P> {
fn clone(&self) -> Self {
Self { value: self.value.clone(), recovered: self.recovered, _policy: PhantomData }
}
}
impl<T: Copy, P> Copy for Recoverable<T, P> {}
impl<T: Default, P> Default for Recoverable<T, P> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: PartialEq, P> PartialEq for Recoverable<T, P> {
fn eq(&self, other: &Self) -> bool {
self.value == other.value
}
}
impl<T: Eq, P> Eq for Recoverable<T, P> {}
impl<T: PartialOrd, P> PartialOrd for Recoverable<T, P> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
self.value.partial_cmp(&other.value)
}
}
impl<T: Ord, P> Ord for Recoverable<T, P> {
fn cmp(&self, other: &Self) -> Ordering {
self.value.cmp(&other.value)
}
}
impl<T: Hash, P> Hash for Recoverable<T, P> {
fn hash<H: Hasher>(&self, state: &mut H) {
self.value.hash(state);
}
}
impl<T: Serialize, P> Serialize for Recoverable<T, P> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_newtype_struct(Recoverable::NEWTYPE_NAME, &self.value)
}
}
impl<'de, T: Deserialize<'de>, P: Recover<T>> Deserialize<'de> for Recoverable<T, P> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
deserializer.deserialize_newtype_struct(Recoverable::NEWTYPE_NAME, RecoverableVisitor(PhantomData))
}
}
struct RecoverableVisitor<T, P>(PhantomData<fn() -> (T, P)>);
impl<'de, T: Deserialize<'de>, P: Recover<T>> Visitor<'de> for RecoverableVisitor<T, P> {
type Value = Recoverable<T, P>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a recoverable value")
}
fn visit_newtype_struct<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
match T::deserialize(deserializer) {
Ok(value) => Ok(Recoverable::new(value)),
Err(err) => P::recover(err).map(Recoverable::recovered),
}
}
}
pub type SelfRecoverable<T> = Recoverable<T, T>;
pub struct With<P>(PhantomData<fn() -> P>);
impl<P> With<P> {
pub fn serialize<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
where
T: Serialize + ?Sized,
S: Serializer,
{
Recoverable::<&T, P>::new(value).serialize(serializer)
}
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<T, D::Error>
where
T: Deserialize<'de>,
P: Recover<T>,
D: Deserializer<'de>,
{
Recoverable::<T, P>::deserialize(deserializer).map(Recoverable::into_inner)
}
}
pub fn serialize<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
where
T: Serialize + ?Sized,
S: Serializer,
{
With::<RecoverDefault>::serialize(value, serializer)
}
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<T, D::Error>
where
T: Deserialize<'de> + Default,
D: Deserializer<'de>,
{
With::<RecoverDefault>::deserialize(deserializer)
}